#!/usr/bin/env python3
"""eval_qwen_grounding.py — score a Qwen-VLA LoRA adapter on held-out
grounding rows.

A rollout "succeeds" when the model answers with parseable [u, v] pixel
coordinates within --tol pixels of ground truth. Last line:

    PIPELINE_RESULT {"success_rate": 0.72, "n": 25, "median_err_px": 11.4}

    python eval_qwen_grounding.py --adapter <dir> --holdout <holdout.jsonl> \
        [--n 25] [--tol 40] [--seed 0]
"""
from __future__ import annotations

import argparse
import json
import pathlib
import random
import re
import statistics
import sys

MODEL_ID = "Qwen/Qwen2.5-VL-3B-Instruct"
COORD_RE = re.compile(r"\[\s*(-?\d+(?:\.\d+)?)\s*,\s*(-?\d+(?:\.\d+)?)\s*\]")


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


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--adapter", required=True)
    ap.add_argument("--holdout", required=True)
    ap.add_argument("--n", type=int, default=25)
    ap.add_argument("--tol", type=float, default=40.0)
    ap.add_argument("--seed", type=int, default=0)
    a = ap.parse_args()

    rows = [json.loads(l) for l in open(a.holdout) if l.strip()]
    random.seed(a.seed)
    random.shuffle(rows)
    rows = rows[: a.n]
    if not rows:
        emit("PIPELINE_RESULT", {"ok": False, "message": "empty holdout"})
        sys.exit(1)

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

    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")
    if pathlib.Path(a.adapter, "adapter_config.json").exists():
        model = PeftModel.from_pretrained(model, a.adapter)
    model.eval()
    proc = AutoProcessor.from_pretrained(MODEL_ID, min_pixels=256 * 28 * 28,
                                         max_pixels=768 * 28 * 28)

    hits, errs = 0, []
    for r in rows:
        img = Image.open(r["_img"]).convert("RGB")
        msgs = [{"role": "user", "content": [
            {"type": "image", "image": img},
            {"type": "text", "text": r["instruction"]}]}]
        prompt = proc.apply_chat_template(msgs, tokenize=False,
                                          add_generation_prompt=True)
        enc = proc(text=[prompt], images=[img], return_tensors="pt").to("cuda:0")
        with torch.no_grad():
            out = model.generate(**enc, max_new_tokens=24, do_sample=False)
        text = proc.batch_decode(out[:, enc.input_ids.shape[1]:],
                                 skip_special_tokens=True)[0]
        m = COORD_RE.search(text)
        if not m:
            print(f"  miss (unparseable): {text!r}")
            continue
        u, v = float(m.group(1)), float(m.group(2))
        gu, gv = r["pixel"]
        err = ((u - gu) ** 2 + (v - gv) ** 2) ** 0.5
        errs.append(err)
        ok = err <= a.tol
        hits += ok
        print(f"  {'hit ' if ok else 'miss'} err {err:5.1f}px  {r['instruction'][:50]}")

    emit("PIPELINE_RESULT", {
        "ok": True,
        "success_rate": round(hits / len(rows), 3),
        "n": len(rows),
        "median_err_px": round(statistics.median(errs), 1) if errs else None,
    })


if __name__ == "__main__":
    try:
        main()
    except Exception as e:
        emit("PIPELINE_RESULT", {"ok": False, "message": f"{type(e).__name__}: {e}"})
        raise
