"""Vérifications des identités mathématiques de Kimi K3, en Python pur.

Aucune dépendance. Lancer :  python3 verif_k3.py
"""
import math, random

random.seed(7)

# ---------------------------------------------------------------- algèbre
def zeros(n, m):      return [[0.0] * m for _ in range(n)]
def eye(n):           return [[1.0 if i == j else 0.0 for j in range(n)] for i in range(n)]
def rand(n, m, s=1.0):return [[random.gauss(0, s) for _ in range(m)] for _ in range(n)]
def matmul(A, B):
    n, k, m = len(A), len(B), len(B[0])
    return [[sum(A[i][t] * B[t][j] for t in range(k)) for j in range(m)] for i in range(n)]
def add(A, B):        return [[a + b for a, b in zip(ra, rb)] for ra, rb in zip(A, B)]
def scale(A, c):      return [[a * c for a in r] for r in A]
def outer(u, v):      return [[ui * vj for vj in v] for ui in u]
def diagmul(d, A):    return [[d[i] * A[i][j] for j in range(len(A[0]))] for i in range(len(A))]
def maxabs(A, B):     return max(abs(a - b) for ra, rb in zip(A, B) for a, b in zip(ra, rb))
def l2norm(v):
    n = math.sqrt(sum(x * x for x in v)) or 1.0
    return [x / n for x in v]

def sigmoid(x):       return 1.0 / (1.0 + math.exp(-x))

# ---------------------------------------------------------------- KDA
DK, DV = 6, 5           # dimensions réduites, l'identité est indépendante de la taille

def kda_step(S, k, v, alpha, beta):
    """Un pas de la récurrence KDA (équation 1 du rapport).

    S_t = (I - beta k k^T) Diag(alpha) S_{t-1} + beta k v^T
    """
    S = diagmul(alpha, S)                       # (1) oublier, canal par canal
    kk = matmul(outer(k, k), S)                 # (2) effacer la composante alignée sur k
    S = add(S, scale(kk, -beta))
    return add(S, scale(outer(k, v), beta))     # (3) écrire la nouvelle association

def kda_transition(k, beta, alpha):
    """La matrice M_t = (I - beta k k^T) Diag(alpha) du même pas."""
    D = [[alpha[j] if i == j else 0.0 for j in range(DK)] for i in range(DK)]
    P = add(eye(DK), scale(outer(k, k), -beta))
    return matmul(P, D)

def make_tokens(n):
    toks = []
    for _ in range(n):
        k = l2norm([random.gauss(0, 1) for _ in range(DK)])
        v = [random.gauss(0, 1) for _ in range(DV)]
        # decroissance bornee : g = g_min * sigmoid(z), alpha = exp(g), g_min = -5
        alpha = [math.exp(-5.0 * sigmoid(random.gauss(0, 1))) for _ in range(DK)]
        beta = sigmoid(random.gauss(0, 1))
        toks.append((k, v, alpha, beta))
    return toks

def run(toks, S0):
    S = [r[:] for r in S0]
    for (k, v, a, b) in toks:
        S = kda_step(S, k, v, a, b)
    return S

# --- test 1 : identite de composition KCP -------------------------------
def test_kcp_composition():
    toks = make_tokens(9)
    S0 = rand(DK, DV)                       # etat entrant arbitraire

    direct = run(toks, S0)                  # recurrence complete depuis S0

    from_zero = run(toks, zeros(DK, DV))    # etat genere localement depuis 0
    M = eye(DK)                             # transition cumulee du segment
    for (k, v, a, b) in toks:
        M = matmul(kda_transition(k, b, a), M)
    composed = add(from_zero, matmul(M, S0))

    err = maxabs(direct, composed)
    assert err < 1e-9, f"KCP composition: erreur {err}"
    return err

# --- test 2 : balayage prefixe sur plusieurs rangs ----------------------
def test_kcp_prefix_scan(P=4, per_rank=5):
    segments = [make_tokens(per_rank) for _ in range(P)]

    # reference : une seule recurrence sequentielle sur toute la sequence
    flat = [t for seg in segments for t in seg]
    ref = run(flat, zeros(DK, DV))

    # KCP : chaque rang calcule ses deux fragments LOCALEMENT, sans communication
    frags = []
    for seg in segments:
        S_local = run(seg, zeros(DK, DV))
        M = eye(DK)
        for (k, v, a, b) in seg:
            M = matmul(kda_transition(k, b, a), M)
        frags.append((M, S_local))

    # all-gather puis balayage prefixe associatif
    S = zeros(DK, DV)
    for (M, S_local) in frags:
        S = add(matmul(M, S), S_local)

    err = maxabs(ref, S)
    assert err < 1e-9, f"KCP prefix scan: erreur {err}"
    return err

# ---------------------------------------------------------------- SiTU-GLU
def situ_glu(xg, xu, b1=4.0, b2=25.0):
    gate = b1 * math.tanh(xg / b1) * sigmoid(xg)
    up = b2 * math.tanh(xu / b2)
    return gate * up

def swiglu(xg, xu):
    return (xg * sigmoid(xg)) * xu

def test_situ_bound():
    worst = 0.0
    for _ in range(200000):
        xg = random.uniform(-500, 500)
        xu = random.uniform(-500, 500)
        worst = max(worst, abs(situ_glu(xg, xu)))
    assert worst <= 4.0 * 25.0, f"borne violee : {worst}"
    return worst

