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

DL4J中DataSetIterator.getLabels()返回null的问题求助

问题原因与解决方案

你调用iterator.getLabels()返回null,是因为RecordReaderDataSetIterator默认不会自动生成或读取标签列表,必须手动配置让它读取CSV表头,或者手动传入标签列表。

针对你的代码,有两种解决方式:

方式1:如果CSV文件第一行是表头(标签名称)

修改CSVRecordReader和迭代器的构建代码,让迭代器读取表头作为标签:

int numLinesToSkip = 0;
char delimeter = ',';
int labelIndex = 0;
int numClasses = 9;
int trainBatchSize = 10000;

// 初始化CSVRecordReader,不跳过行(表头在第一行)
RecordReader recordReader = new CSVRecordReader(numLinesToSkip, delimeter);
recordReader.initialize(new FileSplit(new File("myFile.csv")));

List<DataSet> trainingData = new ArrayList<>();
List<DataSet> testingData = new ArrayList<>();

// 构建迭代器时添加.header(true),让它读取第一行作为标签
DataSetIterator iterator = new RecordReaderDataSetIterator.Builder(recordReader, trainBatchSize)
    .classification(labelIndex, numClasses)
    .header(true) // 关键:启用表头读取,自动提取标签
    .build();

while (iterator.hasNext()) {
    DataSet allData = iterator.next(); // 修复:这里要加括号()
    allData.shuffle();
    SplitTestAndTrain testAndTrain = allData.splitTestAndTrain(0.65);
    trainingData.add(testAndTrain.getTrain());
    testingData.add(testAndTrain.getTest());
    System.out.println(iterator.getLabels()); // 现在会输出标签列表
}

方式2:如果CSV没有表头,手动指定标签列表

如果你的CSV没有表头行,需要自己创建标签列表并传入迭代器:

int numLinesToSkip = 0;
char delimeter = ',';
int labelIndex = 0;
int numClasses = 9;
int trainBatchSize = 10000;

RecordReader recordReader = new CSVRecordReader(numLinesToSkip, delimeter);
recordReader.initialize(new FileSplit(new File("myFile.csv")));

List<DataSet> trainingData = new ArrayList<>();
List<DataSet> testingData = new ArrayList<>();

// 手动创建对应9个类别的标签列表
List<String> manualLabels = Arrays.asList("类别1", "类别2", "类别3", "类别4", "类别5", "类别6", "类别7", "类别8", "类别9");

DataSetIterator iterator = new RecordReaderDataSetIterator.Builder(recordReader, trainBatchSize)
    .classification(labelIndex, numClasses)
    .setLabels(manualLabels) // 关键:手动传入标签列表
    .build();

while (iterator.hasNext()) {
    DataSet allData = iterator.next(); // 修复:添加括号()
    allData.shuffle();
    SplitTestAndTrain testAndTrain = allData.splitTestAndTrain(0.65);
    trainingData.add(testAndTrain.getTrain());
    testingData.add(testAndTrain.getTest());
    System.out.println(iterator.getLabels()); // 输出手动设置的标签列表
}

额外注意点

  • 代码里的iterator.next是语法错误,必须改成iterator.next()才能正确获取DataSet。
  • 如果你用的是旧版本DeepLearning4J,header(true)可能需要替换为手动从CSVRecordReader获取表头:List<String> labels = Arrays.asList(recordReader.getLabels());,再用.setLabels(labels)传入迭代器。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 07:54:11