Java实现神经网络时遭遇ArrayIndexOutOfBoundsException:矩阵乘法报错排查求助
Java神经网络矩阵乘法中的ArrayIndexOutOfBoundsException问题解决
我在尝试用Java实现神经网络时遇到了ArrayIndexOutOfBoundsException,错误发生在矩阵相乘的操作中。我已经试过调整循环条件,但问题依然存在,希望能得到帮助排查。
报错信息
Exception in thread "main" java.lang.ArrayIndexOutOfBoundsException
相关代码
驱动类代码
public class Driver { static double [][]X ={ {0, 0}, {1, 0}, {0, 1}, {1, 1} }; static double [][] Y = { {0},{1},{1},{0} }; public static void main(String[] args) { NNMain nn = new NNMain(2,10,1); // 2,10,1 List<Double> output; nn.fit(X, Y, 50000); double [][] inputMatrix = { {0,0},{0,1},{1,0},{1,1} }; for(double d[]:inputMatrix) { output = nn.predict(d); System.out.println(output.toString()); } } }
出错的矩阵乘法方法
public static Matrix multiply(Matrix a, Matrix b) { Matrix temp = new Matrix(a.r, a.c); for (int i = 0; i < temp.r; i++) { for (int j = 0; j < temp.c; j++) { double sum=0; for(int k=0; k < a.c; k++) { sum += a.data[i][k] * b.data[k][j]; // ERROR HERE } temp.data[i][j]=sum; } } return temp; }
问题根源
矩阵乘法的核心规则是:第一个矩阵的列数必须等于第二个矩阵的行数,而结果矩阵的维度应该是「第一个矩阵的行数 × 第二个矩阵的列数」。
你的代码里犯了一个关键错误:创建结果矩阵temp时,错误地使用了a.c作为结果矩阵的列数,而正确的列数应该是b.c。
举个实际场景的例子:当你计算隐藏层到输出层的矩阵乘法时,隐藏层输出矩阵是4×10,输出层权重矩阵是10×1,正确的结果矩阵应该是4×1。但你的代码会创建一个4×10的矩阵,当循环j到1时,访问b.data[k][j]就会超出b的列数(仅为1),直接触发数组越界异常。
解决方案
修改矩阵乘法方法中结果矩阵的创建逻辑,同时增加参数合法性校验,避免后续维度不匹配的问题:
public static Matrix multiply(Matrix a, Matrix b) { // 提前校验矩阵乘法的合法性:a的列数必须等于b的行数 if (a.c != b.r) { throw new IllegalArgumentException("矩阵乘法失败:a的列数必须等于b的行数"); } // 结果矩阵的维度为「a的行数 × b的列数」 Matrix temp = new Matrix(a.r, b.c); for (int i = 0; i < temp.r; i++) { for (int j = 0; j < temp.c; j++) { double sum = 0; // k的遍历范围是a的列数(即b的行数) for (int k = 0; k < a.c; k++) { sum += a.data[i][k] * b.data[k][j]; } temp.data[i][j] = sum; } } return temp; }
额外建议
- 检查
Matrix类的构造方法:确保它正确初始化data数组的大小为行数×列数,示例如下:public class Matrix { int r; // 行数 int c; // 列数 double[][] data; public Matrix(int rows, int cols) { this.r = rows; this.c = cols; this.data = new double[rows][cols]; } } - 确认神经网络各层维度匹配:比如输入层2个神经元、隐藏层10个神经元,那么输入到隐藏层的权重矩阵应为
2×10,隐藏层到输出层的权重矩阵应为10×1,确保调用multiply时始终满足维度要求。
内容的提问来源于stack exchange,提问作者WickedPlayer53294
相关产品推荐
相关产品推荐

