YOLOv5自定义模型运行报错,是否为训练问题?
问题分析与解决方案
核心问题
你在Colab(Linux/Posix环境)训练的best.pt模型,保存时包含了PosixPath类型的路径对象,而Windows系统使用的是WindowsPath,加载时无法实例化PosixPath导致报错。官方预训练的yolov5s.pt没有保存这类系统相关的路径信息,所以能正常加载。
解决方案
方案1:在训练环境(Colab)中修复模型后再下载
在Colab中运行以下代码,将模型里的PosixPath转为字符串后重新保存:
import torch def convert_posix_to_str(obj): """递归将所有PosixPath对象转为字符串""" if isinstance(obj, dict): return {k: convert_posix_to_str(v) for k, v in obj.items()} elif isinstance(obj, list): return [convert_posix_to_str(i) for i in obj] elif hasattr(obj, "__class__") and obj.__class__.__name__ == "PosixPath": return str(obj) else: return obj # 加载训练好的模型 ckpt = torch.load("/content/yolov5/runs/train/exp/weights/best.pt") # 转换路径类型 ckpt = convert_posix_to_str(ckpt) # 保存修复后的模型 torch.save(ckpt, "/content/yolov5/runs/train/exp/weights/best_fixed.pt")
下载best_fixed.pt到本地,替换原模型路径即可正常加载。
方案2:本地临时兼容修复(应急用)
如果已经下载了best.pt,可以在Windows本地推理脚本最顶部添加以下代码,临时兼容PosixPath:
import pathlib # 临时添加PosixPath兼容类,让Windows能解析 class PosixPath(pathlib.PurePosixPath): def __new__(cls, *args, **kwargs): return pathlib.Path(*args, **kwargs) pathlib.PosixPath = PosixPath
之后再执行原加载逻辑即可,这种方法属于临时hack,建议优先用方案1。
方案3:导出为跨平台格式(如ONNX)
在Colab中训练完成后,将模型导出为ONNX格式,跨平台兼容性更好:
!python /content/yolov5/export.py --weights /content/yolov5/runs/train/exp/weights/best.pt --include onnx
下载best.onnx后,本地用以下代码加载推理:
import torch import cv2 from models.common import DetectMultiBackend from utils.general import non_max_suppression, scale_boxes from utils.torch_utils import select_device # 加载ONNX模型 device = select_device('cpu') model = DetectMultiBackend('best.onnx', device=device) stride, names, pt = model.stride, model.names, model.pt imgsz = (640, 640) # 对应训练时的--img参数 model.warmup(imgsz=(1, 3, *imgsz)) # 模型预热 cap = cv2.VideoCapture(0) while True: ret, frame = cap.read() if not ret: break # 图像预处理 img = cv2.resize(frame, imgsz) img = img.transpose((2, 0, 1))[::-1] # HWC转CHW,BGR转RGB img = torch.from_numpy(img).to(device) img = img.float() / 255.0 # 归一化到0-1 if len(img.shape) == 3: img = img[None] # 添加batch维度 # 推理 pred = model(img, augment=False, visualize=False) pred = non_max_suppression(pred, 0.25, 0.45, classes=None, agnostic=False) # 绘制检测结果 for i, det in enumerate(pred): if len(det): det[:, :4] = scale_boxes(img.shape[2:], det[:, :4], frame.shape).round() for *xyxy, conf, cls in reversed(det): label = f'{names[int(cls)]} {conf:.2f}' cv2.rectangle(frame, (int(xyxy[0]), int(xyxy[1])), (int(xyxy[2]), int(xyxy[3])), (0, 255, 0), 2) cv2.putText(frame, label, (int(xyxy[0]), int(xyxy[1])-10), cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0,255,0), 2) cv2.imshow('YOLOv5 Detection', frame) if cv2.waitKey(1) & 0xFF == ord('q'): break cap.release() cv2.destroyAllWindows()
内容的提问来源于stack exchange,提问作者iLuv
相关产品推荐
相关产品推荐

