Backtracking

Backtracking

Choose → explore → undo. Subsets, combinations, permutations, N-Queens, and pruning tips — with plain-language explanation before templates.

Choose · explore · undo

Backtracking builds candidates incrementally. After exploring a choice, undo the mutation so the sibling branch sees a clean state. Prune early when remaining options can’t beat the constraint — that’s the difference between TLE and AC.

Analogy: maze with undo

Backtracking maze
Choose → explore → hit a wall → undo → try another path.

Backtracking is walking a hedge maze: you try a corridor, and if you hit a dead end you walk back and try the next opening. The undo step is what keeps sibling paths clean — forget it and your maze map stays scribbled.

The shape

Backtracking tree
Solution space tree.
def subsets(nums):
    ans, path = [], []
    def dfs(i):
        if i == len(nums):
            ans.append(path[:]); return
        path.append(nums[i]); dfs(i+1); path.pop()  # take
        dfs(i+1)  # skip
    dfs(0); return ans
List<List<Integer>> subsets(int[] nums) {
    List<List<Integer>> ans = new ArrayList<>();
    List<Integer> path = new ArrayList<>();
    dfs(nums, 0, path, ans);
    return ans;
}
void dfs(int[] nums, int i, List<Integer> path, List<List<Integer>> ans) {
    if (i == nums.length) {
        ans.add(new ArrayList<>(path));
        return;
    }
    path.add(nums[i]); dfs(nums, i + 1, path, ans); path.remove(path.size() - 1); // take
    dfs(nums, i + 1, path, ans); // skip
}

Templates to memorize

  • Subsets (take/skip)
  • Combinations (start index)
  • Permutations (used[])
  • Partition / palindrome partition
  • Board search (word search, N-Queens)

Worked trace: subsets

For [1,2]: at index 0 choose take 1 → then take/skip 2 → paths [1,2],[1]; undo; skip 1 → take/skip 2 → [2],[]. The undo after the take branch is what keeps siblings correct.

def subsets(nums):
    ans, path = [], []
    def dfs(i):
        if i == len(nums):
            ans.append(path.copy()); return
        path.append(nums[i]); dfs(i+1); path.pop()
        dfs(i+1)
    dfs(0)
    return ans
List<List<Integer>> subsets(int[] nums) {
    List<List<Integer>> ans = new ArrayList<>();
    List<Integer> path = new ArrayList<>();
    dfs(nums, 0, path, ans);
    return ans;
}
void dfs(int[] nums, int i, List<Integer> path, List<List<Integer>> ans) {
    if (i == nums.length) {
        ans.add(new ArrayList<>(path));
        return;
    }
    path.add(nums[i]); dfs(nums, i + 1, path, ans); path.remove(path.size() - 1);
    dfs(nums, i + 1, path, ans);
}

When to reach for this pattern

Backtracking is DFS on a decision tree. The skill is the undo and the prune — not memorizing twelve templates.

Core template + worked trace

Subsets — include/exclude or loop-from-start:

def subsets(nums: list[int]) -> list[list[int]]:
    out: list[list[int]] = []
    path: list[int] = []

    def dfs(start: int):
        out.append(path.copy())          # record every node
        for i in range(start, len(nums)):
            path.append(nums[i])         # choose
            dfs(i + 1)                   # explore
            path.pop()                   # undo

    dfs(0)
    return out

# nums=[1,2]
# path=[] → record []
# choose 1 → [1] record; choose 2 → [1,2] record; undo; undo
# choose 2 → [2] record; undo
# → [[],[1],[1,2],[2]]
List<List<Integer>> subsets(int[] nums) {
    List<List<Integer>> out = new ArrayList<>();
    List<Integer> path = new ArrayList<>();
    dfs(nums, 0, path, out);
    return out;
}
void dfs(int[] nums, int start, List<Integer> path, List<List<Integer>> out) {
    out.add(new ArrayList<>(path));          // record every node
    for (int i = start; i < nums.length; i++) {
        path.add(nums[i]);                   // choose
        dfs(nums, i + 1, path, out);         // explore
        path.remove(path.size() - 1);        // undo
    }
}

// nums=[1,2]
// path=[] → record []
// choose 1 → [1] record; choose 2 → [1,2] record; undo; undo
// choose 2 → [2] record; undo
// → [[],[1],[1,2],[2]]

