DSA sheet · Binary Search · 2D matrix + binary search on answer

Kth Smallest in Sorted Matrix

The third 2D-matrix question, and it joins two earlier ideas. From Problem 13 we take the staircase walk, which here counts how many numbers are ≤ some value instead of searching for one. From the binary-search-on-answer problems (14–21) we take the habit of binary searching over possible answers, not over positions. The teacher goes step by step: brute force (flatten and sort) → try every value from smallest to biggest with a count function → replace that loop with binary search. She says plainly that if binary search on answer is new to you, watch those videos first. Part 0 below recaps it so this page stands alone.

Every part below follows the same order:
① the question in simple words → ② what the constraints tell us → ③ intuition → ④ building the conditions from examples → ⑤ approach steps → ⑥ code → ⑦ code line by line → ⑧ dry run → ⑨ complexity & remember

Part 0 · Before starting

Why binary search works

Binary search needs a sorted order, or more generally a yes/no question whose answers go no, no, …, no, yes, yes, …, yes as you move right. Then one check at the middle tells you which half to throw away. low and high mark what's still possible, mid is the middle, and we stop when low > high.

mid = low + (high - low) // 2

In Java/C++, low + high can be bigger than the largest int and overflow (wrap to a negative number). Writing low + (high - low) // 2 avoids building that big sum. The teacher uses this safe form here. Python ints never overflow, so both forms work in Python. Also, // rounds down (towards −∞), which keeps mid between low and high even when the values are negative.

The "binary search on answer" pattern

  1. Answer range. Don't search positions. Ask: "what is the smallest the answer could be, and the biggest?" Here the answer is a value in the matrix, so it lies between the smallest value matrix[0][0] and the biggest matrix[n-1][n-1].
  2. A yes/no check for a guess. For a guess x, ask: "are at least k numbers in the matrix ≤ x?" Call the count count(x).
  3. Monotonic (only goes one way): if x makes count(x) ≥ k true, then any bigger x also does, because a bigger x can only include more numbers. So over the range, the check looks like no no no yes yes yes.
  4. Which side to keep. We want the first yes: the smallest x with count(x) ≥ k. If mid says no (count < k) → the answer is bigger → low = mid + 1. If mid says yes → mid might be the answer, but maybe something smaller is too → high = mid - 1. When the loop ends, low is the first yes.
guess x16…22…262728…94
count(x)1…2…234…16
≥ k = 3?no…no…noyesyes…yesanswer = the first "yes" = 27

Matrix words

The example matrix used on this page

          c0   c1   c2   c3
   r0  [  16   28   60   64 ]
   r1  [  22   41   63   80 ]
   r2  [  27   50   66   82 ]
   r3  [  36   78   83   94 ]

   sorted: 16 22 27 28 36 41 50 60 63 64 66 78 80 82 83 94

This is the teacher's 4 × 4 example. Not every cell is readable in the video, so a few inner values are filled in here so that every fact she states holds: the corners 16, 36, 64, 94; column 0 is 16, 22, 27, 36; row 0 starts 16, 28, 60; the counts for 16, 23, 27 and 30 match hers, and the 10th smallest is 64.

Part A · Brute force: flatten and sort

LeetCode 378 · Kth Smallest Element in a Sorted Matrix

1The question in simple words

