File size: 18,107 Bytes
b418a7a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d57da8f
b418a7a
d57da8f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b418a7a
 
 
 
 
 
 
d57da8f
b418a7a
 
 
 
 
 
 
 
d57da8f
b418a7a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d57da8f
b418a7a
 
 
 
 
 
 
 
 
 
 
 
d57da8f
b418a7a
 
 
 
d57da8f
b418a7a
 
d57da8f
b418a7a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d57da8f
b418a7a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d57da8f
b418a7a
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
---

library_name: transformers
pipeline_tag: text-generation
language:
  - zh
  - en
license: other
license_name: baihu-custom-license
license_link: https://huggingface.co/NovaAI6868/BaiHu-V1-Flash/blob/main/LICENSE.custom.md
base_model: Qwen/Qwen3-0.6B-Base
tags:
  - sparse-attention
  - subq
  - ssa
  - long-context
  - commercial-license-required
  - text-generation
---


# BaiHu-V1-Flash

**BaiHu-V1-Flash** is a retrofit of the dense-attention model `Qwen/Qwen3-0.6B-Base` into an
**SSA (Sparse-attention + SubQ)** architecture, obtained by continued pretraining.

- Base model: `Qwen/Qwen3-0.6B-Base` (28 layers / 16 Q heads / 8 KV heads / head_dim 128 / 32K context / tied embeddings)

- Parameters: 598.8M

- Training data: mixed Chinese + English (Fineweb-Edu-Chinese-V2.1 + fineweb-edu, 50/50)

- License: **free for personal use; a paid license is required for commercial use** (see "License" below)



---



## 1. Architecture: SSA (three paths, each with its own softmax, then summed)



Every layer keeps the base model's MLP / RMSNorm weights and replaces full attention with an

SSA layer built from three parallel paths:



| Path | Role | Complexity |

|---|---|---|

| `shared` | every query sees all **completed** blocks through one compressed vector per block | `O(T·T/B)` |

| `local` | dense causal attention over the most recent window | `O(T·w)` |

| `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)` |



### Hyperparameters



| Parameter | Value | Meaning |

|---|---|---|

| `ssa_block_size` | 64 | block size B |

| `ssa_top_k` | 8 | number of blocks selected by the sparse path |

| `ssa_local_blocks` | 2 | local window = 3 × 64 = 192 tokens |

| `ssa_num_subq_heads` | 4 | SubQ heads, r = 16 / 4 = 4 |
| `ssa_router_dim` | 128 | router subspace dimension |
| `ssa_compress_dim` | 128 | block compression dimension |

Only **2.75M** parameters are new (≈0.46% of the model); all other weights are inherited
from the base model.

---

## 2. Training

| Item | Setting |
|---|---|
| Starting point | a conversion checkpoint that is **bit-exact** with the base model (`max\|Δlogit\| = 0.000e+00`) |
| Tokens seen | 5.0M (≈0.25 epoch of the corpus) |
| Sequence length | 512 (must be a multiple of `ssa_block_size = 64`) |
| Effective batch | 8192 tokens (batch 2 × grad_accum 8) |

| Precision | fp32 |

| Optimizer | SGD with momentum 0.9 |

| Learning rate | 5e-4, 50-step warmup, cosine decay to 10% |

| Hardware | single NVIDIA GTX TITAN X (Maxwell, sm_52, 12.9 GB) |
| Throughput | ≈274 tokens/s |

### Critical issues found and fixed during this retrofit

Several defects silently break training and are worth documenting:

1. **Both new branches had identically zero gradients (blocking).** To make the converted
   model bit-exact with the base model, `compress_out` and `router_out` were initialized to
   exactly zero, and the branches were skipped entirely by a gate. The branch output was
   therefore always zero, so the back-propagated gradient was also always zero: all 2.75M
   SSA parameters **stayed frozen for the entire run** and the sparse attention was dead
   code. The fix is a small non-zero initialization.
2. **Routing was non-differentiable.** `top-k` produces hard indices, and indexing is not
   differentiable. If the routing scores are used only to decide *which* blocks to read and
   never enter the softmax, the gradients of `router_q` / `router_k` are **exactly zero** —
   the router can never learn to route. The fix is to feed the selected blocks' scores,
   squashed through `tanh` and gently scaled, into the attention logits as an additive bias.
3. **The shared summary was a sum, not a mean.** Its magnitude grew linearly with the
   prefix, and because that branch is injected at full weight (a single-element softmax has
   probability exactly 1), it swamped the residual stream: hidden states grew from 0.2 to
   about 7 in layer 0 and to about 1900 by layer 27, and validation loss went 3.56 → 11.38.
   The fix is to divide by the token count.
