如何在Deeplearning4j与DataVec中拆分DataSetIterator为训练测试迭代器?(原方法已弃用)
Hey there! I feel your pain—deprecated methods can throw a real wrench into your workflow, especially when you're working with time series data in DL4J and DataVec. Let's walk through the current, recommended ways to split your DataSetIterator into training and test iterators.
1. For Small-to-Medium Datasets: Load Full Data & Split
If your time series dataset fits comfortably in memory, the simplest approach is to first collect all data from your iterator into a single DataSet, split it, then convert the split datasets back into DataSetIterator instances.
Here's how to do it:
// Step 1: Collect all data from the original iterator into one DataSet DataSet fullDataSet = null; while (yourOriginalIterator.hasNext()) { DataSet batch = yourOriginalIterator.next(); if (fullDataSet == null) { fullDataSet = batch; } else { fullDataSet = fullDataSet.merge(batch); } } // Reset the original iterator if you need to reuse it later yourOriginalIterator.reset(); // Step 2: Split the full dataset (e.g., 80% training, 20% test) SplitTestAndTrain split = fullDataSet.splitTestAndTrain(0.8); DataSet trainDataSet = split.getTrain(); DataSet testDataSet = split.getTest(); // Step 3: Convert split datasets back to DataSetIterator int batchSize = yourOriginalIterator.batch(); DataSetIterator trainIterator = new ListDataSetIterator<>(trainDataSet.asList(), batchSize); DataSetIterator testIterator = new ListDataSetIterator<>(testDataSet.asList(), batchSize);
2. For Large Datasets: Use RecordReader with Input Splits
If your dataset is too big to load entirely into memory, go back to the source with DataVec's RecordReader and split your input data directly before creating iterators. This avoids loading everything into RAM at once.
Assuming you're using a CSVRecordReader (adjust for your data format):
// Step 1: Define your full data input split (e.g., a CSV file) InputSplit fullDataSplit = new FileSplit(new File("your-full-time-series-data.csv")); // Step 2: Split the input into training and test splits (80/20 ratio) InputSplit[] splits = fullDataSplit.sample(0.8, 0.2); InputSplit trainSplit = splits[0]; InputSplit testSplit = splits[1]; // Step 3: Create RecordReaders and iterators for each split // Training iterator RecordReader trainReader = new CSVRecordReader(1, ","); // Adjust header lines and delimiter trainReader.initialize(trainSplit); DataSetIterator trainIterator = new RecordReaderDataSetIterator( trainReader, yourBatchSize, yourLabelIndex, yourNumClasses ); // Test iterator RecordReader testReader = new CSVRecordReader(1, ","); testReader.initialize(testSplit); DataSetIterator testIterator = new RecordReaderDataSetIterator( testReader, yourBatchSize, yourLabelIndex, yourNumClasses );
Critical Note for Time Series Data!
Never use random splits for time series data—shuffling breaks the temporal order and causes data leakage (training on future data to predict the past). Instead, split sequentially:
// Step 1: Calculate how many examples go to training int totalExamples = yourOriginalIterator.totalExamples(); int trainSize = (int) (totalExamples * 0.8); // 80% training int batchSize = yourOriginalIterator.batch(); // Step 2: Collect training data (first N examples in order) DataSet trainDataSet = null; int trainBatches = trainSize / batchSize; for (int i = 0; i < trainBatches; i++) { DataSet batch = yourOriginalIterator.next(); trainDataSet = (trainDataSet == null) ? batch : trainDataSet.merge(batch); } // Step 3: Collect remaining examples as test data DataSet testDataSet = null; while (yourOriginalIterator.hasNext()) { DataSet batch = yourOriginalIterator.next(); testDataSet = (testDataSet == null) ? batch : testDataSet.merge(batch); } // Step 4: Convert to iterators DataSetIterator trainIterator = new ListDataSetIterator<>(trainDataSet.asList(), batchSize); DataSetIterator testIterator = new ListDataSetIterator<>(testDataSet.asList(), batchSize);
内容的提问来源于stack exchange,提问作者user14717506

