BaiHu-V1-Flash

BaiHu-V1-Flash is an SSA (Sparse-attention + SubQ) retrofit of Qwen/Qwen3-0.6B-Base, fine-tuned on 1,294 bilingual multi-turn dialogues whose purpose is generalizing from a single worked example — the model is shown one worked rule / format / method in the first turn, and later turns ask it to reuse that rule on a new case it has never seen.

  • 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 (2.75M are SSA-only modules)
  • Architecture: SSA — every layer runs three attention paths (shared / local / sparse-SubQ), each with its own softmax, then summed
  • Training data: 1,294 synthetic dialogues (English 693 / Chinese 601), six transfer types
  • License: free for personal use; a paid license is required for commercial use (see License)

Revision note. This repository previously hosted the base-pretrain checkpoint of the same name (5.0M tokens of continued pretraining, no instruction tuning, no P0 revision). It has been replaced by the checkpoint described here. The two are not interchangeable: this one uses the revised shared branch (§2) and is a supervised fine-tune.


1. Revisions in this release

earlier release (now replaced) this release
Shared path one mean vector for the entire prefix one compressed vector per completed block (P0 revision)
Training continued pretraining, 5.0M tokens of web text supervised fine-tune, 1,294 transfer dialogues (assistant-token loss)
Behaviour base LM follows a rule / format established earlier in the dialogue

Both revisions share the same SSA architecture, base weights and license.

2. Architecture

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 attends over one compressed summary vector per completed 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 head group; real attention is computed only for the selected top-k blocks O(T·k·B)

The P0 revision (this release)

The previous revision compressed the whole prefix into a single mean vector. A query could only see one undifferentiated global average, the single-element softmax made that average inject at full weight, and during 5.0M tokens of continued pretraining the learned gate shrank instead of growing (0.0100 → 0.0129 → 0.0122) — i.e. the optimizer actively suppressed the branch.

This release implements the design the project always documented: one compressed key/value per completed 64-token block, so queries do a genuine softmax over blocks and "which block matters" becomes learnable. Cost is unchanged (O(T·T/B); block summaries are still computed once per layer).

Hyperparameters

Parameter Value Meaning
ssa_block_size 64 block size B
ssa_top_k 8 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 / ssa_compress_dim 128 / 128 router subspace / summary width

3. Training

Data. 1,294 synthetic multi-turn dialogues (en 693 / zh 601), each 6–12 messages: the first user turn gives a worked case or a rule, a later user turn introduces a new instance that can only be handled by reusing it, and the last turn pushes generalization one step further. Six transfer_kinds are covered: rule_induction, analogy_transfer, format_transfer, counterfactual, cross_domain, teaching_loop.

Split (stratified by language × kind, seed 0): train 1,165 sessions / 345.7K tokens (assistant 188.8K), val 129 sessions / 38.2K tokens (assistant 20.5K).

Recipe. Loss is computed on assistant tokens only; sequences are padded to a multiple of the SSA block size (64); rendering uses the tokenizer's own chat template (apply_chat_template, add_generation_prompt=False, i.e. the final assistant turn carries Qwen3's empty <think></think> block).

Item Value
Epochs / steps 3 / 219
Tokens seen 1.15M
Batch 4 × grad-accum 4 (effective 16)
Optimizer AdamW, lr 2e-5, cosine to 10%, 20-step warmup, wd 0.1
Precision bfloat16
Hardware / time RTX 4090D 24GB — 6.9 min, ~2950 tok/s, 19.6GB peak
Val loss / ppl 2.534 → 1.646 → 1.539 → 1.537 / 4.65 (converged after ~2 epochs)

To make the SSA-only modules trainable, the two output projections and the shared gate are re-initialised to a small non-zero scale (0.01) before training — starting from the exact donor checkpoint would freeze all 2.75M of them at zero gradient.

4. Evaluation

4.1 Held-out transfer loss (129 unseen sessions, 20,520 assistant tokens)

Both models evaluated with the identical script, mask and batching.

