#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Vérification arithmétique de Qwen3.8-Max (dépôt Qwen/Qwen3.8-2.4T-A95B).

Aucune dépendance : Python 3 pur, pas de numpy, pas de torch.

Ce script part uniquement du fichier `config.json` publié par Qwen le
12 août 2026 sur Hugging Face, et recalcule :

  1. le nombre total de paramètres        (annoncé : 2,4 T)
  2. le nombre de paramètres actifs       (annoncé : 95 G)
  3. la taille du dépôt en BF16           (observé  : 4,89 To)
  4. le poids du cache clé-valeur         (non annoncé)
  5. l'état récurrent de l'attention linéaire
  6. le matériel minimal pour l'inférence

Chaque résultat est comparé à la valeur annoncée ou observée. Le script
sort un code de retour non nul si une vérification échoue.

Usage :  python3 verif-qwen3-8-max.py
"""

import sys

# --------------------------------------------------------------------------
# 1. La configuration publiée
# --------------------------------------------------------------------------
# Recopiée telle quelle depuis
# https://huggingface.co/Qwen/Qwen3.8-2.4T-A95B/raw/main/config.json
# Seule la liste `layer_types` est résumée : elle contient 92 entrées,
# répétition de ["linear_attention"] * 3 + ["full_attention"], 23 fois.

CONFIG = {
    "hidden_size": 8192,
    "vocab_size": 248320,
    "num_hidden_layers": 92,
    "full_attention_interval": 4,
    # attention complète (gatée)
    "num_attention_heads": 64,
    "num_key_value_heads": 4,
    "head_dim": 256,
    "attn_output_gate": True,
    "partial_rotary_factor": 0.25,
    # attention linéaire (Gated DeltaNet)
    "linear_num_value_heads": 128,
    "linear_num_key_heads": 16,
    "linear_key_head_dim": 128,
    "linear_value_head_dim": 128,
    "linear_conv_kernel_dim": 4,
    # Mixture-of-Experts
    "num_experts": 512,
    "num_experts_per_tok": 10,
    "moe_intermediate_size": 2048,
    "shared_expert_intermediate_size": 2048,
    # divers
    "tie_word_embeddings": False,
    "mtp_num_hidden_layers": 1,
    "max_position_embeddings": 262144,
}

# Valeurs annoncées par Qwen, à confronter au calcul.
ANNONCE_TOTAL = 2.4e12          # « 2.4 trillion parameters »
ANNONCE_ACTIF = 95e9            # « 95 billion activated »
OBSERVE_DEPOT_TO = 4.89         # taille du dépôt affichée par Hugging Face
CONTEXTE_ETENDU = 1_010_000     # « extensible up to 1,010,000 tokens »

G = 10 ** 9
T = 10 ** 12


def titre(txt):
    print()
    print(txt)
    print("-" * len(txt))


def ligne(label, valeur, unite=""):
    print("  {:<46} {:>18} {}".format(label, valeur, unite))


# --------------------------------------------------------------------------
# 2. Décompte des couches
# --------------------------------------------------------------------------
def compte_couches(cfg):
    """23 blocs de (3 linéaires + 1 complète) = 92 couches."""
    intervalle = cfg["full_attention_interval"]
    total = cfg["num_hidden_layers"]
    n_complete = total // intervalle
    n_lineaire = total - n_complete
    return n_lineaire, n_complete


# --------------------------------------------------------------------------
# 3. Paramètres par bloc
# --------------------------------------------------------------------------
def params_attention_complete(cfg):
    """
    Attention complète gatée, façon Qwen3-Next.

    q_proj porte aussi la porte de sortie quand attn_output_gate vaut True :
    la projection produit 2 x (n_heads * head_dim).
    """
    h = cfg["hidden_size"]
    d = cfg["head_dim"]
    dim_q = cfg["num_attention_heads"] * d
    dim_kv = cfg["num_key_value_heads"] * d

    facteur_gate = 2 if cfg["attn_output_gate"] else 1
    q = h * dim_q * facteur_gate
    k = h * dim_kv
    v = h * dim_kv
    o = dim_q * h
    # q_norm / k_norm : d paramètres chacun, négligeables mais comptés
    normes = 2 * d
    return q + k + v + o + normes


def params_attention_lineaire(cfg):
    """
    Gated DeltaNet.

    Une projection d'entrée unique produit q, k, v et la porte z ;
    une seconde produit beta et a (un scalaire par tête de valeur).
    Une convolution causale de rang 4 agit sur q, k et v.
    """
    h = cfg["hidden_size"]
    dim_qk = cfg["linear_num_key_heads"] * cfg["linear_key_head_dim"]
    dim_v = cfg["linear_num_value_heads"] * cfg["linear_value_head_dim"]

    in_proj_qkvz = h * (dim_qk + dim_qk + dim_v + dim_v)
    in_proj_ba = h * (2 * cfg["linear_num_value_heads"])
    conv = (dim_qk + dim_qk + dim_v) * cfg["linear_conv_kernel_dim"]
    out_proj = dim_v * h
    normes = cfg["linear_value_head_dim"]
    return in_proj_qkvz + in_proj_ba + conv + out_proj + normes


def params_moe(cfg):
    """
    Un bloc MoE par couche : 512 experts routés + 1 expert partagé + routeur.
    Chaque expert est un SwiGLU : gate, up, down.
    """
    h = cfg["hidden_size"]
    i = cfg["moe_intermediate_size"]
    par_expert = 3 * h * i
    routes = cfg["num_experts"] * par_expert
    partage = 3 * h * cfg["shared_expert_intermediate_size"]
    routeur = h * cfg["num_experts"]
    return routes + partage + routeur


def params_moe_actifs(cfg):
    """Seuls top-k experts routés + l'expert partagé sont évalués."""
    h = cfg["hidden_size"]
    i = cfg["moe_intermediate_size"]
    par_expert = 3 * h * i
    routes = cfg["num_experts_per_tok"] * par_expert
    partage = 3 * h * cfg["shared_expert_intermediate_size"]
    routeur = h * cfg["num_experts"]
    return routes + partage + routeur


