Text Generation
Transformers
Safetensors
Chinese
English
baihu_ssa
sparse-attention
subq
ssa
long-context
commercial-license-required
conversational
Instructions to use ZichenAI/BaiHu-V1-Flash with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use ZichenAI/BaiHu-V1-Flash with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="ZichenAI/BaiHu-V1-Flash") messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("ZichenAI/BaiHu-V1-Flash", device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use ZichenAI/BaiHu-V1-Flash with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "ZichenAI/BaiHu-V1-Flash" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "ZichenAI/BaiHu-V1-Flash", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/ZichenAI/BaiHu-V1-Flash
- SGLang
How to use ZichenAI/BaiHu-V1-Flash with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "ZichenAI/BaiHu-V1-Flash" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "ZichenAI/BaiHu-V1-Flash", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "ZichenAI/BaiHu-V1-Flash" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "ZichenAI/BaiHu-V1-Flash", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use ZichenAI/BaiHu-V1-Flash with Docker Model Runner:
docker model run hf.co/ZichenAI/BaiHu-V1-Flash
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}
}
```
|