Metric Qwen3-0.6B-Base BaiHu-V1-Flash Change
loss 2.5339 1.5376 −39.3%
ppl 12.603 4.654 −63.1%
ppl (en) 14.836 5.118 −65.5%
ppl (zh) 10.542 4.187 −60.3%

Per transfer type (ppl):

Kind Base BaiHu-V1-Flash
rule_induction 7.27 2.69
analogy_transfer 23.69 9.08
format_transfer 18.58 5.72
counterfactual 9.13 3.67
cross_domain 16.40 6.51
teaching_loop 8.26 2.95

The improvement holds in all 8 slices (2 languages × 6 kinds + overall).

4.2 Generation samples (12 prompts, one per language × kind, greedy)

  • Base: 7/12 outputs collapse into repeated symbols (⚇⚇⚇, ацион, .TRAILING), the rest are off-topic or hallucinated — the base model has never seen this dialogue format.
  • BaiHu-V1-Flash: 0/12 degenerate; it reuses the rule/format from earlier turns and is correct on most samples (B12A, 100, 50, 25, 12.5, …, format rewrites). Arithmetic errors remain — a 0.6B model with 1.3K training dialogues still slips on multi-step math.

4.3 Standard benchmarks (lm-evaluation-harness)

Both checkpoints were scored with lm-evaluation-harness 0.4.13 in one environment, with the same prompts, batch size and dtype (bfloat16), over the entire evaluation split of every task — 2,376 arc_easy / 1,172 arc_challenge / 10,042 hellaswag / 1,838 piqa / 1,267 winogrande examples, with no --limit subsampling — so the two columns are like-for-like. The metrics are the harness's 0-shot numbers, scored in raw-completion mode (no chat template), which is what a base checkpoint supports.

Task Metric Qwen3-0.6B-Base BaiHu-V1-Flash Δ
arc_easy acc_norm 0.5791 ±0.0101 0.5939 ±0.0101 +0.0148
arc_challenge acc_norm 0.3848 ±0.0142 0.3882 ±0.0142 +0.0034
hellaswag acc_norm 0.5385 ±0.0050 0.5507 ±0.0050 +0.0122
piqa acc_norm 0.6997 ±0.0107 0.7084 ±0.0106 +0.0087
winogrande acc 0.5856 ±0.0138 0.6062 ±0.0137 +0.0206
unweighted mean 0.5575 0.5695 +0.0119

How to read this. All five deltas are positive, but each is between 0.3σ and 1.5σ of its own standard error, so the defensible claim is "no regression in general capability", not "the fine-tune made the model smarter". The movement is also not an SSA effect: a control run of the dense Qwen3-0.6B-Base fine-tuned with the identical data and recipe lands within ±0.005 of BaiHu-V1-Flash on every one of these tasks (arc_easy 0.5951, arc_challenge 0.3908, hellaswag 0.5511, piqa 0.7111, winogrande 0.6014). What this release did buy is the transfer behaviour of §4.1–§4.2 — a 63% perplexity drop on held-out dialogues and no degenerate generations.

Caveats, stated plainly:

  • These are English benchmarks. This harness build ships no C-Eval / CMMLU / C3 tasks, so Chinese capability is not measured above; the held-out split of §4.1 (which contains both languages) is the only Chinese-side evidence in this card.
  • The Hub checkpoint stores float32 weights and the harness ran it at bfloat16 (see §5.2). The two agree to within 0.002 on every task listed, so precision is not driving the table.

4.4 Inference cost and speed (RTX 4090D, bfloat16)

Measured with the project's own scripts/bench_resources.py on the card that trained the model (RTX 4090D 24 GB, no other process on it), bfloat16, torch.no_grad(), greedy decoding of 64 tokens after a prefill of 1,024 / 2,048 / 4,096 tokens, identical script for both models.

The checkpoint exactly as shipped (ssa_force_full_window: true — see below), 1,024-token prefill:

