Trees & Advanced Structures

Master binary trees, heaps, and tries for hierarchical data manipulation

Last generated

Lesson 19 of 23 available20 practice questions

SPACED REPETITION · 20 practice questions

Make this lesson stick.

Try 3 questions now. No account needed. Sample answers aren't saved.

Keep the answer ready while the data changes

Foundation: Arrays & Strings ended with a prefix-sum table: prepare it once in O(n), and every range sum costs one subtraction. It also ended with the table's weakness. Change one value and every later entry is stale.

Here is a workload where that weakness bites. An array holds 100,000 numbers. Then 200,000 operations arrive, interleaved: half are updates ("set nums[i] to v") and half are range-sum queries.

  • Scan each range. Updates are free, but one query can add up 100,000 values. 100,000 queries: up to 10¹⁰ additions.
  • Keep prefix sums. Queries are free, but one update can rewrite 100,000 entries. 100,000 updates: up to 10¹⁰ writes.

Each plan makes one operation free and the other ruinous. The way out is to store sums of pieces: the whole array, its two halves, their halves, and so on down to single values. The pieces form a tree of height 17 (2¹⁷ = 131,072 ≥ 100,000). An update changes one value, so it spoils only the pieces that contain it, one per level. A query glues together a few dozen pieces. The 200,000 operations cost under ten million steps instead of ten billion.

That is the big idea of this lesson: arrange what you store as a tree, so that any single operation touches only a few nodes on each level: one root-to-leaf path for a BST or trie lookup, at most two boundary paths for a segment-tree query. The operation then costs O(height). Every structure below answers two questions: what does each node remember, and how do we keep the height small?

Smell Structure Each node remembers Signature problem
Sorted order must survive inserts and deletes Binary search tree one key; smaller keys go left, larger go right Delete Node in a BST
Sorted input turns that tree into a chain Balanced BST a key plus a little balance information Convert Sorted Array to BST
Many strings, questions about prefixes Trie one character step, plus "a word ends here" Implement Trie
Range sum, min or max mixed with updates Segment tree the answer for one range of the array Range Sum Query – Mutable
Only sums (or XOR) mixed with updates Fenwick tree the sum of a block sized by the index's lowest set bit Range Sum Query – Mutable, in half the code

You should be comfortable with Foundation: Arrays & Strings, especially prefix sums and "what costs what", and with recursion from Advanced Recursion & Backtracking: base cases, trusting the recursive call, Python's recursion limit. This lesson builds the structures. Walking a tree in every order and the classic recursive patterns (depth, diameter, path sums, validating a BST, lowest common ancestor) come next, in Binary Tree Patterns. Heaps get one short section here and their own lesson, Heaps & Priority Queues.

Trees in five minutes: words, height, and how LeetCode writes them

LeetCode gives you binary trees as linked nodes:

class TreeNode:
    def __init__(self, val=0, left=None, right=None):
        self.val = val
        self.left = left      # a TreeNode, or None
        self.right = right

A binary tree is either empty (None) or a node holding a value and two binary trees: its left and right subtrees. That recursive definition is why tree code is so often recursive. Handle None, then trust the call on each child.

Here is the tree we will reuse all lesson:

              8               depth 0  (the root)
            /   \
           3     10           depth 1
          / \      \
         1   6      14        depth 2
            / \     /
           4   7   13         depth 3
  • The root is 8, the only node without a parent. A leaf has no children: 1, 4, 7 and 13. Every other node is internal.
  • 3 is the parent of 1 and 6. The ancestors of 7 are 6, 3 and 8. The descendants of 3 are 1, 6, 4 and 7.
  • The depth of a node counts the edges from the root down to it: 7 has depth 3.
  • The height of a node counts the edges on the longest path from it down to a leaf. A leaf has height 0 and 6 has height 1. The height of the tree is the height of its root: 3.

Two ways to count

Predict first: LeetCode 104: Maximum Depth of Binary Tree asks for this tree's "maximum depth". Is the answer 3?

Check your answer

No, it is 4. LeetCode counts nodes on the longest root-to-leaf path, and 8 → 3 → 6 → 4 has four of them. The edge count is 3. Both conventions are common, and mixing them is the classic off-by-one of tree problems.

The two functions differ only in what the empty tree returns:

def height(node):             # edges: empty tree -1, single node 0
    if node is None:
        return -1
    return 1 + max(height(node.left), height(node.right))

def max_depth(node):          # nodes, as LeetCode 104 counts: empty 0, single node 1
    if node is None:
        return 0
    return 1 + max(max_depth(node.left), max_depth(node.right))

On the example they return 3 and 4. Neither is wrong; they answer different questions. This lesson's convention: "height" and "depth" count edges, and "maximum depth" means LeetCode's node count. When a problem defines its own terms, its definition wins, so check it against the examples. Complexity bounds below say O(h), where h is the height. The longest path holds h + 1 nodes, which changes nothing in big-O but everything in a return value.

⚠️ The base case decides the convention. Return 0 for None and you count nodes; return −1 and you count edges. Keep every helper in one solution on the same convention.

Shapes, and the two numbers that matter

  perfect          complete         full             chain
     1                1                1              1
    / \              / \              / \              \
   2   3            2   3            2   3              2
  / \ / \          /                    / \              \
 4  5 6  7        4                    6   7              3
  • Perfect: every level is full. A perfect tree of height h has 2ʰ⁺¹ − 1 nodes.
  • Complete: every level is full except possibly the last, which fills from the left. Here node 2 has only a left child, so the tree is complete but not full.
  • Full: every node has 0 or 2 children. This one is full but not complete, because the last level has a gap on the left.
  • Chain (degenerate): one child per node, a linked list wearing a tree costume. n nodes, height n − 1.

