"""Tworzy zwarte pliki decyzji i wyniki.json dla Strands Decider 2B (protokol.md).

Etap 1 (tylko u autora): z surowych zapisów (glowna-1000.jsonl, scenariusz-1000.jsonl, uzupelnienie.jsonl, wariant-24.jsonl, z pełnymi
rozkładami prawdopodobieństw, niepublikowane) buduje decyzje.jsonl, decyzje-scenariusz.jsonl, decyzje-uzupelnienie.jsonl i decyzje-24-opcje.jsonl.
p_etykiety i suma_p2 liczone z prawdopodobieństw zaokrąglonych do 4 miejsc, tak jak w surowym zapisie.
Etap 2 (każdy): miary z protokołu wyłącznie z opublikowanych plików decyzji; porównanie z Clef-flash z pliku
../clef-pl-2026-10/decyzje.jsonl (ten sam zbiór 1000 par, opublikowany 3.10.2026).
"""
import json
import math
import random
import statistics
from collections import Counter, defaultdict
from pathlib import Path

TU = Path(__file__).parent
ZIARNO = 20261002
ids = sorted({l.strip() for l in open(TU / "ids-proba.txt") if l.strip()}, key=int)
n = len(ids)


def zwarty(surowy, cel):
    rek = [json.loads(l) for l in open(TU / surowy)]
    assert all("pred" in r for r in rek), "rekordy z błędem: opisać jako „nie zmierzono”"
    with open(TU / cel, "w") as f:
        for w in sorted(rek, key=lambda x: (int(x["id"]), x["locale"] != "pl-PL")):
            f.write(json.dumps({
                "id": w["id"], "jezyk": w["locale"], "zdanie": w["utt"], "etykieta": w["gold"],
                "przewidziana": w["pred"], "pewnosc": w["conf"], "pewnosc_modelu": w["confidence_decider"],
                "p_etykiety": w["probs"].get(w["gold"], 0.0),
                "suma_p2": round(sum(p * p for p in w["probs"].values()), 8),
                "tokeny_wejscia": w["input_tokens"], "sekundy_mac": w["czas_s"],
            }, ensure_ascii=False) + "\n")


for surowy, cel in (("glowna-1000.jsonl", "decyzje.jsonl"), ("scenariusz-1000.jsonl", "decyzje-scenariusz.jsonl"),
                    ("uzupelnienie.jsonl", "decyzje-uzupelnienie.jsonl"), ("wariant-24.jsonl", "decyzje-24-opcje.jsonl")):
    if (TU / surowy).exists():
        zwarty(surowy, cel)

trafna = lambda w: w["przewidziana"] == w["etykieta"]
brier = lambda w: w["suma_p2"] - 2 * w["p_etykiety"] + 1


def wczytaj(plik):
    return {(w["id"], w["jezyk"]): w for w in map(json.loads, open(plik))}


def ece(ws, przedzialy=15):
    kosze = defaultdict(list)
    for w in ws:
        kosze[min(int(w["pewnosc"] * przedzialy), przedzialy - 1)].append(w)
    return sum(len(k) / len(ws) * abs(statistics.mean(trafna(x) for x in k) - statistics.mean(x["pewnosc"] for x in k))
               for k in kosze.values())


def makro_f1(ws):
    tp, fp, fn = Counter(), Counter(), Counter()
    for w in ws:
        if trafna(w):
            tp[w["etykieta"]] += 1
        else:
            fp[w["przewidziana"]] += 1
            fn[w["etykieta"]] += 1
    f1 = []
    for k in {w["etykieta"] for w in ws}:
        p = tp[k] / (tp[k] + fp[k]) if tp[k] + fp[k] else 0.0
        r = tp[k] / (tp[k] + fn[k]) if tp[k] + fn[k] else 0.0
        f1.append(2 * p * r / (p + r) if p + r else 0.0)
    return statistics.mean(f1)


def porownanie(a, b, klucze):
    """Sparowana różnica trafności a−b: bootstrap 10 000 losowań (ziarno 20261002), dokładny test McNemara."""
    m = len(klucze)
    r = random.Random(ZIARNO)
    roznice = sorted(sum(a[j] - b[j] for j in (r.randrange(m) for _ in range(m))) / m for _ in range(10_000))
    x = sum(1 for p, e in zip(a, b) if p and not e)
    y = sum(1 for p, e in zip(a, b) if e and not p)
    p_mc = min(1.0, 2 * sum(math.comb(x + y, i) for i in range(min(x, y) + 1)) / 2 ** (x + y)) if x + y else 1.0
    return {"punkty_proc": round(100 * (statistics.mean(a) - statistics.mean(b)), 1),
            "ci95_punkty_proc": [round(100 * roznice[249], 1), round(100 * roznice[9749], 1)],
            "tylko_pierwszy_trafny": x, "tylko_drugi_trafny": y, "mcnemar_p": p_mc}


