You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.04 21:00:49