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