#!/usr/bin/env python3
"""finetune_qwen_vla.py — QLoRA fine-tune Qwen2.5-VL on RX1 grounding data,
pipeline edition.

Adapted from world_model_test/rx1_brain/finetune_vla.py for trainpipe:
  * emits machine-readable progress:  PIPELINE_PROGRESS {"step": N, "loss": X}
  * saves the LoRA adapter after EVERY epoch (crash = lose one epoch, not all)
  * --resume <adapter_dir> continues from a previous adapter
  * --holdout 0.05 splits off an eval set to <out>/holdout.jsonl BEFORE
    training so eval_qwen_grounding.py scores unseen rows
  * last line:  PIPELINE_RESULT {"ok": ..., "steps": N, "final_loss": X}

    python finetune_qwen_vla.py --data <dir-or-jsonl> --out <adapter_dir> \
        --epochs 2 [--lr 1e-4] [--resume <adapter_dir>] [--holdout 0.05]
"""
from __future__ import annotations

import argparse
import json
import pathlib
import random
import sys

MODEL_ID = "Qwen/Qwen2.5-VL-3B-Instruct"
TARGET_MODULES = ["q_proj", "k_proj", "v_proj", "o_proj",
                  "gate_proj", "up_proj", "down_proj"]


def emit(tag: str, payload: dict):
    print(f"{tag} {json.dumps(payload)}", flush=True)


def load_examples(data: str):
    root = pathlib.Path(data)
    f = root if root.is_file() else root / "grounding.jsonl"
    rows = [json.loads(l) for l in open(f) if l.strip()]
    rows = [r for r in rows if r.get("in_view") and r.get("pixel")]
    base = root.parent if root.is_file() else root
    for r in rows:
        r["_img"] = str(base / r["image"])
    return rows


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--data", required=True)
    ap.add_argument("--out", required=True)
    ap.add_argument("--epochs", type=int, default=2)
    ap.add_argument("--lr", type=float, default=1e-4)
    ap.add_argument("--rank", type=int, default=16)
    ap.add_argument("--grad-accum", type=int, default=8)
    ap.add_argument("--max-samples", type=int, default=0)
    ap.add_argument("--holdout", type=float, default=0.05)
    ap.add_argument("--resume", default=None)
    a = ap.parse_args()

    import torch
    from PIL import Image
    from transformers import (AutoProcessor, BitsAndBytesConfig,
                              Qwen2_5_VLForConditionalGeneration)
    from peft import (LoraConfig, PeftModel, get_peft_model,
                      prepare_model_for_kbit_training)

    out_dir = pathlib.Path(a.out)
    out_dir.mkdir(parents=True, exist_ok=True)

    examples = load_examples(a.data)
    random.seed(7)
    random.shuffle(examples)
    if a.max_samples:
        examples = examples[: a.max_samples]

    # Holdout split first — eval must never see training rows. Kept stable
    # across resumes by seeding the shuffle above.
    n_hold = max(5, int(len(examples) * a.holdout))
    holdout, train_rows = examples[:n_hold], examples[n_hold:]
    with open(out_dir / "holdout.jsonl", "w") as fh:
        for r in holdout:
            fh.write(json.dumps({k: v for k, v in r.items() if k != "_img"} | {"_img": r["_img"]}) + "\n")
    emit("PIPELINE_PROGRESS", {"step": 0, "loss": None,
                               "note": f"{len(train_rows)} train / {n_hold} holdout"})

    quant = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_quant_type="nf4",
                               bnb_4bit_compute_dtype=torch.bfloat16,
                               bnb_4bit_use_double_quant=True)
    model = Qwen2_5_VLForConditionalGeneration.from_pretrained(
        MODEL_ID, quantization_config=quant, torch_dtype=torch.bfloat16,
        device_map="cuda:0")
    model = prepare_model_for_kbit_training(model)
    if a.resume and pathlib.Path(a.resume, "adapter_config.json").exists():
        model = PeftModel.from_pretrained(model, a.resume, is_trainable=True)
        emit("PIPELINE_PROGRESS", {"step": 0, "loss": None, "note": f"resumed adapter {a.resume}"})
    else:
        model = get_peft_model(model, LoraConfig(
            r=a.rank, lora_alpha=2 * a.rank, lora_dropout=0.05, bias="none",
            target_modules=TARGET_MODULES, task_type="CAUSAL_LM"))
    proc = AutoProcessor.from_pretrained(MODEL_ID, min_pixels=256 * 28 * 28,
                                         max_pixels=768 * 28 * 28)

    def encode(ex):
        img = Image.open(ex["_img"]).convert("RGB")
        answer = f"[{ex['pixel'][0]:.0f}, {ex['pixel'][1]:.0f}]"
        user = [{"type": "image", "image": img},
                {"type": "text", "text": ex["instruction"]}]
        msgs = [{"role": "user", "content": user},
                {"role": "assistant", "content": answer}]
        full = proc.apply_chat_template(msgs, tokenize=False)
        prompt = proc.apply_chat_template(msgs[:1], tokenize=False,
                                          add_generation_prompt=True)
        enc = proc(text=[full], images=[img], return_tensors="pt")
        plen = proc(text=[prompt], images=[img],
                    return_tensors="pt").input_ids.shape[1]
        labels = enc.input_ids.clone()
        labels[:, :plen] = -100                      # supervise only the answer
        enc["labels"] = labels
        return {k: v.to("cuda:0") for k, v in enc.items()}

    opt = torch.optim.AdamW(
        [p for p in model.parameters() if p.requires_grad], lr=a.lr)
    model.train()
    step, last_loss = 0, None
    for epoch in range(a.epochs):
        random.shuffle(train_rows)
        for i, ex in enumerate(train_rows):
            out = model(**encode(ex))
            (out.loss / a.grad_accum).backward()
            last_loss = out.loss.item()
            if (i + 1) % a.grad_accum == 0:
                opt.step()
                opt.zero_grad()
                step += 1
                if step % 10 == 0:
                    emit("PIPELINE_PROGRESS", {"step": step, "loss": round(last_loss, 4)})
        model.save_pretrained(out_dir)               # per-epoch: crash-safe
        emit("PIPELINE_PROGRESS", {"step": step, "loss": round(last_loss or 0, 4),
                                   "note": f"epoch {epoch + 1}/{a.epochs} saved"})

    emit("PIPELINE_RESULT", {"ok": True, "steps": step,
                             "final_loss": round(last_loss or 0, 4),
                             "adapter": str(out_dir)})


if __name__ == "__main__":
    try:
        main()
    except Exception as e:  # let the orchestrator classify from the log
        emit("PIPELINE_RESULT", {"ok": False, "message": f"{type(e).__name__}: {e}"})
        raise
