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

如何在Deeplearning4j中使用自定义数据模型创建DataSetIterator?

实现自定义股票数据模型的DataSetIterator(Deeplearning4j)

嘿,这个需求我熟——把自定义的股票数据模型转换成Deeplearning4j能用的DataSetIterator,这在量化交易或者股票价格预测的项目里太常见了。我给你一步步拆解,保证你能快速用上:

第一步:给股票数据模型加便捷提取方法

首先,你的自定义Java类已经包含了所有股票相关字段,咱们给它加两个方法,方便快速提取特征向量和预测标签(比如你要预测收盘价的话):

public class StockDataPoint {
    private long timestamp;
    private double open;
    private double close;
    private double high;
    private double low;
    private double volume;
    private double techIndicator1; // 比如RSI
    private double techIndicator2; // 比如MACD

    // 构造函数、Getter、Setter请按需实现

    // 把所有特征转换成double数组,供DL4J使用
    public double[] getFeatureVector() {
        return new double[]{open, high, low, volume, techIndicator1, techIndicator2};
    }

    // 定义模型要预测的标签,这里以收盘价为例
    public double getLabel() {
        return close;
    }
}

第二步:实现自定义DataSetIterator

DL4J的DataSetIterator是接口,咱们直接继承BaseDatasetIterator可以省掉很多模板代码,核心就是实现批次数据的生成逻辑:

import org.nd4j.linalg.dataset.DataSet;
import org.nd4j.linalg.dataset.api.DataSetPreProcessor;
import org.nd4j.linalg.dataset.api.iterator.BaseDatasetIterator;
import org.nd4j.linalg.factory.Nd4j;

import java.util.List;

public class StockDataSetIterator extends BaseDatasetIterator {

    private final List<StockDataPoint> stockDataList;
    private final int featureCount;
    private final int labelCount;
    private int currentCursor = 0;

    // 构造函数:批次大小、总样本数、股票数据列表
    public StockDataSetIterator(int batchSize, int totalSamples, List<StockDataPoint> stockData) {
        super(batchSize, totalSamples);
        this.stockDataList = stockData;
        this.featureCount = stockData.get(0).getFeatureVector().length;
        this.labelCount = 1; // 这里假设只预测一个值(收盘价)
    }

    @Override
    public DataSet next(int batchSize) {
        // 计算当前批次的结束位置,避免越界
        int batchEnd = Math.min(currentCursor + batchSize, totalExamples);
        int sampleCount = batchEnd - currentCursor;

        // 初始化特征和标签数组
        double[][] features = new double[sampleCount][featureCount];
        double[][] labels = new double[sampleCount][labelCount];

        // 填充数据到数组中
        for (int i = currentCursor; i < batchEnd; i++) {
            StockDataPoint dataPoint = stockDataList.get(i);
            features[i - currentCursor] = dataPoint.getFeatureVector();
            labels[i - currentCursor][0] = dataPoint.getLabel();
        }

        // 转换成DL4J的DataSet对象
        DataSet batchDataSet = new DataSet(
                Nd4j.create(features),
                Nd4j.create(labels)
        );

        // 如果有预处理器(比如归一化),应用到批次上
        if (preProcessor != null) {
            preProcessor.preProcess(batchDataSet);
        }

        // 更新游标,准备下一批次
        currentCursor = batchEnd;
        return batchDataSet;
    }

    @Override
    public void reset() {
        currentCursor = 0;
    }

    @Override
    public boolean hasNext() {
        return currentCursor < totalExamples;
    }

    @Override
    public int inputColumns() {
        return featureCount;
    }

    @Override
    public int totalOutcomes() {
        return labelCount;
    }

    @Override
    public void setPreProcessor(DataSetPreProcessor preProcessor) {
        this.preProcessor = preProcessor;
    }

    @Override
    public DataSetPreProcessor getPreProcessor() {
        return preProcessor;
    }

    @Override
    public boolean resetSupported() {
        return true;
    }

    @Override
    public boolean asyncSupported() {
        return true;
    }
}

第三步:使用迭代器接入DL4J网络

现在你已经把JSON转成了List<StockDataPoint>,接下来就可以创建迭代器,喂给你的DL4J模型了:

// 假设你已经通过JSON解析得到了股票数据列表
List<StockDataPoint> stockData = ...;

// 创建自定义迭代器:批次大小设为32,总样本数是数据列表的长度
StockDataSetIterator dataIterator = new StockDataSetIterator(32, stockData.size(), stockData);

// !重要:给数据加归一化(必须做,不然模型训练效果会很差)
NormalizerStandardize normalizer = new NormalizerStandardize();
normalizer.fit(dataIterator); // 先拟合数据分布
dataIterator.setPreProcessor(normalizer); // 绑定到迭代器

// 喂给DL4J网络训练
MultiLayerNetwork yourModel = ...; // 这里是你已经定义好的网络结构
yourModel.fit(dataIterator);

额外:如果是时序场景(比如用LSTM)

如果你的任务是用过去N天的数据预测未来价格,那需要把数据处理成序列格式。这里给你一个快速生成时序样本的方法:

import org.nd4j.linalg.dataset.DataSet;
import org.nd4j.linalg.factory.Nd4j;

import java.util.ArrayList;
import java.util.List;

public class TimeSeriesStockUtils {
    // timeStep:用过去多少天的数据作为输入
    public static List<DataSet> createTimeSeriesSamples(List<StockDataPoint> rawData, int timeStep) {
        List<DataSet> sequenceSamples = new ArrayList<>();

        for (int i = 0; i <= rawData.size() - timeStep - 1; i++) {
            // 提取timeStep天的特征(形状:[时间步长, 特征数])
            double[][] sequenceFeatures = new double[timeStep][rawData.get(0).getFeatureVector().length];
            for (int j = 0; j < timeStep; j++) {
                sequenceFeatures[j] = rawData.get(i + j).getFeatureVector();
            }

            // 标签是timeStep天后的收盘价
            double[] label = new double[]{rawData.get(i + timeStep).getLabel()};

            // 转换成DL4J的序列DataSet
            sequenceSamples.add(new DataSet(
                    Nd4j.create(sequenceFeatures),
                    Nd4j.create(label)
            ));
        }
        return sequenceSamples;
    }
}

之后你可以基于这个序列样本列表,实现SequenceDataSetIterator,或者直接用ListDataSetIterator来包装序列数据,适配LSTM等循环神经网络。

关键注意事项

  • 数据归一化:绝对不能省略,股票数据的数值范围差异大(比如成交量和RSI),归一化能让模型更快收敛。
  • 缺失值处理:如果你的股票数据里有缺失字段,提前用均值、中位数或者插值法填充,不然会导致训练报错。
  • 样本顺序:股票数据是时序数据,训练时不要打乱样本顺序,否则模型学不到时间依赖关系。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 09:06:13