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

Deeplearning4j v1.0.0-M1.1如何加载已保存模型实现新样本分类预测?

Deeplearning4j 鸢尾花模型预测后续操作步骤

首先你需要补充训练阶段的归一化器保存逻辑,预测时必须使用训练集得到的均值、标准差做归一化,否则预测结果会出现偏差:

1. 训练阶段补充保存归一化器代码

在训练代码的normalizer.fit(trainingData)执行完成后,添加以下代码保存归一化器:

// 保存归一化器到本地
File normalizerFile = new File("C:/Projects/deeplearning4j/iris_normalizer.bin");
NormalizerSerializer.getDefault().writeNormalizer(normalizer, normalizerFile);

2. 加载模型时同步加载归一化器

你加载模型的代码旁,添加加载归一化器的逻辑:

// 加载已训练好的模型
File locationToLoad = new File("C:/Projects/deeplearning4j/trained_iris_model.zip");
MultiLayerNetwork model = ModelSerializer.restoreMultiLayerNetwork(locationToLoad);

// 加载对应训练集的归一化器
File normalizerFile = new File("C:/Projects/deeplearning4j/iris_normalizer.bin");
DataNormalization normalizer = NormalizerSerializer.getDefault().restore(normalizerFile);

3. 处理待预测数据并得到结果

你已经初始化了待预测数据的CSVRecordReader,后续操作代码如下:

// 构造无标签数据集迭代器,batchSize设为待预测样本总数即可
int predictBatchSize = 3; // 按实际待预测样本数修改
DataSetIterator predictIterator = new RecordReaderDataSetIterator(recordReader, predictBatchSize);
DataSet predictData = predictIterator.next();

// 用训练集的归一化规则对待预测数据做归一化,禁止调用fit方法
normalizer.transform(predictData);
INDArray predictFeatures = predictData.getFeatures();

// 方案1:直接获取分类标签(返回int数组,元素为0/1/2,对应三类鸢尾花)
int[] predictedLabels = model.predict(predictFeatures);
for (int i = 0; i < predictedLabels.length; i++) {
    System.out.printf("第%d个样本预测分类:%d%n", i+1, predictedLabels[i]);
}

// 方案2:获取每个类别的预测概率,适用于需要输出置信度的场景
INDArray predictProbs = model.output(predictFeatures);
System.out.println("所有样本对应三个类别的预测概率:");
System.out.println(predictProbs);

注意事项

  • 若待预测数据量较大,可以分批从predictIterator取batch处理,不需要一次性加载全部数据
  • 归一化必须使用训练阶段得到的统计量,不能用待预测数据重新执行normalizer.fit(),否则会出现数据泄露导致结果错误

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 16:57:04