So a binary tree with n nodes has a height somewhere between ⌊log₂ n⌋ and n − 1. For a million nodes that is 19 at best and 999,999 at worst. An operation that walks one root-to-leaf path costs O(h), so the same code is instant or hopeless depending on the shape.

How LeetCode writes a tree

LeetCode lists a tree in level order, left to right, with null (Python None) for a missing child. Our tree is [8,3,10,1,6,null,14,null,null,4,7,13]. Only real nodes get child slots, and trailing nulls are dropped. When a test fails, rebuild the tree locally:

from collections import deque

def build_tree(values):
    """LeetCode level-order list -> root. None marks a missing child."""
    if not values or values[0] is None:
        return None
    root = TreeNode(values[0])
    queue = deque([root])
    i = 1
    while queue and i < len(values):
        node = queue.popleft()
        if i < len(values) and values[i] is not None:
            node.left = TreeNode(values[i])
            queue.append(node.left)
        i += 1
        if i < len(values) and values[i] is not None:
            node.right = TreeNode(values[i])
            queue.append(node.right)
        i += 1
    return root

The queue hands out the real nodes in level order, and each one takes the next two list entries as its children.

Complete trees need no pointers at all. Number the nodes 0, 1, 2, … in level order and store them in a plain list. The children of index i sit at 2i + 1 and 2i + 2, and its parent at (i − 1) // 2. A complete tree has no gaps, so no slot is wasted. Heaps use this layout, and so will our segment tree.

index:  0  1  2  3  4  5  6
value:  1  2  3  4  5  6  7          the perfect tree above
children of index 1 (value 2): 3 and 4   -> values 4 and 5
parent of index 5 (value 6):   (5-1)//2 = 2 -> value 3

Your turn: draw [1,null,2,3]. What are its height and its maximum depth, and which node is the only leaf?

Check your answer

1 has no left child and has 2 as its right child. The next entry, 3, is the first child slot of 2, so 3 is 2's left child. The shape is 1 → (right) 2 → (left) 3: height 2, maximum depth 3, and 3 is the only leaf. A common misreading hangs 3 under 1, but 1's two slots were already used by null and 2.

Structure 1: Binary search trees

Smell: the data must stay sorted while you insert and delete, and you keep asking order questions. Is x present? What is the smallest value ≥ x?

A sorted Python list answers those in O(log n) with bisect. Inserts ruin it: list.insert shifts every later element, which is O(n) per insert. If 100,000 inserts each land at the front, they shift 0 + 1 + … + 99,999, close to five billion entries. A binary search tree keeps the order in its links instead of in positions, so an insert only attaches a node.

The BST property: for every node, every key in its left subtree is smaller than its key, and every key in its right subtree is larger. That means every key in the subtree, not just the children. The example tree is a BST. Everything under 3 is smaller than 8, everything under 10 is larger, and the rule holds again inside each subtree.

This lesson stores distinct keys, like a set. Inserting a key that is already present changes nothing. If you need duplicates, keep a count in the node instead of adding a second node, and the strict "smaller" and "larger" stay true everywhere.

Predict first: is [10,5,15,null,null,6,20] a BST? Every parent–child pair is in order: 5 < 10 < 15 and 6 < 15 < 20.

Check your answer

No. 6 sits in 10's right subtree, so it would have to be larger than 10. Checking parent–child pairs misses a clash between a node and a distant ancestor. A correct check carries the allowed range down the tree; Binary Tree Patterns builds it. Here, remember what the property promises: a whole subtree lies on one side.

Search: twenty questions on a fixed script

This is LeetCode 700: Search in a Binary Search Tree:

def search(root, target):
    node = root
    while node is not None and node.val != target:
        node = node.left if target < node.val else node.right
    return node                     # the node holding target, or None
target nodes visited decisions result
6 8, 3, 6 6 < 8 left; 6 > 3 right; equal the node 6
11 8, 10, 14, 13 right; right; left; left; then None None

Why it works. The loop keeps one sentence true: if target is in the tree, it is in the subtree rooted at node. At the start that subtree is the whole tree. When target < node.val, nothing in the right subtree can equal it, and neither can the node itself, so stepping left keeps the sentence true. When node becomes None the subtree is empty, so target is absent. Every step goes one level down: O(h) time and O(1) extra space.

It is the game of twenty questions: each comparison throws away a whole subtree. The analogy stops at who picks the questions. A good player asks the question that halves what is left. A BST asks whatever its shape dictates, and a chain asks "is it 1? is it 2? is it 3?".

Insert: search, then attach where you fell off

def insert(root, val):
    """Insert val (ignore it if present). Return the root of the tree."""
    if root is None:
        return TreeNode(val)            # the empty spot where val belongs
    if val < root.val:
        root.left = insert(root.left, val)
    elif val > root.val:
        root.right = insert(root.right, val)
    return root                         # equal: already present

That solves LeetCode 701: Insert into a Binary Search Tree. Inserting 5 into the example walks 8 → 3 → 6 → 4 and attaches 5 as the right child of 4. A new key always becomes a leaf, and nothing already in the tree moves.

The line root.left = insert(root.left, val) is a pattern worth memorizing: the recursive call returns the root of the subtree it was given, and the caller re-links it. The assignment matters in the call whose child is None: it links in the node that the child call creates and returns. Drop it and the new node is built, returned and thrown away. Deletion uses the same pattern.

Delete: three cases, one trick

A leaf is easy: unhook it. A node with one child: put the child in its place. The interesting case is a node with two children, such as 3. Its replacement must be larger than everything in its left subtree and smaller than everything in its right subtree.

