You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.14 00:22:03