Skip to content

[Backend] Support TMA lowering for arbitrary (swizzled) SMEM layout - #2380

Merged
LeiWang1999 merged 1 commit into
tile-ai:mainfrom
Yongqi-Zhuo:advanced-tma-lowering
Jun 18, 2026
Merged

LeiWang1999 merged 1 commit into
tile-ai:mainfrom
Yongqi-Zhuo:advanced-tma-lowering

Conversation

@Yongqi-Zhuo

@Yongqi-Zhuo Yongqi-Zhuo commented Jun 11, 2026 •

Copy link
Copy Markdown
Collaborator

Summary by CodeRabbit

Release Notes

  • New Features
    • Added CuTe-style hierarchical layout algebra (including swizzle operations, coordinate mapping, and TileLang recovery) with a new Python API.
    • Enhanced CUDA TMA lowering using more consistent bulk-tile mapping and swizzle-aware compatibility.
  • Bug Fixes
    • Corrected MMA swizzle layout generation to match expected swizzle mapping output.
    • Improved TMA barrier sizing and rest/loop routing handling for multi-iteration cases.
  • Tests
    • Added comprehensive unit tests for CuTe semantics, swizzle recovery, and an end-to-end TMA load integration test.

@github-actions

Copy link
Copy Markdown

👋 Hi! Thank you for contributing to the TileLang project.

Please remember to run pre-commit run --all-files in the root directory of the project to ensure your changes are properly linted and formatted. This will help ensure your contribution passes the format check.

We appreciate you taking this step! Our team will review your contribution, and we look forward to your awesome work! 🚀

@coderabbitai

coderabbitai Bot commented Jun 11, 2026 •

Copy link
Copy Markdown
Contributor

Review Change Stack

No actionable comments were generated in the recent review. 🎉

ℹ️ Recent review info
⚙️ Run configuration

Configuration used: defaults

Review profile: CHILL

Plan: Pro

Run ID: d35eb2fc-706f-40a0-8c0d-3bce82807c46

📥 Commits

Reviewing files that changed from the base of the PR and between 3cc4648 and 847eeda.

📒 Files selected for processing (11)
  • examples/gemm/example_gemm_intrinsics.py
  • src/cuda/op/copy.cc
  • src/cuda/transform/producer_consumer_ws.cc
  • src/layout/cute_layout.cc
  • src/layout/cute_layout.h
  • testing/python/kernel/test_tilelang_kernel_gemm_simt.py
  • testing/python/layout/test_tilelang_cute.py
  • tilelang/cuda/intrinsics/layout/mma_layout.py
  • tilelang/layout/__init__.py
  • tilelang/layout/_cute_ffi_api.py
  • tilelang/layout/cute.py
✅ Files skipped from review due to trivial changes (1)
  • tilelang/layout/_cute_ffi_api.py
🚧 Files skipped from review as they are similar to previous changes (7)
  • examples/gemm/example_gemm_intrinsics.py
  • tilelang/layout/init.py
  • src/cuda/transform/producer_consumer_ws.cc
  • tilelang/cuda/intrinsics/layout/mma_layout.py
  • src/layout/cute_layout.h
  • src/cuda/op/copy.cc
  • tilelang/layout/cute.py

📝 Walkthrough

Walkthrough

Introduces a full CuTe-style layout IR (Swizzle, IntTuple, Layout, ComposedLayout) in C++ with symbolic recovery from TileLang layouts via GF(2) sampling and parity-based proof machinery. Wires the IR to Python via TVM FFI. Refactors TMA bulk-copy descriptor construction and IsTmaCompatibleLayout to consume the new IR, and rewrites MMA swizzle index generation to return three-component tuples with XOR-based swizzle arithmetic.

Changes

CuTe Layout IR, TMA Swizzle Recovery, and Consumers

Layer / File(s) Summary
CuTe layout IR header
src/layout/cute_layout.h
Declares all public C++ node types (SwizzleNode, IntTuple variants, LayoutNode, ComposedLayoutNode), inline helpers, template layout builders, and optional TileLang recovery API (LayoutFromTileLang, ComposedLayoutFromTileLang).
CuTe layout IR implementation
src/layout/cute_layout.cc
Implements Swizzle/IntTuple/Layout algebra (Coalesce, RightInverse, Composition, CoalesceX), proof infrastructure (AddrProbe, LowerBitOps, ProveZero, ProveEquivalent), RecoverPlainLayout, the two TileLang recovery entry points, and TVM FFI registration under tl.cute.*.
Python FFI wiring and cute module
tilelang/layout/_cute_ffi_api.py, tilelang/layout/cute.py, tilelang/layout/__init__.py
Initializes the tl.cute FFI namespace; defines Python wrapper classes for Swizzle, IntTuple, Layout, and ComposedLayout with arithmetic overloads, to_python/from_python converters, and layout algebra helpers; exposes tilelang.layout.cute at package import time.
TMA bulk-copy lowering refactor
src/cuda/op/copy.cc
Adds BulkCopyTile/ComputeBulkCopyTile CuTe helpers; refactors Copy::LowerBulk to select swizzle via ComposedLayoutFromTileLang, build TMA descriptor box/stride via CuTe composition and coalescing with hardware constraint enforcement, and generate addresses/coordinates via make_shared_offset/make_tma_coords lambdas with rest-loop support.
IsTmaCompatibleLayout rewrite
src/cuda/transform/producer_consumer_ws.cc
Replaces explicit swizzle-mode and linear-layout structural checks with ComposedLayoutFromTileLang + Recast + IsTMACompatible().
MMA swizzle index generation refactor
tilelang/cuda/intrinsics/layout/mma_layout.py, examples/gemm/example_gemm_intrinsics.py, testing/python/kernel/test_tilelang_kernel_gemm_simt.py
Rewrites get_swizzle_layout to produce three-component indices (col_tile, row_idx, swizzled_col) with new XOR-based swizzled_col logic and removes arith.Analyzer; updates make_mma_swizzle_layout and call sites to splice the full returned tuple.
Python unit and integration tests
testing/python/layout/test_tilelang_cute.py
Adds 40+ tests covering IntTuple arithmetic, layout construction/coalesce/right-inverse/composition, TileLang-to-CuTe recovery across all edge cases (affine, synthetic swizzle sweep, MMA canonical, broadcast, nonpow2 dims), and an end-to-end TMA load CUDA integration test with explicit swizzled shared-memory layout.

Sequence Diagram(s)

sequenceDiagram
  participant Kernel as TileLang Kernel
  participant LowerBulk as Copy::LowerBulk
  participant ComposedLayoutFromTileLang
  participant CuTeAlgebra as CuTe Coalesce/Compose/RightInverse
  participant TMADescriptor as CU_TENSOR_MAP descriptor
  participant GPU as CUDA TMA Unit

  Kernel->>LowerBulk: annotated shared-mem buffer + global tensor
  LowerBulk->>ComposedLayoutFromTileLang: shared TileLang Layout
  ComposedLayoutFromTileLang-->>LowerBulk: ComposedLayout (b_bits, m_base, s_shift, plain layout)
  LowerBulk->>LowerBulk: map swizzle → CU_TENSOR_MAP_SWIZZLE_*
  LowerBulk->>CuTeAlgebra: compose tile→smem and tile→global, RightInverse, Coalesce(max_extent)
  CuTeAlgebra-->>LowerBulk: TMA box dims, global strides, mode routing tables
  LowerBulk->>TMADescriptor: set box, global_shape, global_stride, swizzle mode
  TMADescriptor-->>LowerBulk: descriptor handle
  LowerBulk->>LowerBulk: build make_shared_offset / make_tma_coords lambdas
  LowerBulk->>GPU: emit tma_load(descriptor, smem_addr, tma_coords) in rest_size loop
Loading

Estimated code review effort

🎯 5 (Critical) | ⏱️ ~120 minutes

Suggested reviewers

  • lucifer1004
  • cherichy

