Instructions to use NovaAI6868/BaiHu-V1-Flash with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use NovaAI6868/BaiHu-V1-Flash with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="NovaAI6868/BaiHu-V1-Flash") messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("NovaAI6868/BaiHu-V1-Flash", device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use NovaAI6868/BaiHu-V1-Flash with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "NovaAI6868/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": "NovaAI6868/BaiHu-V1-Flash", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/NovaAI6868/BaiHu-V1-Flash
- SGLang
How to use NovaAI6868/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 "NovaAI6868/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": "NovaAI6868/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 "NovaAI6868/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": "NovaAI6868/BaiHu-V1-Flash", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use NovaAI6868/BaiHu-V1-Flash with Docker Model Runner:
docker model run hf.co/NovaAI6868/BaiHu-V1-Flash
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:
- Both new branches had identically zero gradients (blocking). To make the converted
model bit-exact with the base model,
compress_outandrouter_outwere 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. - Routing was non-differentiable.
top-kproduces 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 ofrouter_q/router_kare exactly zero โ the router can never learn to route. The fix is to feed the selected blocks' scores, squashed throughtanhand gently scaled, into the attention logits as an additive bias. - 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.
- The shared branch needs an explicit gate. With a single-element softmax the
probability is always 1, so the initialization scale of
compress_outcannot 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 getCUDA error: no kernel image is available for execution on the device. Use torch 2.7.1+cu126, which still containssm_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 = 9blocks 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 totop_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
- 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.
- 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). - 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.
- Sparsity only pays off at long sequences. With
top_k=8andblock_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. - 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.
- 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.
- 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
- -
Model tree for NovaAI6868/BaiHu-V1-Flash
Base model
Qwen/Qwen3-0.6B-Base