Two keys qualify: the largest key on the left (the predecessor) and the smallest key on the right (the successor). This code uses the successor: step right once, then left as far as possible. It solves LeetCode 450: Delete Node in a BST.

def delete(root, val):
    """Delete val if present. Return the root of the tree."""
    if root is None:
        return None                               # val is not in the tree
    if val < root.val:
        root.left = delete(root.left, val)
    elif val > root.val:
        root.right = delete(root.right, val)
    else:
        if root.left is None:                     # zero or one child: splice
            return root.right
        if root.right is None:
            return root.left
        succ = root.right                         # two children: find the successor
        while succ.left is not None:
            succ = succ.left
        root.val = succ.val                       # copy it up...
        root.right = delete(root.right, succ.val) # ...and delete it below
    return root

Delete 3. Its successor is 4: right to 6, then left to 4. The node that held 3 now holds 4, and the old 4 is deleted from the right subtree. That second deletion is always an easy case, because the successor has no left child (a left child would be even smaller).

before                     after delete(root, 3)
       8                          8
     /   \                      /   \
    3     10                   4     10
   / \      \                 / \      \
  1   6      14              1   6      14
     / \     /                    \     /
    4   7   13                     7   13

Why it is still a BST. 4 was the smallest key in 3's right subtree. So it is larger than everything on the left (all of which is below 3) and smaller than everything that stays on the right. Cost: one walk down to find the node and one more to its successor, O(h).

The one traversal you need now

Visiting the left subtree, then the node, then the right subtree is an in-order traversal. On a BST it lists the keys in sorted order, because the property holds at every node. That makes it the easiest way to test BST code:

def inorder(node, out):
    if node is not None:
        inorder(node.left, out)
        out.append(node.val)
        inorder(node.right, out)
    return out

inorder(root, []) on the example gives [1, 3, 4, 6, 7, 8, 10, 13, 14]. The other orders and what they are for are in Binary Tree Patterns.

The question a hash set can't answer

A hash set tells you whether 11 is present. A BST also tells you what lies next to 11: the smallest key ≥ 11 (its ceiling) or the largest key ≤ 11 (its floor). Search for 11 and remember the best candidate you pass:

def ceiling(root, x):
    """Smallest key >= x, or None."""
    best = None
    node = root
    while node is not None:
        if node.val >= x:
            best = node.val         # a candidate; anything better is to its left
            node = node.left
        else:
            node = node.right       # too small, and so is its whole left subtree
    return best
node compared with 11 best afterward next move
8 smaller None right
10 smaller None right
14 ≥ 11 14 left
13 ≥ 11 13 left, reaches None, stop

The result is 13. The invariant: the answer is either best or somewhere in the subtree we are about to enter. A node ≥ x is a candidate, and a better one (smaller but still ≥ x) can only be on its left. A node < x rules out itself and its whole left subtree. The cost is O(h). "What is next to x?" is the everyday reason to reach for an ordered structure: bookings, leaderboards, closest-value problems.

Your turn: start from an empty tree and insert 50, 30, 70, 20, 40, 60, 80, 65. Then delete 50. What is the new root, where does 65 end up, and what does ceiling(root, 66) return?

Check your answer

Before the delete, 65 is the right child of 60, which is the left child of 70. Deleting 50 finds its successor by stepping right to 70 and then left to 60. The root becomes 60. Then 60 is deleted from the right subtree. It has one child, 65, which takes its place: 65 becomes the left child of 70. In-order is [20, 30, 40, 60, 65, 70, 80], and ceiling(root, 66) is 70.

Test it like a professional. Compare your tree against a Python set on many small random operations. Small keys force the cases that break BST code: duplicates, deleting absent keys, deleting the root.

import random

def check_bst(trials=500):
    for _ in range(trials):
        root, model = None, set()
        for _ in range(random.randint(0, 30)):
            k = random.randint(-10, 10)
            if random.random() < 0.6:
                root = insert(root, k)
                model.add(k)
            else:
                root = delete(root, k)
                model.discard(k)
            assert inorder(root, []) == sorted(model)
            x = random.randint(-12, 12)
            bigger = [v for v in model if v >= x]
            assert ceiling(root, x) == (min(bigger) if bigger else None)
    print("all matched")

Plot twist: sorted input builds a linked list

Insert 1, 2, 3, 4, 5 into an empty BST. Each key is larger than everything before it, so each one goes right, right, right: the chain from the shapes diagram, height n − 1.

Count the cost. Inserting 1 to n in order, key k walks past k − 1 nodes, for a total of 0 + 1 + … + (n − 1) = n(n − 1)/2 comparisons. For n = 100,000 that is 4,999,950,000: the Foundation lesson's five-billion all-pairs check, now hidden inside a "logarithmic" structure. Random order is kind. In three runs with 100,000 shuffled keys, the heights were 36, 39 and 40, against an ideal of 16. But real inputs are often sorted or nearly sorted: timestamps, IDs, and test cases written to break you.

Predict first: a recursive insert builds a tree from the keys 0, 1, 2, …, 4999 in that order. Python's default recursion limit is 1,000. What happens?

Check your answer

RecursionError, on key 999 in our run. insert recurses once per level, and every new key makes the chain one level deeper. The same code on 5,000 shuffled keys is fine: in ten runs the height stayed between 25 and 31. Recursive tree code is only as safe as the tree is short.

Two fixes, and you want both: walk down with a loop where that is easy, and keep the tree short.

An iterative insert
def insert_iterative(root, val):
    new = TreeNode(val)
    if root is None:
        return new
    node = root
    while True:
        if val < node.val:
            if node.left is None:
                node.left = new
                return root
            node = node.left
        elif val > node.val:
            if node.right is None:
                node.right = new
                return root
            node = node.right
        else:
            return root                 # already present

