pathcosmos/EVAFRILL-Mo-3B
130
1"""2train/dpo.py — Direct Preference Optimization (DPO) training.3 4Native DPO implementation (no TRL dependency) for EVAFRILL-Mo hybrid models.5Supports LoRA adapters for memory-efficient training on single GPU.6 7Launch:8 python train/dpo.py \9 --sft_checkpoint checkpoints/3b_sft_v2/checkpoint-best \10 --dpo_data data/preference/combined_preference.jsonl \11 --config configs/h100_mig/dpo_3b_1gpu.yaml \12 --device cuda:013"""14 15from __future__ import annotations16 17import argparse18import os19import random20import signal21import shutil22import sys23from pathlib import Path24 25import numpy as np26import torch27import torch.nn as nn28import torch.nn.functional as F29from torch.utils.data import DataLoader, RandomSampler30 31torch.backends.cuda.matmul.allow_tf32 = True32torch.backends.cudnn.allow_tf32 = True33torch.set_float32_matmul_precision("high")34 35_PROJECT_ROOT = Path(__file__).resolve().parent.parent36if str(_PROJECT_ROOT) not in sys.path:37 sys.path.insert(0, str(_PROJECT_ROOT))38 39from model import LLM40from model.lora import apply_lora, get_lora_params, merge_lora, save_lora41from data.dpo_dataset import DPODataset, dpo_collate_fn42from train.utils import (43 get_cosine_schedule_with_warmup,44 is_main_process,45 save_checkpoint,46 load_checkpoint,47)48 49 50def parse_args() -> argparse.Namespace:51 parser = argparse.ArgumentParser(description="DPO Training for EVAFRILL-Mo")52 53 # Paths54 parser.add_argument("--sft_checkpoint", type=Path, required=True,55 help="Path to SFT checkpoint directory")56 parser.add_argument("--dpo_data", type=Path, required=True,57 help="Path to preference JSONL data")58 parser.add_argument("--checkpoint_dir", type=Path, default=Path("checkpoints/3b_dpo"),59 help="Output checkpoint directory")60 parser.add_argument("--resume", type=Path, default=None)61 parser.add_argument("--tokenizer", type=Path, default=None)62 parser.add_argument("--log_file", type=Path, default=None)63 parser.add_argument("--config", type=Path, default=None)64 65 # DPO hyperparameters66 parser.add_argument("--beta", type=float, default=0.1, help="DPO temperature")67 parser.add_argument("--max_steps", type=int, default=3000)68 parser.add_argument("--batch_size", type=int, default=1)69 parser.add_argument("--grad_accum", type=int, default=16)70 parser.add_argument("--lr", type=float, default=5e-7)71 parser.add_argument("--weight_decay", type=float, default=0.01)72 parser.add_argument("--warmup_steps", type=int, default=100)73 parser.add_argument("--max_length", type=int, default=1024)74 parser.add_argument("--seed", type=int, default=42)75 76 # LoRA77 parser.add_argument("--use_lora", action="store_true", default=True)78 parser.add_argument("--lora_rank", type=int, default=32)79 parser.add_argument("--lora_alpha", type=float, default=64.0)80 81 # Infra82 parser.add_argument("--device", type=str, default=None)83 parser.add_argument("--save_interval", type=int, default=500)84 parser.add_argument("--log_interval", type=int, default=10)85 parser.add_argument("--num_workers", type=int, default=4)86 87 args, _ = parser.parse_known_args()88 89 # Load YAML config90 if args.config is not None:91 if not args.config.exists():92 raise FileNotFoundError(f"Config not found: {args.config}")93 import yaml94 with open(args.config) as f:95 cfg = yaml.safe_load(f)96 train_cfg = cfg.get("train", {})97 yaml_map = {98 "max_steps": "max_steps", "batch_size": "batch_size",99 "grad_accum_steps": "grad_accum", "lr": "lr",100 "weight_decay": "weight_decay", "warmup_steps": "warmup_steps",101 "beta": "beta", "max_length": "max_length",102 "save_interval": "save_interval", "log_interval": "log_interval",103 "use_lora": "use_lora", "lora_rank": "lora_rank", "lora_alpha": "lora_alpha",104 }105 defaults = {}106 for yk, ak in yaml_map.items():107 if yk in train_cfg:108 defaults[ak] = train_cfg[yk]109 if defaults:110 parser.set_defaults(**defaults)111 112 return parser.parse_args()113 114 115def set_seed(seed: int) -> None:116 random.seed(seed)117 np.random.seed(seed)118 torch.manual_seed(seed)119 torch.cuda.manual_seed_all(seed)120 121 122def compute_log_probs(123 model: nn.Module,124 input_ids: torch.Tensor,125 labels: torch.Tensor,126) -> torch.Tensor:127 """Compute sum of log probabilities over non-masked tokens.128 129 Args:130 model: The LLM model131 input_ids: (B, T) token ids132 labels: (B, T) target ids, -1 for masked positions133 134 Returns:135 (B,) sum of log probs per sample136 """137 with torch.autocast(device_type="cuda", dtype=torch.bfloat16):138 logits, _ = model(input_ids) # (B, T, V)139 140 # Shift: predict next token141 # logits[:, :-1] predicts labels[:, 1:]142 # But our labels already have the shifted targets (same as SFT convention)143 # labels[i] = token_id means input_ids[i] should predict labels[i]144 log_probs = F.log_softmax(logits.float(), dim=-1) # (B, T, V)145 146 # Gather log probs for target tokens147 # For each position, get log_prob of the label token148 mask = labels != -1 # (B, T)149 # Clamp labels for gather (replace -1 with 0, will be masked out)150 safe_labels = labels.clamp(min=0) # (B, T)151 per_token_logps = log_probs.gather(-1, safe_labels.unsqueeze(-1)).squeeze(-1) # (B, T)152 per_token_logps = per_token_logps * mask.float() # zero out masked positions153 154 return per_token_logps.sum(dim=-1) # (B,)155 156 157def dpo_loss(158 policy_chosen_logps: torch.Tensor,159 policy_rejected_logps: torch.Tensor,160 ref_chosen_logps: torch.Tensor,161 ref_rejected_logps: torch.Tensor,162 beta: float = 0.1,163) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:164 """Compute DPO loss.165 166 Returns:167 (loss, chosen_rewards, rejected_rewards)168 """169 chosen_rewards = beta * (policy_chosen_logps - ref_chosen_logps)170 rejected_rewards = beta * (policy_rejected_logps - ref_rejected_logps)171 172 logits = chosen_rewards - rejected_rewards # (B,)173 loss = -F.logsigmoid(logits).mean()174 175 return loss, chosen_rewards.detach().mean(), rejected_rewards.detach().mean()176 177 178def _resolve_tokenizer_path(args: argparse.Namespace) -> Path:179 if args.tokenizer is not None:180 return Path(args.tokenizer)181 ckpt_tok = args.sft_checkpoint / "tokenizer.json"182 if ckpt_tok.exists():183 return ckpt_tok184 default_tok = _PROJECT_ROOT / "tokenizer" / "korean_sp" / "tokenizer.json"185 if default_tok.exists():186 return default_tok187 raise FileNotFoundError("Cannot find tokenizer.json")188 189 190def main() -> None:191 args = parse_args()192 set_seed(args.seed)193 194 # Device setup195 if args.device:196 device = torch.device(args.device)197 elif torch.cuda.is_available():198 device = torch.device("cuda:0")199 else:200 device = torch.device("cpu")201 202 # Validate checkpoint203 if not args.sft_checkpoint.exists():204 raise FileNotFoundError(f"SFT checkpoint not found: {args.sft_checkpoint}")205 206 # Load SFT model as policy207 print(f"Loading SFT model from {args.sft_checkpoint}...")208 model = LLM.from_pretrained(args.sft_checkpoint)209 model.config.use_fp8 = False # H100 MIG: BF16 only210 model = model.to(device=device, dtype=torch.bfloat16)211 212 # Enable gradient checkpointing213 if hasattr(model, 'gradient_checkpointing_enable'):214 model.gradient_checkpointing_enable()215 print("[INFO] Gradient checkpointing enabled")216 217 # Compute reference log probs BEFORE applying LoRA218 # (reference model = SFT model without LoRA)219 # We'll compute ref logps on-the-fly with LoRA disabled via a context manager220 # Actually for simplicity: precompute nothing, just use model without LoRA adapters221 # For LoRA DPO: ref_model is the base (original weights), policy is base + LoRA222 # Since LoRA is initialized to zero, at start policy = ref223 224 # Apply LoRA225 if args.use_lora:226 n_lora_params = apply_lora(model, rank=args.lora_rank, alpha=args.lora_alpha)227 lora_params = get_lora_params(model)228 print(f"[INFO] LoRA: {n_lora_params:,} trainable params")229 else:230 # Full fine-tuning (risky for VRAM)231 lora_params = None232 233 total_params = sum(p.numel() for p in model.parameters())234 trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)235 print(f"Total params: {total_params:,}, Trainable: {trainable_params:,}")236 237 # Tokenizer238 tokenizer_path = _resolve_tokenizer_path(args)239 print(f"Loading tokenizer from {tokenizer_path}")240 from tokenizers import Tokenizer241 tokenizer = Tokenizer.from_file(str(tokenizer_path))242 243 # Dataset244 train_dataset = DPODataset(245 data_path=args.dpo_data,246 tokenizer=tokenizer,247 max_seq_len=args.max_length,248 )249 250 train_loader = DataLoader(251 train_dataset,252 batch_size=args.batch_size,253 sampler=RandomSampler(train_dataset),254 num_workers=args.num_workers,255 pin_memory=True,256 drop_last=True,257 collate_fn=dpo_collate_fn,258 prefetch_factor=2,259 persistent_workers=True,260 )261 262 # Optimizer — only LoRA params if using LoRA263 if lora_params is not None:264 optimizer = torch.optim.AdamW(265 lora_params,266 lr=args.lr,267 betas=(0.9, 0.95),268 weight_decay=args.weight_decay,269 fused=torch.cuda.is_available(),270 )271 else:272 optimizer = torch.optim.AdamW(273 [p for p in model.parameters() if p.requires_grad],274 lr=args.lr,275 betas=(0.9, 0.95),276 weight_decay=args.weight_decay,277 fused=torch.cuda.is_available(),278 )279 280 scheduler = get_cosine_schedule_with_warmup(281 optimizer=optimizer,282 warmup_steps=args.warmup_steps,283 total_steps=args.max_steps,284 )285 286 # Resume287 start_step = 0288 if args.resume is not None:289 start_step, _ = load_checkpoint(args.resume, model, optimizer, scheduler)290 print(f"Resumed from step {start_step}")291 292 # Checkpoint dir293 args.checkpoint_dir.mkdir(parents=True, exist_ok=True)294 295 # Copy tokenizer296 dest_tok = args.checkpoint_dir / "tokenizer.json"297 if not dest_tok.exists():298 shutil.copy2(str(tokenizer_path), str(dest_tok))299 300 # Log file301 log_fh = None302 if args.log_file:303 Path(args.log_file).parent.mkdir(parents=True, exist_ok=True)304 log_fh = open(args.log_file, "a", encoding="utf-8", buffering=1)305 306 def log(msg: str, level: str = "INFO"):307 import datetime308 ts = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")309 line = f"[{ts}] [{level}] {msg}"310 print(line)311 if log_fh:312 log_fh.write(line + "\n")313 314 # Banner315 eff_batch = args.batch_size * args.grad_accum316 log(f"{'='*60}")317 log(f"DPO Training — EVAFRILL-Mo 3B")318 log(f" SFT ckpt: {args.sft_checkpoint}")319 log(f" DPO data: {args.dpo_data} ({len(train_dataset):,} samples)")320 log(f" LoRA: rank={args.lora_rank} alpha={args.lora_alpha}")321 log(f" beta={args.beta}, lr={args.lr:.2e}, eff_batch={eff_batch}")322 log(f" max_steps={args.max_steps}, max_length={args.max_length}")323 log(f" device={device}")324 log(f"{'='*60}")325 326 # Training loop327 import time328 model.train()329 loader_iter = iter(train_loader)330 epoch = 0331 332 def next_batch():333 nonlocal loader_iter, epoch334 try:335 return next(loader_iter)336 except StopIteration:337 epoch += 1338 loader_iter = iter(train_loader)339 return next(loader_iter)340 341 shutdown_requested = False342 def shutdown_handler(signum, frame):343 nonlocal shutdown_requested344 shutdown_requested = True345 log(f"Shutdown signal received ({signum})", "WARN")346 347 signal.signal(signal.SIGHUP, shutdown_handler)348 signal.signal(signal.SIGTERM, shutdown_handler)349 350 t0 = time.perf_counter()351 running_loss = 0.0352 running_chosen_reward = 0.0353 running_rejected_reward = 0.0354 log_step_count = 0355 356 for step in range(start_step, args.max_steps):357 optimizer.zero_grad(set_to_none=True)358 accum_loss = 0.0359 360 for micro in range(args.grad_accum):361 batch = next_batch()362 chosen_ids = batch[0].to(device, dtype=torch.long, non_blocking=True)363 chosen_labels = batch[1].to(device, dtype=torch.long, non_blocking=True)364 rejected_ids = batch[2].to(device, dtype=torch.long, non_blocking=True)365 rejected_labels = batch[3].to(device, dtype=torch.long, non_blocking=True)366 367 # Policy log probs (with LoRA active)368 policy_chosen_logps = compute_log_probs(model, chosen_ids, chosen_labels)369 policy_rejected_logps = compute_log_probs(model, rejected_ids, rejected_labels)370 371 # Reference log probs (LoRA disabled)372 # For LoRA: temporarily set lora scaling to 0373 with torch.no_grad():374 # Save and zero LoRA params375 if args.use_lora:376 saved_B = []377 for m in model.modules():378 from model.lora import LoRALinear379 if isinstance(m, LoRALinear):380 saved_B.append(m.lora_B.data.clone())381 m.lora_B.data.zero_()382 383 ref_chosen_logps = compute_log_probs(model, chosen_ids, chosen_labels)384 ref_rejected_logps = compute_log_probs(model, rejected_ids, rejected_labels)385 386 # Restore LoRA params387 if args.use_lora:388 idx = 0389 for m in model.modules():390 from model.lora import LoRALinear391 if isinstance(m, LoRALinear):392 m.lora_B.data.copy_(saved_B[idx])393 idx += 1394 395 # DPO loss396 loss, chosen_reward, rejected_reward = dpo_loss(397 policy_chosen_logps, policy_rejected_logps,398 ref_chosen_logps, ref_rejected_logps,399 beta=args.beta,400 )401 402 scaled_loss = loss / args.grad_accum403 scaled_loss.backward()404 accum_loss += loss.item()405 406 # Gradient clipping407 grad_norm = torch.nn.utils.clip_grad_norm_(408 [p for p in model.parameters() if p.requires_grad], 1.0409 ).item()410 411 optimizer.step()412 scheduler.step()413 414 avg_loss = accum_loss / args.grad_accum415 running_loss += avg_loss416 running_chosen_reward += chosen_reward.item()417 running_rejected_reward += rejected_reward.item()418 log_step_count += 1419 420 # Shutdown check421 if shutdown_requested:422 log(f"Graceful shutdown at step {step + 1}", "WARN")423 save_checkpoint(model, optimizer, scheduler, step + 1, avg_loss, str(args.checkpoint_dir))424 if args.use_lora:425 save_lora(model, args.checkpoint_dir / f"lora-{step+1:07d}")426 break427 428 # Logging429 if (step + 1) % args.log_interval == 0:430 t1 = time.perf_counter()431 elapsed = t1 - t0432 avg_l = running_loss / log_step_count433 avg_cr = running_chosen_reward / log_step_count434 avg_rr = running_rejected_reward / log_step_count435 margin = avg_cr - avg_rr436 lr = scheduler.get_last_lr()[0]437 mem_gb = torch.cuda.memory_allocated() / 1e9438 439 log(f"step {step+1:>6d} | loss {avg_l:.4f} | "440 f"margin {margin:.4f} (c={avg_cr:.3f} r={avg_rr:.3f}) | "441 f"lr {lr:.2e} | gnorm {grad_norm:.3f} | mem {mem_gb:.1f}GB")442 443 running_loss = 0.0444 running_chosen_reward = 0.0445 running_rejected_reward = 0.0446 log_step_count = 0447 t0 = t1448 449 # Save checkpoint450 if (step + 1) % args.save_interval == 0:451 ckpt_path = save_checkpoint(452 model, optimizer, scheduler, step + 1, avg_loss, str(args.checkpoint_dir)453 )454 if args.use_lora:455 save_lora(model, args.checkpoint_dir / f"lora-{step+1:07d}")456 log(f"Checkpoint saved -> {ckpt_path}")457 458 # Final save459 final_path = save_checkpoint(460 model, optimizer, scheduler, args.max_steps, avg_loss, str(args.checkpoint_dir)461 )462 if args.use_lora:463 save_lora(model, args.checkpoint_dir / "lora-final")464 # Also merge and save merged model465 log("Merging LoRA weights into base model...")466 merge_lora(model)467 model.save_pretrained(args.checkpoint_dir / "checkpoint-merged")468 log(f"Merged model saved -> {args.checkpoint_dir / 'checkpoint-merged'}")469 470 log(f"DPO training complete. Final checkpoint -> {final_path}")471 472 if log_fh:473 log_fh.close()474 475 476if __name__ == "__main__":477 main()478 