4. **The shared branch needs an explicit gate.** With a single-element softmax the
   probability is always 1, so the initialization scale of `compress_out` cannot control the
   injection strength at all (measured: scales from 1e-4 to 0.03 all left the loss at
   exactly 7.3526). A learnable scalar gate, initialized to a small positive value, lets the
   optimizer decide how far to open it.

---

## 3. How to Run Inference

### 3.1 Important: this is a custom architecture

`BaiHu-V1-Flash` uses `model_type: baihu_ssa`, which is **not** in the Transformers
registry. Loading it with a plain `AutoModelForCausalLM.from_pretrained(...)` fails with:

```

ValueError: The checkpoint you are trying to load has model type `baihu_ssa`

but Transformers does not recognize this architecture.

```

You must register the config and model classes first. This is a one-time, three-line step
(see below).

Also note: **this model does not go through `transformers.GenerationMixin`.** SSA owns its
own KV-cache layout (`BaiHuSSACache`) because the cache stores per-block compressed prefix
sums rather than a plain growing key/value tensor. Call the model's own `generate()` method;
beam search is not supported.

### 3.2 Setup

```bash

# The SSA implementation lives in the project repository (not on the Hub),

# because it is a custom architecture.

git clone <this-project-repo> ssa_model

cd ssa_model

uv venv --python 3.11 .venv



# Ampere or newer (RTX 30xx/40xx, A100, ...): any recent torch works.

uv pip install --python .venv/Scripts/python.exe \

    "numpy>=1.26" "transformers>=4.51" "safetensors>=0.4" "torch>=2.6"



# Maxwell / Pascal / Volta (GTX 9xx/10xx, TITAN X, V100, ...): pin torch 2.7.1 —

# see the GPU note below for why `torch>=2.6` would resolve to a build without your kernels.

uv pip install --python .venv/Scripts/python.exe \

    "numpy>=1.26" "transformers>=4.51" "safetensors>=0.4" "torch==2.7.1"

```

> **GPU note (Maxwell and older).** If you are on an NVIDIA Maxwell card such as the
> GTX TITAN X (compute capability sm_52), recent PyTorch wheels no longer ship kernels for

> it: PyTorch **2.8 removed sm_50/sm_60**, and 2.8.0+cu126 only contains `sm_61…sm_90`.

> You will get `CUDA error: no kernel image is available for execution on the device`.

> Use **torch 2.7.1+cu126**, which still contains `sm_50`. Also train/infer in **fp32** on
> Maxwell: measured fp32 5.30 / fp16 4.46 / bf16 3.20 TFLOPS, so half precision is a loss.

### 3.3 Minimal working example

```python

import os

import sys



import torch

from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer



# Make the custom implementation importable (path to the cloned repo's src/)

sys.path.insert(0, os.path.join("ssa_model", "src"))



from baihu_ssa.configuration_baihu_ssa import BaiHuSSAConfig

from baihu_ssa.model_baihu_ssa import BaiHuSSAForCausalLM



# ---- register the custom architecture (REQUIRED) ----

AutoConfig.register("baihu_ssa", BaiHuSSAConfig, exist_ok=True)

AutoModelForCausalLM.register(BaiHuSSAConfig, BaiHuSSAForCausalLM, exist_ok=True)



REPO = "NovaAI6868/BaiHu-V1-Flash"

device = "cuda" if torch.cuda.is_available() else "cpu"

dtype = torch.float32 if device == "cpu" else torch.float32  # see GPU note: fp32



tokenizer = AutoTokenizer.from_pretrained(REPO)

model = AutoModelForCausalLM.from_pretrained(REPO, dtype=dtype).to(device).eval()



prompt = "人工智能的未来是"

input_ids = tokenizer(prompt, return_tensors="pt").input_ids.to(device)



with torch.no_grad():

    out = model.generate(

        input_ids,

        max_new_tokens=32,

        do_sample=False,        # greedy; supported

        # do_sample=True, temperature=0.8, top_k=50,   # sampling also supported

    )

print(tokenizer.decode(out[0], skip_special_tokens=True))

```

### 3.4 `generate()` parameters

The built-in `generate()` is a self-contained decoding loop (greedy or
temperature/top-k sampling). Supported arguments:

| Argument | Default | Meaning |
|---|---|---|
| `max_new_tokens` | 32 | number of tokens to generate |
| `do_sample` | `False` | `False` = greedy, `True` = sample |
| `temperature` | 1.0 | sampling temperature (used when `do_sample=True`) |
| `top_k` | `None` | top-k sampling cutoff (used when `do_sample=True`) |
| `eos_token_id` | `None` | stop early when all sequences emit this token |
| `use_cache` | `True` | keep the SSA cache between steps; **leave on**, decoding is much slower without it |

