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)Every recursive function needs two parts:
- A base case: a simple situation answered directly, without recursion. Here it's
n == 0. - 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))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)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]))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]))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)Python limits how deep recursion can go, to protect your program. You can check the limit:
import sys print(sys.getrecursionlimit())
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)])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))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))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))Base case: when n is 0 or less, the sum is 0.
Otherwise the sum is n plus the sum of everything below it: n + sum_to(n - 1).
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))Any number to the power 0 is 1. That's your base case.
Otherwise multiply base by power(base, exp - 1).
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)Stop with a plain return when n reaches 0.
Print n, call bounce(n - 1), then print n again. The second print runs on the way back up.
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"))Base case: a string of length 0 or 1 is always a palindrome.
If the first and last characters differ, return False. Otherwise check the middle part, text[1:-1].
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_cacheremembers results. - Use recursion for naturally recursive problems such as nested data, and loops for simple repetition.