Skip to content

Commit af58222

Browse files
committed
refactor to expose use_2cta option in T.tcgen05_gemm and remove 2cta pass config
1 parent d174892 commit af58222

11 files changed

Lines changed: 57 additions & 63 deletions

File tree

‎examples/gemm_sm100/gemm_tcgen5mma_ws.py‎

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66
from tilelang.profiler import do_bench
77

88

9-
@tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_DISABLE_2CTA_TCGEN5MMA: True})
9+
@tilelang.jit
1010
def gemm(A, B, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages, use_tma_store=True):
1111
M, N, K = T.const("M, N, K")
1212

@@ -99,8 +99,16 @@ def gemm_2cta(A, B, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype,
9999
if tx < 32: # warp 0: issue tma
100100
for k in T.serial(k_iters):
101101
T.mbarrier_wait_parity(consumed[k % num_stages], ((k // num_stages) & 1) ^ 1)
102-
T.tma_copy(A[by * block_M : (by + 1) * block_M, k * block_K : (k + 1) * block_K], A_shared[k % num_stages, :, :], barrier=loaded[k % num_stages])
103-
T.tma_copy(B[k * block_K : (k + 1) * block_K, (bx * 2 + cta_id) * (block_N // 2) : (bx * 2 + cta_id + 1) * (block_N // 2)], B_shared[k % num_stages, :, :], barrier=loaded[k % num_stages])
102+
T.tma_copy(
103+
A[by * block_M : (by + 1) * block_M, k * block_K : (k + 1) * block_K],
104+
A_shared[k % num_stages, :, :],
105+
barrier=loaded[k % num_stages],
106+
)
107+
T.tma_copy(
108+
B[k * block_K : (k + 1) * block_K, (bx * 2 + cta_id) * (block_N // 2) : (bx * 2 + cta_id + 1) * (block_N // 2)],
109+
B_shared[k % num_stages, :, :],
110+
barrier=loaded[k % num_stages],
111+
)
104112
T.mbarrier_arrive(loaded[k % num_stages], 0) # arrive on leader cta's barrier
105113
elif cta_id == 0 and tx < 64: # Only warp 1 on leader cta issues tcgen5
106114
for k in T.serial(k_iters):
@@ -111,6 +119,7 @@ def gemm_2cta(A, B, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype,
111119
C_tmem,
112120
mbar=consumed[k % num_stages],
113121
clear_accum=k == 0,
122+
use_2cta=True,
114123
)
115124
T.tcgen05_mma_arrive(tmem_full, arrive_2cta=True)
116125

‎examples/gemm_sm100/gemm_tcgen5mma_ws_persistent.py‎

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
from tilelang.profiler import do_bench
88

99

10-
@tilelang.jit(pass_configs={tilelang.PassConfigKey.TL_DISABLE_2CTA_TCGEN5MMA: True})
10+
@tilelang.jit
1111
def gemm_persistent(
1212
A,
1313
B,
@@ -62,7 +62,8 @@ def gemm_persistent(
6262
phase = w * k_blocks + k
6363
T.mbarrier_wait_parity(consumed[phase % num_stages], ((phase // num_stages) & 1) ^ 1)
6464
T.tma_copy(
65-
A[bx * block_M : (bx + 1) * block_M, k * block_K : (k + 1) * block_K], A_shared[phase % num_stages, :, :],
65+
A[bx * block_M : (bx + 1) * block_M, k * block_K : (k + 1) * block_K],
66+
A_shared[phase % num_stages, :, :],
6667
barrier=loaded[phase % num_stages],
6768
)
6869
T.tma_copy(
@@ -192,7 +193,8 @@ def gemm_persistent_2cta(
192193
phase = w * k_blocks + k
193194
T.mbarrier_wait_parity(consumed[phase % num_stages], ((phase // num_stages) & 1) ^ 1)
194195
T.tma_copy(
195-
A[bx * block_M : (bx + 1) * block_M, k * block_K : (k + 1) * block_K], A_shared[phase % num_stages, :, :],
196+
A[bx * block_M : (bx + 1) * block_M, k * block_K : (k + 1) * block_K],
197+
A_shared[phase % num_stages, :, :],
196198
barrier=loaded[phase % num_stages],
197199
)
198200

@@ -224,6 +226,7 @@ def gemm_persistent_2cta(
224226
C_tmem_0,
225227
mbar=consumed[phase % num_stages],
226228
clear_accum=k == 0,
229+
use_2cta=True,
227230
)
228231
else:
229232
T.tcgen05_gemm(
@@ -232,6 +235,7 @@ def gemm_persistent_2cta(
232235
C_tmem_1,
233236
mbar=consumed[phase % num_stages],
234237
clear_accum=k == 0,
238+
use_2cta=True,
235239
)
236240
T.tcgen05_mma_arrive(tmem_full[w & 1], arrive_2cta=True)
237241

‎src/op/builtin.cc‎

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -32,7 +32,6 @@ TVM_REGISTER_PASS_CONFIG_OPTION(kDisableVectorize256, Bool);
3232
TVM_REGISTER_PASS_CONFIG_OPTION(kEnableAsyncCopy, Bool);
3333
TVM_REGISTER_PASS_CONFIG_OPTION(kEnableVectorizePlannerVerbose, Bool);
3434
TVM_REGISTER_PASS_CONFIG_OPTION(kDisableWGMMA, Bool);
35-
TVM_REGISTER_PASS_CONFIG_OPTION(kDisable2CTATcgen5MMA, Bool);
3635
TVM_REGISTER_PASS_CONFIG_OPTION(kDisableShuffleElect, Bool);
3736
TVM_REGISTER_PASS_CONFIG_OPTION(kStorageRewriteDetectInplace, Bool);
3837
TVM_REGISTER_PASS_CONFIG_OPTION(kASTPrintEnable, Bool);

‎src/op/builtin.h‎

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -62,8 +62,6 @@ static constexpr const char *kEnableAsyncCopy = "tl.enable_async_copy";
6262
static constexpr const char *kEnableVectorizePlannerVerbose =
6363
"tl.enable_vectorize_planner_verbose";
6464
static constexpr const char *kDisableWGMMA = "tl.disable_wgmma";
65-
static constexpr const char *kDisable2CTATcgen5MMA =
66-
"tl.disable_2cta_tcgen5mma";
6765
static constexpr const char *kDisableShuffleElect = "tl.disable_shuffle_elect";
6866
static constexpr const char *kDisableLoopUnswitching =
6967
"tl.disable_loop_unswitching";

‎src/op/copy.cc‎

Lines changed: 4 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1701,8 +1701,7 @@ Stmt CopyNode::LowerBulkCopy(const LowerArgs &T, arith::Analyzer *analyzer,
17011701
barrier_base_id = 0;
17021702
// Detect cluster barrier by checking the buffer scope
17031703
if (auto bl = mbar_handle.as<BufferLoadNode>()) {
1704-
is_cluster_barrier =
1705-
bl->buffer.scope() == "shared.cluster_barrier";
1704+
is_cluster_barrier = bl->buffer.scope() == "shared.cluster_barrier";
17061705
}
17071706
} else if (GetIsTmaCopy()) {
17081707
LOG(FATAL) << "T.tma_copy() requires a barrier argument. "
@@ -1809,10 +1808,9 @@ Stmt CopyNode::LowerBulkCopy(const LowerArgs &T, arith::Analyzer *analyzer,
18091808
Stmt expect_stmt =
18101809
Evaluate(Call(DataType::Handle(), mbarrier_expect_tx(),
18111810
{mbar_handle, cluster_total_bytes}));
1812-
PrimExpr rank =
1813-
Call(DataType::Int(32), block_rank_in_cluster(), {});
1814-
barrier_before_tma_stmt = IfThenElse(
1815-
EQ(rank, IntImm(DataType::Int(32), 0)), expect_stmt);
1811+
PrimExpr rank = Call(DataType::Int(32), block_rank_in_cluster(), {});
1812+
barrier_before_tma_stmt =
1813+
IfThenElse(EQ(rank, IntImm(DataType::Int(32), 0)), expect_stmt);
18161814
} else {
18171815
barrier_before_tma_stmt =
18181816
Evaluate(Call(DataType::Handle(), mbarrier_expect_tx(),

‎src/op/operator.h‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -76,8 +76,9 @@ struct LowerArgs {
7676
// Points to the LowerTileOpPass member so copy.cc sees the buffer
7777
// even when created lazily by the AllocMBarrier callback.
7878
Optional<Buffer> *mbarrier_buffer = nullptr;
79-
// Product of cluster_dims (from block annotation). Defaults to 1 (no cluster).
80-
// Used by TMA copy lowering to scale expect_tx bytes for cluster barriers.
79+
// Product of cluster_dims (from block annotation). Defaults to 1 (no
80+
// cluster). Used by TMA copy lowering to scale expect_tx bytes for cluster
81+
// barriers.
8182
int cluster_size = 1;
8283
};
8384

‎src/transform/lower_blackwell_2sm.cc‎

Lines changed: 10 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -20,10 +20,8 @@
2020
#include <tvm/tir/stmt_functor.h>
2121
#include <tvm/tir/transform.h>
2222

23-
#include "../op/builtin.h"
2423
#include "../op/gemm_py.h"
2524
#include "../op/operator.h"
26-
#include "../op/tcgen5_meta.h"
2725
#include "../target/utils.h"
2826

2927
namespace tvm {
@@ -70,36 +68,27 @@ static bool HasValidClusterDimsFor2Cta(const Stmt &body) {
7068
*/
7169
class Tcgen5_2SmLower : public StmtExprMutator {
7270
public:
73-
Tcgen5_2SmLower(Target target, bool cluster_dims_valid)
74-
: target_(std::move(target)), cluster_dims_valid_(cluster_dims_valid) {}
71+
Tcgen5_2SmLower(bool cluster_dims_valid)
72+
: cluster_dims_valid_(cluster_dims_valid) {}
7573
bool has_2sm_tcgen5mma() const { return has_2sm_tcgen5mma_; }
7674

7775
private:
7876
Stmt VisitStmt_(const EvaluateNode *op) final {
7977
if (const CallNode *call = op->value.as<CallNode>()) {
8078
TileOperator tile_op = ParseOperator(ffi::GetRef<Stmt>(op));
81-
if (tile_op.defined()) {
82-
if (Optional<GemmPy> opt_gemm_py = tile_op.as<GemmPy>()) {
83-
const GemmPyNode *node = opt_gemm_py.value().get();
84-
if (node->allowTcgen5Mma(target_)) {
85-
auto [ok, meta] =
86-
GetTCGEN5MMAMeta(node->m_, node->n_, node->k_,
87-
node->a_->dtype, node->c_->dtype);
88-
if (ok && meta.enable_2cta) {
79+
if (tile_op.defined() && tile_op.as<GemmPy>()) {
80+
// Check if the user explicitly requested 2CTA via the use_2cta
81+
// annotation on the Call node (set by T.tcgen05_gemm(use_2cta=True)).
82+
if (call->annotations.count(attr::kUse2Cta)) {
83+
auto val = call->annotations.Get(attr::kUse2Cta).value();
84+
if (const auto *imm = val.as<IntImmNode>()) {
85+
if (imm->value) {
8986
if (!cluster_dims_valid_) {
9087
LOG(WARNING) << "Invalid cluster_dims disables 2CTA "
9188
"TCGEN5MMA, use 1CTA variant instead.";
9289
return StmtExprMutator::VisitStmt_(op);
9390
}
94-
// LOG(INFO) << "Found 2SM TCGEN5MMA!";
9591
has_2sm_tcgen5mma_ = true;
96-
// Annotate the GemmPy CallNode with use_2cta so that
97-
// Python lower code can read it and pass disable_2cta=False.
98-
auto new_annotations = call->annotations;
99-
new_annotations.Set(attr::kUse2Cta, IntImm(DataType::Int(32), 1));
100-
auto new_call = Call(call->dtype, call->op, call->args,
101-
new_annotations, call->span);
102-
return Evaluate(new_call);
10392
}
10493
}
10594
}
@@ -108,7 +97,6 @@ class Tcgen5_2SmLower : public StmtExprMutator {
10897
return StmtExprMutator::VisitStmt_(op);
10998
}
11099

111-
Target target_;
112100
bool cluster_dims_valid_;
113101
bool has_2sm_tcgen5mma_ = false;
114102
};
@@ -145,14 +133,9 @@ tvm::transform::Pass LowerBlackwell2SM() {
145133
if (!opt_target.defined() || !TargetIsSm100(opt_target.value())) {
146134
return f;
147135
}
148-
if (ctx->GetConfig(kDisable2CTATcgen5MMA, Optional<Bool>())
149-
.value_or(false)) {
150-
LOG(INFO) << "2CTA TCGEN5MMA is disabled by pass config";
151-
return f;
152-
}
153136
Stmt body = f->body;
154137
bool cluster_dims_valid = HasValidClusterDimsFor2Cta(body);
155-
Tcgen5_2SmLower lower(opt_target.value(), cluster_dims_valid);
138+
Tcgen5_2SmLower lower(cluster_dims_valid);
156139
body = lower(std::move(body));
157140
if (lower.has_2sm_tcgen5mma()) {
158141
// Annotate block attr for using 2cta tcgen5

‎src/transform/lower_tile_op.cc‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -323,8 +323,8 @@ class LowerTileOpPass : arith::IRMutatorWithAnalyzer {
323323
}
324324
// Extract cluster_size from cluster_dims annotation
325325
if (op->annotations.count("cluster_dims")) {
326-
if (auto arr = op->annotations.Get("cluster_dims")
327-
->try_cast<Array<Integer>>()) {
326+
if (auto arr =
327+
op->annotations.Get("cluster_dims")->try_cast<Array<Integer>>()) {
328328
int sz = 1;
329329
for (auto d : arr.value())
330330
sz *= static_cast<int>(d->value);

‎src/transform/producer_consumer_ws.cc‎

Lines changed: 11 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -2757,8 +2757,7 @@ class ProducerConsumerWSRewriter : public StmtExprMutator {
27572757
bool drop_arrive, bool is_cluster_barrier,
27582758
int cluster_size)
27592759
: barrier_buf_(barrier_buf), barrier_id_(std::move(barrier_id)),
2760-
drop_arrive_(drop_arrive),
2761-
is_cluster_barrier_(is_cluster_barrier),
2760+
drop_arrive_(drop_arrive), is_cluster_barrier_(is_cluster_barrier),
27622761
cluster_size_(cluster_size) {}
27632762

27642763
Stmt VisitStmt_(const EvaluateNode *op) final {
@@ -2773,10 +2772,10 @@ class ProducerConsumerWSRewriter : public StmtExprMutator {
27732772
call->args.size() == 2) {
27742773
PrimExpr new_bytes =
27752774
call->args[1] * IntImm(DataType::Int(32), cluster_size_);
2776-
auto new_call = Call(
2777-
call->dtype, call->op,
2778-
{makeGetBarrier(barrier_buf_, barrier_id_), new_bytes},
2779-
call->annotations, call->span);
2775+
auto new_call =
2776+
Call(call->dtype, call->op,
2777+
{makeGetBarrier(barrier_buf_, barrier_id_), new_bytes},
2778+
call->annotations, call->span);
27802779
PrimExpr rank =
27812780
Call(DataType::Int(32), tl::block_rank_in_cluster(), {});
27822781
return IfThenElse(EQ(rank, IntImm(DataType::Int(32), 0)),
@@ -2936,10 +2935,10 @@ class ProducerConsumerWSRewriter : public StmtExprMutator {
29362935
call->args.size() == 2) {
29372936
PrimExpr new_bytes =
29382937
call->args[1] * IntImm(DataType::Int(32), cluster_size_);
2939-
auto new_call = Call(
2940-
call->dtype, mbarrier_expect_tx(),
2941-
{makeGetBarrier(barrier_buf_, barrier_id_), new_bytes},
2942-
call->annotations, call->span);
2938+
auto new_call =
2939+
Call(call->dtype, mbarrier_expect_tx(),
2940+
{makeGetBarrier(barrier_buf_, barrier_id_), new_bytes},
2941+
call->annotations, call->span);
29432942
PrimExpr rank =
29442943
Call(DataType::Int(32), tl::block_rank_in_cluster(), {});
29452944
return IfThenElse(EQ(rank, IntImm(DataType::Int(32), 0)),
@@ -2996,9 +2995,8 @@ class ProducerConsumerWSRewriter : public StmtExprMutator {
29962995

29972996
// Rebind the producer-side barrier id and finish the stage with a normal
29982997
// barrier arrival. Pure-TMA pipelines do not need cp.async.mbarrier.arrive.
2999-
Stmt rewritten = MergeAdjacentEquivalentIfs(
3000-
TmaForwardBarrierStmtRewriter(barrier_buf_, barrier_id,
3001-
is_cluster_barrier_, cluster_size_)(stmt));
2998+
Stmt rewritten = MergeAdjacentEquivalentIfs(TmaForwardBarrierStmtRewriter(
2999+
barrier_buf_, barrier_id, is_cluster_barrier_, cluster_size_)(stmt));
30023000
if (!append_arrive) {
30033001
return rewritten;
30043002
}

‎tilelang/language/gemm_op.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,7 @@ def _gemm_impl(
3131
k_pack: int = 1,
3232
wg_wait: int = 0,
3333
mbar: BarrierType | None = None,
34+
annotations: dict | None = None,
3435
) -> tir.PrimExpr:
3536
"""Shared GEMM implementation.
3637
@@ -132,6 +133,7 @@ def legalize_arguments(arg: BufferLikeType | tir.Var) -> BufferLikeType:
132133
mbar_arg,
133134
C_coords[0],
134135
C_coords[1],
136+
annotations=annotations,
135137
)
136138

137139

@@ -277,6 +279,7 @@ def tcgen05_gemm(
277279
clear_accum: bool = False,
278280
*,
279281
mbar: BarrierType,
282+
use_2cta: bool = False,
280283
) -> tir.PrimExpr:
281284
"""Explicit Blackwell TCGEN05 GEMM without an implicit wait.
282285
@@ -285,10 +288,14 @@ def tcgen05_gemm(
285288
- it always requests the TCGEN5MMA lowering path
286289
- it never auto-emits an inlined `mbarrier_wait_parity`
287290
291+
When ``use_2cta=True``, the instruction is lowered to the 2CTA variant
292+
which requires ``cluster_dims`` to be ``(2,1,1)`` or ``(1,2,1)``.
293+
288294
If the current target or operand pattern cannot use Blackwell TCGEN5MMA,
289295
compilation fails instead of silently falling back to another GEMM path.
290296
"""
291297

298+
ann = {"use_2cta": int(use_2cta)} if use_2cta else None
292299
return _gemm_impl(
293300
"tl.tileop.tcgen05_gemm_py",
294301
A,
@@ -301,4 +308,5 @@ def tcgen05_gemm(
301308
1,
302309
0,
303310
mbar,
311+
annotations=ann,
304312
)

0 commit comments

Comments
 (0)