You get an n × n matrix where every row and every column is sorted in non-decreasing order, and a number k. Return the k-th smallest number of the whole matrix (counting repeats, so it's the k-th in sorted order, not the k-th distinct).

Example (LeetCode): [[1,5,9],[10,11,13],[12,13,15]], k = 8 → sorted is 1 5 9 10 11 12 13 13 15 → the 8th is 13. On our 4 × 4 matrix with k = 3 → 27.

2What the constraints tell us

3Intuition

Forget the matrix is sorted. Pour all numbers into one list, sort it, and pick position k − 1 (lists start at index 0).

4Building the logic: what the matrix does and doesn't tell us

The teacher first looks at what we know for sure from the sorting:

The simplest fix: flatten and sort.

5Approach steps

  1. Put every number in one list.
  2. Sort it.
  3. Return values[k - 1].

6Code (Python)

brute force: flatten + sort
class Solution:
    def kthSmallest(self, matrix, k):
        values = []
        for row in matrix:
            for v in row:
                values.append(v)
        values.sort()
        return values[k - 1]          # k is 1-based, lists are 0-based

7Code line by line

linewhat it means
values.append(v)Copy all n² numbers out. This is the extra memory.
values.sort()The expensive step: sorting N = n² numbers.
return values[k - 1]"1st smallest" is index 0, so the k-th is index k − 1.

8Dry run (k = 3)

index012345…15
sorted162227283641…94values[3 − 1] = values[2] = 27

9Complexity & remember

Slow, and it uses extra space. What if the interviewer says "no extra space, work in the matrix itself"? And since the matrix is already sorted, they expect us to use that with binary search.

RememberFlatten → sort → index k − 1. Correct, but O(n² log n) time and O(n²) space, and it ignores the sorting we were given.

Part B · Try every value, with a staircase count

1The question

Same question, but now: no extra list, and use the sorted rows and columns.

2What the constraints tell us

3Intuition: a "how many are ≤ x?" game

The teacher frames it as a game. I tell you a number x. You tell me how many numbers in the matrix are ≤ x.

Now try x = 16, 17, 18, … in order. The first x whose count reaches k is the k-th smallest. For k = 3: 16 to 21 give 1, 22 to 26 give 2, and 27 gives 3 → answer 27.

4Building the count function from examples

Where to start: the bottom-left cell (36)

This is the Problem 13 corner idea again. At 36, smaller numbers are up and bigger numbers are right, one direction each. So the code always knows where to go.

Example: count(16) → too big means go up

At 36: 36 > 16. Going right only gives bigger numbers, so go up (row − 1). 27 > 16 → up. 22 > 16 → up. 16 ≤ 16 ✓. Now we're in row 0. Since the column is sorted, everything above this cell (nothing here) is also ≤ 16. So this column adds row + 1 = 0 + 1 = 1 number.

Rule: value > xrow -= 1 (go up to find smaller numbers)

Example: count(27) → why we add row + 1

At 36 > 27 → up. At 27 ≤ 27 ✓, in row 2. The column is sorted top to bottom, so 27 and everything above it (22, 16) are ≤ 27. That's rows 0, 1, 2 → row + 1 = 3 numbers from this column in one go. No need to walk up and check them one by one.

Rule: value ≤ xcount += row + 1, then col += 1 (this column is done, move right)

Example: count(30) → why we can't stop after the first "≤"

At 36 > 30 → up. At 27 ≤ 30 → count += 3. If we returned now we'd say 3, but 28 (column 1) is also ≤ 30. Other columns can still have small numbers, so we must move right and keep counting:

          c0      c1      c2     c3
   r0  [  16    [28]5   [60]6    64 ]
   r1  [  22    [41]4    63      80 ]
   r2  [ [27]2  [50]3    66      82 ]
   r3  [ [36]1   78      83      94 ]

   [x]k = k-th cell visited for count(30)
steprowcolvaluedecisioncount
1303636 > 30 → up0
2202727 ≤ 30 → add 2 + 1, go right3
3215050 > 30 → up3
4114141 > 30 → up3
5012828 ≤ 30 → add 0 + 1, go right4
6026060 > 30 → up → row = −14
end−12row < 0 → stop4

Boundaries: row only goes down (it can reach −1) and col only goes up (it can reach n). So the loop runs while row >= 0 and col < n.

What the count function returns

The teacher has it return True when count ≥ k, not "== k". With repeated values, the count can jump past k in one step (for [[1,2],[2,3]] and k = 2: count(1) = 1, count(2) = 3). So we look for "at least k", never "exactly k". The first x where it becomes True is the answer.

Doubt 1: she says that starting from 64 would give you the "k-th largest", and uses a students example (with 10 students, the 3rd from the front is the "7th" from the back). Is that right?
→ Two small fixes. (1) Position from the back is n − k + 1, so the 3rd of 10 is the 8th from the back. In the matrix, the k-th smallest is the (n² − k + 1)-th largest. (2) You can also count "how many ≤ x" from the top-right: when the value is ≤ x, everything to its left in that row is ≤ x too, so add col + 1 and go down. Otherwise go left. It gives the same count. Both corners work, so pick one and keep it.
same count, started from the top-right corner (also correct)
def count_top_right(matrix, x):
    n = len(matrix)
    row, col, count = 0, n - 1, 0
    while row < n and col >= 0:
        if matrix[row][col] <= x:
            count += col + 1          # this cell and everything left of it
            row += 1
        else:
            col -= 1
    return count

5Approach steps

  1. low = matrix[0][0], high = matrix[n-1][n-1].
  2. For x = low, low + 1, …, high: if count(x) >= k, return x.
  3. count(x): start at bottom-left. Value > x → up. Else → add row + 1 and go right. Stop when outside.

6Code (Python)

try every value (correct, but slow when values are huge)
class Solution:
    def kthSmallest(self, matrix, k):
        n = len(matrix)
        low, high = matrix[0][0], matrix[n - 1][n - 1]
        for x in range(low, high + 1):
            if self.enough(matrix, x, k):
                return x                      # first x with at least k numbers <= x
        return high

    def enough(self, matrix, x, k):
        n = len(matrix)
        row, col = n - 1, 0                   # bottom-left corner
        count = 0
        while row >= 0 and col < n:
            if matrix[row][col] > x:
                row -= 1                      # too big: go up
            else:
                count += row + 1              # this cell and all above it are <= x
                col += 1                      # next column
        return count >= k

7Code line by line

linewhat it means
low, high = matrix[0][0], matrix[n - 1][n - 1]The answer range: smallest and biggest value.
for x in range(low, high + 1):Try every possible value, smallest first. high + 1 because range stops before its end.
if self.enough(matrix, x, k): return xThe first x that has k numbers ≤ it is the k-th smallest.
return highNever reached (x = high always has all n² ≥ k numbers), just a safe ending.
row, col = n - 1, 0Start at the bottom-left: up = smaller, right = bigger.
while row >= 0 and col < n:Stop when we walk off the top or off the right side.
if matrix[row][col] > x: row -= 1Too big → only up can give smaller numbers.
count += row + 1The cell is ≤ x, and the column is sorted, so rows 0 … row in this column are all ≤ x.
col += 1This column is fully counted. Don't return yet: other columns may still have numbers ≤ x.
return count >= k"At least k", so repeated values are handled.

8Dry run (k = 3)

xcount(x)≥ 3?
16 … 211no
22 … 262no
273yes → return 27

That's 12 calls of the count function before we hit 27. With a wide value range this loop would make a huge number of calls, and Part C cuts them down.

9Complexity & remember

One count costs at most 2n − 1 cells

Each step goes up one row or right one column. There are n rows and n columns, so at most about 2n steps, not n². The longest walk on our matrix is for x = 64 (k = 10):

          c0      c1      c2      c3
   r0  [  16      28      60     [64]7 ]
   r1  [  22      41     [63]5   [80]6 ]
   r2  [  27     [50]3   [66]4    82   ]
   r3  [ [36]1   [78]2    83      94   ]

   7 cells = 2n − 1 for n = 4,  count = 4 + 3 + 2 + 1 = 10
Remembercount(x) = staircase from the bottom-left: too big → up; else add row + 1 and go right. The answer is the first x with count(x) ≥ k.

Part C · Optimal: binary search on the answer

1The question

Same. Part B tried 16, 17, 18, … one by one. Can we skip most of them?

2What the constraints tell us

3Intuition

The values 16 … 94 are in order, and the answer to "count(x) ≥ k?" is no for a while and then yes forever after (Part 0). So we don't need to walk from the left. Jump to the middle value, count, and throw away half of the range.

4Building the conditions

count(mid) < kNot enough numbers ≤ mid → mid is too small, and so is everything below it → low = mid + 1.
count(mid) ≥ kmid is big enough. It might be the answer, but a smaller value might be too → remember it by moving high = mid - 1 and keep looking left.

(While speaking she says "greater than k → left" for both cases for a moment, but the code is clear: less than k → right, otherwise → left.)

Why return low and not high?

The teacher explains it through the moment the two pointers cross. Think about what each pointer means while the loop runs:

The loop ends when low = high + 1. Then high sits on the last "no" (a value where low used to be checked and failed), and low sits on the first "yes". The first yes is the answer, so return low. In her small picture: just before crossing, low was at a value that wasn't good enough, then low moved up past high. high ended on that "not good enough" spot, which can't be the answer.

Doubt 1: mid can be a number that isn't in the matrix (like 55). Can the answer end up being a number that isn't in the matrix?
→ No. count(x) only changes at values that are in the matrix. Between 27 and 28, for example, every x has the same count. So the first x where count becomes ≥ k must be exactly a matrix value. Values not in the matrix get thrown away along the way.
Doubt 2: "When to return low or high, when to write < or <=" confuses many people. How do I check?
→ The teacher's advice: dry run it yourself on paper with a small case. Try returning high and watch it fail. The rule above (low = first yes) is what that dry run shows you.

5Approach steps

  1. low = matrix[0][0], high = matrix[n-1][n-1].
  2. While low <= high: mid = low + (high - low) // 2, c = count(mid).
  3. c < k → low = mid + 1. Else → high = mid - 1.
  4. Return low.

6Code (Python)

Kth Smallest in a Sorted Matrix: binary search on answer
class Solution:
    def kthSmallest(self, matrix, k):
        n = len(matrix)
        low, high = matrix[0][0], matrix[n - 1][n - 1]
        while low <= high:
            mid = low + (high - low) // 2
            if self.count_le(matrix, mid) < k:
                low = mid + 1                 # not enough: answer is bigger
            else:
                high = mid - 1                # enough: try smaller
        return low                            # first value with count >= k

    def count_le(self, matrix, x):            # how many numbers are <= x
        n = len(matrix)
        row, col = n - 1, 0                   # bottom-left corner
        count = 0
        while row >= 0 and col < n:
            if matrix[row][col] > x:
                row -= 1
            else:
                count += row + 1
                col += 1
        return count

7Code line by line

linewhat it means
low, high = matrix[0][0], matrix[n - 1][n - 1]Same answer range as Part B.
while low <= high:The Part B for loop is replaced by a binary search loop.
mid = low + (high - low) // 2A guess in the middle of the range (the overflow-safe form).
if self.count_le(matrix, mid) < k:The count function now returns the number itself, and the comparison with k happens here.
low = mid + 1Fewer than k numbers ≤ mid → the k-th smallest is bigger than mid.
else: high = mid - 1At least k → mid is a "yes". Look for an even smaller yes.
return lowPointers have crossed; low is the first yes.
count_le(...)Exactly the staircase count from Part B.

8Dry run (k = 3, answer 27)

steplowhighmidcount(mid)decisionwhat we throw away
116945577 ≥ 3 → high = 5455 … 94
216543544 ≥ 3 → high = 3435 … 54
316342522 < 3 → low = 2616 … 25
426343044 ≥ 3 → high = 2930 … 34
526292733 ≥ 3 → high = 2627 … 29 (27 was a "yes"; low will land on it)
626262622 < 3 → low = 2726
end2726low > high → return low = 27 ✓ (high = 26 was a "no")
step 116…55…94count 7 → keep the left part
step 316…2526…3435 … 94count(25) = 2 → keep 26 … 34
step 6… 25262728 …count(26) = 2 → low = 27, high = 26 → stop

6 rounds of counting instead of 12 in Part B. The gap grows hugely with a wider range: about 14 rounds instead of 10,000 for her range.

9Complexity & remember

RememberRange = [top-left, bottom-right]. Count < k → low = mid + 1, else high = mid − 1. Return low. Count with the bottom-left staircase.

Part D · Revision page

approachideatimespace
A · flatten + sortcopy, sort, take index k − 1O(n² log n)O(n²)
B · try every valuefirst x from low upward with count(x) ≥ kO(range × n)O(1)
C · binary search on answerhalve the value range using count(mid)O(n log(range))O(1)
Problem 13 staircasecount(x) staircase here
starttop-right (or bottom-left)bottom-left (top-right also works)
on "too big"go leftgo up
on "small enough"go downadd row + 1, go right
returnsTrue when equalthe count after walking off the grid
costO(m + n)O(2n)
If you remember only 5 lines 1. The answer is between matrix[0][0] and matrix[n−1][n−1].
2. count(x) = how many numbers ≤ x, using a bottom-left staircase (add row + 1 and go right, or go up).
3. "count(x) ≥ k" is no…no, yes…yes → binary search over values.
4. count < k → low = mid + 1; else → high = mid − 1.
5. Return low: the first value with count ≥ k (always a matrix value).
Mistakes to avoid ✗ checking count == k (fails with repeated values)
✗ returning from the count function at the first "≤" cell (misses other columns, like 28 for x = 30)
✗ returning high at the end (it's the last "no")
✗ starting the count at top-left or bottom-right (two choices, can't decide)
✗ values[k] instead of values[k - 1] in the brute force
✗ trusting the Part B loop on a huge value range (TLE)
test it yourself (paste under any solution above)
matrix = [[16, 28, 60, 64],
          [22, 41, 63, 80],
          [27, 50, 66, 82],
          [36, 78, 83, 94]]
s = Solution()
print(s.kthSmallest(matrix, 3))                                # 27
print(s.kthSmallest(matrix, 10))                               # 64
print(s.kthSmallest([[1, 5, 9], [10, 11, 13], [12, 13, 15]], 8))  # 13
print(s.kthSmallest([[-5]], 1))                                # -5
print(s.kthSmallest([[1, 2], [2, 3]], 3))                      # 2 (repeats)

Based on this video: Kth Smallest Element in a Sorted Matrix | Binary Search on Answer