def test_situ_local_agreement():
    """SiTU-GLU coincide avec SwiGLU au premier ordre autour de l'origine."""
    worst_rel = 0.0
    for _ in range(2000):
        xg = random.uniform(-0.5, 0.5)
        xu = random.uniform(-0.5, 0.5)
        a, b = situ_glu(xg, xu), swiglu(xg, xu)
        if abs(b) > 1e-6:
            worst_rel = max(worst_rel, abs(a - b) / abs(b))
    assert worst_rel < 0.02, f"ecart local {worst_rel}"
    return worst_rel

# ---------------------------------------------------------------- Quantile Balancing
def quantile(xs, q):
    """Quantile d'ordre q par selection sur la liste triee (convention du rapport)."""
    ys = sorted(xs, reverse=True)
    idx = min(len(ys) - 1, int(round((1.0 - q) * len(ys))))
    return ys[idx]

def route(s, b, k):
    """Top-k sur le score biaise ; renvoie (routes, seuil alpha_i)."""
    n = len(b)
    order = sorted(range(n), key=lambda j: -(s[j] + b[j]))
    return order[:k], s[order[k]] + b[order[k]]      # le (k+1)-ieme est le seuil

def _loads(S, b, k, n):
    m = len(S)
    out = [0] * n
    for i in range(m):
        for j in route(S[i], b, k)[0]:
            out[j] += 1
    return out

def test_quantile_balancing(m=4096, n=16, k=2, steps=8, gamma=1e-3):
    """Compare QB a la regle par signe (DeepSeek-V3) sur le meme lot.

    Renvoie l'ecart de charge relatif (%) par pas pour les deux methodes.
    QB n'a AUCUN hyperparametre ; la regle par signe depend de gamma.
    """
    S = [[sigmoid(random.gauss(0, 2)) for _ in range(n)] for _ in range(m)]
    target = m * k / n

    b_qb, b_sg = [0.0] * n, [0.0] * n
    hist_qb, hist_sg = [], []

    for _ in range(steps):
        # --- Quantile Balancing -------------------------------------
        alphas, loads = [], [0] * n
        for i in range(m):
            routes, a = route(S[i], b_qb, k)
            alphas.append(a)
            for j in routes:
                loads[j] += 1
        hist_qb.append(round((max(loads) - min(loads)) / target * 100, 1))

        nb = [-quantile([S[i][j] - alphas[i] for i in range(m)], 1.0 - k / n)
              for j in range(n)]
        mean = sum(nb) / n
        b_qb = [x - mean for x in nb]          # centrage : invariant pour le Top-k

        # --- regle par signe, a pas fixe ------------------------------
        ld = _loads(S, b_sg, k, n)
        hist_sg.append(round((max(ld) - min(ld)) / target * 100, 1))
        avg = sum(ld) / n
        b_sg = [b_sg[j] + gamma * (1 if avg > ld[j] else -1) for j in range(n)]

    # QB doit reduire fortement le desequilibre des le PREMIER pas,
    # sans aucun taux d'apprentissage a regler.
    assert hist_qb[1] < 0.75 * hist_qb[0], f"QB ne converge pas : {hist_qb}"
    assert hist_qb[-1] < hist_qb[0] / 2, f"QB stagne : {hist_qb}"
    return hist_qb, hist_sg, target

# ---------------------------------------------------------------- comptage de parametres
def param_counts():
    d, L, nMLA, nKDA, V = 7168, 93, 24, 69, 163840
    ell, dm, E, k, Ns = 3584, 3072, 896, 16, 2
    qlora, kvlora, qk_nope, qk_rope, vh, H = 1536, 512, 128, 64, 128, 96
    kda_hd, kda_H, Ilarge = 128, 96, 33792

    exp = 3 * ell * dm
    moe = E * exp + (d * ell + ell * d) + 3 * d * (dm * Ns) + d * E
    mla = (d * qlora + qlora * H * (qk_nope + qk_rope)) + d * (kvlora + qk_rope) \
        + kvlora * H * (qk_nope + vh) + H * vh * d + d * H * vh
    kda = 3 * (d * kda_H * kda_hd) + (kda_H * kda_hd * d) + (d * kda_H * kda_hd)
    dense = 3 * d * Ilarge

    total = 92 * moe + nMLA * mla + nKDA * kda + dense + 2 * V * d
    act = 92 * (k * exp + 2 * d * ell + 3 * d * dm * Ns + d * E) \
        + nMLA * mla + nKDA * kda + dense + V * d
    return total, act, exp, moe, mla, kda

# ---------------------------------------------------------------- main
if __name__ == "__main__":
    e1 = test_kcp_composition()
    e2 = test_kcp_prefix_scan()
    w = test_situ_bound()
    r = test_situ_local_agreement()
    hqb, hsg, target = test_quantile_balancing()
    total, act, exp, moe, mla, kda = param_counts()

    print(f"[OK] KCP composition        erreur max = {e1:.2e}")
    print(f"[OK] KCP balayage prefixe   erreur max = {e2:.2e}  (4 rangs)")
    print(f"[OK] SiTU-GLU borne         max|f| = {w:.3f}  <= b1*b2 = 100")
    print(f"[OK] SiTU vs SwiGLU local   ecart relatif max = {r*100:.3f} %")
    print(f"[OK] Quantile Balancing     desequilibre relatif (%), cible={target:.0f}")
    print(f"       QB (sans hyperparam.) {hqb}")
    print(f"       regle par signe       {hsg}")
    print(f"[OK] Parametres             total = {total/1e12:.3f} T (papier 2,78 T)")
    print(f"                            actifs = {act/1e9:.1f} G (papier 104,2 G)")
    print(f"                            1 expert = {exp/1e6:.1f} M | couche MoE = {moe/1e9:.2f} G")
    print(f"                            couche MLA = {mla/1e6:.1f} M | couche KDA = {kda/1e6:.1f} M")
