"""QLoRA fine-tune of Gemma 4 (E2B/E4B) on the SFT set.

    python scripts/train_qlora.py --config config/config.yaml       # (GPU required)

Primary path uses **Unsloth** (day-0 Gemma 4 support, ~2x faster / ~50% less VRAM, and it handles
the E-series multimodal wrapper + known gotchas for you). Only the LoRA adapter is trained; the
4-bit base weights stay frozen — that's what keeps VRAM/energy low. VRAM: E2B ~8-10GB, E4B ~17GB.

The adapter is written to train.output_adapter and merged/quantized by merge_and_quantize.py.

Gemma-4 QLoRA notes:
  - Gemma 4 supports a real system role, so our system/user/assistant SFT messages apply cleanly.
  - Training loss in the 13-15 range is normal for the E-series (multimodal) models.
  - Use bf16 (not fp16) to avoid overflow on the E-series.
"""
from __future__ import annotations

import argparse
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))

from agent.settings import load_config  # noqa: E402


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--config", default="config/config.yaml")
    args = ap.parse_args()
    cfg = load_config(args.config)

    from datasets import load_dataset
    from trl import SFTConfig, SFTTrainer
    from unsloth import FastModel
    from unsloth.chat_templates import get_chat_template

    base_id = cfg.get("model.base_id")
    train_file = str(cfg.resolve("data.sft_train"))
    output_adapter = str(cfg.resolve("train.output_adapter"))
    max_seq_len = int(cfg.get("train.max_seq_len", 2048))

    # --- Load Gemma 4 in 4-bit (QLoRA) ---
    model, tokenizer = FastModel.from_pretrained(
        model_name=base_id,
        max_seq_length=max_seq_len,
        load_in_4bit=True,          # 4-bit NF4 base — the QLoRA memory win
        full_finetuning=False,
    )
    tokenizer = get_chat_template(tokenizer, chat_template="gemma-4")

    # --- Attach LoRA adapters (text projections) ---
    model = FastModel.get_peft_model(
        model,
        r=int(cfg.get("train.lora_r", 16)),
        lora_alpha=int(cfg.get("train.lora_alpha", 32)),
        lora_dropout=float(cfg.get("train.lora_dropout", 0.05)),
        bias="none",
        target_modules=cfg.get("train.target_modules"),
        use_gradient_checkpointing="unsloth",
        random_state=3407,
    )

    # --- Dataset: apply Gemma 4 chat template to the {"messages": [...]} rows ---
    dataset = load_dataset("json", data_files=train_file, split="train")

    def format_chat(example):
        return {
            "text": tokenizer.apply_chat_template(
                example["messages"], tokenize=False, add_generation_prompt=False
            )
        }

    dataset = dataset.map(format_chat, remove_columns=dataset.column_names)

    sft_cfg = SFTConfig(
        output_dir=str(Path(output_adapter).parent / "trainer"),
        num_train_epochs=float(cfg.get("train.epochs", 2)),
        per_device_train_batch_size=int(cfg.get("train.batch_size", 2)),
        gradient_accumulation_steps=int(cfg.get("train.grad_accum", 8)),
        learning_rate=float(cfg.get("train.lr", 2e-4)),
        max_seq_length=max_seq_len,
        bf16=True,                  # E-series: bf16, not fp16
        packing=True,
        logging_steps=5,
        save_strategy="epoch",
        lr_scheduler_type="cosine",
        warmup_ratio=0.03,
        report_to="none",
        dataset_text_field="text",
    )

    trainer = SFTTrainer(model=model, tokenizer=tokenizer, args=sft_cfg, train_dataset=dataset)
    trainer.train()

    model.save_pretrained(output_adapter)
    tokenizer.save_pretrained(output_adapter)
    print(f"Saved LoRA adapter -> {output_adapter}")
    print("Next: python scripts/merge_and_quantize.py --config", args.config)


# ---------------------------------------------------------------------------
# Fallback (no Unsloth): plain transformers + PEFT + bitsandbytes.
# Gemma 4 E-series is multimodal, so load with the image-text-to-text auto class and target the
# TEXT decoder's projections. Prefer the Unsloth path above unless you can't install it.
#
#   from transformers import AutoModelForImageTextToText, AutoProcessor, BitsAndBytesConfig
#   bnb = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_quant_type="nf4",
#                            bnb_4bit_use_double_quant=True, bnb_4bit_compute_dtype=torch.bfloat16)
#   model = AutoModelForImageTextToText.from_pretrained(base_id, quantization_config=bnb,
#                                                       device_map="auto", torch_dtype=torch.bfloat16)
#   processor = AutoProcessor.from_pretrained(base_id)
#   # then LoraConfig(target_modules=[...]) + TRL SFTTrainer as usual.
# ---------------------------------------------------------------------------

if __name__ == "__main__":
    main()
