☰
DJL与Spring集成:Java后端部署深度学习模型的实践指南
2026/10/12 6:01:01 网站建设 项目流程

简介:面向 Java 及 JVM 平台开发者,这份压缩包聚焦 DJL 与 Spring 框架的整合,提供从环境配置、数据预处理到模型训练、评估与推理的完整实践资料,着力解决在 Java 生态中开发深度学习应用时的入门与集成难题。资源共 66 个文件,以 29 个 java 源码文件为主,辅以 12 张 jpg 与 12 张 png 图示说明、6 个 xml 配置、3 个 params 模型参数及 properties、json 等工程文件,便于对照代码理解实现细节。已有 499 人浏览学习。资料覆盖 DJL 核心组件、数据集批量处理、损失函数与优化器选择、模型保存加载,以及与 Spring Boot/RESTful API 集成等关键环节;包内包含按模块整理的代码示例、笔记文档和小规模实战项目,从环境搭建到故障排查均有涉及。无论初学者还是需要迁移技术栈的开发者,都能获得清晰、可操作的参考路径。

1. DJL 与 Spring 集成:Java 后端跑深度学习的低成本路径

标题里的 Fast Deep Java Library,实际指的就是 DJL(Deep Java Library)。这个框架的价值在于:Java 团队不用切到 Python,就能在 JVM 里完成深度学习的训练与推理,而把它和 Spring 工程整合,是最常见也最实用的落地方式——模型加载、推理接口、训练任务都可以交给 Spring 管理,服务还是原来那个 Spring Boot 服务,不需要额外维护 Python 运行时。适合谁读?手里已经有一套 Spring 服务、正在为 AI 功能纠结“要不要引入 Python 技术栈”的开发者,或者已经决定用 DJL、但不确定该怎么把模型生命周期和 Spring 的 Bean 生命周期对齐的人。下面从核心抽象讲起,一路到可复现的依赖、代码、训练闭环,最后列出我在实际整合中踩过的坑。

2. 先理解 DJL 的三个核心抽象,再定 Spring 里的 Bean 边界

2.1 Model、Predictor 与 Criteria:加载与推理的分工

DJL 的抽象层非常克制,跟 Spring 整合之前,只需要先弄清楚三个对象之间的边界:Model是模型的生命周期载体,持有网络结构和参数;Criteria是描述“我要什么样的模型”的工厂参数;Predictor是真正执行推理的会话对象。这三者的线程模型和生命周期都不同,直接决定你在 Spring 里怎么配置 Bean。

Model在 DJL 里是一个 AutoCloseable,加载后可以长期持有、被并发读取,适合做成 Spring 单例 Bean;Predictor不是线程安全的,但创建成本很低,正确用法是用完即关。最典型的加载代码是这样的:

Criteria<Image, Classifications> criteria = Criteria.builder() // 输入输出类型,影响后续 Translator 的选择 .setTypes(Image.class, Classifications.class) // 从 Model Zoo 里加载指定的模型 .optApplication(Application.CV.IMAGE_CLASSIFICATION) .optArtifactId("resnet18") .build(); Model model = criteria.loadModel();

这段代码的关键在于optArtifactId:它指定了模型 Zoo 里的真实模型标识,第一次加载会把网络结构和权重拉到本地缓存目录。如果你的网络环境受限,这个下载经常会让加载卡在进度条阶段,我一般会在加载前先手动把模型文件准备好,再用optModelPath指向本地路径。setTypes决定了后续Predictor.predict方法的签名:声明输入为Image、输出为Classifications,代码就是强类型的,写错类型编译期就报错。

Predictor 的使用方式,绝大多数人第一次都会写错。正确做法是每次推理都立刻创建一个、用完立刻关闭:

try (Predictor<Image, Classifications> predictor = model.newPredictor()) { Classifications result = predictor.predict(image); return result; }

newPredictor的内部逻辑只是组装 Translator 和推理上下文,并不加载模型,所以这个创建动作开销很低,没必要做池化。真正重的是model的加载动作,它要做参数反序列化和网络初始化,这个动作只应该在 Spring 启动时做一次。

2.2 选型理由:为什么 Java 团队选 DJL 而不是自建 Python 推理服务

很多人第一反应是:AI 能力本来就是 Python 的生态,为什么要在 Java 里硬做?我的理由很简单:如果你的团队全是 Java 开发、公司基础设施也围绕 JVM 建设,那么单独为 AI 功能维护一个 Python 微服务、网络通信、模型部署流水线,成本远高于在同一个进程里直接调 DJL。

两者对比下来:

