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

如何正确加载PyTorch的pth模型并与YOLOv5等其他模型结合使用

问题解决方法

你遇到的报错本质是CRAFT官方仓库没有适配PyTorch Hub入口,且公开的pth文件仅保存了权重参数,没有包含模型结构定义,按以下步骤操作即可解决:

1. 获取模型结构定义

将CRAFT官方仓库中的craft.py文件复制到你的本地项目目录下,该文件包含完整的CRAFT网络结构实现。

2. 初始化模型并加载权重

替换你原来加载OCR模型的代码为以下内容:

import torch
import cv2
import numpy as np
from craft import CRAFT

# 初始化模型结构
ocr_model = CRAFT()

# 适配运行设备
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# 读取权重文件
weight_dict = torch.load('runs/ocr/craft_mlt_25k.pth', map_location=device)

# 去除多卡训练保存的权重前缀module.
if list(weight_dict.keys())[0].startswith('module.'):
    weight_dict = {k[7:]: v for k, v in weight_dict.items()}

# 加载权重到模型
ocr_model.load_state_dict(weight_dict)
ocr_model.to(device)
ocr_model.eval() # 切换为推理模式

3. 补充CRAFT输入预处理逻辑

CRAFT对输入图像有标准化要求,不能直接传入原始裁剪图像,需要先做预处理:

def preprocess_craft_input(img, target_size=768):
    # 图像尺寸缩放
    h, w = img.shape[:2]
    scale = target_size / max(h, w)
    resized = cv2.resize(img, (int(w*scale), int(h*scale)), interpolation=cv2.INTER_LINEAR)
    # 通道转换+归一化
    img_norm = resized.astype(np.float32) / 255.
    mean = np.array([0.485, 0.456, 0.406])
    std = np.array([0.229, 0.224, 0.225])
    img_norm = (img_norm - mean) / std
    # 转CHW格式+加batch维度
    img_tensor = torch.from_numpy(img_norm).permute(2, 0, 1).unsqueeze(0)
    return img_tensor.to(device)

4. 完整串联推理代码

# 加载YOLOv5模型
model = torch.hub.load('.', 'custom', path='runs/train/exp2/weights/best.pt', source='local', force_reload=True)

cap = cv2.VideoCapture('../Dataset/test/09-10.mp4')

while cap.isOpened():
    ret, frame = cap.read()
    if not ret:
        break
    # YOLOv5检测数字区域
    results = model(frame)
    crops = results.crop(save=False)    
    for crop in crops:
        if 'number' in crop['label']:
            # 预处理输入
            input_tensor = preprocess_craft_input(crop['im'])
            # CRAFT推理
            with torch.no_grad():
                score_text, score_link = ocr_model(input_tensor)
            # 后续可根据输出的置信度图解析文本框坐标,再裁剪后送入文本识别模型得到数字内容

注意事项

  • CRAFT仅提供文本检测能力,只能定位数字所在的坐标区域,无法直接输出识别到的数字内容。如果需要得到具体的数字结果,需要额外接入CRNN等文本识别模型。
  • 推理过程中需要在torch.no_grad()上下文内运行OCR模型,避免显存溢出。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 21:48:03