Part VI · More Patterns Pattern 20 4 problems

K-way Merge and Two Heaps

A heap answers one question fast: what is the smallest thing right now? K-way merge asks it of K sorted lists at once. Two heaps ask it from both sides of the middle, which gives you a running median.

Pattern 9 used a heap to keep the K best items. It named two other heap moves and left them for later. This page covers both. They share one tool, heapq, and one habit: push tuples with a safe tie-break so Python never has to compare two objects it cannot order.

Contents

  1. When to use
  2. Core idea
  3. The templates
  4. Common mistakes
  5. Merge k Sorted Lists
  6. Kth Smallest Element in a Sorted Matrix
  7. Find Median from Data Stream
  8. Smallest Range Covering Elements from K Lists
  9. Recap

When to use

The trigger. You have several sorted sources and need their items in one global order, or you need the middle of data that keeps arriving. “Merge K sorted”, “k-th smallest across sorted rows”, “median of a stream” and “smallest range that hits every list” are the classic phrasings.

Core idea

K-way merge: a heap of heads

The smallest item overall must be the head of one of the K lists, because each list is sorted. So keep a min-heap with exactly one entry per list: its current head. Pop the smallest, emit it, and push the next item from the same list. The heap never holds more than K entries, so each step costs O(log K). With N items in total, the merge costs O(N log K).
list 0 1 4 7 list 1 2 5 list 2 3 6 9 dashed: already emitted amber: each list's current head (3,2) (4,0) (5,1) min-heap of heads, one entry per list merged so far 1 2 3 pop 3 from list 2, then push 6, the next item of list 2
Figure 20.1 — The heap only ever holds one head per list, and its top is the smallest unused item.

Reading the figure. Each row is one sorted list. Dashed cells are already in the output. The amber cells are the current heads, and the heap holds exactly those three as (value, list). The heap top is the next item to emit. After it leaves, only list 2 moves forward, so only one push follows.

Tie-break with an index. Push (value, list_index, payload), never (value, payload). If two values tie, Python compares the next field. A list index is an int, so the tie ends there. A linked-list node has no <, so comparing two nodes raises TypeError. Only one entry per list sits in the heap at a time, so the index alone makes every entry unique.

Two heaps: one for each half

Split the numbers into a lower half and an upper half. Keep the lower half in a max-heap, so its biggest value is on top. Keep the upper half in a min-heap, so its smallest value is on top. Those two tops sit on either side of the median. Keep the sizes equal, or the lower half one bigger. Then the median is the lower top, or the average of both tops. heapq is min-only, so the max-heap stores negated values.
lower half (max-heap) upper half (min-heap) 3 1 5 15 top = max top = min heapq stores -3, -1 plain min-heap median = (3 + 5) / 2 = 4.0
Figure 20.2 — The two heap tops sit either side of the median, so reading it costs O(1).

Reading the figure. The left tree is the lower half, with its largest value on top. The right tree is the upper half, with its smallest value on top. The amber tops face each other across the dashed median line. Every value on the left is at most every value on the right. So the median comes from the tops alone.

The templates

Template A — merge K sorted lists with a heap of heads
import heapq


def merge_sorted_arrays(arrays: list[list[int]]) -> list[int]:
    """All values from K sorted lists, in one sorted list.

    Example:
        >>> merge_sorted_arrays([[1, 4, 7], [2, 5], [3, 6, 9]])
        [1, 2, 3, 4, 5, 6, 7, 9]
    """
    # One entry per non-empty list: (head value, list index, position).
    # The list index is the tie-break. Position 0 is each list's head.
    heap = [(arr[0], index, 0) for index, arr in enumerate(arrays) if arr]
    heapq.heapify(heap)           # O(K) to build, cheaper than K pushes

    merged: list[int] = []

    # Invariant: heap holds the next unused item of every list not yet empty,
    # so heap[0] is the smallest unused item overall.
    while heap:
        value, index, pos = heapq.heappop(heap)
        merged.append(value)

        # pos + 1 is the next item in the same list. Push it if it exists.
        if pos + 1 < len(arrays[index]):
            heapq.heappush(heap, (arrays[index][pos + 1], index, pos + 1))

    return merged

