NovaAI commited on
Commit
d57da8f
·
verified ·
1 Parent(s): eb3396b

Upload README.md with huggingface_hub

Browse files
Files changed (1) hide show
  1. README.md +163 -10
README.md CHANGED
@@ -98,9 +98,162 @@ Several defects silently break training and are worth documenting:
98
 
99
  ---
100
 
101
- ## 3. Evaluation
102
 
103
- ### 3.1 Language modeling perplexity (validation set, identical windows)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- ### 3.2 Standard benchmarks (lm-evaluation-harness)
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
- ### 3.3 Inference compute and resource usage
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
- ### 3.4 Sparsity
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 3.3: at these sequence lengths the sparse path is not
174
  yet saving meaningful work.
175
 
176
  ---
177
 
178
- ## 4. Known Limitations
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.3). Continuing
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
- ## 5. License
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
- ## 6. Citation
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,