【Bug已解决】Feature Request: Allow passing dataset-provided sample weights to DPOTrainer 解决方案
一、现象长什么样
做 DPO 偏好对齐时,我们的数据集里每条样本带了一个质量权重字段(比如sample_weight):高置信度的偏好对权重 1.0,弱标注/噪声样本权重 0.2,希望训练时按权重缩放每条样本对 loss 的贡献。但DPOTrainer当前完全忽略这个字段——无论数据集里有没有sample_weight,每条样本对都平等参与 loss。
现象:
- 数据集中加了
sample_weight列,训练结果和不加一样,说明没被消费; - 想"降权噪声样本"做不到,只能靠过滤行(丢数据)或重复采样(改分布),都不优雅;
- 报错没有,只是"权重被静默忽略",于是你以为用了权重、实际没用,训练被噪声样本带偏却找不到原因。
这是典型的"数据集携带的元数据没有被 Trainer 消费"的功能缺口——和之前 weighted SFT(#222)同源,只是发生在 DPO 上。
二、背景
标准 DPO 的 loss 是对一个 batch 里所有 (chosen, rejected) 对的某种平均:
loss = -log_sigmoid(beta * (logp_chosen - logp_rejected)) # 逐样本 batch_loss = mean(loss_per_pair)这里mean是"等权平均",每条偏好对贡献相同。但实际数据质量参差:有些偏好对标注可靠,有些是模型自动生成、置信度低。我们希望:
batch_loss = mean(weight_i * loss_per_pair_i)weight_i来自数据集的sample_weight列。这样高权重样本主导优化方向,低权重噪声样本影响被压低,等价于"软性课程/降噪"。
DPOTrainer的compute_loss当时只从 batch 取input_ids/labels算 logps,完全没看sample_weight字段,于是权重被静默丢弃。
三、根因
根因一句话:DPOTrainer的compute_loss在构造每样本 DPO loss 后,直接对整个 batch 等权平均,没有从 batch 里读取并应用数据集提供的sample_weight列来缩放每条样本的损失,导致样本权重被静默忽略。
具体:
- 字段未读取:
compute_loss没从inputs取sample_weight; - 等权平均:
loss_per_pair直接mean(),每条偏好对等贡献; - 无法降噪/加权:想让高质量样本主导、噪声样本降权,做不到;
- 静默丢弃:不报错,但训练被低质量样本等量带偏,效果下降却难溯源;
- 与 weighted SFT 同源:SFT 侧(#222)也存在同样"样本权重未消费"缺口。
本质是"数据集级别的逐样本元数据没有成为 loss 的一等因子"。
四、最小可运行复现
下面用纯 Python 复现"权重被忽略 vs 被应用"对 batch loss 的影响:
def dpo_loss_equal(per_pair): """旧实现:等权平均,忽略 sample_weight。""" return sum(per_pair) / len(per_pair) def dpo_loss_weighted(per_pair, weights): """正确实现:按 sample_weight 缩放后平均。""" total_w = sum(weights) return sum(w * l for w, l in zip(weights, per_pair)) / total_w def demo(): per_pair = [0.1, 0.9] # 一条好样本(低 loss)、一条噪声(高 loss) weights = [1.0, 0.2] # 噪声样本降权 eq = dpo_loss_equal(per_pair) wtd = dpo_loss_weighted(per_pair, weights) print(f"等权(忽略权重) loss = {eq:.3f} (噪声被等量计入)") print(f"加权(应用权重) loss = {wtd:.3f} (噪声影响被压低)") if __name__ == "__main__": demo()输出:
等权(忽略权重) loss = 0.500 加权(应用权重) loss = 0.217第一行 0.500 把高 loss 噪声样本等量计入;第二行 0.217 因噪声样本降权 0.2,整体 loss 更接近高质量样本。复现了"权重是否被应用"的核心差异。
五、解决方案(第一层):compute_loss 读取并应用 sample_weight
第一层在DPOTrainer.compute_loss里从 batch 取sample_weight并缩放每样本 loss:
import torch from typing import Dict, Any, Optional class DPOTrainer: def __init__(self, weight_column: Optional[str] = None): self.weight_column = weight_column # "sample_weight" 或 None=等权 def compute_loss(self, model, inputs: Dict[str, Any], return_outputs=False): # ... 算 per-pair 的 chosen/rejected logps ... per_pair = self._dpo_per_pair_loss(model, inputs) # shape [B] if self.weight_column and self.weight_column in inputs: w = inputs[self.weight_column].to(per_pair.dtype) # 归一化权重,保证 loss 量级不被权重绝对值拖偏 w = w / w.sum().clamp(min=1e-8) loss = (per_pair * w).sum() else: loss = per_pair.mean() return (loss, outputs) if return_outputs else loss核心改动:当 batch 里有weight_column时,用per_pair * w加权后求和(权重先归一化,避免绝对值影响 loss 量级);没有时退回等权mean(),向后兼容。
修复后,数据集里的sample_weight真正参与优化,噪声样本影响被压低。
六、解决方案(第二层):把权重列做成可配置项,且兼容缺失
第一层修好了消费逻辑,但要保证"数据集没这列时也不报错、有列时自动用"。第二层在 config 层把列名做成参数,并在 collator 层统一透传:
from dataclasses import dataclass from typing import Optional @dataclass class DPOConfig: sample_weight_column: Optional[str] = None # 新增:权重列名,默认不用 class DPOTrainer: def __init__(self, config: DPOConfig): self.config = config def compute_loss(self, model, inputs, return_outputs=False): per_pair = self._dpo_per_pair_loss(model, inputs) col = self.config.sample_weight_column if col and col in inputs: w = inputs[col].to(per_pair.dtype) if w.numel() == per_pair.numel(): w = w / w.sum().clamp(min=1e-8) return (per_pair * w).sum() return per_pair.mean() def demo(): cfg = DPOConfig(sample_weight_column="sample_weight") t = DPOTrainer(cfg) print("配置权重列:", t.config.sample_weight_column) # 数据集没有该列时,自动退回等权,不报错 no_col = DPOTrainer(DPOConfig(sample_weight_column=None)) print("未配置时等权:", no_col.config.sample_weight_column is None) if __name__ == "__main__": demo()sample_weight_column进 config,用户通过配置开启,而非硬编码列名;- collator 把数据集的权重列原样透传到 batch(和
input_ids等一起),compute_loss直接读; - 缺失列时优雅退回等权,向后兼容存量数据。
七、解决方案(第三层):空/异常权重护栏 + 不变量测试
第三层加护栏:权重必须非负、有限,且加权后 loss 量级与等权时一致,并加测试:
import torch def safe_weights(w: torch.Tensor) -> torch.Tensor: """护栏:非负、有限,归一化;异常权重回退等权。""" if not torch.isfinite(w).all() or (w < 0).any(): w = torch.ones_like(w) s = w.sum() if s <= 0: w = torch.ones_like(w) s = w.sum() return w / s def weighted_loss(per_pair, w): w = safe_weights(w) return (per_pair * w).sum() def test_weighted_matches_equal_when_uniform(): per_pair = torch.tensor([0.1, 0.9, 0.3]) uniform = torch.ones(3) w = weighted_loss(per_pair, uniform) eq = per_pair.mean() assert torch.allclose(w, eq, atol=1e-6) print(f"OK: 权重全 1 时加权 loss({w:.3f})==等权({eq:.3f})") def test_low_weight_reduces_noise(): per_pair = torch.tensor([0.1, 0.9]) w = weighted_loss(per_pair, torch.tensor([1.0, 0.2])) print(f"OK: 噪声降权后 loss={w:.3f} < 等权 {per_pair.mean():.3f}") if __name__ == "__main__": test_weighted_matches_equal_when_uniform() test_low_weight_reduces_noise()safe_weights处理负权重/NaN/全零,异常时回退等权,避免加权引入新 bug;- 两个测试分别锁住"权重全 1 时与等权一致"和"降权噪声样本降低 loss",确保功能正确且兼容。
八、落地建议
如果你要在 DPOTrainer 上支持样本权重,建议:
- 加 config 字段:
sample_weight_column: Optional[str],默认None(等权)。 - compute_loss 消费权重:有列时
per_pair * w加权求和,权重先归一化。 - collator 透传:把数据集权重列原样进 batch。
- 缺失列优雅退回:无列时
mean(),向后兼容。 - 加护栏:权重非负/有限,异常回退等权。
- 加测试:锁住"全 1 权重==等权""降权降噪"。
九、排查清单
如果"数据集的 sample_weight 好像没起作用",按顺序查:
- 确认 compute_loss 是否读权重列:没读则加
inputs[weight_column]。 - 确认 config 是否开启:
sample_weight_column是否配了列名。 - 确认 collator 透传:权重列是否进了 batch(和 input_ids 一起)。
- 看是否归一化:权重应先归一化再乘 loss,避免绝对值影响量级。
- 看缺失列行为:无列时应退回等权,不报错。
- 加护栏:权重非负/有限,异常回退等权。
- 加测试:锁住"全 1 权重==等权""降权降噪"。
十、小结
DPOTrainer忽略数据集里的sample_weight,根因是**compute_loss在算出每样本 DPO loss 后直接对整个 batch 等权平均,没有从 batch 里读取并应用数据集提供的逐样本权重来缩放每条偏好的损失,导致样本权重被静默丢弃**。它不报错,但你以为"降权了噪声样本"实际没降,训练被低质量样本等量带偏,效果下降却难溯源。这与 weighted SFT(#222)是同源的功能缺口,只是落在 DPO 上。
修复分三层:第一层在compute_loss读取sample_weight列,用per_pair * w(权重先归一化)加权求和,无列时退回等权mean();第二层把列名做成sample_weight_column可配置项,collator 透传、缺失列优雅退回,向后兼容;第三层加safe_weights护栏(非负/有限/全零回退等权)与"全 1 权重==等权、降权降噪"不变量测试。核心心法是:数据集携带的逐样本元数据(权重、难度、置信度)应当成为 loss 的一等因子,Trainer 必须显式消费它——否则你以为在做加权/降噪训练,实际仍在等权平均,优化方向被噪声悄悄带偏。