**Not supported:** beam search (the SSA cache has no `reorder_cache`), and
`transformers` generation utilities such as `logits_processor` / `stopping_criteria`.

### 3.5 Computing perplexity / loss

```python

with torch.no_grad():

    inputs = tokenizer("Your text here", return_tensors="pt").to(device)

    loss = model(**inputs, labels=inputs["input_ids"]).loss

    print("loss", loss.item(), "ppl", torch.exp(loss).item())

```

### 3.6 Throughput you should expect

Measured on a GTX TITAN X (sm_52, 12.9 GB), fp32, prompt 1024 tokens, 64 generated tokens:



| | Qwen3-0.6B-Base | BaiHu-V1-Flash |

|---|---|---|

| Prefill / TTFT | 503 ms | 1431 ms |

| Per-token decode (TPOT) | 36.6 ms | 80.5 ms |

| Decode throughput | 27.3 tok/s | 12.4 tok/s |

| Peak memory (generation) | 9.91 GB | 9.01 GB |



**BaiHu-V1-Flash is currently ~2.2× slower to decode than the base model**, even though the

sparse path reads fewer keys. The reason is implementation, not architecture: the attention

layer loops over query blocks in Python and launches many small kernels, so kernel-launch

overhead dominates the FLOPs saved. Memory use is lower; speed is not yet a win.

See "Known Limitations" below.



### 3.7 Expected output quality (honest warning)



This checkpoint has seen only **5.0M tokens**, so the new branches are far from converged.

Greedy decoding tends to **repeat itself** and it is **worse than the base model on

perplexity** (ppl 82.3 vs 39.4 at 512 tokens). A real greedy sample:



```

prompt: 人工智能的未来是

output: 人工智能的未来是怎样的?人工智能的未来是怎样的?人工智能的未来是怎样的?…

```



Use it to study or continue the SSA retrofit — **do not** expect it to match

`Qwen3-0.6B-Base` as a general-purpose model yet.



### 3.8 Troubleshooting



| Error | Cause | Fix |

|---|---|---|

| `does not recognize this architecture` / `KeyError: 'baihu_ssa'` | custom architecture not registered | call `AutoConfig.register` + `AutoModelForCausalLM.register` as in 3.3 |
| `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) |

| `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) |

| `BaiHuSSACache` errors when calling `model.generate(...)` from GenerationMixin | this model bypasses `GenerationMixin` | call `model.generate(...)` on the BaiHu model itself |



---



## 4. Evaluation



### 4.1 Language modeling perplexity (validation set, identical windows)



| Sequence length | Qwen3-0.6B-Base | BaiHu-V1-Flash |

|---|---|---|

| 512 | 3.6726 / ppl 39.354 | 4.4108 / ppl 82.335 |

| 1024 | 3.4691 / ppl 32.109 | 4.2575 / ppl 70.631 |

| 2048 | 3.0801 / ppl 21.761 | 3.9213 / ppl 50.467 |



### 4.2 Standard benchmarks (lm-evaluation-harness)



| Task | Qwen3-0.6B-Base | BaiHu-V1-Flash | Delta |

|---|---|---|---|

| arc_easy | 0.5550 | 0.6250 | +0.0700 |
| hellaswag | 0.5350 | 0.5200 | -0.0150 |
| piqa | 0.7050 | 0.6950 | -0.0100 |
| winogrande | 0.6300 | 0.6300 | +0.0000 |

### 4.3 Inference compute and resource usage

| Metric | Qwen3-0.6B-Base | BaiHu-V1-Flash |
|---|---|---|
| Parameters (M) | 596.0500 | 598.8000 |
| Prefill peak memory (GB) | 9.6200 | 8.0210 |
| Generation peak memory (GB) | 9.9120 | 9.0050 |
| Prefill latency (s) | 0.5030 | 1.4310 |
| TTFT (ms) | 502.7 | 1431.2 |
| TPOT (ms) | 36.6 | 80.5 |
| Decode throughput (tok/s) | 27.3300 | 12.4200 |
| Attention FLOPs/token (GFLOPs) | 0.1176 | 0.1057 |
| Attention key accesses vs full attention | 1.0000 | 0.8993 |
| GPU utilization mean/max (%) | 93.9 | 43.3 |
| Power mean/max (W) | 179.1 | 124.3 |

Positive findings: peak inference memory is lower (generation 9.005 vs 9.912 GB, −9.2%), and the
model draws less power because it is not compute-bound.

Negative findings, stated plainly:

