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

PyTorch FasterRCNN模型转CoreML遇包装报错,求解决方案

解决FasterRCNN转CoreML时输出列表导致的索引错误问题

问题原因

PyTorch官方的FasterRCNN模型在eval模式下,输入图片(即使是单张)会返回列表结构,列表中的每个元素是对应单张图片的检测结果字典(包含boxes、labels、scores等键)。而CoreML追踪转换要求输出是可直接通过键访问的字典结构,直接使用原模型输出会触发TypeError: list indices must be integers or slices, not str错误。

解决方案:修改模型包装类,调整输出结构

通过自定义包装类,将原模型的列表输出转换为CoreML可识别的字典结构,以下是两种常见场景的实现:

场景1:单张图片输入(最稳妥,适配可变数量检测框)

import torch
import coremltools as ct

class FasterRCNNWrapper(torch.nn.Module):
    def __init__(self, base_model):
        super().__init__()
        self.base_model = base_model
        self.base_model.eval()  # 确保模型处于评估模式
    
    def forward(self, x):
        # 提取列表中第一个元素(对应单张图片的检测结果字典)
        detection_result = self.base_model(x)[0]
        # 返回标准字典结构,CoreML可直接识别键
        return {
            "boxes": detection_result["boxes"],
            "labels": detection_result["labels"],
            "scores": detection_result["scores"]
        }

# 加载你的FasterRCNN模型
original_model = torch.load("your_model.pt")
# 初始化包装模型
wrapped_model = FasterRCNNWrapper(original_model)

# 构造示例输入(格式:[batch_size, channels, height, width])
example_input = torch.rand(1, 3, 224, 224)

# 用追踪模式转换模型
traced_model = torch.jit.trace(wrapped_model, example_input)
coreml_model = ct.convert(
    traced_model,
    inputs=[ct.ImageType(name="input_image", shape=example_input.shape, scale=1/255.0)]
)

# 保存转换后的CoreML模型
coreml_model.save("faster_rcnn_detector.mlmodel")

场景2:批量图片输入(需保证每张图片检测框数量一致)

如果需要处理批量图片,可将所有图片的对应检测结果堆叠为张量:

def forward(self, x):
    batch_results = self.base_model(x)
    # 堆叠所有图片的boxes、labels、scores
    boxes = torch.stack([res["boxes"] for res in batch_results])
    labels = torch.stack([res["labels"] for res in batch_results])
    scores = torch.stack([res["scores"] for res in batch_results])
    return {"boxes": boxes, "labels": labels, "scores": scores}

验证转换结果

转换完成后,可通过CoreMLTools测试模型输出:

# 生成测试输入
test_input = torch.rand(1, 3, 224, 224).numpy()
# 预测
predictions = coreml_model.predict({"input_image": test_input})
# 检查输出键
print(predictions.keys())  # 应输出 dict_keys(['boxes', 'labels', 'scores'])

内容的提问来源于stack exchange,提问作者anandyn02

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 03:29:51