Skip to content

Commit ff94e53

Browse files
committed
Support int4 T.gemm and direct packed cp.async lowering
1 parent 0924dab commit ff94e53

12 files changed

Lines changed: 585 additions & 435 deletions

File tree

Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,60 @@
1+
"""Frontend int4 GEMM example for the T.gemm int4 path.
2+
3+
This file intentionally models the desired TileLang frontend API:
4+
- A/B are declared as T.int4 tensors
5+
- the matmul is expressed with T.gemm(...)
6+
7+
The example compiles the kernel and prints the generated CUDA source.
8+
"""
9+
10+
import tilelang
11+
import tilelang.language as T
12+
13+
tilelang.disable_cache()
14+
15+
16+
def matmul_nt_int4(M, N, K, block_M, block_N, block_K):
17+
@T.prim_func
18+
def main(
19+
A: T.Tensor((M, K), T.int4),
20+
B: T.Tensor((N, K), T.int4),
21+
C: T.Tensor((M, N), T.int32),
22+
):
23+
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
24+
A_shared = T.alloc_shared((block_M, block_K), T.int4)
25+
B_shared = T.alloc_shared((block_N, block_K), T.int4)
26+
C_local = T.alloc_fragment((block_M, block_N), T.int32)
27+
28+
T.clear(C_local)
29+
for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=3):
30+
T.copy(A[by * block_M, ko * block_K], A_shared)
31+
T.copy(B[bx * block_N, ko * block_K], B_shared)
32+
# Frontend expectation: T.gemm should accept int4 operands directly.
33+
T.gemm(A_shared, B_shared, C_local, transpose_B=True)
34+
35+
T.copy(C_local, C[by * block_M, bx * block_N])
36+
37+
return main
38+
39+
40+
def compile_int4_gemm(
41+
M=1024,
42+
N=1024,
43+
K=1024,
44+
block_M=128,
45+
block_N=128,
46+
block_K=64,
47+
):
48+
func = matmul_nt_int4(M, N, K, block_M, block_N, block_K)
49+
kernel = tilelang.compile(func, out_idx=-1)
50+
print("Compilation succeeded.")
51+
print(kernel.get_kernel_source())
52+
return func, kernel
53+
54+
55+
def main():
56+
compile_int4_gemm()
57+
58+
59+
if __name__ == "__main__":
60+
main()

‎src/op/copy.cc‎

Lines changed: 0 additions & 133 deletions
Original file line numberDiff line numberDiff line change
@@ -49,139 +49,6 @@ PrimExpr GetCopyMbarPhaseExpr(const Map<String, ObjectRef> &annotations,
4949
return phase;
5050
}
5151

