如何为依赖整批次计算的损失函数实现梯度累积?
梯度累积下依赖全局批次概率的损失函数正确实现
你的核心问题在于当前代码里的probs_avg是指数移动平均(0.5权重的平滑),并非梯度累积对应的「大批次全局概率平均」,完全不符合损失函数对全局批次的要求。要实现正确的梯度累积,需要把所有小批次的样本概率都保留/累加,直到累积到目标大批次规模后再计算全局平均,进而计算损失。
修正思路
- 初始化累积容器,存储所有小批次的概率总和与对应总样本数
- 每个小批次仅做前向传播,计算当前批次的概率并累加到容器中,不触发梯度计算
- 当累积步数达到设定值时:
- 用累积的总概率和总样本数计算全局平均
- 基于全局平均计算损失
- 反向传播更新梯度
- 重置累积容器与优化器梯度
修正后的代码
# ----------------------- 梯度累积相关初始化 ---------------------- 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
相关产品推荐
相关产品推荐

