| """LoRA DPO on live-fault chosen/rejected pairs. Not GRPO — that reward still pays for submits.""" |
|
|
| from __future__ import annotations |
|
|
| import json |
| from pathlib import Path |
|
|
| from .constants import ( |
| DEFAULT_DPO_ADAPTER_DIR, |
| DEFAULT_DPO_PACK, |
| LORA_ALPHA, |
| LORA_RANK, |
| ) |
| from .sft import _filter_kwargs, _load_pack, lora_target_modules |
|
|
|
|
| def train_dpo( |
| *, |
| pack_path: Path = DEFAULT_DPO_PACK, |
| output_dir: Path = DEFAULT_DPO_ADAPTER_DIR, |
| model_dir: Path, |
| max_seq_len: int = 4096, |
| max_steps: int = 30, |
| lr: float = 5e-6, |
| beta: float = 0.1, |
| per_device_batch_size: int = 1, |
| grad_accum: int = 8, |
| lora_rank: int = LORA_RANK, |
| smoke: bool = False, |
| ) -> Path: |
| from local_eval.cuda_env import apply as apply_cuda |
|
|
| apply_cuda() |
| pack_path = Path(pack_path) |
| output_dir = Path(output_dir) |
| output_dir.mkdir(parents=True, exist_ok=True) |
| rows = _load_pack(pack_path) |
| if not rows: |
| raise ValueError(f"empty dpo pack: {pack_path}") |
| if smoke: |
| rows = rows[:16] |
| max_steps = min(max_steps, 8) |
| max_seq_len = min(max_seq_len, 2048) |
|
|
| import torch |
| from datasets import Dataset |
| from peft import LoraConfig, get_peft_model |
| from transformers import AutoModelForCausalLM, AutoTokenizer |
| from trl import DPOConfig, DPOTrainer |
|
|
| tokenizer = AutoTokenizer.from_pretrained(str(model_dir), trust_remote_code=False) |
| if tokenizer.pad_token is None: |
| tokenizer.pad_token = tokenizer.eos_token |
| tokenizer.padding_side = "left" |
|
|
| dataset = Dataset.from_list( |
| [ |
| { |
| "prompt": row["prompt"], |
| "chosen": row["chosen"], |
| "rejected": row["rejected"], |
| } |
| for row in rows |
| ] |
| ) |
|
|
| model = AutoModelForCausalLM.from_pretrained( |
| str(model_dir), |
| torch_dtype=torch.bfloat16, |
| trust_remote_code=False, |
| attn_implementation="sdpa", |
| ) |
| model.config.use_cache = False |
| if hasattr(model, "enable_input_require_grads"): |
| model.enable_input_require_grads() |
| model = get_peft_model( |
| model, |
| LoraConfig( |
| r=lora_rank, |
| lora_alpha=LORA_ALPHA, |
| lora_dropout=0.05, |
| bias="none", |
| task_type="CAUSAL_LM", |
| target_modules=lora_target_modules(model), |
| ), |
| ) |
| model.print_trainable_parameters() |
|
|
| args_kwargs = dict( |
| output_dir=str(output_dir), |
| bf16=True, |
| learning_rate=lr, |
| per_device_train_batch_size=per_device_batch_size, |
| gradient_accumulation_steps=grad_accum, |
| gradient_checkpointing=True, |
| logging_steps=1, |
| save_steps=max(max_steps, 50), |
| warmup_ratio=0.03, |
| lr_scheduler_type="cosine", |
| report_to=[], |
| max_length=max_seq_len, |
| max_steps=max_steps, |
| beta=beta, |
| remove_unused_columns=False, |
| ) |
| config = DPOConfig(**_filter_kwargs(DPOConfig, args_kwargs)) |
| trainer = DPOTrainer( |
| model=model, |
| ref_model=None, |
| args=config, |
| train_dataset=dataset, |
| processing_class=tokenizer, |
| ) |
| trainer.train() |
| trainer.save_model(str(output_dir)) |
| tokenizer.save_pretrained(str(output_dir)) |
| (output_dir / "dpo-report.json").write_text( |
| json.dumps( |
| { |
| "pack": str(pack_path), |
| "model": str(model_dir), |
| "n": len(rows), |
| "max_steps": max_steps, |
| "max_seq_len": max_seq_len, |
| "lr": lr, |
| "beta": beta, |
| "lora_rank": lora_rank, |
| "smoke": smoke, |
| }, |
| indent=2, |
| ) |
| + "\n" |
| ) |
| print(f"dpo adapter: {output_dir}", flush=True) |
| return output_dir |
|
|