How to find the Kth smallest sum in a sorted MxN matrix

Viewed 1386

I've seen solutions on how to find the Kth smallest element in a sorted matrix, and I've also seen solutions on how to find the Kth smallest sum in two arrays.

But I found a question recently that asks to find the Kth smallest sum in a sorted MxN matrix. The sum must be made up of one element from each row. I'm really struggling develop anything close to a working solution, let alone a brute force solution. Any help would be greatly appreciated!

I thought this would be some kind of a heap problem... But perhaps it is a graph problem? I'm not that great with graphs.

4 Answers

I assume by "sorted MxN matrix", you mean each row of the matrix is sorted. If you already know how to merge 2 rows and take only the first K elements, you can do that same procedure to merge each and every row of the matrix. Ignore the Java conversion between int[] and List, the following code should work.

class Solution {

/**
 * Runtime O(m * k * logk)
 */ 
public int kthSmallest(int[][] mat, int k) {
    List<Integer> row = IntStream.of(mat[0]).boxed().collect(Collectors.toList());
    for (int i = 1; i < mat.length; i++) {
        row = kthSmallestPairs(row, mat[i], k);
    }
    return row.get(k - 1);
}

/**
 * A pair is formed from one num of n1 and one num of n2. Find the k-th smallest sum of these pairs
 * Queue size is maxed at k, hence this method run O(k logk)
 */
List<Integer> kthSmallestPairs(List<Integer> n1, int[] n2, int k) {
    // 0 is n1's num,     1 is n2's num,     2 is n2's index
    Queue<int[]> que = new PriorityQueue<>((a, b) -> a[0] + a[1] - b[0] - b[1]);

    // first pair each num in n1 with the 0-th num of n2. Don't need to do more than k elements because those greater
    // elements will never have a chance
    for (int i = 0; i < n1.size() && i < k; i++) {
        que.add(new int[] {n1.get(i), n2[0], 0});
    }

    List<Integer> res = new ArrayList<>();
    while (!que.isEmpty() && k-- > 0) {
        int[] top = que.remove();
        res.add(top[0] + top[1]);
        // index of n2 is top[2]
        if (top[2] < n2.length - 1) {
            int nextN2Idx = top[2] + 1;
            que.add(new int[] {top[0], n2[nextN2Idx], nextN2Idx});
        }
    }

    return res;
}

}

You can make a minHeap priority queue and save the sums and the corresponding index of rows in it. Then, once you pop the smallest sum so far, you can examine the next candidates for the smallest sum by incrementing index of each row by one.
Here are the data structures that you would need.

typedef pair<int,vector<int>> pi;
priority_queue<pi,vector<pi>,greater<pi>> pq;

You can try the question now, for help I have also added the code that I have written for this problem.

typedef pair<int,vector<int>> pi;

int kthSmallest(vector<vector<int>>& mat, int k) {
    int m=mat.size();
    int n=mat[0].size();
    priority_queue<pi,vector<pi>,greater<pi>> pq;
    int sum=0;
    for(int i=0;i<m;i++)
        sum+=mat[i][0];
    vector<int> v;
    for(int i=0;i<m;i++)
        v.push_back(0);
    pq.push({sum,v});
    int count=1;
    int ans=sum;
    unordered_map<string,int> meep;
    string s;
    for(int i=0;i<m;i++)
        s+="0";
    meep[s]=1;
    while(count<=k)
    {
        ans=pq.top().first;
        v=pq.top().second;
        // cout<<ans<<endl;
        // for(int i=0;i<v.size();i++)
        //     cout<<v[i]<<" ";
        // cout<<endl;
        pq.pop();
        for(int i=0;i<m;i++)
        {
            vector<int> temp;
            sum=0;
            int flag=0;
            string luuul;
            for(int j=0;j<m;j++)
            {
                if(i==j&&v[j]<n-1)
                {
                    sum+=mat[j][v[j]+1];
                    temp.push_back(v[j]+1);
                    luuul+=to_string(v[j]+1);
                }
                else if(i==j&&v[j]==n-1)
                {
                    flag=1;
                    break;
                }
                else
                {
                    sum+=mat[j][v[j]];
                    temp.push_back(v[j]);
                    luuul+=to_string(v[j]);
                }
            }
            if(!flag)
            {
                if(meep[luuul]==0)
                    pq.push({sum,temp});
                meep[luuul]=1;
            }
        }
        // cout<<endl;
        count++;
    }
    return ans;
}

or every row we calculate all possible sums but keep the k smallest. We can use quickselect to do so in linear time.

The complexity below should be: O(n * m * k).

class Solution {
public:
    int kthSmallest(vector<vector<int>>& mat, int k) {
        vector<int> sums = { 0 }, cur = {};

        for (const auto& row : mat) {
            for (const int cel : row) {
                for (const int sum : sums) {
                    cur.push_back(cel + sum);
                }
            }

            int nth = min((int ) cur.size(), k);

            nth_element(cur.begin(), cur.begin() + nth, cur.end());

            sums.clear();
            copy(cur.begin(), cur.begin() + nth, back_inserter(sums));
            cur.clear();
        }

        return *max_element(sums.begin(), sums.end());
    }
};

So the algorithm goes like this:

We know that elements of each row are sorted, so the minimum sum would be given by selecting the 1st element from each row.

We make a set storing {sum,vector of positions of the current elements we've chosen} sorted wrt to the sum.

So for finding the kth smallest sum, we repeat the following steps k-1 times: i) Take the element at the beginning of the set and erase it. ii) Find the next possible combinations with respect to the previous combination.

After exiting the loop return the sum of combination present at the beginning of the set.

The algorithm (using set) is properly explained with dry runs of test case containing all the corner case conditions. Do watch this youtube video by alGOds : https://youtu.be/ZYlVCy_vRp8

Related