维度DJL 同进程集成独立 Python 推理服务
部署产物一个 Spring Boot jarPython 环境 + Web 服务 + 接口协议
推理延迟无网络开销,数据直接进引擎多一跳 HTTP/RPC
类型安全输入输出强类型需要定义接口协议、异常处理
运维负担跟随 Spring 生命周期需要独立监控、日志、升级流程
模型生态需要转 TorchScript / ONNX原生 Python 模型直接跑

选 DJL 的真实痛点也很明确:训练生态比 Python 弱不少,很多新论文的模型根本没有 DJL 实现。所以我的选型建议是——训练用 DJL 处理中小规模模型没问题,超大规模模型我仍然建议训练在 Python 侧完成,然后把模型导出成 TorchScript 或 ONNX,部署推理交给 DJL 和 Spring。这也是 DJL 定位里最成熟的路径,两边都不耽误。

3. 在 Spring Boot 里把模型加载和推理接起来:依赖、Bean 与接口

3.1 Maven 依赖怎么加:engine、native 与 model-zoo

DJL 的依赖结构比普通 Java 库要复杂一点,核心原因是它有引擎层和 native 层。api是通用抽象,pytorch-engine是引擎实现,而 native 二进制通过pytorch-native-auto按平台分类器引入。如果漏了最后这个分类器依赖,工程能编译过,但一运行就报言「找不到 native 库」。

我一般会在 pom 里这样组织,方便统一管理版本:

<properties> <djl.version>0.21.0</djl.version> <!-- 按部署机器的 OS 和架构改 --> <jni.classifier>linux-x86_64</jni.classifier> </properties> <dependencies> <dependency> <groupId>ai.djl</groupId> <artifactId>api</artifactId> <version>${djl.version}</version> </dependency> <dependency> <groupId>ai.djl</groupId> <artifactId>model-zoo</artifactId> <version>${djl.version}</version> </dependency> <dependency> <groupId>ai.djl.pytorch</groupId> <artifactId>pytorch-engine</artifactId> <version>${djl.version}</version> </dependency> <dependency> <groupId>ai.djl.pytorch</groupId> <artifactId>pytorch-native-auto</artifactId> <version>${djl.version}</version> <classifier>${jni.classifier}</classifier> </dependency> </dependencies>

jni.classifier是最容易踩坑的地方:本地开发往往在 Mac 上(osx-aarch64),测试环境是 Linux(linux-x86_64),打包时如果单台机器上只配了一个 classifier,部署到另一平台就会启动失败。我会把它从 pom 里抽出来,放到构建配置里按环境替换,或者在部署文档里明确标注这台机器用什么架构。

需要注意,model-zoo并不是必须项。如果完全用自己的模型,只引入api和pytorch-engine就够了,model-zoo只是为了加载官方预训练模型时能找到模型清单。版本上,DJL 的 API 在 0.20 到 0.22 之间有一些函数签名调整,跨大版本升级时不要直接改版本号就完事,要跑一遍测试用例。

3.2 用 Criteria 加载模型并注册为 Spring Bean

在 Spring 里管理 DJL 模型,核心思路是把“加载一次、长期持有”的模型声明成@Bean,销毁时交给 Spring 调用close()。这样模型生命周期和 Spring 容器完全对齐,服务重启时模型只会被加载一次,不会因为每次请求都重新初始化而把接口拖慢。

一个可以直接抄的配置类:

