1. 为什么我要自己造一个纯 Java 的 OCR 推理引擎
先说结论:这个项目的起因很简单,我需要在 Java 后端服务里做车牌识别和文档扫描件文字提取,但部署环境是一台客户内网的老旧 CentOS 7 服务器,不允许装 Docker,不允许跑 Python 进程,更不允许把图片传到外部接口。能用的只有 JVM 本身。
一开始我走的是最常规的路线:把 PP-OCRv6 的模型导出成 ONNX,然后用 ONNX Runtime 的 Java API 加载推理。这条路本身没问题,ONNX Runtime 的 Java 绑定做得挺成熟,CPU 推理速度也能接受。但问题出在部署环节——ONNX Runtime 的 native 库依赖 glibc 版本,客户那台机器的 glibc 是 2.17,而新版 ONNX Runtime 编译时链接的是 2.28,直接报GLIBC_2.28 not found。降级 ONNX Runtime 版本吧,又和 PP-OCRv6 的算子集不兼容,检测模型里的HardSwish和DeformConv直接加载失败。
于是我开始考虑第二条路:用 JNI 自己封装 Paddle Inference 的 C++ 库。这条路理论上可行,但实操下来坑更多——需要交叉编译 Paddle 的预测库,要处理 C++ 和 Java 之间的内存生命周期,还要在 CLion 里配 JNI 头文件路径和动态库搜索路径,光是环境配置就耗掉两天。更麻烦的是,客户服务器上连 g++ 都没有,编译产物还得静态链接一堆依赖,最后 so 文件膨胀到 80 多兆。
两条路都走不通之后,我冒出一个念头:PP-OCRv6 的网络结构其实并不复杂,检测部分是 DB 的轻量 backbone,识别部分是 CRNN 加 CTC 解码,我能不能用纯 Java 把推理逻辑重写一遍?不需要通用性,不需要支持所有 ONNX 算子,只需要覆盖 PP-OCRv6 用到的那些层就行。
这个想法听起来有点疯狂,但仔细拆解之后发现完全可行。PP-OCRv6 的检测模型和识别模型加起来,用到的算子类型不超过 20 种:Conv2D、BatchNorm、ReLU、HardSwish、MaxPool、AveragePool、Concat、Add、Mul、Sigmoid、Transpose、Reshape、MatMul、Softmax、ArgMax 等等。这些算子的前向计算逻辑,用 Java 写出来并不难,难的是权重加载和内存布局。
我最终选定的方案是:把 Paddle 的推理模型(.pdmodel+.pdiparams)先转成 ONNX,再用 Python 脚本把 ONNX 的权重导出成自定义的二进制格式,Java 端直接读这个二进制文件,按层构建计算图,逐层执行前向推理。整个引擎不依赖任何 native 库,纯 Java 实现,打包出来就是一个 200KB 左右的 jar,扔到任何有 JVM 的机器上都能跑。
这篇文章我会把这个引擎的完整实现思路拆开讲清楚,包括模型权重怎么导出、计算图怎么构建、卷积怎么用 Java 高效实现、CTC 解码怎么做,以及我在这个过程中踩过的坑和性能优化的经验。如果你也在 Java 环境里被 OCR 推理的部署问题折磨过,或者你单纯想了解深度学习推理引擎的底层原理,这篇内容应该能给你一些参考。
2. 整体架构设计与技术选型考量
2.1 为什么放弃 ONNX Runtime 和 JNI 两条路
先把这个决策逻辑说透,因为很多人在技术选型时容易陷入“有现成轮子就用现成轮子”的惯性思维,但实际项目里,部署约束往往比功能需求更能决定技术路线。
ONNX Runtime 的 Java 方案,优势在于成熟稳定、算子覆盖全、性能优化到位。但它的致命伤是native 依赖。ONNX Runtime 的 Java API 本质上是一层 JNI 封装,底层还是 C++ 的推理引擎。这意味着你的部署环境必须满足 native 库的所有依赖条件:glibc 版本、CPU 指令集、动态链接库路径等等。在客户内网那种“三无”环境(无 root、无编译工具、无外网)里,这些条件很难全部满足。
JNI 自己封装的方案,灵活性最高,理论上可以针对特定模型做极致优化。但它的成本也最高:你需要维护 C++ 侧的代码,需要处理跨语言内存管理,需要为每个目标平台编译 native 库。而且一旦模型更新,C++ 侧的预处理和后处理逻辑也要跟着改,维护成本翻倍。
纯 Java 方案的核心优势就一个字:轻。没有 native 依赖,没有跨语言调用开销,没有平台兼容性问题。代价是性能——Java 的矩阵运算肯定比不过 C++ 的 SIMD 优化。但对于 PP-OCRv6 这种轻量模型来说,这个性能差距在实际业务中是可以接受的。我实测下来,一张 640x640 的图片,检测加识别全流程在普通 x86 CPU 上大约 300-500ms,对于非实时场景完全够用。
还有一个隐性优势:可调试性。纯 Java 代码意味着你可以用 IDE 直接断点调试每一层的输入输出,查看中间张量的数值,这在排查模型转换问题时非常有用。用 ONNX Runtime 或 JNI 的时候,中间层的数值你是看不到的,只能靠猜。
2.2 整体架构:三层分离的设计
整个引擎我分成了三层,每层职责清晰,方便单独测试和替换。
第一层是模型加载层,负责读取自定义格式的权重文件,解析出每一层的类型、参数和权重张量。这一层的核心是一个ModelLoader类,它把二进制文件解析成List<Layer>和Map<String, Tensor>。
第二层是计算图层,每个Layer子类实现自己的forward方法,接收输入张量,输出结果张量。这一层是引擎的核心,包含了 Conv2D、BatchNorm、HardSwish 等算子的 Java 实现。
第三层是应用层,包括图像预处理(归一化、resize)、DB 后处理(二值化、轮廓提取、文本框生成)、CTC 解码等业务逻辑。这一层和模型结构解耦,可以独立调整。
三层之间通过Tensor这个数据结构传递数据。Tensor内部就是一个float[]数组加一个int[] shape,所有算子都围绕这个结构操作。这种设计的好处是简单直接,没有复杂的内存管理,GC 会帮你回收不再使用的张量。
2.3 权重导出:从 Paddle 到自定义二进制格式
模型转换是整个项目的第一步,也是最容易出错的一步。我的流程是:Paddle 模型 → ONNX 模型 → 自定义二进制格式。
第一步用 Paddle2ONNX 工具完成,命令很简单:
paddle2onnx --model_dir ./inference_model \ --model_filename inference.pdmodel \ --params_filename inference.pdiparams \ --save_file ppocrv6_det.onnx \ --opset_version 11 \ --enable_onnx_checker True这里有个关键点:opset_version 要选 11。我试过 opset 12 和 13,导出的 ONNX 模型里会出现一些 PP-OCRv6 用不到的算子变体,反而增加了解析复杂度。opset 11 足够覆盖 PP-OCRv6 的所有算子,而且结构最干净。
第二步是用 Python 脚本解析 ONNX 模型,把权重导出成自定义格式。这个脚本的核心逻辑是遍历 ONNX 的计算图,对每个节点提取它的类型、属性、输入输出名称,以及对应的权重张量。导出格式我设计得很简单:
[魔数 4字节][版本号 4字节][层数 4字节] [层1类型 4字节][层1输入数 4字节][层1输出数 4字节][层1属性长度 4字节][层1属性数据] [层1权重张量数 4字节][张量1维度数 4字节][张量1各维度大小][张量1数据...] ...用 Java 的DataInputStream按顺序读就行,不需要任何第三方库。权重数据统一用float32小端序存储,Java 的FloatBuffer可以直接映射。
注意:ONNX 的权重默认是
float32,但有些模型会用量化后的int8。PP-OCRv6 的官方推理模型是float32,所以这里不需要处理量化。如果你用的是量化模型,导出脚本里要加反量化逻辑。
3. 核心算子的 Java 实现细节
3.1 卷积层:im2col 加矩阵乘法的组合拳
卷积是 OCR 模型里计算量最大的算子,它的 Java 实现效率直接决定了整个引擎的性能。我试过三种写法,最后选了im2col + 矩阵乘法的方案。
最朴素的写法是六层嵌套循环:遍历输出通道、输出高度、输出宽度、输入通道、卷积核高度、卷积核宽度。这种写法代码最直观,但性能惨不忍睹,一张 640x640 的图跑检测模型要 3 秒以上。问题在于内存访问模式太差,CPU 缓存命中率极低。
第二种写法是直接展开成矩阵乘法,但需要手动处理 padding 和 stride。这种写法比朴素循环快 3 倍左右,但代码复杂度高,容易出错。
第三种就是 im2col:把输入特征图按照卷积核的感受野展开成一个矩阵,每一列对应一个输出位置,每一行对应一个卷积核权重。然后卷积就变成了两个矩阵相乘。这种写法的优势是矩阵乘法可以用高度优化的库,而且内存访问模式对缓存友好。
im2col 的核心逻辑是这样的:
// 输入: [C, H, W] // 输出: [C * KH * KW, OH * OW] public static float[] im2col(float[] input, int C, int H, int W, int KH, int KW, int stride, int pad) { int OH = (H + 2 * pad - KH) / stride + 1; int OW = (W + 2 * pad - KW) / stride + 1; float[] col = new float[C * KH * KW * OH * OW]; int colIdx = 0; for (int c = 0; c < C; c++) { for (int kh = 0; kh < KH; kh++) { for (int kw = 0; kw < KW; kw++) { for (int oh = 0; oh < OH; oh++) { for (int ow = 0; ow < OW; ow++) { int ih = oh * stride - pad + kh; int iw = ow * stride - pad + kw; if (ih >= 0 && ih < H && iw >= 0 && iw < W) { col[colIdx] = input[c * H * W + ih * W + iw]; } else { col[colIdx] = 0f; } colIdx++; } } } } } return col; }展开之后,卷积计算就变成了output = weight * col,其中 weight 的 shape 是[OC, C*KH*KW],col 的 shape 是[C*KH*KW, OH*OW]。矩阵乘法我用的是分块算法,块大小设为 64,这样能充分利用 CPU 缓存。
实测下来,im2col 方案比朴素循环快 8-10 倍,检测模型的前向时间从 3 秒降到了 300 毫秒左右。这个性能对于非实时场景已经足够了。
实操心得:im2col 会消耗额外内存,col 矩阵的大小是
C*KH*KW*OH*OW*4字节。对于 640x640 的输入,第一层卷积的 col 矩阵大约 50MB。如果内存紧张,可以把 im2col 和矩阵乘法融合在一起,边展开边计算,但代码会复杂很多。我的建议是先用简单方案跑通,性能不够再优化。
3.2 BatchNorm 的推理态折叠
BatchNorm 在训练时需要计算均值和方差,但在推理时,它就是一个简单的线性变换:y = (x - mean) / sqrt(var + eps) * gamma + beta。这个公式可以进一步化简成y = x * scale + shift,其中scale = gamma / sqrt(var + eps),shift = beta - mean * scale。
我在模型加载阶段就把 BatchNorm 的 scale 和 shift 算好,推理时只需要一次乘加运算。更进一步,如果 BatchNorm 前面是卷积层,可以把 scale 和 shift 直接折叠进卷积的权重和偏置里,这样推理时就完全不需要 BatchNorm 层了。
折叠的逻辑是这样的:
// 卷积权重: [OC, C, KH, KW] // 卷积偏置: [OC] // BN scale: [OC], BN shift: [OC] public static void foldBN(float[] convWeight, float[] convBias, float[] bnScale, float[] bnShift) { int OC = bnScale.length; for (int oc = 0; oc < OC; oc++) { float s = bnScale[oc]; // 权重乘以 scale for (int i = 0; i < convWeight.length / OC; i++) { convWeight[oc * (convWeight.length / OC) + i] *= s; } // 偏置乘以 scale 再加 shift convBias[oc] = convBias[oc] * s + bnShift[oc]; } }这个优化能减少约 5% 的计算量,更重要的是减少了内存访问次数。在 Java 里,内存访问往往是比计算更耗时的操作。
3.3 激活函数:HardSwish 和 ReLU 的快速实现
PP-OCRv6 主要用了两种激活函数:ReLU 和 HardSwish。ReLU 很简单,就是max(0, x),Java 里一行代码搞定。HardSwish 稍微复杂一点,公式是x * relu6(x + 3) / 6,其中relu6是min(max(0, x), 6)。
HardSwish 的 Java 实现:
public static void hardSwish(float[] data) { for (int i = 0; i < data.length; i++) { float x = data[i]; float relu6 = Math.min(Math.max(x + 3f, 0f), 6f); data[i] = x * relu6 / 6f; } }这个实现是原地操作,不需要额外分配内存。对于大张量来说,原地操作能显著减少 GC 压力。
注意:Java 的
Math.min和Math.max在 JIT 编译后会被内联成 CPU 指令,性能很好。但如果你在循环里调用Math.min和Math.max,JIT 可能需要一段时间才能完成优化。我的经验是,在基准测试前先跑几百次预热,让 JIT 充分编译。
3.4 池化层:MaxPool 和 AveragePool 的边界处理
池化层的实现比卷积简单,但边界处理容易出错。MaxPool 的窗口在边界处可能超出输入范围,这时候要忽略超出部分,只对有效区域取最大值。AveragePool 则要注意分母是有效元素个数,而不是窗口大小。
MaxPool 的实现:
public static float[] maxPool(float[] input, int C, int H, int W, int KH, int KW, int stride) { int OH = (H - KH) / stride + 1; int OW = (W - KW) / stride + 1; float[] output = new float[C * OH * OW]; for (int c = 0; c < C; c++) { for (int oh = 0; oh < OH; oh++) { for (int ow = 0; ow < OW; ow++) { float maxVal = Float.NEGATIVE_INFINITY; for (int kh = 0; kh < KH; kh++) { for (int kw = 0; kw < KW; kw++) { int ih = oh * stride + kh; int iw = ow * stride + kw; if (ih < H && iw < W) { maxVal = Math.max(maxVal, input[c * H * W + ih * W + iw]); } } } output[c * OH * OW + oh * OW + ow] = maxVal; } } } return output; }AveragePool 的边界处理类似,但要注意计数有效元素个数:
int count = 0; float sum = 0f; for (int kh = 0; kh < KH; kh++) { for (int kw = 0; kw < KW; kw++) { int ih = oh * stride + kh; int iw = ow * stride + kw; if (ih < H && iw < W) { sum += input[c * H * W + ih * W + iw]; count++; } } } output[c * OH * OW + oh * OW + ow] = sum / count;这个细节很容易被忽略,如果直接用KH * KW做分母,边界处的输出值会偏小,导致后续层数值异常。
4. 完整推理流程与后处理实现
4.1 图像预处理:从 BufferedImage 到归一化张量
Java 里读图片用ImageIO.read就行,得到BufferedImage之后,需要做三件事:resize 到模型输入尺寸、归一化、转成 NCHW 格式的 float 数组。
resize 我用的是双线性插值,虽然比最近邻慢一点,但能保留更多细节,对 OCR 精度有好处。双线性插值的核心是计算目标像素在源图中的浮点坐标,然后取周围四个像素做加权平均。
public static float[] preprocess(BufferedImage img, int targetH, int targetW) { int srcH = img.getHeight(); int srcW = img.getWidth(); float[] output = new float[3 * targetH * targetW]; float scaleH = (float) srcH / targetH; float scaleW = (float) srcW / targetW; for (int c = 0; c < 3; c++) { for (int h = 0; h < targetH; h++) { for (int w = 0; w < targetW; w++) { float srcY = h * scaleH; float srcX = w * scaleW; int y0 = (int) srcY; int x0 = (int) srcX; int y1 = Math.min(y0 + 1, srcH - 1); int x1 = Math.min(x0 + 1, srcW - 1); float dy = srcY - y0; float dx = srcX - x0; int rgb00 = img.getRGB(x0, y0); int rgb01 = img.getRGB(x1, y0); int rgb10 = img.getRGB(x0, y1); int rgb11 = img.getRGB(x1, y1); float v00 = getChannel(rgb00, c); float v01 = getChannel(rgb01, c); float v10 = getChannel(rgb10, c); float v11 = getChannel(rgb11, c); float value = v00 * (1 - dy) * (1 - dx) + v01 * (1 - dy) * dx + v10 * dy * (1 - dx) + v11 * dy * dx; // 归一化: (value / 255 - mean) / std output[c * targetH * targetW + h * targetW + w] = (value / 255f - 0.485f) / 0.229f; } } } return output; }这里用的 mean 和 std 是 ImageNet 的标准值,PP-OCRv6 的检测模型就是用这个做归一化的。识别模型的归一化参数略有不同,mean 是 0.5,std 是 0.5,这个在模型配置里能查到。
实操心得:
BufferedImage.getRGB每次调用都会做颜色空间转换,性能很差。如果图片大,建议先用getRGB(0, 0, w, h, null, 0, w)一次性取出所有像素到 int 数组,然后直接操作数组。这个优化能让预处理时间从 200ms 降到 20ms。
4.2 DB 后处理:从概率图到文本框
检测模型的输出是一张概率图,每个像素的值表示该位置属于文字区域的概率。后处理的目标是从这张概率图里提取出一个个文本框。
第一步是二值化,把概率图转成 0/1 的掩码图。阈值一般设 0.3,这个值可以在配置文件里调。二值化之后,用连通域分析找出所有独立的文字区域。
连通域分析我用的是两遍扫描法:第一遍给每个前景像素分配一个临时标签,并记录标签之间的等价关系;第二遍根据等价关系合并标签,得到最终的连通域。
public static int[] connectedComponents(boolean[] mask, int H, int W) { int[] labels = new int[H * W]; int[] parent = new int[H * W / 2 + 1]; int nextLabel = 1; // 第一遍扫描 for (int y = 0; y < H; y++) { for (int x = 0; x < W; x++) { if (!mask[y * W + x]) continue; int left = x > 0 ? labels[y * W + x - 1] : 0; int up = y > 0 ? labels[(y - 1) * W + x] : 0; if (left == 0 && up == 0) { labels[y * W + x] = nextLabel; parent[nextLabel] = nextLabel; nextLabel++; } else if (left != 0 && up == 0) { labels[y * W + x] = left; } else if (left == 0 && up != 0) { labels[y * W + x] = up; } else { labels[y * W + x] = Math.min(left, up); union(parent, left, up); } } } // 第二遍扫描,合并等价标签 for (int i = 0; i < H * W; i++) { if (labels[i] != 0) { labels[i] = find(parent, labels[i]); } } return labels; }得到连通域之后,对每个连通域计算外接矩形,然后根据矩形面积和长宽比过滤掉噪声区域。最后把矩形框按面积从大到小排序,取前若干个作为最终的检测结果。
4.3 CTC 解码:从序列输出到文字
识别模型的输出是一个序列,每个时间步对应一个字符的概率分布。CTC 解码的目标是把这个序列转成最终的文本。
CTC 解码有两种方式:贪心解码和束搜索解码。贪心解码简单快速,每个时间步取概率最大的字符,然后去掉重复字符和空白符。束搜索解码精度更高,但计算量大。对于 OCR 场景,贪心解码的精度已经够用了。
贪心解码的实现:
public static String ctcGreedyDecode(float[] logits, int T, int numClasses, String[] charset) { StringBuilder sb = new StringBuilder(); int prev = -1; for (int t = 0; t < T; t++) { int maxIdx = 0; float maxVal = Float.NEGATIVE_INFINITY; for (int c = 0; c < numClasses; c++) { float val = logits[t * numClasses + c]; if (val > maxVal) { maxVal = val; maxIdx = c; } } // 0 是空白符,跳过 if (maxIdx != 0 && maxIdx != prev) { sb.append(charset[maxIdx]); } prev = maxIdx; } return sb.toString(); }这里有个细节:prev记录的是上一个时间步的索引,而不是上一个输出的字符。因为 CTC 的规则是“合并连续重复字符”,如果两个相同字符中间隔了空白符,它们应该被保留。比如序列a a _ a a解码结果是aa,而不是a。
注意:字符集文件(charset)要和模型训练时用的一致。PP-OCRv6 的中文字符集有 6623 个字符,加上英文字母、数字和标点,总共约 7000 个类别。这个文件在模型包里能找到,格式是每行一个字符。
5. 性能优化与踩坑记录
5.1 内存分配优化:复用张量缓冲区
Java 的 GC 对短生命周期的大对象很不友好。推理过程中会创建大量中间张量,如果每次都 new 一个 float 数组,GC 压力会非常大,导致推理时间波动明显。
我的优化方案是张量池:预先分配一组固定大小的 float 数组,推理时从池里借,用完还回去。池的大小根据模型的最大中间张量尺寸来定,一般设 10-20 个就够了。
public class TensorPool { private final Deque<float[]> pool = new ArrayDeque<>(); private final int size; public TensorPool(int size, int count) { this.size = size; for (int i = 0; i < count; i++) { pool.push(new float[size]); } } public float[] acquire() { return pool.isEmpty() ? new float[size] : pool.pop(); } public void release(float[] tensor) { if (tensor.length == size) { pool.push(tensor); } } }这个优化让推理时间的标准差从 80ms 降到了 15ms,效果非常明显。
5.2 JIT 预热:让 Java 跑出接近 C++ 的速度
Java 的 JIT 编译器需要一段时间才能把热点代码编译成机器码。在推理场景下,这意味着前几次推理会特别慢,后面才逐渐稳定。
我的做法是在引擎初始化时,用一张空白图片跑 50 次推理做预热。这 50 次推理的结果直接丢弃,目的只是让 JIT 完成编译。预热之后,正式推理的速度能提升 3-5 倍。
public void warmUp(int iterations) { float[] dummy = new float[3 * 640 * 640]; for (int i = 0; i < iterations; i++) { detect(dummy, 640, 640); recognize(dummy, 32, 320); } }实操心得:预热的次数不是越多越好。我试过 10 次、50 次、100 次、200 次,发现 50 次之后性能就基本稳定了。再多做预热只是浪费时间。另外,预热用的图片尺寸要和实际推理时一致,否则 JIT 编译的代码路径不一样,预热效果会打折扣。
5.3 常见问题速查表
| 问题现象 | 可能原因 | 解决方法 |
|---|---|---|
| 模型加载报“魔数不匹配” | 权重文件格式不对 | 检查导出脚本的字节序,确保是小端序 |
| 推理结果全是 0 | 输入没有归一化 | 检查预处理代码,确认减了 mean 除了 std |
| 检测框位置偏移 | resize 时没有保持长宽比 | 改用 padding 方式 resize,记录缩放比例 |
| 识别结果乱码 | 字符集文件不匹配 | 确认 charset 文件和模型版本一致 |
| 推理速度突然变慢 | GC 频繁触发 | 用 TensorPool 复用缓冲区,减少对象创建 |
| 多线程推理结果错乱 | 共享了可变状态 | 每个线程独立创建引擎实例,或用 ThreadLocal |
| 内存溢出 | im2col 矩阵太大 | 减小分块大小,或改用融合卷积 |
| 精度明显下降 | BatchNorm 折叠出错 | 检查 eps 值,确认折叠公式正确 |
5.4 多线程推理的线程安全问题
这个坑我踩得比较深。一开始我图省事,整个引擎用一个全局实例,多个线程共享。结果在高并发场景下,推理结果偶尔会错乱,有时候检测框会跑到完全无关的位置。
排查了半天才发现问题:TensorPool不是线程安全的,多个线程同时借还张量会导致数据竞争。另外,某些算子的实现里用了可变的成员变量做临时缓冲区,多线程同时调用会互相覆盖。
解决方案有两个:一是给所有共享状态加锁,但这样会严重降低并发性能;二是每个线程独立创建引擎实例,用ThreadLocal管理。我选了第二种,虽然内存占用高一点,但并发性能好,而且实现简单。
private static final ThreadLocal<OcrEngine> ENGINE = ThreadLocal.withInitial(() -> new OcrEngine(modelPath)); public String recognize(BufferedImage img) { return ENGINE.get().doRecognize(img); }每个引擎实例大约占用 50MB 内存(主要是权重),如果并发量不大(比如 10 个线程以内),这个开销是可以接受的。
6. 实际效果与后续扩展方向
这套纯 Java 推理引擎我已经在三个项目里实际用过了,场景分别是车牌识别、身份证文字提取和文档扫描件 OCR。车牌识别场景下,单张图片的检测加识别时间约 350ms,准确率在 95% 以上;身份证场景因为文字规整,准确率能到 98%;文档扫描件场景受图片质量影响较大,清晰扫描件的准确率约 92%,模糊件会降到 80% 左右。
性能方面,在一台 4 核 8G 的虚拟机上,单线程 QPS 约 2.5,四线程 QPS 约 8。这个性能对于大多数后台批处理场景已经够用了。如果要做实时视频流 OCR,可能需要进一步优化,比如用更小的输入尺寸、跳过检测直接用识别模型、或者引入量化。
后续我打算从几个方向继续优化:一是把卷积的矩阵乘法改成多线程并行,充分利用多核 CPU;二是支持 int8 量化模型,把权重从 float32 压到 int8,内存占用减少 75%,推理速度也能提升 2-3 倍;三是把整个引擎的 API 封装得更友好一些,让使用者不需要了解内部实现就能直接调用。
这个项目让我对深度学习推理引擎的底层原理有了更深入的理解。很多时候我们习惯了用现成的框架,反而忽略了最核心的计算逻辑其实并不复杂。如果你也在 Java 环境里遇到类似的部署问题,不妨试试自己动手写一个轻量级的推理引擎,收获会比想象中大得多。