Mind Lab Toolkit (MinT)
使用 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_MODELQwen/Qwen3-0.6B底座模型
MINT_LORA_RANK16LoRA 秩
MINT_DPO_STEPS3训练步数
MINT_DPO_LR1e-5Adam 学习率

完整案例

端到端的 DPO 实验(含 eval-first 数据切分、留出集 pair_accuracy)见chat-dpo 实践案例

本页目录