Metric Qwen3-0.6B-Base BaiHu-V1-Flash
Parameters (M) 596.0 598.8
Peak memory, prefill (GB) 1.63 1.52
Peak memory, generate (GB) 1.78 2.01
Prefill latency = TTFT (s) 0.025 1.48
TPOT — time per output token (ms) 17.3 68.4
Decode throughput (tok/s) 57.8 14.6
Attention FLOPs per token (GFLOPs) 0.118 0.202
Attention keys read per query (vs full attention) 1.00× 1.72×
GPU utilization mean (%) / power mean (W) 15.5 / 69.8 20.0 / 69.0

How both quantities scale with context (same run, every cell in the order base → ours):

Prefill TTFT (ms) TPOT (ms) Decode (tok/s) Keys read per query Generate peak (GB)
1,024 25 → 1,477 17.3 → 68.4 57.8 → 14.6 1.00× → 1.72× 1.78 → 2.01
2,048 39 → 2,113 19.8 → 78.7 50.4 → 12.7 1.00× → 1.45× 2.35 → 2.82
4,096 72 → 3,610 21.4 → 108.8 46.7 → 9.2 1.00× → 1.24× 3.53 → 4.43

Stated plainly:

  • There is no speed advantage at any length tested. Prefill (TTFT) is 50–59× slower in wall-clock — 1.48 s vs 25 ms at 1 K — decode is 3.9× slower at 1 K and 5.1× at 4 K, and peak memory is comparable (slightly lower for this model during prefill, slightly higher during generation).
  • The sparse path is not actually saving anything here. This checkpoint is trained and released with ssa_force_full_window: true: the dense local window always covers the entire causal prefix, so the shared and sparse paths are additive on top of full attention instead of replacing part of it. Hence keys read per query above 1.00×.
  • Read the 1.72× honestly. It is the measured cost of the shipped configuration, and it is also why the model spends more attention FLOPs per token than the dense base (0.202 vs 0.118 GFLOPs at 1 K) while still being far slower in wall-clock. Sparsity only starts to bite once top_k × block_size is small relative to the context (see the next point).
  • Turning the window shrinkage on is a flag, not a retrain — but it was never validated in that mode, because the weights were trained with the window forced open. For reference, with ssa_force_full_window=false the same weights read 0.90× / 0.54× / 0.29× as many keys at 1 K / 2 K / 4 K and decode at 20.7 / 18.5 / 14.8 tok/s. Even then it stays 2.5–3.2× slower than the dense base, and quality in that mode is unevaluated — treat the numbers as the architecture's ceiling on this implementation, not as a free win.
  • What follows from this is an implementation item, not an architecture one: the per-block Python loop in the SSA layer issues many small kernels, so launch overhead dominates the arithmetic saved (GPU utilization never exceeds ~28%). The sparse path has to be fused before any of this can pay off on real hardware.

4.5 Effect of the P0 revision

Measured with the same data, recipe and seed, P0 vs the old single-vector shared branch:

old shared branch P0 (per-block)
Held-out loss 1.5376 1.5376
Shared gate after SFT 0.01001 0.01001

Under this regime (short dialogues, full causal window in the local path) the shared branch carries almost no load, so the revision is quality-neutral and the gate still does not grow. P0's intended benefit is long-range: it should only show up once the local window is allowed to shrink / contexts are far longer than the 192-token window. That experiment has not been run yet — see Limitations.

5. How to Run Inference

5.1 This is a custom architecture

model_type: baihu_ssa is not in the Transformers registry; a plain AutoModelForCausalLM.from_pretrained(...) fails. Register the config/model classes first (three lines below). Also note the model does not use GenerationMixin — call its own generate() (greedy or temperature/top-k; no beam search).

The implementation lives in the project repository (not on the Hub). This checkpoint requires two revisions of that code: the P0 revision of modeling_baihu_ssa.py (an older copy loads without error but computes a different shared branch), and the dtype-aware loader in model_baihu_ssa.py — before it, from_pretrained(..., dtype=...) was silently ignored and the model always came back float32.

5.2 Minimal working example

import os
import sys

import torch
from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer

sys.path.insert(0, os.path.join("ssa_model", "src"))  # path to the cloned repo's 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 = "ZichenAI/BaiHu-V1-Flash"
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.bfloat16 if device == "cuda" else torch.float32