# --------------------------------------------------------------------------
# 4. Total et actif
# --------------------------------------------------------------------------
def decompte(cfg):
    n_lin, n_full = compte_couches(cfg)
    h = cfg["hidden_size"]

    p_full = params_attention_complete(cfg)
    p_lin = params_attention_lineaire(cfg)
    p_moe = params_moe(cfg)
    p_moe_act = params_moe_actifs(cfg)

    embeddings = cfg["vocab_size"] * h
    tete = 0 if cfg["tie_word_embeddings"] else cfg["vocab_size"] * h
    normes = 2 * h * cfg["num_hidden_layers"] + h

    total = (
        embeddings
        + tete
        + normes
        + n_full * p_full
        + n_lin * p_lin
        + cfg["num_hidden_layers"] * p_moe
    )

    # La couche de prédiction multi-jetons est un bloc supplémentaire complet.
    mtp = cfg["mtp_num_hidden_layers"] * (p_moe + p_full)
    total_avec_mtp = total + mtp

    # Actif : toute l'attention, mais seulement 11 experts sur 513.
    actif_sans_embeddings = (
        tete
        + normes
        + n_full * p_full
        + n_lin * p_lin
        + cfg["num_hidden_layers"] * p_moe_act
    )
    actif_avec_embeddings = actif_sans_embeddings + embeddings

    return {
        "n_lin": n_lin,
        "n_full": n_full,
        "p_full": p_full,
        "p_lin": p_lin,
        "p_moe": p_moe,
        "p_moe_act": p_moe_act,
        "embeddings": embeddings,
        "tete": tete,
        "total": total,
        "mtp": mtp,
        "total_avec_mtp": total_avec_mtp,
        "actif_sans_embeddings": actif_sans_embeddings,
        "actif_avec_embeddings": actif_avec_embeddings,
    }


# --------------------------------------------------------------------------
# 5. Mémoire d'inférence
# --------------------------------------------------------------------------
def octets_cache_kv_par_jeton(cfg):
    """
    Seules les couches d'attention complète produisent un cache clé-valeur.
    Les couches linéaires ont un état de taille fixe (calculé plus bas).
    """
    _, n_full = compte_couches(cfg)
    dim_kv = cfg["num_key_value_heads"] * cfg["head_dim"]
    return n_full * 2 * dim_kv * 2  # K et V, 2 octets en BF16


def octets_etat_lineaire(cfg):
    """
    L'état de DeltaNet est une matrice (d_k x d_v) par tête de valeur,
    conservée en float32. Il ne dépend pas de la longueur du contexte.
    """
    n_lin, _ = compte_couches(cfg)
    par_couche = (
        cfg["linear_num_value_heads"]
        * cfg["linear_key_head_dim"]
        * cfg["linear_value_head_dim"]
        * 4
    )
    dim_qk = cfg["linear_num_key_heads"] * cfg["linear_key_head_dim"]
    dim_v = cfg["linear_num_value_heads"] * cfg["linear_value_head_dim"]
    conv = (2 * dim_qk + dim_v) * cfg["linear_conv_kernel_dim"] * 4
    return n_lin * (par_couche + conv)


