1. VERL框架中的损失函数聚合机制解析
在深度强化学习框架VERL中,损失函数的计算方式直接影响模型训练的稳定性和最终效果。让我们深入剖析这个核心函数agg_loss的实现逻辑。
1.1 函数接口与参数说明
python复制def agg_loss(
loss_mat: torch.Tensor,
loss_mask: torch.Tensor,
loss_agg_mode: str,
dp_size: int = 1,
batch_num_tokens: Optional[int] = None,
global_batch_size: Optional[int] = None,
loss_scale_factor: Optional[int] = None,
):
关键参数解析:
loss_mat: 微批次损失矩阵,形状为(batch_size, response_length)loss_mask: 损失掩码,标识哪些位置是有效tokenloss_agg_mode: 损失聚合模式,支持四种不同计算方式dp_size: 数据并行规模,用于分布式训练场景下的梯度校正
1.2 四种聚合模式详解
1.2.1 token-mean模式
这是语言模型训练中最常用的模式,计算逻辑如下:
python复制if loss_agg_mode == "token-mean":
if batch_num_tokens is None:
batch_num_tokens = loss_mask.sum()
loss = verl_F.masked_sum(loss_mat, loss_mask) / batch_num_tokens * dp_size
数学表达式:
code复制Loss = (Σ所有有效token的损失) / (总有效token数) × dp_size
典型应用场景:
- 预训练(Pre-training)
- 指令微调(SFT)
- 常规语言模型训练
技术要点:
- 每个token对梯度的贡献权重相等
- 长序列自然获得更大权重
- 分布式训练时需乘以dp_size保持梯度一致性
1.2.2 seq-mean-token-sum模式
python复制elif loss_agg_mode == "seq-mean-token-sum":
seq_losses = torch.sum(loss_mat * loss_mask, dim=-1) # 序列内token求和
seq_mask = (torch.sum(loss_mask, dim=-1) > 0).float() # 排除全mask序
