You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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的重载方法加载:

  1. 用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")
  1. 用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框架的封装元数据,不属于模型权重),有两种可靠方式获取:

  1. 提前导出类别名称文件
    在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"
    }
}
  1. 使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.19 03:10:25