52-
// Rewrite scalar global->shared stores into ptx_cp_async calls.
53-
// This rewriter is applied before the global vectorize pass, so each generated
54-
// cp.async call starts with element-wise bytes and can be widened later.
55-
class CPAsyncStoreRewriter : public StmtMutator {
56-
public:
57-
Stmt Rewrite(const Stmt &stmt) { return VisitStmt(stmt); }
58-
59-
bool RewriteSuccess() const {
60-
return rewritten_any_store_ && !failed_on_shared_store_;
61-
}
62-
63-
private:
64-
static bool IsZeroValue(const PrimExpr &e) {
65-
if (auto *b = e.as<BroadcastNode>()) {
66-
return IsZeroValue(b->value);
67-
}
68-
if (auto *f = e.as<FloatImmNode>()) {
69-
return f->value == 0.0f;
70-
}
71-
if (auto *i = e.as<IntImmNode>()) {
72-
return i->value == 0;
73-
}
74-
return false;
75-
}
76-
77-
static const BufferLoadNode *
78-
MatchZeroFillBufferLoad(const PrimExpr &value,
79-
Optional<PrimExpr> *predicate) {
80-
if (const auto *load = value.as<BufferLoadNode>()) {
81-
return load;
82-
}
83-
84-
const auto *call = value.as<CallNode>();
85-
if (!call || !call->op.same_as(builtin::if_then_else()) ||
86-
!IsZeroValue(call->args[2])) {
87-
return nullptr;
88-
}
89-
90-
const BufferLoadNode *load =
91-
MatchZeroFillBufferLoad(call->args[1], predicate);
92-
if (load == nullptr) {
93-
return nullptr;
94-
}
95-
96-
// Nested zero-fill guards only permit issuing cp.async when every guard
97-
// on the path to the load is true.
98-
*predicate =
99-
predicate->defined()
100-
? Optional<PrimExpr>(And(call->args[0], predicate->value()))
101-
: Optional<PrimExpr>(call->args[0]);
102-
return load;
103-
}
104-
105-
Stmt VisitStmt_(const BufferStoreNode *op) final {
106-
if (!IsSharedBuffer(op->buffer)) {
107-
return StmtMutator::VisitStmt_(op);
108-
}
109-
110-
Optional<PrimExpr> predicate = std::nullopt;
111-
// Accept either a direct load or a nested zero-fill guard chain:
112-
// if_then_else(p1, if_then_else(p2, load, 0), 0). Nested predicates are
113-
// combined so the generated cp.async is only issued when all guards hold.
114-
const BufferLoadNode *load = MatchZeroFillBufferLoad(op->value, &predicate);
115-
if (load == nullptr) {
116-
failed_on_shared_store_ = true;
117-
return StmtMutator::VisitStmt_(op);
118-
}
119-
120-
if (!IsGlobalBuffer(load->buffer)) {
121-
failed_on_shared_store_ = true;
122-
return StmtMutator::VisitStmt_(op);
123-
}
124-
int bytes = op->value.dtype().bytes();
125-
int vectorized_lanes = current_vectorized_lanes_;
126-
127-
if (!IsValidCPAsyncTransferBytes(bytes * vectorized_lanes)) {
128-
failed_on_shared_store_ = true;
129-
return StmtMutator::VisitStmt_(op);
130-
}
131-
132-
// Keep pointer metadata in tl.access_ptr form for downstream analysis;
133-
// LowerAccessPtr will translate it to tvm_access_ptr later.
134-
PrimExpr dst_access_ptr =
135-
Call(DataType::Handle(), tvm::tl::access_ptr(),
136-
{
137-
BufferLoad(op->buffer, op->indices),
138-
IntImm(DataType::Int(32), 1), // extent
139-
IntImm(DataType::Int(32), 2) // rw_mask: write
140-
});
141-
PrimExpr src_access_ptr =
142-
Call(DataType::Handle(), tvm::tl::access_ptr(),
143-
{
144-
BufferLoad(load->buffer, load->indices),
145-
IntImm(DataType::Int(32), 1), // extent
146-
IntImm(DataType::Int(32), 1) // rw_mask: read
147-
});
148-
149-
Array<PrimExpr> args{dst_access_ptr, src_access_ptr, PrimExpr(bytes)};
150-
if (predicate.defined()) {
151-
args.push_back(predicate.value());
152-
}
153-
rewritten_any_store_ = true;
154-
return Evaluate(Call(DataType::Handle(), builtin::ptx_cp_async(), args));
155-
}
156-
157-
Stmt VisitStmt_(const ForNode *op) final {
158-
int previous_vectorized_lanes = current_vectorized_lanes_;
159-
if (op->kind == ForKind::kVectorized) {
160-
// Assume vectorized access pattern is contiguous on the vectorized iter.
161-
// This is guaranteed by tl.VectorizeLoop: if an access pattern is not
162-
// vectorizable/contiguous for the chosen iter, it is scalarized instead
163-
// of staying as ForKind::kVectorized.
164-
const auto *extent_imm = op->extent.as<IntImmNode>();
165-
ICHECK(extent_imm)
166-
<< "Vectorized loops must have constant extent, but got "
167-
<< op->extent;
168-
int lanes = static_cast<int>(extent_imm->value);
169-
if (lanes > 1 && current_vectorized_lanes_ <=
170-
std::numeric_limits<int>::max() / lanes) {
171-
current_vectorized_lanes_ *= lanes;
172-
}
173-
}
174-
175-
Stmt stmt = StmtMutator::VisitStmt_(op);
176-
current_vectorized_lanes_ = previous_vectorized_lanes;
177-
return stmt;
178-
}
179-
180-
bool rewritten_any_store_ = false;
181-
bool failed_on_shared_store_ = false;
182-
int current_vectorized_lanes_ = 1;
183-
};
184-
18552
} // namespace
18653

18754
// Constructs a Copy operator node from call arguments and annotations.

‎src/transform/loop_vectorize.cc‎

