Skip to content

Commit 9f25954

Browse files
authored
[Refactor] Refactor Pass InjectFenceProxy (#1863)
* tir: add T.cdiv alias for T.ceildiv * docs: Update InjectFenceProxy documentation and enhance code comments - Clarified the description of `tl.InjectFenceProxy` to specify the transition between generic and async proxy operations. - Improved explanations of the pass's functionality, including state tracking and the handling of TMA store synchronization. - Added details on the new `ProxyStateSet` class and its role in managing proxy states. - Updated usage instructions for proxy hints to include new options for custom operations. - Enhanced test coverage for handling unknown external calls and proxy hint overrides. * fix * enhance * fix * remove tl.proxy_hint * refactor tma_store_arrive and tma_store_wait * refactor * fix * InjectFenceProxy: hoist fence out of pure-async loops * InjectFenceProxy: hoist fence for if/while pure-async regions * Remove LowerTileOp transform test
1 parent 391cf5b commit 9f25954

10 files changed

Lines changed: 1723 additions & 316 deletions

File tree

‎docs/compiler_internals/inject_fence_proxy.md‎

Lines changed: 17 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -1,34 +1,34 @@
11
# InjectFenceProxy Pass
22

3-
`tl.InjectFenceProxy` is a TIR-level transform that keeps the GPU proxy state consistent on NVIDIA Hopper (SM90+) by inserting `fence.proxy.async` instructions when control flow switches from generic memory operations to asynchronous proxy operations.
3+
`tl.InjectFenceProxy` is a TIR-level transform that keeps the GPU proxy state consistent on NVIDIA Hopper (SM90+) by inserting `fence.proxy.async` instructions when execution switches from **generic proxy** memory operations to **async proxy** operations.
44

55
## Why Fences Are Needed
66

7-
Hopper separates memory instructions into generic and asynchronous proxy paths. When an asynchronous instruction (for example, `cp.async` or `tma.load`) issues after generic traffic (like `ldmatrix` or plain buffer stores), the hardware requires a `fence.proxy.async` to guarantee ordering. Missing fences can lead to race conditions or undefined behavior.
7+
Hopper separates memory instructions into generic and asynchronous proxy paths. When an asynchronous instruction (for example, `wgmma`, `tma.load`, or `cp.async.bulk`) issues after generic traffic (like `ldmatrix`, `cp.async`, or **shared-memory** buffer stores), the hardware requires a `fence.proxy.async` to guarantee ordering. Missing fences can lead to race conditions or undefined behavior.
88

99
## What the Pass Does
1010

11-
- Walks every statement in the `PrimFunc`, tracking whether it behaves as a **generic**, **async**, or **neutral** proxy (neutral statements reset the state, such as an explicit fence).
12-
- Automatically lowers `tma_store` intrinsics into the required `arrive`/`wait` handshake so that TMA stores participate correctly in synchronization.
13-
- Injects an explicit `fence.proxy.async` whenever a generic statement is followed by an async statement without an intervening neutral barrier.
11+
- Walks statements in execution order while tracking a (may-)state of the last proxy kind (**generic**, **async**, or **none/reset**). Control-flow joins (e.g. `if`) merge states conservatively.
12+
- Normalizes `tma_store` by ensuring the required `tma_store_arrive` / `tma_store_wait` handshake exists immediately after the store.
13+
- Injects `fence.proxy.async` right before an async-proxy instruction whenever the preceding state can be generic.
1414

15-
The pass is conservative: unknown extern calls are treated as async so that the fence is inserted rather than accidentally omitted.
15+
By default, unknown/external calls do **not** affect proxy state. Opaque calls that may write into **shared memory** (e.g. via `tvm_access_ptr` / `address_of`) are treated as generic proxy traffic so a later async-proxy op will still be fenced.
1616

1717
### Timeline View
1818

1919
```
20-
generic initialize_wgmma_descriptor → generic shared-store → async wgmma
21-
│ │ │
22-
└─ generic proxy ┴─ generic proxy ┴─ async proxy
23-
│ fence inserted here ↑
24-
└──────────────────────────────┘
20+
generic shared-store (or ldmatrix/stmatrix/cp.async) → async op (wgmma / tma / cp.async.bulk)
21+
│ │
22+
└─ generic proxy └─ async proxy
23+
│ fence inserted here ↑
24+
└──────────────────────────────┘
2525
```
2626

27-
The proxy tracker scans the sequence from left to right. The moment it detects a transition from generic to async (between the store and `cp.async` above), it synthesizes a `fence.proxy.async` to reset the hardware proxy state before the async path runs.
27+
The proxy tracker effectively scans the program in execution order. The moment it detects a possible transition from generic to async (between the store and the async op above), it synthesizes a `fence.proxy.async` to reset the hardware proxy state before the async path runs.
2828

2929
## Coverage of Intrinsics
3030

31-
The tracker understands the TileLang intrinsics for TMA load/store, shared-memory MMA (`wgmma`), and TVM/PTX async copy intrinsics (`cp.async` variants). Generic operations currently include `ldmatrix`, `stmatrix`, and descriptor initialization. Other IR nodes (loops, blocks, attributes) receive a proxy kind derived from their bodies so that the analysis survives structured control flow.
31+
The tracker understands the TileLang intrinsics for TMA load/store, shared-memory MMA (`wgmma`), and TVM/PTX SM90 async copy intrinsics (`cp.async.bulk` family). Generic operations currently include `ldmatrix`, `stmatrix`, `cp.async`, and **shared-memory** `BufferStore` statements. Structured control flow (loops, blocks, branches) is handled by propagating and conservatively merging proxy state.
3232

3333
## Usage
3434

@@ -110,4 +110,7 @@ The only change is the `fence_proxy_async` between the generic descriptor setup
110110

111111
## Extending the Pass
112112

113-
If you introduce a new intrinsic that behaves like an async proxy, add it to `IsAsyncIntrinsic` in `src/transform/inject_fence_proxy.cc`. Likewise, extend `IsKnownGeneric` for additional generic operations. When adding new neutral barriers, make sure they set the proxy kind to `kNeutral` so the state resets correctly.
113+
If you introduce a new intrinsic that behaves like an async proxy, add it to `IsAsyncIntrinsic` in `src/transform/inject_fence_proxy.cc`. Likewise, extend `IsKnownGeneric` for additional generic operations.
114+
115+
Most calls default to `"none"` (no proxy-state effect). `IsNonProxyIntrinsic` exists for well-known synchronization / scheduling helpers and to document intent, but it is not required for correctness if an op is neither generic nor async.
116+
For custom/opaque ops, you must lower them into known intrinsics (or manually insert `fence_proxy_async`) if they participate in proxy switching.

‎examples/quickstart.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -76,8 +76,8 @@ def matmul_relu_kernel(
7676
print("Kernel output matches PyTorch reference.")
7777

7878
# 4. Retrieve and inspect the generated CUDA source (optional)
79-
# cuda_source = matmul_relu_kernel.get_kernel_source()
80-
# print("Generated CUDA kernel:\n", cuda_source)
79+
cuda_source = matmul_relu_kernel.get_kernel_source()
80+
print("Generated CUDA kernel:\n", cuda_source)
8181

8282
# 5.Profile latency with kernel
8383
profiler = matmul_relu_kernel.get_profiler(tensor_supply_type=tilelang.TensorSupplyType.Normal)

‎src/op/atomic_add.cc‎

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -570,7 +570,13 @@ Stmt AtomicAddNode::Lower(const LowerArgs &T, arith::Analyzer *analyzer) const {
570570
Evaluate(Call(DataType::Handle(), tma_store(), args, op_annotations));
571571
}
572572

573-
return IfThenElse(EQ(T.thread_var, T.thread_bounds->min), tma_reduce);
573+
Array<Stmt> seq;
574+
seq.reserve(3);
575+
seq.push_back(tma_reduce);
576+
seq.push_back(Evaluate(Call(DataType::Handle(), tma_store_arrive(), {})));
577+
seq.push_back(Evaluate(Call(DataType::Handle(), tma_store_wait(), {})));
578+
return IfThenElse(EQ(T.thread_var, T.thread_bounds->min),
579+
SeqStmt(std::move(seq)));
574580
}
575581
auto simt_loop = MakeSIMTLoop(analyzer);
576582
auto fused_loop = Downcast<For>(ParallelLoopFuser::Fuse(simt_loop));

‎src/op/copy.cc‎

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1425,6 +1425,19 @@ Stmt CopyNode::LowerBulkCopy(const LowerArgs &T, arith::Analyzer *analyzer,
14251425
args.push_back(GetEvictionPolicy());
14261426
tma_copy = Evaluate(Call(DataType::Handle(), op, args));
14271427
}
1428+
1429+
// Bulk TMA stores participate in the cp.async.bulk group mechanism, so we
1430+
// must commit and wait to ensure completion before the store buffer is
1431+
// reused or the kernel exits.
1432+
if (!is_load) {
1433+
Array<Stmt> seq;
1434+
seq.reserve(3);
1435+
seq.push_back(tma_copy);
1436+
seq.push_back(Evaluate(Call(DataType::Handle(), tma_store_arrive(), {})));
1437+
seq.push_back(Evaluate(Call(DataType::Handle(), tma_store_wait(), {})));
1438+
tma_copy = SeqStmt(std::move(seq));
1439+
}
1440+
14281441
tma_copy = IfThenElse(EQ(T.thread_var, T.thread_bounds->min), tma_copy);
14291442

14301443
return tma_copy;
@@ -1502,6 +1515,16 @@ Stmt CopyNode::LowerBulkCopy1D(const LowerArgs &T, arith::Analyzer *analyzer,
15021515
{global_addr, shared_addr, elements * shared_tensor->dtype.bytes(),
15031516
need_reduce, GetEvictionPolicy()}));
15041517
}
1518+
1519+
if (!is_load) {
1520+
Array<Stmt> seq;
1521+
seq.reserve(3);
1522+
seq.push_back(tma_copy);
1523+
seq.push_back(Evaluate(Call(DataType::Handle(), tma_store_arrive(), {})));
1524+
seq.push_back(Evaluate(Call(DataType::Handle(), tma_store_wait(), {})));
1525+
tma_copy = SeqStmt(std::move(seq));
1526+
}
1527+
15051528
tma_copy = IfThenElse(EQ(T.thread_var, T.thread_bounds->min), tma_copy);
15061529
return tma_copy;
15071530
}

0 commit comments

Comments
 (0)