Deeplearning4J RNN训练报错:RNN层需3D输入,实际传入2D
RNN训练维度不匹配异常修复
异常信息
java.lang.IllegalStateException: 3D input expected to RNN layer expected, got 2
需求说明
训练RNN模型,基于多组训练序列预测序列的下一个double类型值,使用随机数据生成特征,将序列最后一个值作为预测标签。
原错误代码
import org.nd4j.linalg.api.ndarray.INDArray; import org.nd4j.linalg.factory.Nd4j; import org.deeplearning4j.nn.multilayer.MultiLayerNetwork; import java.util.Random; import org.deeplearning4j.nn.conf.MultiLayerConfiguration; import org.deeplearning4j.nn.conf.NeuralNetConfiguration; import org.deeplearning4j.nn.conf.layers.LSTM; import org.deeplearning4j.nn.conf.layers.RnnOutputLayer; import org.nd4j.linalg.activations.Activation; import org.nd4j.linalg.dataset.DataSet; import org.nd4j.linalg.lossfunctions.LossFunctions; public class RnnPredictionExample { public static void main(String[] args) { //generate 100 rows of data that have 50 columns/features each DataSet trainingdata = getRandomDataset(100, 51, 1); // Train the RNN model... MultiLayerNetwork trainedModel = trainRnnModel(trainingdata, 50, 10, 1); // generate a sequence, and Perform next value prediction on the sequence double[] inputSequence = randomData(50, 1); double predictedValue = predictNextValue(trainedModel, inputSequence); System.out.println("Predicted Next Value: " + predictedValue); } public static MultiLayerNetwork trainRnnModel(DataSet trainingdataandlabels, int sequenceLength, int numHiddenUnits, int numEpochs) { // ... Create network configuration ... // Create and initialize the network MultiLayerConfiguration config = new NeuralNetConfiguration.Builder() //.seed(123) .list() .layer(new LSTM.Builder() .nIn(1) .nOut(50) .build() ) .layer(new RnnOutputLayer.Builder(LossFunctions.LossFunction.MSE) .activation(Activation.IDENTITY) .nIn(50) .nOut(1) // Set nOut to 1 .build() ) .build(); MultiLayerNetwork net = new MultiLayerNetwork(config); net.init(); for (int i = 0; i < numEpochs; i++) { net.fit(trainingdataandlabels); } return net; } public static double predictNextValue(MultiLayerNetwork trainedModel, double[] inputSequence) { INDArray inputArray = Nd4j.create(inputSequence); INDArray predicted = trainedModel.rnnTimeStep(inputArray); // Predicted value is the last element of the predicted sequence return predicted.getDouble(predicted.length() - 1); } static Random random = new Random(); public static double[] randomData(int length, int rangeMultiplier) { double[] out = new double[length]; for (int i = 0; i < out.length; i++) { out[i] = random.nextDouble() * rangeMultiplier; } return out; } //assumes labes is the last val in each sequence public static DataSet getRandomDataset(int numRows, int lengthEach, int rangeMultiplier) { INDArray training = Nd4j.zeros(numRows, lengthEach - 1); INDArray labels = Nd4j.zeros(numRows, 1); for (int i = 0; i < numRows; i++) { double[] randomData = randomData(lengthEach, rangeMultiplier); for (int j = 0; j < randomData.length - 1; j++) { training.putScalar(new int[]{i, j}, randomData[j]); } labels.putScalar(new int[]{i, 0}, randomData[randomData.length - 1]); } return new DataSet(training, labels); } }
修改后的可行代码
import org.nd4j.linalg.api.ndarray.INDArray; import org.nd4j.linalg.factory.Nd4j; import org.deeplearning4j.nn.multilayer.MultiLayerNetwork; import java.util.Random; import org.deeplearning4j.nn.conf.MultiLayerConfiguration; import org.deeplearning4j.nn.conf.NeuralNetConfiguration; import org.deeplearning4j.nn.conf.layers.LSTM; import org.deeplearning4j.nn.conf.layers.RnnOutputLayer; import org.nd4j.linalg.activations.Activation; import org.nd4j.linalg.dataset.DataSet; import org.nd4j.linalg.lossfunctions.LossFunctions; public class RnnPredictionExample { public static void main(String[] args) { //generate 100 rows of data that have 50 columns/features each DataSet trainingdata = getRandomDataset(100, 51, 1); // Train the RNN model... MultiLayerNetwork trainedModel = trainRnnModel(trainingdata, 50, 10, 1); // generate a sequence, and Perform next value prediction on the sequence double[] inputSequence = randomData(50, 1); double predictedValue = predictNextValue(trainedModel, inputSequence); System.out.println("Predicted Next Value: " + predictedValue); } public static MultiLayerNetwork trainRnnModel(DataSet trainingdataandlabels, int sequenceLength, int numHiddenUnits, int numEpochs) { // ... Create network configuration ... // Create and initialize the network MultiLayerConfiguration config = new NeuralNetConfiguration.Builder() //.seed(123) .list() .layer(new LSTM.Builder() .nIn(50) .nOut(1) .build() ) .layer(new RnnOutputLayer.Builder(LossFunctions.LossFunction.MSE) .activation(Activation.IDENTITY) .nIn(1) .nOut(1) // Set nOut to 1 .build() ) .build(); MultiLayerNetwork net = new MultiLayerNetwork(config); net.init(); for (int i = 0; i < numEpochs; i++) { net.fit(trainingdataandlabels); } return net; } public static double predictNextValue(MultiLayerNetwork trainedModel, double[] inputSequence) { // INDArray inputArray = Nd4j.create(inputSequence); INDArray inputArray = Nd4j.create(inputSequence).reshape(1, inputSequence.length, 1); INDArray predicted = trainedModel.rnnTimeStep(inputArray); // Predicted value is the last element of the predicted sequence return predicted.getDouble(predicted.length() - 1); } static Random random = new Random(); public static double[] randomData(int length, int rangeMultiplier) { double[] out = new double[length]; for (int i = 0; i < out.length; i++) { out[i] = random.nextDouble() * rangeMultiplier; } return out; } //assumes labes is the last val in each sequence public static DataSet getRandomDataset(int numRows, int lengthEach, int rangeMultiplier) { //INDArray training = Nd4j.zeros(numRows, lengthEach - 1); INDArray training = Nd4j.zeros(numRows, lengthEach - 1, 1); //INDArray labels = Nd4j.zeros(numRows, 1); INDArray labels = Nd4j.zeros(numRows, 1, 1); for (int i = 0; i < numRows; i++) { double[] randomData = randomData(lengthEach, rangeMultiplier); for (int j = 0; j < randomData.length - 1; j++) { // training.putScalar(new int[]{i, j}, randomData[j]); training.putScalar(new int[]{i, j, 0}, randomData[j]); } //labels.putScalar(new int[]{i, 0}, randomData[randomData.length - 1]); labels.putScalar(new int[]{i, 0, 0}, randomData[randomData.length - 1]); } return new DataSet(training, labels); } }
核心修复点
- 维度适配:RNN/LSTM层要求输入为3D张量,维度格式为
[批量大小, 序列长度, 特征数],原代码使用2D数组导致维度不匹配,修改训练数据、标签以及预测输入的数组维度为3D - 层维度匹配:修正LSTM层的
nIn和nOut参数,确保与输入特征数、后续层输入维度对应 - 预测输入处理:将单条预测序列reshape为符合要求的3D张量格式
内容的提问来源于stack exchange,提问作者user3059517
相关产品推荐
相关产品推荐

