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
2927namespace tvm {
@@ -70,36 +68,27 @@ static bool HasValidClusterDimsFor2Cta(const Stmt &body) {
7068 */
7169class Tcgen5_2SmLower : public StmtExprMutator {
7270public:
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
7775private:
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
0 commit comments