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

DeepLearning4J数据获取与线性回归模型训练问题咨询

解决自定义CSV数据的线性回归:DataSetIterator配置与模型训练指南

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 09:29:01