使用 MinT
DPO 偏好优化
DPO(直接偏好优化)适用于 chosen / rejected 偏好对数据,训练模型倾向于 chosen 回复。在 MinT 上通过 forward_backward_custom 结合客户端自定义 Bradley-Terry loss 实现。
数据形状
训练数据由 (prompt, chosen, rejected) 三元组构成,用 PreferencePair 表示。每个 pair 展平成两个 Datum,排列为 [chosen₀, rejected₀, chosen₁, rejected₁, …]——偶数下标对应 chosen,奇数对应 rejected。Loss 函数依赖这一顺序假设。
构造 Datum
Prompt token 的 loss weight 为 0,completion token 为 1.0;completion 前加空格后编码,末尾追加 eos_token_id;input 取 all_tokens[:-1],target 取 all_tokens[1:]。
def build_datum(prompt_tokens, completion_text, tokenizer):
completion_tokens = tokenizer.encode(f" {completion_text}", add_special_tokens=False)
completion_tokens.append(tokenizer.eos_token_id)
all_tokens = prompt_tokens + completion_tokens
weights = [0.0] * (len(prompt_tokens) - 1) + [1.0] * len(completion_tokens)
return types.Datum(
model_input=types.ModelInput.from_ints(tokens=all_tokens[:-1]),
loss_fn_inputs={"target_tokens": all_tokens[1:], "weights": weights},
)Bradley-Terry Loss
forward_backward_custom() 先让 datums 过 model 拿到 logprobs,再用 (data, logprobs_list) 调用 Python loss。关键:logprobs 要保持 Tensor,不能提前转成 Python list,否则梯度断裂,反传失败。
import torch, torch.nn.functional as F
def sequence_logprob(logprobs, weights):
return torch.dot(logprobs.flatten().float(),
_to_float_tensor(weights))
def pairwise_preference_loss(data, logprobs_list):
chosen_scores, rejected_scores = [], []
for c_d, r_d, c_lp, r_lp in zip(
data[::2], data[1::2], logprobs_list[::2], logprobs_list[1::2]
):
chosen_scores.append(sequence_logprob(c_lp, c_d.loss_fn_inputs["weights"]))
rejected_scores.append(sequence_logprob(r_lp, r_d.loss_fn_inputs["weights"]))
margins = torch.stack(chosen_scores) - torch.stack(rejected_scores)
loss = -F.logsigmoid(margins).mean()
metrics = {
"loss": float(loss.detach()),
"pair_accuracy": float((margins > 0).float().mean().detach()),
"mean_margin": float(margins.mean().detach()),
}
return loss, metrics训练循环
result = training_client.forward_backward_custom(
data, pairwise_preference_loss
)
fb = result.result()
print(fb.metrics)
training_client.optim_step(types.AdamParams(learning_rate=1e-5)).result()参数
| 参数 | 默认值 | 含义 |
|---|---|---|
MINT_BASE_MODEL | Qwen/Qwen3-0.6B | 底座模型 |
MINT_LORA_RANK | 16 | LoRA 秩 |
MINT_DPO_STEPS | 3 | 训练步数 |
MINT_DPO_LR | 1e-5 | Adam 学习率 |
完整案例
端到端的 DPO 实验(含 eval-first 数据切分、留出集 pair_accuracy)见chat-dpo 实践案例。