[Backend] Cleanup Metal Codegen, split AsyncCopy lowering and common target utils (#2361)

* [Backend] Cleanup Metal Codegen and split common target utils

* [Backend] Preserve PTX async copy pass docs

* [Backend] Split backend CMake source ownership

* [Backend] Split target warp size helpers
This commit is contained in:
Chaofan Lin
2026-06-09 20:20:25 +08:00
committed by GitHub
parent a3f709328d
commit 2913ad3f00
45 changed files with 1525 additions and 527 deletions
+2 -16
View File
@@ -368,31 +368,15 @@ file(GLOB TILE_LANG_SRCS
src/transform/common/*.cc
src/transform/pipeline/*.cc
src/op/*.cc
src/cpu/op/*.cc
src/cuda/op/copy_analysis.cc
src/metal/op/*.cc
src/webgpu/op/*.cc
src/backend/common/target_utils.cc
src/backend/common/codegen/codegen_utils.cc
src/backend/common/codegen/codegen_c_host.cc
src/backend/common/codegen/codegen_c.cc
src/backend/common/codegen/rt_mod_c.cc
# intrin_rule doesn't have system dependency; always compiled regardless of backend
src/cuda/codegen/intrin_rule_cuda.cc
src/rocm/codegen/intrin_rule_hip.cc
)
# Always include CPU-safe runtime helpers
list(APPEND TILE_LANG_SRCS
src/runtime/error_helpers.cc
)
# Metal codegen is pure C++ (no Apple frameworks) and can generate Metal shader
# source on any platform. Always compile it so that "target.build.tilelang_metal"
# is available for cross-compilation on Linux/Windows.
list(APPEND TILE_LANG_SRCS
src/backend/common/codegen/codegen_metal.cc
)
# Track if the user explicitly selected a backend via cache options.
set(TILELANG_BACKEND_USER_SELECTED OFF)
foreach(BACKEND IN LISTS TILELANG_BACKENDS)
@@ -451,9 +435,11 @@ endif()
# Backend-local CMake files own native source lists, stubs, include paths, and
# compile definitions. Top-level CMake only selects and delegates.
include("${CMAKE_CURRENT_SOURCE_DIR}/src/cpu/CMakeLists.txt")
include("${CMAKE_CURRENT_SOURCE_DIR}/src/cuda/CMakeLists.txt")
include("${CMAKE_CURRENT_SOURCE_DIR}/src/rocm/CMakeLists.txt")
include("${CMAKE_CURRENT_SOURCE_DIR}/src/metal/CMakeLists.txt")
include("${CMAKE_CURRENT_SOURCE_DIR}/src/webgpu/CMakeLists.txt")
set(USE_Z3 ON CACHE STRING "Use Z3 SMT solver for TileLang optimizations")
set(USE_PYPI_Z3 ON CACHE BOOL "Use Z3 provided by PyPI z3-solver package")
+9 -324
View File
@@ -1,345 +1,30 @@
/*!
* \file tl/backend/common/target_utils.cc
* \brief helper functions for target attributes.
* \brief Common target helper dispatch.
*/
#include "backend/common/target_utils.h"
#include "support/check.h"
#include <tvm/ir/cast.h>
#include "dlpack/dlpack.h"
#include <tvm/ffi/reflection/registry.h>
namespace tvm {
namespace tl {
using namespace ffi;
bool TargetIsCuda(Target target) {
return target->GetTargetDeviceType() == kDLCUDA;
}
bool TargetIsRocm(Target target) {
return target->GetTargetDeviceType() == kDLROCM;
}
bool TargetIsMetal(Target target) {
return target->GetTargetDeviceType() == kDLMetal;
}
bool TargetIsCPU(Target target) {
return target->GetTargetDeviceType() == kDLCPU;
}
int GetArchInt(Target target) {
auto s = target->GetAttr<String>("arch");
ICHECK(s.has_value());
const std::string arch_str = s.value();
ICHECK(arch_str.size() >= 3);
ICHECK_EQ(arch_str.compare(0, 3, "sm_"), 0)
<< "arch string must start with sm_";
return std::stoi(arch_str.substr(3));
}
bool TargetIsVolta(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetArchInt(target);
return arch >= 70 && arch < 75;
}
bool TargetIsTuring(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetArchInt(target);
return arch >= 75 && arch < 80;
}
bool TargetIsAmpere(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetArchInt(target);
return arch >= 80 && arch < 90;
}
bool TargetIsHopper(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetArchInt(target);
return arch >= 90 && arch < 100;
}
bool TargetIsSm100(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetArchInt(target);
return arch >= 100 && arch <= 110;
}
bool TargetIsSM120(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetArchInt(target);
return arch >= 120 && arch < 130;
}
bool TargetIsCDNA(Target target) {
if (!TargetIsRocm(target))
return false;
if (target->attrs.count("mcpu")) {
std::string mcpu = Downcast<String>(target->attrs.at("mcpu"));
// if mcpu start with "gfx9", it is CDNA
return mcpu.find("gfx9") == 0;
}
return false;
}
bool TargetIsRDNA(Target target) {
if (!TargetIsRocm(target))
return false;
if (target->attrs.count("mcpu")) {
std::string mcpu = Downcast<String>(target->attrs.at("mcpu"));
// gfx11xx, gfx12xx are RDNA architectures
return mcpu.find("gfx11") == 0 || mcpu.find("gfx12") == 0;
}
return false;
}
bool TargetIsGfx950(Target target) {
if (!TargetIsRocm(target))
return false;
if (target->attrs.count("mcpu")) {
std::string mcpu = Downcast<String>(target->attrs.at("mcpu"));
return mcpu.find("gfx950") != std::string::npos;
}
return false;
}
bool TargetHasAsyncCopy(Target target) {
if (TargetIsCuda(target)) {
int arch = GetArchInt(target);
return arch >= 80;
} else if (TargetIsCDNA(target)) {
if (target->attrs.count("mcpu")) {
std::string mcpu = Downcast<String>(target->attrs.at("mcpu"));
if (mcpu.rfind("gfx9", 0) == 0) {
int gfx_version = std::stoi(mcpu.substr(3, 2));
return gfx_version >= 94;
}
return false;
} else {
return false;
}
return TargetCudaHasAsyncCopy(target);
}
return false;
}
bool TargetHasLdmatrix(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetArchInt(target);
return arch >= 75;
}
bool TargetHasStmatrix(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetArchInt(target);
return arch >= 90;
}
bool TargetHasTmem(Target target) {
if (!TargetIsCuda(target))
return false;
return TargetIsSm100(target);
}
bool TargetHasBulkCopy(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetArchInt(target);
return arch >= 90;
}
bool TargetIsCuTeDSL(Target target) {
for (const auto &key : target->keys) {
if (key == "cutedsl")
return true;
if (TargetIsRocm(target)) {
return TargetRocmHasAsyncCopy(target);
}
return false;
}
bool TargetSupportVectorize256(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetArchInt(target);
return arch >= 100;
}
bool TargetHasSMVersionGE(Target target, int version) {
if (!TargetIsCuda(target))
return false;
int arch = GetArchInt(target);
return arch >= version;
}
int TargetGetWarpSize(Target target) {
int res = 32;
if (TargetIsCDNA(target))
res = 64;
return res;
}
bool IsCudaVectorizableFP8(DataType dtype) {
// NOTE: E8M0 is a special type of FP8 which is not handled here
// We only handle FP8 types which can be represented with
// __nv_fp8_interpretation_t here
return dtype.is_float8_e4m3() || dtype.is_float8_e4m3fn() ||
dtype.is_float8_e5m2();
}
bool IsCudaVectorizableCast(DataType from_ty, DataType target_ty) {
// float16 -> float32
if (from_ty.is_float16() && target_ty.is_float() && target_ty.bits() == 32)
return true;
// float32 -> float16
if (from_ty.is_float() && from_ty.bits() == 32 && target_ty.is_float16())
return true;
// bfloat16 -> float32
if (from_ty.is_bfloat16() && target_ty.is_float() && target_ty.bits() == 32)
return true;
// float32 -> bfloat16
if (from_ty.is_float() && from_ty.bits() == 32 && target_ty.is_bfloat16())
return true;
// float32 -> float8 (E4M3/E5M2)
if (from_ty.is_float() && from_ty.bits() == 32 &&
IsCudaVectorizableFP8(target_ty))
return true;
// float8 (E4M3/E5M2) -> float32
if (IsCudaVectorizableFP8(from_ty) && target_ty.is_float() &&
target_ty.bits() == 32)
return true;
// Not implemented for now
// float64(double) -> float8 (E4M3/E5M2)
// if (from_ty.is_float() && from_ty.bits() == 64 &&
// IsCudaVectorizableFP8(target_ty))
// return true;
// float8 (E4M3/E5M2) -> float64(double)
// if (IsCudaVectorizableFP8(from_ty) && target_ty.is_float() &&
// target_ty.bits() == 64)
// return true;
// float8 (E8M0) -> bfloat16
if (from_ty.is_float8_e8m0fnu() && target_ty.is_bfloat16())
return true;
// bfloat16 -> float8 (E8M0)
if (from_ty.is_bfloat16() && target_ty.is_float8_e8m0fnu())
return true;
// float32 -> float8 (E8M0)
if (from_ty.is_float() && from_ty.bits() == 32 &&
target_ty.is_float8_e8m0fnu())
return true;
// float64(double) -> float8 (E8M0)
if (from_ty.is_float() && from_ty.bits() == 64 &&
target_ty.is_float8_e8m0fnu())
return true;
// float4_e2m1fn -> float16
if (from_ty.is_float4_e2m1fn() && target_ty.is_float16())
return true;
// float16 -> float4_e2m1fn
if (from_ty.is_float16() && target_ty.is_float4_e2m1fn())
return true;
// float4_e2m1fn -> float32
if (from_ty.is_float4_e2m1fn() && target_ty.is_float() &&
target_ty.bits() == 32)
return true;
// float32 -> float4_e2m1fn
if (from_ty.is_float() && from_ty.bits() == 32 &&
target_ty.is_float4_e2m1fn())
return true;
// float4_e2m1fn -> float64(double)
if (from_ty.is_float4_e2m1fn() && target_ty.is_float() &&
target_ty.bits() == 64)
return true;
// float64(double) -> float4_e2m1fn
if (from_ty.is_float() && from_ty.bits() == 64 &&
target_ty.is_float4_e2m1fn())
return true;
// float4_e2m1fn -> bfloat16
if (from_ty.is_float4_e2m1fn() && target_ty.is_bfloat16())
return true;
// bfloat16 -> float4_e2m1fn
if (from_ty.is_bfloat16() && target_ty.is_float4_e2m1fn())
return true;
return false;
}
int TargetGetRDNAGeneration(Target target) {
if (!TargetIsRDNA(target))
return 0;
if (target->attrs.count("mcpu")) {
std::string mcpu = Downcast<String>(target->attrs.at("mcpu"));
if (mcpu.rfind("gfx11", 0) == 0)
return 11;
if (mcpu.rfind("gfx12", 0) == 0)
return 12;
}
return 0;
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = reflection;
refl::GlobalDef()
.def("tl.TargetIsCuda",
[](Target target) { return TargetIsCuda(target); })
.def("tl.TargetIsRocm",
[](Target target) { return TargetIsRocm(target); })
.def("tl.TargetIsMetal",
[](Target target) { return TargetIsMetal(target); })
.def("tl.TargetIsVolta",
[](Target target) { return TargetIsVolta(target); })
.def("tl.TargetIsTuring",
[](Target target) { return TargetIsTuring(target); })
.def("tl.TargetIsAmpere",
[](Target target) { return TargetIsAmpere(target); })
.def("tl.TargetIsHopper",
[](Target target) { return TargetIsHopper(target); })
.def("tl.TargetIsSM120",
[](Target target) { return TargetIsSM120(target); })
.def("tl.TargetIsCDNA",
[](Target target) { return TargetIsCDNA(target); })
.def("tl.TargetIsRDNA",
[](Target target) { return TargetIsRDNA(target); })
.def("tl.TargetIsGfx950",
[](Target target) { return TargetIsGfx950(target); })
.def("tl.TargetHasAsyncCopy",
[](Target target) { return TargetHasAsyncCopy(target); })
.def("tl.TargetHasLdmatrix",
[](Target target) { return TargetHasLdmatrix(target); })
.def("tl.TargetHasStmatrix",
[](Target target) { return TargetHasStmatrix(target); })
.def("tl.TargetHasBulkCopy",
[](Target target) { return TargetHasBulkCopy(target); })
.def("tl.TargetGetRDNAGeneration",
[](Target target) { return TargetGetRDNAGeneration(target); })
.def("tl.TargetGetWarpSize",
[](Target target) { return TargetGetWarpSize(target); });
namespace refl = tvm::ffi::reflection;
refl::GlobalDef().def("tl.TargetHasAsyncCopy", [](Target target) {
return TargetHasAsyncCopy(target);
});
}
} // namespace tl
+6 -29
View File
@@ -1,6 +1,6 @@
/*!
* \file tl/backend/common/target_utils.h
* \brief helper functions for target attributes.
* \brief Common entry points for target attribute helpers.
*
*/
@@ -9,38 +9,15 @@
#include <tvm/target/target.h>
#include "cpu/target_utils.h"
#include "cuda/target_utils.h"
#include "metal/target_utils.h"
#include "rocm/target_utils.h"
namespace tvm {
namespace tl {
bool TargetIsCuda(Target target);
bool TargetIsRocm(Target target);
bool TargetIsMetal(Target target);
bool TargetIsCPU(Target target);
bool TargetIsVolta(Target target);
bool TargetIsTuring(Target target);
bool TargetIsAmpere(Target target);
bool TargetIsHopper(Target target);
bool TargetIsSm100(Target target);
bool TargetIsSM120(Target target);
bool TargetIsCDNA(Target target);
bool TargetIsRDNA(Target target);
bool TargetIsGfx950(Target target);
bool TargetHasAsyncCopy(Target target);
bool TargetHasLdmatrix(Target target);
bool TargetHasStmatrix(Target target);
bool TargetHasTmem(Target target);
bool TargetHasBulkCopy(Target target);
bool TargetIsCuTeDSL(Target target);
bool TargetSupportVectorize256(Target target);
int TargetGetWarpSize(Target target);
bool TargetHasSMVersionGE(Target target, int version);
bool IsCudaVectorizableFP8(DataType dtype);
bool IsCudaVectorizableCast(DataType from_ty, DataType target_ty);
int TargetGetRDNAGeneration(Target target);
} // namespace tl
} // namespace tvm
+13
View File
@@ -0,0 +1,13 @@
# CPU backend: source files that are always safe to compile.
#
# CPU target helpers, operator lowering, and C source codegen have no optional
# system dependencies. Keep them owned by the CPU backend instead of listing
# them in the top-level CMake source block.
file(GLOB TILE_LANG_CPU_SRCS
src/cpu/codegen/codegen_c.cc
src/cpu/codegen/rt_mod_c.cc
src/cpu/op/*.cc
src/cpu/target_utils.cc
)
list(APPEND TILE_LANG_SRCS ${TILE_LANG_CPU_SRCS})
@@ -18,7 +18,7 @@
*/
/*!
* \file codegen_c.cc
* \file cpu/codegen/codegen_c.cc
*/
#include "codegen_c.h"
#include "support/check.h"
@@ -18,11 +18,11 @@
*/
/*!
* \file codegen_c.h
* \file cpu/codegen/codegen_c.h
* \brief Generate C code when target is c (CPU).
*/
#ifndef TVM_TL_CODEGEN_C_H_
#define TVM_TL_CODEGEN_C_H_
#ifndef TILELANG_CPU_CODEGEN_CODEGEN_C_H_
#define TILELANG_CPU_CODEGEN_CODEGEN_C_H_
#include "support/check.h"
#include <string>
@@ -123,4 +123,4 @@ private:
} // namespace codegen
} // namespace tvm
#endif // TVM_TL_CODEGEN_C_H_
#endif // TILELANG_CPU_CODEGEN_CODEGEN_C_H_
+26
View File
@@ -0,0 +1,26 @@
/*!
* \file tl/cpu/target_utils.cc
* \brief CPU target attribute helpers.
*/
#include "cpu/target_utils.h"
#include <tvm/ffi/reflection/registry.h>
#include "dlpack/dlpack.h"
namespace tvm {
namespace tl {
bool TargetIsCPU(Target target) {
return target->GetTargetDeviceType() == kDLCPU;
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::GlobalDef().def("tl.TargetIsCPU",
[](Target target) { return TargetIsCPU(target); });
}
} // namespace tl
} // namespace tvm
+19
View File
@@ -0,0 +1,19 @@
/*!
* \file tl/cpu/target_utils.h
* \brief CPU target attribute helpers.
*/
#ifndef TVM_TL_CPU_TARGET_UTILS_H_
#define TVM_TL_CPU_TARGET_UTILS_H_
#include <tvm/target/target.h>
namespace tvm {
namespace tl {
bool TargetIsCPU(Target target);
} // namespace tl
} // namespace tvm
#endif // TVM_TL_CPU_TARGET_UTILS_H_
+20 -5
View File
@@ -1,4 +1,17 @@
# CUDA backend: toolchain, stub libraries, source files, and build configuration.
# CUDA backend: source files, toolchain, stubs, and build configuration.
#
# Some CUDA-owned files are pure C++ registrations or analysis helpers used by
# common passes even when the CUDA toolchain is disabled. Keep those in the
# backend CMake file, but outside the USE_CUDA block.
file(GLOB TILE_LANG_CUDA_ALWAYS_SRCS
src/cuda/codegen/intrin_rule_cuda.cc
src/cuda/op/copy_analysis.cc
src/cuda/target_utils.cc
src/cuda/transform/lower_ptx_async_copy.cc
src/cuda/transform/ptx_async_copy_injector.cc
)
list(APPEND TILE_LANG_SRCS ${TILE_LANG_CUDA_ALWAYS_SRCS})
if(NOT USE_CUDA)
return()
endif()
@@ -138,7 +151,7 @@ if(TILELANG_USE_CUDA_STUBS)
set(CUDA_NVRTC_LIBRARY nvrtc_stub CACHE STRING "NVRTC library to link against" FORCE)
endif()
file(GLOB TILE_LANG_CUDA_SRCS
file(GLOB TILE_LANG_CUDA_ACTIVE_SRCS
src/cuda/runtime.cc
src/cuda/codegen/ptx.cc
src/cuda/codegen/codegen_cuda.cc
@@ -149,9 +162,11 @@ file(GLOB TILE_LANG_CUDA_SRCS
src/cuda/op/*.cc
src/cuda/transform/*.cc
)
list(REMOVE_ITEM TILE_LANG_CUDA_SRCS
"${CMAKE_CURRENT_SOURCE_DIR}/src/cuda/op/copy_analysis.cc")
list(APPEND TILE_LANG_SRCS ${TILE_LANG_CUDA_SRCS})
list(REMOVE_ITEM TILE_LANG_CUDA_ACTIVE_SRCS
"${CMAKE_CURRENT_SOURCE_DIR}/src/cuda/op/copy_analysis.cc"
"${CMAKE_CURRENT_SOURCE_DIR}/src/cuda/transform/ptx_async_copy_injector.cc"
"${CMAKE_CURRENT_SOURCE_DIR}/src/cuda/transform/lower_ptx_async_copy.cc")
list(APPEND TILE_LANG_SRCS ${TILE_LANG_CUDA_ACTIVE_SRCS})
list(APPEND TILE_LANG_INCLUDES ${CUDAToolkit_INCLUDE_DIRS})
link_directories(${CUDAToolkit_LIBRARY_DIR} ${CUDAToolkit_LIBRARY_DIR}/stubs)
+1 -1
View File
@@ -1559,7 +1559,7 @@ void CodeGenTileLangCUDA::VisitExpr_(const CastNode *op, std::ostream &os) {
// To add a new type conversion, you should do the following things:
// 1. Add the new conversion function in tl_templates. (__tl_cvt_xx)
// 2. Add a new if statement like the one below.
// 3. In src/backend/common/target_utils.cc, allow this vectorizable cast.
// 3. In src/cuda/target_utils.cc, allow this vectorizable cast.
// Handle conversion from float16 to float32
if (from_ty.is_float16() && target_ty.is_float() && target_ty.bits() == 32) {
+2 -2
View File
@@ -9,15 +9,15 @@
#include <tvm/ir/cast.h>
#include <tvm/runtime/logging.h>
#include "backend/common/target_utils.h"
#include "cuda/op/copy.h"
#include "cuda/target_utils.h"
#include "cuda/transform/ptx_async_copy_injector.h"
#include "layout/tcgen05_layout.h"
#include "op/builtin.h"
#include "op/utils.h"
#include "transform/common/loop_fusion_utils.h"
#include "transform/loop_partition.h"
#include "transform/loop_vectorize.h"
#include "transform/ptx_async_copy_injector.h"
#include <tvm/tirx/analysis.h>
#include <tvm/tirx/builtin.h>
+3 -3
View File
@@ -8,7 +8,7 @@
#include <tvm/ffi/extra/structural_equal.h>
#include <tvm/runtime/logging.h>
#include "backend/common/target_utils.h"
#include "cuda/target_utils.h"
#include "op/builtin.h"
#include "op/utils.h"
@@ -356,7 +356,7 @@ bool CheckCPAsyncCopyPreconditions(const CopyNode &op) {
bool CheckCPAsyncCopy(const CopyNode &op, Target target,
const LayoutMap &layout_map, arith::Analyzer *analyzer) {
if (!TargetHasAsyncCopy(target)) {
if (!TargetCudaHasAsyncCopy(target)) {
return false;
}
if (!CheckCPAsyncCopyPreconditions(op)) {
@@ -472,7 +472,7 @@ std::string MakeAsyncUnavailableReason(const CopyNode &op, Target target) {
std::ostringstream oss;
if (!target.defined()) {
oss << "T.async_copy requires a defined target.";
} else if (!TargetHasAsyncCopy(target)) {
} else if (!TargetCudaHasAsyncCopy(target)) {
oss << "T.async_copy is only supported on targets with cp.async support "
"(SM80+). Got target="
<< target;
+2 -2
View File
@@ -5,7 +5,7 @@
#include "backend/common/op/finalize_reducer.h"
#include "backend/common/target_utils.h"
#include "cuda/target_utils.h"
#include <sstream>
@@ -17,7 +17,7 @@ using namespace tirx;
namespace cuda {
struct FinalizeReducer : backend::FinalizeReducerLowerer<FinalizeReducer> {
static int WarpSize(Target target) { return TargetGetWarpSize(target); }
static int WarpSize(Target target) { return TargetCudaGetWarpSize(target); }
static std::string MakeBatchAllReduce(std::string reducer,
int reducing_threads, int scale,
+3 -3
View File
@@ -7,7 +7,7 @@
#include "support/check.h"
#include <tvm/runtime/logging.h>
#include "backend/common/target_utils.h"
#include "cuda/target_utils.h"
#include "op/builtin.h"
#include "op/tcgen5_meta.h"
#include "op/utils.h"
@@ -87,7 +87,7 @@ bool AllowTcgen5Mma(const GemmNode &op, Target target) {
bool AllowWgmma(const GemmNode &op, int block_size, Target target) {
tvm::transform::PassContext ctxt = tvm::transform::PassContext::Current();
int warp_size = TargetGetWarpSize(target);
int warp_size = TargetCudaGetWarpSize(target);
int num_warps = block_size / warp_size;
return !ctxt->GetConfig(kDisableWGMMA, Optional<Bool>()).value_or(false) &&
TargetIsHopper(target) && op.m_ >= 64 && num_warps % 4 == 0 &&
@@ -289,7 +289,7 @@ struct Gemm {
static std::pair<int, int>
ComputeWarpPartition(const GemmWarpPolicyNode &policy, int M, int N,
int block_size, Target target, String gemm_inst) {
int num_warps = block_size / TargetGetWarpSize(target);
int num_warps = block_size / TargetCudaGetWarpSize(target);
if (gemm_inst == kCudaTCGEN05) {
policy.m_warp = 1;
policy.n_warp = num_warps;
+3 -3
View File
@@ -6,7 +6,7 @@
#include "op/gemm.h"
#include "support/check.h"
#include "backend/common/target_utils.h"
#include "cuda/target_utils.h"
#include "op/builtin.h"
#include "op/tcgen5_meta.h"
#include "op/utils.h"
@@ -80,7 +80,7 @@ bool AllowTcgen5Mma(const GemmSPNode &op, Target target) {
bool AllowWgmma(const GemmSPNode &op, int block_size, Target target) {
tvm::transform::PassContext ctxt = tvm::transform::PassContext::Current();
int warp_size = TargetGetWarpSize(target);
int warp_size = TargetCudaGetWarpSize(target);
int num_warps = block_size / warp_size;
return !ctxt->GetConfig(kDisableWGMMA, Optional<Bool>()).value_or(false) &&
TargetIsHopper(target) && op.M >= 64 && num_warps % 4 == 0 &&
@@ -289,7 +289,7 @@ struct GemmSP {
static std::pair<int, int>
ComputeWarpPartition(const GemmSPWarpPolicyNode &policy, int M, int N,
int block_size, Target target, String gemm_inst) {
int num_warps = block_size / TargetGetWarpSize(target);
int num_warps = block_size / TargetCudaGetWarpSize(target);
if (gemm_inst == kCudaTCGEN05SP) {
policy.m_warp = 1;
policy.n_warp = num_warps;
+268
View File
@@ -0,0 +1,268 @@
/*!
* \file tl/cuda/target_utils.cc
* \brief CUDA target attribute helpers.
*/
#include "cuda/target_utils.h"
#include <tvm/ffi/reflection/registry.h>
#include <string>
#include "dlpack/dlpack.h"
#include "support/check.h"
namespace tvm {
namespace tl {
namespace {
int GetCudaArchInt(Target target) {
auto s = target->GetAttr<ffi::String>("arch");
ICHECK(s.has_value());
const std::string arch_str = s.value();
ICHECK(arch_str.size() >= 3);
ICHECK_EQ(arch_str.compare(0, 3, "sm_"), 0)
<< "arch string must start with sm_";
return std::stoi(arch_str.substr(3));
}
} // namespace
bool TargetIsCuda(Target target) {
return target->GetTargetDeviceType() == kDLCUDA;
}
bool TargetIsCuTeDSL(Target target) {
for (const auto &key : target->keys) {
if (key == "cutedsl")
return true;
}
return false;
}
bool TargetIsVolta(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= 70 && arch < 75;
}
bool TargetIsTuring(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= 75 && arch < 80;
}
bool TargetIsAmpere(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= 80 && arch < 90;
}
bool TargetIsHopper(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= 90 && arch < 100;
}
bool TargetIsSm100(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= 100 && arch <= 110;
}
bool TargetIsSM120(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= 120 && arch < 130;
}
bool TargetCudaHasAsyncCopy(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= 80;
}
int TargetCudaGetWarpSize(Target target) {
(void)target;
return 32;
}
bool TargetHasLdmatrix(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= 75;
}
bool TargetHasStmatrix(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= 90;
}
bool TargetHasTmem(Target target) {
if (!TargetIsCuda(target))
return false;
return TargetIsSm100(target);
}
bool TargetHasBulkCopy(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= 90;
}
bool TargetSupportVectorize256(Target target) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= 100;
}
bool TargetHasSMVersionGE(Target target, int version) {
if (!TargetIsCuda(target))
return false;
int arch = GetCudaArchInt(target);
return arch >= version;
}
bool IsCudaVectorizableFP8(DataType dtype) {
// NOTE: E8M0 is a special type of FP8 which is not handled here.
// We only handle FP8 types which can be represented with
// __nv_fp8_interpretation_t here.
return dtype.is_float8_e4m3() || dtype.is_float8_e4m3fn() ||
dtype.is_float8_e5m2();
}
bool IsCudaVectorizableCast(DataType from_ty, DataType target_ty) {
// float16 -> float32
if (from_ty.is_float16() && target_ty.is_float() && target_ty.bits() == 32)
return true;
// float32 -> float16
if (from_ty.is_float() && from_ty.bits() == 32 && target_ty.is_float16())
return true;
// bfloat16 -> float32
if (from_ty.is_bfloat16() && target_ty.is_float() && target_ty.bits() == 32)
return true;
// float32 -> bfloat16
if (from_ty.is_float() && from_ty.bits() == 32 && target_ty.is_bfloat16())
return true;
// float32 -> float8 (E4M3/E5M2)
if (from_ty.is_float() && from_ty.bits() == 32 &&
IsCudaVectorizableFP8(target_ty))
return true;
// float8 (E4M3/E5M2) -> float32
if (IsCudaVectorizableFP8(from_ty) && target_ty.is_float() &&
target_ty.bits() == 32)
return true;
// Not implemented for now
// float64(double) -> float8 (E4M3/E5M2)
// if (from_ty.is_float() && from_ty.bits() == 64 &&
// IsCudaVectorizableFP8(target_ty))
// return true;
// float8 (E4M3/E5M2) -> float64(double)
// if (IsCudaVectorizableFP8(from_ty) && target_ty.is_float() &&
// target_ty.bits() == 64)
// return true;
// float8 (E8M0) -> bfloat16
if (from_ty.is_float8_e8m0fnu() && target_ty.is_bfloat16())
return true;
// bfloat16 -> float8 (E8M0)
if (from_ty.is_bfloat16() && target_ty.is_float8_e8m0fnu())
return true;
// float32 -> float8 (E8M0)
if (from_ty.is_float() && from_ty.bits() == 32 &&
target_ty.is_float8_e8m0fnu())
return true;
// float64(double) -> float8 (E8M0)
if (from_ty.is_float() && from_ty.bits() == 64 &&
target_ty.is_float8_e8m0fnu())
return true;
// float4_e2m1fn -> float16
if (from_ty.is_float4_e2m1fn() && target_ty.is_float16())
return true;
// float16 -> float4_e2m1fn
if (from_ty.is_float16() && target_ty.is_float4_e2m1fn())
return true;
// float4_e2m1fn -> float32
if (from_ty.is_float4_e2m1fn() && target_ty.is_float() &&
target_ty.bits() == 32)
return true;
// float32 -> float4_e2m1fn
if (from_ty.is_float() && from_ty.bits() == 32 &&
target_ty.is_float4_e2m1fn())
return true;
// float4_e2m1fn -> float64(double)
if (from_ty.is_float4_e2m1fn() && target_ty.is_float() &&
target_ty.bits() == 64)
return true;
// float64(double) -> float4_e2m1fn
if (from_ty.is_float() && from_ty.bits() == 64 &&
target_ty.is_float4_e2m1fn())
return true;
// float4_e2m1fn -> bfloat16
if (from_ty.is_float4_e2m1fn() && target_ty.is_bfloat16())
return true;
// bfloat16 -> float4_e2m1fn
if (from_ty.is_bfloat16() && target_ty.is_float4_e2m1fn())
return true;
return false;
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::GlobalDef()
.def("tl.TargetIsCuda",
[](Target target) { return TargetIsCuda(target); })
.def("tl.TargetIsVolta",
[](Target target) { return TargetIsVolta(target); })
.def("tl.TargetIsTuring",
[](Target target) { return TargetIsTuring(target); })
.def("tl.TargetIsAmpere",
[](Target target) { return TargetIsAmpere(target); })
.def("tl.TargetIsHopper",
[](Target target) { return TargetIsHopper(target); })
.def("tl.TargetIsSM120",
[](Target target) { return TargetIsSM120(target); })
.def("tl.TargetCudaGetWarpSize",
[](Target target) { return TargetCudaGetWarpSize(target); })
.def("tl.TargetHasLdmatrix",
[](Target target) { return TargetHasLdmatrix(target); })
.def("tl.TargetHasStmatrix",
[](Target target) { return TargetHasStmatrix(target); })
.def("tl.TargetHasBulkCopy",
[](Target target) { return TargetHasBulkCopy(target); });
}
} // namespace tl
} // namespace tvm
+40
View File
@@ -0,0 +1,40 @@
/*!
* \file tl/cuda/target_utils.h
* \brief CUDA target attribute helpers.
*/
#ifndef TVM_TL_CUDA_TARGET_UTILS_H_
#define TVM_TL_CUDA_TARGET_UTILS_H_
#include <tvm/runtime/data_type.h>
#include <tvm/target/target.h>
namespace tvm {
namespace tl {
bool TargetIsCuda(Target target);
bool TargetIsCuTeDSL(Target target);
bool TargetIsVolta(Target target);
bool TargetIsTuring(Target target);
bool TargetIsAmpere(Target target);
bool TargetIsHopper(Target target);
bool TargetIsSm100(Target target);
bool TargetIsSM120(Target target);
bool TargetCudaHasAsyncCopy(Target target);
int TargetCudaGetWarpSize(Target target);
bool TargetHasLdmatrix(Target target);
bool TargetHasStmatrix(Target target);
bool TargetHasTmem(Target target);
bool TargetHasBulkCopy(Target target);
bool TargetSupportVectorize256(Target target);
bool TargetHasSMVersionGE(Target target, int version);
bool IsCudaVectorizableFP8(DataType dtype);
bool IsCudaVectorizableCast(DataType from_ty, DataType target_ty);
} // namespace tl
} // namespace tvm
#endif // TVM_TL_CUDA_TARGET_UTILS_H_
@@ -0,0 +1,57 @@
/*!
* \brief Lower eligible global->shared copies into PTX cp.async.
* \file cuda/transform/lower_ptx_async_copy.cc
*/
#include <tvm/ffi/reflection/registry.h>
#include <tvm/target/target.h>
#include <tvm/tirx/transform.h>
#include "cuda/target_utils.h"
#include "cuda/transform/ptx_async_copy_injector.h"
#include "op/builtin.h"
namespace tvm {
namespace tl {
using namespace tirx;
using namespace tirx::transform;
tvm::transform::Pass LowerPTXAsyncCopy() {
auto pass_func = [=](PrimFunc f, const IRModule &m, const PassContext &ctx) {
auto target_opt = f->GetAttr<Target>(tvm::attr::kTarget);
if (!target_opt.defined()) {
return f;
}
Target target = target_opt.value();
if (!TargetIsCuda(target)) {
return f;
}
if (!TargetCudaHasAsyncCopy(target)) {
// Graceful fallback on older architectures.
return f;
}
bool enable_auto_async_copy =
ctx->GetConfig<Bool>(kEnableAsyncCopy, Bool(true)).value();
auto *n = f.CopyOnWrite();
auto inject_result =
InjectPTXAsyncCopy(n->body, enable_auto_async_copy,
/*async_without_async_commit_wait=*/false);
n->body = inject_result.stmt;
return f;
};
return CreatePrimFuncPass(pass_func, 0, "tl.cuda.transform.LowerPTXAsyncCopy",
{});
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::GlobalDef().def("tl.cuda.transform.LowerPTXAsyncCopy",
LowerPTXAsyncCopy);
}
} // namespace tl
} // namespace tvm
+2 -2
View File
@@ -3,7 +3,7 @@
* \brief Convert shared.tmem buffers to plain shared + ptx init, and do
* coordinate translation (from logical address to physical address)
*/
#include "backend/common/target_utils.h"
#include "cuda/target_utils.h"
#include "op/builtin.h"
#include "support/check.h"
#include "tvm/ir/type.h"
@@ -284,7 +284,7 @@ private:
Array<Stmt> new_body;
ICHECK(target_.defined()) << "LowerSharedTmem requires a bound target";
auto warp_size = TargetGetWarpSize(target_);
auto warp_size = TargetCudaGetWarpSize(target_);
auto thread_var_div_warp_size =
FloorDiv(thread_var_->var, IntImm(thread_var_->var->dtype, warp_size));
new_body.push_back(IfThenElse(EQ(thread_var_div_warp_size, 0),
@@ -1,11 +1,10 @@
/*!
* \brief Lower eligible global->shared copies into PTX cp.async
* \file lower_ptx_async_copy.cc
* \brief Inject eligible global->shared copies into PTX cp.async intrinsics.
* \file cuda/transform/ptx_async_copy_injector.cc
*/
#include "support/check.h"
#include <tvm/runtime/logging.h>
#include <tvm/s_tir/stmt.h>
#include <tvm/target/target.h>
#include <tvm/tirx/analysis.h>
#include <tvm/tirx/builtin.h>
#include <tvm/tirx/expr.h>
@@ -19,12 +18,10 @@
#include <optional>
#include <vector>
#include "../op/builtin.h"
#include "../op/utils.h"
#include "backend/common/target_utils.h"
#include "ptx_async_copy_injector.h"
#include "cuda/transform/ptx_async_copy_injector.h"
#include "op/builtin.h"
#include "op/utils.h"
#include "tir/ir/buffer_common.h"
#include <tvm/tirx/stmt.h>
namespace tvm {
namespace tl {
@@ -700,8 +697,6 @@ private:
bool uncommitted_sync_copies_{false};
};
using namespace tirx::transform;
PTXAsyncCopyInjectResult
InjectPTXAsyncCopy(const Stmt &body, bool enable_auto_async_copy,
bool async_without_async_commit_wait) {
@@ -711,39 +706,5 @@ InjectPTXAsyncCopy(const Stmt &body, bool enable_auto_async_copy,
return {injector.Finalize(injected), injector.InjectedPTXAsyncCopy()};
}
tvm::transform::Pass LowerPTXAsyncCopy() {
auto pass_func = [=](PrimFunc f, const IRModule &m, const PassContext &ctx) {
auto target_opt = f->GetAttr<Target>(tvm::attr::kTarget);
if (!target_opt.defined()) {
return f;
}
Target target = target_opt.value();
if (!TargetIsCuda(target)) {
return f;
}
if (!TargetHasAsyncCopy(target)) {
// Graceful fallback on older architectures.
return f;
}
bool enable_auto_async_copy =
ctx->GetConfig<Bool>(kEnableAsyncCopy, Bool(true)).value();
auto *n = f.CopyOnWrite();
auto inject_result =
InjectPTXAsyncCopy(n->body, enable_auto_async_copy,
/*async_without_async_commit_wait=*/false);
n->body = inject_result.stmt;
return f;
};
return CreatePrimFuncPass(pass_func, 0, "tl.LowerPTXAsyncCopy", {});
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = reflection;
refl::GlobalDef().def("tl.transform.LowerPTXAsyncCopy", LowerPTXAsyncCopy);
}
} // namespace tl
} // namespace tvm
@@ -13,8 +13,8 @@ struct PTXAsyncCopyInjectResult {
/*! \brief Inject PTX cp.async lowering patterns into a statement.
*
* This is the statement-level entrypoint used by other transforms to apply the
* same rewrite as the `tl.LowerPTXAsyncCopy` pass, but scoped to a region
* (e.g., a lowered parallel loop) rather than the whole PrimFunc.
* same rewrite as CUDA PTX async-copy passes, but scoped to a region (e.g.,
* a lowered parallel loop) rather than the whole PrimFunc.
*/
PTXAsyncCopyInjectResult
InjectPTXAsyncCopy(const tvm::tirx::Stmt &body, bool enable_auto_async_copy,
+17 -4
View File
@@ -1,20 +1,33 @@
# Metal backend: source files and build configuration.
#
# Metal source generation and operator lowering are pure C++ and can be used
# for cross-compilation on non-Apple hosts. Only optional runtime integration
# is gated by USE_METAL / APPLE below.
file(GLOB TILE_LANG_METAL_ALWAYS_SRCS
src/metal/codegen/codegen_metal.cc
src/metal/op/*.cc
src/metal/target_utils.cc
)
list(APPEND TILE_LANG_SRCS ${TILE_LANG_METAL_ALWAYS_SRCS})
if(NOT USE_METAL)
return()
endif()
message(STATUS "METAL Backend is enabled")
# FIXME: CIBW failed with backtrace, why???
set(TVM_FFI_USE_LIBBACKTRACE OFF)
if(NOT APPLE)
# On non-Apple platforms USE_METAL=ON enables only codegen (Metal source
# generation) without requiring the Metal/Foundation frameworks.
message(STATUS "Metal backend on non-Apple: enabling codegen-only mode (no Metal runtime)")
set(USE_METAL OFF)
return()
endif()
file(GLOB TILE_LANG_METAL_SRCS
file(GLOB TILE_LANG_METAL_ACTIVE_SRCS
src/metal/codegen/rt_mod_metal.cc
)
list(APPEND TILE_LANG_SRCS ${TILE_LANG_METAL_SRCS})
# FIXME: CIBW failed with backtrace, why???
set(TVM_FFI_USE_LIBBACKTRACE OFF)
list(APPEND TILE_LANG_SRCS ${TILE_LANG_METAL_ACTIVE_SRCS})
@@ -18,7 +18,7 @@
*/
/*!
* \file codegen_metal.cc
* \file metal/codegen/codegen_metal.cc
*/
#include "codegen_metal.h"
@@ -18,11 +18,11 @@
*/
/*!
* \file codegen_metal.h
* \file metal/codegen/codegen_metal.h
* \brief Generate Metal device code.
*/
#ifndef TVM_TARGET_SOURCE_CODEGEN_METAL_H_
#define TVM_TARGET_SOURCE_CODEGEN_METAL_H_
#ifndef TILELANG_METAL_CODEGEN_CODEGEN_METAL_H_
#define TILELANG_METAL_CODEGEN_CODEGEN_METAL_H_
#include <tvm/target/codegen.h>
@@ -71,4 +71,4 @@ private:
} // namespace codegen
} // namespace tvm
#endif // TVM_TARGET_SOURCE_CODEGEN_METAL_H_
#endif // TILELANG_METAL_CODEGEN_CODEGEN_METAL_H_
+4 -5
View File
@@ -2,17 +2,16 @@
* \file rt_mod_metal.cc
* \brief Metal codegen entry point.
*
* Metal codegen is implemented in backend/common/codegen/codegen_metal.cc,
* which handles simdgroup types, intrinsics, and MSL emission. This file exists
* to satisfy the metal/CMakeLists.txt dependency but delegates to the main
* implementation.
* Metal source codegen is implemented in codegen_metal.cc, which handles
* simdgroup types, intrinsics, and MSL emission. This file exists to satisfy
* the metal/CMakeLists.txt dependency.
*/
#include "support/check.h"
namespace tvm {
namespace codegen {
// Metal codegen entry point is in backend/common/codegen/codegen_metal.cc.
// Metal codegen entry point is in codegen_metal.cc.
} // namespace codegen
} // namespace tvm
+2 -2
View File
@@ -5,8 +5,8 @@
#include "op/copy.h"
#include "backend/common/target_utils.h"
#include "metal/op/utils.h"
#include "metal/target_utils.h"
#include "op/utils.h"
#include <tvm/tirx/builtin.h>
@@ -48,7 +48,7 @@ Stmt LowerSIMDGroupCopy(const CopyNode &op, const LowerArgs &T,
PrimExpr dst_col_base = op.dst_range[1]->min;
PrimExpr dst_stride = op.dst->shape[op.dst->shape.size() - 1];
int warp_size = TargetGetWarpSize(T.target);
int warp_size = TargetMetalGetWarpSize(T.target);
const auto *block_size_imm = T.thread_bounds->extent.as<IntImmNode>();
TVM_FFI_ICHECK(block_size_imm)
<< "simdgroup copy requires constant thread bounds";
+2 -2
View File
@@ -5,8 +5,8 @@
#include "op/gemm.h"
#include "backend/common/target_utils.h"
#include "metal/op/utils.h"
#include "metal/target_utils.h"
#include <tvm/runtime/logging.h>
@@ -89,7 +89,7 @@ struct Gemm {
int block_size, Target target, String gemm_inst) {
TVM_FFI_ICHECK(gemm_inst == kMetalSIMDGroup)
<< "Unsupported Metal GEMM instruction: " << gemm_inst;
int num_warps = block_size / TargetGetWarpSize(target);
int num_warps = block_size / TargetMetalGetWarpSize(target);
return ComputeMetalWarpPartition(policy, M, N, num_warps);
}
+34
View File
@@ -0,0 +1,34 @@
/*!
* \file tl/metal/target_utils.cc
* \brief Metal target attribute helpers.
*/
#include "metal/target_utils.h"
#include <tvm/ffi/reflection/registry.h>
#include "dlpack/dlpack.h"
namespace tvm {
namespace tl {
bool TargetIsMetal(Target target) {
return target->GetTargetDeviceType() == kDLMetal;
}
int TargetMetalGetWarpSize(Target target) {
(void)target;
return 32;
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::GlobalDef()
.def("tl.TargetIsMetal",
[](Target target) { return TargetIsMetal(target); })
.def("tl.TargetMetalGetWarpSize",
[](Target target) { return TargetMetalGetWarpSize(target); });
}
} // namespace tl
} // namespace tvm
+20
View File
@@ -0,0 +1,20 @@
/*!
* \file tl/metal/target_utils.h
* \brief Metal target attribute helpers.
*/
#ifndef TVM_TL_METAL_TARGET_UTILS_H_
#define TVM_TL_METAL_TARGET_UTILS_H_
#include <tvm/target/target.h>
namespace tvm {
namespace tl {
bool TargetIsMetal(Target target);
int TargetMetalGetWarpSize(Target target);
} // namespace tl
} // namespace tvm
#endif // TVM_TL_METAL_TARGET_UTILS_H_
+13 -3
View File
@@ -1,4 +1,13 @@
# ROCm backend: toolchain, stub libraries, source files, and build configuration.
# ROCm backend: source files, toolchain, stubs, and build configuration.
#
# ROCm target helpers and intrinsic rules are pure C++ registrations used by
# common code and source generation. Compile them regardless of USE_ROCM.
file(GLOB TILE_LANG_ROCM_ALWAYS_SRCS
src/rocm/codegen/intrin_rule_hip.cc
src/rocm/target_utils.cc
)
list(APPEND TILE_LANG_SRCS ${TILE_LANG_ROCM_ALWAYS_SRCS})
if(NOT USE_ROCM)
return()
endif()
@@ -118,12 +127,13 @@ if(TILELANG_USE_HIP_STUBS)
"HSA runtime library to link against" FORCE)
endif()
file(GLOB TILE_LANG_HIP_SRCS
file(GLOB TILE_LANG_ROCM_ACTIVE_SRCS
src/rocm/codegen/codegen_hip.cc
src/rocm/codegen/rt_mod_hip.cc
src/rocm/op/*.cc
src/rocm/transform/*.cc
)
list(APPEND TILE_LANG_SRCS ${TILE_LANG_HIP_SRCS})
list(APPEND TILE_LANG_SRCS ${TILE_LANG_ROCM_ACTIVE_SRCS})
list(APPEND TILE_LANG_INCLUDES ${ROCM_INCLUDE_DIRS})
# Register stubs for linking and install
+21 -19
View File
@@ -8,12 +8,12 @@
#include <tvm/ir/cast.h>
#include <tvm/runtime/logging.h>
#include "backend/common/target_utils.h"
#include "op/builtin.h"
#include "op/utils.h"
#include "rocm/target_utils.h"
#include "rocm/transform/async_copy_injector.h"
#include "transform/common/loop_fusion_utils.h"
#include "transform/loop_partition.h"
#include "transform/ptx_async_copy_injector.h"
#include <tvm/tirx/builtin.h>
#include <tvm/tirx/transform.h>
@@ -133,13 +133,13 @@ private:
/*should_vectorize=*/true, par_op->LoopLayoutRequiresPaddingGuard());
auto inject_result =
InjectPTXAsyncCopy(lowered_loop, /*enable_auto_async_copy=*/true,
/*async_without_async_commit_wait=*/
no_implicit_commit_wait || GetIsAsyncCopy(op));
Stmt cp_async_loop = inject_result.stmt;
if (!inject_result.injected_ptx_async_copy) {
DLOG(WARNING) << "cp.async rewrite miss for copy src=" << op.src->name
<< " (scope=" << op.src.scope()
InjectROCmAsyncCopy(lowered_loop, /*enable_auto_async_copy=*/true,
/*async_without_async_commit_wait=*/
no_implicit_commit_wait || GetIsAsyncCopy(op));
Stmt async_copy_loop = inject_result.stmt;
if (!inject_result.injected_rocm_async_copy) {
DLOG(WARNING) << "ROCm async-copy rewrite miss for copy src="
<< op.src->name << " (scope=" << op.src.scope()
<< ", dtype=" << op.src->dtype << "), dst=" << op.dst->name
<< " (scope=" << op.dst.scope()
<< ", dtype=" << op.dst->dtype
@@ -149,27 +149,29 @@ private:
if (no_implicit_commit_wait) {
DLOG(WARNING)
<< "Pipeline-managed async copy fallback to normal copy because "
"cp.async rewrite found no eligible global->shared store.";
"ROCm async-copy rewrite found no eligible global->shared "
"store.";
return lowered_loop;
}
if (explicit_async_semantics) {
LOG(FATAL)
<< "Explicit async copy semantics require cp.async lowering, "
"but no eligible global->shared store was rewritten.";
LOG(FATAL) << "Explicit async copy semantics require ROCm async-copy "
"lowering, "
"but no eligible global->shared store was rewritten.";
}
DLOG(WARNING) << "Fallback to normal copy because cp.async rewrite found "
"no eligible global->shared store.";
DLOG(WARNING)
<< "Fallback to normal copy because ROCm async-copy rewrite "
"found no eligible global->shared store.";
return LowerNormalCopy(op, T, analyzer);
}
if (no_implicit_commit_wait) {
return cp_async_loop;
return async_copy_loop;
}
if (GetIsAsyncCopy(op)) {
Stmt commit_group =
Evaluate(Call(DataType::Handle(), builtin::ptx_commit_group(), {}));
return SeqStmt({cp_async_loop, commit_group});
return SeqStmt({async_copy_loop, commit_group});
}
return cp_async_loop;
return async_copy_loop;
}
static bool CheckCPAsyncCopyPreconditions(const CopyNode &op) {
@@ -185,7 +187,7 @@ private:
static bool CheckCPAsyncCopy(const CopyNode &op, Target target,
const LayoutMap &layout_map,
arith::Analyzer *analyzer) {
if (!TargetHasAsyncCopy(target)) {
if (!TargetRocmHasAsyncCopy(target)) {
return false;
}
return CheckCPAsyncCopyPreconditions(op);
+2 -2
View File
@@ -5,7 +5,7 @@
#include "backend/common/op/finalize_reducer.h"
#include "backend/common/target_utils.h"
#include "rocm/target_utils.h"
#include <sstream>
@@ -17,7 +17,7 @@ using namespace tirx;
namespace rocm {
struct FinalizeReducer : backend::FinalizeReducerLowerer<FinalizeReducer> {
static int WarpSize(Target target) { return TargetGetWarpSize(target); }
static int WarpSize(Target target) { return TargetRocmGetWarpSize(target); }
static std::string MakeBatchAllReduce(std::string reducer,
int reducing_threads, int scale,
+2 -2
View File
@@ -7,7 +7,7 @@
#include "support/check.h"
#include <tvm/runtime/logging.h>
#include "backend/common/target_utils.h"
#include "rocm/target_utils.h"
#include <cmath>
#include <limits>
@@ -118,7 +118,7 @@ struct Gemm {
ComputeWarpPartition(const GemmWarpPolicyNode &policy, int M, int N,
int block_size, Target target, ffi::String gemm_inst) {
(void)gemm_inst;
int num_warps = block_size / TargetGetWarpSize(target);
int num_warps = block_size / TargetRocmGetWarpSize(target);
return ComputeDefaultWarpPartition(policy, M, N, num_warps);
}
+105
View File
@@ -0,0 +1,105 @@
/*!
* \file tl/rocm/target_utils.cc
* \brief ROCm target attribute helpers.
*/
#include "rocm/target_utils.h"
#include <tvm/ffi/reflection/registry.h>
#include <tvm/ir/cast.h>
#include <string>
#include "dlpack/dlpack.h"
namespace tvm {
namespace tl {
bool TargetIsRocm(Target target) {
return target->GetTargetDeviceType() == kDLROCM;
}
bool TargetIsCDNA(Target target) {
if (!TargetIsRocm(target))
return false;
if (target->attrs.count("mcpu")) {
std::string mcpu = Downcast<ffi::String>(target->attrs.at("mcpu"));
// if mcpu start with "gfx9", it is CDNA
return mcpu.find("gfx9") == 0;
}
return false;
}
bool TargetIsRDNA(Target target) {
if (!TargetIsRocm(target))
return false;
if (target->attrs.count("mcpu")) {
std::string mcpu = Downcast<ffi::String>(target->attrs.at("mcpu"));
// gfx11xx, gfx12xx are RDNA architectures
return mcpu.find("gfx11") == 0 || mcpu.find("gfx12") == 0;
}
return false;
}
bool TargetIsGfx950(Target target) {
if (!TargetIsRocm(target))
return false;
if (target->attrs.count("mcpu")) {
std::string mcpu = Downcast<ffi::String>(target->attrs.at("mcpu"));
return mcpu.find("gfx950") != std::string::npos;
}
return false;
}
bool TargetRocmHasAsyncCopy(Target target) {
if (!TargetIsCDNA(target))
return false;
if (target->attrs.count("mcpu")) {
std::string mcpu = Downcast<ffi::String>(target->attrs.at("mcpu"));
if (mcpu.rfind("gfx9", 0) == 0) {
int gfx_version = std::stoi(mcpu.substr(3, 2));
return gfx_version >= 94;
}
}
return false;
}
int TargetRocmGetWarpSize(Target target) {
if (TargetIsCDNA(target)) {
return 64;
}
return 32;
}
int TargetGetRDNAGeneration(Target target) {
if (!TargetIsRDNA(target))
return 0;
if (target->attrs.count("mcpu")) {
std::string mcpu = Downcast<ffi::String>(target->attrs.at("mcpu"));
if (mcpu.rfind("gfx11", 0) == 0)
return 11;
if (mcpu.rfind("gfx12", 0) == 0)
return 12;
}
return 0;
}
TVM_FFI_STATIC_INIT_BLOCK() {
namespace refl = tvm::ffi::reflection;
refl::GlobalDef()
.def("tl.TargetIsRocm",
[](Target target) { return TargetIsRocm(target); })
.def("tl.TargetIsCDNA",
[](Target target) { return TargetIsCDNA(target); })
.def("tl.TargetIsRDNA",
[](Target target) { return TargetIsRDNA(target); })
.def("tl.TargetIsGfx950",
[](Target target) { return TargetIsGfx950(target); })
.def("tl.TargetRocmGetWarpSize",
[](Target target) { return TargetRocmGetWarpSize(target); })
.def("tl.TargetGetRDNAGeneration",
[](Target target) { return TargetGetRDNAGeneration(target); });
}
} // namespace tl
} // namespace tvm
+26
View File
@@ -0,0 +1,26 @@
/*!
* \file tl/rocm/target_utils.h
* \brief ROCm target attribute helpers.
*/
#ifndef TVM_TL_ROCM_TARGET_UTILS_H_
#define TVM_TL_ROCM_TARGET_UTILS_H_
#include <tvm/target/target.h>
namespace tvm {
namespace tl {
bool TargetIsRocm(Target target);
bool TargetIsCDNA(Target target);
bool TargetIsRDNA(Target target);
bool TargetIsGfx950(Target target);
bool TargetRocmHasAsyncCopy(Target target);
int TargetRocmGetWarpSize(Target target);
int TargetGetRDNAGeneration(Target target);
} // namespace tl
} // namespace tvm
#endif // TVM_TL_ROCM_TARGET_UTILS_H_
+712
View File
@@ -0,0 +1,712 @@
/*!
* \brief Inject eligible global->shared copies into ROCm async-copy intrinsics.
* \file rocm/transform/async_copy_injector.cc
*/
#include "support/check.h"
#include <tvm/runtime/logging.h>
#include <tvm/s_tir/stmt.h>
#include <tvm/tirx/analysis.h>
#include <tvm/tirx/builtin.h>
#include <tvm/tirx/expr.h>
#include <tvm/tirx/op.h>
#include <tvm/tirx/stmt_functor.h>
#include <tvm/tirx/transform.h>
#include <algorithm>
#include <cstdint>
#include <limits>
#include <optional>
#include <vector>
#include "op/builtin.h"
#include "op/utils.h"
#include "rocm/transform/async_copy_injector.h"
#include "tir/ir/buffer_common.h"
namespace tvm {
namespace tl {
using namespace tirx;
using namespace ffi;
class ROCmAsyncCopyInjector : public StmtMutator {
public:
explicit ROCmAsyncCopyInjector(bool enable_auto_async_copy,
bool async_without_async_commit_wait)
: enable_auto_async_copy_(enable_auto_async_copy),
async_without_async_commit_wait_(async_without_async_commit_wait) {}
bool InjectedROCmAsyncCopy() const { return injected_rocm_async_copy_; }
Stmt Finalize(Stmt body) {
if (!pending_sync_copies_ || UseExplicitAsyncSemantics()) {
pending_sync_copies_ = false;
uncommitted_sync_copies_ = false;
return body;
}
Array<Stmt> seq;
seq.reserve(3);
seq.push_back(body);
AppendSyncVisibility(&seq, uncommitted_sync_copies_);
pending_sync_copies_ = false;
uncommitted_sync_copies_ = false;
return SeqStmt(seq);
}
Stmt VisitStmt_(const AttrStmtNode *op) final {
if (op->attr_key == s_tir::attr::async_scope) {
++explicit_async_scope_depth_;
Stmt body = this->VisitStmt(op->body);
--explicit_async_scope_depth_;
// `async_scope` is a lowering-only marker for cp.async semantics.
return body;
}
return StmtMutator::VisitStmt_(op);
}
Stmt VisitStmt_(const ForNode *op) final {
// Track nested vectorized loop extents so we can decide whether an
// element-wise copy has a legal final cp.async width after later loop
// vectorization:
// for v in T.vectorized(k): tl.ptx_cp_async(dst, src, elem_count)
// => tl.ptx_cp_async(dst_base, src_base, elem_count * k)
//
// TileLang currently uses tl.ptx_cp_async as the async-copy IR marker.
// HIP codegen maps it to ROCm async-copy intrinsics. The final byte width
// is derived later from the access_ptr dtype, so subbyte dtypes such as
// int4/fp4/int2/int1 remain representable here.
int previous_vectorized_lanes = current_vectorized_lanes_;
bool pushed_vectorized_loop = false;
if (op->kind == 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;
active_vectorized_loops_.push_back({op->loop_var, lanes});
pushed_vectorized_loop = true;
}
}
Stmt stmt = StmtMutator::VisitStmt_(op);
if (pushed_vectorized_loop) {
active_vectorized_loops_.pop_back();
}
current_vectorized_lanes_ = previous_vectorized_lanes;
return stmt;
}
Optional<Stmt>
TryInjectROCmAsyncCopy(const BufferLoadNode *load,
const BufferStoreNode *store, bool predicated = false,
const PrimExpr &predicate_value = PrimExpr()) {
// Pipeline:
// 1) Analyze source/destination indices and transfer width eligibility.
// 2) Build the async-copy IR marker with scalar/vectorized base offsets
// when the eventual transfer byte width is representable.
std::optional<CopyIndexInfo> index_info = PrepareCopyIndexInfo(load, store);
if (!index_info.has_value()) {
return Optional<Stmt>();
}
if (index_info->index_lanes == 1) {
if (current_vectorized_lanes_ > 1 &&
!HasContiguousVectorizedOffsets(index_info->src_index,
index_info->dst_index)) {
return Optional<Stmt>();
}
return MakeCPAsyncStmtFromLoads(
store,
/*dst_base_load=*/BufferLoad(store->buffer, store->indices),
/*src_base_load=*/BufferLoad(load->buffer, load->indices),
/*num_elems=*/index_info->per_access_num_elems, predicated,
predicate_value);
}
Optional<Array<PrimExpr>> src_base_indices =
ExtractVectorBaseIndices(load->indices);
Optional<Array<PrimExpr>> dst_base_indices =
ExtractVectorBaseIndices(store->indices);
if (!src_base_indices.defined() || !dst_base_indices.defined()) {
// If we can't extract base indices from vectorized accesses, fall back.
if (predicated) {
LOG(WARNING)
<< "Cannot extract base indices from vectorized accesses for "
"predicated cp.async; falling back to regular buffer store/load";
}
return Optional<Stmt>();
}
return MakeCPAsyncStmtFromLoads(
store,
/*dst_base_load=*/BufferLoad(store->buffer, dst_base_indices.value()),
/*src_base_load=*/BufferLoad(load->buffer, src_base_indices.value()),
/*num_elems=*/index_info->per_access_num_elems, predicated,
predicate_value);
}
Stmt VisitStmt_(const SeqStmtNode *op) final {
if (UseExplicitAsyncSemantics()) {
return StmtMutator::VisitStmt_(op);
}
// Insert commit+wait at statement boundaries to preserve synchronous
// semantics for normal global->shared BufferStore copies.
//
// Important: avoid flushing inside inner loop bodies just because there
// are trailing no-op statements (e.g., Evaluate(0)) after the injected
// cp.async. Instead, treat "pure copy region" statements as part of the
// copy run and only flush right before the next non-copy statement.
Array<Stmt> out;
out.reserve(op->seq.size() + 2);
CopySyncState sync_state{pending_sync_copies_, uncommitted_sync_copies_};
pending_sync_copies_ = false;
uncommitted_sync_copies_ = false;
for (const Stmt &stmt : op->seq) {
VisitedStmtInfo visited_info = VisitAndAnalyzeStmt(stmt);
bool stmt_is_pure_copy_region = visited_info.analysis.is_pure_copy_region;
// Before we execute a non-copy statement, we must preserve synchronous
// semantics for injected cp.async stores by making the data visible.
if (sync_state.open_copy_region && !stmt_is_pure_copy_region) {
AppendSyncVisibility(&out, sync_state.uncommitted_transfers);
sync_state.open_copy_region = false;
sync_state.uncommitted_transfers = false;
}
// If we are carrying uncommitted injected cp.async into an explicit wait,
// ensure they are committed so the wait actually covers them.
if (sync_state.open_copy_region && sync_state.uncommitted_transfers &&
visited_info.analysis.wait > 0) {
out.push_back(MakeCommitGroupStmt());
sync_state.uncommitted_transfers = false;
}
out.push_back(visited_info.visited);
if (visited_info.opens_copy_region) {
sync_state.open_copy_region = true;
sync_state.uncommitted_transfers =
sync_state.uncommitted_transfers ||
visited_info.has_uncommitted_transfers;
}
if (visited_info.analysis.commit > 0) {
// A commit closes the currently open group, so there are no longer any
// uncommitted injected cp.async transfers.
sync_state.uncommitted_transfers = false;
}
if (visited_info.analysis.wait > 0) {
// Any explicit wait serves as a synchronization boundary for injected
// synchronous copies.
sync_state.open_copy_region = false;
sync_state.uncommitted_transfers = false;
}
}
pending_sync_copies_ = sync_state.open_copy_region;
uncommitted_sync_copies_ = sync_state.uncommitted_transfers;
if (out.empty()) {
return Evaluate(0);
}
if (out.size() == 1) {
return out[0];
}
return SeqStmt(out);
}
Stmt VisitStmt_(const IfThenElseNode *op) final {
if (UseExplicitAsyncSemantics()) {
return StmtMutator::VisitStmt_(op);
}
// Treat branches as separate control flow paths. We propagate pending
// synchronous copies into both branches (they occur before the branch),
// but do not let mutations in one branch affect the other.
bool pending_before = pending_sync_copies_;
bool uncommitted_before = uncommitted_sync_copies_;
pending_sync_copies_ = pending_before;
uncommitted_sync_copies_ = uncommitted_before;
Stmt then_case = this->VisitStmt(op->then_case);
bool pending_then = pending_sync_copies_;
bool uncommitted_then = uncommitted_sync_copies_;
bool pending_else = pending_before;
bool uncommitted_else = uncommitted_before;
Optional<Stmt> else_case;
if (op->else_case.defined()) {
pending_sync_copies_ = pending_before;
uncommitted_sync_copies_ = uncommitted_before;
else_case = this->VisitStmt(op->else_case.value());
pending_else = pending_sync_copies_;
uncommitted_else = uncommitted_sync_copies_;
}
pending_sync_copies_ = pending_then || pending_else;
uncommitted_sync_copies_ = uncommitted_then || uncommitted_else;
if (then_case.same_as(op->then_case) &&
(!else_case.defined() || else_case.same_as(op->else_case))) {
return GetRef<Stmt>(op);
}
return IfThenElse(op->condition, then_case, else_case);
}
Stmt VisitStmt_(const BufferStoreNode *store) final {
if (!IsSharedBuffer(store->buffer)) {
return StmtMutator::VisitStmt_(store);
}
// Only lower copies in regions where async-copy rewrite is enabled.
if (!enable_auto_async_copy_) {
return StmtMutator::VisitStmt_(store);
}
Optional<PrimExpr> predicate = std::nullopt;
const BufferLoadNode *load =
MatchZeroFillBufferLoad(store->value, &predicate);
if (load) {
Optional<Stmt> injected = TryInjectROCmAsyncCopy(
load, store, predicate.defined(),
predicate.defined() ? predicate.value() : PrimExpr());
if (injected.defined()) {
injected_rocm_async_copy_ = true;
if (!UseExplicitAsyncSemantics()) {
pending_sync_copies_ = true;
uncommitted_sync_copies_ = true;
}
return injected.value();
}
}
return StmtMutator::VisitStmt_(store);
}
private:
bool UseExplicitAsyncSemantics() const {
return async_without_async_commit_wait_ || explicit_async_scope_depth_ > 0;
}
// A copy candidate represented after flattening source/destination indexing.
struct CopyIndexInfo {
PrimExpr src_index;
PrimExpr dst_index;
int index_lanes{1};
int per_access_num_elems{0};
};
// Synchronization state for injected cp.async runs carried across statements.
struct CopySyncState {
bool open_copy_region{false};
bool uncommitted_transfers{false};
};
struct ActiveVectorizedLoop {
Var loop_var;
int extent;
};
// ---- Copy candidate analysis helpers ----
static bool IsZeroValue(const PrimExpr &expr) {
if (const auto *broadcast = expr.as<BroadcastNode>()) {
return IsZeroValue(broadcast->value);
}
if (const auto *float_imm = expr.as<FloatImmNode>()) {
return float_imm->value == 0.0f;
}
if (const auto *int_imm = expr.as<IntImmNode>()) {
return int_imm->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;
}
*predicate =
predicate->defined()
? Optional<PrimExpr>(And(call->args[0], predicate->value()))
: Optional<PrimExpr>(call->args[0]);
return load;
}
static Optional<PrimExpr>
FlattenToLinearOffset(const Buffer &buf, const Array<PrimExpr> &indices) {
// Convert N-D indices (potentially with axis_separators) into a single
// row-major linear element offset.
Array<PrimExpr> physical = buf.OffsetOf(indices);
Buffer flattened_buf = buf.GetFlattenedBuffer();
if (physical.size() != flattened_buf->shape.size() || physical.empty()) {
return Optional<PrimExpr>();
}
PrimExpr linear = physical[0];
for (size_t i = 1; i < physical.size(); ++i) {
linear = linear * flattened_buf->shape[i] + physical[i];
}
return linear;
}
std::optional<CopyIndexInfo>
PrepareCopyIndexInfo(const BufferLoadNode *load,
const BufferStoreNode *store) {
if (!IsGlobalBuffer(load->buffer)) {
return std::nullopt;
}
Optional<PrimExpr> src_index_opt =
FlattenToLinearOffset(load->buffer, load->indices);
Optional<PrimExpr> dst_index_opt =
FlattenToLinearOffset(store->buffer, store->indices);
if (!src_index_opt.defined() || !dst_index_opt.defined()) {
return std::nullopt;
}
PrimExpr src_index = src_index_opt.value();
PrimExpr dst_index = dst_index_opt.value();
if (src_index->dtype.lanes() != dst_index->dtype.lanes()) {
// Not a straightforward vectorized copy; skip.
return std::nullopt;
}
const int index_lanes = src_index->dtype.lanes();
const int value_lanes = load->dtype.lanes();
if (value_lanes > 1 && index_lanes > 1 && value_lanes != index_lanes) {
// Mismatched vector lane representations; be conservative.
return std::nullopt;
}
const int effective_lanes = std::max(value_lanes, index_lanes);
const int per_access_bits = effective_lanes * load->dtype.bits();
const int total_bits = static_cast<int>(per_access_bits) *
static_cast<int>(current_vectorized_lanes_);
// The async-copy marker is byte-granular. `tl.ptx_cp_async` stores logical
// element counts, but we still need to know that the eventual vectorized
// transfer can map to a legal byte width without over-copying packed
// subbyte data.
if (total_bits % 8 != 0) {
return std::nullopt;
}
const int total_bytes = total_bits / 8;
if (!IsValidCPAsyncTransferBytes(total_bytes)) {
return std::nullopt;
}
CopyIndexInfo info;
info.src_index = src_index;
info.dst_index = dst_index;
info.index_lanes = index_lanes;
info.per_access_num_elems = effective_lanes;
return info;
}
static PrimExpr ExtractVectorBase(const PrimExpr &index) {
if (index.dtype().lanes() == 1) {
return index;
}
if (const auto *broadcast = index.as<BroadcastNode>()) {
return broadcast->value;
}
if (const auto *ramp = index.as<RampNode>()) {
if (!is_one(ramp->stride)) {
return PrimExpr();
}
return ramp->base;
}
const auto *add = index.as<AddNode>();
if (!add) {
return PrimExpr();
}
// Common pattern after flattening a vectorized N-D buffer access:
// (broadcast(base_offset) + ramp(vec_base, 1, lanes))
// or its commuted form:
// (ramp(vec_base, 1, lanes) + broadcast(base_offset))
const PrimExpr &lhs = add->a;
const PrimExpr &rhs = add->b;
if (const auto *lhs_ramp = lhs.as<RampNode>()) {
if (!is_one(lhs_ramp->stride)) {
return PrimExpr();
}
if (const auto *rhs_broadcast = rhs.as<BroadcastNode>()) {
return tirx::Add(lhs_ramp->base, rhs_broadcast->value);
}
}
if (const auto *rhs_ramp = rhs.as<RampNode>()) {
if (!is_one(rhs_ramp->stride)) {
return PrimExpr();
}
if (const auto *lhs_broadcast = lhs.as<BroadcastNode>()) {
return tirx::Add(rhs_ramp->base, lhs_broadcast->value);
}
}
return PrimExpr();
}
static Optional<Array<PrimExpr>>
ExtractVectorBaseIndices(const Array<PrimExpr> &indices) {
Array<PrimExpr> base_indices;
base_indices.reserve(indices.size());
for (const PrimExpr &index : indices) {
PrimExpr base = ExtractVectorBase(index);
if (!base.defined()) {
return Optional<Array<PrimExpr>>();
}
base_indices.push_back(base);
}
return base_indices;
}
static PrimExpr MakeAccessPtrFromLoad(const BufferLoad &base_load, int extent,
int rw_mask) {
return Call(DataType::Handle(), tvm::tl::access_ptr(),
{base_load, IntImm(DataType::Int(32), extent),
IntImm(DataType::Int(32), rw_mask)});
}
static Optional<Stmt>
MakeCPAsyncStmtFromLoads(const BufferStoreNode *store,
const BufferLoad &dst_base_load,
const BufferLoad &src_base_load, int num_elems,
bool predicated, const PrimExpr &predicate_value) {
PrimExpr dst_access_ptr =
MakeAccessPtrFromLoad(dst_base_load, num_elems, /*rw_mask=*/2);
PrimExpr src_access_ptr =
MakeAccessPtrFromLoad(src_base_load, num_elems, /*rw_mask=*/1);
Array<PrimExpr> cp_async_args;
if (predicated) {
cp_async_args = {dst_access_ptr, src_access_ptr, PrimExpr(num_elems),
predicate_value};
} else {
cp_async_args = {dst_access_ptr, src_access_ptr, PrimExpr(num_elems)};
}
return Evaluate(
Call(store->buffer->dtype, tvm::tl::ptx_cp_async(), cp_async_args));
}
static Stmt MakeCommitGroupStmt() {
return Evaluate(Call(DataType::Handle(), builtin::ptx_commit_group(), {}));
}
static Stmt MakeWaitGroupStmt(int n) {
return Evaluate(Call(DataType::Handle(), builtin::ptx_wait_group(),
{IntImm(DataType::Int(32), n)}));
}
// ---- Vectorized-offset contiguity helpers ----
static bool TryGetConstInt64(const PrimExpr &expr, int64_t *value) {
if (const auto *imm = expr.as<IntImmNode>()) {
*value = imm->value;
return true;
}
return false;
}
bool HasUnitStrideForVectorizedLoop(const PrimExpr &expr,
const ActiveVectorizedLoop &loop) {
PrimExpr prev = analyzer_.Simplify(
Substitute(expr, {{loop.loop_var, IntImm(loop.loop_var->dtype, 0)}}));
int64_t stride = 0;
for (int value = 1; value < loop.extent; ++value) {
PrimExpr curr = analyzer_.Simplify(Substitute(
expr, {{loop.loop_var, IntImm(loop.loop_var->dtype, value)}}));
PrimExpr delta = analyzer_.Simplify(curr - prev);
int64_t delta_value = 0;
if (!TryGetConstInt64(delta, &delta_value)) {
return false;
}
if (value == 1) {
stride = delta_value;
} else if (delta_value != stride) {
return false;
}
prev = curr;
}
return stride == 1;
}
bool HasContiguousVectorizedOffsets(const PrimExpr &src_index,
const PrimExpr &dst_index) {
for (const auto &loop : active_vectorized_loops_) {
if (!HasUnitStrideForVectorizedLoop(src_index, loop) ||
!HasUnitStrideForVectorizedLoop(dst_index, loop)) {
return false;
}
}
return true;
}
// ---- Copy-region synchronization analysis helpers ----
struct CopyRegionAnalysis {
bool is_pure_copy_region = true;
int commit = 0;
int wait = 0;
};
struct VisitedStmtInfo {
Stmt visited;
CopyRegionAnalysis analysis;
bool opens_copy_region{false};
bool has_uncommitted_transfers{false};
};
static CopyRegionAnalysis
MergeCopyRegionAnalysis(CopyRegionAnalysis a, const CopyRegionAnalysis &b) {
a.is_pure_copy_region = a.is_pure_copy_region && b.is_pure_copy_region;
a.commit += b.commit;
a.wait += b.wait;
return a;
}
static CopyRegionAnalysis AnalyzeCopyRegion(const Stmt &stmt) {
CopyRegionAnalysis out;
if (!stmt.defined()) {
return out;
}
if (const auto *seq = stmt.as<SeqStmtNode>()) {
for (const Stmt &s : seq->seq) {
out = MergeCopyRegionAnalysis(out, AnalyzeCopyRegion(s));
}
return out;
}
if (const auto *ite = stmt.as<IfThenElseNode>()) {
// Ignore the condition: treat it as pure control flow, and only care
// whether the branches are pure copy regions so we can hoist sync out.
out = MergeCopyRegionAnalysis(out, AnalyzeCopyRegion(ite->then_case));
if (ite->else_case.defined()) {
out = MergeCopyRegionAnalysis(
out, AnalyzeCopyRegion(ite->else_case.value()));
}
return out;
}
if (const auto *eval = stmt.as<EvaluateNode>()) {
if (is_const_int(eval->value)) {
return out;
}
const auto *call = eval->value.as<CallNode>();
if (!call) {
out.is_pure_copy_region = false;
return out;
}
if (call->op.same_as(builtin::ptx_cp_async()) ||
call->op.same_as(tl::ptx_cp_async())) {
return out;
}
if (call->op.same_as(builtin::ptx_commit_group())) {
out.commit += 1;
return out;
}
if (call->op.same_as(builtin::ptx_wait_group())) {
out.wait += 1;
return out;
}
out.is_pure_copy_region = false;
return out;
}
if (stmt.as<BindNode>()) {
CopyRegionAnalysis out;
out.is_pure_copy_region = false;
return out;
}
if (const auto *attr = stmt.as<AttrStmtNode>()) {
return AnalyzeCopyRegion(attr->body);
}
if (const auto *loop = stmt.as<ForNode>()) {
return AnalyzeCopyRegion(loop->body);
}
if (const auto *block = stmt.as<SBlockNode>()) {
if (block->init.defined()) {
out = MergeCopyRegionAnalysis(out,
AnalyzeCopyRegion(block->init.value()));
}
out = MergeCopyRegionAnalysis(out, AnalyzeCopyRegion(block->body));
return out;
}
if (const auto *realize = stmt.as<SBlockRealizeNode>()) {
// Treat the predicate as pure control flow (no side effects). We only
// care whether the realized body is a pure copy region so we can hoist
// the final commit+wait out of sequential loop nests.
const SBlockNode *block = realize->block.get();
if (block->init.defined()) {
out = MergeCopyRegionAnalysis(out,
AnalyzeCopyRegion(block->init.value()));
}
out = MergeCopyRegionAnalysis(out, AnalyzeCopyRegion(block->body));
return out;
}
out.is_pure_copy_region = false;
return out;
}
VisitedStmtInfo VisitAndAnalyzeStmt(const Stmt &stmt) {
pending_sync_copies_ = false;
uncommitted_sync_copies_ = false;
Stmt visited = this->VisitStmt(stmt);
VisitedStmtInfo out;
out.visited = visited;
out.analysis = AnalyzeCopyRegion(visited);
out.opens_copy_region = pending_sync_copies_;
out.has_uncommitted_transfers = uncommitted_sync_copies_;
return out;
}
// ---- Synchronization emission helpers ----
void AppendSyncVisibility(Array<Stmt> *seq, bool include_commit) const {
if (include_commit) {
seq->push_back(MakeCommitGroupStmt());
}
seq->push_back(MakeWaitGroupStmt(0));
}
// Note: AnalyzeCopyRegion replaces both the old `IsPureCopyRegion` and
// `SummarizeAsyncIntrinsics` helpers to avoid redundant traversals.
bool enable_auto_async_copy_{true};
bool async_without_async_commit_wait_{false};
int explicit_async_scope_depth_{0};
int current_vectorized_lanes_{1};
std::vector<ActiveVectorizedLoop> active_vectorized_loops_;
arith::Analyzer analyzer_;
bool injected_rocm_async_copy_{false};
bool pending_sync_copies_{false};
bool uncommitted_sync_copies_{false};
};
ROCmAsyncCopyInjectResult
InjectROCmAsyncCopy(const Stmt &body, bool enable_auto_async_copy,
bool async_without_async_commit_wait) {
ROCmAsyncCopyInjector injector(enable_auto_async_copy,
async_without_async_commit_wait);
Stmt injected = injector(body);
return {injector.Finalize(injected), injector.InjectedROCmAsyncCopy()};
}
} // namespace tl
} // namespace tvm
+24
View File
@@ -0,0 +1,24 @@
#pragma once
#include <tvm/tirx/stmt.h>
namespace tvm {
namespace tl {
struct ROCmAsyncCopyInjectResult {
tvm::tirx::Stmt stmt;
bool injected_rocm_async_copy{false};
};
/*! \brief Inject ROCm async-copy lowering patterns into a statement.
*
* This is the statement-level entrypoint used by other transforms to apply the
* same rewrite as ROCm async-copy lowering, but scoped to a region (e.g., a
* lowered parallel loop) rather than the whole PrimFunc.
*/
ROCmAsyncCopyInjectResult
InjectROCmAsyncCopy(const tvm::tirx::Stmt &body, bool enable_auto_async_copy,
bool async_without_async_commit_wait = false);
} // namespace tl
} // namespace tvm
+4 -3
View File
@@ -23,8 +23,9 @@
#include "../op/gemm_sp.h"
#include "../op/operator.h"
#include "../op/utils.h"
#include "backend/common/target_utils.h"
#include "ptx_async_copy_injector.h"
#include "cpu/target_utils.h"
#include "cuda/target_utils.h"
#include "cuda/transform/ptx_async_copy_injector.h"
#include "arith/ir_mutator_with_analyzer.h"
#include "common/mbarrier.h"
@@ -1480,7 +1481,7 @@ private:
// Only parallel-loop lowering needs PTX cp.async injection. Thread-level
// lowering does not require converting eligible global->shared copies to
// `tir.ptx_cp_async`.
if (TargetIsCuda(target_) && TargetHasAsyncCopy(target_)) {
if (TargetCudaHasAsyncCopy(target_)) {
tvm::transform::PassContext ctx = tvm::transform::PassContext::Current();
bool enable_auto_async_copy =
ctx->GetConfig<Bool>(kEnableAsyncCopy, Bool(true)).value();
+9
View File
@@ -0,0 +1,9 @@
# WebGPU backend: pure C++ operator lowering registrations.
#
# WebGPU currently has no optional native toolchain integration in TileLang, so
# its sources are always safe to compile.
file(GLOB TILE_LANG_WEBGPU_SRCS
src/webgpu/op/*.cc
)
list(APPEND TILE_LANG_SRCS ${TILE_LANG_WEBGPU_SRCS})
@@ -5,7 +5,7 @@ Pre-fix `SharedMemoryAlignmentPlanner` (`merge_shared_memory_allocations.cc`)
gated the 1024-byte alignment of TMA-touched smem buffers on
`TargetIsHopper(target)` (arch in `[90, 100)`). For sm_100 / sm_120 (also
TMA-capable: see `cp.async.bulk.tensor.{1..5}d.shared::cta.global.*` PTX
support gated by `TargetHasBulkCopy(target)` in `src/backend/common/target_utils.cc`)
support gated by `TargetHasBulkCopy(target)` in `src/cuda/target_utils.cc`)
the planner fell back to the global default (16 bytes) and TMA destinations
landed at 16-byte-aligned offsets. `cp.async.bulk.tensor.*` requires the
destination smem pointer to be 128-byte aligned, so on sm_100/sm_120 this
@@ -1,4 +1,4 @@
"""Tests for TileLang `LowerPTXAsyncCopy` transform pass."""
"""Tests for TileLang CUDA `LowerPTXAsyncCopy` transform pass."""
from tilelang import tvm
import tilelang as tl
@@ -48,7 +48,7 @@ def test_lower_ptx_async_copy_rewrites_plain_parallel_copy():
func = before.with_attr("global_symbol", "main").with_attr("target", target)
mod = tvm.IRModule.from_expr(func)
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
calls = _count_calls(mod["main"])
assert calls.get("tl.ptx_cp_async", 0) > 0
@@ -74,7 +74,7 @@ def test_lower_ptx_async_copy_respects_explicit_async_scope():
func = before.with_attr("global_symbol", "main").with_attr("target", target)
mod = tvm.IRModule.from_expr(func)
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
calls = _count_calls(mod["main"])
assert calls.get("tl.ptx_cp_async", 0) > 0
@@ -99,7 +99,7 @@ def test_lower_ptx_async_copy_supports_multi_dim_indices():
func = before.with_attr("global_symbol", "main").with_attr("target", target)
mod = tvm.IRModule.from_expr(func)
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
calls = _count_calls(mod["main"])
assert calls.get("tl.ptx_cp_async", 0) > 0
@@ -125,7 +125,7 @@ def test_lower_ptx_async_copy_rewrites_vectorized_float16_loop():
func = before.with_attr("global_symbol", "main").with_attr("target", target)
mod = tvm.IRModule.from_expr(func)
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
calls = _count_calls(mod["main"])
assert calls.get("tl.ptx_cp_async", 0) > 0
@@ -153,7 +153,7 @@ def test_lower_ptx_async_copy_hoists_sync_out_of_predicated_block():
func = before.with_attr("global_symbol", "main").with_attr("target", target)
mod = tvm.IRModule.from_expr(func)
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
calls = _count_calls(mod["main"])
assert calls.get("tl.ptx_cp_async", 0) > 0
assert calls.get("tirx.ptx_commit_group", 0) > 0
@@ -192,7 +192,7 @@ def test_lower_ptx_async_copy_respects_enable_async_copy_config():
mod = tvm.IRModule.from_expr(func)
with tvm.transform.PassContext(config={tl.PassConfigKey.TL_ENABLE_ASYNC_COPY: False}):
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
calls = _count_calls(mod["main"])
assert calls.get("tl.ptx_cp_async", 0) == 0
@@ -219,7 +219,7 @@ def test_lower_ptx_async_copy_does_not_duplicate_existing_sync():
func = before.with_attr("global_symbol", "main").with_attr("target", target)
mod = tvm.IRModule.from_expr(func)
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
calls = _count_calls(mod["main"])
assert calls.get("tl.ptx_cp_async", 0) > 0
@@ -245,7 +245,7 @@ def test_lower_ptx_async_copy_inserts_commit_before_existing_wait():
func = before.with_attr("global_symbol", "main").with_attr("target", target)
mod = tvm.IRModule.from_expr(func)
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
calls = _count_calls(mod["main"])
assert calls.get("tl.ptx_cp_async", 0) > 0
@@ -271,7 +271,7 @@ def test_lower_ptx_async_copy_keeps_sync_out_of_inner_unrolled_loops_in_pipeline
func = before.with_attr("global_symbol", "main").with_attr("target", target)
mod = tvm.IRModule.from_expr(func)
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
calls = _count_calls(mod["main"])
assert calls.get("tl.ptx_cp_async", 0) > 0
assert calls.get("tirx.ptx_commit_group", 0) > 0
@@ -319,7 +319,7 @@ def test_lower_ptx_async_copy_from_vectorized_loop():
func = before.with_attr("global_symbol", "main").with_attr("target", target)
mod = tvm.IRModule.from_expr(func)
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
calls = _count_calls(mod["main"])
assert calls.get("tl.ptx_cp_async", 0) > 0
@@ -342,7 +342,7 @@ def test_lower_ptx_async_copy_skips_vectorized_broadcast_source():
func = before.with_attr("global_symbol", "main").with_attr("target", target)
mod = tvm.IRModule.from_expr(func)
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
calls = _count_calls(mod["main"])
assert calls.get("tl.ptx_cp_async", 0) == 0
assert calls.get("tirx.ptx_commit_group", 0) == 0
@@ -365,7 +365,7 @@ def test_lower_ptx_async_copy_from_ramp():
func = before.with_attr("global_symbol", "main").with_attr("target", target)
mod = tvm.IRModule.from_expr(func)
mod = tl.transform.LowerPTXAsyncCopy()(mod)
mod = tl.cuda.transform.LowerPTXAsyncCopy()(mod)
print(mod)
calls = _count_calls(mod["main"])
print(calls)
+20
View File
@@ -91,6 +91,25 @@ def LowerLDGSTG():
return _ffi_api.LowerLDGSTG() # type: ignore
def LowerPTXAsyncCopy():
"""Lower eligible global->shared copies into PTX `cp.async` on CUDA.
When enabled (pass config `tl.enable_async_copy`, default True), this pass
may rewrite plain user-written global->shared `BufferStore` patterns (e.g.
SIMT copies in `T.Parallel`) into `tir.ptx_cp_async`, and insert
`tir.ptx_commit_group` + `tir.ptx_wait_group(0)` to preserve synchronous
semantics for normal stores. If explicit commit/wait intrinsics already
exist, the pass avoids duplicating them (and may insert a missing commit
immediately before an existing wait to cover injected `cp.async`).
Returns
-------
fpass : tvm.transform.Pass
The result pass
"""
return _ffi_api.LowerPTXAsyncCopy() # type: ignore
def MarkCudaSyncCalls(have_pdl: bool = False):
"""MarkCudaSyncCalls"""
return _ffi_api.MarkCudaSyncCalls(have_pdl) # type: ignore
@@ -154,6 +173,7 @@ __all__ = [
"LowerHopperIntrin",
"LowerLDGSTG",
"LowerL2Persistent",
"LowerPTXAsyncCopy",
"LowerSharedBarrier",
"LowerSharedTmem",
"MarkCudaSyncCalls",
+1 -1
View File
@@ -139,7 +139,7 @@ def target_is_gfx950(target: Target) -> bool:
def target_get_warp_size(target: Target) -> int:
return _target_ffi_api().TargetGetWarpSize(target)
return _target_ffi_api().TargetRocmGetWarpSize(target)
def target_get_rdna_generation(target: Target) -> int:
-24
View File
@@ -262,30 +262,6 @@ def VectorizeLoop(enable_vectorize: bool = True):
return _ffi_api.VectorizeLoop(enable_vectorize) # type: ignore
def LowerPTXAsyncCopy():
"""Lower eligible global->shared copies into PTX `cp.async` on CUDA.
When enabled (pass config `tl.enable_async_copy`, default True), this pass
may rewrite plain user-written global->shared `BufferStore` patterns (e.g.
SIMT copies in `T.Parallel`) into `tir.ptx_cp_async`, and insert
`tir.ptx_commit_group` + `tir.ptx_wait_group(0)` to preserve synchronous
semantics for normal stores. If explicit commit/wait intrinsics already
exist, the pass avoids duplicating them (and may insert a missing commit
immediately before an existing wait to cover injected `cp.async`).
Returns
-------
fpass : tvm.transform.Pass
The result pass
"""
return _ffi_api.LowerPTXAsyncCopy() # type: ignore
def InjectPTXAsyncCopy():
"""Deprecated alias of `LowerPTXAsyncCopy`."""
return LowerPTXAsyncCopy()
def ConfigIndexBitwidth():
"""Config index bitwidth.