~/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.
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.
You must list every valid combination, subset, ordering or placement, and the input is small.
O(2ⁿ · n)
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:
nup 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.
- Sort
nums, then callbacktrack(0, target)with an emptypath. The arguments arestart, the first index this call may use, andremain, what’s still needed. - Choose: the loop tries each index
ifromstartonwards. Appendnums[i]topath. - Explore: call
backtrack(i + 1, remain - nums[i]). Passingi + 1means 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]. - When
remainreaches 0,pathis an answer. Append a copy of it toresultand return. - Un-choose: when the call returns, pop the value off
path. The loop then tries the next index from the same starting point. - Prune: if
nums[i] > remain, this value overshoots. Becausenumsis sorted, every later value overshoots too, sobreakout of the loop. - Skip repeats: if
i > startandnums[i] == nums[i - 1], the loop has just explored this same value in this same position. Every group it could lead to is already inresult, socontinue. - When the root call returns,
resultholds every group exactly once.
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
iinstead ofi + 1. The loop still never goes back to smaller indices, so each group still appears once. - All subsets, or subsets of size k. Record
pathat every call, not only at the leaves. For sizek, record whenlen(path) == k, and prune when fewer thank - len(path)values are left. - Orderings. Every position may take any unused value, so loop over all indices and keep a
usedarray: mark it when you choose, clear it when you un-choose. With repeated values, sort and skipnums[i]when it equalsnums[i - 1]andnums[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
Trueas 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
You get 40 numbers, each between 1 and 50, and must return how many subsets add up to exactly 500. Which approach fits?
Only the count is needed, and totals repeat across many subsets, so a table of about 40 × 500 entries answers it. Backtracking visits each subset that stays within 500 one by one, and there can be close to 2⁴⁰ of those; pruning cuts only branches that overshoot, so it doesn’t save you. Greedy finds at most one subset, not a count.
-
2
What does this print?
def subsets(nums):result, path = [], []def go(i):if i == len(nums):result.append(path)returnpath.append(nums[i])go(i + 1)path.pop()go(i + 1)go(0)return resultprint(subsets([1, 2]))result.append(path)stores the same list object four times, not four snapshots. The search finishes with every value popped, so that one list is empty, and all four entries show it. Appendingpath[:]would give the intended[[1, 2], [1], [2], []]. -
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]:continuepath.append(nums[i])go(i + 1)path.pop()go(0)return outprint(len(distinct_subsets([2, 1, 2])), len(distinct_subsets([3, 3, 3])))Every call records its
path, so every distinct subset is counted once, the empty one included.[1, 2, 2]has 6: [], [1], [1, 2], [1, 2, 2], [2], [2, 2].[3, 3, 3]has 4: no 3s, one, two or three.8 8is what you’d get without the skip, counting 2³ subsets of positions; the skip only drops a repeat at the same level, so[3, 3]and[3, 3, 3]still appear. -
4
This should return every distinct group from
nums(each value used at most once) that adds up totarget. Forgroups([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[:])returnfor i in range(start, len(nums)):if i > 0 and nums[i] == nums[i - 1]:continueif nums[i] > remain:breakpath.append(nums[i])backtrack(i + 1, remain - nums[i])path.pop()backtrack(0, target)return resultRepeats should be skipped only among the choices of one loop. With
i > 0, the call that already holds[1, 2]starts at the second 2, sees that it equals the value before it, and skips it, so[1, 2, 2]is never built.breakis right becausenumsis sorted: once a value is too big, so is everything after it. Recursing withiwould allow a single 2 to be used twice. -
5
Listing all orderings of
ndistinct items takes aboutn! · nsteps. Roughly how large cannbe if you have about 10⁸ simple steps?10! · 10 is about 3.6 × 10⁷, which fits; 11! · 11 is about 4.4 × 10⁸, which doesn’t. Subsets grow as 2ⁿ, so they reach about 20 in the same budget, but n! grows far faster: 20! is about 2.4 × 10¹⁸.
Practice problems
Solve these right here, in Python, C++ or Java. Tests run as you go.