It uses O(1) extra space at any height. The O(h) time is still there: on sorted input it survives, but it is still slow.

Rotations: the move that restores balance

A rotation re-hangs three subtrees around two nodes. Here x moves up and y moves down:

        y                         x
       / \     rotate_right(y)   / \
      x   C    ------------->   A   y
     / \       <-------------      / \
    A   B      rotate_left(x)     B   C
def rotate_right(y):
    x = y.left
    y.left = x.right      # B moves across: its keys lie between x and y
    x.right = y
    return x              # the new root of this subtree

def rotate_left(x):
    y = x.right
    x.right = y.left
    y.left = x
    return y

Read both pictures in order: A, x, B, y, C, then A, x, B, y, C. A rotation never changes the in-order sequence, so the BST property survives. Of the three subtrees, only B changes parent, and its keys are exactly the ones between x and y. What does change is height: A's side rises one level and C's side sinks one.

A left-leaning chain 3 → 2 → 1 becomes 2 with children 1 and 3 after rotate_right at 3. Its height drops from 2 to 1. A zig-zag (3 with left child 1, and 1 with right child 2) needs two rotations: rotate_left at 1 straightens it into the chain, then rotate_right at 3 finishes the job.

Self-balancing trees: what they promise

You will rarely write one in an interview, but you should know what they guarantee. Both kinds below check a local rule after each insert or delete and repair it with rotations:

AVL tree Red-black tree
Rule at every node, the two subtree heights differ by at most 1 no red node has a red child, and every path from a node down to None passes the same number of black nodes
Height bound about 1.44 log₂ n at most 2 log₂(n + 1)
Rotations per insert at most one single or double rotation at most 2
Rotations per delete possibly one at each of O(log n) ancestors at most 3, plus recolorings
Where you meet it lookup-heavy in-memory sets Java TreeMap, most C++ std::map implementations, the Linux kernel's CPU scheduler

Both give O(log n) search, insert and delete in the worst case. For n = 100,000 the height bounds are about 24 for AVL and 33 for red-black, against a height of 99,999 for the chain.

B-trees take the same idea to disk. A node holds hundreds of sorted keys and has hundreds of children, sized to fill one disk page, so a few levels cover hundreds of millions of keys. Every level costs one page read, which is why database indexes, such as PostgreSQL's default index and MySQL InnoDB's, are B-tree variants rather than binary trees.

In Python there is no built-in balanced BST. Your options:

  • bisect on a sorted list: O(log n) search, but O(n) insert and delete. Fine when n is small or inserts are rare.
  • sortedcontainers.SortedList, a third-party package that LeetCode's Python environment provides. It adds, removes and bisects in about O(log n). Inside it is a list of sorted sublists, not a tree, but it answers the same order questions.
  • Your own BST, when the problem is to build one.

Your turn: given a sorted list, build a BST of minimum height. This is LeetCode 108: Convert Sorted Array to Binary Search Tree. Which key should be the root so that half the keys go each way?

Check your answer

The middle one, and then the same rule inside each half:

def sorted_to_bst(nums):
    def build(lo, hi):                  # builds nums[lo..hi], inclusive
        if lo > hi:
            return None
        mid = (lo + hi) // 2
        return TreeNode(nums[mid], build(lo, mid - 1), build(mid + 1, hi))
    return build(0, len(nums) - 1)

The two halves differ in size by at most one, so the height is ⌊log₂ n⌋ (checked for every n from 1 to 199). Passing indices instead of slices matters. nums[:mid] copies, and copying about n elements on each of the log₂ n levels costs O(n log n). With indices, each key becomes one node: O(n) time, O(log n) recursion stack, plus the O(n) nodes you return. When all the keys are known up front, build balanced directly and never rotate.

The cousin that only knows its minimum

A binary heap is a complete binary tree in the array layout from earlier, with a weaker rule than a BST: every parent is ≤ its children (in a min-heap). There is no left/right order, only "parents first". That is enough to keep the minimum at index 0, and not enough for anything else:

Question Min-heap Balanced BST
What is the smallest? O(1) O(log n)
Insert, or remove the smallest O(log n) O(log n)
Is x present? O(n) O(log n)
Smallest key ≥ x? O(n) O(log n)
All keys in sorted order O(n log n), by popping O(n), by in-order

A heap is a plain Python list driven by heapq: no pointers, small constants, and O(n) to build from a list. If the only question is "what is the smallest (or largest) right now?", use a heap. As soon as you need "what is next to x?", you need order: a BST or a sorted list. heapq is a min-heap; for a max-heap, push negated keys, or on Python 3.14+ use heapq.heappush_max and its siblings. Sifting, heapify, top-k and the two-heap median are in Heaps & Priority Queues.

A node can remember more than its key. Store in each BST node the size of its subtree, and you can find the k-th smallest key in O(h). If the left subtree holds exactly k − 1 keys, the answer is this node. If it holds more, go left. Otherwise go right and look for the (k − left size − 1)-th key there. Insert and delete fix the sizes along the path they walk, which is still O(h). Keep that idea, a node that stores a summary of its subtree, in mind: segment trees are built entirely out of it.

Structure 2: Tries

Smell: many strings, and the questions are about prefixes. Does any word start with "compu"? Which words do? Or you extend a string one character at a time, and after each character you must know whether you are still on the way to some word.

A hash set of 100,000 words answers "is 'computer' a word?" in expected O(L) for a word of length L. It cannot answer "does any word start with 'compu'?" without scanning all the words or storing every prefix of every word as its own string. A sorted list with bisect can find the first word ≥ "compu" and check it, in O(L log n) per question. The trie's real advantage is incremental: going from "comp" to "compu" is one step down, one dictionary lookup, no matter how long the prefix already is.