tok = AutoTokenizer.from_pretrained(REPO)
model = AutoModelForCausalLM.from_pretrained(REPO, dtype=dtype).to(device).eval()

messages = [
    {"role": "user", "content": "What is the rule behind this sequence? 1, 4, 9, 16, 25"},
    {"role": "assistant", "content": "The gaps between consecutive terms are 3, 5, 7, 9 "
                                     "(+2 each time), so the n-th term is n squared."},
    {"role": "user", "content": "Using that same rule, what comes after 36 and 49?"},
]
ids = tok.apply_chat_template(messages, add_generation_prompt=True, tokenize=True,
                              enable_thinking=False)
ids = ids["input_ids"] if not isinstance(ids, list) else ids   # transformers 5.x returns a dict here
with torch.no_grad():
    out = model.generate(torch.tensor([ids], device=device), max_new_tokens=64,
                         do_sample=False, eos_token_id=tok.convert_tokens_to_ids("<|im_end|>"))
text = tok.decode(out[0, len(ids):], skip_special_tokens=True)
if "<think>" in text:            # non-thinking mode emits an empty think block
    text = text.split("</think>")[-1]
print(text.strip())

Training used add_generation_prompt=True, enable_thinking=False (the generation prompt then ends with <|im_start|>assistant\n<think>\n\n</think>\n\n). Keep that format for best results.

The Hub checkpoint stores float32 weights (2.4 GB). Passing dtype=torch.bfloat16, as above, loads them at 1.20 GB and leaves the (float32) rotary tables untouched; omitting dtype gives the 2.40 GB float32 model. Both score identically on the benchmarks in §4.3.

5.3 Troubleshooting

Error Cause Fix
does not recognize this architecture custom architecture not registered AutoConfig.register + AutoModelForCausalLM.register as in §5.2
output is repetitive garbage wrong chat format (e.g. a raw prompt with no ChatML wrapper) render with the tokenizer's chat template, enable_thinking=False
CUDA error: no kernel image is available PyTorch without kernels for an old GPU (Maxwell, sm_52) pin torch 2.7.1+cu126 (2.8 dropped sm_50/sm_60)

6. Known Limitations

  1. 0.6B scale. Multi-step arithmetic is still unreliable; the SFT teaches the behaviour (reuse the earlier rule/format, answer directly), not new reasoning ability.
  2. The shared branch is still barely used (gate ≈ 0.0101 after SFT, same as the old revision). Whether P0 pays off can only be decided in a long-context / window-shrink regime.
  3. No inference speedup — in fact a slowdown. As shipped, the model is 3.9–5.1× slower to decode and ~50× slower to prefill than the dense base, and reads more attention keys, because the released configuration forces the dense window open (see §4.4). Use it to study the architecture, not to serve traffic.
  4. Trained on synthetic dialogues. The 1,294 conversations are model-generated; style and coverage are limited to the six transfer kinds and the topics they cover. Evaluation above is on a held-out slice of the same distribution — not a general capability claim.
  5. Custom architecture, no llama.cpp/GGUF support (SSA attention is not implemented there).
  6. Not an instruct model at large. It is a 0.6B research model; expect terse answers and occasional arithmetic slips.

7. License

Custom license (LICENSE.custom.md in this repository):

  • Personal / non-commercial use: free, including running, modifying and publicly distributing derivative models under the same license.
  • Commercial use requires a paid license — contact novaweb6868@outlook.com.
  • Metadata: license: other, tag commercial-license-required.

8. Citation

@misc{baihu_v1_flash,
  title  = {BaiHu-V1-Flash: an SSA (sparse-attention + SubQ) retrofit of Qwen3-0.6B-Base, fine-tuned on bilingual transfer dialogues},
  author = {ZichenAI},
  year   = {2026},
  url    = {https://huggingface.co/ZichenAI/BaiHu-V1-Flash}
}
Downloads last month
253
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 ZichenAI/BaiHu-V1-Flash

Finetuned
(719)
this model