Repository navigation
perf(ArcAttention): port vLLM rms_norm_kernel — fuse the V4 per-head Q RMSNorm (7 launches -> 1) - #184
perf(ArcAttention): port vLLM rms_norm_kernel — fuse the V4 per-head Q RMSNorm (7 launches -> 1)#184heydryft wants to merge 2 commits into
Conversation
…o one launch Ports vLLM's csrc/activation_kernels.cu `silu_and_mul_clamp` (and SGLang's deepseek_v4/silu_and_mul_masked_post_quant.cuh `silu_and_mul<kApplySwigluLimit>`), both Apache-2.0, attributed in the file header. Replaces the 8-launch candle chain (2 casts + 3 clamp binaries + silu + mul + cast, plus five 1-element H2D copies for the scalar clamp operands) with a single kernel, on both the routed-expert and shared-expert paths. Compiled by the dedicated IEEE (no fast-math) builder so it stays bit-identical to candle-kernels. Also replaces sinkhorn.cu's vacuous `#if defined(__USE_FAST_MATH__)` #error — nvcc 12.4 defines no such macro in either pass — with assert_ieee_kernel_flags() in build.rs. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
…SNorm Structural port of vLLM csrc/layernorm_kernels.cu:108-181 rms_norm_kernel and SGLang fused_add_rmsnorm.cuh:57 (both Apache-2.0, attributed in the header): one launch, in-kernel reduction, no intermediate tensors. Replaces seven candle launches per attention layer (43x per decode token): sqr -> fast_sum -> affine(1/n,0) -> affine(1,eps) -> recip -> sqrt -> broadcast_mul. Arithmetic is candle's, not upstream's: the chain accumulates the 512-term sum of squares in BF16 and the kernel reproduces that exactly, including fast_sum's pairwise reduction order and the host-side bf16 down-conversion of the 1/n and eps constants. ARC_QNORM_F32_ACC=1 exposes the float accumulator for a future quality A/B but is off by default and not bit-identical. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Code Metrics Report━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ Language Files Lines Code Comments Blanks ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ C Header 5 305 210 52 43 CSS 2 1181 1036 34 111 CUDA 73 25015 17952 4311 2752 Dockerfile 1 39 22 8 9 JavaScript 16 3546 2676 482 388 Jinja2 7 694 656 5 33 JSON 74 4600 4597 0 3 Makefile 1 6 5 0 1 Metal Shading Lan| 33 12224 9431 1142 1651 PowerShell 1 300 227 30 43 Python 145 15139 12482 811 1846 Shell 39 9777 6538 2596 643 Plain Text 4 3801 0 2479 1322 TOML 33 1498 1294 54 150 YAML 3 25 23 2 0 ───────────────────────────────────────────────────────────────────────────────── HTML 4 2687 2604 43 40 |- CSS 2 543 479 37 27 |- JavaScript 1 1233 1215 12 6 (Total) 4463 4298 92 73 ───────────────────────────────────────────────────────────────────────────────── Jupyter Notebooks 4 122 83 23 16 |- Markdown 1 60 30 22 8 |- Python 1 122 113 1 8 (Total) 304 226 46 32 ───────────────────────────────────────────────────────────────────────────────── Markdown 203 44777 0 34735 10042 |- BASH 72 1654 1202 331 121 |- C 3 17 17 0 0 |- CUDA 2 84 56 16 12 |- JSON 18 708 708 0 0 |- PowerShell 1 1 1 0 0 |- Python 23 1008 787 113 108 |- Rust 66 2051 1716 77 258 |- TOML 6 207 164 0 43 |- YAML 5 41 36 5 0 (Total) 50548 4687 35277 10584 ───────────────────────────────────────────────────────────────────────────────── Rust 672 328294 283104 16426 28764 |- Markdown 490 28532 471 24619 3442 (Total) 356826 283575 41045 32206 ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ Total 1320 490291 349935 88466 51890 ━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ |
791eba0 to
9fa561f
Compare
Retargeted at the integration branchBase changed: The queue is being restructured to the shape the owner asked for: one PR open against Two things had to land on
This PR was not closed and is not considered stale. An audit of the queue found the overwhelming majority of it to be real work that was never merged, not noise. What you need to do: rebase onto 📚 Stack order — this is the TOPThis PR used to be based on its parent PR's branch, which is exactly the pattern the Both are now retargeted at Merge the bottom PR FIRST, then this one. Until the bottom lands, this PR's diff against the integration branch will also contain the bottom's commits. 🔴
|
Stacked on #183 — merge that first (D20).
Structural port of vLLM's
csrc/layernorm_kernels.cu:108-181rms_norm_kerneland SGLang'sfused_add_rmsnorm.cuh:57: one block per row, the sum of squares reduced inside the kernel, the normalised row written back in the same launch, no intermediate tensors. Both Apache-2.0; notice in the header.V4's per-head Q RMSNorm ran as seven candle launches per attention layer, 43x per decode token:
Measured on H200, DeepSeek-V4-Flash qtip2 UQFF, b=1 decode, 200 tokens
Same binary;
ARC_NO_FUSED_QNORM=1selects the old chain (SwiGLU fusion from #183 on in both legs, so this isolates the qnorm).-0.5 ms/token, and the fused leg is faster in all three pairs. nsys, per decode token: -258.0 kernel launches (exactly 6 x 43 layers, i.e. 7 candle kernels replaced by 1), -86.0 H2D copies, -344.0 device allocations.
Cumulative with #183 vs. the unfused baseline: 58.146 -> 54.203 ms/token, -744.7 launches / -572.0 H2D / -1313.3 allocations per token.
Bit-identity
The arithmetic is candle's, not upstream's. Three things a naive rewrite gets wrong, all pinned here:
sum_keepdimrunsfast_sumwithblock_dim = min(1024, n).next_power_of_two(), one element per thread, then a pairwise tree — not a sequential accumulation. This is the same trap that made the firstsinkhornkernel fail its H200 A/B.affinedown-converts1/nandepsto bf16 on the host, so they are passed in as raw u16 bit patterns computed with the samefrom_f64.On-GPU A/B vs the exact candle chain: bit-identical over 33,440 elements across the real
[64, 512]shape and a non-power-of-two[7, 96]shape that exercises the identity padding, with a negative control (1 bf16 ULP -> reported difference). End-to-end, every model leg produced the same SHA-256 over 200 generated tokens as the baseline.Surfaced, not shipped
Accumulating 512 squares with an 8-bit mantissa carries roughly
sqrt(512) * 2^-8~ 9% relative error in the norm. That is a pre-existing quality defect in the code being replaced, not one this port introduces — and it is preserved here on purpose so this PR is a pure perf change.ARC_QNORM_F32_ACC=1exposes a float accumulator so it can be measured in a perplexity A/B; it is off by default and is not bit-identical. Worth a separate change.