Skip to content

Add CuTeDSL MXFP8 K-groups scale rearrange kernel - #4704

Open
alexsamardzic wants to merge 5 commits into
gh/alexsamardzic/12/headfrom
gh/alexsamardzic/13/head
Open

alexsamardzic wants to merge 5 commits into
gh/alexsamardzic/12/headfrom
gh/alexsamardzic/13/head

Conversation

@alexsamardzic

@alexsamardzic alexsamardzic commented Aug 5, 2026 •

Copy link
Copy Markdown
Collaborator

(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:

pytest -q test/prototype/moe_training/test_kernels.py::test_mx_block_rearrange_2d_K_groups
Benchmarking results, for M,K derived from some realistic models
M x TotalK Groups Triton us CuTeDSL us CuTeDSL GB/s CuTeDSL speedup
2048x131072 8 33.87 18.35 918.0 1.85
4096x131072 8 46.12 29.24 1152.0 1.58
7168x131072 8 78.32 49.08 1201.1 1.60
8192x131072 8 131.11 47.33 1423.4 2.77
14336x131072 8 137.75 94.96 1241.5 1.45
28672x131072 8 273.66 155.57 1515.7 1.76
2048x262144 8 98.37 30.73 1093.9 3.20
4096x262144 8 120.86 48.13 1397.2 2.51
7168x262144 8 145.37 68.77 1711.1 2.11
8192x262144 8 277.34 98.25 1368.8 2.82
14336x262144 8 457.50 180.11 1306.7 2.54
28672x262144 8 501.25 315.47 1492.0 1.59
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 1000
Benchmarking results, for M,K being powers of two
M x TotalK Groups Triton us CuTeDSL us CuTeDSL GB/s CuTeDSL speedup
128x1024 8 4.10 6.14 2.0 0.67
128x2048 8 4.10 6.14 3.3 0.67
128x4096 8 6.14 6.15 6.0 1.00
128x8192 8 6.14 6.15 11.3 1.00
128x16384 8 10.25 6.15 22.0 1.67
128x32768 8 14.33 6.16 43.2 2.32
128x65536 8 16.38 8.19 64.5 2.00
128x131072 8 24.58 8.01 131.4 3.07
128x262144 8 68.48 8.19 256.5 8.36
128x524288 8 172.07 8.31 505.3 20.71
256x1024 8 4.10 6.14 4.0 0.67
256x2048 8 4.10 6.15 6.7 0.67
256x4096 8 6.14 6.15 12.0 1.00
256x8192 8 8.19 6.15 22.6 1.33
256x16384 8 8.20 6.15 44.0 1.33
256x32768 8 16.38 6.15 86.6 2.66
256x65536 8 20.49 8.19 129.0 2.50
256x131072 8 37.25 8.19 257.0 4.55
256x262144 8 96.22 8.20 512.6 11.74
256x524288 8 108.29 12.29 683.5 8.81
512x1024 8 4.10 6.14 8.0 0.67
512x2048 8 4.10 6.15 13.3 0.67
512x4096 8 6.14 6.15 24.0 1.00
512x8192 8 8.19 6.15 45.3 1.33
512x16384 8 8.19 6.77 79.8 1.21
512x32768 8 14.33 7.37 144.6 1.95
512x65536 8 28.66 8.19 258.0 3.50
512x131072 8 34.66 8.20 513.7 4.23
512x262144 8 83.94 12.30 683.6 6.83
512x524288 8 122.76 18.44 910.7 6.66
1024x1024 8 4.10 6.14 16.0 0.67
1024x2048 8 4.10 6.15 26.7 0.67
1024x4096 8 6.14 6.15 48.0 1.00
1024x8192 8 6.15 6.15 90.6 1.00
1024x16384 8 6.15 8.19 132.0 0.75
1024x32768 8 12.29 8.20 259.9 1.50
1024x65536 8 19.61 10.24 412.8 1.92
1024x131072 8 49.13 12.29 685.5 4.00
1024x262144 8 49.13 18.43 912.1 2.67
1024x524288 8 134.06 28.80 1166.3 4.66
2048x1024 8 4.10 6.14 32.0 0.67
2048x2048 8 4.10 6.14 53.3 0.67
2048x4096 8 6.14 6.15 96.0 1.00
2048x8192 8 6.15 8.19 136.0 0.75
2048x16384 8 10.24 8.20 263.8 1.25
2048x32768 8 16.38 10.24 416.1 1.60
2048x65536 8 22.54 14.11 599.1 1.60
2048x131072 8 24.74 18.47 912.0 1.34
2048x262144 8 51.18 26.65 1261.7 1.92
2048x524288 8 151.52 53.30 1260.3 2.84
4096x1024 8 4.10 6.14 64.0 0.67
4096x2048 8 4.11 6.15 106.6 0.67
4096x4096 8 6.15 8.20 143.9 0.75
4096x8192 8 7.53 8.20 271.8 0.92
4096x16384 8 12.43 10.24 422.5 1.21
4096x32768 8 16.39 12.28 693.8 1.34
4096x65536 8 50.99 18.44 917.1 2.77
4096x131072 8 59.44 30.71 1096.9 1.94
4096x262144 8 192.62 52.66 1277.0 3.66
4096x524288 8 276.38 79.92 1681.1 3.46
8192x1024 8 4.10 6.15 127.9 0.67
8192x2048 8 6.14 8.19 160.0 0.75
8192x4096 8 8.20 10.25 230.3 0.80
8192x8192 8 12.29 12.01 371.2 1.02
8192x16384 8 14.34 14.33 603.5 1.00
8192x32768 8 30.72 20.50 831.2 1.50
8192x65536 8 51.21 28.19 1199.6 1.82
8192x131072 8 78.98 51.28 1313.9 1.54
8192x262144 8 185.87 88.94 1512.1 2.09
8192x524288 8 424.00 153.41 1751.5 2.76
16384x1024 8 6.14 8.19 192.0 0.75
16384x2048 8 6.15 10.25 255.8 0.60
16384x4096 8 10.25 14.52 325.0 0.71
16384x8192 8 13.24 16.42 542.9 0.81
16384x16384 8 18.43 22.62 765.0 0.81
16384x32768 8 72.50 36.81 925.7 1.97
16384x65536 8 116.69 59.55 1135.8 1.96
16384x131072 8 149.05 84.49 1594.8 1.76
16384x262144 8 360.45 190.50 1411.9 1.89
16384x524288 8 570.32 354.23 1517.1 1.61

[ghstack-poisoned]
@alexsamardzic

alexsamardzic commented Aug 5, 2026 •

Copy link
Copy Markdown
Collaborator Author

@pytorch-bot

pytorch-bot Bot commented Aug 5, 2026 •

Copy link
Copy Markdown

🔗 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 Failures

As of commit 0c8b8e3 with merge base a701b6a (image):
💚 Looks good so far! There are no failures yet. 💚

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Aug 5, 2026
@alexsamardzic alexsamardzic added module: training quantize_ api training flow moe labels Aug 5, 2026
@alexsamardzic
alexsamardzic changed the base branch from gh/alexsamardzic/12/head to main August 5, 2026 19:07
@alexsamardzic
alexsamardzic changed the base branch from main to gh/alexsamardzic/12/head August 6, 2026 07:42
Comment thread scripts/validate_k_groups_cutedsl.py Outdated
@@ -0,0 +1,157 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

split between test and benchmarks

(1024, 4097, 8),
],
)
def test_cutedsl_mx_block_rearrange_2d_K_groups(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

prefix with _

@vkuzo

vkuzo commented Aug 13, 2026

Copy link
Copy Markdown
Contributor

@claude review

@claude

claude Bot commented Aug 13, 2026 •

Copy link
Copy Markdown

Claude finished @vkuzo's task in 2m 39s —— View job


Claude Code Review

  • Read changed files
  • Review kernel implementation
  • Review tests & validate script
  • Address reviewer's inline comments

Reviewed the CuTeDSL K-groups rearrange kernel, its test, and the validate script against the origin/gh/alexsamardzic/12/head base. Overall the kernel is well-structured (nice use of the 128-bit vectorized fast path with a byte-wise _load4 fallback, TMA bulk S2G store, and swizzled smem layout), and the benchmarks show solid speedups over Triton. Comments below, grouped by the three inline threads plus some additional findings.

Re: @vkuzo's inline comments

1. scripts/validate_k_groups_cutedsl.py — split between test and benchmarks ✅ agree.
This script does two jobs that already have homes in the repo:

  • Correctness (torch.equal vs the reference) belongs in test/prototype/moe_training/test_kernels.py — which the new test already covers.
  • Benchmarking vs Triton belongs in benchmarks/prototype/moe_training/mxfp8/bench_mx_block_rearrange_2d_K_groups.py, which already exists and benchmarks torch + Triton for this exact op. Adding a cutedsl column there (and optionally the CUDA-graph timing path) keeps everything discoverable and avoids a top-level scripts/ file that imports from benchmarks.utils.

I'd suggest deleting scripts/validate_k_groups_cutedsl.py and folding its CUDA-graph benchmark helper into the existing bench file.
Fix this →

2. test_kernels.py:471 — can this be combined with an existing test? ✅ agree, with one caveat.
The new test_cutedsl_mx_block_rearrange_2d_K_groups mirrors test_triton_mx_block_rearrange_2d_K_groups. They can be combined by parametrizing over the implementation, but note the two tests do not build their inputs the same way today:

  • Triton test: generate_jagged_offs(n_groups, total_k, multiple_of=block_size) then // block_size, and scales from to_mx(randn(...)).
  • CuteDSL test: scale_cols = ceil_div(total_k, block_size), offsets generated directly over scale_cols with multiple_of=1, and scales from arange().remainder(251).

To combine cleanly, unify the input/offset construction first, then add an impl param (e.g. ("triton", ...), ("cutedsl", ...)) sharing one reference comparison. The multiple_of=1 case (4097) is a good addition worth keeping in the merged parametrization.

3. cutedsl_rearrange_2d_k_groups.py:444 — prefix with _.
This is the public wrapper mx_block_rearrange_2d_k_groups_cutedsl. Two things to note if you rename it to _mx_block_rearrange_2d_k_groups_cutedsl:

  • It's re-exported in torchao/prototype/moe_training/kernels/mxfp8/__init__.py and imported by the test — both need updating.
  • The sibling mx_block_rearrange_2d_m_groups_cutedsl (m-groups file, line 645) is currently not underscore-prefixed, so consider whether both should be made consistent in the same pass.

Additional findings

Output is torch.zeros every call (perf). _..._impl allocates the output with torch.zeros((padded_rows, padded_cols)). The kernel already explicitly zero-fills padding within each active chunk's smem tile, but padding rows (rows ≥ rows up to padded_rows) and the inter-group column padding rely on this initial zeroing. Since blocks with active == 0 skip the store entirely, those output regions are never written by the kernel and genuinely need the pre-zero. Worth a brief comment noting the kernel depends on zeros (not empty) so a future change doesn't switch it and silently produce garbage in padding.

register_fake ignores chunk_width (line 428) — correct, since it only affects the launch config, not shape. Good.

Empty groups. With multiple_of=1, generate_jagged_offs can produce zero-size groups (end == prev). The scan handles this (chunks_in_group = ceil_div(0, CHUNK_WIDTH) == 0) and _group_start_after_padding adds ceil_div(0,4)*4 == 0, so it looks correct — just flagging it's exercised by the 4097 case and worth keeping covered.

Grid over-provisioning. max_active_col_chunks = ceil_div(cols, chunk_width) + num_groups launches up to num_groups extra column-chunk blocks to cover per-group chunk rounding. Those extra blocks find active == 0 and no-op on the store side, but still allocate smem and run the scan. Fine for correctness; minor waste that scales with num_groups. A short comment explaining the + num_groups term would help future readers.

Naming. File/function use lowercase k_groups while the Triton/torch API uses 2d_K_groups (capital K). This matches the existing m_groups CuteDSL file, so it's internally consistent — just noting the mismatch with the Triton naming in case you want uniformity.

No correctness bugs spotted in the load/store/swizzle logic on inspection; the input_col + N < group_end guards and the smem zero-fills line up with the reference layout, and the tests assert exact torch.equal.
· branch gh/alexsamardzic/13/head

[ghstack-poisoned]
@alexsamardzic

Copy link
Copy Markdown
Collaborator Author

All comments addressed.

[ghstack-poisoned]
[ghstack-poisoned]
[ghstack-poisoned]

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. module: training quantize_ api training flow moe

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants