MediumBinary Search

Kth Smallest Element in a Sorted Matrix โ€” Solution

Problem

Given an nร—n matrix where every row and every column is sorted in ascending order, find the kth smallest value in the matrix. Each row is sorted left to right, and each column is sorted top to bottom โ€” but the rows themselves are not sorted relative to each other.

For example:

matrix = [ [ 1, 5, 9], [10, 11, 13], [12, 13, 15] ] k = 8
  • Input: the matrix above, k = 8
  • Output: 13
  • Explanation: the sorted order is 1, 5, 9, 10, 11, 12, 13, 13, 15 โ€” the 8th element is 13.

Intuition

The matrix's row-sorted structure lets us treat each row as its own sorted stream, which a min-heap can merge efficiently. But the column-sorted structure adds something stronger: if we pick any value mid, we can count exactly how many matrix elements are โ‰ค mid in O(n) time using a bottom-left traversal. That count function turns the problem into a binary search over the value range โ€” find the smallest matrix value where count โ‰ฅ k.

Approach 1 โ€” Min-Heap

Treat each row as a sorted stream. Seed a min-heap with the first element of every row, then pop k times, refilling from the next column each time.

  1. Push (matrix[row][0], row, 0) for every row into a min-heap.
  2. Pop from the heap k times.
  3. After each pop (value, row, col), if col + 1 < n, push (matrix[row][col + 1], row, col + 1).
  4. The value from the kth pop is the answer.
1import heapq
2from typing import List
3
4class Solution:
5    def kthSmallest(self, matrix: List[List[int]], k: int) -> int:
6        n = len(matrix)
7        min_heap = [(matrix[row][0], row, 0) for row in range(n)]
8        heapq.heapify(min_heap)
9
10        value = 0
11        for _ in range(k):
12            value, row, col = heapq.heappop(min_heap)
13            if col + 1 < n:
14                heapq.heappush(min_heap, (matrix[row][col + 1], row, col + 1))
15
16        return value

Time: O(k log n) โ€” k heap pops, each costing O(log n) with at most n elements in the heap.
Space: O(n) โ€” the heap holds exactly one element per row at any time.

Approach 2 โ€” Binary Search on Value Range

Binary search over the value range [matrix[0][0], matrix[n-1][n-1]]. For a candidate mid, count how many elements are โ‰ค mid by traversing from the bottom-left corner. Find the smallest matrix value where that count reaches k.

  1. Set lo = matrix[0][0], hi = matrix[n-1][n-1].
  2. While lo < hi, compute mid = lo + (hi - lo) / 2.
  3. Count elements โ‰ค mid using the bottom-left traversal:
    • Start at (row = n-1, col = 0).
    • If matrix[row][col] <= mid: every element in this column from row 0 to row is also โ‰ค mid, so add row + 1 to count and move right.
    • Otherwise move up.
  4. If count < k: lo = mid + 1. Otherwise: hi = mid.
  5. Return lo โ€” it converges to the smallest real matrix element where count โ‰ฅ k.
1from typing import List
2
3class Solution:
4    def kthSmallest(self, matrix: List[List[int]], k: int) -> int:
5        n = len(matrix)
6
7        def count_less_equal(target: int) -> int:
8            count = 0
9            row, col = n - 1, 0  # start bottom-left to exploit both sort orders
10            while row >= 0 and col < n:
11                if matrix[row][col] <= target:
12                    count += row + 1  # entire column prefix row 0..row is also <= target
13                    col += 1
14                else:
15                    row -= 1
16            return count
17
18        lo, hi = matrix[0][0], matrix[n - 1][n - 1]
19        while lo < hi:
20            mid = lo + (hi - lo) // 2
21            if count_less_equal(mid) < k:
22                lo = mid + 1
23            else:
24                hi = mid
25
26        return lo

Time: O(n log(max โˆ’ min)) โ€” log(max โˆ’ min) binary search steps, each with an O(n) count traversal.
Space: O(1) โ€” no auxiliary data structures beyond a few variables.

Complexity Summary

ApproachTimeSpaceWhen to use
Min-HeapO(k log n)O(n)When k is small relative to nยฒ โ€” fewer pops means less work
Binary SearchO(n log(max โˆ’ min))O(1)When k is large or memory is constrained; preferred for large matrices

Common Mistakes

  • Counting row instead of row + 1 in the binary search count step: the column prefix from index 0 to row contains row + 1 elements, not row.
  • Skipping the >= left boundary check on binary search direction: when count_less_equal(mid) == k, setting lo = mid + 1 instead of hi = mid misses that mid itself might be the answer.
  • Not guarding col + 1 < n before pushing to the heap: pushing an out-of-bounds column causes an index error or incorrect results on the last column.
  • Trusting that lo converges to a real matrix value without understanding why: mid in the binary search may not exist in the matrix, but the invariant guarantees lo lands on a real element โ€” confusing this leads to returning a value that isn't in the matrix.
  • Using (lo + hi) / 2 instead of lo + (hi - lo) / 2 when matrix values are large integers โ€” risks overflow in Java and C++.

Related Problems

Ready to practice? Try it on SkillFlow

Adaptive problems, AI follow-up interviews, and a skill score that shows exactly where you need to improve.

Practice This Problem โ†’