In production code, list(heapq.merge(*arrays)) does the same thing lazily. Write the loop in the interview, then name heapq.merge.

Template B — two heaps around the middle
def running_medians(stream: list[int]) -> list[float]:
    """The median after each new number arrives.

    Example:
        >>> running_medians([5, 15, 1, 3])
        [5.0, 10.0, 5.0, 4.0]
    """
    lower: list[int] = []         # max-heap of the small half, stored negated
    upper: list[int] = []         # min-heap of the large half
    medians: list[float] = []

    # Invariant at the top of each pass: every value in lower is at most every
    # value in upper, and len(lower) is len(upper) or len(upper) + 1.
    for value in stream:
        # Step 1: push into lower (negate it, max-heap). Then move lower's
        # largest across, so upper gets the right value and order holds.
        heapq.heappush(lower, -value)
        heapq.heappush(upper, -heapq.heappop(lower))   # - undoes the negation

        # Step 2: rebalance. Upper may now be one too big. Move its smallest
        # back, negated, so lower is never smaller than upper.
        if len(upper) > len(lower):
            heapq.heappush(lower, -heapq.heappop(upper))

        # Step 3: read the tops. lower[0] is the negated max of the low half.
        if len(lower) > len(upper):
            medians.append(float(-lower[0]))           # odd count: lower top
        else:
            # Even count: average the two tops. / 2 is true division.
            medians.append((-lower[0] + upper[0]) / 2)

    return medians
before adding 1 lower upper 5 15 median 10.0 1. push 1 into lower lower upper 1 5 15 sizes 2 and 1 1. move lower's max up lower upper 1 5 15 sizes 1 and 2 2. move 5 back down lower upper 1 5 15 median 5.0
Figure 20.3 — Push into lower, move its max up, then rebalance. Order between the halves can never break.

Reading the figure. Each column is one moment while 1 is added. Blue cells are the lower half and violet cells the upper half. The amber cell is the value that just moved. The move up always sends the largest value of the lower half, so the order between halves holds. The move back only fixes the sizes.

The push-then-move dance in step 1 is the trick. It avoids comparing the new value with either top by hand, and it can never break the order between halves.

Common mistakes

The problems

1. Merge k Sorted Lists Hard

Problem

You are given k sorted linked lists. Merge them into one sorted linked list and return its head.

The idea

Seed a min-heap with the head node of every non-empty list. Pop the smallest node, link it onto the result, and push that node’s next. A dummy head node removes the “is the result empty yet” special case. The list index is the tie-break, so two nodes are never compared.

Solution

from __future__ import annotations

import heapq


class MergeNode:
    """A singly linked list node for this page's merge problems."""

    def __init__(self, val: int = 0, next: MergeNode | None = None) -> None:
        self.val = val            # default 0: only the dummy head uses it
        self.next = next


def build_merge_list(values: list[int]) -> MergeNode | None:
    """Linked list holding values, in order. Empty input gives None."""
    head = None
    # Build from the back, so each new node points at the one built before.
    for value in reversed(values):
        head = MergeNode(value, head)
    return head


def merge_list_values(head: MergeNode | None) -> list[int]:
    """The values of a linked list, as a Python list."""
    out: list[int] = []
    while head is not None:
        out.append(head.val)
        head = head.next
    return out


