二分查找求解矩阵中位数时答案错误,请求排查问题
问题描述
给定一个行排序的m * n大小的矩阵mat,其中m和n分别为矩阵的行数和列数。你的任务是找到并返回该矩阵的中位数。注意:m和n始终为奇数。
示例:
输入: n = 5, m = 5 mat = [ [ 1, 5, 7, 9, 11 ], [ 2, 3, 4, 8, 9 ], [ 4, 11, 14, 19, 20 ], [ 6, 10, 22, 99, 100 ], [ 7, 15, 17, 24, 28 ] ] 输出: 10
解释:将矩阵元素按排序顺序排列为数组后如下:
1 2 3 4 4 5 6 7 7 8 9 9 10 11 11 14 15 17 19 20 22 24 28 99 100
中位数为10,位于索引12处,因总元素数为25,索引12恰好是中间位置,故答案为10。
我尝试使用二分查找解决该问题,以下是编写的代码,但仅能通过部分测试用例,其余用例返回错误答案,且答案与正确值仅相差1。
import java.util.Arrays; public final class Solution { public static int findMedian(int matrix[][], int m, int n) { // Write your code here int[] pos= findPos(matrix,m,n); int req=(m*n)/2; while (pos[0]<=pos[1]) { int mid = (pos[0]+pos[1])/2; int smallEquals = blackbox(matrix,m,n,mid); if(smallEquals>=req){ pos[1]=mid-1; }else{ pos[0]=mid+1; } } return pos[0]; } public static int blackbox(int matrix[][], int m, int n, int mid){ int cnt=0; for (int i = 0; i < m; i++) { int low=0; int high=n-1; while(low<=high){ int mid1=(low+high)/2; if(matrix[i][mid1]<mid){ low=mid1+1; }else{ high=mid1-1; } } cnt+=low; } return cnt; } public static int[] findPos(int matrix[][], int m, int n){ int min=matrix[m-1][n-1]; int max=matrix[0][0]; for(int i=0; i<m; i++){ if(matrix[i][0]<min){ min=matrix[i][0]; } if(matrix[i][n-1]>max){ max=matrix[i][n-1]; } } return new int[]{min,max}; } }
错误原因分析
中位数位置计算错误:
你使用req=(m*n)/2,但对于奇数个元素,中位数是第(m*n+1)/2小的元素(1-based计数)。比如9个元素时,中位数是第5个元素,而非第4个,错误的位置计算会导致二分逻辑偏离目标。元素计数逻辑偏差:
你的blackbox函数计算的是小于mid的元素数量,但判断中位数需要的是小于等于mid的元素数量。当mid恰好等于中位数时,小于mid的元素数会小于目标位置,导致代码错误地认为需要增大mid,最终返回比正确值大1的结果。二分条件不匹配:
由于计数逻辑是“小于mid”,原有的smallEquals>=req条件无法正确对应中位数的位置判断,导致边界情况(如矩阵中存在大量重复元素)处理失效。
修正后的代码
import java.util.Arrays; public final class Solution { public static int findMedian(int matrix[][], int m, int n) { int[] pos = findPos(matrix, m, n); // 修正:计算第k小的元素(1-based),对应奇数个元素的中位数位置 int k = (m * n + 1) / 2; while (pos[0] <= pos[1]) { // 避免整数溢出的写法 int mid = pos[0] + (pos[1] - pos[0]) / 2; int countLessOrEqual = countElementsLessOrEqual(matrix, m, n, mid); if (countLessOrEqual >= k) { // 中位数在mid左侧或就是mid,缩小右边界 pos[1] = mid - 1; } else { // 中位数在mid右侧,缩小左边界 pos[0] = mid + 1; } } return pos[0]; } // 修正:计算矩阵中小于等于mid的元素总数 public static int countElementsLessOrEqual(int matrix[][], int m, int n, int mid) { int cnt = 0; for (int[] row : matrix) { int idx = Arrays.binarySearch(row, mid); if (idx < 0) { // 未找到mid,返回的是第一个大于mid的位置的负数索引转换 idx = -idx - 1; } else { // 找到mid时,需要定位到最后一个等于mid的位置 while (idx < n && row[idx] == mid) { idx++; } } cnt += idx; } return cnt; } // 优化初始化逻辑,更直观 public static int[] findPos(int matrix[][], int m, int n) { int min = matrix[0][0]; int max = matrix[0][n - 1]; for (int i = 1; i < m; i++) { if (matrix[i][0] < min) { min = matrix[i][0]; } if (matrix[i][n - 1] > max) { max = matrix[i][n - 1]; } } return new int[]{min, max}; } }
关键修正点说明
- 目标位置修正:将
req改为(m*n+1)/2,准确指向中位数对应的第k小元素位置。 - 计数逻辑修正:
countElementsLessOrEqual函数计算小于等于mid的元素数,使用Arrays.binarySearch简化实现,并处理了重复元素的情况。 - 二分条件匹配:根据小于等于mid的元素数与k的关系调整边界,确保正确收敛到中位数。
内容的提问来源于stack exchange,提问作者Great412
相关产品推荐
相关产品推荐