🐇 A bunny hopped through layout land,
Where XOR swizzles were once unplanned—
Now IntTuples hop left and right,
ComposedLayouts shining bright,
And TMA maps with CuTe in hand! 🎉

🚥 Pre-merge checks | ✅ 4 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 23.79% which is insufficient. The required threshold is 80.00%. Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (4 passed)
Check name Status Explanation
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.
Title check ✅ Passed The title 'Support TMA lowering for arbitrary (swizzled) SMEM layout' accurately describes the main change across the PR, which extends TMA lowering to handle swizzled shared memory layouts, as evidenced by multiple files modifying copy operations, layout computations, and compatibility checks.
Linked Issues check ✅ Passed Check skipped because no linked issues were found for this pull request.
Out of Scope Changes check ✅ Passed Check skipped because no linked issues were found for this pull request.

✏️ Tip: You can configure your own custom pre-merge checks in the settings.

✨ Finishing Touches
🧪 Generate unit tests (beta)
  • Create PR with unit tests

Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out.

❤️ Share

Comment @coderabbitai help to get the list of available commands and usage tips.

@Yongqi-Zhuo
Yongqi-Zhuo force-pushed the advanced-tma-lowering branch from d62e9c7 to 75b65d2 Compare June 11, 2026 20:19
@Yongqi-Zhuo

Copy link
Copy Markdown
Collaborator Author

@regression-perf

@github-actions

Copy link
Copy Markdown

Performance Regression Test Report

Triggered by: @Yongqi-Zhuo
Workflow run: https://github.com/tile-ai/tilelang/actions/runs/27375053576

Results

File Original Latency Current Latency Speedup
example_mhc_pre 0.143004 0.147065 0.972386
example_dequant_gemm_fp4_hopper 0.716627 0.729351 0.982554
example_mha_sink_fwd_bhsd_sliding_window 0.0125145 0.0127113 0.984517
example_vertical_slash_sparse_attn 0.166225 0.167902 0.990016
example_tilelang_gemm_fp8_2xAcc 0.0907618 0.0916655 0.990142
example_topk 28.3166 28.5449 0.992001
example_fusedmoe_tilelang 0.0951073 0.0957638 0.993145
sparse_mla_fwd 0.0822961 0.0828427 0.993401
example_gemm_intrinsics 0.0253009 0.0254279 0.995007
sparse_mla_bwd 0.227706 0.228774 0.995332
example_mla_decode 0.317294 0.318764 0.995389
example_mha_inference 0.0621668 0.0623403 0.997217
example_dequant_gemm_bf16_fp4_hopper 0.396303 0.397136 0.997903
example_gemm 0.0170869 0.0171221 0.997942
example_tilelang_gemm_fp8 0.238845 0.239231 0.998388
example_tilelang_sparse_gqa_decode_varlen_indice 0.0117446 0.0117632 0.998421
example_mha_sink_bwd_bhsd_sliding_window 0.0379248 0.0379801 0.998544
block_sparse_attn_tilelang 0.00673224 0.00674049 0.998775
fp8_lighting_indexer 0.0228922 0.0229149 0.999009
example_elementwise_add 0.112944 0.113018 0.999345
example_tilelang_block_sparse_attn 0.00724329 0.00724789 0.999366
example_group_per_split_token_cast_to_fp8 0.00759987 0.00760016 0.999962
example_dequant_gemm_w4a8 3.82638 3.82619 1.00005
example_dequant_gemv_fp16xint4 0.0269809 0.0269761 1.00018
example_warp_specialize_gemm_barrierpipe_stage2 0.0295191 0.0295081 1.00037
example_tilelang_nsa_fwd 0.00528357 0.00527864 1.00093
example_gqa_sink_bwd_bhsd 0.0295935 0.0295549 1.00131
example_gqa_decode 0.0411168 0.0410624 1.00132
sparse_mla_fwd_pipelined 0.0592159 0.0591339 1.00139
example_warp_specialize_gemm_copy_0_gemm_1 0.0269204 0.0268757 1.00166
example_tilelang_gemm_splitk_vectorize_atomicadd 0.782353 0.780319 1.00261
example_blocksparse_gemm 0.0137579 0.0137174 1.00296
example_linear_attn_fwd 0.0285275 0.0284383 1.00314
example_tilelang_gemm_splitk 0.770004 0.767582 1.00316
example_per_token_cast_to_fp8 0.00651869 0.00649574 1.00353
example_tilelang_nsa_decode 0.00552887 0.0055089 1.00363
example_gqa_sink_bwd_bhsd_sliding_window 0.0180009 0.0179352 1.00366
example_mha_sink_fwd_bhsd 0.0127119 0.0126646 1.00373
example_dynamic 0.49892 0.496428 1.00502
example_tilelang_sparse_gqa_decode_varlen_mask 0.0127524 0.0126876 1.00511
topk_selector 0.0414984 0.0412809 1.00527
example_mhc_post 0.106842 0.106212 1.00592
example_gemv 0.20242 0.201195 1.00609
example_linear_attn_bwd 0.117829 0.116688 1.00978
example_mha_sink_bwd_bhsd 0.0525958 0.0519408 1.01261
example_dequant_gemm_bf16_mxfp4_hopper 0.366242 0.361017 1.01447
example_warp_specialize_gemm_copy_1_gemm_0 0.0194373 0.0190662 1.01946
example_warp_specialize_gemm_softpipe_stage2 0.0195237 0.0190239 1.02628
example_convolution 0.912849 0.876647 1.0413
example_convolution_autotune 0.737906 0.690159 1.06918

Artifacts

  • regression_result.png (speedup plot) is attached as a workflow artifact. Download it from the workflow run page above.

@Yongqi-Zhuo
Yongqi-Zhuo force-pushed the advanced-tma-lowering branch 2 times, most recently from ddf357e to f3088ec Compare June 13, 2026 15:14
@Yongqi-Zhuo

Copy link
Copy Markdown
Collaborator Author

@regression-perf

@github-actions

Copy link
Copy Markdown

Performance Regression Test Report

Triggered by: @Yongqi-Zhuo
Workflow run: https://github.com/tile-ai/tilelang/actions/runs/27470615557

Results

