~/techniques/backtracking

Recursion & backtracking

List every subset, ordering or combination by building one choice at a time, and undoing each choice before trying the next.

what

Walk the tree of choices depth-first with one shared path: choose, explore, un-choose. Cut any branch that can't lead to an answer.

use when

You must list every valid combination, subset, ordering or placement, and the input is small.

time

O(2ⁿ · n)

space

O(n)

You’ll recognise it when

  • The question says every, all or list: all subsets, all orderings, every way to split, place or fill something.
  • The input is tiny: n up to about 20 for subsets, or about 10 for orderings. Those limits are a hint that exponential time is expected.
  • You build an answer piece by piece, and you can tell early that a partial answer is hopeless: it already went over a budget, or two queens already attack each other.
  • Equal values in the input must not produce the same answer twice.

If the question only wants how many ways, or the best one, and the same smaller question keeps coming up, dynamic programming is usually much faster.

The idea

Think of trying on outfits. You pick a shirt, then trousers to go with it, then shoes. When you’ve seen that combination, you take the shoes off and try the next pair. When you’ve tried every pair, you take the trousers off too and try the next trousers. You never undress completely to start over: you undo only the last choice.

Every “list them all” problem is a tree of choices. The root is “nothing chosen yet”, each edge adds one choice, and each leaf is a finished candidate. Backtracking walks that tree depth-first while keeping a single shared path of the choices made so far. Each step down has three parts:

path.append(choice) # choose
backtrack(...) # explore everything that starts with this path
path.pop() # un-choose: path is exactly as it was before

The un-choose step is what lets one list serve the whole tree. The second trick is pruning: if a partial path can never become an answer, return at once and skip its whole subtree.

How it works

The running example: from nums = [2, 1, 3, 2], find every group of values that adds up to target = 5, using each value at most once. The two 2s are interchangeable, so [2, 3] must appear only once.

Scroll through the steps and the graphic follows along. Each circle is one call and shows its remain, the amount still needed; each edge is the value chosen. You can also press play, or edit the input: try 1, 1, 1, 2 with target 3.

  1. Sort nums, then call backtrack(0, target) with an empty path. The arguments are start, the first index this call may use, and remain, what’s still needed.
  2. Choose: the loop tries each index i from start onwards. Append nums[i] to path.
  3. Explore: call backtrack(i + 1, remain - nums[i]). Passing i + 1 means each value is used at most once, and every group is built in sorted order, so [1, 3] is never built again as [3, 1].
  4. When remain reaches 0, path is an answer. Append a copy of it to result and return.
  5. Un-choose: when the call returns, pop the value off path. The loop then tries the next index from the same starting point.
  6. Prune: if nums[i] > remain, this value overshoots. Because nums is sorted, every later value overshoots too, so break out of the loop.
  7. Skip repeats: if i > start and nums[i] == nums[i - 1], the loop has just explored this same value in this same position. Every group it could lead to is already in result, so continue.
  8. When the root call returns, result holds every group exactly once.
loading backtracking…

Why it’s correct: each call is entered with path holding exactly the choices on the way down, and it returns with path unchanged. So every branch starts from a clean slate. Building groups in index order produces each set of positions once, and skipping a repeated value at the same level removes the only other way to build the same group of values.

def groups_with_sum(nums, target):
"""Every distinct group of values from nums, each value used at most
once, that adds up to target. Values are positive."""
nums = sorted(nums)
result, path = [], []
def backtrack(start, remain):
if remain == 0:
result.append(path[:])
return
for i in range(start, len(nums)):
if i > start and nums[i] == nums[i - 1]:
continue # same value, same subtree
if nums[i] > remain:
break # sorted: the rest are bigger too
path.append(nums[i])
backtrack(i + 1, remain - nums[i])
path.pop()
backtrack(0, target)
return result
#include <algorithm>
#include <vector>
using namespace std;
void backtrack(const vector<int>& nums, int start, int remain,
vector<int>& path, vector<vector<int>>& result) {
if (remain == 0) {
result.push_back(path);
return; // (push_back copies path)
}
for (int i = start; i < (int)nums.size(); i++) {
if (i > start && nums[i] == nums[i - 1]) {
continue; // same value, same subtree
}
if (nums[i] > remain) {
break; // sorted: the rest are bigger too
}
path.push_back(nums[i]);
backtrack(nums, i + 1, remain - nums[i], path, result);
path.pop_back();
}
}
vector<vector<int>> groupsWithSum(vector<int> nums, int target) {
sort(nums.begin(), nums.end());
vector<int> path;
vector<vector<int>> result;
backtrack(nums, 0, target, path, result);
return result;
}
import java.util.*;
class GroupsWithSum {
private int[] nums;
private final List<Integer> path = new ArrayList<>();
private final List<List<Integer>> result = new ArrayList<>();
List<List<Integer>> find(int[] values, int target) {
nums = values.clone();
Arrays.sort(nums);
backtrack(0, target);
return result;
}
private void backtrack(int start, int remain) {
if (remain == 0) {
result.add(new ArrayList<>(path));
return;
}
for (int i = start; i < nums.length; i++) {
if (i > start && nums[i] == nums[i - 1]) {
continue; // same value, same subtree
}
if (nums[i] > remain) {
break; // sorted: the rest are bigger too
}
path.add(nums[i]);
backtrack(i + 1, remain - nums[i]);
path.remove(path.size() - 1);
}
}
}

