"""Direct preference optimisation on a LoRA adapter, with TRL's DPOTrainer.

Purpose: the course's reference preference-tuning run. Loads a preference dataset of
    prompt, chosen and rejected, trains an adapter against a frozen reference copy of
    the starting model, evaluates on a held-out split after every epoch, and appends
    one run record to the lab notebook. Beta and the learning rate are the two dials
    the lesson asks you to move, so they are arguments rather than constants.
Platform: spark, strix, nvidia (CUDA or ROCm), and mac on PyTorch's MPS backend in
    float32 with a 1B to 2B model. mlx-lm ships no preference trainer, so Track M
    takes the PyTorch path for this lab.
Minimum memory: 16 GB
Assumes: torch, transformers, trl, peft and datasets installed in the active
    environment; make-preference-pairs.py has been run so that <data-dir>/train.jsonl
    and <data-dir>/valid.jsonl exist; runlog.py sits next to this file.

Usage: python3 train-dpo.py --data-dir pairs --output-dir runs/dpo-style --labbook labbook.md
       python3 train-dpo.py --data-dir pairs --model Qwen/Qwen3-1.7B \
           --adapter runs/sft-my-format --beta 0.1 --lr 1e-5 --labbook labbook.md
       python3 train-dpo.py --data-dir pairs --loss-type ipo --beta 0.05 --labbook labbook.md
"""

from __future__ import annotations

import argparse
import json
import time
from pathlib import Path
from typing import Any

import torch
from datasets import load_dataset
from peft import LoraConfig
from transformers import AutoTokenizer
from trl import DPOConfig, DPOTrainer

import runlog

DEFAULT_TARGETS = ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"]

# Every value TRL's DPOConfig documents for loss_type. The lesson's table covers the
# five the course discusses; the rest are here so the script does not have to be
# edited to try one.
LOSS_TYPES = [
    "sigmoid", "hinge", "ipo", "exo_pair", "nca_pair", "robust", "bco_pair", "sppo_hard",
    "aot", "aot_unpaired", "apo_zero", "apo_down", "discopop", "sft", "sigmoid_norm",
]


def pick_device() -> str:
    if torch.cuda.is_available():
        return "cuda"
    mps = getattr(torch.backends, "mps", None)
    if mps is not None and mps.is_available():
        return "mps"
    return "cpu"


def use_bf16(device: str, requested: str) -> bool:
    if requested == "fp32":
        return False
    if requested == "bf16":
        return True
    return device == "cuda" and torch.cuda.is_bf16_supported()


