简介:Java 与 Python 混合调用 YOLO ONNX 模型完成视频目标检测的完整工程,面向需要将深度学习能力接入 Java 服务端、处理 RTSP/RTMP 实时视频流的开发者,覆盖 YOLOv5、YOLOv7、YOLOv8 等主流模型,解决跨语言推理、数据格式转换和检测结果可视化等实际问题。压缩包共 68 个文件、约 272MB,包括 17 个 Java 源码、5 个 ONNX 模型、1 个 Python 推理脚本,以及演示视频、效果截图、GIF 动图和说明文档,目录结构清晰,便于直接阅读与二次开发。已有 265 人学习。工程完整展示了从视频流获取、帧预处理、模型调用到后处理的全链路:Java 端负责解析 RTSP/RTMP 视频流,对视频帧进行缩放、归一化、填补操作并转换为 Numpy 数组,再通过 JNI 等机制调用 Python 脚本加载 ONNX 模型执行目标检测;返回结果后,Java 端再做置信度过滤、识别框绘制等后处理。配合演示视频和预览图,可快速验证同一套方案在不同 YOLO 版本下的运行效果,适合正在做目标检测落地的中级以上 Java 开发者参考,能显著减少跨语言集成的踩坑时间。
1. Java 调 Python YOLO ONNX 模型做视频检测:先看清这套架构的边界
Java 后台要接视频目标检测,算法同事丢来三个 .pt 权重,yolov5s、yolov7、yolov8n,说“你拿去集成吧”。直接拿 Java 写一堆 letterbox 和 NMS 的代价,比想象中大得多。把 PyTorch 模型导出成 ONNX,再让 Python 侧用 ONNX Runtime 常驻推理,Java 只负责取帧、发送和画框,这套「Java 调用 Python YOLO ONNX 模型进行视频目标检测与识别」的管线,是我做下来最不容易翻车的组合。它适合手里已有现成 YOLO 权重、主系统是 Java、又不想把模型后处理逻辑重写一遍的团队。识别这一步,最终落在类别 ID 到名称的映射上,Java 端拿回 JSON 画框标字即可。
2. 把 PyTorch 权重导出成 ONNX:YOLOv5/v7/v8 的导出命令与输出校验
2.1 三个版本的导出命令差异
先说环境前提:Python 3.8 以上,装好 onnxruntime 和 torch。没装 Python 的先去官网装一个,安装时记得把 Add to PATH 勾上,否则后面 Java 的 ProcessBuilder 会报找不到 python 命令。预训练权重在官方 release 页可以直接下载,yolov8 用命令行会自动拉取。
YOLOv5 的导出在仓库根目录执行:
# yolov5 仓库根目录 python export.py --weights yolov5s.pt --include onnx --opset 12注意 v5 的 export.py 有个--grid参数。不加它,导出的是三个尺度原始特征图,后处理要自己写 grid 偏移和 anchor 先验;加上它,直接输出解码后的 1x25200x85。我在 ORT 里用后一种,省事。
YOLOv7 的导出:
python export.py --weights yolov7.pt --grid --simplify --img-size 640 640v7 的--end2end参数是给 TensorRT 做内置 NMS 用的,ONNX Runtime 里用不上,我一般不碰。--simplify会走一遍 onnx-simplifier,把多余的 Shape、Gather 节点清掉。
YOLOv8 用 ultralytics 的 CLI:
yolo export model=yolov8n.pt format=onnx opset=12 imgsz=640v8 导出默认输出 1x84x8400,84 是 4 个坐标加 80 个类别分数,8400 是三个尺度中心点总数,没有 objectness 分支。部分新版本支持transposed=True,输出变成 1x8400x84,两者只有维度顺序差异,后面校验脚本能一眼看出来。
2.2 用 ONNX Runtime 校验输出:输入输出名与 shape
导出后先别急着写后处理,用一段脚本把输入输出名、shape 打出来。onnx 只是模型文件格式,onnxruntime 是真正干活的推理引擎,这两者常被混着说,但命令和报错信息完全不一样。
import onnxruntime as ort import numpy as np import cv2 sess = ort.InferenceSession("yolov8n.onnx", providers=["CPUExecutionProvider"]) inp = sess.get_inputs()[0] out = sess.get_outputs()[0] print("input:", inp.name, inp.shape, inp.type) print("output:", out.name, out.shape, out.type) # 单帧冒烟测试 frame = cv2.imread("frame.jpg") frame = cv2.resize(frame, (640, 640)) x = frame[:, :, ::-1].transpose(2, 0, 1)[None].astype(np.float32) / 255.0 pred = sess.run([out.name], {inp.name: x})[0] print("pred shape:", pred.shape)这段脚本做了三件事:确认输入名和 shape;确认输出名和 shape;验证预处理链路。frame[:, :, ::-1]是把 OpenCV 读进来的 BGR 转成 RGB,transpose(2, 0, 1)[None]是把 HWC 转成 NCHW,除以 255 是归一化。跑出来如果是(1, 84, 8400)就是 v8,(1, 25200, 85)就是带 decode 的 v5/v7。第一次跑先不碰后处理,这一步就能暴露八成格式问题。
2.3 导出参数三个必调项:opset、simplify、动态轴
opset 选 12 到 14 比较稳妥,ORT 1.11 以上都能跑。opset 开太高,旧版本 ORT 会直接报 UnsupportedOperator,这时候最容易走玄学排障路线。v5 导出自带 simplify,v8 导出完建议自己再跑一次 onnxsim,去掉多余节点后输出名会更干净,Java/Python 两侧对模型文件也更有底。
动态轴我一般不开。固定 640x640 输入,模型加载后每次推理的内存布局一致,后处理里 grid 也是写死的,比动态 shape 省心。动态输入每次 shape 变化时 ORT 第一次推理有额外开销,Java 端也难预估单帧耗时。如果确实要跑多种分辨率,我宁愿多导出几个固定尺寸的 onnx 文件,按场景切换,不搞一个万能动态模型。
3. Java 和 Python 的三种联调方式:选型对比与最小可跑通的常驻服务
3.1 三种联调方式的取舍:进程启动、常驻服务、Java 直调
Java 调 Python 里的 ONNX 模型,常见做法有三种:
| 方案 | 优点 | 缺点 | 适用场景 |
|---|---|---|---|
| 每帧启动 ProcessBuilder 跑 python 脚本 | 实现最简单,代码只有几行 | Python 解释器启动加模型加载要 1 到 3 秒,帧率掉到个位数 | 离线批量测试 |
| Python 常驻进程 + 本机 Socket 通信 | 一次加载多次推理,后处理改完重启进程就生效 | 要处理帧序列化和进程生命周期 | 生产环境,项目选型推荐 |
| Java 直调 ONNX Runtime Java API | 吞吐最高,没有跨进程拷贝 | letterbox、NMS、三种版本解码全要在 Java 里重写 | 纯 Java 团队且模型不再频繁换 |
ONNX Runtime 确实有 Java 绑定,但 YOLO 后处理细节不少。v5 要算 objectness 乘类别分数,v8 直接取类别最大值,NMS 还要自己实现或者调 OpenCV 的接口。模型一换,Java 代码就得跟着发版,这个维护成本会持续叠加。我选第二种:Java 管视频流和业务,Python 管模型推理,两边职责清楚,模型更新只换 onnx 文件、重启 Python 进程,Java 端一行不动。
3.2 Python 端 ONNX Runtime 常驻推理服务
Python 端起一个 Socket 服务,监听 127.0.0.1。协议很简单:先收 4 字节大端长度,再收完整 JPEG 帧;返回结果用同样方式。绑定本机回环地址就够了,不暴露给外部。
import socket import struct import threading import json import numpy as np import cv2 import onnxruntime as ort class YoloDetector: def __init__(self, onnx_path): self.session = ort.InferenceSession(onnx_path, providers=["CPUExecutionProvider"]) self.input_name = self.session.get_inputs()[0].name self.input_size = 640 self.lock = threading.Lock() def preprocess(self, jpg_bytes): img = cv2.imdecode(np.frombuffer(jpg_bytes, np.uint8), cv2.IMREAD_COLOR) h, w = img.shape[:2] scale = min(self.input_size / w, self.input_size / h) nw, nh = int(w * scale), int(h * scale) resized = cv2.resize(img, (nw, nh)) canvas = np.full((self.input_size, self.input_size, 3), 114, np.uint8) dw, dh = (self.input_size - nw) // 2, (self.input_size - nh) // 2 canvas[dh:dh + nh, dw:dw + nw] = resized blob = canvas[:, :, ::-1].transpose(2, 0, 1)[None].astype(np.float32) / 255.0 return blob, scale, dw, dh def decode(self, pred, conf_thres, iou_thres): pred = np.squeeze(pred) if pred.shape[0] > pred.shape[1]: pred = pred.T # v5/v7: (N, 85),其中第 5 列是 objectness # v8: (N, 84),没有 objectness,直接取类别最大值 if pred.shape[1] == 85: obj = pred[:, 4:5] cls_score = pred[:, 5:] scores = (obj * cls_score).max(axis=1) cls = cls_score.argmax(axis=1) else: cls_score = pred[:, 4:] scores = cls_score.max(axis=1) cls = cls_score.argmax(axis=1) keep = scores > conf_thres boxes = pred[keep, :4] scores = scores[keep] cls = cls[keep] if len(scores) == 0: return [] x1 = boxes[:, 0] - boxes[:, 2] / 2 y1 = boxes[:, 1] - boxes[:, 3] / 2 x2 = boxes[:, 0] + boxes[:, 2] / 2 y2 = boxes[:, 1] + boxes[:, 3] / 2 boxes = np.stack([x1, y1, x2, y2], axis=1) idx = cv2.dnn.NMSBoxes(boxes.tolist(), scores.tolist(), conf_thres, iou_thres) if len(idx) == 0: return [] idx = np.array(idx).flatten() return boxes[idx], scores[idx], cls[idx] def infer(self, jpg_bytes, conf=0.25, iou=0.45): blob, scale, dw, dh = self.preprocess(jpg_bytes) with self.lock: pred = self.session.run(None, {self.input_name: blob})[0] boxes, scores, cls = self.decode(pred, conf, iou) result = [] for (x1, y1, x2, y2), score, c in zip(boxes, scores, cls): result.append({ "x1": round(float((x1 - dw) / scale), 2), "y1": round(float((y1 - dh) / scale), 2), "x2": round(float((x2 - dw) / scale), 2), "y2": round(float((y2 - dh) / scale), 2), "score": round(float(score), 4), "cls": int(c) }) return resultdecode 里 85 和 84 的区分是这套代码同时吃三种模型的关键。v5/v7 的 model shape 是 25200x85,85 里第 5 列是 objectness,所以分数要拿 objectness 乘类别概率;v8 是 8400x84,直接对类别维度取最大值就行。
服务端主循环:
def handle(conn): while True: header = conn.recv(4) if len(header) < 4: break length = struct.unpack(">I", header)[0] buf = b"" while len(buf) < length: chunk = conn.recv(length - len(buf)) if not chunk: break buf += chunk dets = detector.infer(buf) payload = json.dumps(dets).encode() conn.sendall(struct.pack(">I", len(payload)) + payload) srv = socket.socket(socket.AF_INET, socket.SOCK_STREAM) srv.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) srv.bind(("127.0.0.1", 9010)) srv.listen(5) while True: conn, _ = srv.accept() threading.Thread(target=handle, args=(conn,), daemon=True).start()长度前缀必须用大端>I,和 Java 端 DataOutputStream.writeInt 默认行为一致。收帧时不能只 recv 一次就假设拿完整包,TCP 是流协议,分包和粘包都会遇到,循环收满再解析是底线。session.run 外面加了锁,多线程连接同时进来时不会出现偶发崩溃。
3.3 Java 端启动 Python 进程与 Socket 客户端封装
Java 端用 ProcessBuilder 把 Python 服务拉起来,之后所有检测请求都走同一个 Socket 连接。
public class YoloClient implements Closeable { private final DataInputStream in; private final DataOutputStream out; public YoloClient(String host, int port) throws IOException { Socket socket = new Socket(host, port); in = new DataInputStream(new BufferedInputStream(socket.getInputStream())); out = new DataOutputStream(new BufferedOutputStream(socket.getOutputStream())); } public List<DetectedBox> detect(byte[] jpeg) throws IOException { out.writeInt(jpeg.length); // 大端,与 Python struct.pack(">I") 对齐 out.write(jpeg); out.flush(); int len = in.readInt(); byte[] buf = new byte[len]; in.readFully(buf); // readFully 保证拿完整包 return parseJson(new String(buf, StandardCharsets.UTF_8)); } public static void startPythonService() throws IOException { ProcessBuilder pb = new ProcessBuilder( "python", "/opt/yolo/server.py", "--model", "/opt/yolo/yolov8n.onnx", "--port", "9010"); pb.redirectErrorStream(true); pb.redirectOutput(new File("/var/log/yolo-service.log")); pb.start(); // 轮询端口,等服务真正起来 for (int i = 0; i < 20; i++) { try (Socket s = new Socket("127.0.0.1", 9010)) { return; } catch (IOException e) { Thread.sleep(500); } } throw new IOException("python service not ready"); } }detect 方法里那个out.flush()很容易漏。BufferedOutputStream 会把 writeInt 和 write 都缓存住,不 flush 数据就躺在缓冲区,Python 端一直等不到请求。readFully 则是处理半包的后悔药,普通 read 一次读不满 JSON 就得自己拼,readFully 内部帮你解决。
生产部署时我更推荐把 Python 服务单独拉起,交给 systemd 托管,Java 只当客户端,进程生命周期不绑在一起。ProcessBuilder 拉起床的唯一好处是演示环境下够简单。
4. Java 端视频取帧与结果回填:JavaCV 管线、坐标还原与画框
4.1 JavaCV 取流与 JPEG 编码
JavaCV 的 FFmpegFrameGrabber 能同时处理本地 MP4 和 RTSP 摄像头流。取帧后转成 Mat,再编码成 JPEG 发给 Python,这是目前最顺的链路。
FFmpegFrameGrabber grabber = new FFmpegFrameGrabber("rtsp://admin:password@192.168.1.10:554/stream1"); grabber.setOption("rtsp_transport", "tcp"); // UDP 在弱网下疯狂丢包 grabber.start(); OpenCVFrameConverter.ToMat toMat = new OpenCVFrameConverter.ToMat(); MatOfByte buf = new MatOfByte(); long intervalMs = 100; // 每秒最多 10 帧 while (true) { long t0 = System.currentTimeMillis(); Frame frame = grabber.grabImage(); if (frame == null) break; Mat bgr = toMat.convert(frame); Imgcodecs.imencode(".jpg", bgr, buf); // JPEG 质量由 ENCODE_PARAM 控制 byte[] jpeg = buf.toArray(); List<DetectedBox> dets = client.detect(jpeg); drawBoxes(bgr, dets); // 按节拍取帧,不依赖 grabber 内部帧率 long elapsed = System.currentTimeMillis() - t0; if (elapsed < intervalMs) Thread.sleep(intervalMs - elapsed); // 业务处理完记得 release }RTSP 传输用 TCP 而不是默认 UDP,视频流在弱网下不会变成满屏马赛克,检测框也不会跟着乱飘。不要依赖 setFrameRate 来节流,那个参数在流输入场景下并不控制取帧节奏,手动按 100 毫秒间隔 sleep 更可靠。
JPEG 质量是个容易忽略的坑。JDK 自带 ImageIO 不开放质量参数,默认压缩率偏高,小目标会被压糊。用 OpenCV 的 imencode 可以控制:
MatOfInt params = new MatOfInt( Imgcodecs.IMWRITE_JPEG_QUALITY, 85, Imgcodecs.IMWRITE_JPEG_OPTIMIZE, 1); Imgcodecs.imencode(".jpg", bgr, buf, params);质量 85 是我常用的值,再低检测精度开始下降,再高编码开销变大,对检测结果没有额外收益。
4.2 坐标还原与检测类别回填
Python 端返回的 x1、y1、x2、y2 已经是原始视频帧坐标,因为 server.py 里用(x - pad) / scale做过 letterbox 逆变换。Java 端拿到后直接画框,不需要再乘任何缩放系数。这里的分工必须写清楚:还原坐标是 Python 端的事,Java 端只负责展示和存储。
private void drawBoxes(Mat bgr, List<DetectedBox> dets) { for (DetectedBox b : dets) { Rect rect = new Rect((int) b.x1, (int) b.y1, (int) (b.x2 - b.x1), (int) (b.y2 - b.y1)); Imgproc.rectangle(bgr, rect, new Scalar(0, 255, 200), 2); String label = COCO_CLASSES[b.cls] + " " + b.score; Imgproc.putText(bgr, label, new Point(b.x1, Math.max(b.y1 - 6, 20)), Imgproc.FONT_HERSHEY_SIMPLEX, 0.6, new Scalar(0, 255, 200), 1); } }COCO_CLASSES 是 80 个类别名的静态数组,网上到处都能找到完整清单。如果是自定义数据集,自己训练时会在 data.yaml 里保留类别名,写个启动时读取的逻辑替换掉硬编码数组即可。putText 的 y 坐标要防止框贴顶时文字画到画面外,所以取了max(y1 - 6, 20)。
4.3 取帧参数参考表
| 参数 | 推荐值 | 说明 |
|---|---|---|
| rtsp_transport | tcp | 避免信道丢包导致检测目标缺损 |
| 抽帧间隔 | 100ms | 与单帧推理耗时匹配,留出余量 |
| JPEG 质量 | 85 | 低于 70 小目标容易漏检 |
| 输入尺寸 | 640x640 | 固定尺寸,ORT 推理消耗可控 |
| 服务地址 | 127.0.0.1:9010 | 本机回环,不占外网端口 |
取帧管线里丢一两帧没关系。检测任务本来就按节拍抽帧,不用每帧都送推理。如果推理耗时 80 毫秒,抽帧间隔设 100 毫秒最合理;设成 33 毫秒只会让请求在 Python 端排队,内存占用慢慢涨上去。
5. 联调避坑记录:五个把方案拖垮的常见问题
5.1 letterbox pad 没还原:框整体朝右下飘
现象:检测框位置大体对,但越往右下角偏得越狠,小目标偏得尤其明显。
原因:输入模型之前做了 letterbox,把原图缩放并填充到 640x640。模型输出的坐标是在 640x640 坐标系里的,如果不减 pad、不除以缩放比,直接当成原图坐标画框,必然偏移。padding 在右边和下边,所以框都往右下飘。
解决:preprocess 阶段把 scale、dw、dh 三个值带出来,解码后逐框反算。代码里就是(x - dw) / scale这一步,反算完再 clamp 到原图宽高范围内。
5.2 每帧都起 Python 进程:2 FPS 的瓶颈
现象:Java 端每次都 new 一个 ProcessBuilder 跑python detect.py,单帧耗时 2 秒以上,CPU 直接打满。
原因:Python 解释器启动要几百毫秒,onnxruntime 加载模型又要几百毫秒到一两秒,这还没算推理本身。每帧重来一次,时间全部花在启动上。
解决:改成常驻进程。模型加载只做一次,Socket 复用,Java 端只是发送 JPEG 和收 JSON。实测从 2 FPS 直接跳到 10 FPS 以上,瓶颈才回到真正的推理耗时上。
5.3 Java 与 Python 的半包和字节序:联调变黑匣子
现象:Python 单独跑一张图有结果,Java 发过去之后要么解析报错,要么返回空列表,而且不是每次必现。
原因:两个隐藏问题叠在一起。一是字节序,Java DataOutputStream.writeInt 默认大端,Python 端如果用了struct.unpack("<I")小端解析,长度值就完全错乱;二是半包,TCP 流可能一次只到一半数据,直接用 recv 一次拿不全。
解决:统一用大端>I,Java 端 readInt 天然对齐。Python 端用 while 循环收满 length 字节再解析,Java 端用in.readFully()。这两个坑都踩平之后,联调就不会再出玄学问题。
5.4 换 YOLOv8 后硬编码维度直接炸
现象:v5 跑得好好的,换成 yolov8.onnx 之后返回全空,或者 Java 端解析 JSON 时数组越界。
原因:v5 输出是 1x25200x85,v8 输出是 1x84x8400。如果解码代码里写死 25200 或者写死 84,换模型必炸。更隐蔽的是自定义数据集,v8 训练 10 个类别时输出是 1x14x8400,写死 84 一样出错。
解决:解码函数里动态推导。pred.shape[1] == 85走 v5/v7 分支,否则走 v8 分支;类别数从pred.shape[1] - 5或pred.shape[1] - 4推导,不写死。换模型时只改 onnx 路径,代码不用动。
5.5 内存只涨不降:Mat 没 release 的连锁反应
现象:服务跑一晚上,Java 进程内存从 500MB 涨到 3GB,最后 OOM 被杀。
原因:OpenCV 的 Java 绑定不像 Java 对象那样自动回收,Mat 底层是 native 内存。每帧取出来 new 一个 Mat,用完不 release,native 内存就持续累积。Python 端如果请求积压,队列里的 numpy 数组也会跟着堆积。
解决:Mat 用完放 finally 里bgr.release();JavaCV 的 grabber 在 finally 里调用 stop 和 close。Python 端 Socket 接收缓冲区设置上限,检测请求跟不上就主动丢帧,不堆队列。定位内存问题时用jcmd GC.class_histogram看 native 内存,别只看堆。
6. 阈值、线程数与 int8 量化:让这套检测管线持续产出的验证清单
6.1 参数表:conf、iou、线程与抽帧
| 参数 | 建议值 | 说明 |
|---|---|---|
| conf_thres | 0.25 | 误报多时调到 0.4,框数量明显减少 |
| iou_thres | 0.45 | 密集遮挡场景调 0.3,抑制重叠框 |
| intra_op_num_threads | CPU 核数一半 | 留一半线程给取帧和业务 |
| 抽帧间隔 | 单帧耗时的 1.2 倍 | 防止请求在 Python 端排队 |
| JPEG 质量 | 85 | 质量再高对检测结果无提升 |
线程数不要盲目拉满。onnxruntime 的 intra_op_num_threads 设成 CPU 核数时,每帧推理快了,但 Java 端取帧线程会被挤到等 CPU,整体吞吐反而下降。留一半核给业务线程,是多次压测后的折中。
6.2 int8 量化与联调一致性验证
跑通之后如果 CPU 资源紧张,可以给模型做 int8 量化:
from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic("yolov8n.onnx", "yolov8n_int8.onnx", weight_type=QuantType.QUInt8)量化后的模型体积缩小到四分之一左右,CPU 推理能快 20% 到 40%。代价是精度下降,小目标漏检率会上升。所以原模型文件先备份,量化模型单独存,上线前用同一段视频对比两版输出。这也是我给自己留的后悔药:Java 端不用改任何代码,Python 服务换个 onnx 路径重启就行。
验证联调一致性,最直接的办法是让 Python 单独跑一遍视频,Java 联调再跑一遍,比对同帧的框:
def iou(a, b): x1 = max(a["x1"], b["x1"]); y1 = max(a["y1"], b["y1"]) x2 = min(a["x2"], b["x2"]); y2 = min(a["y2"], b["y2"]) inter = max(0, x2 - x1) * max(0, y2 - y1) area_a = (a["x2"] - a["x1"]) * (a["y2"] - a["y1"]) area_b = (b["x2"] - b["x1"]) * (b["y2"] - b["y1"]) return inter / (area_a + area_b - inter + 1e-6) def match_rate(dets_py, dets_java): matched = sum(any(iou(a, b) > 0.5 for b in dets_java) for a in dets_py) return matched / max(len(dets_py), 1)匹配率 95% 以上就算通过,百分之几的差异来自浮点精度和 JPEG 重编码,不影响业务判断。这套方案我前后接过三个项目,从 v5 迁到 v8 都是只换 onnx 文件,Java 侧代码始终没动过。每个项目最后我都会留一段测试视频和这个比对脚本,下次联调直接跑,省掉大量反复确认的黑匣子时间。希望帮到你。
本文还有配套的精品资源,点击获取