PyTorch RuntimeError:二次反向传播失败,数据集中途扩容引发
问题分析与解决方案
核心问题原因
你遇到的第一个报错本质是前4层的计算图被多次反向传播复用:
- 变量
out是模型前4层的输出,它携带完整的计算图(关联输入data和前4层的所有参数)。 - 每个子批次的
out_bm是从out切片得到的,因此每个子批次的loss.backward()都会尝试回溯到前4层的计算图。 - 第一次执行
loss.backward()时,PyTorch默认会释放计算图的中间缓存以节省内存;第二次反向传播时,缓存已被清空,自然触发报错。
而设置retain_graph=True后出现的inplace操作错误,是因为保留计算图后,后续的optimizer.step()(参数inplace更新)或mixup_data中的inplace操作修改了计算图依赖的变量版本,导致第二次反向传播时变量状态不匹配。
正确修改思路
你需要将所有子批次的损失累加,只执行一次反向传播和参数更新。这样前4层的计算图只会被回溯一次,既避免了重复使用计算图的问题,也符合梯度累积的训练逻辑(等效于用8*batch_size的大批次训练)。
修改后的代码
for batch_idx, (data, target, _) in enumerate(train_loader): data = data.to(device) target_ohe = F.one_hot(target, args.num_classes) target_ohe = target_ohe.to(device) # Forward pass first part out = model(data, depth = args.depth, pass_part = 'first') # Mixup data out, target_ohe_full, _ = mixup_data.mixup_data(args.method, out, target_ohe, epoch, batch_idx, args.n_fraction, args.g_fraction, args.betaf_alpha, device) # Forward pass second part num_batch_mixup = int(np.ceil(out.shape[0]/args.batch_size)) total_loss = 0.0 # 初始化总损失 for batch_mixup_idx in range(num_batch_mixup): out_bm = out[args.batch_size*batch_mixup_idx : args.batch_size*(batch_mixup_idx+1)] output = model(out_bm, depth = args.depth, pass_part = 'second') target_ohe = target_ohe_full[args.batch_size*batch_mixup_idx : args.batch_size*(batch_mixup_idx+1)] loss = criterion(output, target_ohe) total_loss += loss # 累加子批次损失 # 仅执行一次反向传播和参数更新 total_loss.backward() if args.grad_clip: nn.utils.clip_grad_value_(parameters = model.parameters(), clip_value = args.grad_clip) optimizer.step() optimizer.zero_grad() scheduler.step()
额外说明
如果你确实需要每个子批次单独更新参数(不推荐,等效于用小批次多次更新,前4层会被多次更新),则需要在生成out后先剥离计算图(out = out.detach()),但这样前4层的参数无法获得梯度更新,显然不符合你的训练目标,因此优先选择损失累加的方案。
内容的提问来源于stack exchange,提问作者Liisjak
相关产品推荐
相关产品推荐

