如何加载SegFormer PyTorch checkpoint并可视化预测结果?
SegFormer模型加载与预测最简实现方案
核心问题解决思路
加载checkpoint时的state_dict键不匹配,大多是因为训练时用了PyTorch Lightning等框架,保存的权重键带有model.前缀,或是模型结构与训练时存在细微差异。下面是适配该场景的最简可运行代码:
完整代码实现
1. 导入依赖库
import torch import torchvision.transforms as transforms from PIL import Image import matplotlib.pyplot as plt # 注意:这里要导入你训练时用的SegFormer模型定义类 # 如果是Roboflow官方代码,可能是从他们的封装模块导入,或是你自定义的SegFormer类 from your_model_module import SegFormer
2. 模型初始化与权重加载
# 初始化模型:参数必须和训练时完全一致(类别数、输入尺寸等) num_classes = 2 # 替换为你的数据集类别总数 model = SegFormer(num_classes=num_classes) # 加载checkpoint并处理键不匹配问题 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') checkpoint = torch.load('best.ckpt', map_location=device) # 处理权重键:移除可能存在的"model."前缀 if 'model' in checkpoint: state_dict = {k.replace('model.', ''): v for k, v in checkpoint['model'].items()} elif 'state_dict' in checkpoint: state_dict = {k.replace('model.', ''): v for k, v in checkpoint['state_dict'].items()} else: state_dict = checkpoint # 加载权重,strict=False允许部分非核心键不匹配 model.load_state_dict(state_dict, strict=False) model.eval() model.to(device)
3. 图片预处理(与训练时保持一致)
# 完全复用训练时的transform配置,示例如下 transform = transforms.Compose([ transforms.Resize((512, 512)), # 替换为训练时的输入尺寸 transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])
4. 预测与结果转换
# 加载测试图片 img_path = "test_image.jpg" # 替换为你的测试图片路径 original_img = Image.open(img_path).convert('RGB') img_tensor = transform(original_img).unsqueeze(0).to(device) # 添加batch维度 # 执行预测 with torch.no_grad(): outputs = model(img_tensor) # 将logits转为类别掩码 mask = torch.argmax(outputs, dim=1).squeeze(0).cpu().numpy()
5. 掩码叠加显示
# 定义类别对应颜色(根据你的类别数调整) class_colors = [(0, 0, 0), (255, 0, 0)] # 示例:背景黑色、目标红色 # 生成彩色掩码图 mask_img = Image.new('RGB', mask.shape[::-1]) pixels = mask_img.load() for y in range(mask.shape[0]): for x in range(mask.shape[1]): pixels[x, y] = class_colors[mask[y, x]] # 叠加掩码到原图(透明度50%) resized_original = original_img.resize((512, 512)) overlay_img = Image.blend(resized_original, mask_img, alpha=0.5) # 显示结果 plt.figure(figsize=(15, 5)) plt.subplot(131) plt.imshow(original_img) plt.title('原图') plt.axis('off') plt.subplot(132) plt.imshow(mask_img) plt.title('分割掩码') plt.axis('off') plt.subplot(133) plt.imshow(overlay_img) plt.title('叠加效果') plt.axis('off') plt.show()
关键注意事项
- 模型定义必须一致:导入的
SegFormer类要和训练时的代码完全相同,包括层数、通道数、类别数等所有参数。 - 预处理严格匹配:Resize尺寸、Normalize的均值方差必须和训练阶段完全一致,否则会导致预测结果失效。
- 权重键处理:如果上述键替换仍无法解决问题,可以打印
state_dict.keys()和model.state_dict().keys()对比差异,针对性调整键名。
内容的提问来源于stack exchange,提问作者Jesper Andersen
相关产品推荐
相关产品推荐

