[v1][feature] add dpo trainer (#10544)

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
codingma
2026-06-26 15:32:10 +08:00
committed by GitHub
parent b7615dbdc9
commit 9c0b4b3835
7 changed files with 859 additions and 4 deletions

View File

@@ -14,6 +14,7 @@
import os
from dataclasses import dataclass, field
from typing import Literal
from uuid import uuid4
from .arg_utils import BatchingStrategy, PluginConfig, get_plugin_config
@@ -115,6 +116,30 @@ class TrainingArguments:
default=1,
metadata={"help": "Log metrics every N optimizer steps."},
)
pref_loss: Literal["sigmoid", "orpo", "simpo"] = field(
default="sigmoid",
metadata={"help": "The type of DPO loss to use."},
)
pref_beta: float = field(
default=0.1,
metadata={"help": "The beta parameter in the preference loss."},
)
pref_ftx: float = field(
default=0.0,
metadata={"help": "The supervised fine-tuning loss coefficient in DPO training."},
)
simpo_gamma: float = field(
default=0.5,
metadata={"help": "The target reward margin term in SimPO loss."},
)
dpo_label_smoothing: float = field(
default=0.0,
metadata={"help": "The robust DPO label smoothing parameter in cDPO that should be between 0 and 0.5."},
)
ld_alpha: float | None = field(
default=None,
metadata={"help": "Alpha parameter from LD-DPO, controls weighting of verbose token log-probabilities."},
)
def __post_init__(self) -> None:
self.dist_config = get_plugin_config(self.dist_config)