Pull requests / #783
#783 perf(cuda): fused decode/verify/MTP kernels and graph launch reductions on 0.1.39
closed · @stuchapin909 · 0 评论 · 在 GitHub 查看
BenchmarksMulti-GPUAMD / HIPNVIDIA / CUDAModels & quantsSecurityWindows
描述
## Summary
This is the follow-up to #646 rebased cleanly onto `0.1.39` (`6f32ec0`). Since #646 merged only the initial commit on that branch (`884c14e`), this PR brings in the remaining verify-window, MTP, and batch-decode kernel and CUDA-graph fusions split into **4 self-contained commits by subsystem**.
Measured apples-to-apples on 2x RTX 3090 24 GB (`IQ2_XS`, `--layer-split` `25/23`, **100% of all 24,576 experts resident in VRAM on both baseline and PR**):
| Metric (2x RTX 3090, `IQ2_XS`, 100% VRAM Resident) | Pre-#646 (100% Resident) | `0.1.39` (#646 `884c14e`) | **`0.1.39` + This PR (#783)** | Pure Kernel/Graph Delta |
| :--- | :---: | :---: | :---: | :---: |
| **Verify Window Latency (`lru_code`, `T = 5..8`)** | `15.92 ms/win` | `~15.10 ms/win` | **`13.07 ms/win`** (`13.07–15.02` across suite) | **-13.4% vs `0.1.39`** (-17.9% vs pre-#646) |
| **3-Prompt Benchmark (`lru_code` / `skiplist` / `prose`)** | `164.1 / 156.4 / 151.0 tok/s` | `~170.2 / 162.0 / 156.5 tok/s` | **`186.6 / 177.1 / 170.4 tok/s`** | **+9% to +10% vs `0.1.39`** (+13.7% vs pre-#646) |
| **Live Coding Decode (53k–62k context, 4,790 tokens)** | — | — | **`210.7 – 217.5 tok/s`** | Sustained high-acceptance coding session |
*(Note: Compared to `0.1.38`'s default `26/22` split where GPU 0 held `491/512` experts at `18.31 ms/win` / `149.2 tok/s`, combining `0.1.39`'s `25/23` all-resident split with #646 + #783 reduces verify window time from `18.31 ms/win` to `13.07 ms/win`.)*
All changes are scoped strictly to `include/strata/{core,kernels}` and `src/{core,kernels/cuda}` (32 files), preserve **0-ULP bitwise parity** across all 25 GPU parity tests, and maintain full HIP (`gfx906` wave64 + RDNA wave32) compatibility.
---
## Commits & What's Included
### 1. Multi-token GR / RMSNorm / post-op kernels and exact-`T` specializations (`fused_gr.cu`, `gr.cu`, `native_gr_norm.cu`, `native_gr_postops.cu`, `native_bf16.cu`, `quantize_act.cu`)
- **`fused_gr.cu`**:
- Adds pre-unpacked BF16 8-wide dot helpers (`Bf16x8`, `unpack8`, `dot8u`, `dot8u_ptr`) so weight tiles are unpacked once across tokens.
- Templates `gr_up_multi_kernel<MAX_T, EXACT_T>` and `gr_down_staged_kernel<MAX_T, EXACT_T>` to eliminate dynamic loop bounds for `T = 1..8` (`STRATA_GR_DOWN_MAX4=0` disables the `MAX_T=4` specialization for `ct <= 4`).
- Preserves `f16_from_f32`, `STRATA_DP4A`, and `STRATA_LDG` from `f166564` for HIP/ROCm portability.
- **`gr.cu` (`gr_read_multi`, `gr_write_multi`)**: Multi-token hyper-connection read and write across `n_tok` tokens in one launch.
- **`native_gr_norm.cu` (`native_gr_rms_norm_weighted_multi`)**: Multi-token weighted RMSNorm across `n_tok` rows.
- **`native_gr_postops.cu` (`native_gr_pre_gated_multi`, `native_gr_post_multi`)**: Multi-token pre-gated mixing and residual post-ops.
- **`native_bf16.cu` (`bf16_f32_mmvf_multi_kernel<BLOCK_SIZE, NT, EXACT_T>`)**: Compile-time token-count dispatch in `bf16_gemv_fp32_mmvf_multi`.
- **`quantize_act.cu`**: Warp-shuffle reduction in `quantize_q8_0_scaled_kernel`.
### 2. IQ/MMVQ sub-warp & 2-row kernels and MoE / shared-expert fusions (`iq_kernels.cu`, `native_mmvq.cu`, `native_moe.cu`, `native_router.cu`, `s2_expert_grouped.cu`, `shared_expert.cu`, `elementwise.cu`)
- **`iq_kernels.cu`**:
- Adds 16-lane sub-warp `row_dot_80_sub16<TY, NC, EXACT_N, STAGE_GRID>` and templates `mmvq_multi_kernel<TY, NC, STAGE_GRID, EXACT_N>` (`launch_mmvq_multi_nc`).
- Adds fused gate+up handling (`gu_split`) in `native_gu_multi_kernel` alongside `0.1.39`'s `down_iq4nl_multi_kernel`.
- **`native_mmvq.cu`**: Adds `SmallTraits<Q40Block, 4>`, `SmallTraits<Q50Block, 4>`, `SmallTraits<Q80Block, 8>`, and `SmallTraits<IQ4NLBlock, 4>` 2-row-per-warp (`ROWS = 2`) multi-column MMVQ specialization in `launch_multi_n` when `!multi_exact`.
- **`native_moe.cu` (`combine_k10_vec4`)**: `float4`-vectorized $k=10$ expert output combination in `native_moe_combine` and `native_moe_combine_multi`.
- **`native_router.cu` (`route_multi`, `native_router_top10_multi`)**: Multi-token top-10 router kernel across `n_tok` tokens in a single launch.
- **`s2_expert_grouped.cu` (`swiglu_quantize_q8_0_scaled_kernel`)**: Fuses SwiGLU and scaled `Q8_0` quantization in `moe_grouped_s2` while writing back `gate_up[idx]` so `s2_expert_grouped_parity` remains 100% bitwise identical (`STRATA_OLD_GROUPED=1` selects the previous path).
- **`shared_expert.cu` (`sigmoid_scale_rows_vec4_kernel`, `launch_sigmoid_scale_rows`)**: `float4`-vectorized sigmoid gating in `shared_expert` and `shared_expert_multi`.
- **`elementwise.cu` (`copy_rows_from_mapped_kernel<ZERO_HITS>`)**: Allows `copy_rows_from_mapped` (`zero_hits = false`) to skip redundant zero-fill when all experts are resident.
### 3. Batched KV append, QSA indexer/RoPE fusion, and parallel `resident_plan_kernel` (`kv_q4.cu`, `kv_q8.cu`, `native_qsa_indexer.cu`, `native_rope.cu`, `verify_kernels.cu`, `fused_gdn.cu`, `qsa_select.cu`, `kv_stream.cu`)
- **`kv_q4.cu` (`kv_append_q4_steps`) & `kv_q8.cu` (`kv_append_q8_steps`)**: Append K and V across `n_tok` tokens in a single 3D grid launch.
- **`native_qsa_indexer.cu` (`native_qsa_indexer_append_steps`)**: Multi-token QSA indexer append in a single launch.
- **`native_rope.cu` (`norm_rope_kernel<TAB>`, `native_qsa_rms_norm_rope`)**: Fuses per-head QSA RMSNorm and RoPE application.
- **`verify_kernels.cu`**:
- **`resident_plan_kernel<ALL_RESIDENT>`**: Keeps `kResidentPlanMax = 128` (`f945515`) and replaces thread-0's serial $O(n^2)$ bubble sort with parallel shared-memory grouping and a single packed `(my_cnt << 16) | (is_first ? 1 : 0)` prefix scan using `__shfl_up_sync(0xffffffffu, pref, d)`. Using width-32 `__shfl_up_sync` without `__ballot_sync` keeps the scan identical on both 32-lane warps (CUDA / RDNA) and 64-lane wavefronts (`gfx906`).
- Templates `gdn_ab_multi_kernel<MAX_T, EXACT_T>` and `gdn_step_norm_multi_kernel<ALL_OUT>`.
### 4. Wire multi-token kernels into `Verifier` and `MtpDrafter` CUDA graphs (`verify.cpp`, `mtp.cpp`, `verify.hpp`)
- **`verify.cpp`**:
- Replaces per-token loops in `Verifier::record_window` with `native_gr_rms_norm_weighted_multi`, `native_gr_pre_gated_multi`, `native_gr_post_multi`, `native_qsa_rms_norm_rope`, `kv_append_q8_steps`, `kv_append_q4_steps`, `native_qsa_indexer_append_steps`, and `native_router_top10_multi`.
- Skips redundant `copy_rows_from_mapped` zero-fill on the all-resident `direct_parts` path.
- Note on environment flags: `STRATA_FUSE_HEAD_GR=1` (fused LM-head GR read) and `STRATA_VERIFY_PRELAUNCH=1` (cross-stage prelaunch) are both **opt-in (default off)** so default `--layer-split` commit and stage ordering remain identical to `0.1.39` (`deee447`).
- **`mtp.cpp`**:
- Dispatches `gr_read_multi`, `gr_write_multi`, `native_gr_rms_norm_weighted_multi`, `native_qsa_rms_norm_rope`, `kv_append_q8_steps`, `kv_append_q4_steps`, `native_qsa_indexer_append_steps`, and `native_router_top10_multi` inside `MtpDrafter::record_forward`.
---
## Verification
### GPU Parity Suite (CTest — 25/25 Passed)
Built and ran the GPU parity test suite on Windows 11 / CUDA 13.2 (2x NVIDIA GeForce RTX 3090 24 GB):
- `mmvq_multi_parity` — Passed
- `iq_multi_parity` — Passed
- `iq_parity_fixtures` / `iq_parity` — Passed
- `native_grouped_parity` — Passed
- `shared_expert_parity` — Passed
- `s2_expert_grouped_parity` — Passed (bitwise identical across all 22 checks)
- `s2_gemv_parity` / `s2_gemv_q8_parity` — Passed
- `gr_parity` — Passed
- `gdn_parity` — Passed
- `qsa_parity` / `qsa_topk_parity` — Passed
- `rope_parity` — Passed
- `quantize_act_parity` — Passed
- `route_window_parity` / `router_top10_parity` — Passed
- `elementwise_parity` — Passed
- `bf16_gemv_parity` — Passed
- `sampler_parity` — Passed
- `cvec_parity` — Passed
- `kv_q8_parity` / `kv_q4_parity` / `kv_hybrid_parity` / `kv_stream_parity` — Passed
站内延伸阅读
链到安装、模型与版本说明,便于 SEO/GEO,非官方 issue 正文。