Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 12 additions & 12 deletions examples/deepseek_mla/example_mla_decode_ws.py
Original file line number Diff line number Diff line change
Expand Up @@ -219,17 +219,17 @@ def main_split(
T.ptx_cp_async(
T.access_ptr(KV_shared_0_l[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[bid, kv_indices, cur_kv_head, 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(KV_shared_0_r[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[bid, kv_indices, cur_kv_head, dim // 2 + 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(K_tail_shared_0[r * 16 + (tx - 256) // 8, (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(K_pe[bid, kv_indices, cur_kv_head, (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.cp_async_barrier_noinc(bar_k_0_ready[0])

Expand All @@ -241,17 +241,17 @@ def main_split(
T.ptx_cp_async(
T.access_ptr(KV_shared_1_l[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[bid, kv_indices, cur_kv_head, 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(KV_shared_1_r[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[bid, kv_indices, cur_kv_head, dim // 2 + 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(K_tail_shared_1[r * 16 + (tx - 256) // 8, (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(K_pe[bid, kv_indices, cur_kv_head, (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.cp_async_barrier_noinc(bar_k_1_ready[0])

Expand Down Expand Up @@ -467,17 +467,17 @@ def main_no_split(
T.ptx_cp_async(
T.access_ptr(KV_shared_0_l[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[bid, kv_indices, cur_kv_head, 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(KV_shared_0_r[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[bid, kv_indices, cur_kv_head, dim // 2 + 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(K_tail_shared_0[r * 16 + (tx - 256) // 8, (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(K_pe[bid, kv_indices, cur_kv_head, (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.cp_async_barrier_noinc(bar_k_0_ready[0])

Expand All @@ -489,17 +489,17 @@ def main_no_split(
T.ptx_cp_async(
T.access_ptr(KV_shared_1_l[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[bid, kv_indices, cur_kv_head, 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(KV_shared_1_r[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[bid, kv_indices, cur_kv_head, dim // 2 + 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(K_tail_shared_1[r * 16 + (tx - 256) // 8, (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(K_pe[bid, kv_indices, cur_kv_head, (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.cp_async_barrier_noinc(bar_k_1_ready[0])

Expand Down
12 changes: 6 additions & 6 deletions examples/deepseek_v32/sparse_mla_fwd_pipelined.py
Original file line number Diff line number Diff line change
Expand Up @@ -271,17 +271,17 @@ def main(
T.ptx_cp_async(
T.access_ptr(KV_shared_0_l[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, indices_local, g_i, 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(KV_shared_0_r[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, indices_local, g_i, D // 2 + 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(K_tail_shared_0[r * 16 + (tx - 256) // 8, (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, indices_local, g_i, D + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.cp_async_barrier_noinc(bar_k_0_ready[0])

Expand All @@ -296,17 +296,17 @@ def main(
T.ptx_cp_async(
T.access_ptr(KV_shared_1_l[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, indices_local, g_i, 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(KV_shared_1_r[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, indices_local, g_i, D // 2 + 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(K_tail_shared_1[r * 16 + (tx - 256) // 8, (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, indices_local, g_i, D + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.cp_async_barrier_noinc(bar_k_1_ready[0])

Expand Down
12 changes: 6 additions & 6 deletions examples/deepseek_v32/sparse_mla_fwd_seesaw.py
Original file line number Diff line number Diff line change
Expand Up @@ -222,18 +222,18 @@ def main(
T.ptx_cp_async(
T.access_ptr(KV_shared_0_l[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, index, g_i, 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(KV_shared_0_r[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, index, g_i, D // 2 + 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
# tail_dim (64) needs only one iter of 8 elems per 8 collaborating threads
T.ptx_cp_async(
T.access_ptr(K_tail_shared_0[r * 16 + (tx - 256) // 8, (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, index, g_i, D + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.cp_async_barrier_noinc(bar_k_0_ready[0])

Expand All @@ -253,17 +253,17 @@ def main(
T.ptx_cp_async(
T.access_ptr(KV_shared_1_l[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, index, g_i, 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(KV_shared_1_r[r * 16 + (tx - 256) // 8, 64 * u + (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, index, g_i, D // 2 + 64 * u + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.ptx_cp_async(
T.access_ptr(K_tail_shared_1[r * 16 + (tx - 256) // 8, (tx - 256) % 8 * 8], "w", 8),
T.access_ptr(KV[b_i, index, g_i, D + (tx - 256) % 8 * 8], "r", 8),
16,
8,
)
T.cp_async_barrier_noinc(bar_k_1_ready[0])

Expand Down
60 changes: 60 additions & 0 deletions examples/gemm_int4/example_tilelang_gemm_int4.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,60 @@
"""Frontend int4 GEMM example for the T.gemm int4 path.

This file intentionally models the desired TileLang frontend API:
- A/B are declared as T.int4 tensors
- the matmul is expressed with T.gemm(...)

The example compiles the kernel and prints the generated CUDA source.
"""

import tilelang
import tilelang.language as T

tilelang.disable_cache()


def matmul_nt_int4(M, N, K, block_M, block_N, block_K):
@T.prim_func
def main(
A: T.Tensor((M, K), T.int4),
B: T.Tensor((N, K), T.int4),
C: T.Tensor((M, N), T.int32),
):
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
A_shared = T.alloc_shared((block_M, block_K), T.int4)
B_shared = T.alloc_shared((block_N, block_K), T.int4)
C_local = T.alloc_fragment((block_M, block_N), T.int32)

T.clear(C_local)
for ko in T.Pipelined(T.ceildiv(K, block_K), num_stages=3):
T.copy(A[by * block_M, ko * block_K], A_shared)
T.copy(B[bx * block_N, ko * block_K], B_shared)
# Frontend expectation: T.gemm should accept int4 operands directly.
T.gemm(A_shared, B_shared, C_local, transpose_B=True)

T.copy(C_local, C[by * block_M, bx * block_N])

return main


def compile_int4_gemm(
M=1024,
N=1024,
K=1024,
block_M=128,
block_N=128,
block_K=64,
):
func = matmul_nt_int4(M, N, K, block_M, block_N, block_K)
kernel = tilelang.compile(func, out_idx=-1)
print("Compilation succeeded.")
print(kernel.get_kernel_source())
return func, kernel


def main():
compile_int4_gemm()


if __name__ == "__main__":
main()
4 changes: 2 additions & 2 deletions src/op/builtin.h
Original file line number Diff line number Diff line change
Expand Up @@ -443,8 +443,8 @@ TVM_DLL const Op &ptx_cp_async_barrier_noinc();
/*!
* \brief TileLang intrinsic for PTX async copy from global to shared memory
*
* ptx_cp_async(dst_access_ptr, src_access_ptr, bytes)
* ptx_cp_async(dst_access_ptr, src_access_ptr, bytes, predicate)
* ptx_cp_async(dst_access_ptr, src_access_ptr, num_elems)
* ptx_cp_async(dst_access_ptr, src_access_ptr, num_elems, predicate)
*
*/
TVM_DLL const Op &ptx_cp_async();
Expand Down
133 changes: 0 additions & 133 deletions src/op/copy.cc
Original file line number Diff line number Diff line change
Expand Up @@ -49,139 +49,6 @@ PrimExpr GetCopyMbarPhaseExpr(const Map<String, ObjectRef> &annotations,
return phase;
}

// Rewrite scalar global->shared stores into ptx_cp_async calls.
// This rewriter is applied before the global vectorize pass, so each generated
// cp.async call starts with element-wise bytes and can be widened later.
class CPAsyncStoreRewriter : public StmtMutator {
public:
Stmt Rewrite(const Stmt &stmt) { return VisitStmt(stmt); }

bool RewriteSuccess() const {
return rewritten_any_store_ && !failed_on_shared_store_;
}

private:
static bool IsZeroValue(const PrimExpr &e) {
if (auto *b = e.as<BroadcastNode>()) {
return IsZeroValue(b->value);
}
if (auto *f = e.as<FloatImmNode>()) {
return f->value == 0.0f;
}
if (auto *i = e.as<IntImmNode>()) {
return i->value == 0;
}
return false;
}

static const BufferLoadNode *
MatchZeroFillBufferLoad(const PrimExpr &value,
Optional<PrimExpr> *predicate) {
if (const auto *load = value.as<BufferLoadNode>()) {
return load;
}

const auto *call = value.as<CallNode>();
if (!call || !call->op.same_as(builtin::if_then_else()) ||
!IsZeroValue(call->args[2])) {
return nullptr;
}

const BufferLoadNode *load =
MatchZeroFillBufferLoad(call->args[1], predicate);
if (load == nullptr) {
return nullptr;
}

// Nested zero-fill guards only permit issuing cp.async when every guard
// on the path to the load is true.
*predicate =
predicate->defined()
? Optional<PrimExpr>(And(call->args[0], predicate->value()))
: Optional<PrimExpr>(call->args[0]);
return load;
}

Stmt VisitStmt_(const BufferStoreNode *op) final {
if (!IsSharedBuffer(op->buffer)) {
return StmtMutator::VisitStmt_(op);
}

Optional<PrimExpr> predicate = std::nullopt;
// Accept either a direct load or a nested zero-fill guard chain:
// if_then_else(p1, if_then_else(p2, load, 0), 0). Nested predicates are
// combined so the generated cp.async is only issued when all guards hold.
const BufferLoadNode *load = MatchZeroFillBufferLoad(op->value, &predicate);
if (load == nullptr) {
failed_on_shared_store_ = true;
return StmtMutator::VisitStmt_(op);
}

if (!IsGlobalBuffer(load->buffer)) {
failed_on_shared_store_ = true;
return StmtMutator::VisitStmt_(op);
}
int bytes = op->value.dtype().bytes();
int vectorized_lanes = current_vectorized_lanes_;

if (!IsValidCPAsyncTransferBytes(bytes * vectorized_lanes)) {
failed_on_shared_store_ = true;
return StmtMutator::VisitStmt_(op);
}

// Keep pointer metadata in tl.access_ptr form for downstream analysis;
// LowerAccessPtr will translate it to tvm_access_ptr later.
PrimExpr dst_access_ptr =
Call(DataType::Handle(), tvm::tl::access_ptr(),
{
BufferLoad(op->buffer, op->indices),
IntImm(DataType::Int(32), 1), // extent
IntImm(DataType::Int(32), 2) // rw_mask: write
});
PrimExpr src_access_ptr =
Call(DataType::Handle(), tvm::tl::access_ptr(),
{
BufferLoad(load->buffer, load->indices),
IntImm(DataType::Int(32), 1), // extent
IntImm(DataType::Int(32), 1) // rw_mask: read
});

Array<PrimExpr> args{dst_access_ptr, src_access_ptr, PrimExpr(bytes)};
if (predicate.defined()) {
args.push_back(predicate.value());
}
rewritten_any_store_ = true;
return Evaluate(Call(DataType::Handle(), builtin::ptx_cp_async(), args));
}

Stmt VisitStmt_(const ForNode *op) final {
int previous_vectorized_lanes = current_vectorized_lanes_;
if (op->kind == ForKind::kVectorized) {
// Assume vectorized access pattern is contiguous on the vectorized iter.
// This is guaranteed by tl.VectorizeLoop: if an access pattern is not
// vectorizable/contiguous for the chosen iter, it is scalarized instead
// of staying as ForKind::kVectorized.
const auto *extent_imm = op->extent.as<IntImmNode>();
ICHECK(extent_imm)
<< "Vectorized loops must have constant extent, but got "
<< op->extent;
int lanes = static_cast<int>(extent_imm->value);
if (lanes > 1 && current_vectorized_lanes_ <=
std::numeric_limits<int>::max() / lanes) {
current_vectorized_lanes_ *= lanes;
}
}

Stmt stmt = StmtMutator::VisitStmt_(op);
current_vectorized_lanes_ = previous_vectorized_lanes;
return stmt;
}

bool rewritten_any_store_ = false;
bool failed_on_shared_store_ = false;
int current_vectorized_lanes_ = 1;
};

} // namespace

// Constructs a Copy operator node from call arguments and annotations.
Expand Down
Loading
Loading