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

DL4J实现线性回归遇问题:拟合y=2x+300求指导

帮你排查DL4J线性拟合的问题

嘿,作为DL4J新手,你选了个超棒的入门案例,但代码里有几个关键细节没处理对,导致结果不达预期。咱们一个个拆解:

1. 输入维度与偏置的重复设置

你设置了numInputs = 2,还在训练数据里手动加了1作为输入特征——这其实是画蛇添足!DL4J的OutputLayer默认会自动添加偏置(bias)参数,你手动输入的1给网络多了一个冗余输入,让网络学习的是:

y_scaled = w1*(x/scale) + w2*1 + b

完全偏离了咱们预期的线性关系。正确做法是:

  • 把numInputs改成1
  • 训练数据只传缩放后的x值,去掉手动加的1:
    new DataSet(
        Nd4j.create(new double[]{x / scale}),
        Nd4j.create(new double[]{y / scale})
    )
    
  • 网络配置里的输出层nIn对应改成1:
    .layer(0, new OutputLayer.Builder(LossFunctions.LossFunction.MSE)
        .activation(Activation.IDENTITY)
        .nIn(1).nOut(numOutputs).build())
    

2. 训练轮数严重不足

你设置的nEpochs = 1,意味着网络只过了一遍训练数据就结束了。哪怕是最简单的线性拟合,也需要多轮迭代让参数(权重、偏置)收敛到正确值。建议把nEpochs改成50~100,比如:

public final int nEpochs = 50;

3. 其他小细节优化

  • iterations参数在DL4J新版里已经过时,直接删掉就行,不影响训练逻辑;
  • 数据缩放的逻辑没问题,把x缩到[-0.2,0.2]、y缩到[0.56,1],反而有助于SGD更快收敛,不用改动。

修正后的核心代码片段

参数调整

public final int numInputs = 1;
public final int nEpochs = 50;

训练数据生成

public DataSetIterator generateTrainingData() {
    List<DataSet> list = new ArrayList<>();
    for (int i = 0; i < batchSize; i++) {
        double x = rng.nextDouble() * maxX * (rng.nextBoolean() ? 1 : -1);
        double y = y(x);
        list.add(
            new DataSet(
                Nd4j.create(new double[]{x / scale}),
                Nd4j.create(new double[]{y / scale})
            )
        );
    }
    return new ListDataSetIterator(list, batchSize);
}

网络配置

public MultiLayerConfiguration createConf() {
    return new NeuralNetConfiguration.Builder()
        .seed(seed)
        .optimizationAlgo(OptimizationAlgorithm.STOCHASTIC_GRADIENT_DESCENT)
        .learningRate(learningRate)
        .weightInit(WeightInit.XAVIER)
        .updater(new Nesterovs(0.9))
        .list()
        .layer(0, new OutputLayer.Builder(LossFunctions.LossFunction.MSE)
            .activation(Activation.IDENTITY)
            .nIn(numInputs).nOut(numOutputs).build())
        .pretrain(false).backprop(true).build();
}

改完这些再运行测试,你应该就能看到预测值和真实值几乎重合啦!

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 07:37:55