Permutations — used[] mask:

def permute(nums: list[int]) -> list[list[int]]:
    out, path = [], []
    used = [False] * len(nums)

    def dfs():
        if len(path) == len(nums):
            out.append(path.copy())
            return
        for i, x in enumerate(nums):
            if used[i]:
                continue
            used[i] = True
            path.append(x)
            dfs()
            path.pop()
            used[i] = False

    dfs()
    return out
List<List<Integer>> permute(int[] nums) {
    List<List<Integer>> out = new ArrayList<>();
    List<Integer> path = new ArrayList<>();
    boolean[] used = new boolean[nums.length];
    dfs(nums, used, path, out);
    return out;
}
void dfs(int[] nums, boolean[] used, List<Integer> path, List<List<Integer>> out) {
    if (path.size() == nums.length) {
        out.add(new ArrayList<>(path));
        return;
    }
    for (int i = 0; i < nums.length; i++) {
        if (used[i]) continue;
        used[i] = true;
        path.add(nums[i]);
        dfs(nums, used, path, out);
        path.remove(path.size() - 1);
        used[i] = false;
    }
}

Edge cases & common bugs

Complexity — say it aloud

Interview talk track

You: ‘I’ll DFS decisions: choose an element, recurse from the next index, then undo. I record a copy of the path at every node for subsets. Complexity is O(n·2ⁿ) because we emit each subset.’

Practice set

  • Subsets / Subsets II
  • Permutations / Permutations II
  • Combination Sum / Combination Sum II
  • Letter Combinations of a Phone Number
  • Generate Parentheses
  • N-Queens
  • Word Search
  • Palindrome Partitioning

Harder follow-up

Harder variant: Combination Sum with unlimited reuse — pass i (not i+1) so the same index can be reused; prune when remaining sum < 0. Speak the reuse rule so they don’t think it’s a bug.

Pattern bank: more questions + efficient solutions

Five backtracking templates. Always undo mutations; prune when the partial candidate is illegal.

Q: Subsets: return all subsets of a distinct integer array (power set).

# Subsets — O(n * 2^n)
def subsets(nums):
    out = []
    path = []
    def dfs(i):
        if i == len(nums):
            out.append(path[:])
            return
        dfs(i + 1)                 # skip
        path.append(nums[i])       # take
        dfs(i + 1)
        path.pop()
    dfs(0)
    return out
// Subsets — O(n * 2^n)
List<List<Integer>> subsets(int[] nums) {
    List<List<Integer>> out = new ArrayList<>();
    dfs(nums, 0, new ArrayList<>(), out);
    return out;
}
void dfs(int[] nums, int i, List<Integer> path, List<List<Integer>> out) {
    if (i == nums.length) { out.add(new ArrayList<>(path)); return; }
    dfs(nums, i + 1, path, out);
    path.add(nums[i]);
    dfs(nums, i + 1, path, out);
    path.remove(path.size() - 1);
}

Q: Combination Sum: given distinct candidates and a target, return all unique combinations that sum to target (reuse allowed; order within a combo does not matter).

# Combination Sum — backtrack with reuse
def combination_sum(candidates, target: int):
    out, path = [], []
    candidates = sorted(candidates)
    def dfs(start, remain):
        if remain == 0:
            out.append(path[:])
            return
        for i in range(start, len(candidates)):
            c = candidates[i]
            if c > remain:
                break
            path.append(c)
            dfs(i, remain - c)   # reuse i
            path.pop()
    dfs(0, target)
    return out
// Combination Sum — backtrack with reuse
List<List<Integer>> combinationSum(int[] candidates, int target) {
    Arrays.sort(candidates);
    List<List<Integer>> out = new ArrayList<>();
    dfs(candidates, 0, target, new ArrayList<>(), out);
    return out;
}
void dfs(int[] a, int start, int remain, List<Integer> path, List<List<Integer>> out) {
    if (remain == 0) { out.add(new ArrayList<>(path)); return; }
    for (int i = start; i < a.length; i++) {
        if (a[i] > remain) break;
        path.add(a[i]);
        dfs(a, i, remain - a[i], path, out);
        path.remove(path.size() - 1);
    }
}

Q: Permutations: return all permutations of a distinct integer array.

