#!/usr/bin/env python3
"""K3 架构图解的教学计算。Python 3.10+，依赖 NumPy。

运行：python K3_图解计算示例.py
写出完整结果：python K3_图解计算示例.py --output results.json

本脚本只验证公式、容量算账与自造数据上的 QB 更新；不加载模型权重。
QB 对应 Kimi K3 技术报告 §2.3.3、附录 C–D：
https://arxiv.org/html/2607.24653v2
矩阵与容量配置：
https://huggingface.co/moonshotai/Kimi-K3/blob/main/config.json
独立 QB 反例来源：
https://nor-blog.pages.dev/posts/2026-09-02-quantile-balancing/
"""
from __future__ import annotations

import argparse
import itertools
import json
from pathlib import Path

import numpy as np

SEED = 20260909


def sigmoid(x: np.ndarray | float) -> np.ndarray:
    a = np.asarray(x, dtype=np.float64)
    return np.exp(-np.logaddexp(0.0, -a))


def softmax(x: np.ndarray, axis: int = -1) -> np.ndarray:
    a = np.asarray(x, dtype=np.float64)
    ex = np.exp(a - np.max(a, axis=axis, keepdims=True))
    return ex / ex.sum(axis=axis, keepdims=True)


def formula_checks() -> dict:
    rng = np.random.default_rng(SEED)
    # 非方形尺寸可更容易发现左右乘与转置错误。
    t, h, latent, dk, dv, dr, out = 6, 3, 5, 2, 3, 2, 7
    c, r = rng.normal(size=(t, latent)), rng.normal(size=(t, dr))
    uk, uv = rng.normal(size=(h, dk, latent)), rng.normal(size=(h, dv, latent))
    qc, qr = rng.normal(size=(h, dk)), rng.normal(size=(h, dr))
    kc = np.einsum('hkl,tl->htk', uk, c)
    v = np.einsum('hvl,tl->htv', uv, c)
    logits_a = (np.einsum('hk,htk->ht', qc, kc) + qr @ r.T) / np.sqrt(dk + dr)
    transformed_q = np.einsum('hkl,hk->hl', uk, qc)
    logits_b = (transformed_q @ c.T + qr @ r.T) / np.sqrt(dk + dr)
    # 查询位置设为第 5 个，最后一个位置被因果掩码挡住。
    logits_a[:, -1] = logits_b[:, -1] = -np.inf
    a, b = softmax(logits_a), softmax(logits_b)
    head_a = np.einsum('ht,htv->hv', a, v)
    head_b = np.einsum('hvl,hl->hv', uv, b @ c)
    gate, wo = sigmoid(rng.normal(size=(h, dv))), rng.normal(size=(out, h * dv))
    ya, yb = wo @ (gate * head_a).ravel(), wo @ (gate * head_b).ravel()
    np.testing.assert_allclose(ya, yb, atol=1e-12, rtol=1e-12)

    state = rng.normal(size=(4, 3))
    key = rng.normal(size=4); key /= np.linalg.norm(key)
    value = rng.normal(size=3)
    decay, beta = rng.uniform(.2, .9, 4), .7
    retained = decay[:, None] * state
    corrected = retained + beta * np.outer(key, value - retained.T @ key)
    factorized = (np.eye(4) - beta * np.outer(key, key)) @ np.diag(decay) @ state + beta * np.outer(key, value)
    np.testing.assert_allclose(corrected, factorized, atol=1e-12, rtol=1e-12)

    def rotation(angle: float) -> np.ndarray:
        return np.array([[np.cos(angle), -np.sin(angle)], [np.sin(angle), np.cos(angle)]])
    theta, m, n = np.pi / 3, 4, 7
    relative = rotation(m * theta).T @ rotation(n * theta)
    np.testing.assert_allclose(relative, rotation((n-m)*theta), atol=1e-12)

    a40 = 40.
    silu_product = a40**2 * sigmoid(a40)
    situ_product = 4 * np.tanh(a40 / 4) * sigmoid(a40) * 25 * np.tanh(a40 / 25)
    return {
        'MLA含输出门_最大绝对误差': float(np.max(np.abs(ya-yb))),
        'KDA两种递推形式_最大绝对误差': float(np.max(np.abs(corrected-factorized))),
        'RoPE相对旋转_最大绝对误差': float(np.max(np.abs(relative-rotation((n-m)*theta)))),
        'a40_SwiGLU中间乘积': float(silu_product),
        'a40_SiTU_GLU中间乘积': float(situ_product),
        '门控函数数值': [dict(a=x, sigmoid=float(sigmoid(x)), SiLU=float(x*sigmoid(x)), SiTU=float(4*np.tanh(x/4)*sigmoid(x))) for x in [-2., 0., 2., 8.]],
    }


