"""LoRA fine-tune. Runs on the rented GPU, not on the collection box.

Loss is computed on the assistant turn only. Without that the model spends most of
its capacity learning to reproduce briefs, which is the half we already have for free.

  python train/lora.py                       # 3 epochs, exports GGUF when done
  python train/lora.py --epochs 2 --lr 1e-4  # second run, if the first overfits

See train/README.md for the pod setup.
"""
import argparse
import json
import os

os.environ.setdefault("UNSLOTH_RETURN_LOGITS", "0")

from unsloth import FastLanguageModel                      # noqa: E402  (must precede transformers)
from unsloth.chat_templates import train_on_responses_only  # noqa: E402
from datasets import load_dataset                           # noqa: E402
from trl import SFTTrainer, SFTConfig                       # noqa: E402


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", default="unsloth/Qwen3-4B-Instruct-2507",
                    help="fallback if 4B won't fit locally: unsloth/Llama-3.2-3B-Instruct")
    ap.add_argument("--train", default="data/train.jsonl")
    ap.add_argument("--val", default="data/val.jsonl")
    ap.add_argument("--out", default="out")
    ap.add_argument("--epochs", type=float, default=3)
    ap.add_argument("--lr", type=float, default=2e-4)
    ap.add_argument("--rank", type=int, default=16)
    ap.add_argument("--seq", type=int, default=2048)
    ap.add_argument("--batch", type=int, default=2)
    ap.add_argument("--accum", type=int, default=8)
    ap.add_argument("--quants", default="q4_k_m,q5_k_m",
                    help="GGUF quantisations to export; empty to skip export")
    args = ap.parse_args()

    model, tokenizer = FastLanguageModel.from_pretrained(
        model_name=args.model,
        max_seq_length=args.seq,
        load_in_4bit=True,
        dtype=None,
    )

    model = FastLanguageModel.get_peft_model(
        model,
        r=args.rank,
        lora_alpha=args.rank * 2,
        lora_dropout=0.0,
        bias="none",
        target_modules=["q_proj", "k_proj", "v_proj", "o_proj",
                        "gate_proj", "up_proj", "down_proj"],
        use_gradient_checkpointing="unsloth",
        random_state=3407,
    )

    ds = load_dataset("json", data_files={"train": args.train, "val": args.val})

    def to_text(batch):
        return {"text": [tokenizer.apply_chat_template(m, tokenize=False)
                         for m in batch["messages"]]}

    ds = ds.map(to_text, batched=True, remove_columns=ds["train"].column_names)
    print(f"train {len(ds['train'])} / val {len(ds['val'])}")
    print("--- one formatted example ---")
    print(ds["train"][0]["text"][:900])

    trainer = SFTTrainer(
        model=model,
        tokenizer=tokenizer,
        train_dataset=ds["train"],
        eval_dataset=ds["val"],
        args=SFTConfig(
            output_dir=f"{args.out}/checkpoints",
            dataset_text_field="text",
            max_seq_length=args.seq,
            packing=False,          # packing and completion-only loss don't mix
            per_device_train_batch_size=args.batch,
            gradient_accumulation_steps=args.accum,
            num_train_epochs=args.epochs,
            learning_rate=args.lr,
            warmup_ratio=0.05,
            lr_scheduler_type="cosine",
            optim="adamw_8bit",
            weight_decay=0.01,
            logging_steps=10,
            eval_strategy="steps",
            eval_steps=50,
            save_strategy="epoch",
            save_total_limit=2,
            bf16=True,
            seed=3407,
            report_to="none",
        ),
    )

    # Mask everything up to the assistant turn. Qwen and Llama-3 both use ChatML-style
    # markers here; if you swap to a base model with a different template, print one
    # formatted example above and copy the exact delimiters.
    trainer = train_on_responses_only(
        trainer,
        instruction_part="<|im_start|>user\n",
        response_part="<|im_start|>assistant\n",
    )

    stats = trainer.train()
    print(stats)

    model.save_pretrained(f"{args.out}/adapter")
    tokenizer.save_pretrained(f"{args.out}/adapter")
    print(f"adapter -> {args.out}/adapter")

    metrics = trainer.evaluate()
    print("final eval:", metrics)
    with open(f"{args.out}/metrics.json", "w") as fh:
        json.dump({"args": vars(args), "eval": metrics}, fh, indent=2)

    for q in [q.strip() for q in args.quants.split(",") if q.strip()]:
        print(f"exporting GGUF {q} ...")
        model.save_pretrained_gguf(f"{args.out}/gguf-{q}", tokenizer,
                                   quantization_method=q)
        print(f"  -> {args.out}/gguf-{q}")


if __name__ == "__main__":
    main()