File Original Latency Current Latency Speedup
example_dequant_gemm_fp4_hopper 0.714993 0.724542 0.986821
example_dynamic 0.493846 0.498481 0.990701
example_mla_decode 0.313546 0.316364 0.991094
example_tilelang_gemm_fp8_2xAcc 0.0910377 0.0917162 0.992602
example_topk 28.2377 28.4417 0.992831
example_mha_inference 0.0627534 0.06318 0.993247
example_gqa_sink_bwd_bhsd 0.0295266 0.0297026 0.994074
sparse_mla_fwd_pipelined 0.0587608 0.0590753 0.994677
example_dequant_gemm_bf16_mxfp4_hopper 0.363279 0.36487 0.995638
example_gemm_intrinsics 0.0252776 0.025386 0.995727
sparse_mla_fwd 0.0823226 0.0825489 0.997258
example_mha_sink_bwd_bhsd 0.0517626 0.0518689 0.997951
example_elementwise_add 0.112912 0.113042 0.998849
example_tilelang_gemm_splitk 0.769031 0.769917 0.998849
example_dequant_gemv_fp16xint4 0.026959 0.0269836 0.999085
example_linear_attn_bwd 0.116756 0.116851 0.999182
example_per_token_cast_to_fp8 0.00650473 0.00650874 0.999383
example_tilelang_sparse_gqa_decode_varlen_indice 0.0117518 0.0117574 0.999522
example_gemv 0.201174 0.201242 0.99966
example_blocksparse_gemm 0.0137462 0.0137493 0.999769
example_gqa_sink_bwd_bhsd_sliding_window 0.0179448 0.0179449 0.999993
example_dequant_gemm_w4a8 3.82673 3.82666 1.00002
example_tilelang_sparse_gqa_decode_varlen_mask 0.0127454 0.012741 1.00035
example_convolution_autotune 0.735246 0.734926 1.00044
example_gemm 0.0171065 0.0170912 1.0009
fp8_lighting_indexer 0.0229259 0.0229014 1.00107
example_warp_specialize_gemm_copy_0_gemm_1 0.0270955 0.0270661 1.00109
example_tilelang_gemm_splitk_vectorize_atomicadd 0.793755 0.792861 1.00113
example_mhc_post 0.106632 0.106501 1.00123
example_warp_specialize_gemm_copy_1_gemm_0 0.0195646 0.0195385 1.00133
example_tilelang_block_sparse_attn 0.00724467 0.00723491 1.00135
example_convolution 0.915694 0.914437 1.00137
example_group_per_split_token_cast_to_fp8 0.00761172 0.00760082 1.00143
example_fusedmoe_tilelang 0.0954199 0.0952743 1.00153
example_gqa_decode 0.0411001 0.0410306 1.00169
example_linear_attn_fwd 0.0284992 0.0284453 1.0019
example_warp_specialize_gemm_barrierpipe_stage2 0.0295362 0.0294494 1.00295
example_tilelang_nsa_decode 0.00553415 0.00551754 1.00301
topk_selector 0.0413605 0.041232 1.00312
sparse_mla_bwd 0.232584 0.231712 1.00376
example_tilelang_nsa_fwd 0.00529075 0.00527001 1.00394
example_mha_sink_bwd_bhsd_sliding_window 0.038162 0.0380107 1.00398
example_mha_sink_fwd_bhsd_sliding_window 0.012584 0.0125265 1.00459
block_sparse_attn_tilelang 0.00673763 0.0067047 1.00491
example_warp_specialize_gemm_softpipe_stage2 0.0195509 0.0194358 1.00593
example_mha_sink_fwd_bhsd 0.0127136 0.0126376 1.00601
example_tilelang_gemm_fp8 0.240134 0.237771 1.00994
example_mhc_pre 0.146432 0.144515 1.01326
example_dequant_gemm_bf16_fp4_hopper 0.39763 0.391747 1.01502
example_vertical_slash_sparse_attn 0.167798 0.164883 1.01768

Artifacts

  • regression_result.png (speedup plot) is attached as a workflow artifact. Download it from the workflow run page above.

@Yongqi-Zhuo
Yongqi-Zhuo force-pushed the advanced-tma-lowering branch from f3088ec to 704c267 Compare June 15, 2026 04:22
@Yongqi-Zhuo

Copy link
Copy Markdown
Collaborator Author

@regression-perf

@github-actions

Copy link
Copy Markdown

Performance Regression Test Report

Triggered by: @Yongqi-Zhuo
Workflow run: https://github.com/tile-ai/tilelang/actions/runs/27524466521

Results

File Original Latency Current Latency Speedup
example_topk 28.1869 28.6164 0.984991
example_gemm_intrinsics 0.0252206 0.0254069 0.992667
example_tilelang_gemm_fp8_2xAcc 0.0901584 0.0908202 0.992713
example_dequant_gemm_bf16_fp4_hopper 0.395873 0.39846 0.993508
example_dequant_gemm_bf16_mxfp4_hopper 0.363622 0.365963 0.993604
example_tilelang_gemm_fp8 0.238538 0.240031 0.99378
example_mha_sink_bwd_bhsd_sliding_window 0.0383019 0.0385341 0.993974
topk_selector 0.0412984 0.0415472 0.994011
example_warp_specialize_gemm_copy_1_gemm_0 0.0194612 0.0195607 0.994915
example_mha_sink_fwd_bhsd_sliding_window 0.0125213 0.0125584 0.997042
example_gqa_sink_bwd_bhsd_sliding_window 0.0178957 0.017946 0.997195
example_elementwise_add 0.113119 0.113318 0.998243
example_fusedmoe_tilelang 0.0951681 0.0953289 0.998313
example_gemm 0.0171156 0.0171285 0.99925
example_tilelang_gemm_splitk 0.768391 0.768889 0.999353
example_dynamic 0.497132 0.497352 0.999557
example_dequant_gemv_fp16xint4 0.0269813 0.0269852 0.999856
sparse_mla_fwd 0.0825907 0.0826022 0.999861
example_mla_decode 0.316261 0.316302 0.999872
example_linear_attn_bwd 0.116437 0.116431 1.00005
example_linear_attn_fwd 0.0284592 0.0284533 1.00021
block_sparse_attn_tilelang 0.00673823 0.00673313 1.00076
example_mhc_post 0.106624 0.106542 1.00077
sparse_mla_fwd_pipelined 0.0591694 0.0591238 1.00077
example_tilelang_nsa_decode 0.00551784 0.0055074 1.00189
example_tilelang_nsa_fwd 0.00528697 0.00527552 1.00217
example_per_token_cast_to_fp8 0.00651011 0.00649482 1.00235
fp8_lighting_indexer 0.0229453 0.0228882 1.0025
example_gqa_decode 0.0411546 0.0410362 1.00289
sparse_mla_bwd 0.228597 0.227923 1.00296
example_gqa_sink_bwd_bhsd 0.0297876 0.0296971 1.00305
example_tilelang_sparse_gqa_decode_varlen_indice 0.0117962 0.0117578 1.00326
example_dequant_gemm_w4a8 3.82638 3.81227 1.0037
example_mhc_pre 0.145094 0.144554 1.00373
example_blocksparse_gemm 0.0137712 0.0137136 1.00421
example_tilelang_block_sparse_attn 0.00727468 0.00724418 1.00421
example_tilelang_sparse_gqa_decode_varlen_mask 0.012804 0.0127442 1.0047
example_warp_specialize_gemm_softpipe_stage2 0.0195355 0.0194298 1.00544
example_convolution 0.920249 0.914792 1.00596
example_convolution_autotune 0.732975 0.728286 1.00644
example_gemv 0.202622 0.201244 1.00685
example_mha_sink_fwd_bhsd 0.0127091 0.0126211 1.00697
example_mha_inference 0.0628724 0.0623925 1.00769
example_warp_specialize_gemm_copy_0_gemm_1 0.027247 0.0270306 1.00801
example_group_per_split_token_cast_to_fp8 0.00765111 0.00758002 1.00938
example_warp_specialize_gemm_barrierpipe_stage2 0.0296885 0.0293851 1.01033
example_tilelang_gemm_splitk_vectorize_atomicadd 0.79176 0.782033 1.01244
example_vertical_slash_sparse_attn 0.167853 0.165255 1.01572
example_mha_sink_bwd_bhsd 0.0528743 0.0520261 1.0163
example_dequant_gemm_fp4_hopper 0.724449 0.710129 1.02017

Artifacts

  • regression_result.png (speedup plot) is attached as a workflow artifact. Download it from the workflow run page above.

@Yongqi-Zhuo
Yongqi-Zhuo force-pushed the advanced-tma-lowering branch 6 times, most recently from 588c33f to c7585e6 Compare June 17, 2026 16:53
@Yongqi-Zhuo

Copy link
Copy Markdown
Collaborator Author

@regression-perf

@github-actions

Copy link
Copy Markdown

Performance Regression Test Report

Triggered by: @Yongqi-Zhuo
Workflow run: https://github.com/tile-ai/tilelang/actions/runs/27705603172

Results

