Ternarization_instructions_27b_fr.md
6.3 KB · 117 lines · markdown Raw
1 # Instructions de ternarisation et d’export vers GGUF/TQ2_0 — Qwen3.8-27B
2
3 Modèle : `/mnt/nfs_share/Qwen3_8/`, classe `JiRackDeltaNet_27b.py`
4 (hybride 3:1 Gated-DeltaNet / Gated-Attention, `model_type: qwen3_5`).
5
6 ## Statut (mis à jour le 2026-08-28)
7
8 ✅ **Fait :**
9
10 - Conversion HF → format personnalisé, tests de toute la pipeline
11 - **Bloc MTP implémenté et intégré** (voir section 0 ci-dessous)
12 - Conversion GGUF réussie (pt → safetensors → GGUF, bf16)
13 - Smoke test réussi : le modèle répond correctement (« Paris » à « What is the capital of France? »)
14 - **Perplexité de référence : PPL = 5.4957 ± 0.12** (wikitext-100k, c=2048, bf16, lambda=0.0)
15 - Fichier : `/mnt/nfs_share/Qwen3_8/wiki_100k.txt` (utiliser le même fichier pour la comparaison après ternarisation)
16 - Journal : `/mnt/nfs_share/Qwen3_8/ppl_baseline_bf16.log` (dernière ligne, 11 chunks)
17
18 ⏳ **Prochaines étapes :** warmup QAT lambda 0.0→1.0, matérialisation de la ternarisation, export TQ2_0, perplexité finale
19
20 Checkpoint actuel : `/mnt/nfs_share/Qwen3_8/qwen38_27b_checkpoint_migrated.pt` (53,79 Go, lambda=0.0)
21
22 - Contient les poids MTP (transférés depuis le checkpoint HF d’origine), mais n’est pas encore ternarisé.
23 - Il sera converti en un nouveau fichier `.pt` avec prise en charge native de MTP via le script `convert_qwen38_to_jirack.py` (version mise à jour).
24
25 Mécanisme de quantification (identique au modèle 32B) : `BitLinear` calcule gamma comme **absmean par tenseur** ;
26 les poids sont projetés sur exactement trois valeurs (`-gamma, 0, +gamma`), ce qui garantit un aller-retour TQ2_0 sans perte.
27
28 **Différence par rapport au modèle 32B :** deux ajustements architecturaux — voir les sections 0 et 1.
29
30 ---
31
32 ## 0. Bloc MTP (Medusa Token Prediction, couche 64)
33
34 **De quoi s’agit-il :** la nouvelle architecture llama.cpp (version ≥ ca3d5a3e1, tag b10665) exige
35 une couche MTP pour les modèles de type `qwen3_5`. MTP est un décodeur spéculatif (couche draft)
36 destiné à accélérer la génération par lots, à la manière de Medusa. Implémenté dans `JiRackDeltaNet_27b.py`.
37
38 **Ce qui a déjà été fait :**
39
40 - La classe `MTPBlock` est implémentée dans le code (15 tenseurs : fc, normes).
41 - Les poids MTP ont été portés depuis le checkpoint HF original de Qwen3.8-27B.
42 - Ils ont été intégrés dans le fichier GGUF lors de la conversion (`qwen38_27b_jirack_baseline_bf16.gguf`).
43 - MTP **n’est pas** inclus dans le processus de ternarisation (il reste en pleine précision, ~100 Mo de poids).
44
45 **Pourquoi nous ne ternarisons pas MTP :**
46
47 - La couche MTP n’est pas indispensable à la génération standard (elle n’est utilisée que lorsque le drapeau de décodage spéculatif est activé).
48 - Ternariser les 400 couches principales suffit pour atteindre la compression cible.
49 - Conserver MTP en pleine précision préserve la qualité du décodeur spéculatif pour les utilisateurs qui choisissent de l’activer.
50
51 **Pour le warmup QAT et la matérialisation :** il suffit d’ignorer les couches `mtp.*` lors de l’application
52 du patch `BitLinear` — c’est déjà géré par `apply_bitlinear_patch(skip_mtp=True)`.
53
54 ---
55
56 ## 1. Ce qui est ternarisé et ce qui ne l’est pas (TROIS EXCEPTIONS)
57
58 Par défaut (`include_gate_proj=False` dans `apply_bitlinear_patch`),
59 les 400 couches linéaires suivantes sont patchées :
60
61 - `linear_attn.in_proj_qkv`, `linear_attn.in_proj_z`, `linear_attn.out_proj`
62 - `self_attn.q_proj`, `k_proj`, `v_proj`, `o_proj`
63 - `mlp.gate_proj`, `up_proj`, `down_proj`
64
65 **Les éléments suivants ne sont PAS patchés et doivent rester en pleine précision :**
66
67 - `linear_attn.in_proj_b` et `linear_attn.in_proj_a` — ce sont les scalaires de porte
68 pour DeltaNet (projections dt/A-gate) ; ils contrôlent la stabilité de la récurrence. C’était une décision délibérée : le risque de déstabiliser la formule de récurrence elle-même
69 par une quantification ternaire trop grossière a été jugé plus grand que l’économie
70 potentielle.
71 - Toutes les couches de normalisation (`RMSNorm`, `RMSNormGated`), `conv1d` (depthwise,
72 4 taps), `dt_bias`, `A_log`.
73 - `embed_tokens`, `lm_head`.
74
75 **Implication pour l’export GGUF :** lors de la quantification vers GGUF, il faut indiquer
76 explicitement le type de tenseur pour ces couches (`--tensor-type` ou
77 l’équivalent dans votre version de `llama-quantize`), afin que
78 `in_proj_b`/`in_proj_a`, conv1d et toutes les couches de normalisation restent en f16/f32
79 au lieu d’être soumises au réglage général TQ2_0 d’un passage de quantification sur tout le fichier. Sans spécification explicite, certains scripts de quantification
80 appliquent un seul type à tous les tenseurs éligibles — ce qui casserait
81 silencieusement le mécanisme de récurrence.
82
83 ## 2. Risque architectural — DÉJÀ VÉRIFIÉ ✅
84
85 Gated DeltaNet est une architecture hybride encore récente dans llama.cpp. Le risque tenait au fait
86 que son support est moins mature que celui de l’attention Qwen2 dans les modèles 32B.
87
88 **Vérifié (2026-08-28) :**
89
90 - Conversion GGUF (`pt` → `safetensors` → `qwen38_27b_jirack_baseline_bf16.gguf`) :
91 réussie, 866 tenseurs, 54,6 Go
92 - Smoke test (llama-cli avec chat-template) : le modèle répond correctement
93 (« The capital of France is Paris » avec couche de raisonnement)
94 - Perplexité de référence sur wikitext-100k : **PPL = 5.4957 ± 0.12**
95 (11 chunks de 2048 tokens, bf16, lambda=0.0)
96
97 **Conclusion :** l’architecture fonctionne correctement dans llama.cpp. Toute dégradation
98 supplémentaire pendant la quantification ternaire/TQ2_0 viendra du processus de quantification lui-même,
99 et non de bugs d’implémentation.
100
101 **Note pour plus tard :** si la perplexité explose de façon anormale (>>6.5) pendant TQ2_0,
102 ce n’est pas un risque architectural — cela indique un problème de warmup QAT
103 ou de couches mal exclues.
104
105 ## 3. Entraînement QAT (lorsque la quantification ternaire est requise)
106
107 1. Charger le checkpoint :
108
109 ```python
110 model = JiRackQwen38ForCausalLM.load_checkpoint(
111 "/data/qwen38_27b_checkpoint_migrated.pt", dtype=torch.bfloat16)
112 ```
113
114 2. Augmenter progressivement le lambda de warmup de 0.0 à 1.0 (comme pour le modèle 32B,
115 en utilisant `set_lambda` sur tous les modules qui prennent en charge cette méthode).
116 3. Sauvegarder le checkpoint entraîné.
117