ZichenAI commited on
Commit
e282eae
·
verified ·
1 Parent(s): d57da8f

Add files using upload-large-folder tool

Browse files
Files changed (3) hide show
  1. README.md +264 -390
  2. config.json +79 -79
  3. model.safetensors +1 -1
README.md CHANGED
@@ -1,390 +1,264 @@
1
- ---
2
- library_name: transformers
3
- pipeline_tag: text-generation
4
- language:
5
- - zh
6
- - en
7
- license: other
8
- license_name: baihu-custom-license
9
- license_link: https://huggingface.co/NovaAI6868/BaiHu-V1-Flash/blob/main/LICENSE.custom.md
10
- base_model: Qwen/Qwen3-0.6B-Base
11
- tags:
12
- - sparse-attention
13
- - subq
14
- - ssa
15
- - long-context
16
- - commercial-license-required
17
- - text-generation
18
- ---
19
-
20
- # BaiHu-V1-Flash
21
-
22
- **BaiHu-V1-Flash** is a retrofit of the dense-attention model `Qwen/Qwen3-0.6B-Base` into an
23
- **SSA (Sparse-attention + SubQ)** architecture, obtained by continued pretraining.
24
-
25
- - Base model: `Qwen/Qwen3-0.6B-Base` (28 layers / 16 Q heads / 8 KV heads / head_dim 128 / 32K context / tied embeddings)
26
- - Parameters: 598.8M
27
- - Training data: mixed Chinese + English (Fineweb-Edu-Chinese-V2.1 + fineweb-edu, 50/50)
28
- - License: **free for personal use; a paid license is required for commercial use** (see "License" below)
29
-
30
- ---
31
-
32
- ## 1. Architecture: SSA (three paths, each with its own softmax, then summed)
33
-
34
- Every layer keeps the base model's MLP / RMSNorm weights and replaces full attention with an
35
- SSA layer built from three parallel paths:
36
-
37
- | Path | Role | Complexity |
38
- |---|---|---|
39
- | `shared` | every query sees all **completed** blocks through one compressed vector per block | `O(T·T/B)` |
40
- | `local` | dense causal attention over the most recent window | `O(T·w)` |
41
- | `sparse` (**SubQ**) | only 4 of 16 query heads produce block scores, shared across the whole head group; real attention is computed only for the selected top-k blocks | `O(T·k·B)` |
42
-
43
- ### Hyperparameters
44
-
45
- | Parameter | Value | Meaning |
46
- |---|---|---|
47
- | `ssa_block_size` | 64 | block size B |
48
- | `ssa_top_k` | 8 | number of blocks selected by the sparse path |
49
- | `ssa_local_blocks` | 2 | local window = 3 × 64 = 192 tokens |
50
- | `ssa_num_subq_heads` | 4 | SubQ heads, r = 16 / 4 = 4 |
51
- | `ssa_router_dim` | 128 | router subspace dimension |
52
- | `ssa_compress_dim` | 128 | block compression dimension |
53
-
54
- Only **2.75M** parameters are new (≈0.46% of the model); all other weights are inherited
55
- from the base model.
56
-
57
- ---
58
-
59
- ## 2. Training
60
-
61
- | Item | Setting |
62
- |---|---|
63
- | Starting point | a conversion checkpoint that is **bit-exact** with the base model (`max\|Δlogit\| = 0.000e+00`) |
64
- | Tokens seen | 5.0M (≈0.25 epoch of the corpus) |
65
- | Sequence length | 512 (must be a multiple of `ssa_block_size = 64`) |
66
- | Effective batch | 8192 tokens (batch 2 × grad_accum 8) |
67
- | Precision | fp32 |
68
- | Optimizer | SGD with momentum 0.9 |
69
- | Learning rate | 5e-4, 50-step warmup, cosine decay to 10% |
70
- | Hardware | single NVIDIA GTX TITAN X (Maxwell, sm_52, 12.9 GB) |
71
- | Throughput | ≈274 tokens/s |
72
-
73
- ### Critical issues found and fixed during this retrofit
74
-
75
- Several defects silently break training and are worth documenting:
76
-
77
- 1. **Both new branches had identically zero gradients (blocking).** To make the converted
78
- model bit-exact with the base model, `compress_out` and `router_out` were initialized to
79
- exactly zero, and the branches were skipped entirely by a gate. The branch output was
80
- therefore always zero, so the back-propagated gradient was also always zero: all 2.75M
81
- SSA parameters **stayed frozen for the entire run** and the sparse attention was dead
82
- code. The fix is a small non-zero initialization.
83
- 2. **Routing was non-differentiable.** `top-k` produces hard indices, and indexing is not
84
- differentiable. If the routing scores are used only to decide *which* blocks to read and
85
- never enter the softmax, the gradients of `router_q` / `router_k` are **exactly zero** —
86
- the router can never learn to route. The fix is to feed the selected blocks' scores,
87
- squashed through `tanh` and gently scaled, into the attention logits as an additive bias.
88
- 3. **The shared summary was a sum, not a mean.** Its magnitude grew linearly with the
89
- prefix, and because that branch is injected at full weight (a single-element softmax has
90
- probability exactly 1), it swamped the residual stream: hidden states grew from 0.2 to
91
- about 7 in layer 0 and to about 1900 by layer 27, and validation loss went 3.56 → 11.38.
92
- The fix is to divide by the token count.
93
- 4. **The shared branch needs an explicit gate.** With a single-element softmax the
94
- probability is always 1, so the initialization scale of `compress_out` cannot control the
95
- injection strength at all (measured: scales from 1e-4 to 0.03 all left the loss at
96
- exactly 7.3526). A learnable scalar gate, initialized to a small positive value, lets the
97
- optimizer decide how far to open it.
98
-
99
- ---
100
-
101
- ## 3. How to Run Inference
102
-
103
- ### 3.1 Important: this is a custom architecture
104
-
105
- `BaiHu-V1-Flash` uses `model_type: baihu_ssa`, which is **not** in the Transformers
106
- registry. Loading it with a plain `AutoModelForCausalLM.from_pretrained(...)` fails with:
107
-
108
- ```
109
- ValueError: The checkpoint you are trying to load has model type `baihu_ssa`
110
- but Transformers does not recognize this architecture.
111
- ```
112
-
113
- You must register the config and model classes first. This is a one-time, three-line step
114
- (see below).
115
-
116
- Also note: **this model does not go through `transformers.GenerationMixin`.** SSA owns its
117
- own KV-cache layout (`BaiHuSSACache`) because the cache stores per-block compressed prefix
118
- sums rather than a plain growing key/value tensor. Call the model's own `generate()` method;
119
- beam search is not supported.
120
-
121
- ### 3.2 Setup
122
-
123
- ```bash
124
- # The SSA implementation lives in the project repository (not on the Hub),
125
- # because it is a custom architecture.
126
- git clone <this-project-repo> ssa_model
127
- cd ssa_model
128
- uv venv --python 3.11 .venv
129
-
130
- # Ampere or newer (RTX 30xx/40xx, A100, ...): any recent torch works.
131
- uv pip install --python .venv/Scripts/python.exe \
132
- "numpy>=1.26" "transformers>=4.51" "safetensors>=0.4" "torch>=2.6"
133
-
134
- # Maxwell / Pascal / Volta (GTX 9xx/10xx, TITAN X, V100, ...): pin torch 2.7.1 —
135
- # see the GPU note below for why `torch>=2.6` would resolve to a build without your kernels.
136
- uv pip install --python .venv/Scripts/python.exe \
137
- "numpy>=1.26" "transformers>=4.51" "safetensors>=0.4" "torch==2.7.1"
138
- ```
139
-
140
- > **GPU note (Maxwell and older).** If you are on an NVIDIA Maxwell card such as the
141
- > GTX TITAN X (compute capability sm_52), recent PyTorch wheels no longer ship kernels for
142
- > it: PyTorch **2.8 removed sm_50/sm_60**, and 2.8.0+cu126 only contains `sm_61…sm_90`.
143
- > You will get `CUDA error: no kernel image is available for execution on the device`.
144
- > Use **torch 2.7.1+cu126**, which still contains `sm_50`. Also train/infer in **fp32** on
145
- > Maxwell: measured fp32 5.30 / fp16 4.46 / bf16 3.20 TFLOPS, so half precision is a loss.
146
-
147
- ### 3.3 Minimal working example
148
-
149
- ```python
150
- import os
151
- import sys
152
-
153
- import torch
154
- from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer
155
-
156
- # Make the custom implementation importable (path to the cloned repo's src/)
157
- sys.path.insert(0, os.path.join("ssa_model", "src"))
158
-
159
- from baihu_ssa.configuration_baihu_ssa import BaiHuSSAConfig
160
- from baihu_ssa.model_baihu_ssa import BaiHuSSAForCausalLM
161
-
162
- # ---- register the custom architecture (REQUIRED) ----
163
- AutoConfig.register("baihu_ssa", BaiHuSSAConfig, exist_ok=True)
164
- AutoModelForCausalLM.register(BaiHuSSAConfig, BaiHuSSAForCausalLM, exist_ok=True)
165
-
166
- REPO = "NovaAI6868/BaiHu-V1-Flash"
167
- device = "cuda" if torch.cuda.is_available() else "cpu"
168
- dtype = torch.float32 if device == "cpu" else torch.float32 # see GPU note: fp32
169
-
170
- tokenizer = AutoTokenizer.from_pretrained(REPO)
171
- model = AutoModelForCausalLM.from_pretrained(REPO, dtype=dtype).to(device).eval()
172
-
173
- prompt = "人工智能的未来是"
174
- input_ids = tokenizer(prompt, return_tensors="pt").input_ids.to(device)
175
-
176
- with torch.no_grad():
177
- out = model.generate(
178
- input_ids,
179
- max_new_tokens=32,
180
- do_sample=False, # greedy; supported
181
- # do_sample=True, temperature=0.8, top_k=50, # sampling also supported
182
- )
183
- print(tokenizer.decode(out[0], skip_special_tokens=True))
184
- ```
185
-
186
- ### 3.4 `generate()` parameters
187
-
188
- The built-in `generate()` is a self-contained decoding loop (greedy or
189
- temperature/top-k sampling). Supported arguments:
190
-
191
- | Argument | Default | Meaning |
192
- |---|---|---|
193
- | `max_new_tokens` | 32 | number of tokens to generate |
194
- | `do_sample` | `False` | `False` = greedy, `True` = sample |
195
- | `temperature` | 1.0 | sampling temperature (used when `do_sample=True`) |
196
- | `top_k` | `None` | top-k sampling cutoff (used when `do_sample=True`) |
197
- | `eos_token_id` | `None` | stop early when all sequences emit this token |
198
- | `use_cache` | `True` | keep the SSA cache between steps; **leave on**, decoding is much slower without it |
199
-
200
- **Not supported:** beam search (the SSA cache has no `reorder_cache`), and
201
- `transformers` generation utilities such as `logits_processor` / `stopping_criteria`.
202
-
203
- ### 3.5 Computing perplexity / loss
204
-
205
- ```python
206
- with torch.no_grad():
207
- inputs = tokenizer("Your text here", return_tensors="pt").to(device)
208
- loss = model(**inputs, labels=inputs["input_ids"]).loss
209
- print("loss", loss.item(), "ppl", torch.exp(loss).item())
210
- ```
211
-
212
- ### 3.6 Throughput you should expect
213
-
214
- Measured on a GTX TITAN X (sm_52, 12.9 GB), fp32, prompt 1024 tokens, 64 generated tokens:
215
-
216
- | | Qwen3-0.6B-Base | BaiHu-V1-Flash |
217
- |---|---|---|
218
- | Prefill / TTFT | 503 ms | 1431 ms |
219
- | Per-token decode (TPOT) | 36.6 ms | 80.5 ms |
220
- | Decode throughput | 27.3 tok/s | 12.4 tok/s |
221
- | Peak memory (generation) | 9.91 GB | 9.01 GB |
222
-
223
- **BaiHu-V1-Flash is currently ~2.2× slower to decode than the base model**, even though the
224
- sparse path reads fewer keys. The reason is implementation, not architecture: the attention
225
- layer loops over query blocks in Python and launches many small kernels, so kernel-launch
226
- overhead dominates the FLOPs saved. Memory use is lower; speed is not yet a win.
227
- See "Known Limitations" below.
228
-
229
- ### 3.7 Expected output quality (honest warning)
230
-
231
- This checkpoint has seen only **5.0M tokens**, so the new branches are far from converged.
232
- Greedy decoding tends to **repeat itself** and it is **worse than the base model on
233
- perplexity** (ppl 82.3 vs 39.4 at 512 tokens). A real greedy sample:
234
-
235
- ```
236
- prompt: 人工智能的未来是
237
- output: 人工智能的未来是怎样的?人工智能的未来是怎样的?人工智能的未来是怎样的?…
238
- ```
239
-
240
- Use it to study or continue the SSA retrofit — **do not** expect it to match
241
- `Qwen3-0.6B-Base` as a general-purpose model yet.
242
-
243
- ### 3.8 Troubleshooting
244
-
245
- | Error | Cause | Fix |
246
- |---|---|---|
247
- | `does not recognize this architecture` / `KeyError: 'baihu_ssa'` | custom architecture not registered | call `AutoConfig.register` + `AutoModelForCausalLM.register` as in 3.3 |
248
- | `CUDA error: no kernel image is available for execution on the device` | installed PyTorch has no kernels for your GPU (Maxwell/sm_52) | install torch 2.7.1+cu126 or older (2.8 dropped sm_50/sm_60) |
249
- | `AttributeError: ... has no attribute 'tie_weights'` | an external tool treating the model as `PreTrainedModel` | use a version of the implementation that defines `tie_weights()` (already fixed upstream) |
250
- | `BaiHuSSACache` errors when calling `model.generate(...)` from GenerationMixin | this model bypasses `GenerationMixin` | call `model.generate(...)` on the BaiHu model itself |
251
-
252
- ---
253
-
254
- ## 4. Evaluation
255
-
256
- ### 4.1 Language modeling perplexity (validation set, identical windows)
257
-
258
- | Sequence length | Qwen3-0.6B-Base | BaiHu-V1-Flash |
259
- |---|---|---|
260
- | 512 | 3.6726 / ppl 39.354 | 4.4108 / ppl 82.335 |
261
- | 1024 | 3.4691 / ppl 32.109 | 4.2575 / ppl 70.631 |
262
- | 2048 | 3.0801 / ppl 21.761 | 3.9213 / ppl 50.467 |
263
-
264
- ### 4.2 Standard benchmarks (lm-evaluation-harness)
265
-
266
- | Task | Qwen3-0.6B-Base | BaiHu-V1-Flash | Delta |
267
- |---|---|---|---|
268
- | arc_easy | 0.5550 | 0.6250 | +0.0700 |
269
- | hellaswag | 0.5350 | 0.5200 | -0.0150 |
270
- | piqa | 0.7050 | 0.6950 | -0.0100 |
271
- | winogrande | 0.6300 | 0.6300 | +0.0000 |
272
-
273
- ### 4.3 Inference compute and resource usage
274
-
275
- | Metric | Qwen3-0.6B-Base | BaiHu-V1-Flash |
276
- |---|---|---|
277
- | Parameters (M) | 596.0500 | 598.8000 |
278
- | Prefill peak memory (GB) | 9.6200 | 8.0210 |
279
- | Generation peak memory (GB) | 9.9120 | 9.0050 |
280
- | Prefill latency (s) | 0.5030 | 1.4310 |
281
- | TTFT (ms) | 502.7 | 1431.2 |
282
- | TPOT (ms) | 36.6 | 80.5 |
283
- | Decode throughput (tok/s) | 27.3300 | 12.4200 |
284
- | Attention FLOPs/token (GFLOPs) | 0.1176 | 0.1057 |
285
- | Attention key accesses vs full attention | 1.0000 | 0.8993 |
286
- | GPU utilization mean/max (%) | 93.9 | 43.3 |
287
- | Power mean/max (W) | 179.1 | 124.3 |
288
-
289
- Positive findings: peak inference memory is lower (generation 9.005 vs 9.912 GB, −9.2%), and the
290
- model draws less power because it is not compute-bound.
291
-
292
- Negative findings, stated plainly:
293
-
294
- - **Decode is 2.2× slower** (12.42 vs 27.33 tok/s) and **prefill is 2.8× slower**
295
- (TTFT 1431 vs 503 ms), despite the sparse path reading fewer keys. The current
296
- implementation loops over query blocks in Python and issues many small kernels, so
297
- launch overhead dominates the FLOPs saved. **The sparse attention does not yet pay off
298
- on this hardware.**
299
- - **Attention key accesses are still 89.9% of full attention** at this sequence length.
300
- The reason is structural: the local window already covers 3 blocks (192 tokens) and the
301
- sparse path reads up to `top_k + 1 = 9` blocks from a grid that only has 16 blocks at
302
- 1024 tokens, so the selected set is almost the whole grid. Sparsity only becomes a real
303
- saving once the sequence is long relative to `top_k × block_size` (i.e. well beyond
304
- 10k tokens).
305
- - **Language modeling perplexity is clearly worse than the base model** at every length
306
- tested (ppl 82.3 vs 39.4 at 512; 50.5 vs 21.8 at 2048). This is the honest cost of
307
- shrinking the dense local window from the full prefix to 192 tokens while the new
308
- long-range branches are still very weakly trained.
309
- - **Standard benchmarks are roughly neutral but not better**: arc_easy improves
310
- (+0.070 acc_norm), winogrande is unchanged, while hellaswag (−0.015) and piqa (−0.010)
311
- regress slightly.
312
-
313
- ### 4.4 Sparsity
314
-
315
- Two measurements are reported because they answer different questions and are **not**
316
- interchangeable:
317
-
318
- | Scenario | Average blocks selected per query | Key access ratio vs full attention |
319
- |---|---|---|
320
- | Chunked forward over a 512-token validation window | 0.38 | 0.0938 |
321
- | Prefill of 1024 tokens (steady state) | up to `top_k` | 0.8993 |
322
-
323
- The first number averages over *all* query blocks including the early ones, which have no
324
- completed blocks available to select and therefore read nothing through the sparse path.
325
- The second is the steady-state ratio for later queries, and it is the one that matters for
326
- efficiency — see the note in section 4.3: at these sequence lengths the sparse path is not
327
- yet saving meaningful work.
328
-
329
- ---
330
-
331
- ## 5. Known Limitations
332
-
333
- 1. **Trained for very little.** Only 5.0M tokens (≈0.25 epoch). The new branches have not
334
- converged; ppl is well above the base model and decode is slower (see 3.6 and 4.1). Continuing
335
- to 200M tokens or more is required before the SSA layers can genuinely take over
336
- long-range modeling.
337
- 2. **The shared branch does not appear to help and was actively suppressed by the
338
- optimizer.** Its learnable gate *decreased* over training (0.0100 → 0.0129 at step 100 →
339
- 0.0122 at step 610) instead of growing, meaning the optimizer found the single
340
- prefix-mean summary not worth injecting. This is the single most important thing to
341
- change next: replace it with **one compressed vector per block** (same `O(T·T/B)` cost,
342
- far more information retained).
343
- 3. **Sparse attention is not yet a net win on this hardware.** It reads fewer keys but runs
344
- 2.2× slower because the implementation loops over query blocks in Python and issues many
345
- small kernels. It needs kernel-level batching (or a fused implementation) before the
346
- sparsity can translate into speed.
347
- 4. **Sparsity only pays off at long sequences.** With `top_k=8` and `block_size=64`, the
348
- sparse path can read up to 9 blocks = 576 tokens; at 1024 tokens the grid only has 16
349
- blocks, so the selected set covers most of the context and the local window already
350
- covers the rest. Real savings require sequences well beyond 10k tokens.
351
- 5. **Routing quality is not fully validated.** The distribution of selected top-k blocks
352
- should be checked for degeneration (e.g. always selecting the same blocks). The
353
- non-differentiable-routing bug that would have made this *impossible* to learn has been
354
- fixed (see section 2), but the learned policy has not been analyzed in detail.
355
- 6. **Block size B = 64 was not ablated.** Limited by memory and the Windows WDDM watchdog on
356
- this machine; B ∈ {32, 64, 128} should be swept on a larger GPU.
357
- 7. **Document-boundary packing.** The current packing strategy places multiple documents in
358
- one sequence (separated by EOS).
359
-
360
- ---
361
-
362
- ## 6. License
363
-
364
- **Free for personal use; a paid license is required for commercial use.**
365
-
366
- - ✅ Personal study, research, teaching, hobby projects: **free**, no application required
367
- - ✅ Academic research with public publication: **free** (please cite the source)
368
- - 💰 Internal company/studio use, paid API/SaaS, product integration, client deliverables:
369
- **commercial license required**
370
-
371
- **Commercial licensing contact: novaweb6868@outlook.com**
372
-
373
- Full terms: [LICENSE.custom.md](./LICENSE.custom.md).
374
-
375
- This model is an architectural retrofit of `Qwen/Qwen3-0.6B-Base` (Apache License 2.0).
376
- This license governs only the newly added portions and does not alter the upstream
377
- component's original license.
378
-
379
- ---
380
-
381
- ## 7. Citation
382
-
383
- ```bibtex
384
- @misc{baihu-v1-flash,
385
- title = {BaiHu-V1-Flash: An SSA (Sparse-attention + SubQ) Retrofit of Qwen3-0.6B-Base},
386
- author = {NovaAI6868},
387
- year = {2026},
388
- url = {https://huggingface.co/NovaAI6868/BaiHu-V1-Flash}
389
- }
390
- ```
 
