如何用Java从PyTorch模型文件中提取类别名称映射?
问题描述
我拿到了一个PyTorch模型文件和一些目标检测结果,检测结果仅提供识别目标的编号,但我需要从模型文件中获取对应的类别名称。
我找到的Python实现代码如下:
model = DetectMultiBackend(weights, device=device, dnn=dnn, data=data, fp16=half) stride, names, pt = model.stride, model.names, model.pt
我确定需要获取其中的names数组,但我使用的是Java而非Python。我查看了ai.djl.pytorch.engine.PtModel,但未找到类似的编号与名称映射。
甚至DeepJavaLibrary(DJL)似乎无法加载纯.pt文件,测试代码如下:
String fname = "/tmp/yolov5s.pt"; { PtEngine engine = (PtEngine) Engine.getEngine("PyTorch"); Model model = engine.newModel("bacon", null); model.load(new File(fname).toPath()); Block block = model.getBlock(); System.out.println(block); }
报错信息:
Exception in thread "main" ai.djl.engine.EngineException: PytorchStreamReader failed locating file constants.pkl: file not found at ai.djl.pytorch.jni.PyTorchLibrary.moduleLoad(Native Method) at ai.djl.pytorch.jni.JniUtils.loadModule(JniUtils.java:1550) at ai.djl.pytorch.engine.PtModel.load(PtModel.java:90) at ai.djl.Model.load(Model.java:110) at project.pictureServer.PyTorchFile.main(PyTorchFile.java:37)
请问使用Java和PyTorch模型文件实现编号到类别名称映射的正确方法是什么?
解决方案
一、解决DJL加载.pt文件的报错问题
你遇到的constants.pkl找不到错误,是因为DJL默认期望加载DJL格式导出的PyTorch模型,而非原生.pt文件(比如YOLOv5直接导出的模型)。要加载原生PyTorch模型,需先将模型转为TorchScript格式,再用DJL的重载方法加载:
- 用Python将模型转为TorchScript格式
import torch # 加载你的模型,这里以YOLOv5为例 model = torch.hub.load('ultralytics/yolov5', 'yolov5s', pretrained=True) model.eval() # 生成脚本化模型 traced_model = torch.jit.trace(model, torch.randn(1, 3, 640, 640)) traced_model.save("yolov5s_scripted.pt")
- 用DJL加载转换后的模型
import ai.djl.pytorch.engine.PtEngine; import ai.djl.Model; import ai.djl.ModelLoadOptions; import java.nio.file.Paths; public class ModelLoader { public static void main(String[] args) throws Exception { String fname = "/tmp/yolov5s_scripted.pt"; PtEngine engine = (PtEngine) Engine.getEngine("PyTorch"); Model model = engine.newModel("yolov5", null); ModelLoadOptions options = new ModelLoadOptions(); options.setMapLocation("cpu"); // 可根据设备调整为"cuda" model.load(Paths.get(fname), options); } }
二、获取类别名称映射
原生PyTorch模型本身不存储names数组(这是YOLO框架的封装元数据,不属于模型权重),有两种可靠方式获取:
- 提前导出类别名称文件
在Python中提取names数组并保存为JSON,再在Java中读取:
import json model = torch.hub.load('ultralytics/yolov5', 'yolov5s', pretrained=True) with open("class_names.json", "w") as f: json.dump(model.names, f)
Java读取示例:
import com.google.gson.Gson; import java.io.FileReader; public class ClassNameReader { public static void main(String[] args) throws Exception { Gson gson = new Gson(); String[] classNames = gson.fromJson(new FileReader("class_names.json"), String[].class); // 编号对应数组下标,比如编号0取classNames[0] System.out.println(classNames[0]); // 输出"person" } }
- 使用DJL的YOLOv5预定义模块
如果是YOLOv5系列模型,DJL的预训练封装内置了类别映射:
import ai.djl.Model; import ai.djl.repository.zoo.Criteria; import ai.djl.repository.zoo.ZooModel; import ai.djl.modality.cv.Image; import ai.djl.modality.cv.output.DetectedObjects; public class YoloClassName { public static void main(String[] args) throws Exception { Criteria<Image, DetectedObjects> criteria = Criteria.builder() .setTypes(Image.class, DetectedObjects.class) .optModelUrls("djl://ai.djl.pytorch/yolov5s") .optEngine("PyTorch") .build(); try (ZooModel<Image, DetectedObjects> model = criteria.loadModel()) { // 读取内置的类别文件 String[] classNames = model.getArtifact("classes.txt").readAllLines().toArray(new String[0]); System.out.println(classNames[0]); } } }
三、额外说明
- 若是自定义训练的YOLOv5模型,
names数组来自训练时的data.yaml文件,直接提取该文件中的类别列表转成JSON/文本,在Java中读取即可,无需从模型文件提取。 - DJL加载原生PyTorch模型时,必须确保模型是TorchScript格式,否则会加载失败。
内容的提问来源于stack exchange,提问作者Mutant Bob
相关产品推荐
相关产品推荐

