Find the Kth Smallest Sum of a Matrix With Sorted Rows — Hard Problem & Solution
mat has m rows, each sorted in non-decreasing order. An array sum picks exactly one element from every row and adds them up.
- Difficulty: Hard
- Topics: Arrays, Matrix, Binary Search, Heap
- Asked at: Amazon, Google, Uber
- Time limit: 2 s
- Memory limit: 256 MB
- Languages: JavaScript, TypeScript, Python, Java, C++, C, C#, Go, Kotlin, Swift, Rust, PHP and Ruby
Problem statement
mat has m rows, each sorted in non-decreasing order. An array sum picks exactly one element from every row and adds them up.
Return the k-th smallest array sum among all n^m possibilities.
Example 1
Input: mat = [[1,3,11],[2,4,6]], k = 5
Output: 7
Explanation: The smallest sums are 3, 5, 7, 7, 9 — the fifth is 7 (1+6 and 3+4 both give 7).
Example 2
Input: mat = [[1,3,11],[2,4,6]], k = 9
Output: 17
Example 3
Input: mat = [[1,10,10],[1,4,5],[2,3,6]], k = 7
Output: 9
Constraints
m == mat.lengthn == mat[i].length1 <= m, n <= 401 <= mat[i][j] <= 50001 <= k <= min(200, n^m)Each row of mat is sorted in non-decreasing order.
How to solve Find the Kth Smallest Sum of a Matrix With Sorted Rows
Fold the rows in one by one. After absorbing a row, keep only the k smallest partial sums — the rest can never grow into one of the k smallest totals, because every remaining row only adds non-negative amounts.
Approach
- Start with the single partial sum
0. - For each row, form every
partial + element, sort, and truncate tokentries. - After the last row, the answer is the
k-th entry.
Why it works
If a partial sum is not among the k smallest at some stage, then at least k partial sums are no larger, and each of them extends to a total no larger than this one's best extension (the remaining rows add the same minimum to all of them). So at least k totals beat it, and it cannot be the k-th smallest. Each fold handles at most k · n candidates, bounding the whole computation.
Complexity
- Time —
O(m · k · n · log(k · n)) - Space —
O(k · n)
Pitfalls
- Trying to enumerate all
n^mcombinations overflows any counter, let alone memory. - Truncating before sorting keeps the wrong
kcandidates. - A min-heap over
(sum, indices)also works but needs a visited set to avoid re-expanding states.
Reference solution
Python
from typing import List
def kthSmallest(mat: List[List[int]], k: int) -> int:
cur = [0]
for row in mat:
nxt = []
for c in cur:
for v in row:
nxt.append(c + v)
nxt.sort()
cur = nxt[:k]
return cur[k - 1]JavaScript
var kthSmallest = function(mat, k) {
var cur = [0];
for (var r = 0; r < mat.length; r++) {
var next = [];
for (var i = 0; i < cur.length; i++) {
for (var j = 0; j < mat[r].length; j++) next.push(cur[i] + mat[r][j]);
}
next.sort(function(a, b) { return a - b; });
cur = next.slice(0, k);
}
return cur[k - 1];
};Also on the editorial tab: C, C#, C++, Go, Java, Kotlin, PHP, Ruby, Rust, Swift, TypeScript.