Repository navigation
[Kernel] Add tuned W8A8 block-FP8 GEMM configs for gfx950 (MI350X) - #57492
mustafayildirim wants to merge 9 commits into
Conversation
GLM-5.3-Flash MLA projection shapes (N,K): 3072x4096, 4096x1536, 4096x256, 512x4096. Swept per M-bucket on MI350X (ROCm 7.2.3) against _w8a8_triton_block_scaled_mm; replaces default tile configs. Serving GLM-5.3-Flash 8xMI350X TP=8: TTFT -39%/-13%, TPOT -22%/-13% (decode-heavy / long-context profiles) vs default configs. Signed-off-by: Mustafa YILDIRIM <mustafa@character.ai>
simondanielsson
left a comment
There was a problem hiding this comment.
Thanks for adding this!
FYI: ROCm/aiter#5602
Rename config files to AMD_Instinct_MI350X (add gfx950 to _ROCM_DEVICE_ID_NAME_MAP so get_device_name_as_file_name() resolves it instead of falling back to amd-smi market_name) and pretty-print the JSONs with indent=4. Signed-off-by: Mustafa YILDIRIM <mustafa@character.ai>
|
Both addressed in 1ad2a5b:
Re aiter#5602: complementary rather than overlapping — that one tunes aiter's own GEMM backends (FlyDSL/ASM/Torch, BF16 + A8W8) on MI355X, while these JSONs feed vllm's triton |
|
Thanks for the pointer! Looked at ROCm/aiter#5602 — that one covers gfx950 A8W8 block-scale + BF16 configs for the GLM-5.3 Flash projection shapes, on MI355X (256 CUs). This PR is MI350X (304 CUs) W8A8 per-tensor configs from tuning runs on our fleet, so they complement rather than overlap — different SKU, different scale scheme, and the vllm-side lookup needs the device-specific entries regardless. Happy to cross-check shape coverage against the 1,826 captured shapes from that PR if useful. |
Purpose
Add tuned W8A8 block-FP8 triton GEMM configs for gfx950 (AMD Instinct MI350X,
device_name=0x75b0). These cover the four MLA projection GEMM shapes exercised by GLM-5.3-Flash (Glm5NextForConditionalGeneration); the configs directory currently has zero entries for gfx950, so the kernel falls back to default tile configs with a "Performance might be sub-optimal!" warning at boot.Shapes added (N, K): (3072, 4096), (4096, 1536), (4096, 256), (512, 4096) — each with 12 M buckets spanning decode (M=1..256) and prefill (M=1024..16384).
Methodology
Swept triton launch configs (BLOCK_SIZE_M/N/K, GROUP_SIZE_M, num_warps, num_stages) per M bucket against
_w8a8_triton_block_scaled_mmon MI350X (ROCm 7.2.3, vllm nightly 0bfc7a1),triton.testing.do_benchtiming, selecting the fastest valid config per bucket. JSONs follow the existingget_w8a8_block_fp8_configs()lookup format.Validation
Serving GLM-5.3-Flash on 8x MI350X (TP=8, breakable cudagraphs, aiter MoE/MLA), identical
vllm bench serveruns before/after:Boot log confirms all 32 "Using default W8A8 Block FP8 kernel config" warnings replaced by "Using configuration from ..." lookups. Full correctness suite (mixed chunked-prefill/decode repro, 32K-524K long-context ladder) passes with zero memory-access faults.
Test Plan
N/A — data-only change (JSON config files); no code paths modified.