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

基于Deeplearning4J的Java比特币价格预测代码输出异常值问题排查

问题:基于Deeplearning4J的比特币价格预测输出异常低结果

我编写了一段基于Deeplearning4J的Java代码,输入包含比特币每日美元收盘价完整时间序列的数组,用于预测次日比特币价格。预期预测结果与当前价格不会相差过大,但实际得到了低于20美元的异常值。以下是核心代码、Maven依赖配置、CSV读取逻辑:

核心Java代码

public static void main(String[] args) {
    
    // Time series array
    double[] data = readCsv();

    // Hyperparameters
    int sequenceLength = 100; // The length of your sequences
    int numHiddenNodes = 256; // The number of hidden nodes
    int numEpochs = 10; // The number of epochs

    // Create input and output arrays
    INDArray input = Nd4j.create(1, 1, data.length - sequenceLength);
    INDArray output = Nd4j.create(1, 1, data.length - sequenceLength);
    for (int i = 0; i < data.length - sequenceLength; i++) {
        input.putScalar(new int[]{0, i % sequenceLength, 0}, data[i]);
        output.putScalar(new int[]{0, 0, i % sequenceLength}, data[i + sequenceLength]);
    }

    // Configure network
    MultiLayerConfiguration conf = new NeuralNetConfiguration.Builder()
        .weightInit(WeightInit.XAVIER)
        .updater(new Nadam())
        .list()
        .layer(0, new LSTM.Builder().nIn(1).nOut(numHiddenNodes)
            .activation(Activation.TANH).build())
        .layer(1, new RnnOutputLayer.Builder(LossFunctions.LossFunction.MSE)
            .activation(Activation.IDENTITY).nIn(numHiddenNodes).nOut(1).build())
        .build();

    MultiLayerNetwork net = new MultiLayerNetwork(conf);
    net.init();
    net.setListeners(new ScoreIterationListener(100));
    
    // Train the network
    for (int epoch = 0; epoch < numEpochs; epoch++) {
        net.fit(input, output);
    }

    // Initialize the input for next prediction with the last sequenceLength number of values from the input INDArray
    INDArray nextInput = Nd4j.create(new double[] {data[data.length - sequenceLength]}, new int[]{1, 1, 1});

    // Use trained model to predict the next value
    INDArray predicted = net.rnnTimeStep(nextInput);

    // Print out predicted value
    System.out.println("Predicted: " + predicted.getDouble(0));
}

Maven依赖配置

<dependency>
       <groupId>org.deeplearning4j</groupId>
       <artifactId>deeplearning4j-core</artifactId>
       <version>1.0.0-M2.1</version>
   </dependency>
   <dependency>
       <groupId>org.nd4j</groupId>
       <artifactId>nd4j-native-platform</artifactId>
       <version>1.0.0-M2.1</version>
   </dependency>
   <dependency>
       <groupId>org.apache.commons</groupId>
       <artifactId>commons-csv</artifactId>
       <version>1.8</version>
   </dependency>
   <dependency>
       <groupId>ch.qos.logback</groupId>
       <artifactId>logback-classic</artifactId>
       <version>1.2.3</version>
   </dependency>

CSV读取逻辑

private static double[] readCsv() {
    String csvFile = "/BTCDaily.csv"; // File in the resources folder
    List<Double> priceList = new ArrayList<>();

    try {
        URL url = App.class.getResource(csvFile);
        Reader in = new FileReader(url.toURI().getPath());
        Iterable<CSVRecord> records = CSVFormat.DEFAULT.withFirstRecordAsHeader().parse(in);
        for (CSVRecord record : records) {
            String price = record.get("Price");
            price = price.replace(",", "");  // remove commas
            priceList.add(Double.parseDouble(price));
        }
    } catch (Exception e) {
        e.printStackTrace();
    }

    return reverse(priceList.stream().mapToDouble(Double::doubleValue).toArray());
}

public static double[] reverse(double[] array) {
    int length = array.length;
    double[] reversed = new double[length];
    for (int i = 0; i < length; i++) {
        reversed[i] = array[length - i - 1];
    }
    return reversed;
}

问题排查与修复方案

1. 数据未做归一化处理

