如何正确加载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
相关产品推荐
相关产品推荐

