Ternarization_instructions_27b_ru.md
13.5 KB · 213 lines · markdown Raw
1 # Инструкция по тернаризации и экспорту в GGUF/TQ2_0 -- Qwen3.8-27B
2
3 Модель: `/mnt/nfs_share/Qwen3_8/`, класс `JiRackDeltaNet_27b.py`
4 (гибрид 3:1 Gated-DeltaNet / Gated-Attention, `model_type: qwen3_5`).
5
6 ## Статус (обновлено 2026-08-28)
7
8 ✅ **Готово:**
9 - Конвертация HF → hand-rolled формат, full-pipeline testing
10 - **MTP-блок реализован и встроен** (см. раздел 0 ниже)
11 - GGUF-конвертация успешна (pt → safetensors → GGUF, bf16)
12 - Smoke-test пройден: модель отвечает корректно ("Paris" на "What is the capital of France?")
13 - **Baseline перплексия: PPL = 5.4957 ± 0.12** (wikitext-100k, c=2048, bf16, lambda=0.0)
14 - Файл: `/mnt/nfs_share/Qwen3_8/wiki_100k.txt` (используй этот же для сравнения после терн)
15 - Лог: `/mnt/nfs_share/Qwen3_8/ppl_baseline_bf16.log` (финальная строка, 11 чанков)
16
17 ⏳ **Вперёд:** QAT warmup lambda 0.0→1.0, материализация тернарки, TQ2_0-экспорт, финальная перплексия
18
19 Текущий чекпоинт: `/mnt/nfs_share/Qwen3_8/qwen38_27b_checkpoint_migrated.pt` (53.79GB, lambda=0.0)
20 - Содержит MTP веса (привитые из оригинального HF чекпоинта), но ещё не терн.
21 - Будет перегнан скриптом `convert_qwen38_to_jirack.py` (обновлённая версия) в новый `.pt` с нативным MTP.
22
23 Механизм квантования (как в 32B): `BitLinear` считает gamma как **per-tensor absmean**,
24 веса → ровно три значения (`-gamma, 0, +gamma`), TQ2_0 round-trip без потерь.
25 **Отличие от 32B:** две поправки в архитектуре -- см. разделы 0 и 1.
26
27 ---
28
29 ## 0. MTP-блок (Medusa Token Prediction, layer 64)
30
31 **Что это:** Новая архитектура llama.cpp (версия ≥ ca3d5a3e1, tag b10665) требует
32 MTP-слой для моделей типа `qwen3_5`. MTP -- это спекулятивный декодер (draft-layer),
33 улучшение для батчевой генерации, аналогично Medusa. Реализация в `JiRackDeltaNet_27b.py`.
34
35 **Что уже сделано:**
36 - Класс `MTPBlock` реализован в коде, 15 тензоров (fc, norms)
37 - MTP веса привиты из оригинального Qwen3.8-27B HF-чекпоинта
38 - Встроены в GGUF при конвертации (`qwen38_27b_jirack_baseline_bf16.gguf`)
39 - MTP **не** включена в тернаризацию (остаётся в full precision, ~100MB весов)
40
41 **Почему не тернируем MTP:**
42 - MTP-слой несущественен для основной генерации (используется только с флагом спекулятивки)
43 - Тернаризация всех 400 основных слоёв достаточна для целевого сжатия
44 - Full-precision MTP сохраняет качество спекулятивного декодера для тех, кто его включит
45
46 **Для QAT warmup и материализации:** просто пропускайте `mtp.*` слои при применении
47 `BitLinear` патча -- это уже сделано в `apply_bitlinear_patch(skip_mtp=True)`.
48
49 ---
50
51 ## 1. Что тернаризуется, а что нет (ТРИ ИСКЛЮЧЕНИЯ)
52
53 По умолчанию (`include_gate_proj=False` в `apply_bitlinear_patch`)
54 патчатся 400 линейных слоёв:
55 - `linear_attn.in_proj_qkv`, `linear_attn.in_proj_z`, `linear_attn.out_proj`
56 - `self_attn.q_proj`, `k_proj`, `v_proj`, `o_proj`
57 - `mlp.gate_proj`, `up_proj`, `down_proj`
58
59 **НЕ патчатся и должны остаться full precision:**
60 - `linear_attn.in_proj_b` и `linear_attn.in_proj_a` -- гейтовые скаляры
61 DeltaNet (dt/A-gate проекции), управляют устойчивостью рекуррентности.
62 Решение осознанное: риск дестабилизировать саму рекуррентную формулу
63 через грубое тернарное квантование посчитали выше потенциальной
64 экономии.
65 - Все нормализации (`RMSNorm`, `RMSNormGated`), `conv1d` (depthwise,
66 4 tap), `dt_bias`, `A_log`.
67 - `embed_tokens`, `lm_head`.
68
69 **Следствие для GGUF-экспорта:** при квантовании в GGUF нужно явно
70 указывать тип тензора для этих слоёв (через `--tensor-type` или
71 эквивалент в используемой версии `llama-quantize`), чтобы
72 `in_proj_b`/`in_proj_a`, conv1d и все нормы остались в f16/f32, а НЕ
73 попали под общий TQ2_0 при квантовании всего файла разом. Если этого не
74 сделать явно, часть скриптов квантования применяет один тип ко всем
75 подходящим тензорам без разбора -- и это тихо сломает рекуррентность.
76
77 ## 2. Архитектурный риск -- УЖЕ ПРОВЕРЕН ✅
78
79 Gated DeltaNet -- гибридная и новая для llama.cpp архитектура. Риск был в том, что
80 поддержка менее вызвана, чем Qwen2-attention в 32B.
81
82 **Проверено (2026-08-28):**
83 - GGUF-конвертация (`pt` → `safetensors` → `qwen38_27b_jirack_baseline_bf16.gguf`):
84 успешно, 866 тензоров, 54.6GB
85 - Smoke-test (llama-cli с chat-template): модель отвечает корректно
86 ("The capital of France is Paris" с reasoning layer)
87 - Baseline перплексия на wikitext-100k: **PPL = 5.4957 ± 0.12**
88 (11 чанков по 2048 токенов, bf16, lambda=0.0)
89
90 **Вывод:** архитектура в llama.cpp работает правильно. Дальнейшая деградация
91 при терн/TQ2_0 будет от самого квантования, а не от bug-ов в реализации.
92
93 **На будущее:** если при TQ2_0 перплексия подскочит аномально (>>6.5), это не
94 архитектурный риск -- это признак проблемы в QAT warmup или неправильно исключённых слоёв.
95
96 ## 3. QAT-обучение (когда есть заказ на тернарку)
97
98 1. Загрузить чекпоинт:
99 ```python
100 model = JiRackQwen38ForCausalLM.load_checkpoint(
101 "/data/qwen38_27b_checkpoint_migrated.pt", dtype=torch.bfloat16)
102 ```
103 2. Прогнать warmup lambda от 0.0 до 1.0 постепенно (аналогично 32B, через
104 `set_lambda` на всех модулях с этим методом).
105 3. Сохранить обученный чекпоинт через `save_checkpoint(...)` -- lambda и
106 список патченных слоёв уже сохраняются этим методом.
107
108 ## 4. Обратный экспорт в HF-формат (lambda должна быть 1.0)
109
110 Та же логика, что у 32B (материализовать тернарные веса в обычные
111 float-тензоры под настоящими HF-именами `qwen3_5`), но со строгим
112 условием: слои из списка "НЕ патчатся" в разделе 1 копируются как есть,
113 без тернаризации, остальные 400 -- материализуются как `-gamma/0/+gamma`.
114
115 ```python
116 model.eval()
117 model.set_lambda(1.0)
118 hf_out = {}
119 # embed_tokens, lm_head, model.norm -- как есть, без изменений
120 # per layer:
121 # input_layernorm, post_attention_layernorm -- как есть
122 # linear_attn.dt_bias, A_log, conv1d.weight, norm.weight -- как есть
123 # linear_attn.in_proj_b, in_proj_a -- КАК ЕСТЬ (full precision!)
124 # linear_attn.in_proj_qkv, in_proj_z, out_proj -- материализовать ternary
125 # self_attn.q/k/v/o_proj, q_norm, k_norm -- q/k/v/o тернарные, normы как есть
126 # mlp.gate_proj, up_proj, down_proj -- материализовать ternary
127 #
128 # материализация тернарного слоя (тот же приём, что в 32B):
129 # w = lin.weight.float()
130 # gamma = w.abs().mean().clamp(min=lin.eps)
131 # w_ternary = torch.clamp(torch.round(w / gamma), -1, 1) * gamma
132 ```
133
134 Дальше сохранить как HF-чекпоинт с `config.json` под `model_type: qwen3_5`,
135 `architectures: ["Qwen3_5ForCausalLM"]`, со всеми константами из
136 `JiRackDeltaNet_27b.py` (HIDDEN_SIZE=5120, VOCAB_SIZE=248320,
137 NUM_LAYERS=64, layer_types-паттерн и т.д.)
138
139 ## 5. GGUF-конвертация
140
141 ```bash
142 python convert_hf_to_gguf.py qwen38_27b_ternary_hf/ \
143 --outfile qwen38_27b_ternary_f16.gguf --outtype f16
144 ```
145 Перед этим шагом проверить версию `llama.cpp`/`convert_hf_to_gguf.py` на
146 поддержку `model_type: qwen3_5` -- поддержка гибридных Qwen3.5-моделей
147 могла появиться позже, чем у обычных Qwen2/Qwen3, и не факт что она уже
148 есть в используемой версии.
149
150 ```bash
151 ./llama-quantize qwen38_27b_ternary_f16.gguf \
152 qwen38_27b_ternary_tq2_0.gguf TQ2_0 \
153 --tensor-type "*.in_proj_b.weight=F16" \
154 --tensor-type "*.in_proj_a.weight=F16"
155 # (плюс аналогичные исключения для conv1d/dt_bias/A_log/норм, если
156 # инструмент квантования иначе трогает их по умолчанию -- проверить
157 # фактическое поведение конкретной версии llama-quantize)
158 ```
159
160 ## 6. Проверка
161
162 - Перплексия базовой f16/Q4 GGUF **против оригинальной HF-модели** --
163 обязательно, отдельно от эффекта тернаризации (см. раздел 2).
164 - Перплексия TQ2_0 GGUF против f16 GGUF той же (уже тернарной) модели --
165 чтобы отделить деградацию от warmup от возможных потерь GGUF round-trip.
166 - Несовпадение схемы активаций (per-token int8 у нас vs q8_K по блокам в
167 llama.cpp) -- тот же источник погрешности, что и в 32B, ожидаемый и
168 неизбежный.
169
170 ## 7. Чек-лист перед выкладкой
171
172 - [ ] Обычный (не тернарный) GGUF проверен на корректность DeltaNet-слоёв
173 - [ ] Warmup lambda 0->1 завершён, финальная lambda = 1.0 в чекпоинте
174 - [ ] Обратный экспорт сделан с сохранением полного списка "неприкасаемых" слоёв
175 - [ ] `convert_hf_to_gguf.py` поддерживает `qwen3_5` и прошёл без missing/unexpected keys
176 - [ ] `llama-quantize` в TQ2_0 с явным разделением типов (gate-слои остались f16)
177 - [ ] Перплексия TQ2_0 сравнена и с f16-тернарной версией, и с оригинальной HF-моделью
178
179 ## 8. `config.json` для HF-папки (УТОЧНЕНИЕ ПРИ ЭКСПОРТЕ)
180
181 **Ошибка в предыдущих версиях (исправлено 2026-08-28):**
182
183 Qwen3.8-27B в HF используют **вложенную мультимодальную структуру**
184 (`Qwen3_5ForConditionalGeneration`), где реальные параметры модели лежат
185 в подполе `text_config` (типа `qwen3_5_text`). Плоская конфигурация
186 привела к ошибке конвертации:
187 ```
188 RuntimeError: shape '[16, 2, 1, 1]' is invalid for input of size 48
189 ```
190 (конвертер молча подставлял default value_heads=32 вместо реального 48).
191
192 **Правильный подход:**
193 1. Скачать оригинальный `config.json` из Qwen3.8-27B на HF
194 2. Извлечь поле `text_config` (это и есть реальная text-модель, а не wrapper)
195 3. Установить `architectures: ["Qwen3_5ForCausalLM"]` (или то, что требует llama.cpp)
196 4. Сохранить как `/mnt/nfs_share/Qwen3_8/config.json`
197
198 **Текущий файл** (на сервере): `/mnt/nfs_share/Qwen3_8/config.json` --
199 это уже исправленная вложенная версия, которая работает (проверено конвертацией).
200
201 **Старые версии** (как для информации):
202 - `config_27b_flat_broken` -- плоская попытка, привела к ошибке shape (backup на сервере)
203 - `config_27b.json`, `config_27b_v2.json` -- промежуточные версии (obsolete)
204
205 **При терн/HF-экспорте:** используй текущий `config.json` как шаблон, просто убедись
206 что `text_config` содержит всё нужное (hidden_size, layer_types, etc.).
207
208 **Ручная проверка:**
209 ```bash
210 python3 -c "import json; c=json.load(open('config.json')); print(c['text_config']['hidden_size'])"
211 # Должно быть 5120
212 ```
213