Skip to content

feat(rocm): wvSplitK skinny GEMM for decode M<=4 — the #487 decode GEMM lever - #506

Draft
VikashLoomba wants to merge 1 commit into
mudler:mainfrom
VikashLoomba:row/ROCM-SKINNY-GEMM
Draft

feat(rocm): wvSplitK skinny GEMM for decode M<=4 — the #487 decode GEMM lever#506
VikashLoomba wants to merge 1 commit into
mudler:mainfrom
VikashLoomba:row/ROCM-SKINNY-GEMM

Conversation

@VikashLoomba

Copy link
Copy Markdown
Contributor

Row

BACKEND-ROCM — M5-adjacent decode perf, the RDNA3 skinny-GEMM path. Issue #487 (decode M=1 GEMMs on 128-tile rocBLAS), coordinating with @joral (gfx1200 generic line). This board is gfx1100 (RDNA3) — the only arch that can test the RDNA3-specific config.

What changed

Ports vLLM's wvSplitK_hf_sml_ (csrc/rocm/skinny_gemms.cu, de-torched) into NEW src/vt/rocm/rocm_skinny_gemm.hip and routes decode-skinny MatmulBT (M in 1..4, bf16, K%8==0, activation fits the 64KB LDS stage) to it in rocm_matmul_hipblaslt.hip, ahead of the default-off naive GEMV and the rocBLAS tile path. VT_ROCM_SKINNY=0 rolls back to BLAS for A/B. New cross-device case (decode-skinny MatmulBT, bf16, M∈{1,4}) gates it vs the CPU oracle.

Evidence (4× RX 7900 XTX gfx1100, ROCm 7.14, Release)

Kernel-level A/B vs the current rocBLAS path (same buffers, back-to-back, many iters):

shape rocBLAS wvSplitK speedup
qkv 5120×1024 33.1 us 13.2 us 2.52x
o_proj 1024×2048 18.3 us 5.3 us 3.47x
mlp gate/up 3072×1024 10.9 us 6.1 us 1.78x
lm_head 151936×1024 1238.2 us 340.1 us 3.64x

In-engine decode (Qwen3-0.6B, 128-token steady state): 88.1 vs 70.4 tok/s (+25%) with VT_ROCM_SKINNY on vs off. Qwen3.5-0.8B GDN model output unchanged.

Correctness: ported kernel standalone-validated vs CPU reference across real decode shapes (NMSE ~3e-6, zero bad outputs); new in-tree cross-device case 8/8 green; full ctest zero-delta vs base (same 7 pre-existing host/lane failures).

Speed claims

  • The numbers above were run on this board under $GPU_LOCK; recorded here with the repro recipe. (Kernel bench + in-engine A/B, same binary.)

Honest gaps

  • gfx1100/RDNA3 only. gfx1200 (RDNA4) and gfx9 (wave64) paths in the donor (MFMA variants) are NOT ported — those need their own boards. joral's gfx1200 line is the generic one.
  • The donor's bf16 path has no HW dot on gfx1x (unpacks to f32 mul-add), mirrored exactly; a dot2-f16 path exists for fp16 but our decode is bf16.
  • CuCount is the device MP count (donor passes it in); the sweep found CU=40-96 all correct and near-best on this board, but a per-shape autotune is a follow-on.

…routing (mudler#487)

Ports vLLM's wvSplitK_hf_sml_ (csrc/rocm/skinny_gemms.cu:351-573) — the
split-K, LDS-staged, CU-count-aware skinny GEMM that wins M<=4 shapes — and
routes decode-skinny MatmulBT (M 1..4, bf16, K%8==0, activation fits the 64KB
LDS stage) to it instead of the 128x128-macro-tile rocBLAS GEMM. gfx1100
(GFX1X/wave32), bf16. VT_ROCM_SKINNY=0 restores the BLAS path for A/B.

Measured on 4x RX 7900 XTX (gfx1100), ROCm 7.14, Release:
- kernel-level vs the current rocBLAS tile path, same buffers back-to-back:
  qkv(5120x1024) 2.52x, o_proj(1024x2048) 3.47x, mlp_gateup(3072x1024) 1.78x,
  lm_head(151936x1024) 3.64x (the issue's worst single case)
- in-engine decode, Qwen3-0.6B 128-token steady state: 88.1 vs 70.4 tok/s
  (+25%) with VT_ROCM_SKINNY on vs off; 0.8B GDN model output unchanged
- cross-device: new decode-skinny MatmulBT case green (8/8, NMSE vs CPU)

FOLLOWING_AGENTS_PROTOCOL

Following-Agents-Protocol: true
AI-Assisted: true
Assisted-by: pi:kimi-k3 [pi]
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant