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

Deeplearning4j报错:org.nd4j.linalg.factory.Nd4j.randomFactory为空

解决Deeplearning4j中Nd4j.randomFactory为空的NullPointerException问题

问题描述

运行以下Deeplearning4j示例代码时抛出空指针异常:

package application.test;

import java.io.File;
import java.io.IOException;
import java.util.ArrayList;
import java.util.List;

import org.deeplearning4j.nn.conf.MultiLayerConfiguration;
import org.deeplearning4j.nn.conf.NeuralNetConfiguration;
import org.deeplearning4j.nn.conf.layers.DenseLayer;
import org.deeplearning4j.nn.conf.layers.OutputLayer;
import org.deeplearning4j.nn.multilayer.MultiLayerNetwork;
import org.deeplearning4j.nn.weights.WeightInit;
import org.deeplearning4j.util.ModelSerializer;
import org.nd4j.jita.conf.Configuration;
import org.nd4j.jita.conf.Configuration.ExecutionModel;
import org.nd4j.linalg.activations.Activation;
import org.nd4j.linalg.api.buffer.DataType;
import org.nd4j.linalg.api.ndarray.INDArray;
import org.nd4j.linalg.dataset.DataSet;
import org.nd4j.linalg.dataset.api.iterator.BaseDatasetIterator;
import org.nd4j.linalg.dataset.api.iterator.DataSetIterator;
import org.nd4j.linalg.dataset.api.iterator.fetcher.DataSetFetcher;
import org.nd4j.linalg.factory.Nd4j;
import org.nd4j.linalg.learning.config.Adam;
import org.nd4j.linalg.lossfunctions.LossFunctions.LossFunction;
import org.nd4j.linalg.factory.RandomFactory;


public class Main {

    public static void main(String[] args) {
        int numInputs = 3;
        int numOutputs = 2;
        int numHiddenNodes = 4;
        String modelFilePath = "C:\\Users\\Léon\\Desktop\\model.zip";

        MultiLayerNetwork model = loadOrCreateModel(modelFilePath, numInputs, numOutputs, numHiddenNodes);

        // 假设你有输入和输出数据
        double[][] inputData = {{0.1, 0.2, 0.3}, {0.4, 0.5, 0.6}}; // 示例数据
        double[][] outputData = {{1, 0}, {0, 1}}; // 分类标签

        int batchSize = 1; // 训练批次大小
        int numEpochs = 10; // 训练轮数

        DataSetIterator trainingData = prepareTrainingData(inputData, outputData, batchSize);

        // 仅当模型未加载时训练并保存
        File file = new File(modelFilePath);
        if (!file.exists()) {
            trainNetwork(model, trainingData, numEpochs);
            saveModel(model, modelFilePath);
        }

        // 此处可使用模型进行预测或评估
    }

    private static MultiLayerNetwork loadOrCreateModel(String modelFilePath, int numInputs, int numOutputs, int numHiddenNodes) {
        File file = new File(modelFilePath);
        if (file.exists()) {
            try {
                return ModelSerializer.restoreMultiLayerNetwork(file);
            } catch (IOException e) {
                e.printStackTrace();
            }
        }
        return createModel(numInputs, numOutputs, numHiddenNodes);
    }

    private static MultiLayerNetwork createModel(int numInputs, int numOutputs, int numHiddenNodes) {
        MultiLayerConfiguration conf =  new NeuralNetConfiguration.Builder()
                .updater(new Adam())
                .weightInit(WeightInit.XAVIER)
                .list().layer(0, new DenseLayer.Builder().nIn(numInputs).nOut(numHiddenNodes).activation(Activation.RELU).build())
                .layer(1, new OutputLayer.Builder(LossFunction.NEGATIVELOGLIKELIHOOD).activation(Activation.SOFTMAX).build())
                .build();

        MultiLayerNetwork model = new MultiLayerNetwork(conf);
        model.init();
        
        return model;
    }

    private static DataSetIterator prepareTrainingData(double[][] inputData, double[][] outputData, int batchSize) {
        INDArray inputNDArray = Nd4j.create(inputData);
        INDArray outputNDArray = Nd4j.create(outputData);

        DataSet dataSet = new DataSet(inputNDArray, outputNDArray);
        List<DataSet> listDataSet = new ArrayList<>();
        listDataSet.add(dataSet);

        return new BaseDatasetIterator(batchSize, listDataSet.size(), (DataSetFetcher) listDataSet.iterator());
    }

    private static void trainNetwork(MultiLayerNetwork model, DataSetIterator trainingData, int numEpochs) {
        for (int i = 0; i < numEpochs; i++) {
            model.fit(trainingData);
        }
    }

    private static void saveModel(MultiLayerNetwork model, String filePath) {
        File file = new File(filePath);
        try {
            ModelSerializer.writeModel(model, file, true);
        } catch (IOException e) {
            e.printStackTrace();
        }
    }
}

抛出异常:

Exception in thread "main" java.lang.NullPointerException: Cannot invoke "org.nd4j.linalg.factory.RandomFactory.getRandom()" because "org.nd4j.linalg.factory.Nd4j.randomFactory" is null

已尝试排查缺失ND4J依赖、降低ND4J版本(当前使用M2.1),均未解决问题。

解决方案

1. 修正ND4J依赖配置

ND4J需根据运行环境选择对应后端依赖,确保添加正确的后端包:

  • CPU环境:添加nd4j-native或nd4j-native-platform(自动适配系统)
  • GPU环境(CUDA):添加nd4j-cuda-XX-platform(XX对应CUDA版本,如11.8)

Maven依赖示例(CPU平台):

<dependency>
    <groupId>org.nd4j</groupId>
    <artifactId>nd4j-native-platform</artifactId>
    <version>2.1.0</version>
</dependency>
<dependency>
    <groupId>org.deeplearning4j</groupId>
    <artifactId>deeplearning4j-core</artifactId>
    <version>2.1.0</version>
</dependency>

2. 修复错误的数据迭代器类型转换

prepareTrainingData方法中,将listDataSet.iterator()强制转换为DataSetFetcher属于错误用法,会导致数据加载异常,间接引发ND4J初始化问题。替换为正确的DataSetIterator实现:

private static DataSetIterator prepareTrainingData(double[][] inputData, double[][] outputData, int batchSize) {
    INDArray inputNDArray = Nd4j.create(inputData);
    INDArray outputNDArray = Nd4j.create(outputData);

    DataSet dataSet = new DataSet(inputNDArray, outputNDArray);
    // 使用ListDataSetIterator替代错误的BaseDatasetIterator用法
    return new ListDataSetIterator<>(Collections.singletonList(dataSet), batchSize);
}

需导入org.deeplearning4j.datasets.iterator.impl.ListDataSetIterator和java.util.Collections。

3. 补充输出层缺失配置

创建OutputLayer时未指定输出神经元数量nOut,会导致模型初始化失败,补充该配置:

.layer(1, new OutputLayer.Builder(LossFunction.NEGATIVELOGLIKELIHOOD)
        .nOut(numOutputs) // 添加输出神经元数量
        .activation(Activation.SOFTMAX)
        .build())

4. 显式初始化ND4J(可选)

在程序入口处添加ND4J初始化代码,确保底层环境加载完成:

public static void main(String[] args) {
    // 显式初始化ND4J
    Nd4j.getRandom();
    // 后续原有代码...
}

验证步骤

  1. 替换上述错误代码片段
  2. 更新Maven依赖并重新构建项目
  3. 删除已存在的model.zip文件(如果有),重新运行程序

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 06:24:58