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

# 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

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

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.

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

@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}
}
Downloads last month
-
Safetensors
Model size
0.6B params
Tensor type
F32
ยท
Inference Providers NEW
This model isn't deployed by any Inference Provider. ๐Ÿ™‹ Ask for provider support

Model tree for NovaAI6868/BaiHu-V1-Flash

Finetuned
(718)
this model