本文目录导读:

- 案例一:手写数字识别(CNN 卷积神经网络)
- 案例二:情感分析(LSTM 循环神经网络)
- 案例三:鸢尾花分类(多层感知机 MLP)—— 最简单入门
- 案例四:基于 Spring Boot 的在线推理服务(生产环境集成)
- 如何选择与学习建议
这里为你整理了几个 Deeplearning4j (DL4J) 的典型实战案例,从基础的图像分类到复杂的时序预测,并附上核心代码逻辑和思路。
DL4J 是 JVM 生态中非常成熟的深度学习框架,特别适合与 Java 后端(如 Spring Boot)集成,或者处理大规模分布式训练。
手写数字识别(CNN 卷积神经网络)
场景:经典 MNIST 数据集,识别 0-9 的手写数字,这是入门 DL4J 的 Hello World。
核心思路:
- 加载数据:使用
MnistDataSetIterator自动下载并分批处理数据。 - 构建模型:使用
ConvolutionalNetwork或ComputationGraph,包含卷积层、池化层、全连接层。 - 训练:设置迭代次数、优化器。
- 评估:输出准确率。
核心代码逻辑:
// 1. 数据加载(分为训练集和测试集)
int batchSize = 64;
DataSetIterator mnistTrain = new MnistDataSetIterator(batchSize, true, 12345);
DataSetIterator mnistTest = new MnistDataSetIterator(batchSize, false, 12345);
// 2. 构建神经网络配置
MultiLayerConfiguration conf = new NeuralNetConfiguration.Builder()
.seed(123)
.optimizationAlgo(OptimizationAlgorithm.STOCHASTIC_GRADIENT_DESCENT)
.updater(new Adam(0.001))
.list()
// 输入层:28x28 单通道
.layer(0, new ConvolutionLayer.Builder(5, 5)
.nIn(1) // 单通道
.nOut(20) // 20个滤波器
.activation(Activation.RELU)
.build())
.layer(1, new SubsamplingLayer.Builder(SubsamplingLayer.PoolingType.MAX)
.kernelSize(2, 2)
.build())
.layer(2, new DenseLayer.Builder()
.activation(Activation.RELU)
.nOut(500)
.build())
// 输出层:10个类别
.layer(3, new OutputLayer.Builder(LossFunctions.LossFunction.NEGATIVELOGLIKELIHOOD)
.nOut(10)
.activation(Activation.SOFTMAX)
.build())
.setInputType(InputType.convolutionalFlat(28, 28, 1)) // 关键:指定输入形状
.build();
MultiLayerNetwork model = new MultiLayerNetwork(conf);
model.init();
// 3. 训练(迭代 1 个 epoch)
for (int i = 0; i < 1; i++) {
model.fit(mnistTrain);
}
// 4. 评估
Evaluation eval = model.evaluate(mnistTest);
System.out.println("Accuracy: " + eval.accuracy());
情感分析(LSTM 循环神经网络)
场景:利用 IMDB 电影评论数据,判断评论是正面还是负面,这涉及到 NLP 和序列数据处理。
核心思路:
- 数据预处理:将文本转化为词向量(Word2Vec 或 预训练的 GloVe)。
- 嵌入层:使用
EmbeddingSequenceLayer将词索引映射为向量。 - LSTM 层:处理序列数据,捕捉上下文信息。
- 输出层:二分类(Sigmoid)。
核心代码逻辑:
// 假设已经将评论转化为词索引序列,且词向量大小为 50(通过 WordVectorSerializer 加载)
WordVectors wordVectors = WordVectorSerializer
.loadGoogleModel(new File("path/to/GoogleNews-vectors-negative300.bin"), true);
// 构建数据集迭代器(TextDataSetIterator 或自定义)
// 核心是生成 INDArray 序列
MultiLayerConfiguration conf = new NeuralNetConfiguration.Builder()
.updater(new Adam(0.002))
.list()
// 嵌入层:输入词汇表大小,输出向量维度
.layer(0, new EmbeddingSequenceLayer.Builder()
.nIn(30000) // 词汇表大小
.nOut(100) // 嵌入维度
.build())
// LSTM 层:隐藏单元数 128
.layer(1, new GravesLSTM.Builder()
.nIn(100) // 来自嵌入层
.nOut(128)
.activation(Activation.TANH)
.build())
// RNN 输出层(二分类)
.layer(2, new RnnOutputLayer.Builder(LossFunctions.LossFunction.MCXENT)
.nIn(128)
.nOut(2) // 正/负
.activation(Activation.SOFTMAX)
.build())
.build();
// 训练时,label 是 [batchSize, 2, sequenceLength] 的 one-hot 向量
鸢尾花分类(多层感知机 MLP)—— 最简单入门
场景:经典的 4 个特征(花瓣长度等),预测 3 种鸢尾花类型,适合理解 DL4J 数据流。
核心思路:
读取标准 CSV 文件,构建 RecordReaderDataSetIterator,喂给全连接网络。
核心代码逻辑:
// 1. 读取 CSV
RecordReader rr = new CSVRecordReader(0, ","); // 跳过头行
rr.initialize(new FileSplit(new File("iris.csv")));
DataSetIterator iterator = new RecordReaderDataSetIterator.Builder(rr, batchSize)
.classification(4, 3) // 第4列是标签,共3类
.build();
// 2. 构建模型(全连接层)
MultiLayerConfiguration conf = new NeuralNetConfiguration.Builder()
.seed(42)
.activation(Activation.RELU)
.weightInit(WeightInit.XAVIER)
.updater(new Adam(0.01))
.list()
.layer(0, new DenseLayer.Builder().nIn(4).nOut(10).build()) // 输入4维特征
.layer(1, new DenseLayer.Builder().nIn(10).nOut(10).build())
.layer(2, new OutputLayer.Builder(LossFunctions.LossFunction.NEGATIVELOGLIKELIHOOD)
.nIn(10).nOut(3)
.activation(Activation.SOFTMAX)
.build())
.build();
MultiLayerNetwork model = new MultiLayerNetwork(conf);
model.init();
model.fit(iterator);
基于 Spring Boot 的在线推理服务(生产环境集成)
场景:将训练好的 DL4J 模型部署为 REST API,接收图片/文本,返回预测结果,这是 DL4J 在企业应用中的最大亮点。
核心思路:
- 在 Java 后端加载序列化好的模型文件(
.zip格式)。 - 将对传入的原始数据进行预处理(归一化、缩放等)。
- 创建
INDArray,调用model.output()。 - 将结果打包为 JSON 返回给前端。
核心代码逻辑(Service 层):
@Service
public class PredictionService {
private MultiLayerNetwork model;
@PostConstruct
public void init() throws IOException {
// 加载模型(每次启动时加载一次)
model = ModelSerializer.restoreMultiLayerNetwork(
new File("/models/my_model.zip"));
}
public double[] predict(float[] features) {
// 1. 将特征转为 INDArray(1行 N列)
INDArray inputArray = Nd4j.create(features).reshape(1, features.length);
// 2. 执行推理
INDArray outputArray = model.output(inputArray, false);
// 3. 返回概率数组
return outputArray.toDoubleVector();
}
}
如何选择与学习建议
| 模型类型 | 适用场景 | 使用的 DL4J API |
|---|---|---|
| DenseLayer | 结构化数据(表格类) | 逻辑回归、多分类 |
| ConvolutionLayer | 图像、视频识别 | 人脸识别、OCR、医学影像 |
| LSTM / GRU | 时间序列、自然语言 | 股票预测、机器翻译、情感分析 |
| ComputationGraph | 多输入/多输出 | 多任务学习、注意力机制 |
补充建议:
- 数据加载:DL4J 的数据管道的核心是
DataSetIterator,理解RecordReader(CSV、JSON)和ImageRecordReader是必备技能。 - 性能:如果使用 CPU 训练,建议开启
NativeBlas并配置合理的线程数;如果是生产环境,建议使用 ONNX 转换或直接使用 JavaCPP 加速。 - 调试:训练时关注
ModelListener(如ScoreIterationListener)打印的损失值,观察是否收敛。
如果你有特定的需求(NLP 的命名实体识别、目标检测),可以告诉我具体方向,我再为你补充更详细的案例代码。