---
library_name: transformers
pipeline_tag: text-generation
language:
- zh
- en
license: other
license_name: baihu-custom-license
license_link: https://huggingface.co/ZichenAI/BaiHu-V1-Flash/blob/main/LICENSE.custom.md
base_model: Qwen/Qwen3-0.6B-Base
base_model_relation: finetune
tags:
- sparse-attention
- subq
- ssa
- long-context
- supervised-fine-tuning
- transfer-learning
- commercial-license-required
- text-generation
---
# BaiHu-V1-Flash
**BaiHu-V1-Flash** is an SSA (Sparse-attention + SubQ) retrofit of `Qwen/Qwen3-0.6B-Base`,
fine-tuned on **1,294 bilingual multi-turn dialogues** whose purpose is **generalizing from a
single worked example** — the model is shown one worked rule / format / method in the first
turn, and later turns ask it to reuse that rule on a *new* case it has never seen.
- 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** (2.75M are SSA-only modules)
- Architecture: SSA — every layer runs three attention paths (shared / local / sparse-SubQ), each with its own softmax, then summed
- Training data: 1,294 synthetic dialogues (English 693 / Chinese 601), six transfer types
- License: **free for personal use; a paid license is required for commercial use** (see [License](#license))
> **Revision note.** This repository previously hosted the base-pretrain checkpoint of the
> same name (5.0M tokens of continued pretraining, no instruction tuning, no P0 revision).
> It has been **replaced** by the checkpoint described here. The two are not interchangeable:
> this one uses the revised shared branch (§2) and is a supervised fine-tune.
---
## 1. Revisions in this release
| | earlier release (now replaced) | **this release** |
|---|---|---|
| Shared path | one mean vector for the **entire prefix** | **one compressed vector per completed block** (P0 revision) |
| Training | continued pretraining, 5.0M tokens of web text | supervised fine-tune, 1,294 transfer dialogues (assistant-token loss) |
| Behaviour | base LM | follows a rule / format established earlier in the dialogue |
Both revisions share the same SSA architecture, base weights and license.
## 2. Architecture
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 attends over **one compressed summary vector per completed 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 head group; real attention is computed only for the selected top-k blocks | `O(T·k·B)` |
### The P0 revision (this release)
The previous revision compressed the **whole prefix into a single mean vector**. A query
could only see one undifferentiated global average, the single-element softmax made that
average inject at full weight, and during 5.0M tokens of continued pretraining the learned
gate **shrank instead of growing** (0.0100 → 0.0129 → 0.0122) — i.e. the optimizer actively
suppressed the branch.
This release implements the design the project always documented: **one compressed key/value
per completed 64-token block**, so queries do a genuine softmax over blocks and "which block
matters" becomes learnable. Cost is unchanged (`O(T·T/B)`; block summaries are still computed
once per layer).
### Hyperparameters
| Parameter | Value | Meaning |
|---|---|---|
| `ssa_block_size` | 64 | block size B |
| `ssa_top_k` | 8 | 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` / `ssa_compress_dim` | 128 / 128 | router subspace / summary width |
## 3. Training
**Data.** 1,294 synthetic multi-turn dialogues (en 693 / zh 601), each 6–12 messages: the
first user turn gives a worked case or a rule, a later user turn introduces a new instance
that can only be handled by reusing it, and the last turn pushes generalization one step
further. Six `transfer_kind`s are covered: `rule_induction`, `analogy_transfer`,
`format_transfer`, `counterfactual`, `cross_domain`, `teaching_loop`.
Split (stratified by language × kind, seed 0): **train 1,165 sessions / 345.7K tokens**
(assistant 188.8K), **val 129 sessions / 38.2K tokens** (assistant 20.5K).
**Recipe.** Loss is computed on assistant tokens only; sequences are padded to a multiple of
the SSA block size (64); rendering uses the tokenizer's own chat template
(`apply_chat_template`, `add_generation_prompt=False`, i.e. the final assistant turn carries
Qwen3's empty `` block).
| Item | Value |
|---|---|
| Epochs / steps | 3 / 219 |
| Tokens seen | 1.15M |
| Batch | 4 × grad-accum 4 (effective 16) |
| Optimizer | AdamW, lr 2e-5, cosine to 10%, 20-step warmup, wd 0.1 |
| Precision | bfloat16 |
| Hardware / time | RTX 4090D 24GB — **6.9 min**, ~2950 tok/s, 19.6GB peak |
| Val loss / ppl | 2.534 → 1.646 → 1.539 → **1.537 / 4.65** (converged after ~2 epochs) |
To make the SSA-only modules trainable, the two output projections and the shared gate are
re-initialised to a small non-zero scale (0.01) before training — starting from the exact
donor checkpoint would freeze all 2.75M of them at zero gradient.
## 4. Evaluation
### 4.1 Held-out transfer loss (129 unseen sessions, 20,520 assistant tokens)
Both models evaluated with the identical script, mask and batching.
| Metric | Qwen3-0.6B-Base | **BaiHu-V1-Flash** | Change |
|---|---|---|---|
| loss | 2.5339 | **1.5376** | −39.3% |
| **ppl** | **12.603** | **4.654** | **−63.1%** |
| ppl (en) | 14.836 | 5.118 | −65.5% |
| ppl (zh) | 10.542 | 4.187 | −60.3% |
Per transfer type (ppl):
| Kind | Base | BaiHu-V1-Flash |
|---|---|---|
| rule_induction | 7.27 | 2.69 |
| analogy_transfer | 23.69 | 9.08 |
| format_transfer | 18.58 | 5.72 |
| counterfactual | 9.13 | 3.67 |
| cross_domain | 16.40 | 6.51 |
| teaching_loop | 8.26 | 2.95 |
The improvement holds in **all 8 slices** (2 languages × 6 kinds + overall).
### 4.2 Generation samples (12 prompts, one per language × kind, greedy)
- **Base**: **7/12 outputs collapse into repeated symbols** (`⚇⚇⚇`, `ацион`,
`.TRAILING`), the rest are off-topic or hallucinated — the base model has never seen this
dialogue format.
- **BaiHu-V1-Flash**: **0/12 degenerate**; it reuses the rule/format from earlier turns and
is correct on most samples (`B12A`, `100, 50, 25, 12.5, …`, format rewrites). Arithmetic
errors remain — a 0.6B model with 1.3K training dialogues still slips on multi-step math.
### 4.3 Standard benchmarks (lm-evaluation-harness)
Both checkpoints were scored with `lm-evaluation-harness` **0.4.13** in one environment, with
the same prompts, batch size and dtype (`bfloat16`), over the **entire evaluation split** of
every task — 2,376 arc_easy / 1,172 arc_challenge / 10,042 hellaswag / 1,838 piqa / 1,267
winogrande examples, with no `--limit` subsampling — so the two columns are like-for-like. The
metrics are the harness's 0-shot numbers, scored in raw-completion mode (no chat template),
which is what a base checkpoint supports.
| Task | Metric | Qwen3-0.6B-Base | **BaiHu-V1-Flash** | Δ |
|---|---|---|---|---|
| arc_easy | acc_norm | 0.5791 ±0.0101 | **0.5939 ±0.0101** | +0.0148 |
| arc_challenge | acc_norm | 0.3848 ±0.0142 | **0.3882 ±0.0142** | +0.0034 |
| hellaswag | acc_norm | 0.5385 ±0.0050 | **0.5507 ±0.0050** | +0.0122 |
| piqa | acc_norm | 0.6997 ±0.0107 | **0.7084 ±0.0106** | +0.0087 |
| winogrande | acc | 0.5856 ±0.0138 | **0.6062 ±0.0137** | +0.0206 |
| unweighted mean | | 0.5575 | **0.5695** | +0.0119 |
**How to read this.** All five deltas are positive, but each is between 0.3σ and 1.5σ of its
own standard error, so the defensible claim is **"no regression in general capability"**, not
"the fine-tune made the model smarter". The movement is also not an SSA effect: a control run
of the *dense* `Qwen3-0.6B-Base` fine-tuned with the identical data and recipe lands within
±0.005 of BaiHu-V1-Flash on every one of these tasks (arc_easy 0.5951, arc_challenge 0.3908,
hellaswag 0.5511, piqa 0.7111, winogrande 0.6014). What this release did buy is the transfer
behaviour of §4.1–§4.2 — a 63% perplexity drop on held-out dialogues and no degenerate
generations.
Caveats, stated plainly:
- These are **English** benchmarks. This harness build ships no C-Eval / CMMLU / C3 tasks, so
Chinese capability is not measured above; the held-out split of §4.1 (which contains both
languages) is the only Chinese-side evidence in this card.
- The Hub checkpoint stores float32 weights and the harness ran it at `bfloat16` (see §5.2).
The two agree to within 0.002 on every task listed, so precision is not driving the table.
### 4.4 Inference cost and speed (RTX 4090D, bfloat16)
Measured with the project's own `scripts/bench_resources.py` on the card that trained the model
(RTX 4090D 24 GB, no other process on it), bfloat16, `torch.no_grad()`, greedy decoding of 64
tokens after a prefill of 1,024 / 2,048 / 4,096 tokens, identical script for both models.
**The checkpoint exactly as shipped** (`ssa_force_full_window: true` — see below), 1,024-token
prefill:
| Metric | Qwen3-0.6B-Base | BaiHu-V1-Flash |
|---|---|---|
| Parameters (M) | 596.0 | 598.8 |
| Peak memory, prefill (GB) | 1.63 | **1.52** |
| Peak memory, generate (GB) | 1.78 | 2.01 |
| Prefill latency = TTFT (s) | 0.025 | 1.48 |
| TPOT — time per output token (ms) | 17.3 | 68.4 |
| Decode throughput (tok/s) | 57.8 | 14.6 |
| Attention FLOPs per token (GFLOPs) | 0.118 | 0.202 |
| Attention keys read per query (vs full attention) | 1.00× | 1.72× |
| GPU utilization mean (%) / power mean (W) | 15.5 / 69.8 | 20.0 / 69.0 |
How both quantities scale with context (same run, every cell in the order *base → ours*):
| Prefill | TTFT (ms) | TPOT (ms) | Decode (tok/s) | Keys read per query | Generate peak (GB) |
|---|---|---|---|---|---|
| 1,024 | 25 → 1,477 | 17.3 → 68.4 | 57.8 → 14.6 | 1.00× → 1.72× | 1.78 → 2.01 |
| 2,048 | 39 → 2,113 | 19.8 → 78.7 | 50.4 → 12.7 | 1.00× → 1.45× | 2.35 → 2.82 |
| 4,096 | 72 → 3,610 | 21.4 → 108.8 | 46.7 → 9.2 | 1.00× → 1.24× | 3.53 → 4.43 |
Stated plainly:
- **There is no speed advantage at any length tested.** Prefill (TTFT) is 50–59× slower in
wall-clock — 1.48 s vs 25 ms at 1 K — decode is 3.9× slower at 1 K and 5.1× at 4 K, and peak
memory is comparable (slightly *lower* for this model during prefill, slightly higher during
generation).
- **The sparse path is not actually saving anything here.** This checkpoint is trained *and
released* with `ssa_force_full_window: true`: the dense local window always covers the entire
causal prefix, so the shared and sparse paths are **additive on top of full attention** instead
of replacing part of it. Hence keys read per query above 1.00×.
- **Read the 1.72× honestly.** It is the measured cost of the shipped configuration, and it is
also why the model spends more attention FLOPs per token than the dense base (0.202 vs 0.118
GFLOPs at 1 K) while still being far slower in wall-clock. Sparsity only starts to bite once
`top_k × block_size` is small relative to the context (see the next point).
- **Turning the window shrinkage on is a flag, not a retrain** — but it was never validated in
that mode, because the weights were trained with the window forced open. For reference, with
`ssa_force_full_window=false` the same weights read **0.90× / 0.54× / 0.29×** as many keys at
1 K / 2 K / 4 K and decode at 20.7 / 18.5 / 14.8 tok/s. Even then it stays **2.5–3.2× slower**
than the dense base, and quality in that mode is **unevaluated** — treat the numbers as the
architecture's ceiling on this implementation, not as a free win.
- What follows from this is an implementation item, not an architecture one: the per-block Python
loop in the SSA layer issues many small kernels, so launch overhead dominates the arithmetic
saved (GPU utilization never exceeds ~28%). The sparse path has to be fused before any of this
can pay off on real hardware.
### 4.5 Effect of the P0 revision
Measured with the same data, recipe and seed, P0 vs the old single-vector shared branch:
| | old shared branch | P0 (per-block) |
|---|---|---|
| Held-out loss | 1.5376 | 1.5376 |
| Shared gate after SFT | 0.01001 | 0.01001 |
Under this regime (short dialogues, full causal window in the local path) the shared branch
carries almost no load, so the revision is quality-neutral and the gate still does not grow.
P0's intended benefit is long-range: it should only show up once the local window is allowed
to shrink / contexts are far longer than the 192-token window. That experiment has **not**
been run yet — see Limitations.
## 5. How to Run Inference
### 5.1 This is a custom architecture
`model_type: baihu_ssa` is **not** in the Transformers registry; a plain
`AutoModelForCausalLM.from_pretrained(...)` fails. Register the config/model classes first
(three lines below). Also note the model **does not** use `GenerationMixin` — call its own
`generate()` (greedy or temperature/top-k; no beam search).
The implementation lives in the project repository (not on the Hub). **This checkpoint
requires two revisions** of that code: the P0 revision of `modeling_baihu_ssa.py` (an older
copy loads without error but computes a different shared branch), and the dtype-aware loader
in `model_baihu_ssa.py` — before it, `from_pretrained(..., dtype=...)` was silently ignored
and the model always came back float32.
### 5.2 Minimal working example
```python
import os
import sys
import torch
from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer
sys.path.insert(0, os.path.join("ssa_model", "src")) # path to the cloned repo's 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 = "ZichenAI/BaiHu-V1-Flash"
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.bfloat16 if device == "cuda" else torch.float32
tok = AutoTokenizer.from_pretrained(REPO)
model = AutoModelForCausalLM.from_pretrained(REPO, dtype=dtype).to(device).eval()
messages = [
{"role": "user", "content": "What is the rule behind this sequence? 1, 4, 9, 16, 25"},
{"role": "assistant", "content": "The gaps between consecutive terms are 3, 5, 7, 9 "
"(+2 each time), so the n-th term is n squared."},
{"role": "user", "content": "Using that same rule, what comes after 36 and 49?"},
]
ids = tok.apply_chat_template(messages, add_generation_prompt=True, tokenize=True,
enable_thinking=False)
ids = ids["input_ids"] if not isinstance(ids, list) else ids # transformers 5.x returns a dict here
with torch.no_grad():
out = model.generate(torch.tensor([ids], device=device), max_new_tokens=64,
do_sample=False, eos_token_id=tok.convert_tokens_to_ids("<|im_end|>"))
text = tok.decode(out[0, len(ids):], skip_special_tokens=True)
if "" in text: # non-thinking mode emits an empty think block
text = text.split("")[-1]
print(text.strip())
```
Training used `add_generation_prompt=True, enable_thinking=False` (the generation prompt then
ends with `<|im_start|>assistant\n\n\n\n\n`). Keep that format for best results.
The Hub checkpoint stores **float32** weights (2.4 GB). Passing `dtype=torch.bfloat16`, as
above, loads them at **1.20 GB** and leaves the (float32) rotary tables untouched; omitting
`dtype` gives the 2.40 GB float32 model. Both score identically on the benchmarks in §4.3.
### 5.3 Troubleshooting
| Error | Cause | Fix |
|---|---|---|
| `does not recognize this architecture` | custom architecture not registered | `AutoConfig.register` + `AutoModelForCausalLM.register` as in §5.2 |
| 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` |
| `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) |
## 6. Known Limitations
1. **0.6B scale.** Multi-step arithmetic is still unreliable; the SFT teaches the *behaviour*
(reuse the earlier rule/format, answer directly), not new reasoning ability.
2. **The shared branch is still barely used** (gate ≈ 0.0101 after SFT, same as the old
revision). Whether P0 pays off can only be decided in a long-context / window-shrink regime.
3. **No inference speedup — in fact a slowdown.** As shipped, the model is 3.9–5.1× slower to
decode and ~50× slower to prefill than the dense base, and reads *more* attention keys,
because the released configuration forces the dense window open (see §4.4). Use it to study
the architecture, not to serve traffic.
4. **Trained on synthetic dialogues.** The 1,294 conversations are model-generated; style and
coverage are limited to the six transfer kinds and the topics they cover. Evaluation above
is on a held-out slice of the same distribution — not a general capability claim.
5. **Custom architecture, no llama.cpp/GGUF support** (SSA attention is not implemented there).
6. **Not an instruct model at large.** It is a 0.6B research model; expect terse answers and
occasional arithmetic slips.
## 7. License
Custom license (`LICENSE.custom.md` in this repository):
- **Personal / non-commercial use: free**, including running, modifying and publicly
distributing derivative models under the same license.
- **Commercial use requires a paid license** — contact **novaweb6868@outlook.com**.
- Metadata: `license: other`, tag `commercial-license-required`.
## 8. Citation
```bibtex
@misc{baihu_v1_flash,
title = {BaiHu-V1-Flash: an SSA (sparse-attention + SubQ) retrofit of Qwen3-0.6B-Base, fine-tuned on bilingual transfer dialogues},
author = {ZichenAI},
year = {2026},
url = {https://huggingface.co/ZichenAI/BaiHu-V1-Flash}
}
```