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

PyTorch从Checkpoint加载GradScaler后执行scaler.step(optimizer)出现张量形状不匹配RuntimeError问题

解决PyTorch AMP中scaler.step(optimizer)的张量形状不匹配问题

首先,从错误栈可以明确:问题出在Adam优化器更新动量的核心步骤——exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1),这里exp_avg是优化器缓存的动量张量(本应和模型参数同形状),但它的形状(32)和当前参数的梯度张量形状(64)不匹配。本质原因是加载的优化器state_dict与当前模型的参数在形状/数量上不兼容,下面是具体的排查和解决步骤:

1. 检查优化器的初始化参数

你的代码里写的是optimizer = optim.Adam(parameters, lr, betas),这里的parameters必须严格对应model1.parameters()。如果传入的是其他参数集合,优化器的缓存张量会基于错误的参数生成,加载后必然和当前模型参数形状不匹配。

修复方式:
确保初始化optimizer时传入模型的参数:

optimizer = optim.Adam(model1.parameters(), lr=lr, betas=betas)

2. 确认模型结构完全一致

如果保存checkpoint时的model1结构和当前resume时的模型有差异(比如某个卷积层的通道数从32改成64、新增/删除了网络层),哪怕model.load_state_dict()没有报错(比如用了strict=False),也会导致优化器的缓存张量和当前模型参数形状不匹配。

排查方式:

  • 对比保存checkpoint时的模型代码和当前代码,确保所有层的维度、数量完全一致;
  • 加载模型后,打印参数形状验证:
# 加载模型后执行
for name, param in model1.named_parameters():
    print(f"{name}: {param.shape}")

3. 验证优化器state_dict与模型参数的匹配性

加载checkpoint后,检查优化器的参数组和模型参数数量是否一致:

# 加载checkpoint后执行
param_count_model = len(list(model1.parameters()))
param_count_optimizer = len(optimizer.state_dict()['param_groups'][0]['params'])
print(f"模型参数数量:{param_count_model},优化器缓存参数数量:{param_count_optimizer}")

如果两个数字不一致,说明优化器加载的state_dict和当前模型不匹配,需要重新用当前模型参数初始化优化器,或者重新保存正确的checkpoint。

4. 修正学习率调度器的step时机

你的代码在每个迭代(iter)都调用scheduler.step(),但LambdaLR是epoch级别的调度器,应该放在epoch循环的末尾,每个epoch调用一次。虽然这不是直接导致形状不匹配的原因,但错误的step时机可能引发其他训练异常。

修复方式:

for epoch in trange(epoch_resume, config['epochs']+1, desc='Epochs'):
    for content_image, style_image in tqdm(dataloader, desc='Dataloader'):
        # ... 训练逻辑 ...
        scaler.step(optimizer)
        scaler.update()
        optimizer.zero_grad()
    # 每个epoch结束后更新学习率
    scheduler.step()

5. 确保设备一致性

如果保存checkpoint时模型在CUDA上,当前resume时要确保模型和优化器都移动到同一设备:

# 初始化后移动到设备
model1 = model1().to(device)
optimizer = optim.Adam(model1.parameters(), lr=lr, betas=betas)
# 加载checkpoint后若有设备不匹配问题,手动移动优化器状态
for state in optimizer.state.values():
    for k, v in state.items():
        if isinstance(v, torch.Tensor):
            state[k] = v.to(device)

按照上述步骤排查,最可能的原因是优化器初始化时传入了错误的参数,或者模型结构发生了未察觉的变化。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 17:12:32