Learn Python / Recursion

Recursion

A recursive function is a function that calls itself. It solves a problem by breaking it into a smaller version of the same problem, and keeps going until the problem is small enough to answer directly.

A first example: countdown

def countdown(n):
    if n == 0:
        print("Liftoff!")
        return
    print(n)
    countdown(n - 1)

countdown(5)
Output

Every recursive function needs two parts:

  1. A base case: a simple situation answered directly, without recursion. Here it's n == 0.
  2. A recursive case: the function calls itself with a smaller problem, here n - 1, moving towards the base case.

Factorial

The factorial of n, written n!, is n × (n − 1) × … × 1. So 5! = 5 × 4 × 3 × 2 × 1 = 120.

Notice that 5! = 5 × 4!. The factorial is defined in terms of a smaller factorial, which makes it a natural fit for recursion:

def factorial(n):
    if n <= 1:              # base case
        return 1
    return n * factorial(n - 1)   # recursive case

print(factorial(5))
print(factorial(10))
Output

Here's what happens for factorial(4):

factorial(4) = 4 * factorial(3)
             = 4 * (3 * factorial(2))
             = 4 * (3 * (2 * factorial(1)))
             = 4 * (3 * (2 * 1))
             = 24

Each call waits for the one below it to return, then finishes its own multiplication.

Seeing the calls

Printing with indentation makes the calls visible:

def factorial(n, depth=0):
    indent = "  " * depth
    print(f"{indent}factorial({n})")
    if n <= 1:
        result = 1
    else:
        result = n * factorial(n - 1, depth + 1)
    print(f"{indent}-> {result}")
    return result

factorial(4)
Output

Recursion on strings and lists

Recursion works well whenever data has a "first item plus the rest" shape:

def reverse(text):
    if text == "":
        return ""
    return reverse(text[1:]) + text[0]

def total(numbers):
    if not numbers:
        return 0
    return numbers[0] + total(numbers[1:])

print(reverse("python"))
print(total([4, 8, 15, 16, 23, 42]))
Output

Nested data

Recursion really shines with data that contains smaller copies of itself, like lists inside lists:

def flatten(items):
    flat = []
    for item in items:
        if isinstance(item, list):
            flat.extend(flatten(item))   # recurse into the inner list
        else:
            flat.append(item)
    return flat

print(flatten([1, [2, 3], [4, [5, [6, 7]]], 8]))
Output

Doing this with plain loops is awkward, because you don't know in advance how deep the nesting goes.

Forgetting the base case

Without a base case, or with one that's never reached, the function calls itself forever: 3, 2, 1, 0, -1, -2… Python stops it with a RecursionError:

def countdown(n):
    countdown(n - 1)    # no base case, so this never stops!

countdown(3)
Output

Python limits how deep recursion can go, to protect your program. You can check the limit:

import sys
print(sys.getrecursionlimit())
Output

Fibonacci and the cost of repeated work

In the Fibonacci sequence, each number is the sum of the two before it: 0, 1, 1, 2, 3, 5, 8, 13… A direct recursive version is short:

def fib(n):
    if n < 2:
        return n
    return fib(n - 1) + fib(n - 2)

print([fib(i) for i in range(10)])
Output

But it's slow for bigger numbers, because it recalculates the same values again and again. fib(30) makes over a million calls. Remembering results that are already known, called memoization, fixes that:

from functools import lru_cache

@lru_cache
def fib(n):
    if n < 2:
        return n
    return fib(n - 1) + fib(n - 2)

print(fib(30))
print(fib(100))
Output

The @lru_cache line is a decorator. You'll learn how those work in the Decorators lesson.

Recursion or a loop?

Anything recursive can also be written with a loop, and for simple counting a loop is usually clearer and faster:

def factorial_loop(n):
    result = 1
    for i in range(2, n + 1):
        result *= i
    return result

print(factorial_loop(5))
Output

Reach for recursion when the problem is naturally recursive: nested data, trees, folders inside folders, or divide-and-conquer algorithms.

Exercises

Exercise 1: Sum to n

Write a recursive function sum_to(n) that returns 1 + 2 + ... + n. sum_to(10) should be 55.

def sum_to(n):
    pass

print(sum_to(10))
Output

Exercise 2: Power

Write a recursive power(base, exp) without using **. power(2, 10) should be 1024.

def power(base, exp):
    pass

print(power(2, 10))
print(power(3, 0))
Output

Exercise 3: Count down and up

Write bounce(n) that prints the numbers from n down to 1 and then back up to n. Hint: put a print both before and after the recursive call. For bounce(3) it should print 3 2 1 1 2 3 (one per line).

def bounce(n):
    pass

bounce(3)
Output

Exercise 4: Palindrome

Write a recursive is_palindrome(text): a string is a palindrome if its first and last characters match and the middle part is a palindrome. Empty strings and single characters are palindromes.

def is_palindrome(text):
    pass

print(is_palindrome("racecar"))
print(is_palindrome("python"))
Output

Summary

  • A recursive function calls itself on a smaller version of the problem.
  • Every recursive function needs a base case that stops the recursion.
  • Each call waits for the call below it, and the results combine on the way back up.
  • Missing or unreachable base cases cause a RecursionError.
  • Naive recursion can repeat work. functools.lru_cache remembers results.
  • Use recursion for naturally recursive problems such as nested data, and loops for simple repetition.