ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

使用 MXNet Java Inference API 进行 SSD 多目标检测推理实战

使用 MXNet Java Inference API 进行 SSD 多目标检测推理实战 使用 MXNet Java Inference API 进行 SSD 多目标检测推理实战【免费下载链接】mxnetLightweight, Portable, Flexible Distributed/Mobile Deep Learning with Dynamic, Mutation-aware Dataflow Dep Scheduler; for Python, R, Julia, Scala, Go, Javascript and more项目地址: https://gitcode.com/gh_mirrors/mxne/mxnet本教程演示如何基于 Apache MXNet 的 Java Inference API加载一个在 Pascal VOC 2012 数据集上预训练好的 Single Shot DetectorSSD目标检测模型并对其执行端到端的图像推理。通过本教程你将掌握在 Java 工程中完成模型与数据下载、输入形状与 DataDesc 声明、CPU/GPU 上下文选择、推理调用以及检测结果坐标解析的完整技术链路并可以直接运行一个可复现的目标检测 Java 程序。背景Java Inference API 与 SSD 模型MXNet 官方为 Java 提供了一套高效、易用的推理 API它是 Scala Infer API 的扩展提供模型加载与推理能力目标是让已有模型能够快速部署到 Java 生态的生产环境中详见 Java Guide。本教程使用的推理入口ObjectDetector正是该 API 面向目标检测场景封装的高层组件。本教程使用的模型是ResNet50 SSD结构以 ResNet50 作为骨干网络提取图像特征SSD 检测头在多个尺度的特征图上进行密集目标框回归与分类。模型在Pascal VOC 2012数据集上训练可检测以下 20 个类别aeroplane、bicycle、bird、boat、bottle、bus、car、cat、chair、cow、diningtable、dog、horse、motorbike、person、pottedplant、sheep、sofa、train、tvmonitor从仓库变更记录看MXNet 对 Java 推理生态持续投入例如新增 Java Image APINEWS.md 中记录的 MXNET-1180、Scala/Java 绘制边界框的 Image APIMXNET-1285以及 Java Predictor 与 Object Detector API 的单元测试MXNET-1263说明ObjectDetector是一个经过测试验证的正式 API 组件。前置条件与资源准备完成本教程需要IntelliJ IDEA 中的 MXNet Java 环境可选建议先完成 Java 环境的搭建确保工程能够依赖org.apache.mxnet相关构件。wget 工具用于下载模型产物与测试图片。SSD 模型产物包含三部分文件——符号图文件-symbol.json、权重参数文件-0000.params以及类别标签文件synset.txt。测试图片用于推理的样例图片。下载模型产物在终端中执行以下脚本将模型文件下载到/tmp/resnet50_ssd目录data_path/tmp/resnet50_ssd mkdir -p $data_path wget https://s3.amazonaws.com/model-server/models/resnet50_ssd/resnet50_ssd_model-symbol.json -P $data_path wget https://s3.amazonaws.com/model-server/models/resnet50_ssd/resnet50_ssd_model-0000.params -P $data_path wget https://s3.amazonaws.com/model-server/models/resnet50_ssd/synset.txt -P $data_path下载测试图片image_path/tmp/resnet50_ssd/images mkdir -p $image_path cd $image_path wget https://cloud.githubusercontent.com/assets/3307514/20012567/cbb60336-a27d-11e6-93ff-cbc3f09f5c9e.jpg -O dog.jpg wget https://cloud.githubusercontent.com/assets/3307514/20012563/cbb41382-a27d-11e6-92a9-18dab4fd1ad3.jpg -O person.jpgdog.jpg用于复现本教程结尾的推理结果person.jpg可作为额外的验证样本检测person类别是否被正确识别。编写 ObjectDetectionTutorial.java按照 IntelliJ IDEA 中 MXNet Java 的工程搭建方式在同一个JavaMXNet工程中新建一个空类ObjectDetectionTutorial.java。下面按步骤拆解代码的各个组成部分。1. 定义模型与图片路径在main函数中声明模型文件前缀与输入图片路径。注意modelPathPrefix是不含扩展名的前缀ObjectDetector会根据该前缀自动拼接-symbol.json与-0000.paramsString modelPathPrefix /tmp/resnet50_ssd/resnet50_ssd_model; String inputImagePath /tmp/resnet50_ssd/images/dog.jpg;2. 选择运行上下文CPU / GPU推理可在 CPU 或 GPU若有 GPU 机器上运行通过向ListContext中添加对应上下文实现。下面的示例选择 CPUprivate static ListContext getContext() { ListContext ctx new ArrayList(); ctx.add(Context.cpu()); // Choosing CPU Context here return ctx; }若机器配备 GPU可将Context.cpu()替换为Context.gpu()指定设备号MXNet 的 Java Inference API 会自动将计算调度到对应设备。3. 声明模型输入形状与 DataDescSSD 模型的输入是一个四维张量通过Shape与DataDesc描述Shape inputShape new Shape(new int[] {1, 3, 512, 512}); ListDataDesc inputDescriptors new ArrayListDataDesc(); inputDescriptors.add(new DataDesc(data, inputShape, DType.Float32(), NCHW));输入形状可以这样解读batch size 为 1一次推理一张图、3 个 RGB 通道、图像高与宽均为 512。DataDesc的四个参数分别是输入名data与模型符号图中的输入节点名对应、形状、数据类型Float32以及数据排布NCHW通道在前符合 MXNet 默认布局。排布顺序必须与训练时一致否则预处理与推理结果会出错。4. 加载图片并执行推理BufferedImage img ObjectDetector.loadImageFromFile(inputImagePath); ObjectDetector objDet new ObjectDetector(modelPathPrefix, inputDescriptors, context, 0); ListListObjectDetectorOutput output objDet.imageObjectDetect(img, 3);这里发生了三件关键的事ObjectDetector.loadImageFromFile(...)静态方法从文件系统加载图片为BufferedImage构造ObjectDetector时传入模型前缀、输入描述符、上下文列表以及阈值参数0可将其理解为过滤低置信度检测框的最低分数数值越小保留的候选框越多imageObjectDetect(img, 3)表示返回置信度最高的前 3 个检测对象返回结构为嵌套列表外层对应输入图像内层是该图像上检测到的对象集合每个ObjectDetectorOutput包含类别名、概率与归一化坐标。5. 输出结果解析完整的类代码如下package mxnet; import org.apache.mxnet.infer.javaapi.ObjectDetector; import org.apache.mxnet.infer.javaapi.ObjectDetectorOutput; import org.apache.mxnet.javaapi.Context; import org.apache.mxnet.javaapi.DType; import org.apache.mxnet.javaapi.DataDesc; import org.apache.mxnet.javaapi.Shape; import java.awt.image.BufferedImage; import java.util.ArrayList; import java.util.Arrays; import java.util.List; public class ObjectDetectionTutorial { public static void main(String[] args) { String modelPathPrefix /tmp/resnet50_ssd/resnet50_ssd_model; String inputImagePath /tmp/resnet50_ssd/images/dog.jpg; ListContext context getContext(); Shape inputShape new Shape(new int[] {1, 3, 512, 512}); ListDataDesc inputDescriptors new ArrayListDataDesc(); inputDescriptors.add(new DataDesc(data, inputShape, DType.Float32(), NCHW)); BufferedImage img ObjectDetector.loadImageFromFile(inputImagePath); ObjectDetector objDet new ObjectDetector(modelPathPrefix, inputDescriptors, context, 0); ListListObjectDetectorOutput output objDet.imageObjectDetect(img, 3); printOutput(output, inputShape); } private static ListContext getContext() { ListContext ctx new ArrayList(); ctx.add(Context.cpu()); return ctx; } private static void printOutput(ListListObjectDetectorOutput output, Shape inputShape) { StringBuilder outputStr new StringBuilder(); int width inputShape.get(3); int height inputShape.get(2); for (ListObjectDetectorOutput ele : output) { for (ObjectDetectorOutput i : ele) { outputStr.append(Class: i.getClassName() \n); outputStr.append(Probabilties: i.getProbability() \n); ListFloat coord Arrays.asList(i.getXMin() * width, i.getXMax() * height, i.getYMin() * width, i.getYMax() * height); StringBuilder sb new StringBuilder(); for (float c: coord) { sb.append(, ).append(c); } outputStr.append(Coord: sb.substring(2) \n); } } System.out.println(outputStr); } }在结果解析部分ObjectDetectorOutput返回的坐标是归一化的取值在 0~1 之间因此乘以输入图像的宽、高这里均为 512即可换算回像素坐标。需要说明的是示例代码中坐标换算为xMin*width, xMax*height, yMin*width, yMax*height从代码结构看后三项的宽高引用存在错位xMax/yMin应分别乘以宽、高由于本示例宽高恰好都为 512所以不会影响最终数值若你的模型输入宽高不等建议按xMin*width, xMax*width, yMin*height, yMax*height进行修正。编译与运行在工程根目录下执行 Maven 构建同时拉取运行期依赖mvn clean install dependency:copy-dependencies构建完成后target目录下会生成javaMXNet-1.0-SNAPSHOT.jar。随后从工程根目录运行主类java -cp target/javaMXNet-1.0-SNAPSHOT.jar:target/dependency/* mxnet.ObjectDetectionTutorial注意-cp中同时包含了生成的 jar 与dependency/*下的全部依赖MXNet Java 运行时、Scala 与原生库桥接等缺一不可。预期推理输出对dog.jpg执行上述程序可得到类似如下的输出Class: car Probabilties: 0.99847263 Coord:312.21335, 72.02908, 456.01443, 150.66176 Class: bicycle Probabilties: 0.9047381 Coord:155.9581, 149.96365, 383.83694, 418.94516 Class: dog Probabilties: 0.82268167 Coord:Coord:83.82356, 179.14001, 206.63783, 476.78754模型在样例图片中同时检出了car置信度约 0.998、bicycle约 0.905与dog约 0.823每个结果附带一组坐标对应原图中该对象所在区域的边界框左上角与右下角像素位置。imageObjectDetect(img, 3)的返回顺序即按置信度从高到低排列因此在多目标场景下靠前的检测通常更可靠。从这些坐标可以看到推理 API 的输出语义类别 概率 边界框这正好对应 SSD 模型多尺度密集检测 NMS 后处理的输出形式Java API 已把符号图推理、后处理与坐标归一化封装在ObjectDetector内部使用者只需关心业务层的结果消费。进一步探索若想深入了解 MXNet Java 推理 API 的整体能力与 Javadoc可参考 Java Guide 与 Java 教程索引本文所用ObjectDetector与 Scala 版 Infer API 同源可对照阅读 Scala Inference 教程 理解底层Predictor、DataDesc与Shape的对应关系仓库变更记录NEWS.md中关于 Java Image APIMXNET-1180、Scala/Java 绘制边界框MXNET-1285与 Object Detector 单元测试MXNET-1263的条目可作为深入理解该 API 实现与演进历史的线索。【免费下载链接】mxnetLightweight, Portable, Flexible Distributed/Mobile Deep Learning with Dynamic, Mutation-aware Dataflow Dep Scheduler; for Python, R, Julia, Scala, Go, Javascript and more项目地址: https://gitcode.com/gh_mirrors/mxne/mxnet创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表