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

Scala中Deeplearning4J SameDiff:INDArray转DataSetIterator适配fit方法问题

问题:Scala 3.2中SameDiff训练线性回归时的DataSetIterator适配问题

我在Scala 3.2里成功跑通了SameDiff文档里的基础线性回归示例,但想用输入和标签INDArray训练时卡壳了——没法把INDArray转换成fit方法需要的正确DataSetIterator。选SameDiff而不是现成层,是为了后续用它优化EM算法、双曲嵌入这类自定义损失函数。

我的代码

package main
import org.nd4j.autodiff.samediff._
import org.nd4j.linalg.factory.Nd4j
import scala.jdk.CollectionConverters._
import org.nd4j.linalg.api.buffer.DataType
import org.nd4j.weightinit.impl.XavierInitScheme
import javax.xml.crypto.Data
import org.nd4j.autodiff.samediff.TrainingConfig
import org.nd4j.linalg.learning.config.Adam
import org.deeplearning4j.datasets.iterator.INDArrayDataSetIterator
import org.nd4j.common.primitives.Pair

@main def hello: Unit =

        val nIn = 4
        val nOut = 2

        val sd = SameDiff.create()

        //First: Let's create our placeholders. Shape: [minibatch, in/out]
        val input = sd.placeHolder("input", DataType.FLOAT, -1, nIn)
        val labels = sd.placeHolder("labels", DataType.FLOAT, -1, 1)

        //Second: let's create our variables
        val weights = sd.`var`("weights", new XavierInitScheme('c', nIn, nOut), DataType.FLOAT, nIn,nOut)
        val bias = sd.`var`("bias")

        //And define our forward pass:
        val out = input.mmul(weights).add(bias)    //Note: it's broadcast add here

        //And our loss function (done manually here for the purposes of this example):
        val difference = labels.sub(out)
        val sqDiff = sd.math().square(difference)
        val mse = sqDiff.mean("mse")

        //Let's create some mock data for this example:
        val minibatch = 10
        Nd4j.getRandom().setSeed(12345)
        val inputArr = Nd4j.rand(minibatch, nIn)
        val labelArr = Nd4j.rand(minibatch, nOut)
        println(labelArr)

        val placeholderData = Map("input" -> inputArr, "labels" -> labelArr).asJava

        //Execute forward pass:
        val loss = sd.output(placeholderData, "mse").get("mse")
        println("MSE: " + loss)

        //Calculate gradients:
        val gradMap = sd.calculateGradients(placeholderData, "weights", "bias").asScala.toMap
        System.out.println("Weights gradient:")
        System.out.println(gradMap.get("weights"))
        System.out.println("Bias gradient:")
        System.out.println(gradMap.get("bias"))

        val config = TrainingConfig.builder().
            l2(1e-4).
            updater(Adam(1e-3)).
            dataSetFeatureMapping("input").
            dataSetLabelMapping("labels").
            build()

        sd.setTrainingConfig(config)
        val data = INDArrayDataSetIterator(List(Pair(inputArr,labelArr)).asJava,32)
        val hist = sd.fit(data) //Can't find overloaded method

当前SBT依赖

另外我不确定该用哪些SBT依赖,当前配置的依赖如下:

libraryDependencies += "org.nd4j" % "nd4j" % "1.0.0-M2.1"
libraryDependencies += "org.deeplearning4j" % "deeplearning4j-core" % "1.0.0-M2.1"
libraryDependencies += "org.nd4j" % "nd4j-native-platform" % "1.0.0-M2.1"

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 07:05:09