比特币价格数值范围极大(从早期几美元到数万美元),直接用原始值训练LSTM会导致模型难以收敛,甚至输出异常值。LSTM对输入数据的尺度极度敏感,必须先将数据缩放到[0,1]或[-1,1]区间,预测后再反归一化得到真实价格。

修复代码示例:

// 归一化处理
MinMaxScaler scaler = new MinMaxScaler(0, 1);
INDArray dataArray = Nd4j.create(data);
dataArray = scaler.fitTransform(dataArray);
double[] normalizedData = dataArray.toDoubleVector();

// 预测时反归一化
INDArray predictedScaled = net.output(nextInput);
INDArray predicted = scaler.inverseTransform(predictedScaled);
System.out.println("Predicted: " + predicted.getDouble(0));

2. 输入输出数据构造逻辑错误

当前代码构造input和output的索引逻辑完全错误:

  • i % sequenceLength会反复覆盖序列的前100个位置,无法生成长度为100的连续时序样本。
  • Deeplearning4J中LSTM的输入格式应为[样本数, 特征数, 序列长度],每个样本对应一段连续的历史序列,输出为该序列的下一个价格。

修复代码示例:

int numSamples = data.length - sequenceLength;
INDArray input = Nd4j.create(numSamples, 1, sequenceLength);
INDArray output = Nd4j.create(numSamples, 1);

for (int i = 0; i < numSamples; i++) {
    // 填充单特征的连续序列
    for (int j = 0; j < sequenceLength; j++) {
        input.putScalar(new int[]{i, 0, j}, data[i + j]);
    }
    // 对应输出为序列的下一个值
    output.putScalar(new int[]{i, 0}, data[i + sequenceLength]);
}

3. 预测输入序列不完整

当前预测仅传入单个数据点,而LSTM需要完整的sequenceLength长度的历史序列才能捕捉时序规律,单一点无法提供足够信息,必然导致异常输出。

修复代码示例:

// 构造完整的预测输入序列(最后100个历史数据)
INDArray nextInput = Nd4j.create(1, 1, sequenceLength);
for (int i = 0; i < sequenceLength; i++) {
    nextInput.putScalar(new int[]{0, 0, i}, data[data.length - sequenceLength + i]);
}
// 使用output方法获取预测结果
INDArray predicted = net.output(nextInput);

4. 训练轮数不足

仅训练10轮对于时序预测任务远远不够,LSTM通常需要数百甚至数千轮训练才能收敛。同时建议添加学习率衰减策略,避免模型训练后期震荡。

修复代码示例:

// 配置学习率及衰减策略
MultiLayerConfiguration conf = new NeuralNetConfiguration.Builder()
    .weightInit(WeightInit.XAVIER)
    .updater(new Nadam(0.001)) // 设置初始学习率
    .learningRateSchedule(new MapSchedule(ScheduleType.EPOCH, Collections.singletonMap(100, 0.0001))) // 100轮后衰减学习率
    .list()
    .layer(0, new LSTM.Builder().nIn(1).nOut(numHiddenNodes).activation(Activation.TANH).build())
    .layer(1, new RnnOutputLayer.Builder(LossFunctions.LossFunction.MSE).activation(Activation.IDENTITY).nIn(numHiddenNodes).nOut(1).build())
    .build();

// 增加训练轮数
int numEpochs = 200;
for (int epoch = 0; epoch < numEpochs; epoch++) {
    net.fit(input, output);
}

5. 模型结构过于简单

单一层LSTM难以捕捉比特币价格复杂的波动规律,可添加多层LSTM及Dropout层防止过拟合,提升模型拟合能力。

修复代码示例:

MultiLayerConfiguration conf = new NeuralNetConfiguration.Builder()
    .weightInit(WeightInit.XAVIER)
    .updater(new Nadam(0.001))
    .list()
    .layer(0, new LSTM.Builder().nIn(1).nOut(numHiddenNodes).activation(Activation.TANH).build())
    .layer(1, new Dropout.Builder(0.2).build()) // 添加Dropout抑制过拟合
    .layer(2, new LSTM.Builder().nIn(numHiddenNodes).nOut(numHiddenNodes/2).activation(Activation.TANH).build())
    .layer(3, new RnnOutputLayer.Builder(LossFunctions.LossFunction.MSE).activation(Activation.IDENTITY).nIn(numHiddenNodes/2).nOut(1).build())
    .build();

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 07:55:02