def merge_k_lists(lists: list[MergeNode | None]) -> MergeNode | None:
    """Merge k sorted linked lists into one sorted linked list.

    Args:
        lists: Heads of sorted linked lists. Any of them may be None.

    Returns:
        The head of the merged list, or None if every list is empty.

    Example:
        >>> heads = [build_merge_list([1, 4, 5]), build_merge_list([1, 3, 4]),
        ...          build_merge_list([2, 6])]
        >>> merge_list_values(merge_k_lists(heads))
        [1, 1, 2, 3, 4, 4, 5, 6]
        >>> merge_k_lists([None, None]) is None
        True
    """
    # (value, list index, node). The index breaks ties, so heapq never
    # has to compare two MergeNode objects, which would raise TypeError.
    heap: list[tuple[int, int, MergeNode]] = []
    for index, head in enumerate(lists):
        if head is not None:      # skip empty lists: they have no head
            heap.append((head.val, index, head))
    heapq.heapify(heap)           # O(k) build

    dummy = MergeNode()           # placeholder before the real head
    tail = dummy                  # last node of the merged list so far

    # Invariant: dummy.next .. tail is sorted and holds every popped node.
    # heap holds the current head of every list not yet used up.
    while heap:
        _, index, node = heapq.heappop(heap)   # _ : value, already in node
        tail.next = node
        tail = node

        # Same index: the next node comes from the same list, and that list
        # has no other entry in the heap, so the tie-break stays unique.
        if node.next is not None:
            heapq.heappush(heap, (node.next.val, index, node.next))

    return dummy.next             # skip the placeholder

Walkthrough

Lists [1, 4, 5], [1, 3, 4], [2, 6]. Heap entries shown as (value, index):

step heap, smallest first output seed (1,0) (1,1) (2,2) pop (1,0) (1,1) (2,2) (4,0) 1 pop (1,1) (2,2) (3,1) (4,0) 1 1 pop (2,2) (3,1) (4,0) (6,2) 1 1 2 pop (3,1) (4,0) (4,1) (6,2) 1 1 2 3 4 more pops empty 1 1 2 3 4 4 5 6
Figure 20.4 — Each pop emits the smallest head and pushes the next node from the same list.

Reading the figure. Each row is the state after one step. Heap entries are (value, list index), smallest first. The amber entry is the next to pop. The violet entry was just pushed from the list that lost its head. Notice row 2: two entries tie on value 1, and the list index settles the order. Green cells are the merged output.

TimeO(N log k)SpaceO(k) for the heap

Edge cases to raise

Say this out loud: “The next node overall is the head of one of the k lists, so I keep a heap of just the heads. Each pop costs log k. I store the list index as a tie-break, so Python never tries to compare two nodes.”

2. Kth Smallest Element in a Sorted Matrix Medium

Problem

An n × n matrix has every row and every column sorted in ascending order. Return the k-th smallest element, counting duplicates.

The idea

Each row is a sorted list, so this is a K-way merge over n rows that stops early. Seed the heap with the first cell of each row. Pop k - 1 times, pushing the cell to the right each time. The heap top is then the k-th smallest. Only the first min(n, k) rows can matter, because the k-th smallest cannot sit below row k.

Solution

def kth_smallest_in_matrix(matrix: list[list[int]], k: int) -> int:
    """The k-th smallest value in a matrix with sorted rows and columns.

    Args:
        matrix: A non-empty n x n matrix. Rows and columns ascend.
        k: Rank to find, 1 <= k <= n * n.

    Returns:
        The k-th smallest value, duplicates counted.

    Example:
        >>> kth_smallest_in_matrix([[1, 5, 9], [10, 11, 13], [12, 13, 15]], 8)
        13
        >>> kth_smallest_in_matrix([[-5]], 1)
        -5
    """
    n = len(matrix)

    # Seed with column 0 of each row: (value, row, col). Rows past k cannot
    # hold the answer, so min(n, k) rows is enough. row is the tie-break.
    heap = [(matrix[row][0], row, 0) for row in range(min(n, k))]
    heapq.heapify(heap)

    # Pop k - 1 times: after that the smallest k - 1 are gone and the top
    # of the heap is the k-th smallest.
    # Invariant: heap holds the next unused cell of every seeded row.
    for _ in range(k - 1):
        _, row, col = heapq.heappop(heap)
        # col + 1 is the next cell to the right. Push it if the row has one.
        if col + 1 < n:
            heapq.heappush(heap, (matrix[row][col + 1], row, col + 1))

    return heap[0][0]             # heap[0]: smallest entry. [0]: its value

Walkthrough

Matrix rows [1, 5, 9], [10, 11, 13], [12, 13, 15], k = 8:

