train_dpo

DPO fine-tune on top of marola's SFT adapter — Layer 3 in README.md (MIP-0025 §4.3).

Continues training from an SFT LoRA adapter (train_lora.py's output) using the preference pairs build_dpo_dataset.py wrote (finetune/data/dpo_pairs.jsonl: {"prompt", "chosen", "rejected"}, one pair per real Reviewer.scala reject/revise decision). Same preset ladder and the same per-preset out/<preset>/ layout as train_lora.py; an adapter trained on a different base is refused before anything loads, rather than after the base weights are in RAM.

STATUS: written against the documented peft/trl DPOTrainer API. See finetune/README.md's Layer 3 section for whether a real run's evidence has landed yet.

Usage: python build_dataset.py # SFT dataset (Layers 1+2) python train_lora.py --preset tiny --no-4bit --epochs 3 # SFT -> out/tiny/adapter python build_dpo_dataset.py # DPO pairs (Layer 3) python train_dpo.py --preset tiny --no-4bit --epochs 1 # DPO -> out/tiny/dpo-adapter

  1"""DPO fine-tune on top of marola's SFT adapter — Layer 3 in README.md (MIP-0025 §4.3).
  2
  3Continues training from an SFT LoRA adapter (train_lora.py's output) using the preference pairs
  4build_dpo_dataset.py wrote (finetune/data/dpo_pairs.jsonl: {"prompt", "chosen", "rejected"}, one
  5pair per real Reviewer.scala reject/revise decision). Same preset ladder and the same per-preset
  6`out/<preset>/` layout as train_lora.py; an adapter trained on a different base is refused before
  7anything loads, rather than after the base weights are in RAM.
  8
  9STATUS: written against the documented peft/trl DPOTrainer API. See finetune/README.md's Layer 3
 10section for whether a real run's evidence has landed yet.
 11
 12Usage:
 13    python build_dataset.py                                  # SFT dataset (Layers 1+2)
 14    python train_lora.py --preset tiny --no-4bit --epochs 3   # SFT -> out/tiny/adapter
 15    python build_dpo_dataset.py                               # DPO pairs (Layer 3)
 16    python train_dpo.py --preset tiny --no-4bit --epochs 1    # DPO -> out/tiny/dpo-adapter
 17"""
 18
 19from __future__ import annotations
 20
 21import argparse
 22from pathlib import Path
 23
 24from train_lora import (
 25    PRESETS,
 26    check_adapter_base,
 27    check_base,
 28    default_out,
 29    latest_checkpoint,
 30    mark_base,
 31    ollama_base_for,
 32    preset_for,
 33)
 34
 35
 36def main() -> None:
 37    ap = argparse.ArgumentParser()
 38    ap.add_argument(
 39        "--preset",
 40        choices=sorted(PRESETS),
 41        default=None,
 42        help="must match the SFT adapter's preset — tiny/small/base, see train_lora.py --help",
 43    )
 44    ap.add_argument("--base", default=None, help="explicit HF model id; overrides --preset")
 45    ap.add_argument(
 46        "--resume",
 47        action="store_true",
 48        help="continue from the newest checkpoint in --out if one exists. Trainer checkpoints "
 49        "carry optimizer, scheduler, RNG and step state, so this resumes mid-epoch rather than "
 50        "restarting the epoch — what makes a multi-day or interrupted run survivable",
 51    )
 52    ap.add_argument(
 53        "--save-steps",
 54        type=int,
 55        default=200,
 56        help="checkpoint every N steps; with --resume this bounds what a crash costs",
 57    )
 58    ap.add_argument(
 59        "--sft-adapter",
 60        default=None,
 61        help="the SFT LoRA adapter to continue training from, i.e. train_lora.py's --out "
 62        "(default: finetune/out/<preset>/adapter, the same per-preset layout)",
 63    )
 64    ap.add_argument("--data", default=str(Path(__file__).parent / "data" / "dpo_pairs.jsonl"))
 65    ap.add_argument(
 66        "--out",
 67        default=None,
 68        help="where the DPO adapter goes (default: finetune/out/<preset>/dpo-adapter)",
 69    )
 70    ap.add_argument("--epochs", type=int, default=1)
 71    ap.add_argument("--lr", type=float, default=5e-6)
 72    ap.add_argument("--beta", type=float, default=0.1, help="DPO KL-penalty strength")
 73    ap.add_argument("--max-length", type=int, default=1024)
 74    ap.add_argument(
 75        "--no-4bit", action="store_true", help="skip bitsandbytes 4-bit (CPU / no CUDA)"
 76    )
 77    args = ap.parse_args()
 78    if not args.base:
 79        args.base = PRESETS[args.preset or "tiny"]["hf"]
 80    preset_name = args.preset or preset_for(args.base)
 81    if not args.sft_adapter:
 82        args.sft_adapter = str(default_out(args.base, "adapter"))
 83    if not args.out:
 84        args.out = str(default_out(args.base, "dpo-adapter"))
 85    out = Path(args.out)
 86    if not Path(args.sft_adapter).exists():
 87        raise SystemExit(
 88            f"no SFT adapter at {args.sft_adapter} — run train_lora.py first: MIP-0025 task 5 is "
 89            "SFT (Layers 1+2) then DPO (Layer 3) continuing from it, not DPO trained from scratch"
 90        )
 91    # Both guards before torch: the adapter must belong to this base, and this directory must not
 92    # already hold another base's DPO run.
 93    check_adapter_base(Path(args.sft_adapter), args.base)
 94    check_base(out, args.base)
 95    print(
 96        f"base model: {args.base} ({preset_name or 'off-table'}) — Ollama FROM for "
 97        f"Modelfile.adapter: {ollama_base_for(args.base)} — continuing DPO from SFT adapter "
 98        f"{args.sft_adapter}, writing {out}"
 99    )
100
101    import torch
102    from datasets import load_dataset
103    from peft import PeftModel
104    from transformers import AutoModelForCausalLM, AutoTokenizer
105    from trl import DPOConfig, DPOTrainer
106
107    tok = AutoTokenizer.from_pretrained(args.base)
108    tok.pad_token = tok.pad_token or tok.eos_token
109
110    model_kwargs: dict = {
111        "torch_dtype": torch.bfloat16 if torch.cuda.is_available() else torch.float32
112    }
113    if not args.no_4bit and torch.cuda.is_available():
114        from transformers import BitsAndBytesConfig
115
116        model_kwargs["quantization_config"] = BitsAndBytesConfig(
117            load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_compute_dtype=torch.bfloat16
118        )
119        model_kwargs["device_map"] = "auto"
120    base_model = AutoModelForCausalLM.from_pretrained(args.base, **model_kwargs)
121    model = PeftModel.from_pretrained(base_model, args.sft_adapter, is_trainable=True)
122
123    data = load_dataset("json", data_files={"train": args.data})["train"]
124
125    cfg = DPOConfig(
126        output_dir=args.out,
127        num_train_epochs=args.epochs,
128        learning_rate=args.lr,
129        beta=args.beta,
130        per_device_train_batch_size=1,
131        gradient_accumulation_steps=2,
132        logging_steps=1,
133        max_length=args.max_length,
134        save_strategy="steps",
135        save_steps=args.save_steps,
136        bf16=torch.cuda.is_available(),
137        report_to=[],
138    )
139    trainer = DPOTrainer(model=model, args=cfg, train_dataset=data, processing_class=tok)
140    resume = latest_checkpoint(out) if args.resume else None
141    print(f"resuming from {resume}" if resume else "training from scratch (no checkpoint found)")
142    mark_base(out, args.base)
143    trainer.train(resume_from_checkpoint=resume)
144    trainer.save_model(args.out)
145    print(
146        f"DPO adapter saved to {args.out} — convert with llama.cpp's convert_lora_to_gguf.py, "
147        "then see Modelfile.adapter"
148    )
149
150
151if __name__ == "__main__":
152    main()
def main() -> None:
 37def main() -> None:
 38    ap = argparse.ArgumentParser()
 39    ap.add_argument(
 40        "--preset",
 41        choices=sorted(PRESETS),
 42        default=None,
 43        help="must match the SFT adapter's preset — tiny/small/base, see train_lora.py --help",
 44    )
 45    ap.add_argument("--base", default=None, help="explicit HF model id; overrides --preset")
 46    ap.add_argument(
 47        "--resume",
 48        action="store_true",
 49        help="continue from the newest checkpoint in --out if one exists. Trainer checkpoints "
 50        "carry optimizer, scheduler, RNG and step state, so this resumes mid-epoch rather than "
 51        "restarting the epoch — what makes a multi-day or interrupted run survivable",
 52    )
 53    ap.add_argument(
 54        "--save-steps",
 55        type=int,
 56        default=200,
 57        help="checkpoint every N steps; with --resume this bounds what a crash costs",
 58    )
 59    ap.add_argument(
 60        "--sft-adapter",
 61        default=None,
 62        help="the SFT LoRA adapter to continue training from, i.e. train_lora.py's --out "
 63        "(default: finetune/out/<preset>/adapter, the same per-preset layout)",
 64    )
 65    ap.add_argument("--data", default=str(Path(__file__).parent / "data" / "dpo_pairs.jsonl"))
 66    ap.add_argument(
 67        "--out",
 68        default=None,
 69        help="where the DPO adapter goes (default: finetune/out/<preset>/dpo-adapter)",
 70    )
 71    ap.add_argument("--epochs", type=int, default=1)
 72    ap.add_argument("--lr", type=float, default=5e-6)
 73    ap.add_argument("--beta", type=float, default=0.1, help="DPO KL-penalty strength")
 74    ap.add_argument("--max-length", type=int, default=1024)
 75    ap.add_argument(
 76        "--no-4bit", action="store_true", help="skip bitsandbytes 4-bit (CPU / no CUDA)"
 77    )
 78    args = ap.parse_args()
 79    if not args.base:
 80        args.base = PRESETS[args.preset or "tiny"]["hf"]
 81    preset_name = args.preset or preset_for(args.base)
 82    if not args.sft_adapter:
 83        args.sft_adapter = str(default_out(args.base, "adapter"))
 84    if not args.out:
 85        args.out = str(default_out(args.base, "dpo-adapter"))
 86    out = Path(args.out)
 87    if not Path(args.sft_adapter).exists():
 88        raise SystemExit(
 89            f"no SFT adapter at {args.sft_adapter} — run train_lora.py first: MIP-0025 task 5 is "
 90            "SFT (Layers 1+2) then DPO (Layer 3) continuing from it, not DPO trained from scratch"
 91        )
 92    # Both guards before torch: the adapter must belong to this base, and this directory must not
 93    # already hold another base's DPO run.
 94    check_adapter_base(Path(args.sft_adapter), args.base)
 95    check_base(out, args.base)
 96    print(
 97        f"base model: {args.base} ({preset_name or 'off-table'}) — Ollama FROM for "
 98        f"Modelfile.adapter: {ollama_base_for(args.base)} — continuing DPO from SFT adapter "
 99        f"{args.sft_adapter}, writing {out}"
100    )
101
102    import torch
103    from datasets import load_dataset
104    from peft import PeftModel
105    from transformers import AutoModelForCausalLM, AutoTokenizer
106    from trl import DPOConfig, DPOTrainer
107
108    tok = AutoTokenizer.from_pretrained(args.base)
109    tok.pad_token = tok.pad_token or tok.eos_token
110
111    model_kwargs: dict = {
112        "torch_dtype": torch.bfloat16 if torch.cuda.is_available() else torch.float32
113    }
114    if not args.no_4bit and torch.cuda.is_available():
115        from transformers import BitsAndBytesConfig
116
117        model_kwargs["quantization_config"] = BitsAndBytesConfig(
118            load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_compute_dtype=torch.bfloat16
119        )
120        model_kwargs["device_map"] = "auto"
121    base_model = AutoModelForCausalLM.from_pretrained(args.base, **model_kwargs)
122    model = PeftModel.from_pretrained(base_model, args.sft_adapter, is_trainable=True)
123
124    data = load_dataset("json", data_files={"train": args.data})["train"]
125
126    cfg = DPOConfig(
127        output_dir=args.out,
128        num_train_epochs=args.epochs,
129        learning_rate=args.lr,
130        beta=args.beta,
131        per_device_train_batch_size=1,
132        gradient_accumulation_steps=2,
133        logging_steps=1,
134        max_length=args.max_length,
135        save_strategy="steps",
136        save_steps=args.save_steps,
137        bf16=torch.cuda.is_available(),
138        report_to=[],
139    )
140    trainer = DPOTrainer(model=model, args=cfg, train_dataset=data, processing_class=tok)
141    resume = latest_checkpoint(out) if args.resume else None
142    print(f"resuming from {resume}" if resume else "training from scratch (no checkpoint found)")
143    mark_base(out, args.base)
144    trainer.train(resume_from_checkpoint=resume)
145    trainer.save_model(args.out)
146    print(
147        f"DPO adapter saved to {args.out} — convert with llama.cpp's convert_lora_to_gguf.py, "
148        "then see Modelfile.adapter"
149    )