File Original Latency Current Latency Speedup
example_linear_attn_fwd 0.0288481 0.0608715 0.473918
example_tilelang_nsa_fwd 0.00547886 0.00906023 0.604715
example_tilelang_nsa_decode 0.00556938 0.0089825 0.620026
example_elementwise_add 0.173212 0.27441 0.631216
block_sparse_attn_tilelang 0.0156604 0.022562 0.694108
example_dequant_gemm_w4a8 3.48024 4.61515 0.754089
example_gqa_sink_bwd_bhsd_sliding_window 0.0296813 0.0373217 0.795283
example_mhc_post 0.245294 0.289523 0.847235
example_tilelang_sparse_gqa_decode_varlen_indice 0.0290699 0.0342161 0.849598
example_convolution 2.23147 2.58806 0.862216
example_dequant_gemm_bf16_mxfp4_hopper 0.679015 0.780142 0.870374
example_mha_sink_bwd_bhsd 0.107778 0.121243 0.888944
example_dequant_gemm_fp4_hopper 1.42619 1.59906 0.891892
example_mha_sink_fwd_bhsd 0.0358678 0.0390396 0.918756
fp8_lighting_indexer 0.0253181 0.0273498 0.925717
example_vertical_slash_sparse_attn 0.357736 0.378747 0.944527
example_mhc_pre 0.323288 0.341974 0.94536
example_blocksparse_gemm 0.0154417 0.0158759 0.972649
example_mha_sink_bwd_bhsd_sliding_window 0.0886675 0.090262 0.982335
example_dequant_gemv_fp16xint4 0.0692528 0.0704823 0.982556
sparse_mla_fwd_pipelined 0.0640364 0.0649697 0.985634
example_mha_inference 0.128109 0.129494 0.989301
example_warp_specialize_gemm_copy_1_gemm_0 0.0468199 0.0470823 0.994426
example_tilelang_block_sparse_attn 0.011139 0.0111911 0.995342
sparse_mla_bwd 0.489196 0.488638 1.00114
sparse_mla_fwd 0.0840441 0.0838171 1.00271
example_topk 28.4739 28.3715 1.00361
example_linear_attn_bwd 0.310535 0.308314 1.0072
topk_selector 0.0752025 0.0740041 1.01619
example_gqa_sink_bwd_bhsd 0.0792224 0.0760411 1.04184
example_dynamic 1.37224 1.30987 1.04761
example_warp_specialize_gemm_softpipe_stage2 0.0494282 0.0467604 1.05705
example_warp_specialize_gemm_copy_0_gemm_1 0.0599127 0.055936 1.07109
example_dequant_gemm_bf16_fp4_hopper 1.08491 0.96555 1.12362
example_warp_specialize_gemm_barrierpipe_stage2 0.039732 0.0345956 1.14847
example_mla_decode 0.885889 0.727658 1.21745
example_convolution_autotune 1.68163 1.35944 1.237
example_mha_sink_fwd_bhsd_sliding_window 0.0501023 0.0392483 1.27655
example_per_token_cast_to_fp8 0.0133504 0.0100713 1.32558
example_tilelang_gemm_fp8_2xAcc 0.135772 0.0803289 1.6902
example_gqa_decode 0.0879916 0.0437858 2.00959
example_gemv 0.412593 0.202677 2.03571
example_tilelang_gemm_fp8 0.516199 0.249297 2.07062
example_tilelang_sparse_gqa_decode_varlen_mask 0.0467421 0.0205352 2.2762
example_tilelang_gemm_splitk_vectorize_atomicadd 1.83942 0.790183 2.32783
example_tilelang_gemm_splitk 1.85807 0.76762 2.42057
example_group_per_split_token_cast_to_fp8 0.0213835 0.00844416 2.53235
example_fusedmoe_tilelang 0.276863 0.0959193 2.88642

Artifacts

  • regression_result.png (speedup plot) is attached as a workflow artifact. Download it from the workflow run page above.

@Yongqi-Zhuo

Copy link
Copy Markdown
Collaborator Author

@regression-perf

@github-actions

Copy link
Copy Markdown

Performance Regression Test Report

Triggered by: @Yongqi-Zhuo
Workflow run: https://github.com/tile-ai/tilelang/actions/runs/27708927143

Results

File Original Latency Current Latency Speedup
example_warp_specialize_gemm_barrierpipe_stage2 0.0411873 0.065007 0.633583
topk_selector 0.0782948 0.119472 0.655338
example_tilelang_sparse_gqa_decode_varlen_indice 0.0274068 0.0411026 0.66679
example_tilelang_gemm_splitk 1.36217 1.99577 0.682528
example_tilelang_gemm_splitk_vectorize_atomicadd 1.54131 1.86092 0.828254
example_gqa_sink_bwd_bhsd 0.0583556 0.0697172 0.837033
example_warp_specialize_gemm_copy_1_gemm_0 0.0362884 0.0430421 0.843092
example_per_token_cast_to_fp8 0.0099558 0.0117462 0.847574
example_dequant_gemm_w4a8 3.85383 4.53657 0.849502
example_linear_attn_bwd 0.257778 0.295276 0.873009
example_gqa_sink_bwd_bhsd_sliding_window 0.0336014 0.0373724 0.899097
sparse_mla_bwd 0.567438 0.628254 0.903198
example_vertical_slash_sparse_attn 0.384208 0.418994 0.916975
fp8_lighting_indexer 0.0248161 0.0264778 0.937243
example_convolution 2.30123 2.39524 0.960751
example_elementwise_add 0.296355 0.304827 0.972208
example_tilelang_nsa_decode 0.00555691 0.0056131 0.98999
example_dequant_gemm_fp4_hopper 1.7526 1.75387 0.999275
example_topk 28.3484 28.3454 1.0001
example_dynamic 1.30303 1.29916 1.00298
sparse_mla_fwd_pipelined 0.0629631 0.0627624 1.0032
sparse_mla_fwd 0.0847349 0.0843251 1.00486
example_dequant_gemm_bf16_mxfp4_hopper 0.902297 0.892303 1.0112
example_mha_sink_fwd_bhsd_sliding_window 0.0357496 0.0352358 1.01458
example_mha_sink_bwd_bhsd 0.139779 0.13726 1.01835
example_mha_inference 0.130354 0.127624 1.02139
example_tilelang_gemm_fp8_2xAcc 0.150381 0.14657 1.026
example_fusedmoe_tilelang 0.254131 0.247668 1.02609
example_blocksparse_gemm 0.0156683 0.0149504 1.04802
example_mhc_pre 0.355273 0.332406 1.06879
example_gemv 0.490861 0.459119 1.06914
example_mha_sink_fwd_bhsd 0.0394602 0.035621 1.10778
example_group_per_split_token_cast_to_fp8 0.018573 0.0163693 1.13462
example_tilelang_sparse_gqa_decode_varlen_mask 0.0317203 0.02786 1.13856
example_dequant_gemv_fp16xint4 0.0860375 0.0748948 1.14878
example_tilelang_gemm_fp8 0.594601 0.515106 1.15433
example_mha_sink_bwd_bhsd_sliding_window 0.129632 0.110514 1.17298
example_mla_decode 0.942366 0.78403 1.20195
example_warp_specialize_gemm_softpipe_stage2 0.0585226 0.0470971 1.24259
example_convolution_autotune 1.85071 1.43686 1.28802
example_dequant_gemm_bf16_fp4_hopper 0.945499 0.709905 1.33187
example_mhc_post 0.287917 0.208685 1.37967
example_gqa_decode 0.0665255 0.0451674 1.47287
block_sparse_attn_tilelang 0.0192617 0.0121248 1.58861
example_tilelang_nsa_fwd 0.00910429 0.00552611 1.6475
example_linear_attn_fwd 0.0531891 0.0287661 1.84902
example_tilelang_block_sparse_attn 0.0147894 0.0074694 1.98
example_warp_specialize_gemm_copy_0_gemm_1 0.0599396 0.0281628 2.12833

Artifacts

  • regression_result.png (speedup plot) is attached as a workflow artifact. Download it from the workflow run page above.

@Yongqi-Zhuo
Yongqi-Zhuo force-pushed the advanced-tma-lowering branch from c7585e6 to 4394ac9 Compare June 17, 2026 19:18
@Yongqi-Zhuo Yongqi-Zhuo changed the title [WIP] [DO NOT REVIEW] [Backend] Support TMA lowering for arbitrary (swizzled) SMEM layout [Backend] Support TMA lowering for arbitrary (swizzled) SMEM layout Jun 17, 2026
@Yongqi-Zhuo
Yongqi-Zhuo marked this pull request as ready for review June 17, 2026 19:19

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: 4394ac97bf

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

