CoolFace
Modelpublic

pathcosmos/EVAFRILL-Mo-3B

sourceHugging Facemitupdated 6mo agoView on Hugging Face
1likes30downloads
dpo.py478 linesDownload Raw Back to scripts
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