Skip to content

[Feature] Support TMA store in T.tma_copy() - #1981

Merged
LeiWang1999 merged 2 commits into
mainfrom
feature/tma-copy-store-support
Mar 27, 2026
Merged

LeiWang1999 merged 2 commits into
mainfrom
feature/tma-copy-store-support

Conversation

@LeiWang1999

@LeiWang1999 LeiWang1999 commented Mar 27, 2026 •

Copy link
Copy Markdown
Member

Summary

  • Extend T.tma_copy() to support TMA stores (shared -> global) with user-managed synchronization
  • Make the barrier parameter optional — it is required for TMA loads but not needed for stores
  • For stores, emit only tma_store + tma_store_arrive (no tma_store_wait), so users can batch multiple stores and call T.tma_store_wait() explicitly

Test plan

  • Existing T.tma_copy() load tests pass (test_tma_copy_pipeline_2/3_stages)
  • New T.tma_copy() store tests in test_tilelang_language_tma_copy.py pass (test_tma_copy_store_pipeline_2/3_stages)
  • Dedicated TMA store unit tests in test_tilelang_language_tma_store.py pass (test_tma_store_2/3_stages)
  • Correctness verified via assert_allclose against torch.matmul reference
  • Generated kernel source verified to contain tma_store_arrive

Summary by CodeRabbit

  • New Features

    • Made barrier optional for TMA copy; barriers required for TMA loads but not for TMA stores.
    • Added pipelined TMA store-capable GEMM kernels and builders.
  • Bug Fixes

    • Store-side synchronization now emits wait only when appropriate; store wait accepts an explicit integer argument.
  • Documentation

    • Clarified synchronization semantics for TMA loads vs stores.
  • Tests

    • Added CUDA-gated tests for 2- and 3-stage TMA store pipelines.

…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
@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 Mar 27, 2026 •

Copy link
Copy Markdown
Contributor
📝 Walkthrough

Walkthrough

Lowering and API changes make TMA store synchronization explicit: tma_copy() (stores) now emits only tma_store_arrive() and omits the implicit wait; copy() retains arrive+wait. tma_store_wait() now takes an explicit int argument, and tma_copy(..., barrier=None) makes barrier optional. Tests for store-capable TMA pipelines were added.

Changes

Cohort / File(s) Summary
TMA store lowering
src/op/copy.cc
Conditional emission: emit tma_store_wait() only when lowering non-TMA-copy paths; tma_copy() lowering now omits the wait.
Atomic TMA store lowering
src/op/atomic_add.cc
tma_store_wait() lowering updated to pass explicit IntImm(..., 0) argument.
Builtins / codegen for tma_store_wait
src/op/builtin.cc, src/target/codegen_cuda.cc, src/target/codegen_cutedsl.cc
Builtin signature changed to accept one input; CUDA/CuTeDSL codegen now reads op->args[0] and emits tma_store_wait with that immediate as the compile-time/count argument.
TileLang copy API
tilelang/language/copy_op.py
tma_copy barrier parameter made optional (barrier=None) and now attaches barrier annotation only if provided; docstring updated to reflect different sync rules for loads vs stores.
TileLang builtins wrappers
tilelang/language/builtin.py
fence_proxy_async, tma_store_arrive now take no args; tma_store_wait signature changed to count: int = 0 and forwards that exact count.
Tests — TMA copy/store pipelines
testing/python/language/test_tilelang_language_tma_copy.py, testing/python/language/test_tilelang_language_tma_store.py
Added store-capable GEMM builders and tests (2-stage and 3-stage) exercising explicit T.tma_copy store path followed by T.tma_store_wait() and validating correctness.

Estimated code review effort

🎯 4 (Complex) | ⏱️ ~60 minutes

Possibly related PRs

Suggested reviewers

  • cherichy

Poem

🐰
A hop, a copy, arrive in stride,
The wait now waits where users decide.
Stores send word, but leave the rest,
Pipelines hum, the kernels rest.
🥕✨

🚥 Pre-merge checks | ✅ 2 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 53.85% which is insufficient. The required threshold is 80.00%. Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (2 passed)
Check name Status Explanation
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.
Title check ✅ Passed The title '[Feature] Support TMA store in T.tma_copy()' directly and clearly describes the main feature being added: extending T.tma_copy() to support TMA store operations. It accurately captures the primary objective of the changeset.

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

✨ Finishing Touches
📝 Generate docstrings
  • Create stacked PR
  • Commit on current branch
🧪 Generate unit tests (beta)
  • Create PR with unit tests
  • Commit unit tests in branch feature/tma-copy-store-support

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.

@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: 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 only tma_store_arrive doesn'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

📥 Commits

Reviewing files that changed from the base of the PR and between dc60ab8 and 7f9a770.

📒 Files selected for processing (4)
  • src/op/copy.cc
  • testing/python/language/test_tilelang_language_tma_copy.py
  • testing/python/language/test_tilelang_language_tma_store.py
  • tilelang/language/copy_op.py

Comment thread tilelang/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)

@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.

🧹 Nitpick comments (2)
src/target/codegen_cutedsl.cc (1)

641-643: Add the same tma_store_wait arg guards in CuTeDSL codegen.

This code path also assumes op->args[0] is present and IntImm. 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: Validate tma_store_wait arg 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

📥 Commits

Reviewing files that changed from the base of the PR and between 7f9a770 and ea6af5b.

📒 Files selected for processing (6)
  • src/op/atomic_add.cc
  • src/op/builtin.cc
  • src/op/copy.cc
  • src/target/codegen_cuda.cc
  • src/target/codegen_cutedsl.cc
  • tilelang/language/builtin.py

@LeiWang1999
LeiWang1999 merged commit bdf436d into main Mar 27, 2026
5 of 6 checks passed
@LeiWang1999
LeiWang1999 deleted the feature/tma-copy-store-support branch April 14, 2026 06:06
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.

1 participant