Deeplearning4j案例

wen java案例 2

本文目录导读:

Deeplearning4j案例

  1. 案例一:手写数字识别(CNN 卷积神经网络)
  2. 案例二:情感分析(LSTM 循环神经网络)
  3. 案例三:鸢尾花分类(多层感知机 MLP)—— 最简单入门
  4. 案例四:基于 Spring Boot 的在线推理服务(生产环境集成)
  5. 如何选择与学习建议

这里为你整理了几个 Deeplearning4j (DL4J) 的典型实战案例,从基础的图像分类到复杂的时序预测,并附上核心代码逻辑和思路。

DL4J 是 JVM 生态中非常成熟的深度学习框架,特别适合与 Java 后端(如 Spring Boot)集成,或者处理大规模分布式训练。


手写数字识别(CNN 卷积神经网络)

场景:经典 MNIST 数据集,识别 0-9 的手写数字,这是入门 DL4J 的 Hello World。

核心思路

  1. 加载数据:使用 MnistDataSetIterator 自动下载并分批处理数据。
  2. 构建模型:使用 ConvolutionalNetworkComputationGraph,包含卷积层、池化层、全连接层。
  3. 训练:设置迭代次数、优化器。
  4. 评估:输出准确率。

核心代码逻辑

// 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 和序列数据处理。

核心思路

  1. 数据预处理:将文本转化为词向量(Word2Vec 或 预训练的 GloVe)。
  2. 嵌入层:使用 EmbeddingSequenceLayer 将词索引映射为向量。
  3. LSTM 层:处理序列数据,捕捉上下文信息。
  4. 输出层:二分类(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 在企业应用中的最大亮点。

核心思路

  1. 在 Java 后端加载序列化好的模型文件(.zip 格式)。
  2. 将对传入的原始数据进行预处理(归一化、缩放等)。
  3. 创建 INDArray,调用 model.output()
  4. 将结果打包为 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 的命名实体识别、目标检测),可以告诉我具体方向,我再为你补充更详细的案例代码。

抱歉,评论功能暂时关闭!