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

PyTorch CNN模型每次推理结果不一致问题求助

PyTorch CNN推理同一批图像结果不一致问题排查

问题描述

我有一个训练好的PyTorch CNN模型,用于对7种不同的block进行分类。当对同一批图像执行推理时,每次运行都会得到不同的结果,示例如下:

=====================
Run 1
pred    truth   match
=====================
block_1 block_0 no
block_1 block_1 yes
block_1 block_1 yes
block_1 block_2 no
block_1 block_2 no
block_1 block_2 no
block_1 block_3 no
block_6 block_3 no
block_6 block_3 no
block_1 block_4 no
block_1 block_4 no
block_1 block_5 no
block_1 block_5 no
block_1 block_6 no
block_1 block_6 no

=====================
Run 2 
pred    truth   match
=====================
block_3 block_0 no
block_3 block_1 no
block_3 block_1 no
block_3 block_2 no
block_3 block_2 no
block_6 block_2 no
block_3 block_3 yes
block_3 block_3 yes
block_3 block_3 yes
block_3 block_4 no
block_3 block_4 no
block_3 block_5 no
block_3 block_5 no
block_3 block_6 no
block_3 block_6 no

相关代码:

模型加载

def load_model(model, model_dir='models', model_file_name='blocks.pt'):
    model_path = os.path.join(model_dir, model_file_name)
    model.load_state_dict(torch.load(model_path), strict=False)
    return model

预测函数

def prediction(model, device, batch_input):
    model.to(device)
    model.eval()

    data = batch_input.to(device)
    output = model(data)

    prob = F.softmax(output, dim=1)
    pred_prob = prob.data.max(dim=1)[0]
    pred_index = prob.data.max(dim=1)[1]

    return pred_index.cpu().numpy(), pred_prob.cpu().numpy()

图像变换(与训练时一致)

def image_common_transforms(mean=(0.4515, 0.3976, 0.3339), std=(0.3639, 0.3361, 0.3224)):
    common_transforms = transforms.Compose([
        transforms.Resize(256),
        transforms.CenterCrop(224),
        transforms.ToTensor(),
        transforms.Normalize(mean, std)
    ])
    return common_transforms

我原本以为model.eval()会阻止所有随机行为,是不是哪里理解错了?


问题原因及解决方法

1. strict=False导致模型参数加载不完整

这是最可能的原因:你在加载模型时使用了strict=False,这个参数会跳过模型定义与权重文件中不匹配的参数键。如果你的模型包含BatchNorm层,它依赖running_mean和running_var这两个训练阶段积累的统计量做归一化;如果这些统计量没有被正确加载(比如权重文件中没有对应键,或者模型结构有变更),每次启动时这些值会用随机初始化的默认值,即使调用了model.eval(),推理时也会用这些随机统计量,导致结果波动。

解决:

  • 优先改为strict=True加载模型,确保所有参数(包括BatchNorm的统计量)都被正确加载:
    model.load_state_dict(torch.load(model_path), strict=True)
    
  • 如果必须用strict=False,需要手动核对未加载的参数,确保BatchNorm的running_mean和running_var能从权重文件中读取。

2. 优化预测函数的模式控制

虽然你调用了model.eval(),但可以进一步优化:

  • 用torch.no_grad()包裹推理过程,关闭梯度计算(避免不必要的资源消耗,同时确保模型处于纯推理状态);
  • 避免每次预测都重复移动模型到设备(多次移动可能导致微小精度变化),建议加载模型后一次性移到目标设备。

修改后的预测函数:

# 加载模型后先移到目标设备
model = load_model(model)
model = model.to(device)

def prediction(model, device, batch_input):
    model.eval()  # 固定为评估模式
    with torch.no_grad():
        data = batch_input.to(device)
        output = model(data)
        prob = F.softmax(output, dim=1)
        pred_prob, pred_index = prob.max(dim=1)  # 简化写法
    return pred_index.cpu().numpy(), pred_prob.cpu().numpy()

3. 排查其他随机因素

  • 检查模型中是否有自定义随机层(比如自定义Dropout变体),这类层可能不受model.eval()控制;
  • 确认数据加载流程没有随机操作(你的图像变换都是确定性的,这部分没问题);
  • 确保没有其他代码在预测前后调用model.train(),覆盖评估模式。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 02:29:56