"""Fine-tune a 4B to 8B model with QLoRA: a 4-bit frozen base and trainable LoRA adapters.

Purpose: the project's training run. Loads the base in 4-bit NormalFloat with double
         quantisation so that an eight-billion-parameter model fits a 16 GB machine, attaches
         adapters to every linear layer in QLoRA's own style, trains with an evaluation after
         each epoch and early stopping, reports peak memory, and appends a run record to the
         lab notebook. Passing --no-4bit trains the same configuration with a bfloat16 base,
         which is what Track M and any machine without bitsandbytes use.
Platform: spark, nvidia (CUDA) and strix (ROCm; bitsandbytes lists gfx1151 among its ROCm wheel
          targets from ROCm 6.4.4, read 2026-09-09). Not mac: bitsandbytes' installation page
          shows Apple silicon only in its CPU build table, so Track M runs --no-4bit here, or
          trains against an already-quantised MLX model with mlx_lm.lora instead.
Minimum memory: 16 GB. Part 11's arithmetic: a 4-bit base is about 0.55 bytes per parameter,
          the adapter carries 16 bytes per trainable parameter, and activations depend on your
          batch size and sequence length.
Assumes: torch, transformers, trl, peft and datasets installed, plus bitsandbytes unless
         --no-4bit is used; make-domain-dataset.py has been run so that data/train.jsonl and
         data/valid.jsonl exist; sftlog.py sits next to this file.

Usage: python3 train-qlora.py --model Qwen/Qwen3-8B --data-dir data \
           --output-dir runs/domain-qwen3-8b --labbook labbook.md
       python3 train-qlora.py --model Qwen/Qwen3-4B --no-4bit --rank 32 --alpha 64
"""
from __future__ import annotations

import argparse
import importlib.util
import json
import time
from pathlib import Path

import torch
from datasets import load_dataset
from peft import LoraConfig
from transformers import AutoTokenizer, BitsAndBytesConfig, EarlyStoppingCallback
from trl import SFTConfig, SFTTrainer

import sftlog


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 check_4bit_available(device: str) -> None:
    """Fail with the reason rather than with a stack trace three minutes into a download."""
    if importlib.util.find_spec("bitsandbytes") is None:
        raise SystemExit(
            "bitsandbytes is not installed, so --load-in-4bit cannot work.\n"
            "  CUDA and ROCm: pip install bitsandbytes, then re-run.\n"
            "  macOS: the installation page lists Apple silicon only under its CPU builds, so\n"
            "         there is no 4-bit GPU path here. Re-run with --no-4bit and a smaller\n"
            "         model, or use mlx_lm.lora against an already-quantised MLX model."
        )
    if device != "cuda":
        raise SystemExit(
            f"the active device is {device!r}, and 4-bit loading with bitsandbytes is documented\n"
            "for CUDA and ROCm devices, both of which PyTorch reports as 'cuda'.\n"
            "Re-run with --no-4bit, and choose a model your memory tier can hold at bfloat16."
        )


