DeepLearning4J数据获取与线性回归模型训练问题咨询
Hey there! Let's work through this step by step—building a minimal neural net for linear regression with your custom CSV data and fixing that DataSetIterator snag. I'll assume you're using Deeplearning4j since DataSetIterator is a core class there, but the general logic applies to other frameworks too with minor tweaks.
1. First: Fix Your CSV Format for Easy Reading
Your CSV uses ; to separate samples (like 1,2; 3,4; 5,6), but most data iterators expect one sample per line, with commas splitting features and labels. Let's fix that quickly:
You can run a tiny Python script to reformat the file:
with open('test.csv', 'r') as f: raw_content = f.read() # Replace "; " with newlines to put each sample on its own line processed_content = raw_content.replace('; ', '\n') with open('processed_test.csv', 'w') as f: f.write(processed_content)
Now your file will look like this (perfect for the iterator):
1,2 3,4 5,6 ...
2. Configure DataSetIterator Correctly
Now let's set up the iterator to pull data properly. We'll use CSVRecordReader and RecordReaderDataSetIterator to handle parsing:
import org.deeplearning4j.datasets.datavec.RecordReaderDataSetIterator; import org.datavec.api.records.reader.impl.csv.CSVRecordReader; import org.datavec.api.split.FileSplit; import java.io.File; // Initialize reader: skip 0 lines (no header), split columns with commas CSVRecordReader recordReader = new CSVRecordReader(0, ','); recordReader.initialize(new FileSplit(new File("processed_test.csv"))); // Build the iterator: // - Batch size: 10 (adjust based on your data size) // - Number of features: 1 (each input is a single number) // - Number of labels: 1 (each label is the input + 1) DataSetIterator dataIterator = new RecordReaderDataSetIterator( recordReader, 10, 1, 1 );
If you really don't want to reformat the CSV, you could mess with custom delimiters, but one sample per line is way less error-prone.
3. Build the Minimal Linear Regression Neural Net
Linear regression doesn't need fancy hidden layers—just an output layer with a linear activation function (since we're predicting a continuous value) and mean squared error (MSE) loss:
import org.deeplearning4j.nn.conf.MultiLayerConfiguration; import org.deeplearning4j.nn.conf.NeuralNetConfiguration; import org.deeplearning4j.nn.conf.layers.OutputLayer; import org.deeplearning4j.nn.multilayer.MultiLayerNetwork; import org.nd4j.linalg.activations.Activation; import org.nd4j.linalg.lossfunctions.LossFunctions; import org.nd4j.linalg.learning.config.Sgd; MultiLayerConfiguration config = new NeuralNetConfiguration.Builder() .seed(123) // Seed for reproducibility .updater(new Sgd(0.01)) // Stochastic gradient descent with learning rate 0.01 .list() // Output layer: maps 1 input feature to 1 output value .layer(new OutputLayer.Builder(LossFunctions.LossFunction.MSE) .nIn(1) .nOut(1) .activation(Activation.IDENTITY) // Linear activation for regression .build()) .build(); MultiLayerNetwork model = new MultiLayerNetwork(config); model.init();
That's it—no hidden layers needed here, since we're learning a simple linear mapping y = x + 1.
4. Train the Model to Fit Your Data
Now let's train the model using our iterator. We'll run through the data multiple times (epochs) and watch the loss drop:
int numEpochs = 100; for (int epoch = 0; epoch < numEpochs; epoch++) { model.fit(dataIterator); dataIterator.reset(); // Reset iterator to go through data again next epoch // Print loss every 10 epochs to track progress if (epoch % 10 == 0) { System.out.printf("Epoch %d - MSE Loss: %.4f%n", epoch, model.score()); } }
You should see the MSE loss decrease over time and stabilize near 0—this means the model is learning the y = x + 1 pattern.
5. Test the Model's Predictions
Once training is done, let's test with new data to see how well it works:
import org.nd4j.linalg.api.ndarray.INDArray; import org.nd4j.linalg.factory.Nd4j; // Test with input x=10 (expected output 11) INDArray testInput = Nd4j.create(new double[]{10}, new int[]{1, 1}); INDArray prediction = model.output(testInput); System.out.printf("Input: %.0f%n", testInput.getDouble(0)); System.out.printf("Predicted Output: %.4f%n", prediction.getDouble(0)); System.out.println("Expected Output: 11");
If training went well, the predicted value will be extremely close to 11. You can also batch-test multiple inputs to verify the model's consistency.
内容的提问来源于stack exchange,提问作者Niclas Schwalbe