A trie stores the words as a tree of characters, so words that share a prefix share its nodes:

             (root)           words: car, cart, cat, do, dog
            /      \          * marks is_word = True
           c        d
           |        |
           a        o*
          / \       |
         r*  t*     g*
         |
         t*

Every root-to-node path spells a prefix, and the root is the empty prefix. Whether a word ends at a node is a separate fact, stored in the node's is_word flag. "car" and "cart" share three nodes, and "do" ends at a node that also has a child. The five words hold 15 characters in 8 nodes below the root.

Predict first: a trie holds only "apple". What do search("app") and starts_with("app") return?

Check your answer

False and True. Walking a, p, p succeeds in both cases, because those nodes exist for "apple". The difference is the flag: the node after the second p has is_word = False. A trie without the flag cannot tell a word from the prefix of a word.

This is LeetCode 208: Implement Trie (Prefix Tree), where the last method is called startsWith:

class TrieNode:
    def __init__(self):
        self.children = {}      # character -> TrieNode
        self.is_word = False    # does a stored word end exactly here?

class Trie:
    def __init__(self):
        self.root = TrieNode()  # the empty prefix

    def insert(self, word):
        node = self.root
        for ch in word:
            if ch not in node.children:
                node.children[ch] = TrieNode()
            node = node.children[ch]
        node.is_word = True

    def find_node(self, s):
        """The node reached by spelling s, or None."""
        node = self.root
        for ch in s:
            node = node.children.get(ch)
            if node is None:
                return None
        return node

    def search(self, word):
        node = self.find_node(word)
        return node is not None and node.is_word

    def starts_with(self, prefix):
        return self.find_node(prefix) is not None

Cost. Each operation takes one step per character, so O(L) time for a string of length L (each step is an expected O(1) dictionary lookup). Space is at most one node per inserted character, O(total characters), and less when words share prefixes. For lowercase words a common variant stores children as a 26-slot list: indexing is faster, but every node carries 26 slots, most of them empty.

The characters do not have to be letters. Insert each integer as its 31 bits, highest bit first, and you get a binary trie of height 31, which can answer "which stored number agrees with x on the longest run of leading bits?" in 31 steps.

Autocomplete: walk, then collect

def words_with_prefix(trie, prefix):
    """All stored words that start with prefix, in sorted order."""
    out = []
    def collect(node, path):
        if node.is_word:
            out.append("".join(path))
        for ch in sorted(node.children):
            path.append(ch)
            collect(node.children[ch], path)
            path.pop()                  # undo, so path is right for the next child
    start = trie.find_node(prefix)
    if start is not None:
        collect(start, list(prefix))
    return out

On the example trie, the prefix "ca" gives ['car', 'cart', 'cat'] and "d" gives ['do', 'dog']. Cost: O(p) to reach the prefix node, then one visit per node below it. Reaching "compu" takes five steps however big the dictionary is, but listing its completions costs as much as the completions themselves. "Autocomplete in O(prefix length)" is true only of finding the spot.

A node can remember a summary here too. Store in each node how many words pass through it, and "how many words start with p?" takes O(p) with no collecting. To delete a word, clear the flag at its last node; you may also remove, on the way back up, nodes that are left with no children and no flag.

Your turn: LeetCode 648: Replace Words. Given the roots ["cat", "bat", "rat"], replace every word of "the cattle was rattled by the battery" with the shortest root that is a prefix of it, keeping words that have none. Write shortest_root(trie, word).

Check your answer

Walk the word through the trie and stop at the first node whose flag is set:

def shortest_root(trie, word):
    node = trie.root
    for i, ch in enumerate(word):
        node = node.children.get(ch)
        if node is None:
            return word                 # no root is a prefix of this word
        if node.is_word:
            return word[:i + 1]         # the first flag is the shortest root
    return word

The sentence becomes "the cat was rat by the bat". Each word costs O(its length) whatever the number of roots. Checking every root against every word would multiply the two.

Boss level: Word Search II

LeetCode 212: Word Search II. Given a grid of letters and a list of words, return every word that can be spelled along a path of horizontally or vertically adjacent cells, using each cell at most once per word.

board:  o a a n        words: oath, pea, eat, rain
        e t a e
        i h k r        answer: oath, eat
        i f l v

Searching once per word repeats the same grid walks for every word with the same start. Constraint Satisfaction solves the one-word version with backtracking. A trie lets one walk serve every word at once: the search carries a trie node, and each step to a neighboring cell is one step down the trie. If the trie has no child for the next letter, no word continues this way, and the walk stops.

def find_words(board, words):
    trie = Trie()
    for w in words:
        trie.insert(w)
    rows, cols = len(board), len(board[0])
    found = []

    def dfs(r, c, node, path):
        ch = board[r][c]
        child = node.children.get(ch)
        if child is None:
            return                          # no word starts with path + ch
        path.append(ch)
        if child.is_word:
            found.append("".join(path))
            child.is_word = False           # report each word once
        board[r][c] = "#"                   # in use on the current path
        for nr, nc in ((r + 1, c), (r - 1, c), (r, c + 1), (r, c - 1)):
            if 0 <= nr < rows and 0 <= nc < cols and board[nr][nc] != "#":
                dfs(nr, nc, child, path)
        board[r][c] = ch                    # free the cell for other paths
        path.pop()
        if not child.children and not child.is_word:
            del node.children[ch]           # nothing left to find below: prune it

    for r in range(rows):
        for c in range(cols):
            dfs(r, c, trie.root, [])
    return found

