大模型之基于TRL的DPO对齐训练实战篇
2026/7/31 12:02:01 网站建设 项目流程

1、概念说明

(1)TRL

TRL = Transformer Reinforcement Learning,它是huggingface提供的专门给Transformer大模型做对齐训练的工具库,把大模型对齐常用算法封装好了。可以完成:

  • SFT 监督微调
  • RM 奖励模型训练
  • PPO(RLHF)
  • DPO / IPO / KTO 偏好对齐

核心Trainer类:

  • SFTTrainer:监督微调 SFT,做 SFT 训练;
  • RewardTrainer:训练奖励模型 RM;
  • PPOTrainer:传统 RLHF 的 PPO 训练器,需要 RM,实现 PPO 强化学习循环,会在线 generate 采样;
  • DPOTrainer:实现标准 DPO 算法

DPOTrainer的底层工作:

  • 读取prompt/chosen/rejected,按照 tokenizer 的 chat_template 拼接输入;
  • 自动计算 policy 模型、ref 参考模型对chosen、rejected完整回答的序列联合对数概率;
  • 实现完整 DPO loss;
  • 自动处理 ref 模型冻结,ref 不计算梯度;
  • 支持 LoRA 训练(peft),不用全量微调;
  • 自动计算监控指标:rewards/chosen、rewards/rejected、奖励差值,训练日志直接打印;
  • 支持梯度累积、bf16、8bit 优化器等。

总结:TRL 是 HuggingFace 开源的大模型对齐库,封装了 SFT、奖励模型、PPO、DPO 等对齐算法。我们做 DPO 直接使用DPOTrainer,它内部已经实现 DPO 损失函数,负责计算 policy/ref 模型的序列对数概率、冻结 ref 模型、训练循环,极大降低对齐代码开发量。

2、数据准备

医疗循证DPO数据集

还是医疗相关的DPO数据集,因为之前已经训练了一个医疗相关的SFT模型。

总共有1400条数据,

每条数据包含三个字段:

  • prompt(string): 医学问题或查询
  • chosen(string): 高质量回答(作为偏好目标)
  • rejected(string): 低质量回答(作为负样本)

示例格式:

{ "prompt": "在阿尔茨海默病与溃疡性结肠炎患者中,PPARG 和 NOS2 作为共同基因,是否通过调控巨噬细胞和小胶质细胞极化参与疾病的发生发展?", "chosen": "从目前的人类与动物实验证据来看,PPARG 和 NOS2 很有可能作为共同炎症枢纽基因,通过调控巨噬细胞/小胶质细胞的极化状态参与阿尔茨海默病和溃疡性结肠炎的发生发展...", "rejected": "这是一个非常具体且专业的问题,涉及到两种疾病的共同机制。根据现有的生物医学研究,我们可以进行一个基于科学逻辑的推理和分析..." }

这个数据集的标注逻辑:
chosen:循证严谨,区分证据等级,承认哪些地方证据不足
rejected:过于绝对,把未完全证实的假说当成板上钉钉结论,产生误导。

共有1400条样本集,人工切分成2部分,med_dpo_answer_train.jsonl和med_dpo_answer_test.jsonl

3、训练dpo的lora

代码:

