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
|
Download README.md from ZichenAI/BaiHu-V1-Flash: direct link, hf CLI and curl.
- Browser
- Download file 18.1 kB
-
https://huggingface.co/ZichenAI/BaiHu-V1-Flash/resolve/main/README.md
- Command line
-
hf download hf://ZichenAI/BaiHu-V1-Flash/README.md
-
curl -L -o README.md https://huggingface.co/ZichenAI/BaiHu-V1-Flash/resolve/main/README.md
18.1 kB
| 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} | |
| } | |
| ``` | |