Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
Prev Previous commit
Next Next commit
Change tl.ptx_cp_async to use num_elems semantics
  • Loading branch information
LeiWang1999 committed Apr 20, 2026
commit 52a696593bcf961d6e219a34d10e7378a7e2edef
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
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
73 changes: 68 additions & 5 deletions src/target/codegen_cuda.cc
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@
#include <tvm/tir/op.h>

#include <cmath>
#include <cstdint>
#include <optional>
#include <string>
#include <utility>
#include <vector>
Expand All @@ -26,6 +28,69 @@ using namespace ffi;

namespace {

bool IsValidCPAsyncTransferBytes(int64_t bytes) {
return bytes == 4 || bytes == 8 || bytes == 16;
}

std::optional<DataType> GetAccessPtrElementType(const PrimExpr &expr) {
const auto *ptr_call = expr.as<CallNode>();
if (ptr_call == nullptr) {
return std::nullopt;
}
if (ptr_call->op.same_as(builtin::address_of())) {
const auto *buffer_load = ptr_call->args[0].as<BufferLoadNode>();
ICHECK(buffer_load) << "address_of arg must be BufferLoad";
return buffer_load->buffer->dtype;
}
if (ptr_call->op.same_as(builtin::tvm_access_ptr())) {
ICHECK(!ptr_call->args.empty());
return ptr_call->args[0].dtype();
}
if (ptr_call->op.same_as(tl::access_ptr())) {
ICHECK_EQ(ptr_call->args.size(), 3U)
<< "tl.access_ptr expects 3 args: (BufferLoad, extent, rw_mask)";
const auto *buffer_load = ptr_call->args[0].as<BufferLoadNode>();
ICHECK(buffer_load) << "tl.access_ptr arg0 must be BufferLoad";
return buffer_load->buffer->dtype;
}
return std::nullopt;
}

int GetTileLangCPAsyncTransferBytes(const CallNode *op) {
ICHECK(op->args.size() == 3 || op->args.size() == 4)
<< "tl::ptx_cp_async expects 3 or 4 arguments (dst_access_ptr, "
"src_access_ptr, num_elems, [predicate])";
const auto *num_elems_imm = op->args[2].as<IntImmNode>();
ICHECK(num_elems_imm) << "tl::ptx_cp_async num_elems must be IntImm, but got "
<< op->args[2];
int64_t num_elems = num_elems_imm->value;
ICHECK_GT(num_elems, 0);

auto dst_elem_type = GetAccessPtrElementType(op->args[0]);
auto src_elem_type = GetAccessPtrElementType(op->args[1]);
ICHECK(dst_elem_type.has_value() && src_elem_type.has_value())
<< "tl::ptx_cp_async expects address_of, tl.access_ptr, or "
"tvm_access_ptr operands";

int64_t dst_total_bits =
num_elems * dst_elem_type.value().bits() * dst_elem_type.value().lanes();
int64_t src_total_bits =
num_elems * src_elem_type.value().bits() * src_elem_type.value().lanes();
ICHECK_EQ(dst_total_bits, src_total_bits)
<< "tl::ptx_cp_async requires src/dst transfer widths to match, but got "
<< dst_total_bits << " vs " << src_total_bits << " bits";
ICHECK_EQ(dst_total_bits % 8, 0)
<< "tl::ptx_cp_async requires byte-aligned transfers, but got "
<< dst_total_bits << " bits";

int64_t total_bytes = dst_total_bits / 8;
ICHECK(IsValidCPAsyncTransferBytes(total_bytes))
<< "tl::ptx_cp_async requires a final PTX byte width in {4, 8, 16}, but "
"got "
<< total_bytes;
return static_cast<int>(total_bytes);
}

bool CanEmitPackedX2Math(DataType t) {
int lanes = t.lanes();
if (lanes < 2 || lanes % 2 != 0) {
Expand Down Expand Up @@ -1999,14 +2064,12 @@ void CodeGenTileLangCUDA::VisitExpr_(const CallNode *op, std::ostream &os) {
}
} else if (op->op.same_as(tl::ptx_cp_async())) {
// TileLang version: args[0] = dst_access_ptr, args[1] = src_access_ptr,
// args[2] = bytes, args[3] = predicate (optional)
ICHECK(op->args.size() == 3 || op->args.size() == 4)
<< "tl::ptx_cp_async expects 3 or 4 arguments (dst_access_ptr, "
"src_access_ptr, bytes, [predicate])";
// args[2] = num_elems, args[3] = predicate (optional)
int total_bytes = GetTileLangCPAsyncTransferBytes(op);

std::string dst = this->PrintExpr(op->args[0]);
std::string src = this->PrintExpr(op->args[1]);
std::string size = this->PrintExpr(op->args[2]);
std::string size = std::to_string(total_bytes);

this->PrintIndent();
if (op->args.size() == 3) {
Expand Down
73 changes: 68 additions & 5 deletions src/target/codegen_cutedsl.cc
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@
#include <tvm/tir/op.h>

#include <cmath>
#include <cstdint>
#include <optional>
#include <string>
#include <utility>
#include <vector>
Expand Down Expand Up @@ -64,6 +66,69 @@ void ReplaceAll(std::string &str, const std::string &from,
}
}

bool IsValidCPAsyncTransferBytes(int64_t bytes) {
return bytes == 4 || bytes == 8 || bytes == 16;
}

std::optional<DataType> GetAccessPtrElementType(const PrimExpr &expr) {
const auto *ptr_call = expr.as<CallNode>();
if (ptr_call == nullptr) {
return std::nullopt;
}
if (ptr_call->op.same_as(builtin::address_of())) {
const auto *buffer_load = ptr_call->args[0].as<BufferLoadNode>();
ICHECK(buffer_load) << "address_of arg must be BufferLoad";
return buffer_load->buffer->dtype;
}
if (ptr_call->op.same_as(builtin::tvm_access_ptr())) {
ICHECK(!ptr_call->args.empty());
return ptr_call->args[0].dtype();
}
if (ptr_call->op.same_as(tl::access_ptr())) {
ICHECK_EQ(ptr_call->args.size(), 3U)
<< "tl.access_ptr expects 3 args: (BufferLoad, extent, rw_mask)";
const auto *buffer_load = ptr_call->args[0].as<BufferLoadNode>();
ICHECK(buffer_load) << "tl.access_ptr arg0 must be BufferLoad";
return buffer_load->buffer->dtype;
}
return std::nullopt;
}

int GetTileLangCPAsyncTransferBytes(const CallNode *op) {
ICHECK(op->args.size() == 3 || op->args.size() == 4)
<< "tl::ptx_cp_async expects 3 or 4 arguments (dst_access_ptr, "
"src_access_ptr, num_elems, [predicate])";
const auto *num_elems_imm = op->args[2].as<IntImmNode>();
ICHECK(num_elems_imm) << "tl::ptx_cp_async num_elems must be IntImm, but got "
<< op->args[2];
int64_t num_elems = num_elems_imm->value;
ICHECK_GT(num_elems, 0);

auto dst_elem_type = GetAccessPtrElementType(op->args[0]);
auto src_elem_type = GetAccessPtrElementType(op->args[1]);
ICHECK(dst_elem_type.has_value() && src_elem_type.has_value())
<< "tl::ptx_cp_async expects address_of, tl.access_ptr, or "
"tvm_access_ptr operands";

int64_t dst_total_bits =
num_elems * dst_elem_type.value().bits() * dst_elem_type.value().lanes();
int64_t src_total_bits =
num_elems * src_elem_type.value().bits() * src_elem_type.value().lanes();
ICHECK_EQ(dst_total_bits, src_total_bits)
<< "tl::ptx_cp_async requires src/dst transfer widths to match, but got "
<< dst_total_bits << " vs " << src_total_bits << " bits";
ICHECK_EQ(dst_total_bits % 8, 0)
<< "tl::ptx_cp_async requires byte-aligned transfers, but got "
<< dst_total_bits << " bits";

int64_t total_bytes = dst_total_bits / 8;
ICHECK(IsValidCPAsyncTransferBytes(total_bytes))
<< "tl::ptx_cp_async requires a final PTX byte width in {4, 8, 16}, but "
"got "
<< total_bytes;
return static_cast<int>(total_bytes);
}

} // namespace

CodeGenTileLangCuTeDSL::CodeGenTileLangCuTeDSL() {
Expand Down Expand Up @@ -429,14 +494,12 @@ void CodeGenTileLangCuTeDSL::VisitExpr_(const CallNode *op,
}
} else if (op->op.same_as(tl::ptx_cp_async())) {
// TileLang version: args[0] = dst_access_ptr, args[1] = src_access_ptr,
// args[2] = bytes, args[3] = predicate (optional)
ICHECK(op->args.size() == 3 || op->args.size() == 4)
<< "tl::ptx_cp_async expects 3 or 4 arguments (dst_access_ptr, "
"src_access_ptr, bytes, [predicate])";
// args[2] = num_elems, args[3] = predicate (optional)
int total_bytes = GetTileLangCPAsyncTransferBytes(op);

std::string dst = PrintExpr_(op->args[0]);
std::string src = PrintExpr_(op->args[1]);
std::string size = PrintExpr_(op->args[2]);
std::string size = std::to_string(total_bytes);

if (op->args.size() == 3) {
this->PrintIndent();
Expand Down
Loading
Loading