@Configuration public class DjlModelConfig { @Bean(destroyMethod = "close") public Model imageClassificationModel() throws ModelException, IOException { Criteria<Image, Classifications> criteria = Criteria.builder() .setTypes(Image.class, Classifications.class) .optApplication(Application.CV.IMAGE_CLASSIFICATION) .optArtifactId("resnet18") // 明确使用 CPU,避免开发机没有 CUDA 时启动报错 .optDevice(Device.cpu()) .build(); return criteria.loadModel(); } }

这里要重点看destroyMethod = "close"这个属性。Spring 容器关闭时,它会自动调用Model.close(),释放原生内存和线程资源。如果你漏掉这个配置,开发环境可能没什么感觉,生产环境反复发布后就会出现堆外内存持续增长的问题——因为旧模型对象没有被正确回收其 native 资源。

optDevice(Device.cpu())是我被坑过之后习惯性加上的:开发机上没有 CUDA 的话,不指定设备会默认尝试 GPU,报错信息不够直观。线上有 GPU 时,再把这里改成Device.gpu()。

3.3 Controller 里做推理:Predictor 每次新建,不要在请求间共享

模型是单例,但 Predictor 绝对不能用单例。原因在于 Predictor 内部持有输入数据的中间状态,多个线程同时对它调用predict会互相污染结果。我的习惯是把它封装在一个 Service 里,每次调用都新建,让 Controller 层不感知这些细节:

@Service public class ClassifyService { private final Model model; public ClassifyService(Model model) { this.model = model; } public Classifications classify(Image image) throws TranslateException { try (Predictor<Image, Classifications> predictor = model.newPredictor()) { return predictor.predict(image); } } }

Controller 只负责解析请求和转换输入格式:

@RestController public class ClassifyController { private final ClassifyService classifyService; public ClassifyController(ClassifyService classifyService) { this.classifyService = classifyService; } @PostMapping("/classify") public Map<String, Double> classify(@RequestParam("file") MultipartFile file) throws IOException, TranslateException { Image image = ImageFactory.getInstance().fromInputStream(file.getInputStream()); Classifications result = classifyService.classify(image); // 取 Top-3,方便前端展示置信度 return result.topK(3).stream() .collect(Collectors.toMap( entry -> entry.getKey(), entry -> entry.getValue())); } }

topK(3)拿到的列表顺序是从高到低,直接用Collectors.toMap转成 JSON 输出。这里注意,Classifications的toString已经很好看了,但接口返回时我还是建议只输出 Top-N,避免把一个几千类的全量概率都吐给前端。

ImageFactory.getInstance().fromInputStream会读取完整图片字节流,如果图片很大,这个操作本身就有一定耗时。对于高并发的图片上传接口,我会在接入层加一个文件大小上限校验,防止大图把线程池堵死。

4. 用 Spring 管训练流程:数据集、训练循环与模型回存

4.1 准备数据集:ImageFolder 与目录结构约定

DJL 的ImageFolder数据集约定非常直观:一类一个子目录,子目录名就是类别名。结构像/data/train/cat/xxx.jpg、/data/train/dog/xxx.jpg这样,调用prepare()之后它会自动扫描子目录并建立类别索引。这个设计让训练数据集的组织和 Spring 工程的资源目录一样有章可循。

构建训练数据集的代码:

Dataset trainingDataset = ImageFolder.builder() // 指向训练数据根目录 .setRepositoryPath(Paths.get("/data/images/train")) // 只扫描一层子目录 .optMaxDepth(1) // 随机裁剪缩放,兼顾数据增强 .addTransform(new RandomResizedCrop(112, 112)) // 转成 NDArray,归一化也在这里做 .addTransform(new ToTensor()) .build(); trainingDataset.prepare();

optMaxDepth(1)很关键,它限制采样时只往下找一层子目录,否则会把类别目录里的子目录也当成类别。RandomResizedCrop是训练集专用的增广方式,验证集不应该用它,验证集一般只做Resize和ToTensor。所以训练和验证要分别构建两个ImageFolder实例。

prepare()会在第一次调用时扫描整个目录树并建立索引,数据量大时耗时明显。这个动作不适合放在训练循环里反复执行,我在 Spring 工程里会把它放到训练任务的初始化阶段,只在任务启动时调用一次。

4.2 构建网络与训练配置:用 Trainer 跑训练循环

DJL 的 Block API 长得很像 PyTorch 和 Keras 的混合体。一个简单的分类网络可以这样搭:

Block block = new SequentialBlock() // 把 112x112x3 展平成一维向量 .add(Blocks.batchFlattenBlock(112 * 112 * 3)) .add(Linear.builder().setUnits(64).build()) .add(Activation::relu) .add(Linear.builder().setUnits(10).build());

这是一个全连接网络,适合快速验证流程是否通。真实任务里一般会换成卷积模块或直接用 DJL Model Zoo 里的残差网络做迁移学习,但训练代码结构完全一样。注意最后输出 10 个单元,对应 10 个类别,这个数字要和数据集的类别数一致。

训练配置用DefaultTrainingConfig,把损失函数、优化器、验证数据集挂进去:

DefaultTrainingConfig config = new DefaultTrainingConfig(new SoftmaxCrossEntropyLoss()) .optOptimizer(Optimizer.adam().optLearningRate(0.001f).build()) .optDevices(Device.cpu()) .optValidateDataset(validationDataset); // 每次训练都新建模型实例,避免重复加载残留 try (Model model = Model.newInstance("image-classifier")) { model.setBlock(block); try (Trainer trainer = model.newTrainer(config)) { // 初始化权重,需要明确输入形状 trainer.initialize(new Shape(1, 3, 112, 112)); // fit 内部会按 epoch 循环数据,并在结束时调用 validate EasyTrain.fit(trainer, 10, trainingDataset); model.save(Paths.get("build/model"), "image-classifier"); } }

trainer.initialize(new Shape(1, 3, 112, 112))的1是 batch 维度,DJL 在 initialize 时只看形状不看具体数据。EasyTrain.fit这个封装会替你做 epoch 循环和 batch 迭代,如果配置里加了optValidateDataset,它会在每个 epoch 结束时用验证集算一次准确率,打印到日志里。

训练超参数就两个方向:学习率太大训练震荡,太小收敛慢。我习惯先把optLearningRate放到 0.001 跑一个 epoch 看 loss 曲线,不收敛再往下调一个量级。epoch 数不是越大越好,后面验证集准确率不再上升时,就该考虑提前停止。

4.3 模型保存与加载闭环:训练产物回到 Spring Bean

model.save(Paths.get("build/model"), "image-classifier")会把网络结构和参数文件写到一个目录下。这个目录就是训练和推理之间的交接物,我一般会把目录路径写进 Spring 的配置文件,让加载模型的地方和训练产物的落点保持一致。

从保存目录加载到一个新的 Spring Bean,只需要把 Criteria 的加载来源从 Model Zoo 改成模型路径:

@Bean(destroyMethod = "close") public Model trainedModel() throws ModelException, IOException { Criteria<Image, Classifications> criteria = Criteria.builder() .setTypes(Image.class, Classifications.class) .optModelPath(Paths.get("build/model/image-classifier")) .optTranslator(ImageTranslator.builder() .optResize(112, 112) .build()) .build(); return criteria.loadModel(); }

这里有一个容器上下文的问题:训练任务通常不会放在 Web 请求里跑,因为训练耗时长、占用资源高,一个请求把它启动起来,接口会直接超时。我的做法是在 Spring Boot 启动阶段用一个独立的组件来触发训练,训练完成后,推理用的 Bean 再被注入到 Controller。整个链路是:

  • 训练组件读配置里的数据路径和超参数 → 训练并保存模型
  • 训练完成后,保存路径变成模型加载 Bean 的输入
  • 推理接口运行期间不关心模型是哪里来的,只管调predict

这个闭环的好处是,模型更新只需要替换训练产物目录下的文件,再重启服务,训练和部署就完成了交接。

5. 集成中的 5 个常见问题与排查:从依赖冲突到内存泄漏

5.1 启动报 UnsatisfiedLinkError,native 库加载失败

现象:Spring Boot 启动时一直正常,加载模型时突然抛UnsatisfiedLinkError,提示找不到 jni 相关的符号。

原因:pytorch-native-auto的 classifier 和实际操作系统不匹配。最常见的是开发机是 Mac ARM,打包的 classifier 是linux-x86_64,部署到 Linux ARM 机器上直接失败。另一个原因是多引擎同时引入——既加了 PyTorch 又加了 ONNX Runtime 的 native,两个引擎的 JNI 符号互相干扰。

解决:先确认部署环境架构,再回看 pom 里的jni.classifier。我把这个值拆到环境变量里之后,再没因为架构问题翻车过。多引擎场景下,如果必须共存,就给不同的引擎设置独立的 ClassLoader 隔离,或者干脆只保留一个引擎。

5.2 并发一上来,推理结果变成乱序或直接报错

现象:单线程测试一切正常,用压测工具跑 20 个并发线程,结果出现IllegalStateException,或者返回的类别跟输入图片对不上。

原因:典型的 Predictor 共享。某个同学把Predictor当成 Bean 注入到了 Service 里,所有线程共用一个实例,内部状态被互相覆盖。

解决:严格执行“Predictor 每次新建、用完即 close”。如果用ThreadLocal复用 Predictor,一定要记得在线程池任务结束后清理,否则线程长期存活时 Predictor 的 native 资源无法及时释放。最简单的排查方法:在所有用到predict的地方搜一下,除了newPredictor(),有没有其他地方持有了 Predictor 对象的引用。

5.3 本地 IDE 跑得好好的,打 jar 部署后模型找不到

现象:java -jar启动后,加载模型时抛FileNotFoundException,但同一个构建产物在 IDE 里运行完全正常。

原因:Spring Boot 的 fat jar 把模型文件压缩在 jar 包里,Paths.get("build/model")是文件系统路径,根本访问不到 jar 内部资源。

解决:模型不打进 jar,而是放到独立目录,通过配置项指定绝对路径。如果必须打包进 jar,加载前先用工具类把资源复制到临时目录,再让 DJL 读取临时目录文件。经验是:模型文件几百 MB,打进 jar 会让启动时解压很久,独立目录对后续模型更新也更友好。

5.4 训练跑完内存不见回落,多次训练后直接 OOM

现象:训练任务循环跑完,应用还在运行,但 JVM 堆外内存持续上升,几次训练后容器直接 OOM。

原因:DJL 的 NDArray 和 Batch 对象都在堆外分配原生内存,它们不归 JVM GC 管。训练循环里如果某个Batch没有显式关闭,每一轮都会泄漏一部分原生内存。

解决:训练循环里对Batch使用 try-with-resources,或者迭代完后手动调用batch.close()。我用EasyTrain.fit时也会在训练方法外层套 try-with-resources,确保 Trainer、Model 的 close 一定执行。排查时可以开启 DJL 的 NDArray 泄漏检测日志,它会告诉你哪条创建路径没有关闭。

5.5 自训练模型部署后,分类结果是一串数字而不是类别名

现象:用ImageFolder训练、保存、再加载的模型,推理返回的 Top-3 是像0: 0.982这样的数字标签。

原因:DJL 在保存自定义模型时不会自动把类别名表写进模型文件。模型推理时只输出类别索引,ImageTranslator找不到 synset,就直接显示索引号。

解决:为模型配置一个自定义 Translator,在toClassifications时用你训练时那份类别列表映射索引到类名。

public class CustomImageTranslator extends ImageTranslator { private final List<String> synset; public CustomImageTranslator(List<String> synset) { super(ImageTranslator.builder().build()); this.synset = synset; } @Override public Classifications toClassifications(NDArray array) { NDArray probabilities = array.softmax(0); List<String> classNames = new ArrayList<>(); for (int i = 0; i < synset.size(); i++) { classNames.add(synset.get(i)); } return new Classifications(classNames, probabilities); } }

这段代码的关键是把训练时的classNames顺序保存下来,作为配置项传给 Translator。softmax(0)是把原始输出转成概率分布,如果训练损失函数里已经带了 softmax,这里要避免二次计算,直接取原始数组。这个坑很容易被忽视,因为模型 Zoo 里官方模型都自带 synset,自训练模型不会自动带。

6. 上线前该做的验证:预热、并发压测与线程参数

模型加载成功后,第一次推理通常会比后续慢一个量级,因为引擎懒加载算子和内存池。直接在线上让第一个用户承担这个延迟,体验很差。我的做法是在 Spring 启动完成后主动做一次预热推理,用一个固定的小图触发全部初始化路径:

@Component public class ModelWarmupRunner implements ApplicationRunner { private final ClassifyService classifyService; public ModelWarmupRunner(ClassifyService classifyService) { this.classifyService = classifyService; } @Override public void run(ApplicationArguments args) throws Exception { Image placeholder = ImageFactory.getInstance() .fromUrl("https://resources.djl.ai/images/0.png"); classifyService.classify(placeholder); } }

预热代码里随便传一张图就行,目标是把Predictor.newPredictor、引擎初始化和内存分配全部触发一次。如果团队不方便访问外网图片,就用代码生成一张纯色的小图,效果完全相同。

压测时格外注意 CPU 线程参数。DJL 的 PyTorch 引擎默认会占用机器的所有计算资源,如果 Spring 服务同时还处理其他业务请求,两者会抢 CPU,导致接口响应时间抖动。我的习惯是通过配置把引擎的计算线程数显式设置成 CPU 核数的一半,给业务线程池留出余地。这个参数不调对,压测曲线会出现神奇的毛刺,看起来像玄学,其实是底层算子把线程吃满了。

并发验证也有一个小技巧:先用 20 线程跑 100 次请求,观察结果是否稳定、响应时间是否平滑,再逐步加大到 50、100 线程。如果并发上去之后出现超时增长但不是错误,说明每线程新建 Predictor 的逻辑是合理的;如果直接报错,回去查 Predictor 或者 native 库的线程模型。我习惯把每次推理耗时记录到日志里,至少保留 p95 和 p99 两个指标,DJL 推理很多时候受制于 CPU 调度而不是模型计算,没有指标就看不出来。

最后说一个我自己的教训:第一次做 DJL 集成时,我把 Predictor 当成普通 Service 顺手注册成了@Bean,压测一上并发立刻翻车,排查了大半天。从那以后我养成了一个习惯——每个 DJL 对象都用 try-with-resources 管起来,并且在代码注释里写明生命周期。DJL 这种东西,显式管理资源会给你省下很多半夜排查的时间。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询