1
+ ---
2
+ library_name: transformers
3
+ pipeline_tag: text-generation
4
+ language:
5
+ - zh
6
+ - en
7
+ license: other
8
+ license_name: baihu-custom-license
9
+ license_link: https://huggingface.co/ZichenAI/BaiHu-V1-Flash/blob/main/LICENSE.custom.md
10
+ base_model: Qwen/Qwen3-0.6B-Base
11
+ base_model_relation: finetune
12
+ tags:
13
+ - sparse-attention
14
+ - subq
15
+ - ssa
16
+ - long-context
17
+ - supervised-fine-tuning
18
+ - transfer-learning
19
+ - commercial-license-required
20
+ - text-generation
21
+ ---
22
+
23
+ # BaiHu-V1-Flash
24
+
25
+ **BaiHu-V1-Flash** is an SSA (Sparse-attention + SubQ) retrofit of `Qwen/Qwen3-0.6B-Base`,
26
+ fine-tuned on **1,294 bilingual multi-turn dialogues** whose purpose is **举一反三** — the
27
+ model is shown one worked rule / format / method in the first turn, and later turns ask it to
28
+ reuse that rule on a *new* case it has never seen.
29
+
30
+ - Base model: `Qwen/Qwen3-0.6B-Base` (28 layers / 16 Q heads / 8 KV heads / head_dim 128 / 32K context / tied embeddings)
31
+ - Parameters: **598.8M** (2.75M are SSA-only modules)
32
+ - Architecture: SSA — every layer runs three attention paths (shared / local / sparse-SubQ), each with its own softmax, then summed
33
+ - Training data: 1,294 synthetic dialogues (English 693 / Chinese 601), six transfer types
34
+ - License: **free for personal use; a paid license is required for commercial use** (see [License](#license))
35
+
36
+ > **Revision note.** This repository previously hosted the base-pretrain checkpoint of the
37
+ > same name (5.0M tokens of continued pretraining, no instruction tuning, no P0 revision).
38
+ > It has been **replaced** by the checkpoint described here. The two are not interchangeable:
39
+ > this one uses the revised shared branch (§2) and is a supervised fine-tune.
40
+
41
+ ---
42
+
43
+ ## 1. Revisions in this release
44
+
45
+ | | earlier release (now replaced) | **this release** |
46
+ |---|---|---|
47
+ | Shared path | one mean vector for the **entire prefix** | **one compressed vector per completed block** (P0 revision) |
48
+ | Training | continued pretraining, 5.0M tokens of web text | supervised fine-tune, 1,294 transfer dialogues (assistant-token loss) |
49
+ | Behaviour | base LM | follows a rule / format established earlier in the dialogue |
50
+
51
+ Both revisions share the same SSA architecture, base weights and license.
52
+
53
+ ## 2. Architecture
54
+
55
+ Every layer keeps the base model's MLP / RMSNorm weights and replaces full attention with
56
+ an SSA layer built from three parallel paths:
57
+
58
+ | Path | Role | Complexity |
59
+ |---|---|---|
60
+ | `shared` | every query attends over **one compressed summary vector per completed block** | `O(T·T/B)` |
61
+ | `local` | dense causal attention over the most recent window | `O(T·w)` |
62
+ | `sparse` (**SubQ**) | only 4 of 16 query heads produce block scores, shared across the head group; real attention is computed only for the selected top-k blocks | `O(T·k·B)` |
63
+
64
+ ### The P0 revision (this release)
65
+
66
+ The previous revision compressed the **whole prefix into a single mean vector**. A query
67
+ could only see one undifferentiated global average, the single-element softmax made that
68
+ average inject at full weight, and during 5.0M tokens of continued pretraining the learned
69
+ gate **shrank instead of growing** (0.0100 → 0.0129 → 0.0122) — i.e. the optimizer actively
70
+ suppressed the branch.
71
+
72
+ This release implements the design the project always documented: **one compressed key/value
73
+ per completed 64-token block**, so queries do a genuine softmax over blocks and "which block
74
+ matters" becomes learnable. Cost is unchanged (`O(T·T/B)`; block summaries are still computed
75
+ once per layer).
76
+
77
+ ### Hyperparameters
78
+
79
+ | Parameter | Value | Meaning |
80
+ |---|---|---|
81
+ | `ssa_block_size` | 64 | block size B |
82
+ | `ssa_top_k` | 8 | blocks selected by the sparse path |
83
+ | `ssa_local_blocks` | 2 | local window = 3 × 64 = 192 tokens |
84
+ | `ssa_num_subq_heads` | 4 | SubQ heads, r = 16 / 4 = 4 |
85
+ | `ssa_router_dim` / `ssa_compress_dim` | 128 / 128 | router subspace / summary width |
86
+
87
+ ## 3. Training
88
+
89
+ **Data.** 1,294 synthetic multi-turn dialogues (en 693 / zh 601), each 6–12 messages: the
90
+ first user turn gives a worked case or a rule, a later user turn introduces a new instance
91
+ that can only be handled by reusing it, and the last turn pushes generalization one step
92
+ further. Six `transfer_kind`s are covered: `rule_induction`, `analogy_transfer`,
93
+ `format_transfer`, `counterfactual`, `cross_domain`, `teaching_loop`.
94
+
95
+ Split (stratified by language × kind, seed 0): **train 1,165 sessions / 345.7K tokens**
96
+ (assistant 188.8K), **val 129 sessions / 38.2K tokens** (assistant 20.5K).
97
+
98
+ **Recipe.** Loss is computed on assistant tokens only; sequences are padded to a multiple of
99
+ the SSA block size (64); rendering uses the tokenizer's own chat template
100
+ (`apply_chat_template`, `add_generation_prompt=False`, i.e. the final assistant turn carries
101
+ Qwen3's empty `<think></think>` block).
102
+
103
+ | Item | Value |
104
+ |---|---|
105
+ | Epochs / steps | 3 / 219 |
106
+ | Tokens seen | 1.15M |
107
+ | Batch | 4 × grad-accum 4 (effective 16) |
108
+ | Optimizer | AdamW, lr 2e-5, cosine to 10%, 20-step warmup, wd 0.1 |
109
+ | Precision | bfloat16 |
110
+ | Hardware / time | RTX 4090D 24GB — **6.9 min**, ~2950 tok/s, 19.6GB peak |
111
+ | Val loss / ppl | 2.534 → 1.646 → 1.539 → **1.537 / 4.65** (converged after ~2 epochs) |
112
+
113
+ To make the SSA-only modules trainable, the two output projections and the shared gate are
114
+ re-initialised to a small non-zero scale (0.01) before training — starting from the exact
115
+ donor checkpoint would freeze all 2.75M of them at zero gradient.
116
+
117
+ ## 4. Evaluation
118
+
119
+ ### 4.1 Held-out transfer loss (129 unseen sessions, 20,520 assistant tokens)
120
+
121
+ Both models evaluated with the identical script, mask and batching.
122
+
123
+ | Metric | Qwen3-0.6B-Base | **BaiHu-V1-Flash** | Change |
124
+ |---|---|---|---|
125
+ | loss | 2.5339 | **1.5376** | −39.3% |
126
+ | **ppl** | **12.603** | **4.654** | **−63.1%** |
127
+ | ppl (en) | 14.836 | 5.118 | −65.5% |
128
+ | ppl (zh) | 10.542 | 4.187 | −60.3% |
129
+
130
+ Per transfer type (ppl):
131
+
132
+ | Kind | Base | BaiHu-V1-Flash |
133
+ |---|---|---|
134
+ | rule_induction | 7.27 | 2.69 |
135
+ | analogy_transfer | 23.69 | 9.08 |
136
+ | format_transfer | 18.58 | 5.72 |
137
+ | counterfactual | 9.13 | 3.67 |
138
+ | cross_domain | 16.40 | 6.51 |
139
+ | teaching_loop | 8.26 | 2.95 |
140
+
141
+ The improvement holds in **all 8 slices** (2 languages × 6 kinds + overall).
142
+
143
+ ### 4.2 Generation samples (12 prompts, one per language × kind, greedy)
144
+
145
+ - **Base**: **7/12 outputs collapse into repeated symbols** (`⚇⚇⚇`, `ацион`,
146
+ `.TRAILING`), the rest are off-topic or hallucinated — the base model has never seen this
147
+ dialogue format.
148
+ - **BaiHu-V1-Flash**: **0/12 degenerate**; it reuses the rule/format from earlier turns and
149
+ is correct on most samples (`B12A`, `100, 50, 25, 12.5, …`, format rewrites). Arithmetic
150
+ errors remain — a 0.6B model with 1.3K training dialogues still slips on multi-step math.
151
+
152
+ ### 4.3 Effect of the P0 revision
153
+
154
+ Measured with the same data, recipe and seed, P0 vs the old single-vector shared branch:
155
+
156
+ | | old shared branch | P0 (per-block) |
157
+ |---|---|---|
158
+ | Held-out loss | 1.5376 | 1.5376 |
159
+ | Shared gate after SFT | 0.01001 | 0.01001 |
160
+
161
+ Under this regime (short dialogues, full causal window in the local path) the shared branch
162
+ carries almost no load, so the revision is quality-neutral and the gate still does not grow.
163
+ P0's intended benefit is long-range: it should only show up once the local window is allowed
164
+ to shrink / contexts are far longer than the 192-token window. That experiment has **not**
165
+ been run yet — see Limitations.
166
+
167
+ ## 5. How to Run Inference
168
+
169
+ ### 5.1 This is a custom architecture
170
+
171
+ `model_type: baihu_ssa` is **not** in the Transformers registry; a plain
172
+ `AutoModelForCausalLM.from_pretrained(...)` fails. Register the config/model classes first
173
+ (three lines below). Also note the model **does not** use `GenerationMixin` — call its own
174
+ `generate()` (greedy or temperature/top-k; no beam search).
175
+
176
+ The implementation lives in the project repository (not on the Hub). **This checkpoint
177
+ requires the P0 revision** of `modeling_baihu_ssa.py` — an older copy loads without error but
178
+ computes a different shared branch.
179
+
180
+ ### 5.2 Minimal working example
181
+
182
+ ```python
183
+ import os
184
+ import sys
185
+
186
+ import torch
187
+ from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer
188
+
189
+ sys.path.insert(0, os.path.join("ssa_model", "src")) # path to the cloned repo's src/
190
+
191
+ from baihu_ssa.configuration_baihu_ssa import BaiHuSSAConfig
192
+ from baihu_ssa.model_baihu_ssa import BaiHuSSAForCausalLM
193
+
194
+ # ---- register the custom architecture (REQUIRED) ----
195
+ AutoConfig.register("baihu_ssa", BaiHuSSAConfig, exist_ok=True)
196
+ AutoModelForCausalLM.register(BaiHuSSAConfig, BaiHuSSAForCausalLM, exist_ok=True)
197
+
198
+ REPO = "ZichenAI/BaiHu-V1-Flash"
199
+ device = "cuda" if torch.cuda.is_available() else "cpu"
200
+ dtype = torch.bfloat16 if device == "cuda" else torch.float32
201
+
202
+ tok = AutoTokenizer.from_pretrained(REPO)
203
+ model = AutoModelForCausalLM.from_pretrained(REPO, dtype=dtype).to(device).eval()
204
+
205
+ messages = [
206
+ {"role": "user", "content": "这个数列的规律是什么?1、4、9、16、25"},
207
+ {"role": "assistant", "content": "相邻两项的差是 3、5、7、9,每次加 2,所以第 n 项是 n 的平方。"},
208
+ {"role": "user", "content": "按同样的规律,下一项是多少?36、49 之后呢?"},
209
+ ]
210
+ ids = tok.apply_chat_template(messages, add_generation_prompt=True, tokenize=True,
211
+ enable_thinking=False)
212
+ ids = ids["input_ids"] if not isinstance(ids, list) else ids # transformers 5.x returns a dict here
213
+ with torch.no_grad():
214
+ out = model.generate(torch.tensor([ids], device=device), max_new_tokens=64,
215
+ do_sample=False, eos_token_id=tok.convert_tokens_to_ids("<|im_end|>"))
216
+ text = tok.decode(out[0, len(ids):], skip_special_tokens=True)
217
+ if "<think>" in text: # non-thinking mode emits an empty think block
218
+ text = text.split("</think>")[-1]
219
+ print(text.strip())
220
+ ```
221
+
222
+ Training used `add_generation_prompt=True, enable_thinking=False` (the generation prompt then
223
+ ends with `<|im_start|>assistant\n<think>\n\n</think>\n\n`). Keep that format for best results.
224
+
225
+ ### 5.3 Troubleshooting
226
+
227
+ | Error | Cause | Fix |
228
+ |---|---|---|
229
+ | `does not recognize this architecture` | custom architecture not registered | `AutoConfig.register` + `AutoModelForCausalLM.register` as in §5.2 |
230
+ | output is repetitive garbage | wrong chat format (e.g. a raw prompt with no ChatML wrapper) | render with the tokenizer's chat template, `enable_thinking=False` |
231
+ | `CUDA error: no kernel image is available` | PyTorch without kernels for an old GPU (Maxwell, sm_52) | pin torch 2.7.1+cu126 (2.8 dropped sm_50/sm_60) |
232
+
233
+ ## 6. Known Limitations
234
+
235
+ 1. **0.6B scale.** Multi-step arithmetic is still unreliable; the SFT teaches the *behaviour*
236
+ (reuse the earlier rule/format, answer directly), not new reasoning ability.
237
+ 2. **The shared branch is still barely used** (gate ≈ 0.0101 after SFT, same as the old
238
+ revision). Whether P0 pays off can only be decided in a long-context / window-shrink regime.
239
+ 3. **Trained on synthetic dialogues.** The 1,294 conversations are model-generated; style and
240
+ coverage are limited to the six transfer kinds and the topics they cover. Evaluation above
241
+ is on a held-out slice of the same distribution — not a general capability claim.
242
+ 4. **Custom architecture, no llama.cpp/GGUF support** (SSA attention is not implemented there).
243
+ 5. **Not an instruct model at large.** It is a 0.6B research model; expect terse answers and
244
+ occasional arithmetic slips.
245
+
246
+ ## 7. License
247
+
248
+ Custom license (`LICENSE.custom.md` in this repository):
249
+
250
+ - **Personal / non-commercial use: free**, including running, modifying and publicly
251
+ distributing derivative models under the same license.
252
+ - **Commercial use requires a paid license** — contact **novaweb6868@outlook.com**.
253
+ - Metadata: `license: other`, tag `commercial-license-required`.
254
+
255
+ ## 8. Citation
256
+
257
+ ```bibtex
258
+ @misc{baihu_v1_flash,
259
+ title = {BaiHu-V1-Flash: an SSA (sparse-attention + SubQ) retrofit of Qwen3-0.6B-Base, fine-tuned on bilingual transfer dialogues},
260
+ author = {ZichenAI},
261
+ year = {2026},
262
+ url = {https://huggingface.co/ZichenAI/BaiHu-V1-Flash}
263
+ }
264
+ ```
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
config.json CHANGED
@@ -1,79 +1,79 @@
1
- {
2
- "architectures": [
3
- "BaiHuSSAForCausalLM"
4
- ],
5
- "attention_bias": false,
6
- "attention_dropout": 0.0,
7
- "bos_token_id": 151643,
8
- "dtype": "bfloat16",
9
- "eos_token_id": 151643,
10
- "head_dim": 128,
11
- "hidden_act": "silu",
12
- "hidden_size": 1024,
13
- "initializer_range": 0.02,
14
- "intermediate_size": 3072,
15
- "layer_types": [
16
- "full_attention",
17
- "full_attention",
18
- "full_attention",
19
- "full_attention",
20
- "full_attention",
21
- "full_attention",
22
- "full_attention",
23
- "full_attention",
24
- "full_attention",
25
- "full_attention",
26
- "full_attention",
27
- "full_attention",
28
- "full_attention",
29
- "full_attention",
30
- "full_attention",
31
- "full_attention",
32
- "full_attention",
33
- "full_attention",
34
- "full_attention",
35
- "full_attention",
36
- "full_attention",
37
- "full_attention",
38
- "full_attention",
39
- "full_attention",
40
- "full_attention",
41
- "full_attention",
42
- "full_attention",
43
- "full_attention"
44
- ],
45
- "max_position_embeddings": 32768,
46
- "max_window_layers": 28,
47
- "model_type": "baihu_ssa",
48
- "num_attention_heads": 16,
49
- "num_hidden_layers": 28,
50
- "num_key_value_heads": 8,
51
- "pad_token_id": null,
52
- "rms_norm_eps": 1e-06,
53
- "rope_parameters": {
54
- "rope_theta": 1000000,
55
- "rope_type": "default"
56
- },
57
- "sliding_window": null,
58
- "ssa_block_size": 64,
59
- "ssa_component_init": 0.01,
60
- "ssa_component_seed": 1234,
61
- "ssa_compress_dim": 128,
62
- "ssa_donor_model": "Qwen3-0.6B-Base",
63
- "ssa_donor_revision": null,
64
- "ssa_force_full_window": false,
65
- "ssa_local_blocks": 2,
66
- "ssa_num_subq_heads": 4,
67
- "ssa_router_bias_scale": 0.1,
68
- "ssa_router_dim": 128,
69
- "ssa_router_init_scale": 0.01,
70
- "ssa_shared_init_scale": 0.01,
71
- "ssa_shared_kv": true,
72
- "ssa_top_k": 8,
73
- "ssa_use_router_key": true,
74
- "tie_word_embeddings": true,
75
- "transformers_version": "5.17.0",
76
- "use_cache": true,
77
- "use_sliding_window": false,
78
- "vocab_size": 151936
79
- }
 
1
+ {
2
+ "architectures": [
3
+ "BaiHuSSAForCausalLM"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "bos_token_id": 151643,
8
+ "dtype": "bfloat16",
9
+ "eos_token_id": 151643,
10
+ "head_dim": 128,
11
+ "hidden_act": "silu",
12
+ "hidden_size": 1024,
13
+ "initializer_range": 0.02,
14
+ "intermediate_size": 3072,
15
+ "layer_types": [
16
+ "full_attention",
17
+ "full_attention",
18
+ "full_attention",
19
+ "full_attention",
20
+ "full_attention",
21
+ "full_attention",
22
+ "full_attention",
23
+ "full_attention",
24
+ "full_attention",
25
+ "full_attention",
26
+ "full_attention",
27
+ "full_attention",
28
+ "full_attention",
29
+ "full_attention",
30
+ "full_attention",
31
+ "full_attention",
32
+ "full_attention",
33
+ "full_attention",
34
+ "full_attention",
35
+ "full_attention",
36
+ "full_attention",
37
+ "full_attention",
38
+ "full_attention",
39
+ "full_attention",
40
+ "full_attention",
41
+ "full_attention",
42
+ "full_attention",
43
+ "full_attention"
44
+ ],
45
+ "max_position_embeddings": 32768,
46
+ "max_window_layers": 28,
47
+ "model_type": "baihu_ssa",
48
+ "num_attention_heads": 16,
49
+ "num_hidden_layers": 28,
50
+ "num_key_value_heads": 8,
51
+ "pad_token_id": null,
52
+ "rms_norm_eps": 1e-06,
53
+ "rope_parameters": {
54
+ "rope_theta": 1000000,
55
+ "rope_type": "default"
56
+ },
57
+ "sliding_window": null,
58
+ "ssa_block_size": 64,
59
+ "ssa_component_init": 0.01,
60
+ "ssa_component_seed": 1234,
61
+ "ssa_compress_dim": 128,
62
+ "ssa_donor_model": "Qwen3-0.6B-Base",
63
+ "ssa_donor_revision": null,
64
+ "ssa_force_full_window": true,
65
+ "ssa_local_blocks": 2,
66
+ "ssa_num_subq_heads": 4,
67
+ "ssa_router_bias_scale": 0.1,
68
+ "ssa_router_dim": 128,
69
+ "ssa_router_init_scale": 0.01,
70
+ "ssa_shared_init_scale": 0.01,
71
+ "ssa_shared_kv": true,
72
+ "ssa_top_k": 8,
73
+ "ssa_use_router_key": true,
74
+ "tie_word_embeddings": true,
75
+ "transformers_version": "5.17.0",
76
+ "use_cache": true,
77
+ "use_sliding_window": false,
78
+ "vocab_size": 151936
79
+ }
model.safetensors CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:40f5a9f413e8ce888f0e44e8bf3c18acd9d711cf2d0c76482096d43ec473c4c6
3
  size 2395267832
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:e869217a7b844cf394980908f70312b97634fcd8f3fac5f264967e65da323fae
3
  size 2395267832