如何解决PyTorch中输入类型与权重类型不匹配的RuntimeError问题
解决RuntimeError: Input type and weight type mismatch问题
从你的错误信息和代码来看,问题的核心是你的模型权重被意外移回了CPU,但输入数据仍在GPU上,导致两者设备不匹配。下面是具体的排查和解决步骤:
1. 移除导致模型回退到CPU的关键代码
在你的classify函数中有一行致命代码:
model.cpu()
这行代码会把整个模型的权重从GPU移到CPU,后续再用GPU上的输入img进行推理时,就会出现Input type (torch.cuda.FloatTensor) and weight type (torch.FloatTensor)的不匹配错误。直接删除这行代码即可解决最根本的问题。
修复后的classify函数片段:
def classify(image): img = process_image(image) print('Image processed') pred = model(img)[0] pred = non_max_suppression(pred, 0.4, 0.5, classes = [0, 1, 2], agnostic = None ) if classify: pred = apply_classifier(pred, modelc, img, im0s) print("1 ", pred) # 移除 model.cpu() 这一行 classification = torch.cat(pred)[:, -1] if len(classification) == 0: return None index = int(classification[0]) print(names[index]) return names[index]
2. 统一模型加载逻辑,避免重复初始化
你的detect函数中每次调用都会重新加载模型:
def detect(frame): # ... 其他代码 model = attempt_load(file_path, map_location = device) model.to(device).eval()
这种重复加载不仅浪费资源,还容易导致模型设备状态混乱(比如多次加载后可能出现部分层在CPU、部分在GPU的情况)。建议全局只加载一次模型,然后在函数间共享:
步骤1:在主代码开头全局加载YOLO模型
# 全局初始化YOLO模型 device = select_device() # 加载人脸检测用的YOLO模型 yolo_det_path = 'weights/yolov5s.pt' yolo_det_model = attempt_load(yolo_det_path, map_location=device) yolo_det_model.to(device).eval() yolo_det_names = yolo_det_model.module.names if hasattr(yolo_det_model, 'module') else yolo_det_model.names # 加载分类用的YOLO模型(mask.pt) yolo_cls_path = 'weights/mask.pt' yolo_cls_model = attempt_load(yolo_cls_path, map_location=device) yolo_cls_model.to(device).eval() yolo_cls_names = yolo_cls_model.module.names if hasattr(yolo_cls_model, 'module') else yolo_cls_model.names
步骤2:修改detect和classify函数,传入全局模型
def detect(frame, model, names): img = process_image(frame) pred = model(img)[0] pred = non_max_suppression(pred, 0.4, 0.5, classes = [0], agnostic = None) gn = torch.tensor(frame.shape)[[1,0,1,0]] for i, det in enumerate(pred): det[:,:4] = scale_coords(img.shape[2:], det[:,:4], frame.shape).round() for *xyxy, conf, cls in reversed(det): xywh = (xyxy2xywh(torch.tensor(xyxy).view(1, 4)) / gn).view(-1).tolist() label = '%s %.2f' % (names[int(cls)], conf) if label is not None and (label.split())[0] == 'person': plot_one_box(xyxy, frame, label=label, color=[75, 125, 2], line_thickness=1) def classify(image, model, names): img = process_image(image) print('Image processed') pred = model(img)[0] pred = non_max_suppression(pred, 0.4, 0.5, classes = [0, 1, 2], agnostic = None ) if classify: pred = apply_classifier(pred, modelc, img, im0s) print("1 ", pred) classification = torch.cat(pred)[:, -1] if len(classification) == 0: return None index = int(classification[0]) print(names[index]) return names[index]
步骤3:在主循环中调用修改后的函数
# 主循环中的调用部分 if nrof_faces > 0: detect(frame, yolo_det_model, yolo_det_names) label = classify(frame, yolo_cls_model, yolo_cls_names) # ... 后续逻辑
3. 验证设备一致性(可选但推荐)
如果还是不确定问题,可以在推理前加入设备检查代码,确保输入和模型在同一设备上:
# 在classify或detect函数中添加 print(f"Model device: {next(model.parameters()).device}") print(f"Input device: {img.device}")
如果输出不一致,手动强制把模型移到目标设备:
model = model.to(device)
额外注意事项
- 确保你的
process_image函数中,输入张量的类型和模型一致:你已经用img.float()转成了float32,和默认的模型权重类型匹配,这部分没问题。 - 如果你之前尝试过半精度推理(
model.half()),要确保输入也转成半精度:img = img.half(),但你现在注释了这部分代码,所以不需要处理。
内容的提问来源于stack exchange,提问作者harry r
相关产品推荐
相关产品推荐

