"""Uzupełnienie: pozostałe 1974 pary części testowej MASSIVE i cała część testowa (2974 pary).

Analiza główna (1000 par z ids-proba.txt, wyniki.json) zostaje bez zmian. Ten skrypt liczy te same miary dla par
spoza próby głównej i dla całej części testowej oraz sprawdza na nowych parach analizy po fakcie z próby głównej:
odsetek błędnych decyzji przy progach pewności, pasmo pewności 0,90–0,95, podział na rzadkie i częste intencje
(rzadkie = najwyżej 11 przykładów w próbie głównej, to samo kryterium co w README i przygotuj.py) oraz rozkład wyniku
Briera na decyzje błędne i trafne.

Etap 1 (tylko u autora): z surowego pelny.jsonl tworzy zwarty decyzje-uzupelnienie.jsonl (format jak decyzje.jsonl).
Etap 2 (każdy): wyniki-uzupelnienie.json z decyzje.jsonl i decyzje-uzupelnienie.jsonl.
"""
import json
import math
import random
import statistics
from collections import Counter, defaultdict
from pathlib import Path

TU = Path(__file__).parent
ZIARNO = 20261002
PROGI = (0.8, 0.85, 0.9, 0.95, 0.97)

glowne = [l.strip() for l in open(TU / "ids-proba.txt") if l.strip()]
pelne = [l.strip() for l in open(TU / "ids-pelny.txt") if l.strip()]
reszta = sorted(set(pelne) - set(glowne), key=int)
assert set(glowne) <= set(pelne) and len(reszta) == len(pelne) - len(glowne)

# Etap 1: surowy zapis -> decyzje-uzupelnienie.jsonl (pomijany, gdy surowego zapisu nie ma)
if not (TU / "pelny.jsonl").exists():
    print("Brak pelny.jsonl: pomijam etap 1 i liczę wyniki z istniejącego decyzje-uzupelnienie.jsonl.")
else:
    zbior = set(reszta)
    surowe = [w for w in map(json.loads, open(TU / "pelny.jsonl")) if w["id"] in zbior]
    assert len(surowe) == 2 * len(reszta) and all(w["status"] == "ok" for w in surowe)
    with open(TU / "decyzje-uzupelnienie.jsonl", "w") as f:
        for w in sorted(surowe, 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"], "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["sekundy"],
            }, ensure_ascii=False) + "\n")

# Etap 2
wszystkie = [json.loads(l) for f in ("decyzje.jsonl", "decyzje-uzupelnienie.jsonl") for l in open(TU / f)]
po = {l: {d["id"]: d for d in wszystkie if d["jezyk"] == l} for l in ("pl-PL", "en-US")}
assert all(set(po[l]) == set(pelne) for l in po)
trafna = lambda w: w["przewidziana"] == w["etykieta"]
licznosc = Counter(po["pl-PL"][i]["etykieta"] for i in glowne)
rzadkie = {k for k, v in licznosc.items() if v <= 11}
brier = lambda w: w["suma_p2"] - 2 * w["p_etykiety"] + 1


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 bledy_od_progu(ws, prog=0.9):
    pewne = [w for w in ws if w["pewnosc"] >= prog]
    return sum(not trafna(w) for w in pewne) / len(pewne)


