Text Generation
Transformers
Safetensors
Chinese
English
baihu_ssa
sparse-attention
subq
ssa
long-context
supervised-fine-tuning
transfer-learning
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)# pip install -U transformers accelerate # 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
NovaAI commited on
Upload README.md with huggingface_hub
Browse files
README.md
CHANGED
|
@@ -98,9 +98,162 @@ Several defects silently break training and are worth documenting:
|
|
| 98 |
|
| 99 |
---
|
| 100 |
|
| 101 |
-
## 3.
|
| 102 |
|
| 103 |
-
### 3.1
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 104 |
|
| 105 |
| Sequence length | Qwen3-0.6B-Base | BaiHu-V1-Flash |
|
| 106 |
|---|---|---|
|
|
@@ -108,7 +261,7 @@ Several defects silently break training and are worth documenting:
|
|
| 108 |
| 1024 | 3.4691 / ppl 32.109 | 4.2575 / ppl 70.631 |
|
| 109 |
| 2048 | 3.0801 / ppl 21.761 | 3.9213 / ppl 50.467 |
|
| 110 |
|
| 111 |
-
###
|
| 112 |
|
| 113 |
| Task | Qwen3-0.6B-Base | BaiHu-V1-Flash | Delta |
|
| 114 |
|---|---|---|---|
|
|
@@ -117,7 +270,7 @@ Several defects silently break training and are worth documenting:
|
|
| 117 |
| piqa | 0.7050 | 0.6950 | -0.0100 |
|
| 118 |
| winogrande | 0.6300 | 0.6300 | +0.0000 |
|
| 119 |
|
| 120 |
-
###
|
| 121 |
|
| 122 |
| Metric | Qwen3-0.6B-Base | BaiHu-V1-Flash |
|
| 123 |
|---|---|---|
|
|
@@ -157,7 +310,7 @@ Negative findings, stated plainly:
|
|
| 157 |
(+0.070 acc_norm), winogrande is unchanged, while hellaswag (−0.015) and piqa (−0.010)
|
| 158 |
regress slightly.
|
| 159 |
|
| 160 |
-
###
|
| 161 |
|
| 162 |
Two measurements are reported because they answer different questions and are **not**
|
| 163 |
interchangeable:
|
|
@@ -170,15 +323,15 @@ interchangeable:
|
|
| 170 |
The first number averages over *all* query blocks including the early ones, which have no
|
| 171 |
completed blocks available to select and therefore read nothing through the sparse path.
|
| 172 |
The second is the steady-state ratio for later queries, and it is the one that matters for
|
| 173 |
-
efficiency — see the note in section
|
| 174 |
yet saving meaningful work.
|
| 175 |
|
| 176 |
---
|
| 177 |
|
| 178 |
-
##
|
| 179 |
|
| 180 |
1. **Trained for very little.** Only 5.0M tokens (≈0.25 epoch). The new branches have not
|
| 181 |
-
converged; ppl is well above the base model and decode is slower (see 3.
|
| 182 |
to 200M tokens or more is required before the SSA layers can genuinely take over
|
| 183 |
long-range modeling.
|
| 184 |
2. **The shared branch does not appear to help and was actively suppressed by the
|
|
@@ -206,7 +359,7 @@ yet saving meaningful work.
|
|
| 206 |
|
| 207 |
---
|
| 208 |
|
| 209 |
-
##
|
| 210 |
|
| 211 |
**Free for personal use; a paid license is required for commercial use.**
|
| 212 |
|
|
@@ -225,7 +378,7 @@ component's original license.
|
|
| 225 |
|
| 226 |
---
|
| 227 |
|
| 228 |
-
##
|
| 229 |
|
| 230 |
```bibtex
|
| 231 |
@misc{baihu-v1-flash,
|
|
|
|
| 98 |
|
| 99 |
---
|
| 100 |
|
| 101 |
+
## 3. How to Run Inference
|
| 102 |
|
| 103 |
+
### 3.1 Important: this is a custom architecture
|
| 104 |
+
|
| 105 |
+
`BaiHu-V1-Flash` uses `model_type: baihu_ssa`, which is **not** in the Transformers
|
| 106 |
+
registry. Loading it with a plain `AutoModelForCausalLM.from_pretrained(...)` fails with:
|
| 107 |
+
|
| 108 |
+
```
|
| 109 |
+
ValueError: The checkpoint you are trying to load has model type `baihu_ssa`
|
| 110 |
+
but Transformers does not recognize this architecture.
|
| 111 |
+
```
|
| 112 |
+
|
| 113 |
+
You must register the config and model classes first. This is a one-time, three-line step
|
| 114 |
+
(see below).
|
| 115 |
+
|
| 116 |
+
Also note: **this model does not go through `transformers.GenerationMixin`.** SSA owns its
|
| 117 |
+
own KV-cache layout (`BaiHuSSACache`) because the cache stores per-block compressed prefix
|
| 118 |
+
sums rather than a plain growing key/value tensor. Call the model's own `generate()` method;
|
| 119 |
+
beam search is not supported.
|
| 120 |
+
|
| 121 |
+
### 3.2 Setup
|
| 122 |
+
|
| 123 |
+
```bash
|
| 124 |
+
# The SSA implementation lives in the project repository (not on the Hub),
|
| 125 |
+
# because it is a custom architecture.
|
| 126 |
+
git clone <this-project-repo> ssa_model
|
| 127 |
+
cd ssa_model
|
| 128 |
+
uv venv --python 3.11 .venv
|
| 129 |
+
|
| 130 |
+
# Ampere or newer (RTX 30xx/40xx, A100, ...): any recent torch works.
|
| 131 |
+
uv pip install --python .venv/Scripts/python.exe \
|
| 132 |
+
"numpy>=1.26" "transformers>=4.51" "safetensors>=0.4" "torch>=2.6"
|
| 133 |
+
|
| 134 |
+
# Maxwell / Pascal / Volta (GTX 9xx/10xx, TITAN X, V100, ...): pin torch 2.7.1 —
|
| 135 |
+
# see the GPU note below for why `torch>=2.6` would resolve to a build without your kernels.
|
| 136 |
+
uv pip install --python .venv/Scripts/python.exe \
|
| 137 |
+
"numpy>=1.26" "transformers>=4.51" "safetensors>=0.4" "torch==2.7.1"
|
| 138 |
+
```
|
| 139 |
+
|
| 140 |
+
> **GPU note (Maxwell and older).** If you are on an NVIDIA Maxwell card such as the
|
| 141 |
+
> GTX TITAN X (compute capability sm_52), recent PyTorch wheels no longer ship kernels for
|
| 142 |
+
> it: PyTorch **2.8 removed sm_50/sm_60**, and 2.8.0+cu126 only contains `sm_61…sm_90`.
|
| 143 |
+
> You will get `CUDA error: no kernel image is available for execution on the device`.
|
| 144 |
+
> Use **torch 2.7.1+cu126**, which still contains `sm_50`. Also train/infer in **fp32** on
|
| 145 |
+
> Maxwell: measured fp32 5.30 / fp16 4.46 / bf16 3.20 TFLOPS, so half precision is a loss.
|
| 146 |
+
|
| 147 |
+
### 3.3 Minimal working example
|
| 148 |
+
|
| 149 |
+
```python
|
| 150 |
+
import os
|
| 151 |
+
import sys
|
| 152 |
+
|
| 153 |
+
import torch
|
| 154 |
+
from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer
|
| 155 |
+
|
| 156 |
+
# Make the custom implementation importable (path to the cloned repo's src/)
|
| 157 |
+
sys.path.insert(0, os.path.join("ssa_model", "src"))
|
| 158 |
+
|
| 159 |
+
from baihu_ssa.configuration_baihu_ssa import BaiHuSSAConfig
|
| 160 |
+
from baihu_ssa.model_baihu_ssa import BaiHuSSAForCausalLM
|
| 161 |
+
|
| 162 |
+
# ---- register the custom architecture (REQUIRED) ----
|
| 163 |
+
AutoConfig.register("baihu_ssa", BaiHuSSAConfig, exist_ok=True)
|
| 164 |
+
AutoModelForCausalLM.register(BaiHuSSAConfig, BaiHuSSAForCausalLM, exist_ok=True)
|
| 165 |
+
|
| 166 |
+
REPO = "NovaAI6868/BaiHu-V1-Flash"
|
| 167 |
+
device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 168 |
+
dtype = torch.float32 if device == "cpu" else torch.float32 # see GPU note: fp32
|
| 169 |
+
|
| 170 |
+
tokenizer = AutoTokenizer.from_pretrained(REPO)
|
| 171 |
+
model = AutoModelForCausalLM.from_pretrained(REPO, dtype=dtype).to(device).eval()
|
| 172 |
+
|
| 173 |
+
prompt = "人工智能的未来是"
|
| 174 |
+
input_ids = tokenizer(prompt, return_tensors="pt").input_ids.to(device)
|
| 175 |
+
|
| 176 |
+
with torch.no_grad():
|
| 177 |
+
out = model.generate(
|
| 178 |
+
input_ids,
|
| 179 |
+
max_new_tokens=32,
|
| 180 |
+
do_sample=False, # greedy; supported
|
| 181 |
+
# do_sample=True, temperature=0.8, top_k=50, # sampling also supported
|
| 182 |
+
)
|
| 183 |
+
print(tokenizer.decode(out[0], skip_special_tokens=True))
|
| 184 |
+
```
|
| 185 |
+
|
| 186 |
+
### 3.4 `generate()` parameters
|
| 187 |
+
|
| 188 |
+
The built-in `generate()` is a self-contained decoding loop (greedy or
|
| 189 |
+
temperature/top-k sampling). Supported arguments:
|
| 190 |
+
|
| 191 |
+
| Argument | Default | Meaning |
|
| 192 |
+
|---|---|---|
|
| 193 |
+
| `max_new_tokens` | 32 | number of tokens to generate |
|
| 194 |
+
| `do_sample` | `False` | `False` = greedy, `True` = sample |
|
| 195 |
+
| `temperature` | 1.0 | sampling temperature (used when `do_sample=True`) |
|
| 196 |
+
| `top_k` | `None` | top-k sampling cutoff (used when `do_sample=True`) |
|
| 197 |
+
| `eos_token_id` | `None` | stop early when all sequences emit this token |
|
| 198 |
+
| `use_cache` | `True` | keep the SSA cache between steps; **leave on**, decoding is much slower without it |
|
| 199 |
+
|
| 200 |
+
**Not supported:** beam search (the SSA cache has no `reorder_cache`), and
|
| 201 |
+
`transformers` generation utilities such as `logits_processor` / `stopping_criteria`.
|
| 202 |
+
|
| 203 |
+
### 3.5 Computing perplexity / loss
|
| 204 |
+
|
| 205 |
+
```python
|
| 206 |
+
with torch.no_grad():
|
| 207 |
+
inputs = tokenizer("Your text here", return_tensors="pt").to(device)
|
| 208 |
+
loss = model(**inputs, labels=inputs["input_ids"]).loss
|
| 209 |
+
print("loss", loss.item(), "ppl", torch.exp(loss).item())
|
| 210 |
+
```
|
| 211 |
+
|
| 212 |
+
### 3.6 Throughput you should expect
|
| 213 |
+
|
| 214 |
+
Measured on a GTX TITAN X (sm_52, 12.9 GB), fp32, prompt 1024 tokens, 64 generated tokens:
|
| 215 |
+
|
| 216 |
+
| | Qwen3-0.6B-Base | BaiHu-V1-Flash |
|
| 217 |
+
|---|---|---|
|
| 218 |
+
| Prefill / TTFT | 503 ms | 1431 ms |
|
| 219 |
+
| Per-token decode (TPOT) | 36.6 ms | 80.5 ms |
|
| 220 |
+
| Decode throughput | 27.3 tok/s | 12.4 tok/s |
|
| 221 |
+
| Peak memory (generation) | 9.91 GB | 9.01 GB |
|
| 222 |
+
|
| 223 |
+
**BaiHu-V1-Flash is currently ~2.2× slower to decode than the base model**, even though the
|
| 224 |
+
sparse path reads fewer keys. The reason is implementation, not architecture: the attention
|
| 225 |
+
layer loops over query blocks in Python and launches many small kernels, so kernel-launch
|
| 226 |
+
overhead dominates the FLOPs saved. Memory use is lower; speed is not yet a win.
|
| 227 |
+
See "Known Limitations" below.
|
| 228 |
+
|
| 229 |
+
### 3.7 Expected output quality (honest warning)
|
| 230 |
+
|
| 231 |
+
This checkpoint has seen only **5.0M tokens**, so the new branches are far from converged.
|
| 232 |
+
Greedy decoding tends to **repeat itself** and it is **worse than the base model on
|
| 233 |
+
perplexity** (ppl 82.3 vs 39.4 at 512 tokens). A real greedy sample:
|
| 234 |
+
|
| 235 |
+
```
|
| 236 |
+
prompt: 人工智能的未来是
|
| 237 |
+
output: 人工智能的未来是怎样的?人工智能的未来是怎样的?人工智能的未来是怎样的?…
|
| 238 |
+
```
|
| 239 |
+
|
| 240 |
+
Use it to study or continue the SSA retrofit — **do not** expect it to match
|
| 241 |
+
`Qwen3-0.6B-Base` as a general-purpose model yet.
|
| 242 |
+
|
| 243 |
+
### 3.8 Troubleshooting
|
| 244 |
+
|
| 245 |
+
| Error | Cause | Fix |
|
| 246 |
+
|---|---|---|
|
| 247 |
+
| `does not recognize this architecture` / `KeyError: 'baihu_ssa'` | custom architecture not registered | call `AutoConfig.register` + `AutoModelForCausalLM.register` as in 3.3 |
|
| 248 |
+
| `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) |
|
| 249 |
+
| `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) |
|
| 250 |
+
| `BaiHuSSACache` errors when calling `model.generate(...)` from GenerationMixin | this model bypasses `GenerationMixin` | call `model.generate(...)` on the BaiHu model itself |
|
| 251 |
+
|
| 252 |
+
---
|
| 253 |
+
|
| 254 |
+
## 4. Evaluation
|
| 255 |
+
|
| 256 |
+
### 4.1 Language modeling perplexity (validation set, identical windows)
|
| 257 |
|
| 258 |
| Sequence length | Qwen3-0.6B-Base | BaiHu-V1-Flash |
|
| 259 |
|---|---|---|
|
|
|
|
| 261 |
| 1024 | 3.4691 / ppl 32.109 | 4.2575 / ppl 70.631 |
|
| 262 |
| 2048 | 3.0801 / ppl 21.761 | 3.9213 / ppl 50.467 |
|
| 263 |
|
| 264 |
+
### 4.2 Standard benchmarks (lm-evaluation-harness)
|
| 265 |
|
| 266 |
| Task | Qwen3-0.6B-Base | BaiHu-V1-Flash | Delta |
|
| 267 |
|---|---|---|---|
|
|
|
|
| 270 |
| piqa | 0.7050 | 0.6950 | -0.0100 |
|
| 271 |
| winogrande | 0.6300 | 0.6300 | +0.0000 |
|
| 272 |
|
| 273 |
+
### 4.3 Inference compute and resource usage
|
| 274 |
|
| 275 |
| Metric | Qwen3-0.6B-Base | BaiHu-V1-Flash |
|
| 276 |
|---|---|---|
|
|
|
|
| 310 |
(+0.070 acc_norm), winogrande is unchanged, while hellaswag (−0.015) and piqa (−0.010)
|
| 311 |
regress slightly.
|
| 312 |
|
| 313 |
+
### 4.4 Sparsity
|
| 314 |
|
| 315 |
Two measurements are reported because they answer different questions and are **not**
|
| 316 |
interchangeable:
|
|
|
|
| 323 |
The first number averages over *all* query blocks including the early ones, which have no
|
| 324 |
completed blocks available to select and therefore read nothing through the sparse path.
|
| 325 |
The second is the steady-state ratio for later queries, and it is the one that matters for
|
| 326 |
+
efficiency — see the note in section 4.3: at these sequence lengths the sparse path is not
|
| 327 |
yet saving meaningful work.
|
| 328 |
|
| 329 |
---
|
| 330 |
|
| 331 |
+
## 5. Known Limitations
|
| 332 |
|
| 333 |
1. **Trained for very little.** Only 5.0M tokens (≈0.25 epoch). The new branches have not
|
| 334 |
+
converged; ppl is well above the base model and decode is slower (see 3.6 and 4.1). Continuing
|
| 335 |
to 200M tokens or more is required before the SSA layers can genuinely take over
|
| 336 |
long-range modeling.
|
| 337 |
2. **The shared branch does not appear to help and was actively suppressed by the
|
|
|
|
| 359 |
|
| 360 |
---
|
| 361 |
|
| 362 |
+
## 6. License
|
| 363 |
|
| 364 |
**Free for personal use; a paid license is required for commercial use.**
|
| 365 |
|
|
|
|
| 378 |
|
| 379 |
---
|
| 380 |
|
| 381 |
+
## 7. Citation
|
| 382 |
|
| 383 |
```bibtex
|
| 384 |
@misc{baihu-v1-flash,
|