1 pop 1 5 pop 2 9 pop 3 r0 10 pop 4 11 pop 5 13 pop 7 r1 12 pop 6 13 answer 15 r2 Seed: column 0, so 1, 10, 12. Each pop pushes the cell to its right. Pops 1 to 7 take 1, 5, 9, 10, 11, 12, 13. The heap top is now row 2’s 13. That is the 8th smallest.
Figure 20.5 — Seven pops remove the seven smallest cells, and the heap top is the 8th.

Reading the figure. Each cell shows its value and the pop that removed it. Blue cells were popped. The order snakes across rows, because the heap always holds the next cell of each row. The green cell is on top of the heap after k - 1 = 7 pops. The grey 15 is never touched.

TimeO(min(n, k) + k log min(n, k))SpaceO(min(n, k))

Edge cases to raise

Say this out loud: “Each row is a sorted list, so this is a K-way merge that stops after k pops. I only seed the first min(n, k) rows, because the answer cannot be deeper than row k.”

3. Find Median from Data Stream Hard

Problem

Design a class with add_num(num), which adds a number from a stream, and find_median(), which returns the median of every number so far. Both should be fast for millions of calls.

The idea

Sorting on every call is O(n log n). Inserting into a sorted list is O(n). Two heaps give O(log n) to add and O(1) to read. The lower half lives in a max-heap of negated values. The upper half lives in a min-heap. Lower is allowed one extra item, so with an odd count the median is the lower top.

Solution

class MedianFinder:
    """Running median with two heaps.

    lower is a max-heap (values negated) holding the smaller half.
    upper is a min-heap holding the larger half. len(lower) is always
    len(upper) or len(upper) + 1, and every lower value <= every upper value.

    Example:
        >>> finder = MedianFinder()
        >>> finder.add_num(1)
        >>> finder.add_num(2)
        >>> finder.find_median()
        1.5
        >>> finder.add_num(3)
        >>> finder.find_median()
        2.0
    """

    def __init__(self) -> None:
        self.lower: list[int] = []    # negated values: lower[0] is -(max of low half)
        self.upper: list[int] = []    # plain values: upper[0] is min of high half

    def add_num(self, num: int) -> None:
        """Add num to the stream in O(log n)."""
        # Push into lower, negated so heapq's min acts as a max.
        heapq.heappush(self.lower, -num)
        # Move lower's largest to upper. - undoes the negation. This keeps
        # every lower value <= every upper value, whatever num was.
        heapq.heappush(self.upper, -heapq.heappop(self.lower))

        # Upper may now hold one more than lower. Move its smallest back,
        # negated, so lower is equal or exactly 1 bigger.
        if len(self.upper) > len(self.lower):
            heapq.heappush(self.lower, -heapq.heappop(self.upper))

    def find_median(self) -> float:
        """Median of every number added so far, in O(1).

        Raises:
            ValueError: If no number has been added.
        """
        if not self.lower:            # lower is empty only when both are
            raise ValueError("no numbers added yet")

        # Odd count: lower holds the extra one, and its top is the median.
        # [0] is the heap top, and - turns the stored negative back.
        if len(self.lower) > len(self.upper):
            return float(-self.lower[0])

        # Even count: average the two middle values. / 2 is true division,
        # so 1 and 2 give 1.5, not 1.
        return (-self.lower[0] + self.upper[0]) / 2

Walkthrough

Add 1, 2, 3. Heaps shown as plain values:

add 1 lower upper 1 empty median 1.0 1 went up, then came back add 2 lower upper 1 2 median (1 + 2) / 2 = 1.5 2 moved up to upper add 3 lower upper 1 2 3 median 2.0 3 went up, 2 came back
Figure 20.6 — Lower keeps the extra item on odd counts, so its top is the median.

Reading the figure. Each column is the state after one add_num. Blue cells are the lower half and violet the upper half, shown sorted. The amber cells are the two heap tops, the only values find_median reads. With an odd count, lower is one bigger and its top is the answer.

add_numO(log n)find_medianO(1)SpaceO(n)

Edge cases to raise

Say this out loud: “I keep the smaller half in a max-heap and the larger half in a min-heap, with the lower one allowed one extra. Every new number goes through lower to upper, then I rebalance. The median is one top or the average of both.”