# Permutations — O(n * n!)
def permute(nums):
    out = []
    def dfs(i):
        if i == len(nums):
            out.append(nums[:])
            return
        for j in range(i, len(nums)):
            nums[i], nums[j] = nums[j], nums[i]
            dfs(i + 1)
            nums[i], nums[j] = nums[j], nums[i]
    dfs(0)
    return out
// Permutations — O(n * n!)
List<List<Integer>> permute(int[] nums) {
    List<List<Integer>> out = new ArrayList<>();
    dfs(nums, 0, out);
    return out;
}
void dfs(int[] nums, int i, List<List<Integer>> out) {
    if (i == nums.length) {
        List<Integer> row = new ArrayList<>();
        for (int x : nums) row.add(x);
        out.add(row);
        return;
    }
    for (int j = i; j < nums.length; j++) {
        swap(nums, i, j);
        dfs(nums, i + 1, out);
        swap(nums, i, j);
    }
}
void swap(int[] a, int i, int j) { int t = a[i]; a[i] = a[j]; a[j] = t; }

Q: N-Queens: place n queens on an n×n board so no two attack; return all distinct board configurations (or count solutions).

# N-Queens — backtrack with col/diag sets
def solve_n_queens(n: int):
    cols, diag, anti = set(), set(), set()
    board = [["."] * n for _ in range(n)]
    out = []
    def dfs(r):
        if r == n:
            out.append(["".join(row) for row in board])
            return
        for c in range(n):
            if c in cols or (r - c) in diag or (r + c) in anti:
                continue
            cols.add(c); diag.add(r - c); anti.add(r + c)
            board[r][c] = "Q"
            dfs(r + 1)
            board[r][c] = "."
            cols.remove(c); diag.remove(r - c); anti.remove(r + c)
    dfs(0)
    return out
// N-Queens — backtrack with col/diag sets
List<List<String>> solveNQueens(int n) {
    List<List<String>> out = new ArrayList<>();
    char[][] board = new char[n][n];
    for (char[] row : board) Arrays.fill(row, '.');
    dfs(0, n, new HashSet<>(), new HashSet<>(), new HashSet<>(), board, out);
    return out;
}
void dfs(int r, int n, Set<Integer> cols, Set<Integer> diag, Set<Integer> anti,
         char[][] board, List<List<String>> out) {
    if (r == n) {
        List<String> snap = new ArrayList<>();
        for (char[] row : board) snap.add(new String(row));
        out.add(snap);
        return;
    }
    for (int c = 0; c < n; c++) {
        if (cols.contains(c) || diag.contains(r - c) || anti.contains(r + c)) continue;
        cols.add(c); diag.add(r - c); anti.add(r + c);
        board[r][c] = 'Q';
        dfs(r + 1, n, cols, diag, anti, board, out);
        board[r][c] = '.';
        cols.remove(c); diag.remove(r - c); anti.remove(r + c);
    }
}

Q: Palindrome Partitioning: partition a string so every substring is a palindrome; return all possible partitions.

# Palindrome Partitioning
def partition(s: str):
    out, path = [], []
    n = len(s)
    def is_pal(l, r):
        while l < r:
            if s[l] != s[r]:
                return False
            l += 1
            r -= 1
        return True
    def dfs(start):
        if start == n:
            out.append(path[:])
            return
        for i in range(start, n):
            if is_pal(start, i):
                path.append(s[start : i + 1])
                dfs(i + 1)
                path.pop()
    dfs(0)
    return out
// Palindrome Partitioning
List<List<String>> partition(String s) {
    List<List<String>> out = new ArrayList<>();
    dfs(s, 0, new ArrayList<>(), out);
    return out;
}
void dfs(String s, int start, List<String> path, List<List<String>> out) {
    if (start == s.length()) { out.add(new ArrayList<>(path)); return; }
    for (int i = start; i < s.length(); i++) {
        if (isPal(s, start, i)) {
            path.add(s.substring(start, i + 1));
            dfs(s, i + 1, path, out);
            path.remove(path.size() - 1);
        }
    }
}
boolean isPal(String s, int l, int r) {
    while (l < r) if (s.charAt(l++) != s.charAt(r--)) return false;
    return true;
}

45-minute pattern drill

← Lattice