mirror of
https://github.com/tile-ai/tilelang.git
synced 2026-10-05 07:42:33 +08:00
[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:
+2
-16
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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_
|
||||
@@ -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
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
@@ -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>
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
@@ -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;
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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),
|
||||
|
||||
+5
-44
@@ -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
|
||||
+2
-2
@@ -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,
|
||||
@@ -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_
|
||||
@@ -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
|
||||
|
||||
@@ -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";
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -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);
|
||||
|
||||
@@ -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
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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_
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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();
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user