4. Smallest Range Covering Elements from K Lists Hard

Problem

You have k sorted lists of integers. Find the smallest range [a, b] that includes at least one number from every list. A range is smaller if b - a is smaller, or if the widths tie and a is smaller.

The idea

Pick one number from each list. The range they cover runs from their min to their max. To shrink it, the only useful move is to raise the min, because raising anything else can only widen it. So keep a heap of the current pick from each list, which gives the min in O(1), and track the max by hand. Pop the min, record the range if it is the best yet, and replace it with the next number from the same list. Stop when any list runs out, because then no pick covers every list.

Solution

def smallest_range(nums: list[list[int]]) -> list[int]:
    """Smallest [lo, hi] holding at least one value from every list.

    Args:
        nums: k non-empty lists, each sorted ascending.

    Returns:
        [lo, hi] for the smallest such range. Ties go to the smaller lo.

    Raises:
        ValueError: If nums is empty or any list is empty.

    Example:
        >>> smallest_range([[4, 10, 15, 24, 26], [0, 9, 12, 20], [5, 18, 22, 30]])
        [20, 24]
        >>> smallest_range([[1, 2, 3], [1, 2, 3]])
        [1, 1]
    """
    if not nums or any(not row for row in nums):
        raise ValueError("need at least one list, and no empty lists")

    # One pick per list: (value, list index, position). Position 0 is the
    # list's smallest value. The list index is the tie-break.
    heap = [(row[0], index, 0) for index, row in enumerate(nums)]
    heapq.heapify(heap)
    # The heap gives the min. The max has to be tracked by hand.
    current_max = max(row[0] for row in nums)       # [0]: each list's head

    # Start with the infinite range, so the first real range always wins.
    best_lo, best_hi = float("-inf"), float("inf")

    # Invariant: heap holds exactly one value from every list, and
    # current_max is the largest of them. So [heap[0][0], current_max]
    # covers every list.
    while True:
        low, index, pos = heapq.heappop(heap)

        # Strict <: on a tie keep the earlier range, whose lo is smaller,
        # because lo only grows as the loop runs.
        if current_max - low < best_hi - best_lo:
            best_lo, best_hi = low, current_max

        # pos + 1 == len: this list has no next value, so no later pick can
        # cover it. Every remaining range has been tried.
        if pos + 1 == len(nums[index]):
            return [best_lo, best_hi]

        nxt = nums[index][pos + 1]                  # pos + 1: next in this list
        heapq.heappush(heap, (nxt, index, pos + 1))
        current_max = max(current_max, nxt)         # the new pick may be the max

Walkthrough

Lists [4, 10, 15, 24, 26], [0, 9, 12, 20], [5, 18, 22, 30]. Picks shown as a set:

[0, 5] w 5 best, width 5 [4, 9] w 5 tie, keep [0, 5] [5, 10] w 5 [9, 18] w 9 [10, 18] w 8 [12, 18] w 6 [15, 20] w 5 [18, 24] w 6 [20, 24] w 4 new best, width 4 0 5 10 15 20 25 30
Figure 20.7 — Raising the minimum each step slides the window right, and [20, 24] is the narrowest it ever gets.

Reading the figure. Each bar is the range of the current picks, one row per step, top to bottom. The white dots are the three picks, one from each list. The left dot is always the heap’s min, and it is the one replaced next. Green bars are a new best. The amber bar ties on width and loses on the left end. The search stops after the last row, because list 1 has run out.

TimeO(N log k)SpaceO(k)

Edge cases to raise

Say this out loud: “I hold one number from each list. The only move that can shrink the range is raising the min, so I pop it from a heap and replace it with the next value from its list. I track the max by hand, and I stop when any list runs out.”

Recap

The six things to carry forward

Where this goes next

Pattern 21, Weighted Shortest Paths, puts the same heap of (cost, tie-break, node) tuples at the centre of Dijkstra’s algorithm. There the heap holds the frontier of a graph instead of the heads of K lists.


← 19 — Bit Manipulation 21 — Weighted Shortest Paths →