"""Pomiar Clef-flash na MASSIVE 1.1 (test, pl-PL i en-US). Wynik: JSONL, wznawialny.

Użycie: python pomiar.py --ids ids-proba.txt --out pelny.jsonl [--batch 8]
(dane MASSIVE 1.1 rozpakowane do katalogu 1.1/data obok skryptu)
"""
import argparse
import json
import sys
import time
from pathlib import Path

import torch
from huggingface_hub import snapshot_download

REWIZJA = "17f0b0ad64efb65d273590632833508766b2aae6"
DANE = Path(__file__).parent / "1.1" / "data"
INSTRUKCJA = "Which intent does the user's request express?"


def wczytaj(lokalizacja):
    return {r["id"]: r for r in map(json.loads, open(DANE / f"{lokalizacja}.jsonl")) if r["partition"] == "test"}


def schemat(intencje):
    return {"type": "choice", "instructions": INSTRUKCJA, "criteria": {i: i.replace("_", " ") for i in intencje}}


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--ids", required=True)
    parser.add_argument("--out", required=True)
    parser.add_argument("--batch", type=int, default=1)
    args = parser.parse_args()

    path = snapshot_download("Cloudflare/clef-flash", revision=REWIZJA)
    sys.path.insert(0, path)
    from joint_schema_model import collate_records, encode_record, load_release_model, systemone_answer

    model, processor = load_release_model(path, device="mps")
    wszystkie = sorted({json.loads(l)["intent"] for l in open(DANE / "pl-PL.jsonl")})
    pytanie = schemat(wszystkie)
    dane = {"pl-PL": wczytaj("pl-PL"), "en-US": wczytaj("en-US")}

    ids = [l.strip() for l in open(args.ids) if l.strip()]
    out = Path(args.out)
    zrobione = {(r["id"], r["locale"]) for r in map(json.loads, open(out))} if out.exists() else set()
    kolejka = [(rid, loc) for rid in ids for loc in ("pl-PL", "en-US") if (rid, loc) not in zrobione]

    with open(out, "a") as plik:
        for start in range(0, len(kolejka), args.batch):
            porcja = kolejka[start:start + args.batch]
            rekordy = [dane[loc][rid] for rid, loc in porcja]
            t0 = time.time()
            try:
                zakodowane = [
                    encode_record(processor.tokenizer, {"state": r["utt"], "questions": {"intent": pytanie}})
                    for r in rekordy
                ]
                with torch.inference_mode():
                    logity = model(collate_records(zakodowane, processor.tokenizer.pad_token_id, torch.device("mps")))
                torch.mps.synchronize()
                wyniki = []
                for enc, log in zip(zakodowane, logity):
                    probs = dict(zip(enc.questions[0].option_ids, log[0].float().softmax(-1).tolist()))
                    a = systemone_answer(pytanie, probs)
                    wyniki.append(dict(pred=a["choice"], conf=a["confidence"], probs=a["probabilities"],
                                       input_tokens=len(enc.input_ids), status="ok"))
            except Exception as blad:  # cała porcja do kategorii „nie zmierzono”
                wyniki = [dict(status="blad", blad=repr(blad)[:300])] * len(porcja)
            sekundy = round((time.time() - t0) / len(porcja), 3)
            for (rid, loc), r, w in zip(porcja, rekordy, wyniki):
                wiersz = {"id": rid, "locale": loc, "utt": r["utt"], "gold": r["intent"], **w,
                          "sekundy": sekundy, "batch": len(porcja)}
                plik.write(json.dumps(wiersz, ensure_ascii=False) + "\n")
            plik.flush()
            if (start // args.batch) % 25 == 0:
                print(f"{start + len(porcja)}/{len(kolejka)}", flush=True)


if __name__ == "__main__":
    main()
