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

如何为依赖整批次计算的损失函数实现梯度累积?

梯度累积下依赖全局批次概率的损失函数正确实现

你的核心问题在于当前代码里的probs_avg是指数移动平均(0.5权重的平滑),并非梯度累积对应的「大批次全局概率平均」,完全不符合损失函数对全局批次的要求。要实现正确的梯度累积,需要把所有小批次的样本概率都保留/累加,直到累积到目标大批次规模后再计算全局平均,进而计算损失。

修正思路

  • 初始化累积容器,存储所有小批次的概率总和与对应总样本数
  • 每个小批次仅做前向传播,计算当前批次的概率并累加到容器中,不触发梯度计算
  • 当累积步数达到设定值时:
    1. 用累积的总概率和总样本数计算全局平均
    2. 基于全局平均计算损失
    3. 反向传播更新梯度
    4. 重置累积容器与优化器梯度

修正后的代码

# ----------------------- 梯度累积相关初始化 ----------------------
global_step = 0
accumulation_steps = 4  # 按需求设置要累积的小批次数量
probs_accumulator = None  # 存储累积的概率总和
total_samples = 0
# ------------------------------------------------------------------

for step, (x, y, idx) in enumerate(metric_logger.log_every(data_loader_u, print_freq, header)):
    x = x.to(device, non_blocking=True)
    y = y.to(device, non_blocking=True)
    
    with torch.cuda.amp.autocast():
        logits = model(x)
        probs_all = F.softmax(logits, dim=-1)  # [batch_size, num_classes]
        
        # 累积当前批次的概率总和与样本数
        if probs_accumulator is None:
            probs_accumulator = probs_all.sum(0)  # [num_classes]
            total_samples = x.size(0)
        else:
            probs_accumulator += probs_all.sum(0)
            total_samples += x.size(0)
    
    # ---------------------- 梯度累积触发逻辑 ----------------------
    if (step + 1) % accumulation_steps == 0:
        global_step += 1
        with torch.cuda.amp.autocast():
            # 计算大批次的全局概率平均
            probs_avg = probs_accumulator / total_samples
            loss_x = -(torch.log(probs_avg)).mean()
            # 损失除以累积步数,保证梯度规模与直接大批次训练一致
            loss = loss_x / accumulation_steps
        
        # 反向传播与优化更新
        loss.backward()
        optimizer.step()
        optimizer.zero_grad()
        
        # 重置累积容器
        probs_accumulator = None
        total_samples = 0
    
    torch.cuda.synchronize()

关键说明

  • 概率累积方式:用总样本数加权的总和计算全局平均,和小批次平均再取平均的数学结果一致,但存储总和更节省内存。
  • 损失缩放的必要性:梯度累积相当于把大批次损失拆分为小批次逐步累积,因此每个小批次对应的损失要除以累积步数,保证最终梯度规模和直接大批次训练匹配。
  • 内存优化:累积过程中仅做前向传播和概率累加,不计算损失梯度,避免不必要的内存占用。
  • 混合精度适配:保持autocast上下文包裹前向传播和损失计算,确保混合精度训练正常运行。

内容的提问来源于stack exchange,提问作者Phan Trang

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 05:40:12