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
相关产品推荐
相关产品推荐

