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

Deeplearning4j模型推理阶段内存占用过高问题求助

Deeplearning4j神经网络推理内存占用过高问题

环境配置详情

  • Deeplearning4j版本:deeplearning4j-core:1.0.0-M2
  • 后端:ND4J CPU后端
  • 操作系统:Mac
  • Java版本:已测试Java 15、Java 21

模型配置

  • 输入节点数:50
  • 隐藏节点数:420
  • 输出节点数:7
  • 网络深度:10层

问题描述

调用output方法时内存占用约16.33GB,远超模型规模应有的消耗,即使减少隐藏节点数量也无法解决该问题。

推理相关代码片段

Object[] rollout(State stato) {
    INDArray p;
    int v;
    if (trained) {
        INDArray inputData = stato.toINDArray();
        INDArray[] out;
        out = model.output(inputData);
        p = out[0];
        v = out[1].getInt(0);
    } else {
        // 模型未训练时的逻辑:随机选择值,给每个子节点分配最大概率
        Random r = new Random();
        p = Nd4j.ones(Board.N).mul(Integer.MAX_VALUE);
        v = Math.abs(r.nextInt()) % MAXREWWARD;
    }
    Object[] out = {p, v};
    return out;
}

神经网络配置代码(复刻AlphaZero)

DeepLearning(int numInputs, int numOutputs, String name, int M, int N, int X, int nndepth) throws IOException {
    this.numInputs = numInputs;
    this.numHiddenNodes = M*N*10;
    this.numOutputs = numOutputs;
    DeepLearning.M = M;
    DeepLearning.N = N;
    DeepLearning.X = X;
    this.name = name;
    File file = new File("./model" + name + M + "." + N + "." + X + "." + ".zip");
    if (file.exists()) {
        trained = true;
        this.model = ComputationGraph.load(file, true);
    } else {
        ComputationGraphConfiguration.GraphBuilder graphBuilder = new NeuralNetConfiguration.Builder()
                .seed(System.currentTimeMillis())
                .weightInit(WeightInit.RELU)
                .l2(1e-4)
                .updater(new Adam(learningRate))
                .graphBuilder();
        graphBuilder.addInputs("input")
                .setInputTypes(InputType.feedForward(numInputs));
        String lastLayer = "input";
        for (int i = 0; i < nndepth; i++) {
            graphBuilder.addLayer("torso_" + i,
                    new DenseLayer.Builder()
                            .nIn(i == 0 ? numInputs : numHiddenNodes)
                            .nOut(numHiddenNodes)
                            .activation(Activation.RELU)
                            .build(),
                    lastLayer);
            lastLayer = "torso_" + i;
        }

        graphBuilder.addLayer("policy_dense",
                new DenseLayer.Builder()
                        .nIn(numHiddenNodes)
                        .nOut(numHiddenNodes)
                        .activation(Activation.RELU)
                        .build(),
                lastLayer);
        graphBuilder.addLayer("policy_output",
                new OutputLayer.Builder(LossFunctions.LossFunction.NEGATIVELOGLIKELIHOOD)
                        .nIn(numHiddenNodes)
                        .nOut(numOutputs)
                        .activation(Activation.SOFTMAX)
                        .build(),
                "policy_dense");

        graphBuilder.addLayer("value_dense",
                new DenseLayer.Builder()
                        .nIn(numHiddenNodes)
                        .nOut(numHiddenNodes)
                        .activation(Activation.RELU)
                        .build(),
                lastLayer);
        graphBuilder.addLayer("value_output",
                new OutputLayer.Builder(LossFunctions.LossFunction.MSE)
                        .nIn(numHiddenNodes)
                        .nOut(1)
                        .activation(Activation.IDENTITY)
                        .build(),
                "value_dense");

        graphBuilder.setOutputs("policy_output", "value_output");

        ComputationGraphConfiguration conf = graphBuilder.build();
        model = new ComputationGraph(conf);
    }
    model.init();
    System.out.println("参数数量: " + model.numParams());
    this.myReplay = new ReplayBuffer();
}

private String addDenseLayer(ComputationGraphConfiguration.GraphBuilder graphBuilder, String inputLayer, String layerName, int nIn, int nOut, Activation activation) {
    graphBuilder.addLayer(layerName, new DenseLayer.Builder()
            .nIn(nIn)
            .nOut(nOut)
            .activation(activation)
            .build(), inputLayer);
    return layerName;
}

已尝试的排查措施(均未解决)

  • 使用工作区管理内存
  • 采用原地操作减少内存分配
  • 降低推理批处理大小
  • 检查内存泄漏与垃圾回收情况

性能分析结果

计算时间主要消耗在ComputationGraph的output调用中,内存使用情况如下:
内存使用情况

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 18:44:56