You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何加载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()

关键注意事项

  1. 模型定义必须一致:导入的SegFormer类要和训练时的代码完全相同,包括层数、通道数、类别数等所有参数。
  2. 预处理严格匹配:Resize尺寸、Normalize的均值方差必须和训练阶段完全一致,否则会导致预测结果失效。
  3. 权重键处理:如果上述键替换仍无法解决问题,可以打印state_dict.keys()和model.state_dict().keys()对比差异,针对性调整键名。

内容的提问来源于stack exchange,提问作者Jesper Andersen

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.14 02:46:03