"""Teaching calculations for K3 Architecture Illustrated. Python 3.10+ and NumPy.

Run: python K3_Calculation_Examples.py
Write full results: python K3_Calculation_Examples.py --output results.json

This script checks formulas, capacity calculations, and QB updates on synthetic data.
It does not load model weights.
QB: Kimi K3 Technical Report, section 2.3.3 and appendices C–D:
https://arxiv.org/html/2607.24653v2
Matrix and capacity configuration:
https://huggingface.co/moonshotai/Kimi-K3/blob/main/config.json
Independent QB counterexample:
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)
    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(0.2, 0.9, 4), 0.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.0
    silu_product = a40 ** 2 * sigmoid(a40)
    situ_product = 4 * np.tanh(a40 / 4) * sigmoid(a40) * 25 * np.tanh(a40 / 25)
    return {'MLA_with_output_gate_max_abs_error': float(np.max(np.abs(ya - yb))), 'KDA_equivalent_updates_max_abs_error': float(np.max(np.abs(corrected - factorized))), 'RoPE_relative_rotation_max_abs_error': float(np.max(np.abs(relative - rotation((n - m) * theta)))), 'a40_SwiGLU_intermediate_product': float(silu_product), 'a40_SiTU_GLU_intermediate_product': float(situ_product), 'gate_function_values': [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, 0.0, 2.0, 8.0]]}

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_width_expert_matrix_parameters': full, 'latent_expert_matrix_parameters': latent, 'active_routed_matrix_parameters_per_token': 8 * full, 'added_projection_parameters_per_layer': adapters, 'two_shared_experts_matrix_parameters': 2 * full, 'KDA_state_scalars_all_layers_single_sequence': state, 'MLA_cache_scalars_per_position_per_layer': 512 + 64, 'expanded_KV_reference_scalars_per_position_per_layer': head * (192 + 128), 'MLA_prefill_MACs_per_position_pair': head * (192 + 128), 'MLA_latent_decode_MACs_per_position_pair': head * (576 + 512), 'GQA8_cache_scalars_per_position_per_layer': 8 * (128 + 128), 'GQA8_MACs_per_position_pair': head * (128 + 128), 'hybrid_vs_93_layer_MLA_crossover_tokens': {'MLA_and_KDA_2_bytes': state / ((93 - 24) * 576), 'MLA_2_bytes_KDA_4_bytes': 2 * state / ((93 - 24) * 576)}, 'main_attention_state_bytes_at_one_million_tokens': {'all_93_MLA_layers_2_bytes': 93 * 576 * 1000000 * 2, 'hybrid_KDA_2_bytes': 24 * 576 * 1000000 * 2 + state * 2, 'hybrid_KDA_4_bytes': 24 * 576 * 1000000 * 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('Expected a 2D score array and 0 < k < expert count.')
    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:
        updated = -np.sort(margin, axis=0)[int((1 - k / n) * count)]
    else:
        if bins < 1:
            raise ValueError('bins must be a positive integer.')
        if np.any(scores < 0) or np.any(scores > 1):
            raise ValueError('This histogram range requires scores in [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, 0.8, 0.55, 0.35, 0.15, 0, -0.1, -0.25, -0.4, -0.6, -0.8, -1, -1.3, -1.6])
    curves: dict[str, list[float]] = {}
    for method in ['QB_exact', 'QB_histogram_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_exact':
                bias = qb_update(scores, bias, k)
            elif method == 'QB_histogram_1000':
                bias = qb_update(scores, bias, k, bins=1000)
            else:
                u = 0.02 if method.endswith('0.02') else 0.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_exact']) - curves['QB_histogram_1000'])))
    assert np.isclose(diff, 0.02734375)
    return {'settings': {'experts': n, 'topk': k, 'tokens_per_batch': tokens, 'seed': SEED}, 'displayed_steps': indices, 'displayed_MaxVio': {name: [v[i] for i in indices] for name, v in curves.items()}, 'max_exact_vs_histogram_MaxVio_difference': diff, 'complete_curves': curves}

def counterexample_check() -> dict:
    scores = np.array([[-3, 1, 0], [1, 3, -3], [-1.9, -2, 2.0]])
    beta = np.array([0.0, 1.0, 0.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.0) and np.isclose(best, 4.0)
    return {'dual_fixed_point_value': dual, 'best_one_to_one_assignment_value': float(best), 'note': 'This counterexample uses a specified quantile convention to show the limits of a general convergence claim; it does not predict training behavior in a real model.'}

def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output', type=Path, help='Optional JSON output path; an existing output file will be overwritten.')
    args = parser.parse_args()
    result = {'formula_checks': formula_checks(), 'capacity_and_main_matrices': capacity_checks(), 'QB_teaching_simulation': qb_examples(), 'QB_counterexample': 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()
