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

