Pull requests / #603
#603 qsa_select: a 1,024-thread top-k past the register kernel's reach (CUDA) - a 243K-token prompt +26% on an RTX 3060
closed · @asp345 · 0 评论 · 在 GitHub 查看
BenchmarksSetup & installAMD / HIPNVIDIA / CUDAModels & quantsWindows
描述
## What On CUDA the register kernel (`block_topk_reg_kernel`) holds 33,792 blocks (~135K cells). Past that, `qsa_block_topk` takes the original kernel (`block_topk_kernel`): 256 threads, each reading its own stretch of the scores on all six passes, and one shared histogram updated with atomics. The dispatch follows the capacity (`--max-context`), so with `--max-context 262144`: - on a card without the sm_75 active-bound dispatch (#512), every prompt batch and every decode window takes the original kernel, whatever the prompt's length; - on sm_75, every prompt batch past 33,792 blocks and every decode window does. This adds one kernel for that range, `block_topk_wide_kernel`: the register kernel's layout (1,024 threads, a histogram per warp, warp scans) with the keys read from memory instead of registers. - **The four radix passes.** Thread `t` reads blocks `t`, `t + 1024`, ...: a histogram has no order, so a warp reads 32 neighbours at a time. - **The count and the output.** Warp `w` holds `per` rows of 32 consecutive blocks, and the cells are written in the order (warp, row, lane). The per-warp counts give each warp its first output position and the tied cells before it; a row without a selected block is skipped. - **The selection rule is the original kernel's:** radix threshold, ties to the lowest index, cells ascending. The ids are identical. - **Dispatch.** Where the dispatch took the original kernel because the blocks do not fit the registers, it now launches the new kernel. One exception: a call that names fewer than 4,608 blocks (a prompt batch below ~18K cells) keeps the original kernel, whose fixed cost is lower there (kernel table). The register kernel's range, the sm_75 active-bound dispatch, the sm_90+ cluster path and HIP are unchanged. HIP does not compile the kernel: I have no AMD card to test it on. - **One launch per call, no scratch memory, no capacity limit.** A decode window (no block count) takes it too. - **A/B switch.** `STRATA_TOPK_OLD=1` keeps the original kernel, as before. - **Test.** `qsa_topk_parity` (new, next to the other kernel parity tests) compares the dispatched top-k with `qsa_block_topk_ref`. ## Relation to #575 #575 covers the same range with a split: the register kernel on two halves of the blocks, then a merge kernel. I built #575 (`cb29a2a`) and this branch on the same base (`99f3dbd`, 0.1.38) and measured both on the same RTX 3060; the numbers are in the tables below. - On this card the kernel here is faster at every size in the kernel timing, and in the prompt speed from 32K tokens up. At 16K tokens the runs of the two overlap. - The split has a fixed cost of about 0.8 ms per 256-query call on this card. Below ~32K cells that is more than the original kernel takes, and a card without the sm_75 dispatch sends every prompt batch through it under a long `--max-context`: the 8K, 16K and 32K prompts spend more time in the selection than stock does. - The split covers the prompt path up to two register ranges (~270K cells). This kernel has no limit and also serves decode windows. I have no RTX 3090, so I cannot say how the two compare on #575's test card. ## Measured **Setup:** RTX 3060 12 GB (sm_86) on PCIe x16, Ryzen 5 5600X, 48 GB DDR4-2933, Debian 13, CUDA 13.4, base `99f3dbd` (0.1.38). Qwen3.8-Flash-Next IQ3_XXS on that one card, `--resident-experts --max-context 262144 --kv int8 --kv-resident 32768 --prefill auto --spec 4`. **Kernel** (`qsa_topk_parity CONTEXT 256 50 CAPACITY`: 256 queries with their block count, capacity 262,144 cells, the last two rows 524,288; ms per call). "Register" is #512's dispatch forced on this card with `STRATA_TOPK_ACTIVE_ANY=1`; it stops at 33,792 blocks, and its column is from a separate run of the same tool. | Context (cells) | Original | Register (#512) | #575 | This PR | This PR against the original | | ---: | ---: | ---: | ---: | ---: | ---: | | 4,096 | 0.071 | 0.271 | 0.805 | 0.071 | the same kernel | | 8,192 | 0.084 | 0.272 | 0.811 | 0.084 | the same kernel | | 16,384 | 0.199 | 0.259 | 0.790 | 0.199 | the same kernel | | 20,480 | 0.351 | 0.257 | 0.785 | 0.272 | 1.3x | | 32,768 | 0.845 | 0.271 | 0.805 | 0.328 | 2.6x | | 65,536 | 2.235 | 0.329 | 0.832 | 0.440 | 5.1x | | 98,304 | 4.205 | 0.617 | 1.149 | 0.724 | 5.8x | | 131,072 | 4.703 | 1.036 | 1.585 | 0.971 | 4.8x | | 163,840 | 6.876 | | 1.673 | 1.159 | 5.9x | | 196,608 | 11.26 | | 1.715 | 1.346 | 8.4x | | 262,144 | 16.09 | | 2.479 | 1.714 | 9.4x | | 393,216 | 26.02 | | 25.96 | 2.489 | 10.5x | | 524,288 | 37.53 | | 37.59 | 3.240 | 11.6x | Between ~20K and ~110K cells the register kernel is faster than this one by up to 0.15 ms per call; I left #512's restriction to sm_75 as it is. A decode window's call (4 queries, no block count, capacity 262,144): 0.031 -> 0.029 ms at 2,100 cells, 0.031 -> 0.028 at 8,192, 0.054 -> 0.034 at 32,768, 0.176 -> 0.083 at 131,072, 0.309 -> 0.127 at 262,144. #575 leaves these calls on the original kernel. **Prompts** (`tools/needle_bench.py --depths 50`, a fresh engine per run, `STRATA_PREFILL_TIMING=1` in every run; "original" is this branch with `STRATA_TOPK_OLD=1`). Several values in a cell are separate runs. Where a prompt reused a 16,384- or 32,768-token checkpoint of the previous prompt in the same engine, the tokens read are in brackets; the reuse is the same in every arm. Prompt speed, tok/s: | Prompt tokens | Original | #575 | This PR | This PR against the original | | ---: | ---: | ---: | ---: | ---: | | 16,471 | 1004.5 / 1016.2 | 988.8 / 984.4 | 989.0 / 1013.6 / 1015.3 | level | | 31,970 | 1043.3 / 1055.5 | 1034.1 / 1031.5 | 1060.1 / 1060.5 / 1059.6 | +1% | | 61,333 | 932.9 | 957.2 | 1013.8 | +8.7% | | 92,086 [75,702] | 919.5 | 972.9 | 985.4 | +7.2% | | 122,755 [106,371] | 904.7 / 904.2 | 975.1 / 971.6 | 992.7 / 989.6 / 988.9 | +9.5% | | 152,505 [119,737] | 845.0 | 941.3 | 963.9 | +14.1% | | 185,729 [169,345] | 819.7 | 933.8 | 954.7 | +16.5% | | 243,176 | 743.3 | 909.2 | 933.7 / 934.2 | +25.6% | The 243K prompt takes 260 s instead of 327 s. The selection phase of the prompt path (`qsa select` in the timing line: the block scores and the top-k), ms: | Prompt tokens | Original | #575 | This PR | | ---: | ---: | ---: | ---: | | 7,884 | 51 | 265 / 268 | 51 | | 16,471 | 195 | 653 / 652 | 195 | | 31,970 | 900 / 901 | 1,538 / 1,536 | 731 / 730 / 733 | | 61,333 | 4,061 | 3,835 | 2,428 | | 92,086 [75,702] | 10,366 | 6,622 | 5,071 | | 122,755 [106,371] | 20,019 / 20,037 | 11,285 / 11,280 | 9,034 / 9,046 / 9,052 | | 152,505 [119,737] | 30,807 | 16,030 | 13,431 | | 185,729 [169,345] | 50,055 | 24,448 | 20,761 | | 243,176 | 96,326 | 40,579 | 35,544 / 35,557 | The 7,884-token prompt is the first of each engine and its tok/s varies from run to run in every arm (600 to 925 overall), so it is in the selection table only. With #512's dispatch forced on this card (`STRATA_TOPK_ACTIVE_ANY=1`, one run) the selection takes 96, 286, 810 and 8,842 ms at 7,884, 16,471, 31,970 and 122,755 tokens, and the 122,755-token prompt reads at 998.2 tok/s. **Other setups** (the same model and flags, `STRATA_TOPK_OLD=1` against the default, tok/s of the prompt): | Setup | Prompt tokens | Original | This PR | Change | | --- | ---: | ---: | ---: | ---: | | One RTX 3060 on PCIe x4 | 122,754 | 423.0 | 441.5 | +4.4% | | | 243,176 | 378.4 | 419.5 | +10.9% | | Both RTX 3060, layer split auto, `--mmap-experts` | 122,754 | 863.8 / 867.5 | 890.5 / 873.7 | +1.9% | | | 243,176 | 807.1 | 871.7 | +8.0% | **Decode** (the x16 card, one long prompt, then 512 tokens generated three times, greedy, tok/s): | Context (tokens) | Original | This PR | | ---: | ---: | ---: | | 122,731 | 35.5 / 33.3 / 36.0 | 35.2 / 37.1 / 36.3 | | 243,153 | 33.4 / 35.7 / 35.3 | 34.2 / 36.5 / 35.6 | Repeated requests of one prompt do not always produce the same text, also within one arm, so only requests with identical text compare like for like. At 243K the first two requests produced identical text in both arms: 15.3 -> 14.9 s and 14.3 -> 14.0 s for the 512 tokens. At 122K no request pair has identical text, and the arms differ by less than the requests of one arm do. ## Test - `qsa_topk_parity --selftest` on both cards: 9 context / capacity pairs (9K; both sides of the register kernel's reach at 135,160 and 135,176 cells; 2,100, 9K and 32K under a 262K capacity; 262K; 300K under 524K; 524K), each called as a decode window calls it (8 queries, no block count) and as the prompt path does (64 queries with their block count), each with continuous scores, 64 score levels and equal scores. All 18 cases: identical ids. It also passes with `STRATA_TOPK_ACTIVE_ANY=1`, where a counted call past 33,792 blocks takes the new kernel. - `qsa_topk_active_parity` (#512's test) passes on both cards; with this change its call without a bound takes the new kernel at the 262,144 capacity. - The kernel table's runs: 256 of 256 queries identical to the reference at every size, also with 64 score levels and with equal scores. - Needles: found in all 59 prompts of the runs above, in every arm. - `decode_cluster_parity` skips on these cards (below sm_90). Not run: a HIP build (the kernel and its dispatch branch are inside `#if !defined(__HIPCC__)`), and a card other than the RTX 3060. 🤖 Generated with [Claude Code](https://claude.com/claude-code)
站内延伸阅读
链到安装、模型与版本说明,便于 SEO/GEO,非官方 issue 正文。