ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch inplace操作报错:原因、排查与修复实战

PyTorch inplace操作报错:原因、排查与修复实战 我翻了翻自己之前训练模型时留的记录发现那个“one of the variables needed for gradient computation has been modified by an inplace operation”的报错几乎每个用PyTorch的人都会撞上。这个报错看着像天书实际背后逻辑很简单你在某个地方对张量做了原地修改打断了反向传播要用的计算图记录。这篇文章就把我踩过的坑、排查的思路和几种有效的修复方式全部摊开讲有代码、有步骤、有经验希望能帮你少走弯路。1. 先把这个报错信息掰开揉碎1.1 报错文字到底在说什么先看一条典型的报错信息RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation: [torch.FloatTensor [4, 4]], which is output 0 of AddBackward0, is at version 2; expected version 1 instead. Hint: enable anomaly detection to find the operation that was triggered without a gradient input, if you want to find the operation that caused the problem.拆开看它其实告诉了你三件事某个变量被inplace操作修改了这里说的是[torch.FloatTensor [4, 4]]也就是一个形状为4x4的浮点张量。它是某个反向传播节点的输出报错里提到的AddBackward0说明这个张量是加法操作的输出而这个加法操作在反向传播时需要用到它的旧值。版本号对不上is at version 2; expected version 1 instead。这句话最关键。PyTorch每个张量内部都会记一个版本号每做一次inplace修改版本号就加1。反向传播想拿的是version 1的值结果发现已经变成version 2了。简单说自动微分机制把前向计算的每一步都记下来了反向传播时要按记录倒推。你中途把某个张量原地改了记录不对了梯度自然算不出来所以就直接报错。1.2 自动微分为什么对inplace操作这么敏感PyTorch的自动微分核心就是一张动态计算图。前向传播时每个张量怎么来的、经过什么运算整个链条都被记录下来。反向传播时loss对每个参数的梯度就是沿着这条链反过来一点点算的。举个例子import torch x torch.randn(4, 4, requires_gradTrue) y x * 2 # 乘法操作创建一个新张量y z y 1 # 加法操作创建一个新张量z loss z.sum() # 求和 loss.backward() # 反向传播求梯度这段代码没问题。计算图大致是x → MulBackward → y → AddBackward → z → SumBackward → loss。反向传播的时候系统需要根据x的值来算梯度。那如果在中间加一个inplace操作呢x torch.randn(4, 4, requires_gradTrue) y x * 2 y.add_(1) # 注意这个是inplace操作原地给y加1 loss y.sum() loss.backward() # 这里会报错y.add_(1)把y原地加了1。问题在于前向计算记录里y是x*2的结果反向传播算梯度时需要读y的旧值。你把它原地改了版本号从1变成2系统发现数据对不上干脆报错。我的理解是这不是PyTorch在刁难你而是在保护你。原地修改张量可能会导致梯度计算出错而PyTorch不知道你想怎么处理这种情况所以宁可报错也不给你一个可能错误的梯度。这个设计思路本身是合理的。1.3 一个很容易混淆的概念叶子节点还有个经常一起出现的报错长这样RuntimeError: a leaf Variable that requires grad is being used in an inplace operation.这是inplace报错的另一个变体区别在于出问题的是叶子节点。叶子节点是指直接创建、不依赖其他张量的张量比如模型的参数或者你手动创建的requires_gradTrue的张量。对叶子节点做inplace操作PyTorch直接禁止没有任何商量的余地。原因很简单优化器更新参数时也是通过inplace操作来做的比如param.data - lr * param.grad如果前向传播过程中你把参数原地改了那优化器更新时用的可能就是改过的参数梯度、动量、学习率调度全都乱套了。2. 实际操作里哪些场景最容易触发这个报错2.1 激活函数里的inplace参数最经典的案例就是把ReLU的inplace参数设为Trueimport torch.nn as nn class MyModel(nn.Module): def __init__(self): super().__init__() self.relu nn.ReLU(inplaceTrue) # 这里就是坑 self.fc nn.Linear(16, 10) def forward(self, x): x self.fc(x) x self.relu(x) return xinplaceTrue的含义是ReLU直接在输入张量上做操作把小于0的元素置为0不新建张量。这样省内存速度也能快一点。问题在于如果输入x是某个计算图的中间结果而且这个计算图的后续反向传播还需要x的原始值那就会触发报错。实操中我建议在大部分情况下都别开inplaceTrue。省的那点内存和速度跟排查问题花费的时间比起来完全不划算。等模型训练稳定了、确实需要优化显存占用时再针对性开启。类似的激活函数还有LeakyReLU(inplaceTrue)、ELU(alpha1.0, inplaceTrue)等等遇到这类参数就多留个心眼。2.2 对中间张量做了原地修改写自定义损失函数或者自定义层的时候这种情况特别多。def custom_loss(pred, target): diff pred - target diff.clamp_(min0) # 原地clamp可能在反向传播时报错 return diff.sum()diff.clamp_(min0)把diff原地截断到不小于0。问题是diff是pred - target的结果反向传播时需要计算对pred的梯度而梯度计算依赖diff的原始输出。你把它改了梯度就出问题了。再看一个更隐蔽的def forward(self, x): x self.conv1(x) x F.relu(x) x[0, 0] 0 # 这种索引赋值也是inplace操作只是不显眼 x self.conv2(x) return xx[0, 0] 0看起来很正常但它是通过__setitem__实现的inplace操作。如果x后续参与梯度计算一样会报错。2.3 索引赋值与mask操作有些场景你会想对符合某些条件的元素做赋值比如grad torch.zeros_like(score) mask score 0 grad[mask] 1 # 这里inplace修改了grad而grad是计算图的一部分如果score需要求梯度而grad又参与了计算图的构建那grad[mask] 1就可能触发inplace报错。一个常用的修复思路是改用torch.wheregrad torch.where(score 0, torch.ones_like(score), torch.zeros_like(score))但torch.where在梯度反传时也有自己的行为特征需要具体看需求。关键是理解凡是带下划线结尾的方法、凡是索引赋值都是inplace操作都可能干扰反向传播。我把常见的inplace操作整理成了一张速查表方便你快速对照操作类别inplace写法常见触发场景加减乘除x.add_(1)、x.mul_(2)、x.div_(3)自定义loss中修改差值、归一化处理截断限制x.clamp_(min0)、x.relu_()损失函数做截断、梯度裁剪索引赋值x[mask] value、x[i, j] 0mask操作、特定位置赋值归一化x.softmax_(dim1)某些激活或归一化层其他x.copy_(y)、x.resize_()、x.zero_()数据预处理、缓冲区修改3. 实战排查三步定位到问题代码报错信息只告诉你出了问题但具体在哪一行需要自己去定位。我总结了一套比较有效的排查流程。3.1 第一步打开异常检测模式报错信息里其实给了提示enable anomaly detection...。PyTorch提供了一个开关torch.autograd.set_detect_anomaly(True)把这个开关放到训练脚本的开头PyTorch会在报错时额外输出更详细的堆栈信息直接指向触发inplace操作的代码位置。但注意这个开关会拖慢训练速度还会增加内存占用。我一般只在训练刚开始阶段或定位问题时开启定位完就关掉。如果你用的是PyTorch 2.x也可以这样from torch import autograd with autograd.detect_anomaly(): loss model(x) loss.backward()这样只在特定代码块内开启异常检测影响范围更小。3.2 第二步比对版本号缩小范围如果异常检测不够用还有一个黑科技手动查看张量的_version属性。x torch.randn(4, 4, requires_gradTrue) print(初始版本, x._version) # 输出通常是0 y x * 2 print(乘法后x版本, x._version) # 仍然是0乘法不是原地操作 y.add_(1) print(inplace修改y版本, y._version) # 变成1了_version属性是个内部计数器每次inplace操作都会加1。你可以把可疑的中间张量的_version打印出来看它在哪个操作之后变了基本就能锁定问题位置。在复杂的模型里我通常会写一个辅助函数来快速定位def check_version(tensor, name): current_version tensor._version if check_version.last_versions.get(name, current_version) ! current_version: print(f张量 {name} 的版本从 {check_version.last_versions[name]} 变成了 {current_version}) check_version.last_versions[name] current_version然后在模型forward的各个关键节点调用这个函数看哪个节点之后版本号变了。3.3 第三步从模型中间结果入手排查如果堆栈信息太混乱另一个办法是一层一层地测试模型。比如你有一个三层的模型class MyModel(nn.Module): def __init__(self): super().__init__() self.layer1 nn.Linear(32, 16) self.layer2 nn.Linear(16, 8) self.layer3 nn.Linear(8, 1) def forward(self, x): x self.layer1(x) x self.layer2(x) x self.layer3(x) return x训练时报inplace错误但不知道怎么定位。可以先用一个简单的输入跑前向传播然后把中间结果存下来逐个检查是否能正常求梯度model MyModel() x torch.randn(4, 32, requires_gradTrue) x1 model.layer1(x) x2 model.layer2(x1) x3 model.layer3(x2) for name, t in [(x, x), (x1, x1), (x2, x2), (x3, x3)]: print(name, t._version)如果某个中间张量的版本号和预期不一致检查这个张量产生了什么操作。如果各张量的版本号都正常那就说明问题出在loss计算或者反向传播过程中。实际排查中这个报错往往不是前向传播本身触发的而是在backward()的时候才暴露出来。因为inplace操作记录在计算图里只有反向传播时系统才会发现版本对不上。所以有时候用户会觉得“前向传播没问题啊报错在backward”其实是前向传播里已经埋下了雷只是引爆时间在backward。4. 有效的解决方案与代码修改模板定位到问题之后怎么改我总结了几套比较通用的方案。4.1 方案一关闭inplace开关最简单粗暴的就是把inplaceTrue改成inplaceFalseself.relu nn.ReLU(inplaceFalse) # 或者直接 nn.ReLU()或者不使用带下划线的原地操作# 修改前 x.clamp_(min0) # 修改后 x x.clamp(min0)这个方案的优点是不会改变模型结构缺点是多了一些临时张量显存占用会稍微增加。经验之谈如果你只是训练一个小模型或者显存还没有紧张到必须精打细算的程度直接用这个方案把inplace全部关掉省心。4.2 方案二克隆张量隔离计算图如果inplace操作暂时去不掉可以在操作之前克隆一份# 修改前 y x * 2 y.add_(1) # 对y做原地加法可能报错 # 修改后 y x * 2 y y.clone() # 关键克隆出一个独立的张量 y.add_(1) # 对克隆后的张量做原地操作不再影响原计算图clone()的作用是从计算图中分离出一个新张量同时对值做了一份拷贝。这样后续对y做inplace操作不会再干扰原来的x * 2节点。有些情况下也可以用detach()效果不同这里需要说清楚clone()会保留梯度路径复制的同时让新张量成为计算图的一部分。detach()会完全脱离计算图新张量不再记录梯度。在inplace报错场景下通常用clone()更安全因为detach()会切断梯度传播可能导致模型无法学习。4.3 方案三改写逻辑避免原地修改这个方案需要改核心逻辑但往往是最一劳永逸的。比如上面提到的diff.clamp_(min0)的场景可以改成# 修改前 diff pred - target diff.clamp_(min0) # inplace操作可能报错 # 修改后 diff (pred - target).clamp(min0) # 非inplace生成新张量再比如你在自定义损失函数里用了log_softmax后想修改结果可以这样改写# 修改前 log_probs F.log_softmax(logits, dim-1) log_probs[range(batch_size), labels] - 0.1 # inplace索引赋值 # 修改后 log_probs F.log_softmax(logits, dim-1) one_hot torch.zeros_like(log_probs) one_hot[range(batch_size), labels] 0.1 log_probs log_probs - one_hot # 非inplace操作核心思想是用“生成新张量”代替“原地修改旧张量”。4.4 优化器更新参数时的处理还有一个非常隐蔽的场景是在自定义优化或梯度裁剪时# 不安全的写法 for param in model.parameters(): param.data - lr * param.grad # param.data的原地修改其实这里PyTorch官方推荐的是用optimizer.step()来做参数更新param.data的操作大概率不会触发inplace报错因为param.data操作的是叶子节点的数据但如果你在更新前对param本身做了原地操作还是会出问题。如果必须手动更新参数比如做特殊的梯度惩罚正确的做法是with torch.no_grad(): for param in model.parameters(): param.data param.data - lr * param.grad注意param.data ...其实是在替换param.data的引用不是对原张量做inplace修改这种写法反而更安全。4.5 使用torch.autograd.set_detect_anomaly(True)之后要怎么做开启异常检测后如果PyTorch给出了具体的堆栈信息比如指向了某个自定义的forward函数那就去检查那个函数里有没有inplace操作。如果有按照前面几种方案去改。需要注意的坑是异常检测模式下PyTorch会记住前向传播的整个历史所以内存消耗会变大大模型训练时可能直接OOM。我建议在脚本里用环境变量控制开关方便随时打开和关闭import os import torch if int(os.environ.get(DEBUG_PYTORCH, 0)): torch.autograd.set_detect_anomaly(True)这样平时训练不开需要定位问题时再设置环境变量DEBUG_PYTORCH1比较灵活。5. 工程化建议从根源上减少inplace报错5.1 制定代码审查清单模型训练代码里有些模式天生就是inplace重灾区。我把它们整理成了清单每次写完代码对照自查一下激活函数用了inplaceTrue改成Falseloss函数里有_结尾的方法比如add_、mul_、clamp_改成非inplace版本张量索引赋值x[mask] value检查x是否参与梯度计算数据预处理里对模型输入做原地归一化、标准化自定义层里对中间结果做了原地操作梯度裁剪时对grad变量做了原地修改注意torch.nn.utils.clip_grad_norm_内部对grad进行原地缩放可能在某些极端情况下出问题这是我自己踩坑之后总结出来的清单不能说100%覆盖所有场景但大部分报错都能覆盖到。5.2 使用PyTorch 2.x的编译优化器PyTorch 2.x引入了torch.compile它会在编译阶段做更多静态分析有些inplace操作在编译模式下会被自动改写model torch.compile(model)但不是所有的inplace问题都能被compile化解。如果你用了torch.compile之后还报inplace错误建议先关掉compile再把代码按常规方式修好最后重新试试compile。这里说一个我实测的体验torch.compile对模型性能的提升是有前提的如果模型本身比较小、或者有大量动态shape操作收益不明显反而会增加编译时间。不要为了炫技而引入compile值不值得跑一遍对比再说。5.3 数据批量测试在正式训练之前写一个小脚本用一小批数据跑一遍前向和反向import torch from torch.utils.data import DataLoader, TensorDataset # 用随机数据模拟 x torch.randn(16, 32) y torch.randint(0, 10, (16,)) dataset TensorDataset(x, y) loader DataLoader(dataset, batch_size4) model MyModel() optimizer torch.optim.Adam(model.parameters(), lr1e-3) loss_fn torch.nn.CrossEntropyLoss() for batch_x, batch_y in loader: pred model(batch_x) loss loss_fn(pred, batch_y) loss.backward() optimizer.step() optimizer.zero_grad() print(OK)如果这段代码顺利跑通说明前向和反向的基本流程没问题可以开始正式训练。如果报错就按照前面的排查步骤来定位。这个习惯我一直在用大大减少了训练中途报错浪费的时间。有些问题等你跑了几百轮才发现那才叫痛苦。预训练一个模型可能要几天如果中途因为inplace报错中断之前的时间都白费了。5.4 把「版本号检查」用起来在你实在找不到问题的情况下可以用一个更底层的招数torch.autograd.graph.saved_tensors_hooks来检查和监控梯度的保存过程。这个接口允许你注册钩子函数在保存中间结果时自动执行一些检查操作import torch def pack_hook(x): print(f保存中间结果: {x.shape}, version: {x._version}) return x def unpack_hook(x): return x with torch.autograd.graph.saved_tensors_hooks(pack_hook, unpack_hook): loss model(x) loss.backward()它会把你模型里所有被保存下来、供反向传播使用的中间张量都打印出来。如果某个张量的版本号异常就能直接看出是谁在搞鬼。不过要提醒的是这个钩子也会拖慢速度同样适合定位问题时用正规训练时建议关掉。6. 聊几个容易搞混的边界情况6.1 inplace操作一定报错吗不一定很多人对这个报错有个误解以为只要用了inplace操作就一定会报错。其实不是。比如x torch.randn(4, 4, requires_gradTrue) x.relu_() # 对叶子节点做inplace这直接就报错了但如果是先detach再操作x torch.randn(4, 4, requires_gradTrue) y x.detach() # y不再需要梯度 y.add_(1) # 不会报错只要这个张量不参与梯度计算怎么inplace都行。报错的条件是被修改的张量在计算图中扮演了需要梯度计算的中间节点。另外如果一个张量已经通过requires_grad_(False)取消了梯度追踪它也可以随便原地修改。所以要不要处理inplace操作取决于这个张量是否参与梯度计算而不是绝对禁止所有inplace操作。6.2 为什么torch.no_grad()里做inplace没问题在torch.no_grad()的上下文里PyTorch不会构建计算图所以inplace操作不涉及梯度记录自然也不会触发报错with torch.no_grad(): x.add_(1) # 不会报错但要注意在torch.no_grad()外面如果你先构建了计算图再进入no_grad对中间张量做inplace还是会报错。因为计算图已经记住了那个张量。我遇到过这种情况有人以为把inplace操作放在no_grad()里就万事大吉了结果还是报错就是因为前面的计算图已经记住了这个张量后面的inplace修改依然会弄脏版本号。6.3 关于torch.utils.checkpoint的坑显存不够时很多人会用到梯度检查点activation checkpointing这个功能也会触发奇怪的inplace报错。原因在于梯度检查点的实现里前向传播会被分成多个段每段用完就丢反向传播时再重新计算。如果你在某个段里做了inplace操作重算的时候可能和其他段的记录对不上。碰到这种情况我的建议是先关掉checkpoint确认报错是否消失如果确实和checkpoint有关检查被checkpoint包裹的模块里有没有inplace操作把它们全部改成非inplace尽量避免在checkpoint段之间共享会被原地修改的张量这个坑比较冷门但一旦踩到会很困扰因为报错信息经常指向无关的代码行让人摸不着头脑。7. 最后的经验之谈做深度学习训练inplace报错本质上是在提醒你你对张量的操作方式跟自动微分的要求产生了矛盾。与其说它是一个bug不如说它是一个约定。理解了这套约定以后写代码自然就知道哪些操作是安全的哪些需要避开。我个人在实际操作中有几个比较固定的习惯所有激活函数一律不开inplace省下的内存远不足以抵消排查问题的成本。所有自定义loss函数只返回新张量绝不做原地截断或原地mask。训练前先跑小批量数据验证让报错发生在训练早期而不是跑了好几个小时之后才爆出来。开启torch.autograd.set_detect_anomaly(True)时一定注意内存大模型慎用。最后再分享一个小技巧如果你实在不想改代码又恰好是单卡训练可以在backward()之前手动调用torch.autograd.graph.saved_tensors_hooks来做一次“软检查”打印所有被保存的张量和版本号变化。它会明确告诉你到底是哪个张量出了问题接下来针对性修复就快多了。希望这篇文章能帮你走出inplace报错的迷宫。遇到这类问题别慌顺着报错信息、版本号和堆栈一层层排查绝大多数都能在半小时内解决。
返回列表