PyTorch Checkpointing报错:全局表示重计算时张量元数据不匹配(含额外采样场景)
我正在用PyTorch搭建模型,通过前向管线计算一个全局表示,之后这个管线会被用到网络后续的额外采样流程中。当我不使用checkpointing、完全重新计算全局表示时,一切正常,梯度也能正确回流。但当我尝试用torch.utils.checkpoint通过反向传播时重计算全局表示来节省内存时,出现了类似这样的运行时错误:
torch.utils.checkpoint.CheckpointError: torch.utils.checkpoint: Recomputed values for the following tensors have different metadata than during the forward pass. tensor at position 34: saved metadata: {'shape': torch.Size([128, 192]), 'dtype': torch.bfloat16, 'device': device(type='mps', index=0)} recomputed metadata: {'shape': torch.Size([128, 128, 192]), 'dtype': torch.float32, 'device': device(type='mps', index=0)} ... (more tensor mismatches follow) ...
我的环境细节:
- 运行在MPS后端(Apple Silicon),使用autocast实现混合精度(bfloat16)
- 全局表示是在一个模块中计算的,之后会输入到额外采样流程里,所以梯度必须能正确回流
- 完全重新计算全局表示(也就是把整个前向跑两遍)效率太低,所以checkpointing是必须的
除此之外,我已经试过一些修复方案,比如把所有原地操作替换成非原地操作,但这些修改并没有解决问题。
另外,我在Gumbel采样流程里用到了下面这行代码:cond_expanded = cond_cont.unsqueeze(1).expand(B, num_samples, -1).reshape(B * num_samples, -1)
我的本意是让条件在多个蒙特卡洛样本上正确广播,但我怀疑这个unsqueeze/expand/reshape的操作序列可能导致了前向保存的张量和反向重计算的张量之间的元数据不匹配。
我猜测这个问题要么和checkpointing与autocast的交互有关,要么是重计算过程中张量维度意外发生了变化。有没有人遇到过类似的问题?或者知道如何在享受checkpointing好处的同时,确保重计算的张量和原前向的张量在形状、dtype和设备上都匹配?任何解决建议,或者能在不牺牲梯度回流的前提下高效节省内存的变通方法,都会非常有帮助。
如果需要更多上下文或代码片段,我可以提供。
(另外或许可以创建一个「torch.utils.checkpoint」标签)
备注:内容来源于stack exchange,提问作者Foenkel