import os os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" import torch from datasets import load_dataset from transformers import AutoModelForCausalLM, AutoTokenizer from trl import DPOTrainer, DPOConfig from peft import LoraConfig # =========路径配置======== model_path = "./merged_sft_qwen7b_med" output_lora_dir = "./med_dpo_lora" train_file = "/root/autodl-tmp/datas/rlhf/med_dpo_answer_train.jsonl" test_file = "/root/autodl-tmp/datas/rlhf/med_dpo_answer_test.jsonl" SYSTEM_PROMPT = "你是专业的医疗咨询助手,回答仅供科普参考,不能替代执业医师面诊,诊疗请遵从线下医生的专业意见。" # 加载tokenizer tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) # =========加载数据集======== ds = load_dataset( "json", data_files={ "train": train_file, "test": test_file } ) def add_system_template(sample): messages = [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": sample["prompt"]} ] full_prompt = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True ) return { "prompt": full_prompt, "chosen": sample["chosen"], "rejected": sample["rejected"] } ds = ds.map(add_system_template) # =========预先过滤超长样本,总token不超过8192======== def filter_fn(sample): full_text = sample["prompt"] + sample["chosen"] + sample["rejected"] tok = tokenizer(full_text, return_length=True) return tok["length"][0] <= 8192 ds = ds.filter(filter_fn) train_dataset = ds["train"].shuffle(seed=42) eval_dataset = ds["test"] #打印校验 print("===检查第一条样本prompt头部===") print(ds["train"][0]["prompt"][:600]) print(f"训练集样本数:{len(train_dataset)}") print(f"验证集样本数:{len(eval_dataset)}") # =========加载完整SFT基座模型======== model = AutoModelForCausalLM.from_pretrained( model_path, torch_dtype=torch.bfloat16, device_map="auto", trust_remote_code=True ) # =========构造LoraConfig对象======== lora_config = LoraConfig( r=16, lora_alpha=32, target_modules=["q_proj", "k_proj", "v_proj", "o_proj"], lora_dropout=0.05, bias="none", task_type="CAUSAL_LM" ) # ========= trl1.9.2 DPOConfig 开启 do_eval=True ========= dpo_args = DPOConfig( output_dir=output_lora_dir, beta=0.25, loss_type="sigmoid", learning_rate=4e-6, num_train_epochs=0.7, # 1100样本,控制不要过高防止过拟合 per_device_train_batch_size=2, gradient_accumulation_steps=4, logging_steps=5, # --------eval相关配置-------- do_eval=True, eval_strategy="steps", eval_steps=10, # 每10训练step跑一次验证集 load_best_model_at_end=True, metric_for_best_model="eval_rewards/accuracies", greater_is_better=True, # 奖励准确率越高越好 save_total_limit=3, # 最多保留3个checkpoint bf16=True, optim="adamw_torch", report_to=[], ) trainer = DPOTrainer( model=model, args=dpo_args, train_dataset=train_dataset, eval_dataset=eval_dataset, peft_config=lora_config, ref_model=None ) #启动训练 trainer.train() # load_best_model_at_end开启后,trainer.model已经是验证集最优权重 trainer.save_model(output_lora_dir) print(f"DPO训练完成,最优LoRA适配器输出至:{output_lora_dir}")

运行环境:

使用32*4=128G的显存运行的,32*2运行不起来,在eval阶段再加载eval模型时会出现OOM。

4、运行指标