def miary(d, klucze):
    out = {}
    for l in ("pl-PL", "en-US"):
        ws = [d[(i, l)] for i in klucze]
        pewne = [w for w in ws if w["pewnosc"] >= 0.9]
        out[l] = {
            "trafnosc_intencji": round(statistics.mean(trafna(w) for w in ws), 4),
            "trafnosc_obszaru_z_intencji": round(statistics.mean(w["przewidziana"].split("_")[0] == w["etykieta"].split("_")[0] for w in ws), 4),
            "makro_f1": round(makro_f1(ws), 4),
            "ece": round(ece(ws), 4),
            "brier": round(statistics.mean(brier(w) for w in ws), 4),
            "srednia_pewnosc": round(statistics.mean(w["pewnosc"] for w in ws), 4),
            "pewnosc_od_0_9": {"decyzje": len(pewne), "bledne": sum(not trafna(w) for w in pewne)},
            "tokeny_srednio": round(statistics.mean(w["tokeny_wejscia"] for w in ws), 2),
            "tokeny_mediana": statistics.median(w["tokeny_wejscia"] for w in ws),
            "sekundy_mac_mediana": round(statistics.median(w["sekundy_mac"] for w in ws), 3),
        }
    pl = [trafna(d[(i, "pl-PL")]) for i in klucze]
    en = [trafna(d[(i, "en-US")]) for i in klucze]
    roz = porownanie(pl, en, klucze)
    roz["ta_sama_decyzja"] = sum(d[(i, "pl-PL")]["przewidziana"] == d[(i, "en-US")]["przewidziana"] for i in klucze)
    return out, roz


glowna = wczytaj(TU / "decyzje.jsonl")
assert len(glowna) == 2 * n
wynik = {"pary": n, "decyzje": len(glowna)}
wynik["jezyki"], wynik["roznica_pl_en"] = miary(glowna, ids)

clef_plik = TU.parent / "clef-pl-2026-10" / "decyzje.jsonl"
if clef_plik.exists():
    clef = wczytaj(clef_plik)
    wynik["clef_flash"] = {}
    for l in ("pl-PL", "en-US"):
        a = [trafna(glowna[(i, l)]) for i in ids]
        b = [trafna(clef[(i, l)]) for i in ids]
        wynik["clef_flash"][l] = {"trafnosc_intencji_clef_flash": round(statistics.mean(b), 4),
                                  "decider_minus_clef_flash": porownanie(a, b, ids),
                                  "tokeny_srednio_clef": round(statistics.mean(clef[(i, l)]["tokeny_wejscia"] for i in ids), 2),
                                  "sekundy_mac_mediana_clef": statistics.median(clef[(i, l)]["sekundy_mac"] for i in ids)}

scen = wczytaj(TU / "decyzje-scenariusz.jsonl")
assert len(scen) == 2 * n
wynik["scenariusz_18_opcji"] = {l: {"trafnosc": round(statistics.mean(trafna(scen[(i, l)]) for i in ids), 4),
                                     "ece": round(ece([scen[(i, l)] for i in ids]), 4),
                                     "srednia_pewnosc": round(statistics.mean(scen[(i, l)]["pewnosc"] for i in ids), 4)}
                                 for l in ("pl-PL", "en-US")}
wynik["scenariusz_18_opcji"]["roznica_pl_en"] = porownanie([trafna(scen[(i, "pl-PL")]) for i in ids],
                                                           [trafna(scen[(i, "en-US")]) for i in ids], ids)

# Wariant z 24 opcjami (zapisany w protokole przed uruchomieniem, ale po analizie głównej: analiza po fakcie)
if (TU / "decyzje-24-opcje.jsonl").exists():
    w24 = wczytaj(TU / "decyzje-24-opcje.jsonl")
    assert len(w24) == 2 * n
    wynik["wariant_24_opcje"] = {l: {"trafnosc_intencji": round(statistics.mean(trafna(w24[(i, l)]) for i in ids), 4),
                                      "ece": round(ece([w24[(i, l)] for i in ids]), 4),
                                      "srednia_pewnosc": round(statistics.mean(w24[(i, l)]["pewnosc"] for i in ids), 4)}
                                  for l in ("pl-PL", "en-US")}
    wynik["wariant_24_opcje"]["roznica_pl_en"] = porownanie([trafna(w24[(i, "pl-PL")]) for i in ids],
                                                            [trafna(w24[(i, "en-US")]) for i in ids], ids)
    wynik["wariant_24_opcje"]["24_minus_60"] = {l: porownanie([trafna(w24[(i, l)]) for i in ids],
                                                              [trafna(glowna[(i, l)]) for i in ids], ids)
                                                for l in ("pl-PL", "en-US")}

# Analizy po fakcie (spoza protokołu)
po_fakcie = {"pole_confidence_modelu": {}, "najczestsze_pomylki_intencji": {}, "najczestsze_pomylki_scenariusza": {}}
for l in ("pl-PL", "en-US"):
    ws = [glowna[(i, l)] for i in ids]
    for prog in (0.8, 0.9):
        p = [w for w in ws if w["pewnosc_modelu"] >= prog]
        po_fakcie["pole_confidence_modelu"].setdefault(l, {})[str(prog)] = {
            "decyzje": len(p), "trafnosc": round(statistics.mean(trafna(w) for w in p), 4)}
    po_fakcie["najczestsze_pomylki_intencji"][l] = [
        [g, p, c] for (g, p), c in Counter((w["etykieta"], w["przewidziana"]) for w in ws if not trafna(w)).most_common(5)]
    ss = [scen[(i, l)] for i in ids]
    po_fakcie["najczestsze_pomylki_scenariusza"][l] = [
        [g, p, c] for (g, p), c in Counter((w["etykieta"], w["przewidziana"]) for w in ss if not trafna(w)).most_common(5)]