def summarise_history(history: list[dict]) -> dict[str, float | int | None]:
    train_losses = [row["loss"] for row in history if "loss" in row]
    evals = [(row.get("epoch"), row["eval_loss"]) for row in history if "eval_loss" 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,
        "evaluations": len(evals),
    }


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    parser.add_argument("--model", default="Qwen/Qwen3-8B",
                        help="base model repository id or local path; an instruct checkpoint")
    parser.add_argument("--data-dir", default="data")
    parser.add_argument("--output-dir", default="runs/domain-qlora")
    parser.add_argument("--no-4bit", dest="load_in_4bit", action="store_false",
                        help="train against a bfloat16 base instead of a 4-bit one")
    parser.set_defaults(load_in_4bit=True)
    parser.add_argument("--epochs", type=float, default=3.0)
    parser.add_argument("--batch-size", type=int, default=1)
    parser.add_argument("--grad-accum", type=int, default=8)
    parser.add_argument("--lr", type=float, default=1e-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=["all-linear"],
                        help="QLoRA's own configuration adapts every linear layer, which PEFT "
                             "expresses as the single value all-linear")
    parser.add_argument("--gradient-checkpointing", action="store_true", default=True)
    parser.add_argument("--no-gradient-checkpointing", dest="gradient_checkpointing",
                        action="store_false")
    parser.add_argument("--early-stopping-patience", type=int, default=2)
    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()
    if args.load_in_4bit:
        check_4bit_available(device)
    bf16 = device == "cuda" and torch.cuda.is_bf16_supported()
    dtype = torch.bfloat16 if bf16 else torch.float32
    method = "qlora" if args.load_in_4bit else "lora"
    print(f"device: {device}   method: {method}   "
          f"compute 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 ({split} split); run make-domain-dataset.py first")
    dataset = load_dataset("json", data_files=files)
    print(f"train examples: {len(dataset['train'])}   "
          f"validation examples: {len(dataset['validation'])}")

    tokenizer = AutoTokenizer.from_pretrained(args.model)
    if tokenizer.chat_template is None:
        raise SystemExit(f"{args.model} has no chat template; choose an instruct checkpoint")

    quantization_config = None
    if args.load_in_4bit:
        # The four keys PEFT's quantisation guide names for QLoRA: 4-bit loading, the NF4
        # data type, double quantisation of the quantisation constants, and bfloat16 for
        # the arithmetic. The frozen base is what gets quantised; the adapters do not.
        quantization_config = BitsAndBytesConfig(
            load_in_4bit=True,
            bnb_4bit_quant_type="nf4",
            bnb_4bit_use_double_quant=True,
            bnb_4bit_compute_dtype=torch.bfloat16,
        )

    targets = args.target_modules[0] if args.target_modules == ["all-linear"] else args.target_modules

    config = SFTConfig(
        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=10,
        max_length=args.max_length,
        packing=False,
        completion_only_loss=True,
        gradient_checkpointing=args.gradient_checkpointing,
        bf16=bf16,
        model_init_kwargs=None if args.load_in_4bit else {"dtype": dtype},
        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="none",
        seed=args.seed,
        data_seed=args.seed,
    )

    peft_config = LoraConfig(
        r=args.rank,
        lora_alpha=args.alpha,
        lora_dropout=args.dropout,
        target_modules=targets,
        bias="none",
        task_type="CAUSAL_LM",
    )

    callbacks = []
    if args.early_stopping_patience > 0:
        callbacks.append(EarlyStoppingCallback(early_stopping_patience=args.early_stopping_patience))

    if device == "cuda":
        torch.cuda.reset_peak_memory_stats()

    # quantization_config plus peft_config is the documented QLoRA path through SFTTrainer:
    # the trainer loads the base with the quantisation applied and wraps it for adapter
    # training, so the model is never held at full precision.
    trainer = SFTTrainer(
        model=args.model,
        args=config,
        train_dataset=dataset["train"],
        eval_dataset=dataset["validation"],
        processing_class=tokenizer,
        peft_config=peft_config,
        quantization_config=quantization_config,
        callbacks=callbacks or None,
    )
    trainer.model.print_trainable_parameters()

    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)
    peak_gb = None
    if device == "cuda":
        peak_gb = round(torch.cuda.max_memory_allocated() / 1e9, 2)
        losses["peak_memory_gb"] = peak_gb
    print(json.dumps(losses, indent=2))
    print(f"adapter saved to {args.output_dir}")
    if peak_gb is not None:
        print(f"peak allocated memory: {peak_gb} GB. Compare it with the estimate you made "
              "from Part 11's arithmetic before the run, and record both.")
    print("\nThe adapter was trained against a 4-bit base." if args.load_in_4bit else
          "\nThe adapter was trained against a bfloat16 base.")
    print("Merge it into the base at bfloat16, never into the quantised copy, and quantise "
          "the merged result afterwards. merge-adapter.py does this.")

    if args.labbook:
        record = sftlog.record(
            labbook=args.labbook,
            lab="part-13/train-qlora",
            model=args.model,
            dataset={
                "path": files["train"],
                "sha256": sftlog.file_sha256(files["train"]),
                "train_examples": len(dataset["train"]),
                "validation_examples": len(dataset["validation"]),
            },
            hyperparameters={
                "method": method,
                "load_in_4bit": args.load_in_4bit,
                "bnb_4bit_quant_type": "nf4" if args.load_in_4bit else None,
                "bnb_4bit_use_double_quant": args.load_in_4bit,
                "rank": args.rank,
                "alpha": args.alpha,
                "dropout": args.dropout,
                "target_modules": targets,
                "epochs": args.epochs,
                "batch_size": args.batch_size,
                "grad_accum": args.grad_accum,
                "effective_batch": args.batch_size * args.grad_accum,
                "learning_rate": args.lr,
                "max_length": args.max_length,
                "gradient_checkpointing": args.gradient_checkpointing,
                "early_stopping_patience": args.early_stopping_patience,
                "compute_precision": "bfloat16" if bf16 else "float32",
                "peak_memory_gb": peak_gb,
                "output_dir": args.output_dir,
            },
            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()
