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