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

如何使用PyTorch预训练.pth模型?以RESA车道检测模型为例

问题解决与方案

为什么调用model.eval()会报错?

你用torch.load加载的是模型的状态字典(state_dict)——这是一个保存了模型权重参数的字典,而非模型实例本身,所以字典对象自然没有eval()这类模型专属方法。

RESA模型的结构说明

RESA的模型定义在仓库的models/resa.py文件中,核心结构如下:

  • 以ResNet50作为骨干网络提取基础图像特征
  • 加入递归特征聚合模块(RESA),通过多轮递归操作聚合不同尺度的特征,强化车道线的特征表达能力
  • 最后连接车道检测头部,输出对应车道线的概率图(适配CULane数据集默认的4车道场景)

正确加载预训练模型的步骤

先初始化模型实例,再加载状态字典,代码示例:

from models.resa import RESA
import torch

# 初始化模型,参数需和预训练模型匹配(骨干网络、车道数等)
model = RESA(backbone='resnet50', num_classes=4)  # CULane默认是4车道,可按需调整

# 加载预训练的状态字典
state_dict = torch.load('culane_resnet50.pth', map_location=torch.device('cpu'))

# 若模型是多GPU训练保存的,需移除键名中的'module.'前缀
if list(state_dict.keys())[0].startswith('module.'):
    state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}

# 将权重加载到模型实例中
model.load_state_dict(state_dict)

# 现在可以正常设置模型为评估模式
model.eval()

在新数据集上做预测的可行方案

完全可以在新数据集上使用该模型,需完成以下步骤:

  1. 对齐数据预处理流程
    必须和CULane训练时的预处理逻辑一致:

    • 将图像resize到(288, 800)(仓库默认输入尺寸)
    • 用ImageNet的均值[0.485, 0.456, 0.406]和方差[0.229, 0.224, 0.225]做归一化
    • 将图像维度从HWC转换为PyTorch要求的CHW格式
  2. 适配新数据集的标注与加载

    • 若新数据集的标注格式和CULane不同(比如CULane是单车道点集,你的数据集是分割掩码),需要参考仓库data/culane.py中的CULane类,编写新数据集的加载逻辑,将标注转换为模型预期的格式(比如每个车道线的锚点概率标签)
  3. 执行预测与后处理
    用torch.no_grad()包裹预测过程避免无效梯度计算,再对输出结果做后处理:

    import cv2
    import numpy as np
    
    # 预处理示例函数
    def preprocess(img_path):
        img = cv2.imread(img_path)
        img = cv2.resize(img, (800, 288))
        img = img / 255.0
        img = (img - np.array([0.485, 0.456, 0.406])) / np.array([0.229, 0.224, 0.225])
        img = img.transpose(2, 0, 1)
        return torch.tensor(img, dtype=torch.float32).unsqueeze(0)
    
    # 加载测试图像并预处理
    input_tensor = preprocess('test_image.jpg')
    
    # 执行预测
    with torch.no_grad():
        output = model(input_tensor)
    
    # 后处理:从输出概率图提取车道线
    threshold = 0.5
    pred_lanes = []
    for lane_prob in output[0]:
        lane_mask = lane_prob > threshold
        # 可复用仓库utils/post_process.py中的逻辑,提取车道线坐标或做拟合
        pred_lanes.append(lane_mask.numpy())
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 12:15:39