def summarise_history(history: list[dict]) -> dict[str, Any]:
    """The numbers a DPO run is judged on, pulled out of a log full of them."""
    train_losses = [row["loss"] for row in history if "loss" in row and "eval_loss" not in row]
    evals = [(row.get("epoch"), row["eval_loss"]) for row in history if "eval_loss" in row]
    accuracies = [row["rewards/accuracies"] for row in history if "rewards/accuracies" in row]
    margins = [row["rewards/margins"] for row in history if "rewards/margins" in row]
    best_epoch, best_eval = min(evals, key=lambda pair: pair[1]) if evals else (None, None)
    return {
        "first_train_loss": round(train_losses[0], 4) if train_losses else None,
        "final_train_loss": round(train_losses[-1], 4) if train_losses else None,
        "final_eval_loss": round(evals[-1][1], 4) if evals else None,
        "best_eval_loss": round(best_eval, 4) if best_eval is not None else None,
        "best_epoch": best_epoch,
        "first_reward_accuracy": round(accuracies[0], 4) if accuracies else None,
        "final_reward_accuracy": round(accuracies[-1], 4) if accuracies else None,
        "final_reward_margin": round(margins[-1], 4) if margins else None,
    }


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    parser.add_argument("--model", default="Qwen/Qwen3-1.7B",
                        help="base model repository id or local path")
    parser.add_argument("--adapter", default=None,
                        help="an existing LoRA adapter to continue, e.g. the fine-tune from Part 13's lab")
    parser.add_argument("--data-dir", default="pairs", help="directory holding train.jsonl and valid.jsonl")
    parser.add_argument("--output-dir", default="runs/dpo")
    parser.add_argument("--beta", type=float, default=0.1,
                        help="how far the policy may move from the reference; higher means less deviation")
    parser.add_argument("--lr", type=float, default=1e-5,
                        help="TRL's DPO default is 1e-6; its documentation suggests about 1e-5 for adapters")
    parser.add_argument("--loss-type", default="sigmoid", choices=LOSS_TYPES)
    parser.add_argument("--label-smoothing", type=float, default=0.0,
                        help="Robust DPO's label-flip probability, in [0.0, 0.5)")
    parser.add_argument("--epochs", type=float, default=1.0)
    parser.add_argument("--batch-size", type=int, default=2)
    parser.add_argument("--grad-accum", type=int, default=4)
    parser.add_argument("--max-length", type=int, default=1024)
    parser.add_argument("--rank", type=int, default=16)
    parser.add_argument("--alpha", type=int, default=32)
    parser.add_argument("--dropout", type=float, default=0.05)
    parser.add_argument("--target-modules", nargs="+", default=DEFAULT_TARGETS)
    parser.add_argument("--precision", choices=["auto", "bf16", "fp32"], default="auto")
    parser.add_argument("--gradient-checkpointing", action="store_true")
    parser.add_argument("--precompute-ref-log-probs", action="store_true",
                        help="score the dataset with the reference model once, then drop it from memory")
    parser.add_argument("--report-to", default="none", choices=["none", "tensorboard"])
    parser.add_argument("--seed", type=int, default=0)
    parser.add_argument("--labbook", default=None)
    parser.add_argument("--notes", default=None)
    args = parser.parse_args()

    device = pick_device()
    bf16 = use_bf16(device, args.precision)
    print(f"device: {device}   precision: {'bfloat16' if bf16 else 'float32'}")

    data_dir = Path(args.data_dir)
    files = {"train": str(data_dir / "train.jsonl"), "validation": str(data_dir / "valid.jsonl")}
    for split, path in files.items():
        if not Path(path).is_file():
            raise SystemExit(f"{path} is missing; run make-preference-pairs.py first ({split} split)")
    dataset = load_dataset("json", data_files=files)
    print(f"preference pairs: {len(dataset['train'])} train, {len(dataset['validation'])} validation")

    columns = set(dataset["train"].column_names)
    for required in ("prompt", "chosen", "rejected"):
        if required not in columns:
            raise SystemExit(
                f"the dataset has no {required!r} column. TRL's DPOTrainer expects a preference "
                f"dataset with prompt, chosen and rejected; found {sorted(columns)}"
            )

    tokenizer = AutoTokenizer.from_pretrained(args.adapter or args.model)
    if tokenizer.chat_template is None:
        raise SystemExit(f"{args.model} has no chat template; pick an instruction-tuned model")

    config = DPOConfig(
        output_dir=args.output_dir,
        num_train_epochs=args.epochs,
        per_device_train_batch_size=args.batch_size,
        per_device_eval_batch_size=args.batch_size,
        gradient_accumulation_steps=args.grad_accum,
        learning_rate=args.lr,
        lr_scheduler_type="cosine",
        warmup_steps=5,
        beta=args.beta,
        loss_type=[args.loss_type],
        label_smoothing=args.label_smoothing,
        max_length=args.max_length,
        precompute_ref_log_probs=args.precompute_ref_log_probs,
        gradient_checkpointing=args.gradient_checkpointing,
        bf16=bf16,
        model_init_kwargs=None if args.adapter else {"dtype": torch.bfloat16 if bf16 else torch.float32},
        eval_strategy="epoch",
        save_strategy="epoch",
        save_total_limit=2,
        load_best_model_at_end=True,
        metric_for_best_model="eval_loss",
        greater_is_better=False,
        logging_steps=5,
        report_to=args.report_to,
        seed=args.seed,
        data_seed=args.seed,
    )

    trainer_kwargs: dict[str, Any] = {}
    if args.adapter:
        from peft import AutoPeftModelForCausalLM  # noqa: PLC0415 - only needed on this branch
        model = AutoPeftModelForCausalLM.from_pretrained(args.adapter, is_trainable=True)
        print(f"continuing the adapter in {args.adapter}; the reference is that model with the "
              f"adapter's starting weights")
    else:
        model = args.model
        trainer_kwargs["peft_config"] = LoraConfig(
            r=args.rank, lora_alpha=args.alpha, lora_dropout=args.dropout,
            target_modules=args.target_modules, bias="none", task_type="CAUSAL_LM",
        )

    trainer = DPOTrainer(
        model=model,
        args=config,
        train_dataset=dataset["train"],
        eval_dataset=dataset["validation"],
        processing_class=tokenizer,
        **trainer_kwargs,
    )

    started = time.time()
    trainer.train()
    elapsed = time.time() - started

    trainer.save_model(args.output_dir)
    tokenizer.save_pretrained(args.output_dir)
    losses = summarise_history(trainer.state.log_history)
    losses["seconds"] = round(elapsed, 1)
    print(json.dumps(losses, indent=2))
    print(f"adapter saved to {args.output_dir}")

    if args.labbook:
        record = runlog.record(
            labbook=args.labbook,
            lab="part-14/train-dpo",
            model=args.adapter or args.model,
            dataset={
                "path": files["train"],
                "sha256": runlog.file_sha256(files["train"]),
                "train_pairs": len(dataset["train"]),
                "validation_pairs": len(dataset["validation"]),
            },
            hyperparameters={
                "method": "dpo-lora",
                "beta": args.beta,
                "loss_type": args.loss_type,
                "label_smoothing": args.label_smoothing,
                "learning_rate": args.lr,
                "epochs": args.epochs,
                "batch_size": args.batch_size,
                "grad_accum": args.grad_accum,
                "effective_batch": args.batch_size * args.grad_accum,
                "max_length": args.max_length,
                "rank": args.rank,
                "alpha": args.alpha,
                "precompute_ref_log_probs": args.precompute_ref_log_probs,
                "precision": "bfloat16" if bf16 else "float32",
                "started_from_adapter": args.adapter,
            },
            seed=args.seed,
            losses=losses,
            scores={},
            config_path=__file__,
            notes=args.notes,
        )
        print(f"recorded run {record['run_id']} in {args.labbook}")


if __name__ == "__main__":
    main()
