Matrice d'architectures — entraîner du « Mistral-like », « Qwen-like », « Kimi-like » avec Forge
Rapport de recherche web (août 2026) + plan d'implémentation basé sur le code actuel
Où en est Forge
Base actuelle : pre-norm RMSNorm/LayerNorm · SwiGLU/GELU · RoPE complet ou positions apprises · MHA/GQA · embeddings liés · MoE top-k (softmax, renormalisé, aux loss) · QAT int8/ternaire · AdamW/Muon · cosine/WSD. Déjà ajoutés cette session : qk_norm, final_softcap, scale_embeddings, n_shared_experts.
Les deltas par famille (chiffres réels des configs HF)
| Famille | Ce qui la distingue | Statut Forge |
|---|---|---|
| Llama 3.2 (1B/3B) | GQA, SwiGLU, tied ≤3B, rope_theta 500k, rope_scaling "llama3" (factor 32, blend low/high freq) | ✅ sauf rope_scaling |
| Mistral 7B | Sliding window 4096 partout ; Ministral : SWA entrelacé full/window | ❌ SWA (kernel) |
| Qwen 2.5 | Biais sur Q/K/V seulement, rope_theta 1e6, GQA 28H/4KV | ❌ attention_bias |
| Qwen 3 (dense) | Remplace le biais par QK-Norm, head_dim découplé (128), tied ≤4B | ✅ qk_norm ; ❌ head_dim découplé |
| Qwen3-30B-A3B (MoE) | 128 experts top-8, pas d'expert partagé, softmax renormalisé, moe_d_ff 768 ≠ d_ff | ✅ sauf moe_intermediate_size séparé |
| Gemma 2 | Local(4096)/global 1:1, softcap attn 50 / final 30, sandwich norm, GeGLU, head_dim 256, emb×√d, query_pre_attn_scalar | ✅ softcap final + emb scale ; ❌ reste |
| Gemma 3 | 5 local : 1 global, window 1024, softcap→QK-norm, rope_theta 10k local / 1M global | ✅ qk_norm ; ❌ patterns par couche |
| DeepSeek V3 | MLA (kv_lora 512, q_lora 1536, nope 128 + rope 64, v 128), 256 routed + 1 shared top-8, routage sigmoïde + équilibrage par biais sans aux loss (γ=0.001), routed_scaling 2.5, 3 premières couches denses, MTP λ=0.3→0.1 | ✅ shared expert ; ❌ MLA, sigmoïde/noaux |
| Kimi K2 | = V3 avec 64 têtes (÷2), 384 experts, couche 0 seule dense, routed_scaling 2.827, MuonClip (Muon + QK-Clip : si logit max > τ=100, W_q,W_k ← ×√(τ/S_max) par tête) | ✅ Muon ; ❌ QK-Clip, reste |
| GPT-OSS | Dense/sliding-128 alterné, attention sinks appris (biais par tête dans le dénominateur softmax), MoE 32/top-4, biais partout, YaRN | ❌ sinks (petit patch kernel) |
| OLMo 2 | Post-norm réordonné (x = x + norm(attn(x))) + QK-norm — la combinaison fait la stabilité |
✅ qk_norm ; ❌ norm_placement |
| SmolLM3 | Llama-like + NoPE 1 couche sur 4 (RoPE sauté) → long contexte 64k+ | ❌ nope_layers (facile) |
| nanochat (Karpathy) | ReLU², head délié zero-init, RMSNorm sans paramètre, qk-norm après RoPE, softcap 15, value embeddings, Muon | ✅ Muon, softcap ; ❌ relu², reste |
Plan d'implémentation priorisé (valeur ÷ coût, ancré dans le code actuel)
Vague 1 — pure plomberie config (pas de kernel, ~1 jour)
rope_scalingtype llama3 : précalculer les inv_freq modifiés (formule ci-dessous) — le kernel RoPE actuel calcule les freqs inline ; passer à une table de freqs précalculée le débloque ET prépare les theta par couche.nope_layers(SmolLM3) : sauterops::ropeselon la couche — unifdansattention.h+ liste en config.attention_bias(Qwen2.5) :Linearsupporte déjà le biais — juste le flag.head_dimdécouplé de d_model/n_heads (Qwen3) :head_dimexplicite en config, wq → n_heads*head_dim.moe_intermediate_sizeséparé +norm_topk_proboptionnel +routed_scaling_factor.relu2comme 3e activation (kernel elementwise trivial + backward).norm_placement: pre | post_olmo | sandwich: recomposition dansTransformerBlock.
Vague 2 — un patch du kernel d'attention (la plus grosse valeur kernel)
8. sliding_window + layer_types par couche (Mistral, Gemma 2/3, GPT-OSS) : dans les kernels flash (scalar + MMA) et la référence, le masque causal devient q−W < k ≤ q ; bonus perf réel : les tuiles KV hors fenêtre sont sautées (borne de boucle), pas juste masquées. Avec rope_theta par type de couche (Gemma 3).
9. attn_logit_softcap (Gemma 2) : cap·tanh(s/cap) sur les scores pré-softmax, dans le kernel (l'op softcap + son backward existent déjà côté élémentwise).
10. Attention sinks (GPT-OSS) : initialiser le (max, somme) du softmax online avec le logit sink appris par tête — quelques lignes dans le kernel flash.
Vague 3 — investissements
11. ✅ Routage sigmoïde + équilibrage sans aux loss (V3/K2) — FAIT : moe_scoring: "sigmoid", sélection top-k sur s+b / gates sur s seul (topk_renorm biaisé), moe_bias_gamma (update ±γ par charge observée après chaque step, compteurs GPU), routed_scaling_factor, moe_d_ff séparé, first_k_dense, moe_norm_topk. Le biais d'équilibrage est checkpointé (param sans gradient).
12. QK-Clip / MuonClip (K2) : monitorer le logit d'attention max par tête (le kernel flash garde déjà le max par ligne — l'exposer), rescaler W_q/W_k si > τ. — restant
13. MLA (V3/K2) : le seul gros chantier kernel — d_k(192) ≠ d_v(128) + RoPE partiel sur une tranche de 64. Praticable d'abord en « décompressé » si le kernel accepte d_k≠d_v. — restant
14. MTP (V3) : bloc supplémentaire + tête partagée, coût trainer modéré. — restant
Formules clés
RoPE scaling llama3 (par composante, old=8192) :
wavelen = 2π/inv_freq
si wavelen < old/high_freq_factor : inchangé
si wavelen > old/low_freq_factor : inv_freq/factor
sinon : blend lisse s=(old/wavelen − low)/(high − low) ; (1−s)·inv_freq/factor + s·inv_freqRoutage V3 : s_i = σ(u·e_i) ; top-k sur s_i + b_i (b pour la sélection seulement) ; g_i = s_i/Σ_topk s_j ; y = x + scaling·Σ g_i·E_i(x) + E_shared(x) ; après chaque step b_i ∓ γ selon surcharge/sous-charge (γ=0.001).
MLA (d=7168) : c_kv = W_dkv·x (→512) ; k_rope = RoPE(W_kr·x) (→64, partagé entre têtes) ; par tête : k_h = [W_ukv,h·c_kv ; k_rope] (192), v_h = W_uv,h·c_kv (128) ; query via c_q (1536) ; scale 1/√192.
Sources
Llama rope_scaling (HF rope_utils) · Mistral 7B · Ministral · Qwen3 · Gemma 3 · DeepSeek-V3 config · V3 report · Kimi K2 config · K2 report (MuonClip) · GPT-OSS · OLMo 2 · SmolLM3 · nanochat gpt.py · Raschka — comparaison d'architectures