"""Chi so phat hien sup do bieu dien cho Audio-JEPA.

Dung trong lop tien huan luyen tiep: goi moi epoch tren MOT batch co dinh,
in ra va luu cung checkpoint. Xem `tien-huan-luyen-tieng-noi/01-ba-quyet-dinh-thiet-ke.md` muc 3.

Diem cot loi: ham loss cua Audio-JEPA la `2 - 2*cos` tren target da chuan hoa
tung token, nen loss GIAM khi bieu dien sup do. Khong duoc dung loss de canh.

Bon chi so bo tro nhau; da kiem bang tensor tong hop (chay `python collapse_metrics.py`):

    truong hop                              std_token  rank_eff  cos_cross  std_utt
    khoe (ngau nhien day du hang)              0.0361     758.7     0.0003   0.0360
    sup do hoan toan (moi token 1 huong)       0.0000     758.7     1.0000   0.0000
    sup do mot phan (hang 8)                   0.0350       8.0    -0.0023   0.0350
    moi utterance giong nhau, token khac       0.0243     159.9     0.5468   0.0001

Doc bang nay ky truoc khi dat nguong:
  - `rank_eff` (tren ma tran DA TRU TRUNG BINH) **khong** bat duoc sup do hoan toan
    — no van 758.7 — vi sau khi tru trung binh chi con nhieu, va nhieu thi day hang.
    Sup do hoan toan do `std_token` va `cos_cross` bat.
  - Nguoc lai, sup do mot phan (hang thap) **chi** co `rank_eff` bat duoc; ba chi so
    kia deu trong nhu binh thuong.
  - `std_utt` la chi so duy nhat bat duoc "moi utterance giong nhau".
  => Phai theo doi ca bon. Bo mot cai la mu mot che do hong.
"""
from __future__ import annotations
import math
import torch
import torch.nn.functional as F

__all__ = ["collapse_metrics", "check_gate", "format_row", "HEADER"]

HEADER = f"{'epoch':>5s} {'std_token':>10s} {'rank_eff':>9s} {'rank_raw':>9s} {'cos_cross':>10s} {'std_utt':>8s}"


def _effective_rank(X: torch.Tensor, eps: float = 1e-12) -> float:
    """exp(entropy cua pho gia tri ky di da chuan hoa). Bang 1 khi hang 1."""
    sv = torch.linalg.svdvals(X.double())
    p = sv / (sv.sum() + eps)
    p = p[p > 0]
    return float(torch.exp(-(p * torch.log(p)).sum()))


@torch.no_grad()
def collapse_metrics(H: torch.Tensor, n_pairs: int = 4096, seed: int = 0) -> dict:
    """H: (B, N, D) dau ra encoder tren mot batch co dinh. B >= 2, N >= 2.

    Tra ve dict 4 chi so + `rank_raw` (hang hieu dung khi KHONG tru trung binh,
    bat duoc sup do ve mot huong chung).
    """
    assert H.dim() == 3, f"can (B,N,D), nhan duoc {tuple(H.shape)}"
    B, N, D = H.shape
    assert B >= 2 and N >= 2, "can it nhat 2 utterance va 2 token"
    H = H.detach().float().cpu()

    Hn3 = F.normalize(H, dim=-1)                 # (B,N,D) da chuan hoa
    Hn = Hn3.reshape(B * N, D)

    std_token = float(Hn.std(dim=0).mean())

    Hf = H.reshape(B * N, D)
    rank_eff = _effective_rank(Hf - Hf.mean(0, keepdim=True))   # sup do chieu
    rank_raw = _effective_rank(Hn)                              # sup do ve 1 huong

    g = torch.Generator().manual_seed(seed)
    b1 = torch.randint(0, B, (n_pairs,), generator=g)
    b2 = (b1 + torch.randint(1, B, (n_pairs,), generator=g)) % B   # bao dam b2 != b1
    n1 = torch.randint(0, N, (n_pairs,), generator=g)
    n2 = torch.randint(0, N, (n_pairs,), generator=g)
    cos_cross = float((Hn3[b1, n1] * Hn3[b2, n2]).sum(-1).mean())

    std_utt = float(F.normalize(H.mean(1), dim=-1).std(dim=0).mean())

    return dict(std_token=std_token, rank_eff=rank_eff, rank_raw=rank_raw,
                cos_cross=cos_cross, std_utt=std_utt)


def format_row(epoch, m: dict) -> str:
    return (f"{epoch:>5} {m['std_token']:10.4f} {m['rank_eff']:9.1f} {m['rank_raw']:9.1f} "
            f"{m['cos_cross']:10.4f} {m['std_utt']:8.4f}")


def check_gate(m: dict, base: dict) -> list[str]:
    """Tra ve danh sach ly do phai DUNG. Rong = di tiep.

    `base` la bo chi so do tren checkpoint goc (epoch 0). Quy tac dung chi
    kich hoat khi vi pham o HAI checkpoint lien tiep — nguoi goi tu giu trang thai.
    """
    bad = []
    if m["rank_eff"] < 0.50 * base["rank_eff"]:
        bad.append(f"rank_eff {m['rank_eff']:.1f} < 50% moc {base['rank_eff']:.1f} (sup do chieu)")
    if m["rank_raw"] < 0.50 * base["rank_raw"]:
        bad.append(f"rank_raw {m['rank_raw']:.1f} < 50% moc {base['rank_raw']:.1f} (sup do ve 1 huong)")
    if m["cos_cross"] > 0.9 or m["cos_cross"] > base["cos_cross"] + 0.3:
        bad.append(f"cos_cross {m['cos_cross']:.3f} (moc {base['cos_cross']:.3f})")
    if m["std_token"] < 0.25 * base["std_token"]:
        bad.append(f"std_token {m['std_token']:.4f} < 25% moc {base['std_token']:.4f}")
    if m["std_utt"] < 0.25 * base["std_utt"]:
        bad.append(f"std_utt {m['std_utt']:.4f} < 25% moc {base['std_utt']:.4f} (moi utterance giong nhau)")
    return bad


if __name__ == "__main__":
    torch.manual_seed(0)
    B, N, D = 64, 128, 768
    v = torch.randn(1, 1, D)
    u = torch.randn(1, 1, D)
    cases = {
        "khoe (ngau nhien day du hang)":       torch.randn(B, N, D),
        "sup do hoan toan (moi token 1 huong)": v.repeat(B, N, 1) + 0.001 * torch.randn(B, N, D),
        "sup do mot phan (hang 8)":             torch.randn(B, N, 8) @ torch.randn(8, D),
        "moi utterance giong nhau, token khac": u + 0.9 * torch.randn(1, N, D).repeat(B, 1, 1)
                                                  + 0.02 * torch.randn(B, N, D),
    }
    base = collapse_metrics(cases["khoe (ngau nhien day du hang)"])
    print(f"{'truong hop':40s} {'std_token':>10s} {'rank_eff':>9s} {'rank_raw':>9s} {'cos_cross':>10s} {'std_utt':>8s}  cong")
    for k, H in cases.items():
        m = collapse_metrics(H)
        bad = check_gate(m, base)
        verdict = "DI TIEP" if not bad else "DUNG: " + "; ".join(x.split(" (")[0] for x in bad)
        print(f"{k:40s} {m['std_token']:10.4f} {m['rank_eff']:9.1f} {m['rank_raw']:9.1f} "
              f"{m['cos_cross']:10.4f} {m['std_utt']:8.4f}  {verdict}")
    print(f"\nmoc ly thuyet trai deu tren mat cau D={D}: std_token = 1/sqrt(D) = {1/math.sqrt(D):.4f}")