Four lines carry the solution:

  • if child is None: return is the pruning. The prefix is dead, so stop.
  • child.is_word = False after reporting. A word can often be spelled along several paths, and clearing its flag reports it once.
  • board[r][c] = "#", restored after the loop, marks the cell as used on the current path only: choose, explore, unchoose, as in Advanced Recursion & Backtracking.
  • del node.children[ch] removes a trie node once every word below it has been found. Without it, a board of a single repeated letter keeps re-walking branches that can no longer produce anything: a 12 × 12 board of as with the words a, aa, …, a¹⁰ took about a second without this line and well under a millisecond with it.

Cost. Building the trie is O(total characters in the words). From each of the R·C cells, a walk can branch 4 ways at its first step and at most 3 after that, since it never steps back onto its own path. With L the length of the longest word, the worst case is O(R·C·4·3ᴸ⁻¹). Separate searches pay a bill like that once per word; the shared walk pays it once, and the two pruning lines usually cut a path after a letter or two. Extra space: the trie, plus a recursion depth of about L.

Structure 3: Segment trees

Smell: range questions (sum, min, max over nums[l..r]) keep arriving while the array keeps changing.

The Foundation lesson's prefix sums fail here twice. An update makes later entries stale, and a range maximum can't be recovered from two prefix readings, because max has no inverse to subtract with. A segment tree fixes both. Each node stores the answer for one range, and a query combines stored answers instead of subtracting them. So any operation where grouping doesn't matter (sum, min, max, gcd) works.

nums = [1, 3, 5, 7, 9, 11]          each box: [range] sum

                   [0..5] 36
                 /           \
         [0..2] 9             [3..5] 27
          /     \              /      \
    [0..1] 4   [2] 5     [3..4] 16   [5] 11
     /    \                /    \
  [0] 1  [1] 3          [3] 7  [4] 9

The root covers everything. Each node splits its range at mid = (lo + hi) // 2 into [lo..mid] and [mid+1..hi], and each leaf holds one value. There are 2n − 1 nodes, and the height is ⌈log₂ n⌉.

Predict first: which boxes does the query [1..4] (3 + 5 + 7 + 9) need?

Check your answer

Three: [1] 3, [2] 5 and [3..4] 16, for a total of 24. No box covers [1..4] exactly, but every range can be tiled by boxes, and the query finds the tiling by starting at the root.

class SegmentTree:
    """Range sums with point updates. Node i has children 2i+1 and 2i+2."""
    def __init__(self, nums):
        self.n = len(nums)
        self.tree = [0] * (4 * self.n)
        if self.n:
            self._build(nums, 0, 0, self.n - 1)

    def _build(self, nums, node, lo, hi):
        if lo == hi:
            self.tree[node] = nums[lo]
            return
        mid = (lo + hi) // 2
        self._build(nums, 2 * node + 1, lo, mid)
        self._build(nums, 2 * node + 2, mid + 1, hi)
        self.tree[node] = self.tree[2 * node + 1] + self.tree[2 * node + 2]

    def update(self, i, value):                  # nums[i] = value
        self._update(0, 0, self.n - 1, i, value)

    def _update(self, node, lo, hi, i, value):
        if lo == hi:
            self.tree[node] = value
            return
        mid = (lo + hi) // 2
        if i <= mid:
            self._update(2 * node + 1, lo, mid, i, value)
        else:
            self._update(2 * node + 2, mid + 1, hi, i, value)
        self.tree[node] = self.tree[2 * node + 1] + self.tree[2 * node + 2]

    def query(self, l, r):                       # sum of nums[l..r], inclusive
        return self._query(0, 0, self.n - 1, l, r)

    def _query(self, node, lo, hi, l, r):
        if r < lo or hi < l:                     # no overlap
            return 0
        if l <= lo and hi <= r:                  # [lo..hi] lies inside [l..r]
            return self.tree[node]
        mid = (lo + hi) // 2                     # partial overlap: ask both halves
        return (self._query(2 * node + 1, lo, mid, l, r) +
                self._query(2 * node + 2, mid + 1, hi, l, r))

Node i's children live at 2i + 1 and 2i + 2, the array layout from the start of the lesson. For many n the tree has gaps in this numbering. For our six values, indices 9 and 10 are never used and the largest index used is 12, so a list of 2n = 12 slots would already overflow. 4n slots are always enough (the worst ratio we measured, for every n up to 3,000, was 3.88n).

Here is query(1, 4) on the example, in the order the calls happen:

node range against [1..4] contributes
0 [0..5] partial asks both halves
1 [0..2] partial asks both halves
3 [0..1] partial asks both halves
7 [0..0] no overlap 0
8 [1..1] inside 3
4 [2..2] inside 5
2 [3..5] partial asks both halves
5 [3..4] inside 16
6 [5..5] no overlap 0

The total is 24. Now update(2, 10): the walk goes [0..5] → [0..2] → [2], sets the leaf to 10, and on the way back recomputes [0..2] = 4 + 10 = 14 and the root = 14 + 27 = 41. query(1, 3) goes from 15 to 20.

Why it is correct. One sentence stays true: tree[node] is the sum of nums over the node's range. The build makes it true bottom-up. An update changes one leaf and recomputes exactly that leaf's ancestors, which are the only nodes whose ranges contain i. A query returns the sum of nums over the part of the node's range that lies in [l..r]: 0 when there is no such part, the stored sum when the whole range lies inside, and otherwise the sum of the two halves' answers.

Why it is fast. An update touches one node per level: O(log n). For a query, look at any one level. Only the boxes that contain position l or position r can overlap the query partially; boxes between them lie inside, and the rest lie outside. So at most two boxes per level ask their children, at most four boxes per level are visited, and a query is O(log n). For n = 100,000 the tree has 18 levels, so a query visits at most 1 + 2 + 4 × 16 = 67 nodes, and random queries do reach 67. The build computes each node once from its children: O(n) time. Space is the 4n slots, plus a recursion stack of only ⌈log₂ n⌉ + 1 frames, so recursion is safe here, unlike on a chain-shaped BST.