Lines changed: 3 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -457,23 +457,9 @@ class VectorizePlanner : public arith::IRMutatorWithAnalyzer {
457457
return arith::IRMutatorWithAnalyzer::VisitExpr_(node);
458458
} else if (node->op.same_as(builtin::ptx_cp_async()) ||
459459
node->op.same_as(tl::ptx_cp_async())) {
460-
// cp.async supports byte sizes 4/8/16. For element-wise calls with small
461-
// byte width (e.g., fp16 => 2 bytes), we rely on vectorization to fold
462-
// multiple calls into one wider cp.async call.
463-
int vectorize_length = 1;
464-
ICHECK_GE(node->args.size(), 3U)
465-
<< "cp.async expects at least 3 arguments, but got " << node->args;
466-
const auto *bytes_imm = node->args[2].as<IntImmNode>();
467-
ICHECK(bytes_imm) << "cp.async byte count must be IntImm, but got "
468-
<< node->args[2];
469-
int bytes = static_cast<int>(bytes_imm->value);
470-
for (int lanes : {16, 8, 4, 2, 1}) {
471-
if (IsValidCPAsyncTransferBytes(bytes * lanes)) {
472-
vectorize_length = lanes;
473-
break;
474-
}
475-
}
476-
buffer_vector_infos_.push_back({Buffer(), vectorize_length, false, {}});
460+
// Existing cp.async is already a final-width async transfer. Do not use
461+
// it to encourage additional loop vectorization.
462+
buffer_vector_infos_.push_back({Buffer(), 1, false, {}});
477463
return arith::IRMutatorWithAnalyzer::VisitExpr_(node);
478464
} else if (node->op == builtin::address_of() ||
479465
node->op == tl::access_ptr()) {

‎src/transform/lower_ptx_async_copy.cc‎

Lines changed: 91 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -66,14 +66,16 @@ class PTXAsyncCopyInjector : public StmtMutator {
6666
}
6767

6868
Stmt VisitStmt_(const ForNode *op) final {
69-
// Track nested vectorized loop extents so we can rewrite element-wise
70-
// copies (e.g. float16 stores) into `tir.ptx_cp_async` with element bytes,
71-
// relying on the later `tl.VectorizeLoop` pass to widen:
72-
// for v in T.vectorized(k): ptx_cp_async(dst, src, elem_bytes)
73-
// => ptx_cp_async(dst_base, src_base, elem_bytes * k)
69+
// Track nested vectorized loop extents so cp.async injection can emit the
70+
// final packed transfer width directly:
71+
// for v in T.vectorized(k): store(load(...))
72+
// => ptx_cp_async(dst_base, src_base, total_transfer_bytes)
7473
//
75-
// This mirrors the logic in `CPAsyncStoreRewriter` used by `T.copy`
76-
// lowering, and avoids duplicating vectorize-loop collapse here.
74+
// cp.async only supports byte widths {4, 8, 16}. Instead of emitting
75+
// element-sized transfers here and relying on a later vectorization pass
76+
// to widen them, TryInjectPTX packs the active vectorized loop directly
77+
// whenever the overall transfer width is legal, then collapses the now
78+
// redundant vectorized loop.
7779
int previous_vectorized_lanes = current_vectorized_lanes_;
7880
bool pushed_vectorized_loop = false;
7981
if (op->kind == ForKind::kVectorized) {
@@ -90,6 +92,13 @@ class PTXAsyncCopyInjector : public StmtMutator {
9092
}
9193
}
9294
Stmt stmt = StmtMutator::VisitStmt_(op);
95+
if (pushed_vectorized_loop) {
96+
if (const auto *loop = stmt.as<ForNode>()) {
97+
if (CanCollapseVectorizedCPAsyncLoop(loop->body, loop->loop_var)) {
98+
stmt = loop->body;
99+
}
100+
}
101+
}
93102
if (pushed_vectorized_loop) {
94103
active_vectorized_loops_.pop_back();
95104
}
@@ -123,11 +132,29 @@ class PTXAsyncCopyInjector : public StmtMutator {
123132
index_info->dst_index)) {
124133
return Optional<Stmt>();
125134
}
135+
if (index_info->collapse_vectorized_loop) {
136+
Optional<Array<PrimExpr>> src_base_indices =
137+
ExtractActiveVectorizedLoopBaseIndices(load->indices);
138+
Optional<Array<PrimExpr>> dst_base_indices =
139+
ExtractActiveVectorizedLoopBaseIndices(store->indices);
140+
if (!src_base_indices.defined() || !dst_base_indices.defined()) {
141+
return Optional<Stmt>();
142+
}
143+
return MakeCPAsyncStmtFromLoads(
144+
store, ptr_info.value(),
145+
/*dst_base_load=*/
146+
BufferLoad(store->buffer, dst_base_indices.value()),
147+
/*src_base_load=*/
148+
BufferLoad(load->buffer, src_base_indices.value()),
149+
/*bytes=*/index_info->total_transfer_bytes, predicated,
150+
predicate_value);
151+
}
126152
return MakeCPAsyncStmtFromLoads(
127153
store, ptr_info.value(),
128154
/*dst_base_load=*/BufferLoad(store->buffer, store->indices),
129155
/*src_base_load=*/BufferLoad(load->buffer, load->indices),
130-
/*bytes=*/index_info->transfer_bytes, predicated, predicate_value);
156+
/*bytes=*/index_info->per_access_transfer_bytes, predicated,
157+
predicate_value);
131158
}
132159

