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
相关产品推荐
相关产品推荐

