Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
Show all changes
39 commits
Select commit Hold shift + click to select a range
9f9677a
support 2-cta alloc, dealloc, umma_arrive
Rachmanino Feb 26, 2026
fbc6bec
draft tma load 2sm support
Rachmanino Feb 26, 2026
4dc2901
draft tcgen5 2sm support
Rachmanino Mar 2, 2026
9b048a4
add draft lower_blackwell_2sm pass
Rachmanino Mar 3, 2026
2fc3064
support threadblock swizzle annotation for cluster launch
Rachmanino Mar 6, 2026
8b27b40
enhance cuda arch restriction for cluster template functions
Rachmanino Mar 6, 2026
ae83639
Change return type in block_rank_in_cluster function from uint32 to i…
Rachmanino Mar 6, 2026
d45865a
fix threadblock swizzle for cluster
Rachmanino Mar 6, 2026
d2d7111
introduce tcgen05.fence::{before, after}_thread_sync
Rachmanino Mar 6, 2026
20f01be
update lower_blackwell_2sm pass
Rachmanino Mar 6, 2026
ebb8b76
update b_continuity calculation in GemmTCGEN5 for compatiblity with 2cta
Rachmanino Mar 6, 2026
150fb9b
use m/n per_cta for offset calculation in 2cta tcgen5 and introduce a…
Rachmanino Mar 6, 2026
6ddea2a
draft refactor of tcgen05mma
Rachmanino Mar 6, 2026
ec9cef9
successfully generated correct 2sm code
Rachmanino Mar 6, 2026
dcc1bae
lint
Rachmanino Mar 6, 2026
d09fc92
fix bug caused by rebase
Rachmanino Mar 9, 2026
a76f221
upd 2sm persistent kernel
Rachmanino Mar 9, 2026
28617da
reorgnize tcgen5mma examples, fix cross-wave bugs, optimize to 1670T
Rachmanino Mar 12, 2026
d229491
lint
Rachmanino Mar 12, 2026
dbc9388
general support for 2cta in GetTCGEN5MMAMeta
Rachmanino Mar 12, 2026
a01e62d
drop the change of make_tcgen05mma_swizzled_layout
Rachmanino Mar 12, 2026
5ddea83
support 2cta for more dtypes and gemm_ts
Rachmanino Mar 12, 2026
d832d14
add check for cluster_dims when lowering 2cta tcgen5
Rachmanino Mar 12, 2026
97beb2c
refactor 2sm conditional lowering
Rachmanino Mar 13, 2026
9572c64
remove unnecessary syncthreads
Rachmanino Mar 13, 2026
15e3048
support Layout B and add maint test
Rachmanino Mar 13, 2026
d901fed
lint
Rachmanino Mar 13, 2026
605d9dc
Restore support for cp.async.mbarrier and cleaning up unused code
Rachmanino Mar 16, 2026
2ec5fe8
Merge branch 'main' into wt/2sm
Rachmanino Mar 20, 2026
4131990
Merge branch 'main' into wt/2sm
Rachmanino Mar 20, 2026
33ac055
fix tma injection pass change
Rachmanino Mar 20, 2026
1caa424
typo
Rachmanino Mar 20, 2026
6c09a53
Merge branch 'main' into wt/2sm
Rachmanino Mar 23, 2026
d174892
support tma_load_2sm lowering after merge main
Rachmanino Mar 23, 2026
6e698b7
refactor to expose use_2cta option in `T.tcgen05_gemm` and remove 2ct…
Rachmanino Mar 23, 2026
3795413
lint
Rachmanino Mar 23, 2026
d42abc8
fix cutedsl codegen for threadblock swizzle
Rachmanino Mar 23, 2026
e7e03b6
lint
Rachmanino Mar 23, 2026
28f4224
upd maint script for tcgen05
Rachmanino Mar 23, 2026
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
lint
  • Loading branch information
Rachmanino committed Mar 18, 2026
commit dcc1bae4f9f48ab56ad3be50c1215b673cea5058
5 changes: 3 additions & 2 deletions examples/gemm_sm100/gemm_tcgen5mma_2sm.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import tilelang
import tilelang.language as T
from tilelang.engine import register_cuda_postproc