133160
Optional<Array<PrimExpr>> src_base_indices =
@@ -147,7 +174,8 @@ class PTXAsyncCopyInjector : public StmtMutator {
147174
store, ptr_info.value(),
148175
/*dst_base_load=*/BufferLoad(store->buffer, dst_base_indices.value()),
149176
/*src_base_load=*/BufferLoad(load->buffer, src_base_indices.value()),
150-
/*bytes=*/index_info->transfer_bytes, predicated, predicate_value);
177+
/*bytes=*/index_info->per_access_transfer_bytes, predicated,
178+
predicate_value);
151179
}
152180

153181
Stmt VisitStmt_(const SeqStmtNode *op) final {
@@ -301,7 +329,9 @@ class PTXAsyncCopyInjector : public StmtMutator {
301329
PrimExpr src_index;
302330
PrimExpr dst_index;
303331
int index_lanes{1};
304-
int transfer_bytes{0};
332+
int per_access_transfer_bytes{0};
333+
int total_transfer_bytes{0};
334+
bool collapse_vectorized_loop{false};
305335
};
306336

307337
// Pointer element type metadata extracted from buffer handle annotations.
@@ -409,9 +439,17 @@ class PTXAsyncCopyInjector : public StmtMutator {
409439
}
410440

411441
const int effective_lanes = std::max(value_lanes, index_lanes);
412-
const int elem_bytes = effective_lanes * load->dtype.bytes();
413-
const int total_bytes = static_cast<int>(elem_bytes) *
414-
static_cast<int>(current_vectorized_lanes_);
442+
const int elem_bits = effective_lanes * load->dtype.bits();
443+
const int total_bits = static_cast<int>(elem_bits) *
444+
static_cast<int>(current_vectorized_lanes_);
445+
// cp.async is byte-granular. We only fold an active vectorized copy into a
446+
// single packed async transfer when the logical payload is exactly
447+
// byte-aligned; otherwise rounding up here would over-copy packed subbyte
448+
// data and change the copy semantics.
449+
if (total_bits % 8 != 0) {
450+
return std::nullopt;
451+
}
452+
const int total_bytes = total_bits / 8;
415453
if (!IsValidCPAsyncTransferBytes(total_bytes)) {
416454
return std::nullopt;
417455
}
@@ -420,10 +458,27 @@ class PTXAsyncCopyInjector : public StmtMutator {
420458
info.src_index = src_index;
421459
info.dst_index = dst_index;
422460
info.index_lanes = index_lanes;
423-
info.transfer_bytes = elem_bytes;
461+
info.per_access_transfer_bytes = (elem_bits + 7) / 8;
462+
info.total_transfer_bytes = total_bytes;
463+
info.collapse_vectorized_loop =
464+
current_vectorized_lanes_ > 1 && index_lanes == 1;
424465
return info;
425466
}
426467

468+
Optional<Array<PrimExpr>>
469+
ExtractActiveVectorizedLoopBaseIndices(const Array<PrimExpr> &indices) {
470+
Array<PrimExpr> base_indices;
471+
base_indices.reserve(indices.size());
472+
for (PrimExpr index : indices) {
473+
for (const auto &loop : active_vectorized_loops_) {
474+
index = analyzer_.Simplify(Substitute(
475+
index, {{loop.loop_var, IntImm(loop.loop_var->dtype, 0)}}));
476+
}
477+
base_indices.push_back(index);
478+
}
479+
return Optional<Array<PrimExpr>>(base_indices);
480+
}
481+
427482
static std::optional<PointerTypeInfo>
428483
PreparePointerTypeInfo(const BufferLoadNode *load,
429484
const BufferStoreNode *store) {
@@ -535,6 +590,28 @@ class PTXAsyncCopyInjector : public StmtMutator {
535590
{IntImm(DataType::Int(32), n)}));
536591
}
537592

593+
static bool IsCPAsyncStmt(const Stmt &stmt) {
594+
const auto *eval = stmt.as<EvaluateNode>();
595+
if (eval == nullptr) {
596+
return false;
597+
}
598+
const auto *call = eval->value.as<CallNode>();
599+
if (call == nullptr) {
600+
return false;
601+
}
602+
return call->op.same_as(builtin::ptx_cp_async()) ||
603+
call->op.same_as(tl::ptx_cp_async());
604+
}
605+
606+
static bool CanCollapseVectorizedCPAsyncLoop(const Stmt &stmt,
607+
const Var &loop_var) {
608+
if (!IsCPAsyncStmt(stmt)) {
609+
return false;
610+
}
611+
return !tir::UsesVar(
612+
stmt, [loop_var](const VarNode *v) { return v == loop_var.get(); });
613+
}
614+
538615
// ---- Vectorized-offset contiguity helpers ----
539616
static bool TryGetConstInt64(const PrimExpr &expr, int64_t *value) {
540617
if (const auto *imm = expr.as<IntImmNode>()) {

0 commit comments

Comments
 (0)