Comment thread src/cuda/op/copy.cc Outdated
Comment on lines +1698 to +1701
if (cute::AsConst(tile_gbasis->shape[i]) > max_box_dim) {
DLOG(WARNING) << "The " << i
<< "-th box dim > 256, fallback to normal copy";
return fallback_to_normal("box dim > 256");

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P1 Badge Split oversized inner TMA boxes instead of falling back

When an unswizzled 2-D/ND T.tma_copy has a contiguous innermost tile larger than 256 elements, such as a 64x512 fp16 shared tile, tile_gbasis->shape[0] hits this branch and fallback_to_normal becomes a fatal error for explicit T.tma_copy. The previous lowering chunked the inner dimension to 256 and emitted a loop, which is valid because CUDA's limit is per boxDim, so this regresses valid TMA copies rather than just rejecting unsupported layouts.

Useful? React with 👍 / 👎.

Comment thread src/cuda/op/copy.cc
// tile.
int64_t inv_size = cute::AsConst(cute::Size(smem_plain_to_tile));
int64_t tile_size = cute::AsConst(cute::Size(tile_to_smem_plain));
ICHECK_EQ(inv_size, tile_size)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P2 Badge Fall back for gapped SMEM layouts instead of ICHECKing

For padded or gapped annotated shared layouts, for example a row-major tile with a padded row stride, RightInverse only covers the contiguous prefix, so this equality fails and aborts compilation. Earlier lowering explicitly fell back to normal copy for padded/unsupported layouts, and ordinary T.copy(..., prefer_instruction="tma") can still be lowered correctly without TMA, so this should use the existing fallback path rather than crashing the compile.

Useful? React with 👍 / 👎.

@Yongqi-Zhuo Yongqi-Zhuo reopened this Jun 17, 2026

@coderabbitai coderabbitai Bot left a comment

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.

Actionable comments posted: 2

🧹 Nitpick comments (2)
src/layout/cute_layout.cc (1)

34-38: 💤 Low value

Consider portability for __builtin_ctzll.

This is a GCC/Clang-specific builtin. If MSVC support is ever needed, a fallback would be required (e.g., _BitScanForward64). Given this is CUDA-focused code where GCC/Clang is standard, this is likely acceptable.

🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@src/layout/cute_layout.cc` around lines 34 - 38, The Log2Exact function uses
__builtin_ctzll, which is a GCC/Clang-specific builtin that is not supported by
MSVC. To improve portability, add conditional compilation directives around the
Log2Exact function implementation to use __builtin_ctzll for GCC/Clang and
provide an alternative implementation using _BitScanForward64 for MSVC, or
alternatively add a comment documenting that this function requires GCC/Clang
and explaining what would be needed for MSVC support if that becomes necessary
in the future.
src/cuda/op/copy.cc (1)

1391-1399: Remove unused makeColumnMajorStrides function.

The function is defined at lines 1391-1399 but never called anywhere in the codebase. This is dead code introduced in the PR.

🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@src/cuda/op/copy.cc` around lines 1391 - 1399, The function
`makeColumnMajorStrides` is defined but never called anywhere in the codebase,
making it dead code. Remove the entire function definition including the
function signature and its body. Search the codebase to confirm there are no
calls to this function before deletion.
🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

Inline comments:
In `@testing/python/layout/test_tilelang_cute.py`:
- Line 48: The function `forward` uses `vars` as a parameter name, which shadows
Python's built-in `vars()` function and is flagged by Ruff linting. Rename the
`vars` parameter in the `forward` function signature to a non-conflicting name
like `args` or `arguments`. Additionally, there is an ambiguous identifier `I`
(referenced at line 117) that should also be renamed to a more descriptive
identifier to avoid Ruff linting errors. These changes will resolve the CI
linting blockers.

In `@tilelang/cuda/intrinsics/layout/mma_layout.py`:
- Around line 234-240: The code at line 234 calculates swizzle_vectors by
dividing swizzle_bytes by 16, but it does not validate that swizzle_bytes
contains an allowed value before this division. Add validation logic before the
swizzle_vectors calculation to ensure swizzle_bytes is one of the supported
values (such as 32, 64, or 128 bytes) that map to valid swizzle vector counts of
2, 4, or 8. If an unsupported swizzle_bytes value is encountered, raise an
appropriate exception with a descriptive error message. This prevents division
by zero or invalid swizzle mappings in the subsequent calculations.

---

Nitpick comments:
In `@src/cuda/op/copy.cc`:
- Around line 1391-1399: The function `makeColumnMajorStrides` is defined but
never called anywhere in the codebase, making it dead code. Remove the entire
function definition including the function signature and its body. Search the
codebase to confirm there are no calls to this function before deletion.

In `@src/layout/cute_layout.cc`:
- Around line 34-38: The Log2Exact function uses __builtin_ctzll, which is a
GCC/Clang-specific builtin that is not supported by MSVC. To improve
portability, add conditional compilation directives around the Log2Exact
function implementation to use __builtin_ctzll for GCC/Clang and provide an
alternative implementation using _BitScanForward64 for MSVC, or alternatively
add a comment documenting that this function requires GCC/Clang and explaining
what would be needed for MSVC support if that becomes necessary in the future.
🪄 Autofix (Beta)

Fix all unresolved CodeRabbit comments on this PR:

  • Push a commit to this branch (recommended)
  • Create a new PR with the fixes

ℹ️ Review info
⚙️ Run configuration

Configuration used: defaults

Review profile: CHILL

Plan: Pro

Run ID: 3be2ebae-8119-48e9-841e-b9e855d81460

📥 Commits

Reviewing files that changed from the base of the PR and between d34109a and 4394ac9.

📒 Files selected for processing (11)
  • examples/gemm/example_gemm_intrinsics.py
  • src/cuda/op/copy.cc
  • src/cuda/transform/producer_consumer_ws.cc
  • src/layout/cute_layout.cc
  • src/layout/cute_layout.h
  • testing/python/kernel/test_tilelang_kernel_gemm_simt.py
  • testing/python/layout/test_tilelang_cute.py
  • tilelang/cuda/intrinsics/layout/mma_layout.py
  • tilelang/layout/__init__.py
  • tilelang/layout/_cute_ffi_api.py
  • tilelang/layout/cute.py

apply(x) = x ^ ((x & yyy) >> s_shift), with yyy = mask << (m_base + s_shift)
and mask = (1 << b_bits) - 1."""

def forward(*vars):

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.

⚠️ Potential issue | 🟡 Minor | ⚡ Quick win

Rename lint-blocking identifiers (vars, I).

forward(*vars) shadows a Python builtin, and I is ambiguous. Ruff flags both as errors, so this can block linted CI.

Suggested rename patch
-    def forward(*vars):
-        addr = intermediate_fn(*vars)
+    def forward(*coords):
+        addr = intermediate_fn(*coords)
@@
-    I = cute.make_identity_layout((2, (2, 2)))
-    assert I.shape == (2, (2, 2))
-    assert [I(c) for c in range(8)] == [(c % 2, (c // 2 % 2, c // 4)) for c in range(8)]
+    identity_layout = cute.make_identity_layout((2, (2, 2)))
+    assert identity_layout.shape == (2, (2, 2))
+    assert [identity_layout(c) for c in range(8)] == [(c % 2, (c // 2 % 2, c // 4)) for c in range(8)]

Also applies to: 117-117

🧰 Tools
🪛 Ruff (0.15.17)

[error] 48-48: Function argument vars is shadowing a Python builtin

(A002)

🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@testing/python/layout/test_tilelang_cute.py` at line 48, The function
`forward` uses `vars` as a parameter name, which shadows Python's built-in
`vars()` function and is flagged by Ruff linting. Rename the `vars` parameter in
the `forward` function signature to a non-conflicting name like `args` or
`arguments`. Additionally, there is an ambiguous identifier `I` (referenced at
line 117) that should also be renamed to a more descriptive identifier to avoid
Ruff linting errors. These changes will resolve the CI linting blockers.

Source: Linters/SAST tools

Comment on lines +234 to +240
swizzle_vectors = int(swizzle_bytes) // 16
col_idx_16B = col_idx // elem_per_16B
col_idx_in_16B = col_idx % elem_per_16B
new_col_idx_16B = col_idx_16B ^ (row_idx % (swizzle_bytes // 16))
return row_idx, ana.simplify(new_col_idx_16B * elem_per_16B + col_idx_in_16B)
col_tile = col_idx_16B // swizzle_vectors
c = col_idx_16B % swizzle_vectors
src = (row_idx % 8) // (8 // swizzle_vectors)
swizzled_col = (c ^ src) * elem_per_16B + col_idx_in_16B

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.

⚠️ Potential issue | 🟡 Minor | ⚡ Quick win

Validate swizzle_bytes domain before deriving swizzle_vectors.

Line 234/239 assume swizzle_bytes maps to valid vector counts (2/4/8), but unsupported values can crash (// 0) or produce invalid swizzle mappings. Add an explicit guard for allowed swizzle sizes and bounds.

Suggested patch
 def get_swizzle_layout(row_idx, col_idx, row_size, dtype: DataType | str, swizzle_bytes=None):
@@
     if swizzle_bytes is None:
         swizzle_bytes = min(128, row_bytes)
+    if swizzle_bytes not in (32, 64, 128):
+        raise ValueError("swizzle_bytes must be one of {32, 64, 128}.")
+    if swizzle_bytes > row_bytes:
+        raise ValueError("swizzle_bytes cannot exceed row size in bytes.")
@@
-    swizzle_vectors = int(swizzle_bytes) // 16
+    swizzle_vectors = swizzle_bytes // 16
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@tilelang/cuda/intrinsics/layout/mma_layout.py` around lines 234 - 240, The
code at line 234 calculates swizzle_vectors by dividing swizzle_bytes by 16, but
it does not validate that swizzle_bytes contains an allowed value before this
division. Add validation logic before the swizzle_vectors calculation to ensure
swizzle_bytes is one of the supported values (such as 32, 64, or 128 bytes) that
map to valid swizzle vector counts of 2, 4, or 8. If an unsupported
swizzle_bytes value is encountered, raise an appropriate exception with a
descriptive error message. This prevents division by zero or invalid swizzle
mappings in the subsequent calculations.

@Yongqi-Zhuo

Copy link
Copy Markdown
Collaborator Author

@regression-perf

@Yongqi-Zhuo
Yongqi-Zhuo force-pushed the advanced-tma-lowering branch from 4394ac9 to 3cc4648 Compare June 17, 2026 19:51
@Yongqi-Zhuo
Yongqi-Zhuo requested a review from SiriusNEO June 17, 2026 19:53

@coderabbitai coderabbitai Bot left a comment

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.

Actionable comments posted: 2

🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

Inline comments:
In `@src/layout/cute_layout.cc`:
- Around line 1458-1490: The code currently commits to the first contiguous run
found in s_at and returns early with std::nullopt if RecoverPlainLayout fails
for that candidate, but it should instead enumerate all candidate contiguous
runs and try each one. Refactor the swizzle candidate selection logic to collect
all valid contiguous runs of target bits that share an s_shift value, then
iterate through each candidate swizzle built from those runs (including
Swizzle::Identity as a fallback), testing each with RecoverPlainLayout and
returning the first one that succeeds. Only return std::nullopt if no candidate
passes validation, avoiding false rejections of valid affine layouts with
non-power-of-two strides.
- Around line 973-1029: The AddrProbe class performs address arithmetic with
narrowing from int64_t to int32_t without overflow checks, risking undefined
behavior on large constant layouts. In the AddrProbe constructor, accumulate
out_strides_ using int64_t (keep acc as int64_t throughout the loop), then add
explicit range checks to verify each accumulated stride value fits in the
int32_t range before storing it in out_strides_. Similarly, in the operator()
method, before performing the address calculation with addr +=
static_cast<int32_t>(*val) * out_strides_[d], add range checks to ensure both
val and the resulting multiplication fit within int32_t bounds, and verify the
accumulated addr doesn't overflow int32_t range before returning it.
🪄 Autofix (Beta)

Fix all unresolved CodeRabbit comments on this PR:

  • Push a commit to this branch (recommended)
  • Create a new PR with the fixes

ℹ️ Review info
⚙️ Run configuration

Configuration used: defaults

Review profile: CHILL

Plan: Pro

Run ID: 765628c6-de2e-4b4e-8a7f-82e095bba332

📥 Commits

Reviewing files that changed from the base of the PR and between 4394ac9 and 3cc4648.

📒 Files selected for processing (11)
  • examples/gemm/example_gemm_intrinsics.py
  • src/cuda/op/copy.cc
  • src/cuda/transform/producer_consumer_ws.cc
  • src/layout/cute_layout.cc
  • src/layout/cute_layout.h
  • testing/python/kernel/test_tilelang_kernel_gemm_simt.py
  • testing/python/layout/test_tilelang_cute.py
  • tilelang/cuda/intrinsics/layout/mma_layout.py
  • tilelang/layout/__init__.py
  • tilelang/layout/_cute_ffi_api.py
  • tilelang/layout/cute.py
✅ Files skipped from review due to trivial changes (1)
  • examples/gemm/example_gemm_intrinsics.py
🚧 Files skipped from review as they are similar to previous changes (8)
  • tilelang/layout/init.py
  • src/cuda/transform/producer_consumer_ws.cc
  • tilelang/layout/_cute_ffi_api.py
  • testing/python/kernel/test_tilelang_kernel_gemm_simt.py
  • tilelang/cuda/intrinsics/layout/mma_layout.py
  • src/cuda/op/copy.cc
  • src/layout/cute_layout.h
  • tilelang/layout/cute.py

Comment thread src/layout/cute_layout.cc
Comment on lines +973 to +1029
// When processing shared tensor layout, all the shapes and strides are int32_t.
// Also, we need to make sure not to mix with int64_t, because otherwise there
// would be a lot of conversion nodes that prevent the analyzer from cancelling
// out the terms.
PrimExpr MakeI32(int x) { return IntImm(DataType::Int(32), x); }

// Probes a TileLang layout's row-major linearized physical address A(x).
// The constant input shape, output strides, forward-index expressions, and
// input placeholders are parsed once in the constructor. operator() folds
// concrete integer coordinates to a constant address; Symbolic() substitutes
// fresh variables for the equivalence proof. valid() is false when any extent
// is non-constant.
class AddrProbe {
public:
explicit AddrProbe(const tvm::tl::Layout &layout) {
ICHECK(layout.defined());
for (const auto &e : layout->InputShape()) {
auto c = as_const_int(e);
ICHECK(c) << "InputShape extent " << e
<< " of a shared tensor layout must be constant";
shape_.push_back(*c);
}
std::vector<int32_t> out_sizes;
for (const auto &e : layout->OutputShape()) {
auto c = as_const_int(e);
ICHECK(c) << "OutputShape extent " << e
<< " of a shared tensor layout must be constant";
out_sizes.push_back(*c);
}
out_strides_.assign(out_sizes.size(), 0);
int32_t acc = 1;
for (int64_t d = static_cast<int64_t>(out_sizes.size()) - 1; d >= 0; --d) {
out_strides_[d] = acc;
acc *= out_sizes[d];
}
forward_index_ = layout->GetForwardIndex();
ICHECK_EQ(forward_index_.size(), out_strides_.size());
ICHECK_EQ(layout->InputDim(), shape_.size());
for (size_t k = 0; k < shape_.size(); ++k)
placeholders_.push_back(InputPlaceholder(k));
}

const std::vector<int32_t> &shape() const { return shape_; }

// Concrete address A(coords); nullopt if it does not fold to a constant.
std::optional<int32_t> operator()(const std::vector<int32_t> &coords) const {
Map<Var, PrimExpr> vmap;
for (size_t k = 0; k < coords.size(); ++k)
vmap.Set(placeholders_[k], MakeI32(coords[k]));
int32_t addr = 0;
for (size_t d = 0; d < forward_index_.size(); ++d) {
PrimExpr e = Substitute(forward_index_[d], vmap);
std::optional<int64_t> val = EvalConstExpr(e);
if (!val)
return std::nullopt;
addr += static_cast<int32_t>(*val) * out_strides_[d];
}

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.

⚠️ Potential issue | 🟠 Major | ⚡ Quick win

Reject out-of-range address arithmetic before narrowing to int32_t.

as_const_int yields int64_t, but extents, output strides, addr, and MakeI32(offset) are narrowed or accumulated as signed int32_t. Large-but-constant layouts can overflow before recovery returns None, which risks UB or proving a wrapped address map. Use checked int64_t accumulation and explicit range checks before creating int32 expressions.

Guard the narrowing points
-PrimExpr MakeI32(int x) { return IntImm(DataType::Int(32), x); }
+int32_t CheckedI32(int64_t value, const char *what) {
+  ICHECK(value >= std::numeric_limits<int32_t>::min() &&
+         value <= std::numeric_limits<int32_t>::max())
+      << what << " is outside int32 range: " << value;
+  return static_cast<int32_t>(value);
+}
+
+PrimExpr MakeI32(int64_t x) {
+  return IntImm(DataType::Int(32), CheckedI32(x, "int32 PrimExpr constant"));
+}
@@
-      shape_.push_back(*c);
+      shape_.push_back(CheckedI32(*c, "InputShape extent"));
@@
-      out_sizes.push_back(*c);
+      out_sizes.push_back(CheckedI32(*c, "OutputShape extent"));
@@
-    int32_t acc = 1;
+    int64_t acc = 1;
     for (int64_t d = static_cast<int64_t>(out_sizes.size()) - 1; d >= 0; --d) {
-      out_strides_[d] = acc;
+      out_strides_[d] = CheckedI32(acc, "OutputShape stride");
+      ICHECK(out_sizes[d] > 0) << "OutputShape extent must be positive";
+      ICHECK_LE(acc, std::numeric_limits<int32_t>::max() / out_sizes[d])
+          << "OutputShape product exceeds int32 range";
       acc *= out_sizes[d];
     }
@@
-    int32_t addr = 0;
+    int64_t addr = 0;
@@
-      addr += static_cast<int32_t>(*val) * out_strides_[d];
+      addr += *val * static_cast<int64_t>(out_strides_[d]);
@@
-    return addr;
+    return CheckedI32(addr, "linearized address");

Also applies to: 1314-1320

🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@src/layout/cute_layout.cc` around lines 973 - 1029, The AddrProbe class
performs address arithmetic with narrowing from int64_t to int32_t without
overflow checks, risking undefined behavior on large constant layouts. In the
AddrProbe constructor, accumulate out_strides_ using int64_t (keep acc as
int64_t throughout the loop), then add explicit range checks to verify each
accumulated stride value fits in the int32_t range before storing it in
out_strides_. Similarly, in the operator() method, before performing the address
calculation with addr += static_cast<int32_t>(*val) * out_strides_[d], add range
checks to ensure both val and the resulting multiplication fit within int32_t
bounds, and verify the accumulated addr doesn't overflow int32_t range before
returning it.

Comment thread src/layout/cute_layout.cc
Comment on lines +1458 to +1490
std::map<int, int> s_at; // swizzle target bit -> s_shift.
for (uint64_t col : cols) {
if (__builtin_popcountll(col) != 2)
continue;
int lo = __builtin_ctzll(col);
int hi = 63 - __builtin_clzll(col);
if (!(weight1 & (uint64_t(1) << lo)))
continue; // two-bit column from a plain stride, not a swizzle.
if (s_at.count(lo))
return std::nullopt; // two sources hitting one target: not a Swizzle.
s_at.emplace(lo, hi - lo);
}
Swizzle swizzle = Swizzle::Identity();
if (!s_at.empty()) {
// Take the lowest contiguous run of target bits that share one s_shift.
int m_base = s_at.begin()->first;
int s_shift = s_at.begin()->second;
int b_bits = 0;
for (auto it = s_at.find(m_base + b_bits);
it != s_at.end() && it->second == s_shift;
it = s_at.find(m_base + b_bits))
++b_bits;
if (s_shift < b_bits)
return std::nullopt; // source and target bit regions must not overlap.
swizzle = Swizzle(b_bits, m_base, s_shift);
}

// The base offset is Sw(A(0)), since Sw is an involution; recover the plain
// layout under that swizzle and offset.
int64_t offset = swizzle->Apply(*A0);
Optional<Layout> plain = RecoverPlainLayout(A, swizzle, offset);
if (!plain.defined())
return std::nullopt;

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.

⚠️ Potential issue | 🟠 Major | ⚡ Quick win

Try proven swizzle candidates instead of committing to the lowest two-bit column.

The detector can false-reject valid affine layouts: an ordinary stride like 3 produces a two-bit column {0,1}, and if another mode has stride 1, weight1 makes it look like Sw<1,0,1>. The selected swizzle then fails RecoverPlainLayout, but the code returns None instead of falling back to identity or trying later real swizzle runs. Avoid returning before proof on duplicate/ambiguous candidates, enumerate candidate contiguous runs, and return the first one that RecoverPlainLayout proves; try identity as the fallback.

Suggested regression shape
def test_plain_nonpow2_stride_does_not_force_swizzle():
    mode = cute.ComposedLayout.from_tilelang(Layout((2, 2), lambda i, j: i * 3 + j))
    assert mode is not None
    assert not mode.swizzle.is_swizzled
    _assert_struct(mode.layout, (2, 2), (3, 1))
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@src/layout/cute_layout.cc` around lines 1458 - 1490, The code currently
commits to the first contiguous run found in s_at and returns early with
std::nullopt if RecoverPlainLayout fails for that candidate, but it should
instead enumerate all candidate contiguous runs and try each one. Refactor the
swizzle candidate selection logic to collect all valid contiguous runs of target
bits that share an s_shift value, then iterate through each candidate swizzle
built from those runs (including Swizzle::Identity as a fallback), testing each
with RecoverPlainLayout and returning the first one that succeeds. Only return
std::nullopt if no candidate passes validation, avoiding false rejections of
valid affine layouts with non-power-of-two strides.

@github-actions

Copy link
Copy Markdown

Performance Regression Test Report

Triggered by: @Yongqi-Zhuo
Workflow run: https://github.com/tile-ai/tilelang/actions/runs/27714594863

Results

File Original Latency Current Latency Speedup
example_tilelang_sparse_gqa_decode_varlen_indice 0.0192836 0.0450475 0.428073
example_group_per_split_token_cast_to_fp8 0.0133991 0.022114 0.605912
example_tilelang_sparse_gqa_decode_varlen_mask 0.0241335 0.0394861 0.611191
example_dequant_gemm_bf16_fp4_hopper 0.914678 1.2378 0.738957
example_mha_sink_fwd_bhsd 0.023708 0.0317614 0.74644
example_mhc_post 0.236215 0.287422 0.821839
example_dequant_gemm_fp4_hopper 1.46776 1.7582 0.834806
example_convolution_autotune 1.4812 1.77295 0.835444
topk_selector 0.0785952 0.0925514 0.849206
example_gqa_sink_bwd_bhsd_sliding_window 0.0292409 0.033421 0.874927
example_warp_specialize_gemm_softpipe_stage2 0.0352908 0.0392877 0.898266
block_sparse_attn_tilelang 0.0158625 0.0174264 0.910255
example_warp_specialize_gemm_copy_1_gemm_0 0.0510575 0.0548575 0.93073
example_mhc_pre 0.306953 0.329438 0.931747
example_tilelang_gemm_splitk_vectorize_atomicadd 1.72688 1.84704 0.934948
example_dequant_gemm_bf16_mxfp4_hopper 1.00604 1.07161 0.938812
example_dequant_gemm_w4a8 3.56046 3.74301 0.951228
fp8_lighting_indexer 0.0253884 0.0265972 0.954551
example_warp_specialize_gemm_barrierpipe_stage2 0.0333419 0.034619 0.963109
example_tilelang_gemm_splitk 1.76805 1.82879 0.966786
example_blocksparse_gemm 0.0150796 0.0155585 0.969223
example_tilelang_gemm_fp8_2xAcc 0.145063 0.147047 0.986507
sparse_mla_fwd 0.0851268 0.0860843 0.988877
example_dynamic 1.30489 1.31885 0.989419
sparse_mla_fwd_pipelined 0.0640339 0.0645908 0.991379
example_elementwise_add 0.295151 0.295595 0.998499
example_mla_decode 0.825546 0.826001 0.99945
example_fusedmoe_tilelang 0.265418 0.264802 1.00233
example_tilelang_gemm_fp8 0.588695 0.584601 1.007
example_tilelang_nsa_decode 0.00562262 0.00557745 1.0081
example_gemv 0.469638 0.464899 1.01019
example_mha_inference 0.135681 0.132986 1.02027
sparse_mla_bwd 0.570878 0.536947 1.06319
example_linear_attn_bwd 0.300264 0.275234 1.09094
example_mha_sink_fwd_bhsd_sliding_window 0.0316793 0.0282601 1.12099
example_dequant_gemv_fp16xint4 0.0777155 0.0693151 1.12119
example_warp_specialize_gemm_copy_0_gemm_1 0.0765051 0.0666473 1.14791
example_gqa_sink_bwd_bhsd 0.0693961 0.0596764 1.16287
example_convolution 2.52454 2.04307 1.23566
example_vertical_slash_sparse_attn 0.403507 0.325178 1.24088
example_mha_sink_bwd_bhsd 0.217627 0.16727 1.30105
example_topk 36.6448 27.7743 1.31938
example_linear_attn_fwd 0.0409997 0.028665 1.4303
example_tilelang_nsa_fwd 0.00899986 0.00553928 1.62473
example_tilelang_block_sparse_attn 0.0203523 0.0111585 1.82392
example_per_token_cast_to_fp8 0.0134016 0.00661986 2.02446
example_mha_sink_bwd_bhsd_sliding_window 0.162281 0.0744829 2.17877
example_gqa_decode 0.134068 0.0441035 3.03984

Artifacts

  • regression_result.png (speedup plot) is attached as a workflow artifact. Download it from the workflow run page above.

@Yongqi-Zhuo
Yongqi-Zhuo force-pushed the advanced-tma-lowering branch from 3cc4648 to 847eeda Compare June 18, 2026 05:45
@Yongqi-Zhuo

Copy link
Copy Markdown
Collaborator Author

@regression-perf

@github-actions

Copy link
Copy Markdown

Performance Regression Test Report

Triggered by: @Yongqi-Zhuo
Workflow run: https://github.com/tile-ai/tilelang/actions/runs/27739471180

Results

File Original Latency Current Latency Speedup
sparse_mla_bwd 0.22824 0.233164 0.978883
example_tilelang_gemm_fp8_2xAcc 0.0783645 0.0794705 0.986083
example_dequant_gemm_bf16_mxfp4_hopper 0.35884 0.361732 0.992005
example_mha_sink_fwd_bhsd 0.012648 0.0127443 0.99244
example_convolution_autotune 0.735308 0.739961 0.993711
example_dequant_gemm_bf16_fp4_hopper 0.396193 0.39869 0.993737
example_linear_attn_bwd 0.118113 0.118708 0.994991
example_mha_sink_bwd_bhsd 0.0520299 0.0522267 0.996233
example_mha_sink_bwd_bhsd_sliding_window 0.0386585 0.0387507 0.997621
example_vertical_slash_sparse_attn 0.165917 0.166224 0.998156
example_warp_specialize_gemm_copy_0_gemm_1 0.027316 0.0273655 0.998192
example_tilelang_gemm_fp8 0.243572 0.243963 0.998396
example_gqa_bwd_tma_reduce_varlen 0.0333119 0.0333647 0.998417
example_mha_bwd_bhsd 0.0313684 0.0314029 0.998901
example_mhc_post 0.106674 0.106778 0.99903
example_tilelang_gemm_splitk 0.76936 0.769752 0.99949
example_gqa_sink_bwd_bhsd 0.0295937 0.0296028 0.999694
example_tilelang_sparse_gqa_decode_varlen_mask 0.0128017 0.0128044 0.999786
example_dequant_gemm_w4a8 1.98006 1.98017 0.999941
example_linear_attn_fwd 0.0283841 0.0283812 1.0001
example_mla_decode 0.316292 0.316256 1.00011
example_tilelang_sparse_gqa_decode_varlen_indice 0.0117631 0.0117617 1.00012
example_fusedmoe_tilelang 0.09528 0.0952657 1.00015
example_dequant_gemv_fp16xint4 0.0270862 0.027082 1.00016
example_mha_sink_fwd_bhsd_sliding_window 0.0125649 0.0125618 1.00025
example_convolution 0.92801 0.927757 1.00027
example_per_token_cast_to_fp8 0.00652299 0.00652 1.00046
block_sparse_attn_tilelang 0.00685534 0.0068521 1.00047
example_group_per_split_token_cast_to_fp8 0.0076316 0.00762701 1.0006
example_gemv 0.202654 0.202494 1.00079
example_warp_specialize_gemm_softpipe_stage2 0.0195716 0.0195539 1.00091
example_elementwise_add 0.113162 0.113057 1.00093
topk_selector 0.0419172 0.0418745 1.00102
example_dynamic 0.520535 0.519986 1.00106
example_tilelang_block_sparse_attn 0.00724402 0.00723344 1.00146
example_mha_fwd_bshd 0.0189277 0.0188998 1.00148
example_warp_specialize_gemm_barrierpipe_stage2 0.029021 0.0289654 1.00192
sparse_mla_fwd_pipelined 0.0592722 0.0591409 1.00222
sparse_mla_fwd 0.0818701 0.0816393 1.00283
example_gqa_sink_bwd_bhsd_sliding_window 0.0179902 0.0179384 1.00289
example_tilelang_nsa_decode 0.0055688 0.00555026 1.00334
example_gqa_decode 0.0415049 0.0413639 1.00341
example_tilelang_gemm_splitk_vectorize_atomicadd 0.793025 0.790326 1.00341
example_mha_fwd_bhsd 0.00918839 0.00915648 1.00348
example_mha_bwd_bshd 0.0311657 0.0310516 1.00367
fp8_lighting_indexer 0.0231092 0.0230196 1.00389
example_warp_specialize_gemm_copy_1_gemm_0 0.0195893 0.0194864 1.00528
example_gqa_fwd_bshd 0.0442435 0.0439955 1.00564
example_blocksparse_gemm 0.0136922 0.0136056 1.00637
example_tilelang_nsa_fwd 0.00543813 0.00540299 1.0065
example_gqa_bwd 0.0329079 0.0326927 1.00658
example_mha_inference 0.0634803 0.0630091 1.00748
example_mha_fwd_varlen 0.0323634 0.0320749 1.00899
example_mhc_pre 0.147143 0.144797 1.0162
example_dequant_gemm_fp4_hopper 0.734945 0.720062 1.02067
example_topk 30.7531 28.1789 1.09135

Artifacts

  • regression_result.png (speedup plot) is attached as a workflow artifact. Download it from the workflow run page above.

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.

2 participants