root@autodl-container-45404b8e68-e4244a4c:~/autodl-tmp/codes/sft# python train_dpo_med.py ===检查第一条样本prompt头部=== <|im_start|>system 你是专业的医疗咨询助手,回答仅供科普参考,不能替代执业医师面诊,诊疗请遵从线下医生的专业意见。<|im_end|> <|im_start|>user 基于最新循证指南的综合治疗管理对于肝硬化患者的临床预后和并发症控制有何影响?<|im_end|> <|im_start|>assistant 训练集样本数:1100 验证集样本数:299 [transformers] `torch_dtype` is deprecated! Use `dtype` instead! Loading weights: 100%|███████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 339/339 [00:02<00:00, 143.71it/s] Dropping fully truncated examples from train dataset: 100%|█████████████████████████████████████████████████████████████████████████████████████████████████| 1100/1100 [00:01<00:00, 781.46 examples/s] Dropping fully truncated examples from eval dataset: 100%|████████████████████████████████████████████████████████████████████████████████████████████████████| 299/299 [00:00<00:00, 754.57 examples/s] [transformers] The tokenizer has new PAD/BOS/EOS tokens that differ from the model config and generation config. The model config and generation config were aligned accordingly, being updated with the tokenizer's values. Updated tokens: {'bos_token_id': None, 'pad_token_id': 151643}. {'loss': '0.6655', 'grad_norm': '22.51', 'learning_rate': '3.835e-06', 'entropy': '1.491', 'num_tokens': '8.12e+04', 'logits/chosen': '0.8174', 'logits/rejected': '1.127', 'mean_token_accuracy': '0.5815', 'rewards/chosen': '-0.118', 'rewards/rejected': '-0.218', 'rewards/accuracies': '0.575', 'rewards/margins': '0.09999', 'logps/chosen': '-1649', 'logps/rejected': '-1457', 'epoch': '0.03636'} {'loss': '0.692', 'grad_norm': '23.94', 'learning_rate': '3.629e-06', 'entropy': '1.462', 'num_tokens': '1.613e+05', 'logits/chosen': '0.8175', 'logits/rejected': '1.096', 'mean_token_accuracy': '0.5936', 'rewards/chosen': '-0.147', 'rewards/rejected': '-0.2055', 'rewards/accuracies': '0.55', 'rewards/margins': '0.05849', 'logps/chosen': '-1611', 'logps/rejected': '-1394', 'epoch': '0.07273'} {'eval_loss': '0.6538', 'eval_runtime': '263.1', 'eval_samples_per_second': '1.136', 'eval_steps_per_second': '0.144', 'eval_entropy': '1.456', 'eval_num_tokens': '1.613e+05', 'eval_logits/chosen': '0.8848', 'eval_logits/rejected': '1.132', 'eval_mean_token_accuracy': '0.5933', 'eval_rewards/chosen': '-0.2268', 'eval_rewards/rejected': '-0.3615', 'eval_rewards/accuracies': '0.6118', 'eval_rewards/margins': '0.1347', 'eval_logps/chosen': '-1583', 'eval_logps/rejected': '-1398', 'epoch': '0.07273'} {'loss': '0.6776', 'grad_norm': '26.73', 'learning_rate': '3.423e-06', 'entropy': '1.461', 'num_tokens': '2.426e+05', 'logits/chosen': '0.8494', 'logits/rejected': '1.165', 'mean_token_accuracy': '0.5897', 'rewards/chosen': '-0.2107', 'rewards/rejected': '-0.3011', 'rewards/accuracies': '0.6', 'rewards/margins': '0.09039', 'logps/chosen': '-1606', 'logps/rejected': '-1445', 'epoch': '0.1091'} {'loss': '0.5729', 'grad_norm': '25.62', 'learning_rate': '3.216e-06', 'entropy': '1.433', 'num_tokens': '3.223e+05', 'logits/chosen': '0.8954', 'logits/rejected': '1.113', 'mean_token_accuracy': '0.5907', 'rewards/chosen': '-0.2589', 'rewards/rejected': '-0.5807', 'rewards/accuracies': '0.725', 'rewards/margins': '0.3218', 'logps/chosen': '-1577', 'logps/rejected': '-1392', 'epoch': '0.1455'} {'eval_loss': '0.5215', 'eval_runtime': '262.6', 'eval_samples_per_second': '1.139', 'eval_steps_per_second': '0.145', 'eval_entropy': '1.455', 'eval_num_tokens': '3.223e+05', 'eval_logits/chosen': '0.8848', 'eval_logits/rejected': '1.131', 'eval_mean_token_accuracy': '0.5935', 'eval_rewards/chosen': '-0.3105', 'eval_rewards/rejected': '-0.7597', 'eval_rewards/accuracies': '0.8158', 'eval_rewards/margins': '0.4492', 'eval_logps/chosen': '-1584', 'eval_logps/rejected': '-1400', 'epoch': '0.1455'} {'loss': '0.4748', 'grad_norm': '19.57', 'learning_rate': '3.01e-06', 'entropy': '1.448', 'num_tokens': '4.024e+05', 'logits/chosen': '0.8937', 'logits/rejected': '1.151', 'mean_token_accuracy': '0.5905', 'rewards/chosen': '-0.4272', 'rewards/rejected': '-0.9954', 'rewards/accuracies': '0.875', 'rewards/margins': '0.5682', 'logps/chosen': '-1604', 'logps/rejected': '-1396', 'epoch': '0.1818'} {'loss': '0.4201', 'grad_norm': '15.62', 'learning_rate': '2.804e-06', 'entropy': '1.436', 'num_tokens': '4.829e+05', 'logits/chosen': '0.931', 'logits/rejected': '1.146', 'mean_token_accuracy': '0.5934', 'rewards/chosen': '-0.5315', 'rewards/rejected': '-1.295', 'rewards/accuracies': '0.875', 'rewards/margins': '0.7631', 'logps/chosen': '-1586', 'logps/rejected': '-1374', 'epoch': '0.2182'} {'eval_loss': '0.3951', 'eval_runtime': '261.7', 'eval_samples_per_second': '1.143', 'eval_steps_per_second': '0.145', 'eval_entropy': '1.454', 'eval_num_tokens': '4.829e+05', 'eval_logits/chosen': '0.8867', 'eval_logits/rejected': '1.134', 'eval_mean_token_accuracy': '0.5934', 'eval_rewards/chosen': '-0.5125', 'eval_rewards/rejected': '-1.343', 'eval_rewards/accuracies': '0.9046', 'eval_rewards/margins': '0.8301', 'eval_logps/chosen': '-1584', 'eval_logps/rejected': '-1402', 'epoch': '0.2182'} {'loss': '0.3736', 'grad_norm': '16.5', 'learning_rate': '2.598e-06', 'entropy': '1.458', 'num_tokens': '5.635e+05', 'logits/chosen': '0.8576', 'logits/rejected': '1.14', 'mean_token_accuracy': '0.5949', 'rewards/chosen': '-0.556', 'rewards/rejected': '-1.466', 'rewards/accuracies': '0.95', 'rewards/margins': '0.91', 'logps/chosen': '-1603', 'logps/rejected': '-1409', 'epoch': '0.2545'} {'loss': '0.3498', 'grad_norm': '16.54', 'learning_rate': '2.392e-06', 'entropy': '1.458', 'num_tokens': '6.434e+05', 'logits/chosen': '0.8618', 'logits/rejected': '1.138', 'mean_token_accuracy': '0.5927', 'rewards/chosen': '-0.726', 'rewards/rejected': '-1.732', 'rewards/accuracies': '0.925', 'rewards/margins': '1.005', 'logps/chosen': '-1599', 'logps/rejected': '-1395', 'epoch': '0.2909'} {'eval_loss': '0.3013', 'eval_runtime': '262.2', 'eval_samples_per_second': '1.14', 'eval_steps_per_second': '0.145', 'eval_entropy': '1.454', 'eval_num_tokens': '6.434e+05', 'eval_logits/chosen': '0.888', 'eval_logits/rejected': '1.136', 'eval_mean_token_accuracy': '0.5935', 'eval_rewards/chosen': '-0.7372', 'eval_rewards/rejected': '-1.949', 'eval_rewards/accuracies': '0.9605', 'eval_rewards/margins': '1.212', 'eval_logps/chosen': '-1585', 'eval_logps/rejected': '-1405', 'epoch': '0.2909'} {'loss': '0.2906', 'grad_norm': '13.11', 'learning_rate': '2.186e-06', 'entropy': '1.433', 'num_tokens': '7.24e+05', 'logits/chosen': '0.8954', 'logits/rejected': '1.148', 'mean_token_accuracy': '0.5962', 'rewards/chosen': '-0.851', 'rewards/rejected': '-2.113', 'rewards/accuracies': '0.95', 'rewards/margins': '1.262', 'logps/chosen': '-1572', 'logps/rejected': '-1436', 'epoch': '0.3273'} {'loss': '0.2298', 'grad_norm': '13.48', 'learning_rate': '1.979e-06', 'entropy': '1.424', 'num_tokens': '8.042e+05', 'logits/chosen': '0.8616', 'logits/rejected': '1.15', 'mean_token_accuracy': '0.5913', 'rewards/chosen': '-0.791', 'rewards/rejected': '-2.354', 'rewards/accuracies': '0.975', 'rewards/margins': '1.563', 'logps/chosen': '-1607', 'logps/rejected': '-1360', 'epoch': '0.3636'} {'eval_loss': '0.2059', 'eval_runtime': '262.4', 'eval_samples_per_second': '1.14', 'eval_steps_per_second': '0.145', 'eval_entropy': '1.454', 'eval_num_tokens': '8.042e+05', 'eval_logits/chosen': '0.8875', 'eval_logits/rejected': '1.139', 'eval_mean_token_accuracy': '0.5934', 'eval_rewards/chosen': '-0.8336', 'eval_rewards/rejected': '-2.586', 'eval_rewards/accuracies': '0.977', 'eval_rewards/margins': '1.752', 'eval_logps/chosen': '-1586', 'eval_logps/rejected': '-1407', 'epoch': '0.3636'} {'loss': '0.222', 'grad_norm': '8.034', 'learning_rate': '1.773e-06', 'entropy': '1.475', 'num_tokens': '8.835e+05', 'logits/chosen': '0.8554', 'logits/rejected': '1.093', 'mean_token_accuracy': '0.5906', 'rewards/chosen': '-1.095', 'rewards/rejected': '-2.744', 'rewards/accuracies': '0.975', 'rewards/margins': '1.65', 'logps/chosen': '-1617', 'logps/rejected': '-1368', 'epoch': '0.4'} {'loss': '0.1913', 'grad_norm': '9.952', 'learning_rate': '1.567e-06', 'entropy': '1.445', 'num_tokens': '9.632e+05', 'logits/chosen': '0.7705', 'logits/rejected': '1.081', 'mean_token_accuracy': '0.594', 'rewards/chosen': '-1.036', 'rewards/rejected': '-3.077', 'rewards/accuracies': '0.975', 'rewards/margins': '2.041', 'logps/chosen': '-1596', 'logps/rejected': '-1351', 'epoch': '0.4364'} {'eval_loss': '0.1529', 'eval_runtime': '262.2', 'eval_samples_per_second': '1.14', 'eval_steps_per_second': '0.145', 'eval_entropy': '1.453', 'eval_num_tokens': '9.632e+05', 'eval_logits/chosen': '0.8913', 'eval_logits/rejected': '1.142', 'eval_mean_token_accuracy': '0.5931', 'eval_rewards/chosen': '-1.181', 'eval_rewards/rejected': '-3.39', 'eval_rewards/accuracies': '0.9803', 'eval_rewards/margins': '2.209', 'eval_logps/chosen': '-1587', 'eval_logps/rejected': '-1410', 'epoch': '0.4364'} {'loss': '0.1591', 'grad_norm': '4.528', 'learning_rate': '1.361e-06', 'entropy': '1.422', 'num_tokens': '1.042e+06', 'logits/chosen': '0.8682', 'logits/rejected': '1.072', 'mean_token_accuracy': '0.6012', 'rewards/chosen': '-1.066', 'rewards/rejected': '-3.329', 'rewards/accuracies': '0.975', 'rewards/margins': '2.263', 'logps/chosen': '-1560', 'logps/rejected': '-1357', 'epoch': '0.4727'} {'loss': '0.1131', 'grad_norm': '7.361', 'learning_rate': '1.155e-06', 'entropy': '1.413', 'num_tokens': '1.122e+06', 'logits/chosen': '0.8699', 'logits/rejected': '1.129', 'mean_token_accuracy': '0.6021', 'rewards/chosen': '-1.194', 'rewards/rejected': '-3.633', 'rewards/accuracies': '1', 'rewards/margins': '2.439', 'logps/chosen': '-1548', 'logps/rejected': '-1421', 'epoch': '0.5091'} {'eval_loss': '0.1244', 'eval_runtime': '262.3', 'eval_samples_per_second': '1.14', 'eval_steps_per_second': '0.145', 'eval_entropy': '1.454', 'eval_num_tokens': '1.122e+06', 'eval_logits/chosen': '0.8879', 'eval_logits/rejected': '1.143', 'eval_mean_token_accuracy': '0.5932', 'eval_rewards/chosen': '-0.9224', 'eval_rewards/rejected': '-3.426', 'eval_rewards/accuracies': '0.9836', 'eval_rewards/margins': '2.503', 'eval_logps/chosen': '-1586', 'eval_logps/rejected': '-1411', 'epoch': '0.5091'} {'loss': '0.1075', 'grad_norm': '6.041', 'learning_rate': '9.485e-07', 'entropy': '1.464', 'num_tokens': '1.201e+06', 'logits/chosen': '0.8079', 'logits/rejected': '1.103', 'mean_token_accuracy': '0.5914', 'rewards/chosen': '-0.9542', 'rewards/rejected': '-3.529', 'rewards/accuracies': '0.975', 'rewards/margins': '2.574', 'logps/chosen': '-1566', 'logps/rejected': '-1403', 'epoch': '0.5455'} {'loss': '0.1837', 'grad_norm': '6.479', 'learning_rate': '7.423e-07', 'entropy': '1.419', 'num_tokens': '1.279e+06', 'logits/chosen': '0.8548', 'logits/rejected': '1.123', 'mean_token_accuracy': '0.6027', 'rewards/chosen': '-1.186', 'rewards/rejected': '-3.573', 'rewards/accuracies': '0.95', 'rewards/margins': '2.387', 'logps/chosen': '-1561', 'logps/rejected': '-1292', 'epoch': '0.5818'} {'eval_loss': '0.1107', 'eval_runtime': '262.5', 'eval_samples_per_second': '1.139', 'eval_steps_per_second': '0.145', 'eval_entropy': '1.455', 'eval_num_tokens': '1.279e+06', 'eval_logits/chosen': '0.8841', 'eval_logits/rejected': '1.141', 'eval_mean_token_accuracy': '0.5932', 'eval_rewards/chosen': '-0.9596', 'eval_rewards/rejected': '-3.634', 'eval_rewards/accuracies': '0.9836', 'eval_rewards/margins': '2.675', 'eval_logps/chosen': '-1586', 'eval_logps/rejected': '-1411', 'epoch': '0.5818'} {'loss': '0.07903', 'grad_norm': '6.533', 'learning_rate': '5.361e-07', 'entropy': '1.464', 'num_tokens': '1.359e+06', 'logits/chosen': '0.8634', 'logits/rejected': '1.161', 'mean_token_accuracy': '0.5882', 'rewards/chosen': '-0.9219', 'rewards/rejected': '-3.809', 'rewards/accuracies': '1', 'rewards/margins': '2.887', 'logps/chosen': '-1603', 'logps/rejected': '-1392', 'epoch': '0.6182'} {'loss': '0.07335', 'grad_norm': '5.129', 'learning_rate': '3.299e-07', 'entropy': '1.433', 'num_tokens': '1.44e+06', 'logits/chosen': '0.8821', 'logits/rejected': '1.128', 'mean_token_accuracy': '0.6019', 'rewards/chosen': '-1.154', 'rewards/rejected': '-4.002', 'rewards/accuracies': '1', 'rewards/margins': '2.847', 'logps/chosen': '-1560', 'logps/rejected': '-1455', 'epoch': '0.6545'} {'eval_loss': '0.09361', 'eval_runtime': '262.4', 'eval_samples_per_second': '1.14', 'eval_steps_per_second': '0.145', 'eval_entropy': '1.455', 'eval_num_tokens': '1.44e+06', 'eval_logits/chosen': '0.8885', 'eval_logits/rejected': '1.142', 'eval_mean_token_accuracy': '0.5931', 'eval_rewards/chosen': '-1.093', 'eval_rewards/rejected': '-4.037', 'eval_rewards/accuracies': '0.9868', 'eval_rewards/margins': '2.943', 'eval_logps/chosen': '-1587', 'eval_logps/rejected': '-1413', 'epoch': '0.6545'} {'loss': '0.07379', 'grad_norm': '4.132', 'learning_rate': '1.237e-07', 'entropy': '1.446', 'num_tokens': '1.52e+06', 'logits/chosen': '0.8279', 'logits/rejected': '1.091', 'mean_token_accuracy': '0.5916', 'rewards/chosen': '-1.101', 'rewards/rejected': '-4.017', 'rewards/accuracies': '1', 'rewards/margins': '2.916', 'logps/chosen': '-1606', 'logps/rejected': '-1405', 'epoch': '0.6909'} {'eval_loss': '0.09319', 'eval_runtime': '262.5', 'eval_samples_per_second': '1.139', 'eval_steps_per_second': '0.145', 'eval_entropy': '1.454', 'eval_num_tokens': '1.552e+06', 'eval_logits/chosen': '0.8901', 'eval_logits/rejected': '1.146', 'eval_mean_token_accuracy': '0.5931', 'eval_rewards/chosen': '-0.9987', 'eval_rewards/rejected': '-4.021', 'eval_rewards/accuracies': '0.9868', 'eval_rewards/margins': '3.022', 'eval_logps/chosen': '-1586', 'eval_logps/rejected': '-1413', 'epoch': '0.7055'} {'train_runtime': '3878', 'train_samples_per_second': '0.199', 'train_steps_per_second': '0.025', 'train_loss': '0.3075', 'epoch': '0.7055'} 100%|█████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████████| 97/97 [1:04:38<00:00, 39.98s/it] DPO训练完成,最优LoRA适配器输出至:./med_dpo_lora