def miary(ids):
    ids = sorted(ids, key=int)
    n = len(ids)
    wynik = {"pary": n, "jezyki": {}}
    for l in ("pl-PL", "en-US"):
        ws = [po[l][i] for i in ids]
        pewne = [w for w in ws if w["pewnosc"] >= 0.9]
        pasmo = [w for w in ws if 0.9 <= w["pewnosc"] < 0.95]
        wynik["jezyki"][l] = {
            "trafnosc_intencji": round(statistics.mean(trafna(w) for w in ws), 4),
            "trafnosc_obszaru": 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),
            "tokeny_srednio": round(statistics.mean(w["tokeny_wejscia"] for w in ws), 2),
            "pewnosc_od_0_9": {"decyzje": len(pewne), "bledne": sum(not trafna(w) for w in pewne)},
            "pewnosc_0_90_do_0_95": {"decyzje": len(pasmo), "trafne": sum(trafna(w) for w in pasmo)},
            "bledy": {"liczba": sum(not trafna(w) for w in ws),
                      "pewnosc_przy_bledach": round(statistics.mean(w["pewnosc"] for w in ws if not trafna(w)), 3),
                      "brier_z_blednych": round(sum(brier(w) for w in ws if not trafna(w)) / n, 4),
                      "brier_z_trafnych": round(sum(brier(w) for w in ws if trafna(w)) / n, 4)},
        }
    pl = [trafna(po["pl-PL"][i]) for i in ids]
    en = [trafna(po["en-US"][i]) for i in ids]
    r = random.Random(ZIARNO)
    roznice = sorted(sum(pl[j] - en[j] for j in (r.randrange(n) for _ in range(n))) / n for _ in range(10_000))
    b = sum(1 for p, e in zip(pl, en) if p and not e)
    c = sum(1 for p, e in zip(pl, en) if e and not p)
    p_mc = min(1.0, 2 * sum(math.comb(b + c, i) for i in range(min(b, c) + 1)) / 2 ** (b + c))
    r = random.Random(ZIARNO)
    rb, rp = [], []
    for _ in range(10_000):
        s = [ids[r.randrange(n)] for _ in range(n)]
        rb.append(statistics.mean(brier(po["pl-PL"][i]) - brier(po["en-US"][i]) for i in s))
        rp.append(bledy_od_progu([po["pl-PL"][i] for i in s]) - bledy_od_progu([po["en-US"][i] for i in s]))
    rb.sort()
    rp.sort()
    # pasmo pewności 0,90–0,95: różnica trafności PL−EN, bootstrap po parach (osobne losowanie z tym samym ziarnem)
    pasmo = lambda w: 0.9 <= w["pewnosc"] < 0.95
    traf_pasma = lambda ws: statistics.mean(trafna(w) for w in ws if pasmo(w))
    r = random.Random(ZIARNO)
    rpas = []
    for _ in range(10_000):
        s = [ids[r.randrange(n)] for _ in range(n)]
        rpas.append(traf_pasma([po["pl-PL"][i] for i in s]) - traf_pasma([po["en-US"][i] for i in s]))
    rpas.sort()
    wynik["roznica_pl_en"] = {
        "punkty_proc": round(100 * (statistics.mean(pl) - statistics.mean(en)), 1),
        "ci95_punkty_proc": [round(100 * roznice[249], 1), round(100 * roznice[9749], 1)],
        "tylko_pl_trafne": b, "tylko_en_trafne": c, "mcnemar_p": float(f"{p_mc:.3g}"),
        "ta_sama_decyzja": sum(po["pl-PL"][i]["przewidziana"] == po["en-US"][i]["przewidziana"] for i in ids),
        "brier": round(statistics.mean(brier(po["pl-PL"][i]) - brier(po["en-US"][i]) for i in ids), 4),
        "brier_ci95": [round(rb[249], 3), round(rb[9749], 3)],
        "bledy_od_0_9_punkty_proc": round(100 * (bledy_od_progu([po["pl-PL"][i] for i in ids]) - bledy_od_progu([po["en-US"][i] for i in ids])), 1),
        "bledy_od_0_9_ci95_punkty_proc": [round(100 * rp[249], 1), round(100 * rp[9749], 1)],
        "trafnosc_w_pasmie_0_90_0_95_punkty_proc": round(100 * (traf_pasma([po["pl-PL"][i] for i in ids]) - traf_pasma([po["en-US"][i] for i in ids])), 1),
        "trafnosc_w_pasmie_ci95_punkty_proc": [round(100 * rpas[249], 1), round(100 * rpas[9749], 1)],
        "intencje_w_parach": len({po["pl-PL"][i]["etykieta"] for i in ids}),
    }
    wynik["progi_pewnosci"] = {str(prog): {l: {"decyzje": sum(po[l][i]["pewnosc"] >= prog for i in ids),
                                                 "bledne": sum(po[l][i]["pewnosc"] >= prog and not trafna(po[l][i]) for i in ids)}
                                             for l in ("pl-PL", "en-US")} for prog in PROGI}
    for nazwa, grupa in (("intencje_rzadkie", [i for i in ids if po["pl-PL"][i]["etykieta"] in rzadkie]),
                         ("intencje_czeste", [i for i in ids if po["pl-PL"][i]["etykieta"] not in rzadkie])):
        wynik[nazwa] = {"pary": len(grupa), **{f"trafne_{l[:2]}": sum(trafna(po[l][i]) for i in grupa) for l in po},
                        "tylko_pl_trafne": sum(trafna(po["pl-PL"][i]) and not trafna(po["en-US"][i]) for i in grupa),
                        "tylko_en_trafne": sum(trafna(po["en-US"][i]) and not trafna(po["pl-PL"][i]) for i in grupa)}
    return wynik


wynik = {"opis": ("Uzupełnienie analizy głównej; analiza główna (wyniki.json, 1000 par) bez zmian. Sprawdzeniem analiz po "
                  "fakcie są wyłącznie wyniki dla pozostałych 1974 par; liczby dla całej części testowej obejmują 1000 par, "
                  "na których te różnice zauważono, więc ich nie sprawdzają."),
         "pozostale_1974_pary": miary(reszta), "cala_czesc_testowa": miary(pelne)}
json.dump(wynik, open(TU / "wyniki-uzupelnienie.json", "w"), ensure_ascii=False, indent=2)
print(json.dumps(wynik, ensure_ascii=False, indent=1))
