[Feature] Support TMA store in T.tma_copy() - #1981
Conversation
…nization T.tma_copy() previously only supported TMA loads (global -> shared) with user-managed barrier synchronization. This extends it to also support TMA stores (shared -> global), where: - The barrier parameter is now optional (not needed for stores) - For stores, only tma_store + tma_store_arrive is emitted (no wait) - Users call T.tma_store_wait() explicitly to synchronize
|
👋 Hi! Thank you for contributing to the TileLang project. Please remember to run We appreciate you taking this step! Our team will review your contribution, and we look forward to your awesome work! 🚀 |
📝 WalkthroughWalkthroughLowering and API changes make TMA store synchronization explicit: Changes
Estimated code review effort🎯 4 (Complex) | ⏱️ ~60 minutes Possibly related PRs
Suggested reviewers
Poem
🚥 Pre-merge checks | ✅ 2 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (2 passed)
✏️ Tip: You can configure your own custom pre-merge checks in the settings. ✨ Finishing Touches📝 Generate docstrings
🧪 Generate unit tests (beta)
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. Comment |
There was a problem hiding this comment.
Actionable comments posted: 1
🧹 Nitpick comments (1)
testing/python/language/test_tilelang_language_tma_store.py (1)
83-87: Tighten the source assertion so this test catches accidental auto-waits.Because Line 50 already adds an explicit
T.tma_store_wait(), this test still passes if lowering regresses and inserts a second implicit wait. Checking onlytma_store_arrivedoesn't pin the new contract; please also assert that the emitted source contains only the single user-authored store-wait invocation.🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@testing/python/language/test_tilelang_language_tma_store.py` around lines 83 - 87, The current test only asserts that "tma_store_arrive" appears but doesn't prevent an extra implicit wait from being generated; update the assertions on kernel_source (the string returned by kernel.get_kernel_source()) to also verify that "tma_store_wait" appears exactly once (e.g., assert kernel_source.count("tma_store_wait") == 1) so the emitted source contains only the single user-authored T.tma_store_wait invocation and no additional implicit waits.
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Inline comments:
In `@tilelang/language/copy_op.py`:
- Around line 218-223: The current logic copies annotations into ann and then
unconditionally sets ann["barrier"] from the barrier param, which overwrites any
caller-provided annotations["barrier"]; change the behavior so that after
creating ann = annotations.copy() if annotations else {}, you only call
_mbar_to_buffer_load(barrier) and assign ann["barrier"] when barrier is not None
AND "barrier" is not already present in ann (i.e., preserve annotations'
precedence); refer to the local variables/parameters annotations, ann, barrier
and the helper _mbar_to_buffer_load to locate and implement this conditional
assignment.
---
Nitpick comments:
In `@testing/python/language/test_tilelang_language_tma_store.py`:
- Around line 83-87: The current test only asserts that "tma_store_arrive"
appears but doesn't prevent an extra implicit wait from being generated; update
the assertions on kernel_source (the string returned by
kernel.get_kernel_source()) to also verify that "tma_store_wait" appears exactly
once (e.g., assert kernel_source.count("tma_store_wait") == 1) so the emitted
source contains only the single user-authored T.tma_store_wait invocation and no
additional implicit waits.
🪄 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: 6a4acf1c-053d-49d8-8892-75a0e2889f35
📒 Files selected for processing (4)
src/op/copy.cctesting/python/language/test_tilelang_language_tma_copy.pytesting/python/language/test_tilelang_language_tma_store.pytilelang/language/copy_op.py
…signatures - Wire tma_store_wait count parameter through the full stack (Python -> C++ builtin registration -> CUDA/CuTeDSL codegen) - Replace vague *args with proper signatures for fence_proxy_async() and tma_store_arrive() (both take no arguments) - Improve docstrings for internal TMA APIs (create_tma_descriptor, tma_load, tma_load_2sm)
There was a problem hiding this comment.
🧹 Nitpick comments (2)
src/target/codegen_cutedsl.cc (1)
641-643: Add the sametma_store_waitarg guards in CuTeDSL codegen.This code path also assumes
op->args[0]is present andIntImm. Guarding count/type keeps diagnostics consistent across backends.Suggested hardening
} else if (op->op.same_as(tl::tma_store_wait())) { - int count = Downcast<IntImm>(op->args[0])->value; + ICHECK_EQ(op->args.size(), 1) << "tma_store_wait expects 1 argument"; + const auto* count_imm = op->args[0].as<IntImmNode>(); + ICHECK(count_imm) << "tma_store_wait expects IntImm count"; + int count = count_imm->value; PrintIndent(); stream << "tl.tma_store_wait(" << count << ")\n";🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@src/target/codegen_cutedsl.cc` around lines 641 - 643, The tma_store_wait emission assumes op->args[0] exists and is an IntImm; update the code around the Downcast<IntImm>(op->args[0])->value usage in codegen_cutedsl.cc to first guard that op->args is non-empty and that op->args[0] is an IntImm (e.g., using a safe type check instead of direct Downcast), and if the guard fails emit a clear diagnostic or fall back to a safe behavior before calling PrintIndent() and writing "tl.tma_store_wait(...)" so diagnostics and behavior match other backends.src/target/codegen_cuda.cc (1)
2037-2039: Validatetma_store_waitarg count/type before use.This path reads
op->args[0]unguarded. Adding explicit checks makes failures deterministic and easier to debug when malformed IR reaches codegen.Suggested hardening
} else if (op->op.same_as(tl::tma_store_wait())) { - int count = Downcast<IntImm>(op->args[0])->value; + ICHECK_EQ(op->args.size(), 1) << "tma_store_wait expects 1 argument"; + const auto* count_imm = op->args[0].as<IntImmNode>(); + ICHECK(count_imm) << "tma_store_wait expects IntImm count"; + int count = count_imm->value; this->PrintIndent(); this->stream << "tl::tma_store_wait<" << count << ">();\n";🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@src/target/codegen_cuda.cc` around lines 2037 - 2039, The code is accessing op->args[0] without validation before Downcast<IntImm>, so add explicit guards that op->args has at least one element and that op->args[0] is an IntImm before using it in the emitted tl::tma_store_wait template; if the checks fail, emit a clear diagnostic (e.g., an ICHECK/LOG + throw or return an error) mentioning the operator and that tma_store_wait's count must be a constant IntImm. Update the block around the Downcast<IntImm>(op->args[0])->value and its emission of "tl::tma_store_wait<...>()" to perform these validations and produce a deterministic, informative failure path instead of undefined behavior.
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Nitpick comments:
In `@src/target/codegen_cuda.cc`:
- Around line 2037-2039: The code is accessing op->args[0] without validation
before Downcast<IntImm>, so add explicit guards that op->args has at least one
element and that op->args[0] is an IntImm before using it in the emitted
tl::tma_store_wait template; if the checks fail, emit a clear diagnostic (e.g.,
an ICHECK/LOG + throw or return an error) mentioning the operator and that
tma_store_wait's count must be a constant IntImm. Update the block around the
Downcast<IntImm>(op->args[0])->value and its emission of
"tl::tma_store_wait<...>()" to perform these validations and produce a
deterministic, informative failure path instead of undefined behavior.
In `@src/target/codegen_cutedsl.cc`:
- Around line 641-643: The tma_store_wait emission assumes op->args[0] exists
and is an IntImm; update the code around the
Downcast<IntImm>(op->args[0])->value usage in codegen_cutedsl.cc to first guard
that op->args is non-empty and that op->args[0] is an IntImm (e.g., using a safe
type check instead of direct Downcast), and if the guard fails emit a clear
diagnostic or fall back to a safe behavior before calling PrintIndent() and
writing "tl.tma_store_wait(...)" so diagnostics and behavior match other
backends.
ℹ️ Review info
⚙️ Run configuration
Configuration used: defaults
Review profile: CHILL
Plan: Pro
Run ID: da32e037-1665-48bc-b2bc-7fd97afaba04
📒 Files selected for processing (6)
src/op/atomic_add.ccsrc/op/builtin.ccsrc/op/copy.ccsrc/target/codegen_cuda.ccsrc/target/codegen_cutedsl.cctilelang/language/builtin.py
Summary
T.tma_copy()to support TMA stores (shared -> global) with user-managed synchronizationbarrierparameter optional — it is required for TMA loads but not needed for storestma_store + tma_store_arrive(notma_store_wait), so users can batch multiple stores and callT.tma_store_wait()explicitlyTest plan
T.tma_copy()load tests pass (test_tma_copy_pipeline_2/3_stages)T.tma_copy()store tests intest_tilelang_language_tma_copy.pypass (test_tma_copy_store_pipeline_2/3_stages)test_tilelang_language_tma_store.pypass (test_tma_store_2/3_stages)assert_allcloseagainst torch.matmul referencetma_store_arriveSummary by CodeRabbit
New Features
barrieroptional for TMA copy; barriers required for TMA loads but not for TMA stores.Bug Fixes
Documentation
Tests