【Bug已解决】Feature Request: Allow passing dataset-provided sample weights to DPOTrainer 解决方案 【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并缩放每样本 lossimport 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_outputsFalse): # ... 算 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(min1e-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_outputsFalse): 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(min1e-8) return (per_pair * w).sum() return per_pair.mean() def demo(): cfg DPOConfig(sample_weight_columnsample_weight) t DPOTrainer(cfg) print(配置权重列, t.config.sample_weight_column) # 数据集没有该列时自动退回等权不报错 no_col DPOTrainer(DPOConfig(sample_weight_columnNone)) 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, atol1e-6) print(fOK: 权重全 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(fOK: 噪声降权后 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 必须显式消费它——否则你以为在做加权/降噪训练实际仍在等权平均优化方向被噪声悄悄带偏。