def capacity_checks() -> dict:
    d, middle, head = 7168, 3072, 96
    full = 3*d*middle
    latent = 3*(d//2)*middle
    assert 8*full == 16*latent
    adapters = 2*d*(d//2)
    state = 69*96*128*128
    return {
        '全宽专家主矩阵参数': full,
        '低维专家主矩阵参数': latent,
        '每token路由专家激活主矩阵参数': 8*full,
        '每层新增降升维主矩阵参数': adapters,
        '两共享专家主矩阵参数': 2*full,
        'KDA全层单序列状态标量数': state,
        'MLA缓存每位置每层标量数': 512+64,
        '展开KV参考路径每位置每层标量数': head*(192+128),
        'MLA_prefill每位置对MAC': head*(192+128),
        'MLA_latent_decode每位置对MAC': head*(576+512),
        'GQA8每位置每层缓存标量数': 8*(128+128),
        'GQA8每位置对MAC': head*(128+128),
        '混合状态容量低于93层全MLA的长度': {
            'MLA与KDA都按2字节': state/((93-24)*576),
            'MLA按2字节_KDA按4字节': 2*state/((93-24)*576),
        },
        '一百万token时注意力主状态字节数': {
            '93层全MLA_2字节': 93*576*1_000_000*2,
            '混合_KDA2字节': 24*576*1_000_000*2+state*2,
            '混合_KDA4字节': 24*576*1_000_000*2+state*4,
        },
    }


def route(scores: np.ndarray, bias: np.ndarray, k: int) -> tuple[np.ndarray, np.ndarray]:
    if scores.ndim != 2 or not 0 < k < scores.shape[1]:
        raise ValueError('需要二维分数数组，且 0 < k < 专家数。')
    biased = scores + bias
    order = np.argsort(-biased, axis=1)
    loads = np.bincount(order[:, :k].ravel(), minlength=scores.shape[1])
    cutoff = biased[np.arange(len(scores)), order[:, k]]
    return loads, cutoff


def qb_update(scores: np.ndarray, bias: np.ndarray, k: int, bins: int | None = None) -> np.ndarray:
    count, n = scores.shape
    _, cutoff = route(scores, bias, k)
    margin = scores - cutoff[:, None]
    if bins is None:
        # 保持 HTML 教学模拟的有限样本顺序统计量约定。
        updated = -np.sort(margin, axis=0)[int((1-k/n)*count)]
    else:
        if bins < 1:
            raise ValueError('bins 必须为正整数。')
        # r 的符号与 margin 相反；这个范围依赖 scores ∈ [0,1]。
        if np.any(scores < 0) or np.any(scores > 1):
            raise ValueError('当前直方图范围要求 scores 在 [0,1] 内。')
        r = -margin
        lo, hi = float(bias.min()-1), float(bias.max()+1)
        updated = np.empty(n)
        q = count*k/n
        for j in range(n):
            hist, edges = np.histogram(r[:, j], bins=bins, range=(lo, hi))
            cumulative = np.cumsum(hist)
            hit = min(int(np.searchsorted(cumulative, q)), bins-1)
            previous = cumulative[hit-1] if hit else 0
            within = (q-previous)/max(int(hist[hit]), 1)
            updated[j] = edges[hit] + np.clip(within, 0, 1)*(edges[hit+1]-edges[hit])
    return updated-updated.mean()


def qb_examples() -> dict:
    n, k, tokens, steps = 16, 2, 4096, 40
    target = tokens*k/n
    mu = np.array([1.6,1.3,1,.8,.55,.35,.15,0,-.1,-.25,-.4,-.6,-.8,-1,-1.3,-1.6])
    curves: dict[str, list[float]] = {}
    for method in ['QB精确', 'QB直方图1000', 'SignSGD_u0.02', 'SignSGD_u0.005']:
        rng = np.random.default_rng(SEED)
        bias = np.zeros(n)
        curve = []
        for _ in range(steps+1):
            scores = sigmoid(mu+rng.standard_normal((tokens, n)))
            loads, _ = route(scores, bias, k)
            curve.append(float((loads.max()-target)/target))
            if method == 'QB精确':
                bias = qb_update(scores, bias, k)
            elif method == 'QB直方图1000':
                bias = qb_update(scores, bias, k, bins=1000)
            else:
                u = .02 if method.endswith('0.02') else .005
                bias = bias+u*np.sign(target-loads)
        curves[method] = curve
    indices = [0,1,3,10,40]
    diff = float(np.max(np.abs(np.array(curves['QB精确'])-curves['QB直方图1000'])))
    assert np.isclose(diff, .02734375)
    return {'设置': {'专家':n, 'topk':k, '每批token':tokens, 'seed':SEED},
            '展示步数':indices, '对应MaxVio':{name:[v[i] for i in indices] for name,v in curves.items()},
            '精确与直方图最大MaxVio差':diff, '完整曲线':curves}


def counterexample_check() -> dict:
    # max-assignment 的对偶写法使用价格 beta；路由偏置 b = -beta。
    scores = np.array([[-3,1,0],[1,3,-3],[-1.9,-2,2.]])
    beta = np.array([0.,1.,0.])
    alpha = np.sort(scores-beta, axis=1)[:, -2]
    next_beta = np.sort(scores-alpha[:, None], axis=0)[-2]
    dual = float(alpha.sum()+beta.sum()+np.maximum(scores-alpha[:, None]-beta, 0).sum())
    best = max(sum(scores[i,j] for i,j in enumerate(p)) for p in itertools.permutations(range(3)))
    np.testing.assert_allclose(next_beta, beta)
    assert np.isclose(dual, 5.) and np.isclose(best, 4.)
    return {'对偶固定点值':dual, '最佳一对一分配值':float(best),
            '说明':'这个指定分位约定的反例展示一般收敛保证的边界，不能直接预测真实模型训练表现。'}


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output', type=Path, help='可选 JSON 输出文件；已有文件将被覆盖。')
    args = parser.parse_args()
    result = {'公式自检':formula_checks(), '容量与主矩阵':capacity_checks(),
              'QB教学模拟':qb_examples(), 'QB反例':counterexample_check()}
    text = json.dumps(result, ensure_ascii=False, indent=2)
    if args.output:
        args.output.parent.mkdir(parents=True, exist_ok=True)
        args.output.write_text(text+'\n', encoding='utf-8')
    print(text)


if __name__ == '__main__':
    main()