wynik["po_fakcie"] = po_fakcie

if (TU / "decyzje-uzupelnienie.jsonl").exists():
    uz = wczytaj(TU / "decyzje-uzupelnienie.jsonl")
    ids_uz = sorted({k[0] for k in uz}, key=int)
    if len(uz) == 2 * len(ids_uz):
        caly = {**glowna, **uz}
        ids_caly = sorted(set(ids) | set(ids_uz), key=int)
        wynik["uzupelnienie"] = {"pary": len(ids_uz)}
        wynik["uzupelnienie"]["jezyki"], wynik["uzupelnienie"]["roznica_pl_en"] = miary(uz, ids_uz)
        wynik["cala_czesc_testowa"] = {"pary": len(ids_caly)}
        wynik["cala_czesc_testowa"]["jezyki"], wynik["cala_czesc_testowa"]["roznica_pl_en"] = miary(caly, ids_caly)

# Po fakcie: skąd wzrost przy 24 opcjach i próg pola confidence 0,9 w zakresie treningu.
# Zestaw 24 opcji odtwarzany tak jak w pomiar.py (random.Random(f"20261002-{id}"), 23 z pozostałych 59 etykiet).
if "wariant_24_opcje" in wynik and (TU / "decyzje-uzupelnienie.jsonl").exists():
    etykiety = sorted({w[k] for plik in ("decyzje.jsonl", "decyzje-uzupelnienie.jsonl") for w in map(json.loads, open(TU / plik))
                       for k in ("etykieta", "przewidziana")})
    assert len(etykiety) == 60
    for l in ("pl-PL", "en-US"):
        rozbicie = Counter()
        for i in ids:
            a, b = glowna[(i, l)], w24[(i, l)]
            opcje = {a["etykieta"], *random.Random(f"{ZIARNO}-{i}").sample([e for e in etykiety if e != a["etykieta"]], 23)}
            assert b["przewidziana"] in opcje
            grupa = "trafne_przy_60" if trafna(a) else (
                "bledne_przy_60_wybrana_w_24" if a["przewidziana"] in opcje else "bledne_przy_60_wybranej_brak_w_24")
            rozbicie[f"{grupa}, {'trafne' if trafna(b) else 'bledne'}_przy_24"] += 1
        pewne = [w24[(i, l)] for i in ids if w24[(i, l)]["pewnosc_modelu"] >= 0.9]
        wynik["wariant_24_opcje"][l]["po_fakcie"] = {
            "pole_confidence_od_0_9": {"decyzje": len(pewne), "trafnosc": round(statistics.mean(trafna(w) for w in pewne), 4)},
            "rozbicie_wzgledem_60": dict(sorted(rozbicie.items()))}

# Rozkład z pytania o 60 opcji zawężony do zestawu 24 opcji (analiza po fakcie). Wymaga pełnych rozkładów
# z surowego zapisu glowna-1000.jsonl (niepublikowany); bez niego przenosimy wynik z poprzedniego wyniki.json.
if "wariant_24_opcje" in wynik:
    if (TU / "glowna-1000.jsonl").exists():
        surowe = {(r["id"], r["locale"]): r for r in map(json.loads, open(TU / "glowna-1000.jsonl"))}
        w24 = wczytaj(TU / "decyzje-24-opcje.jsonl")
        etykiety = sorted(next(iter(surowe.values()))["probs"])
        zaw = {}
        for l in ("pl-PL", "en-US"):
            traf_zaw = []
            for i in ids:
                r = surowe[(i, l)]
                los = random.Random(f"20261002-{i}")
                zestaw = set([r["gold"]] + los.sample([e for e in etykiety if e != r["gold"]], 23))
                traf_zaw.append(max(zestaw, key=lambda e: r["probs"][e]) == r["gold"])
            traf_24 = [trafna(w24[(i, l)]) for i in ids]
            zaw[l] = {"trafnosc_zawezonego_rozkladu_60": round(statistics.mean(traf_zaw), 4),
                      "pytanie_24_minus_zawezony_60": porownanie(traf_24, traf_zaw, ids)}
        wynik["wariant_24_opcje"]["zawezony_rozklad_60"] = zaw
    elif (TU / "wyniki.json").exists():
        stary = json.load(open(TU / "wyniki.json")).get("wariant_24_opcje", {})
        if "zawezony_rozklad_60" in stary:
            wynik["wariant_24_opcje"]["zawezony_rozklad_60"] = stary["zawezony_rozklad_60"]

json.dump(wynik, open(TU / "wyniki.json", "w"), ensure_ascii=False, indent=2)
print(json.dumps(wynik, ensure_ascii=False, indent=1)[:3000])