tilelang.disable_cache()


Expand Down Expand Up @@ -130,7 +131,7 @@ def main(


M, N, K = 2048, 2048, 2048 # FIXME: buggy when size is lager
print(f'M: {M}, N: {N}, K: {K}')
print(f"M: {M}, N: {N}, K: {K}")
block_M, block_N, block_K = 128, 256, 64
in_dtype, out_dtype, accum_dtype = T.bfloat16, T.bfloat16, T.float
num_stages = 0 if block_N >= 256 or block_M >= 256 or block_K >= 256 else 2
Expand All @@ -156,7 +157,7 @@ def main(
c = jit_kernel(a, b)
ref_c = (a @ b).to(torch.bfloat16)
torch.testing.assert_close(c, ref_c, rtol=1e-2, atol=1e-2)
print('ALL CHECK PASSED. ✅')
print("ALL CHECK PASSED. ✅")
profiler = jit_kernel.get_profiler()
latency = profiler.do_bench()
print(f"Latency: {latency} ms")
Expand Down
10 changes: 6 additions & 4 deletions examples/gemm_sm100/gemm_tcgen5mma_ws_2sm.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import tilelang.language as T
from tilelang.profiler import do_bench
from tilelang.engine import register_cuda_postproc

tilelang.disable_cache()


Expand Down Expand Up @@ -151,7 +152,10 @@ def gemm(A, B, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_
for k in T.serial(k_iters):
T.mbarrier_wait_parity(consumed[k % num_stages], ((k // num_stages) & 1) ^ 1)
T.copy(A[bx * block_M : (bx + 1) * block_M, k * block_K : (k + 1) * block_K], A_shared[k % num_stages, :, :])
T.copy(B[k * block_K : (k + 1) * block_K, (by * 2 + cta_id) * (block_N // 2) : (by * 2 + cta_id + 1) * (block_N // 2)], B_shared[k % num_stages, :, :])
T.copy(
B[k * block_K : (k + 1) * block_K, (by * 2 + cta_id) * (block_N // 2) : (by * 2 + cta_id + 1) * (block_N // 2)],
B_shared[k % num_stages, :, :],
)
T.mbarrier_arrive(loaded[k % num_stages], 0) # arrive on leader cta's barrier
elif cta_id == 0 and tx < 64: # Only warp 1 on leader cta issues tcgen5
for k in T.serial(k_iters):
Expand Down Expand Up @@ -190,13 +194,11 @@ def main():
b = torch.randn(K, N, device="cuda", dtype=torch.bfloat16)
print(gemm.get_kernel_source(a, b, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages))
c = gemm(a, b, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages)


ref_c = (a.to(torch.float) @ b.to(torch.float)).to(torch.bfloat16)
# torch.testing.assert_close(c, ref_c, rtol=1e-2, atol=1e-2)
torch.testing.assert_close(c, ref_c, rtol=1e-2, atol=1e-2)
print("All checks passed. ✅")


tl_latency = do_bench(lambda: gemm(a, b, block_M, block_N, block_K, in_dtype, out_dtype, accum_dtype, num_stages), backend="cupti")
torch_latency = do_bench(lambda: a @ b, backend="cupti")
print(f"Tilelang latency: {tl_latency} ms")
Expand Down
3 changes: 2 additions & 1 deletion src/op/builtin.h
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,8 @@ static constexpr const char *kEnableAsyncCopy = "tl.enable_async_copy";
static constexpr const char *kEnableVectorizePlannerVerbose =
"tl.enable_vectorize_planner_verbose";
static constexpr const char *kDisableWGMMA = "tl.disable_wgmma";
static constexpr const char *kDisable2CTATcgen5MMA = "tl.disable_2cta_tcgen5mma";
static constexpr const char *kDisable2CTATcgen5MMA =
"tl.disable_2cta_tcgen5mma";
static constexpr const char *kDisableShuffleElect = "tl.disable_shuffle_elect";
static constexpr const char *kDisableLoopUnswitching =
"tl.disable_loop_unswitching";
Expand Down
7 changes: 4 additions & 3 deletions src/op/gemm_py.cc
Original file line number Diff line number Diff line change
Expand Up @@ -318,9 +318,10 @@ TVM_FFI_STATIC_INIT_BLOCK() {
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::GlobalDef().def(
"tl.get_tcgen5_mma_meta",
[](int M, int N, int K, DataType ab_dtype, DataType c_dtype, bool disable_2cta) {
auto [success, meta] = GetTCGEN5MMAMeta(M, N, K, ab_dtype, c_dtype, disable_2cta);
"tl.get_tcgen5_mma_meta", [](int M, int N, int K, DataType ab_dtype,
DataType c_dtype, bool disable_2cta) {
auto [success, meta] =
GetTCGEN5MMAMeta(M, N, K, ab_dtype, c_dtype, disable_2cta);
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Array<Integer> result;
if (success) {
result.push_back(Integer(meta.atom_m));
Expand Down
5 changes: 3 additions & 2 deletions src/op/tcgen5_meta.h
Original file line number Diff line number Diff line change
Expand Up @@ -19,9 +19,10 @@ struct TCGEN5MMAMeta {
};

inline std::pair<bool, TCGEN5MMAMeta>
GetTCGEN5MMAMeta(int M, int N, int K, DataType ab_dtype, DataType c_dtype, bool disable_2cta = false) {
GetTCGEN5MMAMeta(int M, int N, int K, DataType ab_dtype, DataType c_dtype,
bool disable_2cta = false) {
Comment thread
coderabbitai[bot] marked this conversation as resolved.
// TODO (lei) Currently not all shapes / dtypes are supported for TCGEN5MMA.
// NOTE: In case that other tcgen5mma in the same kernel must use 1-cta,
// NOTE: In case that other tcgen5mma in the same kernel must use 1-cta,
// we should disable 2cta for the current tcgen5mma.
#define FAIL \
return { \
Expand Down
61 changes: 35 additions & 26 deletions src/target/codegen_cuda.cc
Original file line number Diff line number Diff line change
Expand Up @@ -436,7 +436,8 @@ class ClusterInfoExtractor : public tir::StmtVisitor {
cluster_grid_x_ext = cluster_dims[0].as<IntImmNode>()->value;
cluster_grid_y_ext = cluster_dims[1].as<IntImmNode>()->value;
cluster_grid_z_ext = cluster_dims[2].as<IntImmNode>()->value;
ICHECK(cluster_grid_x_ext > 0 && cluster_grid_y_ext > 0 && cluster_grid_z_ext > 0);
ICHECK(cluster_grid_x_ext > 0 && cluster_grid_y_ext > 0 &&
cluster_grid_z_ext > 0);
}
StmtVisitor::VisitStmt(f->body);
}
Expand All @@ -447,10 +448,12 @@ class ClusterInfoExtractor : public tir::StmtVisitor {
int64_t cluster_grid_z_ext = 1;

public:
std::optional<std::tuple<int64_t, int64_t, int64_t>> extract(const PrimFunc &f) {
std::optional<std::tuple<int64_t, int64_t, int64_t>>
extract(const PrimFunc &f) {
this->VisitStmt(f);
if (launch_with_cluster) {
return std::make_tuple(cluster_grid_x_ext, cluster_grid_y_ext, cluster_grid_z_ext);
return std::make_tuple(cluster_grid_x_ext, cluster_grid_y_ext,
cluster_grid_z_ext);
}
return std::nullopt;
}
Expand Down Expand Up @@ -1895,16 +1898,16 @@ void CodeGenTileLangCUDA::VisitExpr_(const CallNode *op, std::ostream &os) {
} else if (op->op.same_as(tl::ptx_init_tensor_memory())) {
std::ostringstream ss;
ss << "tl::tmem_allocate";
if (op->annotations.find("use_2cta") != op->annotations.end()
&& Downcast<Bool>(op->annotations["use_2cta"])->value) {
if (op->annotations.find("use_2cta") != op->annotations.end() &&
Downcast<Bool>(op->annotations["use_2cta"])->value) {
ss << "<true>";
}
print_extern_call_stmt(ss.str());
} else if (op->op.same_as(tl::ptx_deallocate_tensor_memory())) {
std::ostringstream ss;
ss << "tl::tmem_deallocate";
if (op->annotations.find("use_2cta") != op->annotations.end()
&& Downcast<Bool>(op->annotations["use_2cta"])->value) {
if (op->annotations.find("use_2cta") != op->annotations.end() &&
Downcast<Bool>(op->annotations["use_2cta"])->value) {
ss << "<true>";
}
print_extern_call_stmt(ss.str());
Expand All @@ -1917,9 +1920,11 @@ void CodeGenTileLangCUDA::VisitExpr_(const CallNode *op, std::ostream &os) {
this->eviction_policy_names_
[op->args[op->args.size() - 1].as<IntImmNode>()->value];
// Simplify the code by using the default eviction policy
if (op->annotations.find("use_2cta") != op->annotations.end() && Downcast<Bool>(op->annotations["use_2cta"])->value) {
if (op->annotations.find("use_2cta") != op->annotations.end() &&
Downcast<Bool>(op->annotations["use_2cta"])->value) {
if (eviction_policy != "EVICT_NORMAL") {
ss << "tl::tma_load_2sm<tl::CacheHintSm100::" << eviction_policy << ">(";
ss << "tl::tma_load_2sm<tl::CacheHintSm100::" << eviction_policy
<< ">(";
} else {
ss << "tl::tma_load_2sm(";
}
Expand Down Expand Up @@ -2473,15 +2478,13 @@ void CodeGenTileLangCUDA::VisitExpr_(const CallNode *op, std::ostream &os) {
bool enable_2cta = Downcast<Bool>(op->args[14])->value;

auto dtype_enum = tl::codegen::ptx::DTypeFromString(kind_dtype);
std::string ab_type_str =
tl::codegen::ptx::DTypeEnumToString(dtype_enum);
std::string ab_type_str = tl::codegen::ptx::DTypeEnumToString(dtype_enum);

// Currently tcgen05mma_ss<Float16|BFloat16> has use_2cta template param;
// Currently tcgen05mma_ss<Float16|BFloat16> has use_2cta template param;
// others don't. tcgen05mma_ws_ss has no use_2cta.
std::string use_2cta_suffix;
if (!enable_ws &&
(dtype_enum == tl::codegen::ptx::DataType::kFloat16 ||
dtype_enum == tl::codegen::ptx::DataType::kBFloat16)) {
if (!enable_ws && (dtype_enum == tl::codegen::ptx::DataType::kFloat16 ||
dtype_enum == tl::codegen::ptx::DataType::kBFloat16)) {
use_2cta_suffix = std::string(", ") + (enable_2cta ? "true" : "false");
}

Expand Down Expand Up @@ -2564,8 +2567,8 @@ void CodeGenTileLangCUDA::VisitExpr_(const CallNode *op, std::ostream &os) {
need_tcgen05_common_h_ = true;
std::ostringstream ss;
ss << "tl::tcgen05_mma_arrive";
if (op->annotations.find("use_2cta") != op->annotations.end()
&& Downcast<Bool>(op->annotations["use_2cta"])->value) {
if (op->annotations.find("use_2cta") != op->annotations.end() &&
Downcast<Bool>(op->annotations["use_2cta"])->value) {
ss << "<true>";
}
print_extern_call_stmt(ss.str());
Expand Down Expand Up @@ -3378,25 +3381,31 @@ void CodeGenTileLangCUDA::VisitStmt_(const AttrStmtNode *op) {
std::string func_name;
int panel_size = 0;
if (const auto *call = op->value.as<CallNode>()) {
if (call->op.same_as(tir::builtin::tvm_tuple()) && call->args.size() >= 2) {
if (call->op.same_as(tir::builtin::tvm_tuple()) &&
call->args.size() >= 2) {
const auto *name_node = call->args[0].as<StringImmNode>();
const auto *size_node = call->args[1].as<IntImmNode>();
ICHECK(name_node && size_node)
<< "threadblock_swizzle_pattern expects tvm_tuple(device_func, panel_size)";
ICHECK(name_node && size_node) << "threadblock_swizzle_pattern expects "
"tvm_tuple(device_func, panel_size)";
func_name = name_node->value;
panel_size = static_cast<int>(size_node->value);
}
}
ICHECK(!func_name.empty() && panel_size > 0);
if (this->cluster_dims.has_value()) {
auto [cluster_grid_x_ext, cluster_grid_y_ext, cluster_grid_z_ext] = this->cluster_dims.value();
auto [cluster_grid_x_ext, cluster_grid_y_ext, cluster_grid_z_ext] =
this->cluster_dims.value();
ICHECK(cluster_grid_y_ext == 1 && cluster_grid_z_ext == 1)
<< "Only support annotate threadblock swizzle for cluster on X dimension for now!";
ICHECK(panel_size % cluster_grid_x_ext == 0) << "panel_size must be divisible by clusterDim.x";
this->stream << "const dim3 blockIdx = tl::" << func_name << "WithCluster<"
<< panel_size / cluster_grid_x_ext << ", " << cluster_grid_x_ext << ">();\n";
<< "Only support annotate threadblock swizzle for cluster on X "
"dimension for now!";
ICHECK(panel_size % cluster_grid_x_ext == 0)
<< "panel_size must be divisible by clusterDim.x";
this->stream << "const dim3 blockIdx = tl::" << func_name
<< "WithCluster<" << panel_size / cluster_grid_x_ext << ", "
<< cluster_grid_x_ext << ">();\n";
} else {
this->stream << "const dim3 blockIdx = tl::" << func_name << "<" << panel_size << ">();\n";
this->stream << "const dim3 blockIdx = tl::" << func_name << "<"
<< panel_size << ">();\n";
}
this->VisitStmt(op->body);
return;
Expand Down
5 changes: 3 additions & 2 deletions src/tl_templates/cuda/cluster.h
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,9 @@
#include "common.h"

// Config
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900) && \
((__CUDACC_VER_MAJOR__ >= 12) || ((__CUDACC_VER_MAJOR__ == 11) && (__CUDACC_VER_MINOR__ >= 8))))
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900) && \
((__CUDACC_VER_MAJOR__ >= 12) || \
((__CUDACC_VER_MAJOR__ == 11) && (__CUDACC_VER_MINOR__ >= 8))))
#define TILELANG_CLUSTER_ENABLED
#endif

Expand Down
22 changes: 15 additions & 7 deletions src/tl_templates/cuda/common.h
Original file line number Diff line number Diff line change
@@ -1,9 +1,9 @@
#pragma once

#ifndef __CUDACC_RTC__
#include <cuda_runtime.h>
#include <cstdio>
#include <cstdlib>
#include <cuda_runtime.h>
#endif

#include "atomic.h"
Expand Down Expand Up @@ -60,14 +60,22 @@ using int4_t = int4;
} while (0)

#if defined(__CUDA_ARCH__)
#define TILELANG_UNREACHABLE(msg) \
do { printf("%s, %s:%d\n", msg, __FILE__, __LINE__); __trap(); } while(0)
#define TILELANG_UNREACHABLE(msg) \
do { \
printf("%s, %s:%d\n", msg, __FILE__, __LINE__); \
__trap(); \
} while (0)
#elif defined(__CUDACC_RTC__)
#define TILELANG_UNREACHABLE(msg) \
do { __builtin_trap(); } while(0)
#define TILELANG_UNREACHABLE(msg) \
do { \
__builtin_trap(); \
} while (0)
Comment on lines +62 to +72

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🟠 Major

🧩 Analysis chain

🏁 Script executed:

# First, let's look at the file and surrounding context
cat -n src/tl_templates/cuda/common.h | head -100

Repository: tile-ai/tilelang

Length of output: 4549


🏁 Script executed:

# Search for other uses of these macros in the codebase to understand the pattern
rg "__CUDACC_RTC__|__CUDA_ARCH__" --type h --type cpp -B 2 -A 2

Repository: tile-ai/tilelang

Length of output: 7481


🏁 Script executed:

# Check if there are any comments or documentation about this macro
rg "TILELANG_UNREACHABLE|__builtin_trap|printf.*trap" --type h --type cpp -B 3 -A 3

Repository: tile-ai/tilelang

Length of output: 5186


🌐 Web query:

NVRTC __CUDACC_RTC__ __CUDA_ARCH__ macro definitions CUDA runtime compilation

💡 Result:

In NVRTC (CUDA runtime compilation), these macros behave essentially the same as in offline nvcc compilation, with a couple of NVRTC-specific points:

  • __CUDACC_RTC__: Defined by NVRTC (useful to distinguish NVRTC from offline nvcc in #if logic). (docs.nvidia.com)
  • __CUDA_ARCH__: Defined during NVRTC device compilation with the same semantics as nvcc. Its value is tied to the compile target you pass via NVRTC’s --gpu-architecture / -arch option. (docs.nvidia.com)
    • Mapping: if you compile for compute_<version>, then __CUDA_ARCH__ = <version> * 10 (e.g., compute_80 → 800). (docs.nvidia.cn)

Practical NVRTC example:

  • Compile with -arch=compute_80 (or --gpu-architecture=compute_80) ⇒ device code sees __CUDA_ARCH__ == 800. (docs.nvidia.com)

Citations:


Check __CUDACC_RTC__ before __CUDA_ARCH__.

NVRTC defines both macros during device compilation. The current ordering causes the __CUDA_ARCH__ branch to match first, making the __CUDACC_RTC__ branch unreachable. This routes NVRTC builds through the printf path, but <cstdio> is explicitly excluded when __CUDACC_RTC__ is defined (see line 3), causing compilation issues. The RTC-specific __builtin_trap() path must be checked first.

🛠️ Suggested fix
-#if defined(__CUDA_ARCH__)
-#define TILELANG_UNREACHABLE(msg)                                              \
-  do {                                                                         \
-    printf("%s, %s:%d\n", msg, __FILE__, __LINE__);                            \
-    __trap();                                                                  \
-  } while (0)
-#elif defined(__CUDACC_RTC__)
+#if defined(__CUDACC_RTC__)
 `#define` TILELANG_UNREACHABLE(msg)                                              \
   do {                                                                         \
     __builtin_trap();                                                          \
   } while (0)
+#elif defined(__CUDA_ARCH__)
+#define TILELANG_UNREACHABLE(msg)                                              \
+  do {                                                                         \
+    printf("%s, %s:%d\n", msg, __FILE__, __LINE__);                            \
+    __trap();                                                                  \
+  } while (0)
 `#else`
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.

In `@src/tl_templates/cuda/common.h` around lines 62 - 72, The
TILELANG_UNREACHABLE macro currently checks __CUDA_ARCH__ before __CUDACC_RTC__,
so NVRTC builds (which define both) take the printf/__trap branch and fail
because <cstdio> is excluded; reorder the preprocessor checks in
src/tl_templates/cuda/common.h so the __CUDACC_RTC__ check comes first and
defines TILELANG_UNREACHABLE to use __builtin_trap(), leaving the __CUDA_ARCH__
branch after it to retain the printf + __trap() behavior for real device builds.

#else
#define TILELANG_UNREACHABLE(msg) \
do { fprintf(stderr, "%s, %s:%d\n", msg, __FILE__, __LINE__); abort(); } while(0)
#define TILELANG_UNREACHABLE(msg) \
do { \
fprintf(stderr, "%s, %s:%d\n", msg, __FILE__, __LINE__); \
abort(); \
} while (0)
#endif

// using cutlass abs function for half_t
Expand Down
34 changes: 16 additions & 18 deletions src/tl_templates/cuda/copy_sm100.h
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,12 @@
#include <cuda.h>
#endif

#include "barrier.h"
#include "common.h"
#include "cuda_fp8.h"
#include "tcgen_05.h"
#include "tcgen_05_ld.h"
#include "tcgen_05_st.h"
#include "barrier.h"
#include "common.h"

namespace tl {

Expand Down Expand Up @@ -302,8 +302,8 @@ constexpr uint32_t Sm100MmaPeerBitMask = 0xFEFFFFFF;
template <CacheHintSm100 cache_hint = CacheHintSm100::EVICT_NORMAL,
typename BarrierType = uint64_t>
TL_DEVICE void tma_load_2sm(const CUtensorMap &descriptor,
BarrierType &smem_mbar,
void const *const smem_ptr, int32_t const &crd0) {
BarrierType &smem_mbar, void const *const smem_ptr,
int32_t const &crd0) {
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(&descriptor);
// Executed by both CTAs. Set peer bit to 0 so that the
// transaction bytes will update CTA0's barrier.
Expand All @@ -327,9 +327,8 @@ TL_DEVICE void tma_load_2sm(const CUtensorMap &descriptor,
template <CacheHintSm100 cache_hint = CacheHintSm100::EVICT_NORMAL,
typename BarrierType = uint64_t>
TL_DEVICE void tma_load_2sm(const CUtensorMap &descriptor,
BarrierType &smem_mbar,
void const *const smem_ptr, int32_t const &crd0,
int32_t const &crd1) {
BarrierType &smem_mbar, void const *const smem_ptr,
int32_t const &crd0, int32_t const &crd1) {
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(&descriptor);
// Executed by both CTAs. Set peer bit to 0 so that the
// transaction bytes will update CTA0's barrier.
Expand All @@ -353,9 +352,9 @@ TL_DEVICE void tma_load_2sm(const CUtensorMap &descriptor,
template <CacheHintSm100 cache_hint = CacheHintSm100::EVICT_NORMAL,
typename BarrierType = uint64_t>
TL_DEVICE void tma_load_2sm(const CUtensorMap &descriptor,
BarrierType &smem_mbar,
void const *const smem_ptr, int32_t const &crd0,
int32_t const &crd1, int32_t const &crd2) {
BarrierType &smem_mbar, void const *const smem_ptr,
int32_t const &crd0, int32_t const &crd1,
int32_t const &crd2) {
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(&descriptor);
// Executed by both CTAs. Set peer bit to 0 so that the
// transaction bytes will update CTA0's barrier.
Expand All @@ -379,10 +378,9 @@ TL_DEVICE void tma_load_2sm(const CUtensorMap &descriptor,
template <CacheHintSm100 cache_hint = CacheHintSm100::EVICT_NORMAL,
typename BarrierType = uint64_t>
TL_DEVICE void tma_load_2sm(const CUtensorMap &descriptor,
BarrierType &smem_mbar,
void const *const smem_ptr, int32_t const &crd0,
int32_t const &crd1, int32_t const &crd2,
int32_t const &crd3) {
BarrierType &smem_mbar, void const *const smem_ptr,
int32_t const &crd0, int32_t const &crd1,
int32_t const &crd2, int32_t const &crd3) {
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(&descriptor);
// Executed by both CTAs. Set peer bit to 0 so that the
// transaction bytes will update CTA0's barrier.
Expand All @@ -406,10 +404,10 @@ TL_DEVICE void tma_load_2sm(const CUtensorMap &descriptor,
template <CacheHintSm100 cache_hint = CacheHintSm100::EVICT_NORMAL,
typename BarrierType = uint64_t>
TL_DEVICE void tma_load_2sm(const CUtensorMap &descriptor,
BarrierType &smem_mbar,
void const *const smem_ptr, int32_t const &crd0,
int32_t const &crd1, int32_t const &crd2,
int32_t const &crd3, int32_t const &crd4) {
BarrierType &smem_mbar, void const *const smem_ptr,
int32_t const &crd0, int32_t const &crd1,
int32_t const &crd2, int32_t const &crd3,
int32_t const &crd4) {
uint64_t gmem_int_desc = reinterpret_cast<uint64_t>(&descriptor);
// Executed by both CTAs. Set peer bit to 0 so that the
// transaction bytes will update CTA0's barrier.
Expand Down
Loading