Pull requests / #283

#283 Prompt path: BF16 projections get the activation's BF16 remainder too

closed · @sergqwer · 0 コメント · GitHub で見る

BenchmarksNVIDIA / CUDAWindows

本文

## Summary

The prompt path fed BF16-rounded activations (2^-9 relative) to the BF16-weight projections, where decode feeds FP32.
Most of these projections feed discrete choices: the router (top-10), the indexer (block selection), SSM alpha/beta,
the shared-expert gate, the PLE key/value and the hyper-connection. A token whose router logits are close can pick
another expert on the prompt path than decode would.

Each such projection now takes x as `hi + lo`, both BF16 (`lo = bf16(x - hi)`), in two GEMMs summed in FP32. That
gives ~16 mantissa bits, decode's FP32 x to within ~1e-5. The kernels that write the BF16 images (`gr_norm*`,
`gr_silu`, `gr_mix*`, `to_bf16`) write the remainder alongside.

`STRATA_PREFILL_BF16X2` picks the scope: `2` (default) all but the hyper-connection, `1` all, `0` off.

## Measured

Measured on current main (0.1.29, d6708a4) with ISTA-DASLab's IQ2_XS, RTX 5090 32 GB, Ryzen 9 9950X3D, 128 GB DDR5, Windows 11. KL is first-token KL divergence (`STRATA_DUMP_FIRST_LOGITS`, #276); two runs of main itself differ by 0 on 5 of 6 short prompts and by 0.00013 on the sixth. First-token KL against mode 1 (every projection split):

| prompt | main (off) | this PR (mode 2) |
| --- | ---: | ---: |
| 1K | 0.0028 | 0.0014 |
| 2K | 0.0069 (top-1 differs) | 0.0024 (top-1 differs) |
| 4K | 0.0009 | 0.0012 |
| 32K | 0.074 (top-1 differs) | 0.072 (top-1 differs) |

Mode 2 roughly halves the error on the shorter prompts. At 32K both stay 0.07 away from mode 1 and pick another first
token, so the hyper-connection's split is what matters there. Mode 2 is the default because on the NVFP4 fork mode 1
cost 8% of prompt speed (4,360 vs 4,726 tok/s; mode 2 4,726 vs 4,776 off). If 32K accuracy matters more than that,
mode 1 is the better default. It is a one-line change, and your call.

## Switch

`STRATA_PREFILL_BF16X2=0` is main's behaviour, `1` splits every projection.

🤖 Generated with [Claude Code](https://claude.com/claude-code)

https://claude.ai/code/session_01VZy1yKaDDiA8a7svdwaHio

関連リンク

インストール・モデル・リリースへの站内リンク。