如何通过Torch C++ API获取YOLOv5导出TorchScript模型的输入图像尺寸
获取YOLOv5 TorchScript模型的预期输入尺寸(C++ API实现)
以下是几种无需解压模型文件、直接通过Torch C++ API读取输入尺寸的规范方法:
方法1:直接读取模型内置的imgsz属性
YOLOv5导出TorchScript时,会将--imgsz指定的尺寸作为模型属性存储。可以通过Module的attr方法直接获取:
// 加载TorchScript模型 torch::jit::script::Module yolov5_model = torch::jit::load("your_yolov5_model.ts"); // 读取imgsz属性并转换为张量 torch::Tensor imgsz_tensor = yolov5_model.attr("imgsz").toTensor(); // 提取尺寸值(通常为[height, width]或[width, height],对应导出时的设置) std::vector<int64_t> input_size = imgsz_tensor.sizes().vec(); // 打印结果示例 std::cout << "Expected input size: " << input_size[0] << " × " << input_size[1] << std::endl;
方法2:访问模型的extra子模块读取配置
导出的模型会包含extra子模块,对应原模型的配置信息(即你提到的extra/config.txt内容),可直接通过C++ API访问:
// 获取extra子模块 torch::jit::script::Module extra_module = yolov5_model.attr("extra").toModule(); // 从extra模块中读取imgsz torch::Tensor imgsz_tensor = extra_module.attr("imgsz").toTensor(); std::vector<int64_t> input_size = imgsz_tensor.sizes().vec();
这种方法完全对应原配置文件的内容,且依赖TorchScript的模块结构保证稳定性,比解压文件更规范。
方法3:遍历模型属性排查(适配不同YOLOv5版本)
如果遇到属性名差异(不同YOLOv5版本可能命名不同),可以遍历模型的所有属性查找尺寸相关字段:
// 遍历模型所有属性 for (const auto& attr : yolov5_model.named_attributes()) { std::cout << "Attribute name: " << attr.name << std::endl; // 找到类似"imgsz"、"input_size"或存储尺寸张量的属性 }
内容的提问来源于stack exchange,提问作者pasbi
相关产品推荐
相关产品推荐

