Skip to content

Commit 1f2c229

Browse files
committed
move fence_barrier_init inside the election of one thread
1 parent 41a300a commit 1f2c229

2 files changed

Lines changed: 13 additions & 5 deletions

File tree

‎src/transform/lower_hopper_intrin.cc‎

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -142,6 +142,14 @@ class LowerHopperIntrin : public StmtExprMutator {
142142
auto stmts = prefetch_calls_;
143143
stmts.insert(stmts.end(), init_mbarrier_calls_.begin(),
144144
init_mbarrier_calls_.end());
145+
// Add fence_barrier_init inside the if block, after the barrier init
146+
// calls. This ensures the fence is only executed by the thread that
147+
// initializes the barriers, matching the pattern in DeepGEMM/cutlass
148+
if (!init_mbarrier_calls_.empty()) {
149+
stmts.push_back(Evaluate(Call(DataType::Handle(),
150+
tvm::tl::ptx_fence_barrier_init(),
151+
{})));
152+
}
145153
PrimExpr condition;
146154
if (!disable_shuffle_elect_) {
147155
condition = Call(DataType::Bool(), tl_shuffle_elect(), {0});
@@ -159,9 +167,6 @@ class LowerHopperIntrin : public StmtExprMutator {
159167
// with an appropriate sync instruction with the right scope to
160168
// ensure visibility eg. __syncthreads() or a cluster_arrive() +
161169
// cluster_wait()
162-
Stmt mem_fence = Evaluate(Call(
163-
DataType::Handle(), tvm::tl::ptx_fence_barrier_init(), {}));
164-
stmt_seq.push_back(mem_fence);
165170
Stmt mem_sync =
166171
Evaluate(Call(DataType::Handle(), builtin::tvm_storage_sync(),
167172
{StringImm("shared")}));

‎src/transform/lower_shared_barrier.cc‎

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -137,14 +137,17 @@ class SharedBarrierRewriter : public StmtExprMutator {
137137
} else {
138138
condition = EQ(thread_var_->var, 0);
139139
}
140+
// Add fence_barrier_init inside the if block, after the barrier init calls
141+
// This ensures the fence is only executed by the thread that initializes
142+
// the barriers, matching the pattern in DeepGEMM/cutlass
143+
init_mbarrier_calls_.push_back(
144+
Evaluate(Call(DataType::Handle(), ptx_fence_barrier_init(), {})));
140145
new_body.push_back(IfThenElse(condition,
141146
init_mbarrier_calls_.size() == 1
142147
? init_mbarrier_calls_.back()
143148
: SeqStmt(init_mbarrier_calls_),
144149
Stmt()));
145150

146-
new_body.push_back(
147-
Evaluate(Call(DataType::Handle(), ptx_fence_barrier_init(), {})));
148151
new_body.push_back(Evaluate(
149152
Call(DataType::Handle(), builtin::tvm_storage_sync(),
150153
{StringImm(has_cluster_barrier_ ? "cluster" : "shared")})));

0 commit comments

Comments
 (0)