⚠️ The no-overlap return is the operation's identity, not "0". It must be a value that changes nothing when combined: 0 for sums, float('inf') for min, float('-inf') for max. A min-tree that returns 0 there silently beats every positive value.

⚠️ An empty array has no root range. The constructor skips the build, and the caller must not query.

Range updates. "Add 5 to every value in [l..r]" as point updates costs O(log n) per value, O(n log n) for a long range. Lazy propagation stores a pending "add 5" on the O(log n) boxes that tile the range and pushes it down only when a later operation needs to look inside. That makes a range update O(log n). Know that it exists, and write it when a problem mixes range updates with range queries. If all the updates come before the first read, the difference arrays of Prefix, Suffix & Range Updates are simpler.

Your turn: turn the class into a range-minimum tree. Which lines change? Then, starting from [5, 2, 8, 6], call update(1, 9). What do query(0, 3) and query(2, 3) return?

Check your answer

Four lines. The + becomes min(…, …) in _build, in _update and in _query, and the no-overlap return 0 becomes return float('inf'). After the update the values are [5, 9, 8, 6], so query(0, 3) is 5 and query(2, 3) is 6. Forget the identity and query(0, 3) still says 5, because the root lies inside the range and returns at once. But query(2, 3) asks the box [0..1], which lies outside, gets 0 back, and returns min(0, 6) = 0.

Structure 4: Fenwick trees

Smell: you only need sums (or another operation you can undo, such as XOR) with point updates, and you would like a twenty-line class instead of a forty-line one.

A Fenwick tree (also called a binary indexed tree) is a list tree[1..n] in which slot k stores the sum of a block of values that ends at the k-th value. The block's length is the lowest set bit of k, written k & -k. With nums = [5, 2, 9, -3, 4, 1, 8, 6]:

slot k binary k & -k slot k sums the values slot value
1 0001 1 1st 5
2 0010 2 1st–2nd 7
3 0011 1 3rd 9
4 0100 4 1st–4th 13
5 0101 1 5th 4
6 0110 2 5th–6th 5
7 0111 1 7th 8
8 1000 8 1st–8th 32

k & -k isolates the lowest set bit. Python integers behave like two's-complement numbers with unlimited sign bits, so the trick works exactly as it does in C or Java.

Predict first: which slots add up to the sum of the first 7 values, with no value counted twice?

Check your answer

Slots 7, 6 and 4: the 7th value, the 5th–6th, and the 1st–4th, so 8 + 5 + 13 = 26. Look at the indices in binary: 0111 → 0110 → 0100. Each step drops the lowest set bit. That is the whole prefix loop.

class Fenwick:
    """Point add and prefix sums. Callers use 0-based indices; slots are 1-based."""
    def __init__(self, n):
        self.n = n
        self.tree = [0] * (n + 1)       # slot 0 is never used

    def add(self, i, delta):            # nums[i] += delta
        k = i + 1                       # nums[i] is the (i+1)-th value
        while k <= self.n:
            self.tree[k] += delta
            k += k & -k                 # the next slot whose block covers it

    def prefix(self, k):                # sum of the first k values
        s = 0
        while k > 0:
            s += self.tree[k]
            k -= k & -k                 # jump to just before this block
        return s

    def range_sum(self, l, r):          # nums[l..r], inclusive
        return self.prefix(r + 1) - self.prefix(l)

prefix(k) means exactly what prefix[k] meant in the Foundation lesson, the sum of the first k values. So the range formula is the same one: prefix(r + 1) − prefix(l).

Traced on the example:

call slots visited binary result
prefix(7) 7, 6, 4 0111 → 0110 → 0100 8 + 5 + 13 = 26
add(2, 5) (the 3rd value) 3, 4, 8 0011 → 0100 → 1000 each slot gains 5

Why it works. prefix(k) removes the lowest set bit of k at each step: 7 → 6 → 4 → 0. The blocks it reads (the 7th value; the 5th–6th; the 1st–4th) tile the first seven values exactly, with no overlap. Each step clears a 1-bit, so there are at most log₂ n + 1 steps. add climbs the other way. Adding the lowest set bit moves to the next slot whose block contains the same position (3 → 4 → 8), and those are exactly the slots that include the changed value. Both run in O(log n) time, and the whole structure takes n + 1 slots.

Building with n calls to add costs O(n log n). A linear build pushes each slot's total into the next slot that contains it, once:

def fenwick_from_list(nums):
    fw = Fenwick(len(nums))
    for k, x in enumerate(nums, 1):     # k = 1..n
        fw.tree[k] += x
        parent = k + (k & -k)           # the next slot whose block covers slot k's block
        if parent <= fw.n:
            fw.tree[parent] += fw.tree[k]
    return fw

By the time the loop reaches slot k, every smaller slot has already pushed into it, so tree[k] is complete before it is pushed on. O(n) time.

⚠️ add is not set. add(i, v) adds v. To set nums[i] = v, which is what LeetCode 307: Range Sum Query - Mutable asks for, keep a copy of the values, add v − current[i], and update the copy.

⚠️ Slot 0 is a trap. Forget the k = i + 1 shift and add(0, …) never ends: 0 & -0 is 0, so k never moves.

⚠️ Only undoable operations give range answers. A range answer is prefix(r + 1) "minus" prefix(l), and max has no minus. A Fenwick tree can keep prefix maxima if values only ever grow, but for range min or max under arbitrary updates, use a segment tree.

