如何解决PyTorch compile引发的stride断言错误?
使用torch.compile()训练时出现AssertionError(stride不匹配)的解决方法
问题描述
运行PyTorch 2.0教程3.2节代码时,启用torch.compile()后,训练在第一个epoch结束时崩溃,抛出错误:
AssertionError: expected size 64==64, stride 3136==1 at dim=1
不使用torch.compile()时代码可正常运行,已找到相关修复PR但不清楚如何应用。
错误回溯关键片段
File /tmp/torchinductor_isaac-aktam/fi/cfignuiw5cmdhtpeow6axfwfjocu44zvgnrf4rlockxh2k5th7f3.py:5638, in call(args) 5636 del primals_4 5637 buf473 = buf472[0] -> 5638 assert_size_stride(buf473, (s0, 64, 56, 56), (200704, 1, 3584, 64)) 5639 buf474 = buf472[1] 5640 assert_size_stride(buf474, (64, 64, 1, 1), (64, 1, 64, 64)) AssertionError: expected size 64==64, stride 3136==1 at dim=1
解决方法
这个问题是PyTorch Inductor编译器在处理卷积层梯度回传时的张量内存布局不匹配问题,相关修复已合并到PyTorch主分支,可通过以下方式解决:
1. 升级PyTorch到最新稳定版
该修复已包含在PyTorch 2.1及以后的稳定版本中,直接升级即可彻底解决问题:
- 使用pip升级:
pip install --upgrade torch torchvision torchaudio - 使用conda升级(根据你的CUDA版本调整参数):
conda install pytorch torchvision torchaudio pytorch-cuda=11.8 -c pytorch -c nvidia
2. 临时 workaround(无法升级时)
如果暂时无法升级PyTorch,可尝试以下临时方案:
- 切换编译器后端:改用eager后端编译(性能提升有限,但能避免错误):
compiled_model = torch.compile(model, backend="eager") - 确保输入张量连续:在训练步骤中对输入数据调用
.contiguous(),强制张量内存布局连续:def train_step(...): for batch, (X, y) in enumerate(dataloader): X, y = X.to(device).contiguous(), y.to(device).contiguous() # 后续前向传播、损失计算等逻辑 y_pred = model(X) loss = loss_fn(y_pred, y) # ... - 调整编译模式:使用
reduce-overhead模式编译,减少激进优化:compiled_model = torch.compile(model, mode="reduce-overhead")
问题原因
错误源于Inductor编译器在生成反向传播代码时,对卷积层梯度张量的stride预期与实际运行时的张量stride不匹配。相关修复PR调整了卷积梯度计算的内存布局处理逻辑,确保张量的stride符合编译时的预期。
内容的提问来源于stack exchange,提问作者Isaac A
相关产品推荐
相关产品推荐

