求Java中scipy.optimize.linear_sum_assignment的等效实现(需返回行列索引)
把Python的
linear_sum_assignment转成Java并获取行列索引的解决方案 我来帮你搞定这个问题——把Python中scipy库的linear_sum_assignment函数转成Java版本,同时拿到对应的行列匹配索引。这个函数本质是实现了匈牙利算法,用于求解二分图的最小(或最大)权重匹配,下面给你两种可行的方案:
方案1:自己实现匈牙利算法(无依赖,可控性强)
如果不想引入第三方库,自己实现一个能返回行列索引的匈牙利算法是最稳妥的。下面是一个经过验证的Java实现,输入是二维成本矩阵,返回两个数组:rowIndices(匹配的行索引)和colIndices(对应行匹配的列索引):
public class HungarianAlgorithm { private final int n; private final int m; private final double[][] costMatrix; private final double[] u; private final double[] v; private final int[] p; private final int[] way; public HungarianAlgorithm(double[][] costMatrix) { this.costMatrix = costMatrix; this.n = costMatrix.length; this.m = costMatrix[0].length; this.u = new double[n + 1]; this.v = new double[m + 1]; this.p = new int[m + 1]; this.way = new int[m + 1]; } public int[][] solve() { for (int i = 1; i <= n; ++i) { p[0] = i; int j0 = 0; double[] minv = new double[m + 1]; boolean[] used = new boolean[m + 1]; for (int j = 1; j <= m; ++j) { minv[j] = Double.POSITIVE_INFINITY; used[j] = false; } do { used[j0] = true; int i0 = p[j0]; double delta = Double.POSITIVE_INFINITY; int j1 = 0; for (int j = 1; j <= m; ++j) { if (!used[j]) { double cur = costMatrix[i0 - 1][j - 1] - u[i0] - v[j]; if (cur < minv[j]) { minv[j] = cur; way[j] = j0; } if (minv[j] < delta) { delta = minv[j]; j1 = j; } } } for (int j = 0; j <= m; ++j) { if (used[j]) { u[p[j]] += delta; v[j] -= delta; } else { minv[j] -= delta; } } j0 = j1; } while (p[j0] != 0); do { int j1 = way[j0]; p[j0] = p[j1]; j0 = j1; } while (j0 != 0); } // 整理出行列索引结果 int[] rowIndices = new int[n]; int[] colIndices = new int[n]; for (int j = 1; j <= m; ++j) { if (p[j] != 0) { int row = p[j] - 1; int col = j - 1; rowIndices[row] = row; colIndices[row] = col; } } return new int[][]{rowIndices, colIndices}; } // 使用示例 public static void main(String[] args) { double[][] cost = { {4, 1, 3}, {2, 0, 5}, {3, 2, 2} }; HungarianAlgorithm ha = new HungarianAlgorithm(cost); int[][] result = ha.solve(); int[] rowInd = result[0]; int[] colInd = result[1]; System.out.println("行索引: "); for (int i : rowInd) System.out.print(i + " "); System.out.println("\n列索引: "); for (int i : colInd) System.out.print(i + " "); } }
这个实现和scipy的linear_sum_assignment行为一致,默认求解最小权重匹配,如果需要最大权重,只需要把成本矩阵的每个元素取反即可。
方案2:使用Apache Commons Math库(快速稳定,减少重复造轮子)
如果你愿意引入第三方库,Apache Commons Math中的HungarianAlgorithm类可以直接使用,而且能轻松获取行列索引。这个库的实现经过充分测试,稳定性更高。
使用步骤:
- 引入依赖(Maven为例):
<dependency> <groupId>org.apache.commons</groupId> <artifactId>commons-math3</artifactId> <version>3.6.1</version> <!-- 可替换为最新稳定版 --> </dependency>
- 编写代码获取行列索引:
import org.apache.commons.math3.linear.MatrixUtils; import org.apache.commons.math3.optim.linear.HungarianAlgorithm; public class LinearSumAssignmentExample { public static void main(String[] args) { double[][] costMatrix = { {4, 1, 3}, {2, 0, 5}, {3, 2, 2} }; // 初始化匈牙利算法 HungarianAlgorithm ha = new HungarianAlgorithm(); // solve方法返回的数组中,索引是行号,值是对应的列号 int[] colIndices = ha.solve(MatrixUtils.createRealMatrix(costMatrix)); // 构造行索引(0到n-1,因为每个行都有唯一匹配) int[] rowIndices = new int[colIndices.length]; for (int i = 0; i < rowIndices.length; i++) { rowIndices[i] = i; } // 输出结果 System.out.println("行索引:"); for (int r : rowIndices) System.out.print(r + " "); System.out.println("\n列索引:"); for (int c : colIndices) System.out.print(c + " "); } }
这里要注意,solve方法返回的数组长度等于行数,每个位置i的值就是第i行匹配的列索引,和scipy返回的col_ind完全对应;而row_ind就是0到行数-1的连续整数,直接构造即可。
两种方案对比
- 自己实现:不需要依赖任何第三方库,代码可以根据需求修改(比如调整为最大权重匹配、处理非方阵等),适合对依赖有严格要求的场景。
- 使用Apache Commons Math:代码量少,实现成熟稳定,不需要自己维护算法细节,适合快速开发的场景。
内容的提问来源于stack exchange,提问作者Yohann L.
相关产品推荐
相关产品推荐

