如何在笔记本电脑上实时部署.pth模型并显示病害类型与严重程度
解决实时部署ResNet50模型并展示病害类型与严重程度的方案
一、先明确类别映射与模型输出逻辑
首先要确认训练时的类别定义:
- 如果病害类型和严重程度是合并为单个类别(比如
锈病-轻度、炭疽病-重度),需要把预测的类别ID对应到包含两类信息的文本描述; - 如果模型是双分支输出(分别预测病害类型、严重程度),则要调整推理代码,同时获取两个分支的输出结果。
logs.json的作用:这个文件完全有用,可以从中提取训练时的类别名称映射、预处理参数(如归一化均值/标准差)、模型收敛情况,能帮你快速对齐部署与训练的配置,避免因参数不一致导致的预测问题。
二、修改实时画面展示代码
假设你用OpenCV实现实时捕捉,以下是具体修改方案:
基础场景(合并类别)
如果你的模型输出是合并的类别标签,先定义类别解析函数,再修改画面绘制逻辑:
import cv2 import torch from torchvision import transforms # 加载模型(适配笔记本CPU/GPU) model = torch.load('resnet50_model.pth', map_location=torch.device('cpu')) model.eval() # 从logs.json或训练时的标签文件提取类别列表 class_list = ["健康", "锈病-轻度", "锈病-中度", "锈病-重度", "炭疽病-轻度", "炭疽病-重度"] # 图像预处理(必须与训练时一致,可从logs.json提取参数) transform = transforms.Compose([ transforms.ToPILImage(), transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) # 解析类别为病害类型+严重程度 def parse_disease_info(class_name): if class_name == "健康": return "病害类型: 无", "严重程度: 无" disease, severity = class_name.split('-') return f"病害类型: {disease}", f"严重程度: {severity}" cap = cv2.VideoCapture(0) while cap.isOpened(): ret, frame = cap.read() if not ret: break # 推理预测 input_img = transform(frame).unsqueeze(0) with torch.no_grad(): outputs = model(input_img) _, pred_idx = torch.max(outputs, 1) class_name = class_list[pred_idx.item()] # 解析并绘制信息 disease_text, severity_text = parse_disease_info(class_name) cv2.putText(frame, disease_text, (10, 30), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 255, 0), 2) cv2.putText(frame, severity_text, (10, 70), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 0, 255), 2) cv2.imshow('实时病害检测', frame) if cv2.waitKey(1) & 0xFF == ord('q'): break cap.release() cv2.destroyAllWindows()
双分支模型场景
如果模型是双分支输出(分别预测类型和严重程度),修改推理部分即可:
# 替换推理代码块 with torch.no_grad(): # 假设模型返回两个分支的输出 type_output, severity_output = model(input_img) _, type_idx = torch.max(type_output, 1) _, severity_idx = torch.max(severity_output, 1) disease_type = type_class_list[type_idx.item()] severity_level = severity_class_list[severity_idx.item()] # 绘制时直接使用两个变量 cv2.putText(frame, f"病害类型: {disease_type}", (10, 30), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 255, 0), 2) cv2.putText(frame, f"严重程度: {severity_level}", (10, 70), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 0, 255), 2)
三、模型加载失败的排查要点
如果之前加载模型报错,优先检查:
- 确保部署用的框架版本(如PyTorch)与训练时一致;
- 训练时用多GPU的话,加载模型需指定
map_location='cpu'; - 尝试用
load_state_dict加载权重,避免模型结构不匹配:from torchvision.models import resnet50 model = resnet50(num_classes=你的类别总数) model.load_state_dict(torch.load('resnet50_weights.pth', map_location='cpu')) model.eval()
内容的提问来源于stack exchange,提问作者Bon Journey
相关产品推荐
相关产品推荐

