ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

纯 Java 实现 PP-OCRv6 推理引擎,告别 ONNX Runtime 与 JNI

纯 Java 实现 PP-OCRv6 推理引擎,告别 ONNX Runtime 与 JNI 前阵子接了个票据识别的活服务端清一色 Java部署环境又偏封闭不能随便装本地库。我一开始也想过 ONNX Runtime毕竟模型转 ONNX 后调用方便但 ONNX Runtime 的 Java 绑定本质还是 JNI得带 native so/dll版本、系统架构、glibc 都得对齐。更麻烦的是PP-OCRv6 的检测和识别模型如果只靠现成 runtime后处理还得自己写C 侧不好改Java 侧又隔了一层。于是我一横心用纯 Java 写了一个 PP-OCRv6 推理引擎不依赖 ONNX Runtime不写一行 JNI模型权重导出成自定义格式Java 里自己做张量、卷积、BN、激活、CTC 解码和 DB 后处理。整套东西跑在普通 JVM 上Spring Boot 直接引入 jar 就能用。适合 Java 后端、需要内网私有化部署 OCR 的同学也适合想搞明白推理引擎底层到底在算什么的人。1. 整体设计与思路拆解1.1 为什么放弃 ONNX Runtime 和 JNI 这条常规路ONNX Runtime 确实是模型部署的捷径Java 侧也有现成 API但它的底层是 CJava 调用时走 JNI。JNI 本身不是问题问题在于部署环境。Windows 要 dllLinux 要 somacOS 要 dylibDocker 镜像里还得注意 Alpine 的 musl 和 glibc 差异。有一次我在一个老 CentOS 上跑ONNX Runtime 的 so 依赖 GLIBC_2.27系统只有 2.17升级系统不现实换 runtime 版本又遇到算子不支持。那一刻我就明白只要 native 依赖存在部署就多一个不可控变量。纯 Java 推理引擎的最大好处是零 native 依赖。一个 jar 包丢过去JRE 能跑它就能跑。调试也简单堆栈全是 Java 代码哪一层算错了、哪个数组越界IDE 直接断点。性能确实不如 ONNX Runtime 的 C 实现但我的场景是内网票据识别并发不高QPS 个位数准确率和可维护性比极限性能更重要。而且 PP-OCRv6 的模型结构相对固定不需要通用图执行器只实现必要的算子就行代码量可控。1.2 纯 Java 推理引擎的边界只服务 PP-OCRv6 的检测和识别我没有打算写一个通用推理引擎那样工作量太大还会陷入算子兼容的泥潭。这个引擎只支持 PP-OCRv6 的检测模型和识别模型算子清单是固定的Conv2D、DepthwiseConv2D、BatchNorm、ReLU、HardSwish、SE、MaxPool、AvgPool、Resize、Concat、Add、Sigmoid、Softmax、Transpose、Reshape、ArgMax。检测分支用 DB 结构识别分支用 CRNN 或 SVTR 这类序列识别结构输出经过 CTC 解码。图结构不动态前向流程可以按顺序硬编码省掉图优化器的复杂度。这种“窄而深”的设计让我能把精力放在正确性和性能上。比如卷积我不需要支持空洞卷积、分组卷积的任意组合只需要支持 PP-OCRv6 里出现的 3x3、1x1、深度可分离 3x3。BN 在导出阶段就融合进卷积推理时少一层计算。激活函数也就那几种HardSwish 用查表或直接算都行。边界清晰之后测试用例也容易写拿固定输入和 Python 端逐层对比输出误差控制在 1e-4 以内基本就能确认实现正确。1.3 模型格式设计自定义 .jocr 包比 ONNX 更省事既然不用 ONNX Runtime我也没打算在 Java 里解析 ONNX。ONNX 的 protobuf 结构复杂算子属性多解析器写起来不划算。我选择在 Python 侧把 Paddle 权重导出成自定义二进制格式后缀 .jocr。文件开头是 magic number 和版本号接着是张量数量。每个张量按顺序存储名称长度、名称字节、维度数量、各维度大小、数据类型、数据偏移、数据字节。权重统一按小端 float32 写入Python 的 struct.pack 或 numpy tofile 都能做。除了权重.jocr 包里还塞了预处理参数和字典。检测模型的 mean、std、resize 策略识别模型的输入高度、字典文件全部打包在一起。这样做的好处是模型和配置永不脱节。以前用 ONNX 时模型文件、字典、预处理参数经常分开放换个模型忘了换字典识别结果直接乱码。现在一个文件包含所有信息版本管理简单很多。Java 侧用 MappedByteBuffer 读取大模型也不会一次性吃掉堆内存。1.4 模块划分核心层、算子层、模型层、应用层代码分成四块。核心层是 Tensor 和 Shape负责 NCHW 布局的 float 数组、维度计算、索引访问。算子层是 Conv2D、BatchNorm、Activation 等实现每个算子输入输出都是 Tensor不依赖具体模型。模型层负责读取 .jocr、构建检测网络和识别网络、管理权重。应用层是 OcrEngine对外暴露 detectAndRecognize(BufferedImage) 方法内部完成预处理、前向、后处理。这样分层的直接好处是替换模型方便。PP-OCRv6 如果后续出了小改版本只要算子不变我只换 .jocr 文件。如果新增了算子也只在算子层加一个类不会污染应用层。测试时也可以单独测算子拿一个 1x3x4x4 的输入做卷积和 numpy 对比比端到端调试快得多。2. 核心算子与 PP-OCRv6 结构拆解2.1 张量与内存布局NCHW 还是 NHWCPaddle 训练默认 NCHW导出时我保持 NCHW。Java 里用一维 float[] 存储shape 用 int[] 记录stride 手动算。比如 shape 是 [1, 3, 224, 224]索引 (n,c,h,w) 对应 ((n*C c)*H h)*W w。不用多维数组因为 Java 的多维数组是数组的数组访问要两次解引用缓存不友好而且每行都要 new内存碎片多。一维数组配合预计算 stride循环展开后速度明显更好。NCHW 对卷积实现也友好。输入通道、输出通道、卷积核高宽、步长、padding 这些参数确定后可以直接按输出位置循环内层对输入通道和卷积核做乘加。Java 的 JIT 对连续数组访问优化不错只要避免在热循环里做边界检查之外的事情。我还会把常用的 shape 和 stride 缓存到 Tensor 对象里避免每次访问都重新计算。2.2 卷积与 BN 融合推理加速的第一刀PP-OCRv6 的 backbone 里有大量 ConvBN激活。训练时 BN 是独立层推理时完全可以融合。BN 的公式是 y gamma * (x - mean) / sqrt(var eps) beta。把它代入卷积 y Wx b得到新的权重 W W * gamma / sqrt(var eps)新的偏置 b beta - gamma * mean / sqrt(var eps)。融合后卷积直接输出省掉 BN 的减均值、除标准差、乘 gamma、加 beta 四次逐元素操作。对于小模型这可能就是 10% 到 20% 的速度提升。我在 Python 导出阶段做融合Java 侧只看到融合后的 Conv。融合时要注意 eps 和 BN 的 momentum 无关推理只用 moving_mean 和 moving_variance。如果模型里 Conv 没有偏置就令 b0 再算。代码实现上遍历每个 BN 的 weight、bias、running_mean、running_var找到它前面的 Conv 权重按输出通道维度广播计算。导出后最好用 Python 推理一次确认融合前后输出一致再写入 .jocr。2.3 激活函数与 SE 模块HardSwish 和通道注意力PP-OCRv6 的 backbone 很可能用了 PP-LCNet 或 MobileNetV3 风格结构里面少不了 HardSwish 和 SE。HardSwish 的定义是 x * relu6(x 3) / 6Java 里直接算不需要查表因为输入范围不大。relu6 就是 min(max(x3, 0), 6)。SE 模块是通道注意力先对特征图做全局平均池化变成 1x1xC然后经过两个 1x1 卷积中间有 ReLU再 Sigmoid最后乘回原特征图。实现时全局平均池化就是每个通道求均值两个 1x1 卷积可以当成全连接用矩阵乘或逐通道加权。SE 模块的计算量不大但通道顺序容易搞错。我踩过一次坑全局平均池化后我把 C 通道的数据按 HWC 顺序传给全连接结果全错。NCHW 下每个通道是连续的求均值时要按通道循环而不是按空间位置循环。修正后输出就对齐了。HardSwish 在负值区域不是零这点和 ReLU 不同如果导出时误用了 ReLU识别精度会掉得厉害。2.4 检测头与识别头DB 和 CTC 的差异检测模型输出一张概率图尺寸通常是输入的四分之一每个像素表示该点属于文本区域的概率。训练时用了 DB 的近似二值化推理时后处理要做阈值化、膨胀、找轮廓、求外接矩形、unclip 扩展。识别模型输出 T×C 的序列T 是时间步C 是字符类别数包含 blank。CTC 解码时对每个时间步取 argmax然后合并重复字符并去掉 blank最后映射到字典。这两个头的数据布局不同。检测输出是 [1, 1, H, W]识别输出是 [1, T, C] 或 [T, 1, C]。我在 Tensor 里不区分统一按 shape 访问。检测后处理需要遍历 H×W识别后处理需要遍历 T。分开写两个后处理类共用一些工具方法比如二值化、连通域标记。这样结构清晰也方便单独优化。3. 从模型导出到 Java 加载的完整实操3.1 Python 侧导出把 Paddle 权重拆成二进制导出脚本的核心是遍历 state_dict把每个张量写成小端 float32。下面是我用的简化版代码真实项目里还会加 BN 融合和算子图序列化。注意 Paddle 的权重名可能带后缀比如conv2d_0.w_0我一般在导出时重命名成短名Java 侧按层名查权重。import struct import numpy as np import paddle def export_jocr(state_dict, out_path, preprocess_info, dict_lines): with open(out_path, wb) as f: f.write(bJOCR) f.write(struct.pack(I, 1)) # version f.write(struct.pack(I, len(state_dict))) for name, tensor in state_dict.items(): arr tensor.numpy().astype(np.float32) name_bytes name.encode(utf-8) f.write(struct.pack(I, len(name_bytes))) f.write(name_bytes) shape arr.shape f.write(struct.pack(I, len(shape))) for s in shape: f.write(struct.pack(I, s)) f.write(struct.pack(I, arr.size)) f.write(arr.tobytes(orderC)) # 预处理信息用 UTF-8 JSON 追加 info_bytes preprocess_info.encode(utf-8) f.write(struct.pack(I, len(info_bytes))) f.write(info_bytes) dict_bytes \n.join(dict_lines).encode(utf-8) f.write(struct.pack(I, len(dict_bytes))) f.write(dict_bytes)导出后一定要用 Python 端推理一次保存输入和输出。这个输入输出对就是后面 Java 逐层对比的基准。没有基准调试纯 Java 引擎会非常痛苦。3.2 Java 侧读取MappedByteBuffer 和字节序Java 读取时用RandomAccessFile打开文件FileChannel.map映射成MappedByteBuffer设置LITTLE_ENDIAN。按格式依次读 magic、版本、张量数量。每个张量读名称、维度、大小然后根据当前 position 读取 float 数组。因为映射的是文件不用一次性复制到堆里大模型也能加载。下面代码省略异常处理。public class JocrReader { public static MapString, Tensor read(String path) throws IOException { try (RandomAccessFile raf new RandomAccessFile(path, r); FileChannel ch raf.getChannel()) { MappedByteBuffer buf ch.map(FileChannel.MapMode.READ_ONLY, 0, ch.size()); buf.order(ByteOrder.LITTLE_ENDIAN); byte[] magic new byte[4]; buf.get(magic); if (!new String(magic).equals(JOCR)) throw new IOException(bad magic); int version buf.getInt(); int count buf.getInt(); MapString, Tensor map new HashMap(); for (int i 0; i count; i) { int nameLen buf.getInt(); byte[] nb new byte[nameLen]; buf.get(nb); String name new String(nb, StandardCharsets.UTF_8); int dims buf.getInt(); int[] shape new int[dims]; for (int j 0; j dims; j) shape[j] buf.getInt(); int size buf.getInt(); float[] data new float[size]; buf.asFloatBuffer().get(data); buf.position(buf.position() size * 4); map.put(name, new Tensor(data, shape)); } return map; } } }这里有个细节buf.asFloatBuffer().get(data)之后要手动移动原 buffer 的 position因为asFloatBuffer是视图不影响原 position。我一开始忘了移动结果读第二个张量时全错位。3.3 前向执行器硬编码流程比通用图执行器快PP-OCRv6 的检测和识别网络结构固定我没有写通用图执行器而是把前向流程硬编码成方法调用。检测网络按顺序conv1 - bn1 - relu - ... - det_head。识别网络按顺序conv1 - ... - rec_head。每层从权重 map 里取权重和偏置调用算子。这样做的好处是 JIT 能内联方法调用少调试也直观。通用图执行器需要 Map 查找、反射或接口调用热路径上开销不小。当然硬编码的代价是换模型要改代码。但我把每个层封装成类层与层之间只依赖 Tensor改结构时只是调整方法调用顺序不会太麻烦。比如检测网络里有个 Concat我就直接调用Tensor.concat(a, b, 1)不需要解析算子属性。3.4 预处理对齐一个像素都不能差预处理是最容易出错的地方。检测模型通常把输入 resize 到 32 的倍数比如高 640、宽 640同时保持长宽比短边补灰。识别模型把文本行 resize 到高 48宽按比例缩放但最大宽度有限制比如 320 或 640。归一化参数必须和训练一致。PP-OCRv6 检测常用 ImageNet 的 mean 和 std识别常用(img/255 - 0.5) / 0.5。这些参数我从 Python 导出时直接写进 .jocrJava 侧读出来用。Java 里用BufferedImage的getRGB拿像素注意 alpha 通道。如果图片是 PNG 带透明要先合成到白底否则透明区域变黑检测会漏。双线性插值要自己实现因为Graphics2D的缩放算法和 Python 的 cv2.resize 不完全一致。我写了一个resizeBilinear按目标像素反推源坐标取四个邻域加权。实测下来和 cv2 的差异在 1e-3 以内对 OCR 影响可忽略。3.5 检测后处理从概率图到文本框检测输出概率图后先按阈值二值化通常 0.3。然后做连通域标记我用了两遍扫描法第一遍给每个前景像素临时标号记录等价关系第二遍合并等价类得到每个连通域的像素列表。接着对每个连通域求最小外接矩形过滤掉面积太小或长宽比异常的框。DB 后处理还需要 unclip把框向外扩展扩展比例和周长有关公式是offset area * unclip_ratio / perimeter然后对多边形做偏移。纯 Java 做多边形偏移有点麻烦我简化成对矩形向外扩一定像素再按原始比例映射回原图坐标。对于票据这种规则文本效果够用。如果要做任意方向文本就得实现真正的多边形偏移和透视变换。透视变换我写了一个 3x3 矩阵求解用四个点对解线性方程然后对裁剪区域做双线性采样。这部分代码量不小但一次写好就能复用。3.6 识别后处理CTC 贪心解码与字典映射识别输出是 [1, T, C]C 包含 blank。贪心解码很简单对每个 t找 C 维最大值的索引如果索引不是 blank 且和前一个输出不同就加入结果序列。最后把索引映射到字典。字典通常是字符列表blank 的位置要确认。PaddleOCR 的字典有的把 blank 放最后有的不放导出时我把字典和 blank id 一起写进 .jocrJava 侧读配置避免硬编码。public static String ctcGreedy(float[][] logits, String[] dict, int blankId) { StringBuilder sb new StringBuilder(); int prev -1; for (float[] step : logits) { int maxIdx 0; float maxVal step[0]; for (int i 1; i step.length; i) { if (step[i] maxVal) { maxVal step[i]; maxIdx i; } } if (maxIdx ! blankId maxIdx ! prev) { if (maxIdx dict.length) sb.append(dict[maxIdx]); } prev maxIdx; } return sb.toString(); }如果识别结果有重复字符CTC 会合并这是正确的。但中文里有些词确实有连续相同字比如“谢谢”CTC 也能正确输出因为中间会有 blank 分隔。如果发现重复字丢失通常是 blank id 设错了。4. 常见问题与排查技巧实录4.1 识别结果全是空白或乱码最常见的原因是字典不对。我遇到过一次模型导出时字典顺序变了但 Java 侧还在用旧字典结果所有字都偏移一位输出全是乱码。排查方法是拿一张只有单个字的图片看 argmax 的索引再和 Python 端对比。如果索引一致但字不对就是字典问题如果索引不一致就是归一化或输入尺寸问题。另外blank id 设错也会导致输出全空因为所有时间步都被当成 blank 去掉了。4.2 检测框偏移、漏检检测框偏移通常是预处理 resize 后坐标还原错了。比如原图 1000x800resize 到 640x512再补到 640x640映射回原图时忘了减去 padding 的偏移。漏检则可能是阈值太高或者膨胀次数不够。DB 输出的是概率图文本边缘概率低阈值 0.3 比较通用。如果图片对比度低可以试 0.2。还有一个坑Java 的getRGB拿到的 RGB 顺序是 ARGB 打包的 int取红色分量要右移 16 位绿色 8 位蓝色 0 位。搞反了通道检测直接失效。4.3 性能慢从每秒 1 张到每秒 8 张第一版纯 Java 引擎跑一张票据要 1.2 秒太慢。优化后降到 150 毫秒左右。主要做了几件事BN 融合省掉大量逐元素计算卷积循环里把输入通道和输出通道的顺序调整让内层循环连续访问预分配输出 Tensor避免每次 new用float[]而不是Float[]把Math.exp换成查表或近似因为 Sigmoid 和 Softmax 调用频繁。最后开 4 个线程并行处理多张图吞吐量上来了。单张延迟没有本质变化但并发场景下 CPU 利用率满了。4.4 内存溢出im2col 是双刃剑我一开始用 im2col 把卷积转成矩阵乘实现简单但内存爆炸。一个 3x3 卷积在 256 通道上im2col 后的矩阵可能是原特征图的 9 倍。大图直接 OOM。后来改成直接卷积按输出像素循环内层对卷积核乘加。虽然代码多了点但内存稳定。如果非要用 im2col可以分块一次只展开部分输入通道。Java 的垃圾回收对短命大数组不友好频繁分配容易触发 Full GC能复用就复用。4.5 与 Python 输出对不齐逐层对比是唯一出路纯 Java 引擎最怕“看起来对但精度差一点”。我的做法是导出模型时在 Python 端每层都保存输入输出存成 npy。Java 侧写一个调试模式每层算完把结果存成文本或二进制。然后用脚本对比算最大绝对误差和余弦相似度。通常误差会从某一层开始变大那一层就是问题所在。常见问题包括padding 方式不一致、BN 的 eps 不同、激活函数用错、转置维度搞反。逐层对比虽然笨但最快定位。5. 性能实测与部署建议5.1 基准对比纯 Java 和 ONNX Runtime 的差距我在同一台 i7-11800H 上做了对比输入是一张 1200x900 的票据图检测加识别端到端。ONNX Runtime 用 Java APICPU 执行平均 95 毫秒。纯 Java 引擎第一版 1.2 秒优化后 180 毫秒。差距主要在大卷积和矩阵乘C 有 SIMD 和更好的缓存利用。但纯 Java 版本没有 native 依赖部署成本低很多。对于 QPS 低于 10 的场景180 毫秒完全够用。如果并发高可以多实例部署或者后续用 Java Vector API 做 SIMD 加速。5.2 Spring Boot 集成单例引擎和线程池集成方式很简单把 OcrEngine 做成单例 Bean初始化时加载 .jocr 模型。推理方法本身无状态但 Tensor 对象不复用所以线程安全。Spring MVC 的请求线程直接调用即可。如果并发量大建议单独配一个固定大小线程池比如 CPU 核数把 OCR 任务提交进去避免 Tomcat 线程被长时间占用。模型加载一次大约 1 到 2 秒内存占用检测加识别约 200MB堆内存给 1GB 足够。5.3 后续扩展量化和算子优化纯 Java 引擎还有优化空间。INT8 量化可以把权重和激活量化到 8 位内存减半速度提升但需要校准集和量化算子支持。Java Vector API 在 JDK 17 以上可以显式使用 SIMD卷积内层循环改成向量化后速度可能再提升 30%。不过这些都要在保证正确性的前提下做。我的建议是先把 FP32 版本跑稳再考虑量化。毕竟 OCR 对精度敏感量化掉点太多就得不偿失。最后再分享一个小技巧把预处理参数、字典、模型版本一起打进 .jocr 包并且在文件头写一个 CRC32 校验。这样模型文件损坏或者版本对不上时Java 侧启动就能发现不会等到线上识别出乱码才报警。我在实际使用中发现部署环境里最容易出问题的不是算法而是配置和文件版本。一个自包含的模型包能省掉很多扯皮。
返回列表