Repository navigation
Add CuTeDSL MXFP8 K-groups scale rearrange kernel - #4704
alexsamardzic wants to merge 5 commits into
Conversation
|
Stack from ghstack (oldest at bottom): |
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/ao/4704
Note: Links to docs will display an error until the docs builds have been completed. ✅ No FailuresAs of commit 0c8b8e3 with merge base a701b6a ( This comment was automatically generated by Dr. CI and updates every 15 minutes. |
| @@ -0,0 +1,157 @@ | |||
| # Copyright (c) Meta Platforms, Inc. and affiliates. | |||
There was a problem hiding this comment.
split between test and benchmarks
| (1024, 4097, 8), | ||
| ], | ||
| ) | ||
| def test_cutedsl_mx_block_rearrange_2d_K_groups( |
There was a problem hiding this comment.
can this be combined with an existing test
| return scales_tensor.new_empty((padded_rows, padded_cols)) | ||
|
|
||
|
|
||
| def mx_block_rearrange_2d_k_groups_cutedsl( |
|
@claude review |
|
Claude finished @vkuzo's task in 2m 39s —— View job Claude Code Review
Reviewed the CuTeDSL K-groups rearrange kernel, its test, and the validate script against the Re: @vkuzo's inline comments1.
I'd suggest deleting 2.
To combine cleanly, unify the input/offset construction first, then add an 3.
Additional findingsOutput is
Empty groups. With Grid over-provisioning. Naming. File/function use lowercase No correctness bugs spotted in the load/store/swizzle logic on inspection; the |
|
All comments addressed. |
[ghstack-poisoned]
[ghstack-poisoned]
(Replaces #4604 due to incorrect ghstack base branch targeting.)
PR replaces the existing Triton MXFP8 K-groups scale rearrange kernel with a CuTeDSL implementation. It includes pytest coverage against a plain PyTorch reference and a validate/benchmark script for comparing correctness and performance against the current Triton path.
To test:
Benchmarking results, for M,K derived from some realistic models
Benchmarking command to produce the table above
python benchmarks/prototype/moe_training/mxfp8/bench_mx_block_rearrange_2d_K_groups.py \ 2048x131072x8 \ 4096x131072x8 \ 7168x131072x8 \ 8192x131072x8 \ 14336x131072x8 \ 28672x131072x8 \ 2048x262144x8 \ 4096x262144x8 \ 7168x262144x8 \ 8192x262144x8 \ 14336x262144x8 \ 28672x262144x8 \ --multiple-of 1 --cuda-graph-bench --graph-iters 1000Benchmarking results, for M,K being powers of two