PyTorch从Checkpoint加载GradScaler后执行scaler.step(optimizer)出现张量形状不匹配RuntimeError问题
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