# --------------------------------------------------------------------------
# 6. Exécution
# --------------------------------------------------------------------------
def main():
    cfg = CONFIG
    d = decompte(cfg)
    echecs = []

    def verifie(nom, calcule, attendu, tolerance):
        ecart = abs(calcule - attendu) / attendu
        ok = ecart <= tolerance
        print("  [{}] {:<40} écart {:>6.2f} %".format(
            "OK " if ok else "ÉCHEC", nom, 100 * ecart))
        if not ok:
            echecs.append(nom)
        return ecart

    print("=" * 78)
    print("  Vérification de Qwen3.8-Max — Qwen/Qwen3.8-2.4T-A95B")
    print("  Source unique : config.json publié le 12 août 2026")
    print("=" * 78)

    titre("1. Structure des couches")
    ligne("couches au total", d["n_lin"] + d["n_full"])
    ligne("couches à attention linéaire (DeltaNet)", d["n_lin"])
    ligne("couches à attention complète", d["n_full"])
    ligne("ratio linéaire : complète", "{}:1".format(
        d["n_lin"] // d["n_full"]))

    titre("2. Paramètres par bloc")
    ligne("bloc d'attention complète gatée", "{:.1f}".format(d["p_full"] / 1e6), "M")
    ligne("bloc Gated DeltaNet", "{:.1f}".format(d["p_lin"] / 1e6), "M")
    ligne("bloc MoE (513 experts)", "{:.2f}".format(d["p_moe"] / G), "G")
    ligne("bloc MoE, part active (11 experts)", "{:.1f}".format(d["p_moe_act"] / 1e6), "M")
    ligne("table d'embeddings", "{:.3f}".format(d["embeddings"] / G), "G")

    titre("3. Total des paramètres")
    ligne("corps du modèle", "{:.4f}".format(d["total"] / T), "T")
    ligne("couche MTP supplémentaire", "{:.2f}".format(d["mtp"] / G), "G")
    ligne("total avec MTP", "{:.4f}".format(d["total_avec_mtp"] / T), "T")
    ligne("annoncé par Qwen", "{:.1f}".format(ANNONCE_TOTAL / T), "T")
    print()
    verifie("total (hors MTP) vs 2,4 T", d["total"], ANNONCE_TOTAL, 0.02)

    titre("4. Paramètres actifs par jeton")
    ligne("hors table d'embeddings", "{:.2f}".format(d["actif_sans_embeddings"] / G), "G")
    ligne("table d'embeddings incluse", "{:.2f}".format(d["actif_avec_embeddings"] / G), "G")
    ligne("annoncé par Qwen", "{:.0f}".format(ANNONCE_ACTIF / G), "G")
    ligne("part du modèle activée", "{:.2f}".format(
        100 * d["actif_avec_embeddings"] / d["total"]), "%")
    print()
    verifie("actif (embeddings inclus) vs 95 G",
            d["actif_avec_embeddings"], ANNONCE_ACTIF, 0.02)

    titre("5. Taille du dépôt en BF16")
    octets = d["total_avec_mtp"] * 2
    to = octets / 1e12
    ligne("poids seuls, 2 octets par paramètre", "{:.2f}".format(to), "To")
    ligne("taille affichée par Hugging Face", "{:.2f}".format(OBSERVE_DEPOT_TO), "To")
    print()
    verifie("taille du dépôt vs 4,89 To", to, OBSERVE_DEPOT_TO, 0.02)

    titre("6. Mémoire d'inférence")
    kv = octets_cache_kv_par_jeton(cfg)
    etat = octets_etat_lineaire(cfg)
    ligne("cache KV par jeton", "{:.1f}".format(kv / 1024), "KiO")
    ligne("cache KV à 262 144 jetons", "{:.1f}".format(
        kv * cfg["max_position_embeddings"] / 1e9), "Go")
    ligne("cache KV à 1 010 000 jetons", "{:.1f}".format(
        kv * CONTEXTE_ETENDU / 1e9), "Go")
    ligne("état récurrent DeltaNet (constant)", "{:.2f}".format(etat / 1e9), "Go")
    print()

    # Contrefactuel : le même modèle sans attention linéaire.
    kv_tout_complet = kv * cfg["num_hidden_layers"] / compte_couches(cfg)[1]
    ligne("si les 92 couches étaient complètes, à 1 M", "{:.0f}".format(
        kv_tout_complet * CONTEXTE_ETENDU / 1e9), "Go")
    ligne("économie apportée par l'hybride", "{:.0f}".format(
        100 * (1 - kv / kv_tout_complet)), "%")

    titre("7. Matériel minimal")
    for nom, octets_par_param in (("BF16", 2), ("FP8", 1), ("4 bits", 0.5)):
        poids = d["total_avec_mtp"] * octets_par_param
        for gpu, vram in (("H100/H200 80 Go", 80e9), ("B200 180 Go", 180e9)):
            besoin = poids + kv * CONTEXTE_ETENDU + etat
            n = int(besoin / (vram * 0.85)) + 1
            ligne("{} sur {} (contexte 1 M)".format(nom, gpu), n, "GPU")

    print()
    print("=" * 78)
    if echecs:
        print("  RÉSULTAT : {} vérification(s) en échec : {}".format(
            len(echecs), ", ".join(echecs)))
        return 1
    print("  RÉSULTAT : toutes les vérifications passent.")
    print("=" * 78)
    return 0


if __name__ == "__main__":
    sys.exit(main())