Why it’s O(2ⁿ · n)

Without pruning, the tree has one node for every subset of the n values, because each call adds one new value past the last. That’s 2ⁿ calls, and copying an answer costs up to n, so the total is O(2ⁿ · n). Pruning and skipping repeats often cut the tree to a small fraction of that, but they don’t change the worst case: with small values and a large target, few branches get cut.

Orderings grow faster. The first position has n choices, the next n − 1, and so on, so there are n! leaves and the work is O(n! · n).

n subsets, 2ⁿ orderings, n!
5 32 120
10 1,024 3,628,800
15 32,768 1.3 × 10¹²
20 1,048,576 2.4 × 10¹⁸

Extra space is O(n): the path and the recursion stack are never deeper than n. The output itself can be exponentially large.

Common mistakes

Forgetting to un-choose

If you don’t pop after the recursive call, the value stays in path while the loop tries its siblings. Later answers then contain values from branches you already left.

path.append(nums[i]); backtrack(i + 1, remain - nums[i]) # ✗ nums[i] is still in path
path.append(nums[i]); backtrack(i + 1, remain - nums[i]); path.pop() # ✓ path is back as it was

Recording the shared list instead of a copy

result.append(path) stores a reference to the one list that keeps changing. When the search ends, path is empty, so result is full of empty lists. Java has the same trap; C++ push_back(path) copies for you.

result.append(path) # ✗ every entry is the same list, empty at the end
result.append(path[:]) # ✓ a snapshot of path right now

Skipping repeats with i > 0

The check must compare with the previous value in the same loop. With i > 0, the second 2 can never follow the first 2 one level down, so [1, 2, 2] is lost. Skipping only works on sorted input, where equal values are neighbours.

if i > 0 and nums[i] == nums[i - 1]: continue # ✗ also bans the second 2 deeper down
if i > start and nums[i] == nums[i - 1]: continue # ✓ only a repeat at the same level

Recursing with the wrong start

The next call must begin after the value just chosen. start + 1 lets a later call pick a value that comes before i, so the same group turns up in several orders.

backtrack(start + 1, remain - nums[i]) # ✗ builds [1, 3] and [3, 1]
backtrack(i + 1, remain - nums[i]) # ✓ each group in index order, once

Variations

  • Unlimited reuse. If each value may be used any number of times, recurse with i instead of i + 1. The loop still never goes back to smaller indices, so each group still appears once.
  • All subsets, or subsets of size k. Record path at every call, not only at the leaves. For size k, record when len(path) == k, and prune when fewer than k - len(path) values are left.
  • Orderings. Every position may take any unused value, so loop over all indices and keep a used array: mark it when you choose, clear it when you un-choose. With repeated values, sort and skip nums[i] when it equals nums[i - 1] and nums[i - 1] is not in use.
  • Boards and grids. The path is the cells or squares taken so far. Mark a cell on the way down and unmark it on the way up; keep sets of used columns and diagonals for queens, or of used letters for a word search.
  • Stop at the first answer. For puzzles like Sudoku or splitting into equal groups, return True as soon as one branch succeeds and let it bubble up, instead of listing everything.

Check yourself

5 quick questions. Pick an answer to see why it's right or wrong.

  1. 1

    You get 40 numbers, each between 1 and 50, and must return how many subsets add up to exactly 500. Which approach fits?

  2. 2

    What does this print?

    def subsets(nums):
    result, path = [], []
    def go(i):
    if i == len(nums):
    result.append(path)
    return
    path.append(nums[i])
    go(i + 1)
    path.pop()
    go(i + 1)
    go(0)
    return result
    print(subsets([1, 2]))
  3. 3

    What does this print?

    def distinct_subsets(nums):
    nums = sorted(nums)
    out, path = [], []
    def go(start):
    out.append(path[:])
    for i in range(start, len(nums)):
    if i > start and nums[i] == nums[i - 1]:
    continue
    path.append(nums[i])
    go(i + 1)
    path.pop()
    go(0)
    return out
    print(len(distinct_subsets([2, 1, 2])), len(distinct_subsets([3, 3, 3])))
  4. 4

    This should return every distinct group from nums (each value used at most once) that adds up to target. For groups([2, 1, 3, 2], 5) it returns only [[2, 3]] and misses [1, 2, 2]. What’s wrong?

    def groups(nums, target):
    nums = sorted(nums)
    result, path = [], []
    def backtrack(start, remain):
    if remain == 0:
    result.append(path[:])
    return
    for i in range(start, len(nums)):
    if i > 0 and nums[i] == nums[i - 1]:
    continue
    if nums[i] > remain:
    break
    path.append(nums[i])
    backtrack(i + 1, remain - nums[i])
    path.pop()
    backtrack(0, target)
    return result
  5. 5

    Listing all orderings of n distinct items takes about n! · n steps. Roughly how large can n be if you have about 10⁸ simple steps?

Practice problems

Solve these right here, in Python, C++ or Java. Tests run as you go.

All 19 problems on this topic

Further reading

esc