1)rewards/accuracies

含义:偏好判断准确率

2)rewards/margins

含义:奖励差值margin

注意:DPO reward 带 \(\boldsymbol{\beta}\) 缩放,所以 reward 是相对值,不是 0‑1 分数。

margin >0:模型认为 chosen 更好;margin 越大,模型认为好坏差距越大。
rewards/margins:训练集平均 margin
eval_rewards/margins:验证集平均 margin

3)loss

loss:训练集上 DPO 损失的批次均值
eval_loss:验证集上 DPO 损失均值

4)entropy

含义:对每一步解码位置,模型输出词表维度概率分布 p,计算信息熵,训练时对一批样本做平均。

5、测试推理

代码:

import torch from transformers import AutoModelForCausalLM, AutoTokenizer from peft import PeftModel # ==================== 路径配置(与训练脚本完全对齐,无需修改)==================== BASE_MODEL_PATH = "./merged_sft_qwen7b_med" DPO_LORA_PATH = "./med_dpo_lora" # 医学固定系统提示词(和训练一致) SYSTEM_PROMPT = "你是专业的医疗咨询助手,回答仅供科普参考,不能替代执业医师面诊,诊疗请遵从线下医生的专业意见。" # 测试问题:全部为【训练集外】通用医学问题,检测泛化能力 TEST_QUESTIONS = [ "高血压患者日常饮食需要注意哪些事项?", "糖尿病患者如何科学控制餐后血糖?", "肝硬化患者常见的并发症有哪些?日常如何预防?", "慢性支气管炎患者秋冬季节如何养护?", "高血脂长期不控制,会对身体造成哪些危害?" ] # ==================== 加载模型 ==================== print("正在加载SFT基座模型...") tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL_PATH, trust_remote_code=True) base_model = AutoModelForCausalLM.from_pretrained( BASE_MODEL_PATH, torch_dtype=torch.bfloat16, device_map="auto", trust_remote_code=True ) print("正在加载DPO-LoRA微调模型...") dpo_model = PeftModel.from_pretrained(base_model, DPO_LORA_PATH) # ==================== 推理生成函数 ==================== def generate_answer(model, query): messages = [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": query} ] input_text = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True ) inputs = tokenizer(input_text, return_tensors="pt").to("cuda") # 通用稳定生成参数 outputs = model.generate( **inputs, max_new_tokens=1024, temperature=0.7, top_p=0.9, do_sample=True, eos_token_id=tokenizer.eos_token_id ) return tokenizer.decode(outputs[0][inputs["input_ids"].shape[-1]:], skip_special_tokens=True) # ==================== 批量对比测试 ==================== if __name__ == "__main__": for idx, question in enumerate(TEST_QUESTIONS, 1): print(f"\n{'=' * 80}") print(f"【测试问题 {idx}】:{question}") print(f"{'=' * 80}") print("\n[1] 【原始SFT基座模型回答】") print(generate_answer(base_model, question)) print("\n[2] 【SFT + DPO-LoRA 微调模型回答】") print(generate_answer(dpo_model, question))

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询