- **Decode is 2.2× slower** (12.42 vs 27.33 tok/s) and **prefill is 2.8× slower**
  (TTFT 1431 vs 503 ms), despite the sparse path reading fewer keys. The current
  implementation loops over query blocks in Python and issues many small kernels, so
  launch overhead dominates the FLOPs saved. **The sparse attention does not yet pay off

  on this hardware.**
- **Attention key accesses are still 89.9% of full attention** at this sequence length.
  The reason is structural: the local window already covers 3 blocks (192 tokens) and the
  sparse path reads up to `top_k + 1 = 9` blocks from a grid that only has 16 blocks at
  1024 tokens, so the selected set is almost the whole grid. Sparsity only becomes a real
  saving once the sequence is long relative to `top_k × block_size` (i.e. well beyond
  10k tokens).
- **Language modeling perplexity is clearly worse than the base model** at every length
  tested (ppl 82.3 vs 39.4 at 512; 50.5 vs 21.8 at 2048). This is the honest cost of
  shrinking the dense local window from the full prefix to 192 tokens while the new
  long-range branches are still very weakly trained.
- **Standard benchmarks are roughly neutral but not better**: arc_easy improves

  (+0.070 acc_norm), winogrande is unchanged, while hellaswag (−0.015) and piqa (−0.010)
  regress slightly.

### 4.4 Sparsity

Two measurements are reported because they answer different questions and are **not**
interchangeable:

| Scenario | Average blocks selected per query | Key access ratio vs full attention |
|---|---|---|
| Chunked forward over a 512-token validation window | 0.38 | 0.0938 |
| Prefill of 1024 tokens (steady state) | up to `top_k` | 0.8993 |

The first number averages over *all* query blocks including the early ones, which have no
completed blocks available to select and therefore read nothing through the sparse path.
The second is the steady-state ratio for later queries, and it is the one that matters for
efficiency — see the note in section 4.3: at these sequence lengths the sparse path is not
yet saving meaningful work.

---

## 5. Known Limitations

1. **Trained for very little.** Only 5.0M tokens (≈0.25 epoch). The new branches have not
   converged; ppl is well above the base model and decode is slower (see 3.6 and 4.1). Continuing
   to 200M tokens or more is required before the SSA layers can genuinely take over
   long-range modeling.
2. **The shared branch does not appear to help and was actively suppressed by the

   optimizer.** Its learnable gate *decreased* over training (0.0100 → 0.0129 at step 100 →
   0.0122 at step 610) instead of growing, meaning the optimizer found the single
   prefix-mean summary not worth injecting. This is the single most important thing to
   change next: replace it with **one compressed vector per block** (same `O(T·T/B)` cost,
   far more information retained).
3. **Sparse attention is not yet a net win on this hardware.** It reads fewer keys but runs
   2.2× slower because the implementation loops over query blocks in Python and issues many
   small kernels. It needs kernel-level batching (or a fused implementation) before the
   sparsity can translate into speed.
4. **Sparsity only pays off at long sequences.** With `top_k=8` and `block_size=64`, the
   sparse path can read up to 9 blocks = 576 tokens; at 1024 tokens the grid only has 16
   blocks, so the selected set covers most of the context and the local window already
   covers the rest. Real savings require sequences well beyond 10k tokens.
5. **Routing quality is not fully validated.** The distribution of selected top-k blocks
   should be checked for degeneration (e.g. always selecting the same blocks). The
   non-differentiable-routing bug that would have made this *impossible* to learn has been
   fixed (see section 2), but the learned policy has not been analyzed in detail.
6. **Block size B = 64 was not ablated.** Limited by memory and the Windows WDDM watchdog on
   this machine; B ∈ {32, 64, 128} should be swept on a larger GPU.
7. **Document-boundary packing.** The current packing strategy places multiple documents in
   one sequence (separated by EOS).

---

## 6. License

**Free for personal use; a paid license is required for commercial use.**

- ✅ Personal study, research, teaching, hobby projects: **free**, no application required
- ✅ Academic research with public publication: **free** (please cite the source)
- 💰 Internal company/studio use, paid API/SaaS, product integration, client deliverables:
  **commercial license required**

**Commercial licensing contact: novaweb6868@outlook.com**

Full terms: [LICENSE.custom.md](./LICENSE.custom.md).

This model is an architectural retrofit of `Qwen/Qwen3-0.6B-Base` (Apache License 2.0).
This license governs only the newly added portions and does not alter the upstream
component's original license.

---

## 7. Citation

```bibtex

@misc{baihu-v1-flash,

  title  = {BaiHu-V1-Flash: An SSA (Sparse-attention + SubQ) Retrofit of Qwen3-0.6B-Base},

  author = {NovaAI6868},

  year   = {2026},

  url    = {https://huggingface.co/NovaAI6868/BaiHu-V1-Flash}

}

```