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

求Deeplearning4j LSTM时间序列预测Java示例代码

我完全懂你的困扰——Deeplearning4j的官方示例确实偏重于图像处理和分类任务,找个纯Java的时间序列预测LSTM示例真的挺费劲的。刚好我之前做过类似的单变量时间序列预测,整理了一个完整的可运行示例,你可以直接参考:

完整LSTM时间序列预测示例代码
import org.deeplearning4j.datasets.iterator.impl.ListDataSetIterator;
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.deeplearning4j.nn.multilayer.MultiLayerNetwork;
import org.deeplearning4j.nn.weights.WeightInit;
import org.deeplearning4j.optimize.listeners.ScoreIterationListener;
import org.deeplearning4j.util.ModelSerializer;
import org.nd4j.linalg.activations.Activation;
import org.nd4j.linalg.api.ndarray.INDArray;
import org.nd4j.linalg.dataset.DataSet;
import org.nd4j.linalg.dataset.api.iterator.DataSetIterator;
import org.nd4j.linalg.factory.Nd4j;
import org.nd4j.linalg.learning.config.Adam;
import org.nd4j.linalg.lossfunctions.LossFunctions;
import org.nd4j.linalg.dataset.api.preprocessor.NormalizerStandardize;

import java.io.File;
import java.nio.file.Files;
import java.nio.file.Paths;
import java.util.ArrayList;
import java.util.List;
import java.util.stream.Collectors;

public class LstmTimeSeriesPrediction {
    public static void main(String[] args) throws Exception {
        // -------------------------- 1. 读取并预处理数据 --------------------------
        // 读取文本文件(假设每行一个数值)
        List<Double> rawData = Files.readAllLines(Paths.get("your_data.txt"))
                .stream()
                .map(Double::parseDouble)
                .collect(Collectors.toList());

        // 转换为NDArray并归一化(LSTM对数据范围极度敏感,必须做归一化)
        INDArray data = Nd4j.create(rawData.stream().mapToDouble(d -> d).toArray(), new int[]{rawData.size(), 1});
        NormalizerStandardize normalizer = new NormalizerStandardize();
        normalizer.fit(data);
        normalizer.transform(data);

        // 构造时间序列训练集:用前3个值预测第4个(可调整timeSteps参数)
        int timeSteps = 3;
        int inputSize = 1;
        List<DataSet> dataSets = new ArrayList<>();

        for (int i = 0; i < data.rows() - timeSteps; i++) {
            INDArray input = data.get(NDArrayIndex.interval(i, i + timeSteps), NDArrayIndex.all());
            INDArray label = data.get(NDArrayIndex.point(i + timeSteps), NDArrayIndex.all());
            dataSets.add(new DataSet(input, label));
        }

        // 创建批量迭代器
        int batchSize = 8;
        DataSetIterator iterator = new ListDataSetIterator<>(dataSets, batchSize);

        // -------------------------- 2. 构建LSTM网络 --------------------------
        MultiLayerConfiguration config = new NeuralNetConfiguration.Builder()
                .seed(42) // 固定随机种子保证结果可复现
                .updater(new Adam(0.001)) // 用Adam优化器,学习率可根据情况调整
                .weightInit(WeightInit.XAVIER)
                .list()
                .layer(0, new LSTM.Builder()
                        .nIn(inputSize)
                        .nOut(16) // LSTM隐藏层单元数,可调整
                        .activation(Activation.TANH)
                        .build())
                .layer(1, new RnnOutputLayer.Builder(LossFunctions.LossFunction.MSE)
                        .nIn(16)
                        .nOut(inputSize)
                        .activation(Activation.IDENTITY) // 回归任务用恒等激活
                        .build())
                .build();

        MultiLayerNetwork model = new MultiLayerNetwork(config);
        model.init();
        model.setListeners(new ScoreIterationListener(10)); // 每10轮打印一次损失值

        // -------------------------- 3. 训练模型 --------------------------
        int epochs = 100;
        for (int i = 0; i < epochs; i++) {
            iterator.reset();
            model.fit(iterator);
            System.out.println("完成第 " + (i+1) + " 轮训练,当前损失值:" + model.score());
        }

        // -------------------------- 4. 预测下一个数值 --------------------------
        // 取数据集最后一段序列作为输入
        INDArray lastSequence = data.get(NDArrayIndex.interval(data.rows() - timeSteps, data.rows()), NDArrayIndex.all());
        INDArray normalizedPrediction = model.output(lastSequence);

        // 反归一化得到真实数值
        normalizer.revert(normalizedPrediction);
        System.out.println("预测的下一个数值:" + normalizedPrediction.getDouble(0));
    }
}

关键细节说明

  • 归一化是核心:LSTM对输入数据的范围非常敏感,如果跳过归一化步骤,模型几乎不可能收敛。这里用NormalizerStandardize将数据转换为均值为0、方差为1的标准分布。
  • 滑动窗口构造:手动构建了时间序列的输入输出对——每timeSteps个连续数值作为模型输入,下一个数值作为预测标签,你可以根据自己的数据特性调整timeSteps的大小。
  • 网络适配回归任务:因为是数值预测(回归任务),所以输出层用MSE损失函数和恒等激活函数,这和分类任务的交叉熵+Softmax完全不同。
  • 批量训练:用ListDataSetIterator将手动构造的数据集转换为模型可处理的批量迭代器,提升训练效率。

内容的提问来源于stack exchange,提问作者Thorin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 03:43:52