Segment tree Fenwick tree
Range query any operation where grouping doesn't matter: sum, min, max, gcd undoable operations: sum, XOR
Point update O(log n) O(log n)
Build O(n) O(n) with the linear build
Space up to 4n slots n + 1 slots
Range updates lazy propagation possible, with extra tricks
Code about 40 lines about 20 lines

Your turn: with n = 16, which slots does add touch for the 5th value (i = 4), and which slots does prefix(11) read?

Check your answer

add touches 5, 6, 8, 16 (binary 0101 → 0110 → 1000 → 10000). prefix(11) reads 11, 10, 8 (1011 → 1010 → 1000): the 11th value, the 9th–10th, and the 1st–8th, which together tile the first 11 values.

Final round: no label on the problem

Real problems don't say which structure they want. Find the question each operation asks, then pick the structure that answers it by visiting only a few nodes per level of a short tree.

Challenge 1: My Calendar I

Implement book(start, end) for half-open intervals [start, end). If the new booking overlaps no existing booking, store it and return True; otherwise return False. book(10, 20) → True, book(15, 25) → False, book(20, 30) → True. There are at most 1,000 calls. This is LeetCode 729: My Calendar I.

  1. What question does each call ask about the existing bookings?
  2. Which existing bookings can possibly overlap a new one?
  3. What does each call cost?
Hint

Keep the bookings sorted by start. Since they never overlap each other, they are then sorted by end as well. Only the two neighbors of the new start can collide with it.

Check your answer

It is a question about neighbors in sorted order: the last booking that starts at or before start (a floor), and the first one that starts after it.

import bisect

class MyCalendar:
    def __init__(self):
        self.starts = []        # sorted; ends[j] belongs to starts[j]
        self.ends = []

    def book(self, start, end):
        i = bisect.bisect_right(self.starts, start)
        if i > 0 and self.ends[i - 1] > start:              # the previous one runs into us
            return False
        if i < len(self.starts) and self.starts[i] < end:   # the next one starts too soon
            return False
        self.starts.insert(i, start)
        self.ends.insert(i, end)
        return True

Bookings before i − 1 end no later than booking i − 1 starts, so they end by start at the latest. Bookings after i start no earlier than booking i, so if booking i is clear of end, so are they. That leaves two checks. bisect costs O(log n) and list.insert costs O(n), so 1,000 calls cost at most about a million element moves, which is fine. With SortedList or a balanced BST each call is O(log n), which would matter if the limit were 10⁵.

Challenge 2: count the smaller numbers to the right

For each nums[i], count how many later elements are smaller. [5, 2, 6, 1] → [2, 1, 1, 0]. Constraints: n ≤ 10⁵, values between −10⁴ and 10⁴. This is LeetCode 315: Count of Smaller Numbers After Self.

Before peeking: why is the double loop too slow? Scanning from the right, what question do you ask about the values already seen, and which structure answers it while values keep arriving?

Check your answer

The double loop is about n²/2 = 5 × 10⁹ comparisons. Scan from the right instead. When you reach x, the values already seen are exactly the elements after it, and the question is "how many seen values are smaller than x?". That is a prefix count over values, with a +1 update after each step: a Fenwick tree indexed by the value's rank.

def count_smaller(nums):
    rank = {v: r for r, v in enumerate(sorted(set(nums)))}   # value -> 0..m-1
    fw = Fenwick(len(rank))
    out = []
    for x in reversed(nums):
        r = rank[x]
        out.append(fw.prefix(r))    # seen values with rank < r
        fw.add(r, 1)
    return out[::-1]

prefix(r) sums ranks 0 to r − 1, the values strictly smaller than x, so equal values are not counted. Time O(n log n) (the sort plus n Fenwick operations), extra space O(n). Shifting each value by 10⁴ instead of ranking also works here, because the values fit in 20,001 slots. A plain BST that stores subtree sizes also works on random input, but a sorted input makes it a chain and the whole thing O(n²).

Challenge 3: a trap

An array never changes, and 10⁴ queries ask for the sum of nums[l..r]. Segment tree or Fenwick tree?

Check your answer

Neither. With no updates, the Foundation lesson's prefix sums answer each query in O(1) after an O(n) build. This is LeetCode 303: Range Sum Query - Immutable. A tree would work, but it adds a log factor and twenty to forty lines to solve a problem that no longer exists. Trees earn their keep only when the data changes between questions.

Cheat sheet

When the problem… Reach for Cost per operation Key invariant
needs sorted order under inserts and deletes, or "what is next to x?" a BST; in Python, bisect or SortedList O(h); O(log n) when balanced "left subtree < node < right subtree"
feeds sorted or hostile input to a tree a balanced tree, or build from sorted data height O(log n) "a rotation keeps the in-order sequence"
asks about prefixes of many strings a trie O(L) per string "root-to-node path = prefix; the flag marks a word"
mixes range sum, min or max queries with updates a segment tree O(log n); O(n) build "tree[node] = answer for its range"
mixes prefix or range sums with updates a Fenwick tree O(log n) "slot k sums the k & -k values ending at the k-th"
only asks for the current minimum or maximum a heap O(1) peek, O(log n) push and pop "parent ≤ children"
never changes between queries prefix sums O(1) "prefix[k] = sum of the first k"

Before moving on, pick one structure and rebuild it from a blank editor. Say what each node remembers, which sentence every operation must keep true, which nodes an operation visits on each level, and what input makes the tree tall. Then name the final-round problem that the wrong structure would fail.

Next: Binary Tree Patterns for traversal orders and the recursive patterns built on them, Heaps & Priority Queues for the heap in depth, and Prefix, Suffix & Range Updates for range updates on arrays that stay put.