排查MMDetection中_sigmoid_focal_loss原地操作报错问题
检查MMCV内置的Focal Loss实现:你使用的是编译了OPS的MMCV,部分场景下会调用
mmcv.ops.sigmoid_focal_loss这个底层实现而非纯Python版。可以全局搜索代码中是否直接调用了该函数,若存在,尝试替换为纯Python实现,或检查该底层函数是否存在原地修改张量的逻辑。排查数据预处理流程:训练数据增强的transform可能存在原地操作,比如部分自定义或内置的transform在处理图像、gt_bboxes、gt_labels时使用了
inplace=True参数。检查你的dataset和data_pipeline配置,逐一核对每个transform的实现细节。检查其他损失分支的实现:除Focal Loss外,模型的IoU Loss、CrossEntropy Loss等其他损失函数也可能存在原地修改。遍历
mmdet/models/losses/下的所有文件,重点排查是否有x[:] = ...、x += y这类原地赋值操作。排查自定义模块与钩子函数:如果你添加了自定义模型层、训练Hook或优化器扩展,这些部分可能存在原地操作。比如Hook中直接修改模型输出或梯度,或者自定义层使用了
nn.ReLU(inplace=True)这类带inplace参数的激活函数。特征融合与张量操作环节:模型neck或head部分的特征拼接、拆分、融合操作中,可能存在对张量的原地修改。比如使用
torch.cat后直接对结果张量进行原地赋值,或特征融合时使用了inplace操作符。精准定位被修改的变量:利用
torch.autograd.set_detect_anomaly(True)的报错信息,找到被原地修改的变量名称,全局搜索该变量的所有操作,定位到具体的修改步骤。比如报错提示variable 'xxx' was modified inplace,就追踪xxx的定义和后续所有操作,排查是否有xxx.data = ...、xxx.fill_(...)这类原地操作。
内容的提问来源于stack exchange,提问作者user9690322

