[Minor] Enhance JIT kernel and add dev docs (#14570)

This commit is contained in:
DarkSharpness
2025-12-23 22:34:59 +08:00
committed by GitHub
parent d7301c89ba
commit 291f11ae39
13 changed files with 817 additions and 290 deletions

View File

@@ -13,7 +13,6 @@
#include <initializer_list>
#include <optional>
#include <ranges>
#include <source_location>
#include <span>
#include <sstream>
#include <string>
@@ -21,13 +20,21 @@
#include <type_traits>
#include <utility>
#ifdef __CUDACC__
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#endif
namespace host {
namespace stdr = std::ranges;
namespace stdv = std::views;
namespace details {
inline constexpr auto kAnyDeviceID = -1;
inline constexpr auto kAnySize = static_cast<int64_t>(-1);
inline constexpr auto kNullSize = static_cast<int64_t>(-1);
inline constexpr auto kNullDType = static_cast<DLDataTypeCode>(18u);
inline constexpr auto kNullDevice = static_cast<DLDeviceType>(-1);
struct SizeRef;
struct DTypeRef;
struct DeviceRef;
@@ -37,7 +44,7 @@ struct dtype_trait {};
template <std::integral T>
struct dtype_trait<T> {
inline static constexpr auto value = DLDataType{
inline static constexpr DLDataType value = {
.code = std::is_signed_v<T> ? DLDataTypeCode::kDLInt : DLDataTypeCode::kDLUInt,
.bits = static_cast<std::uint8_t>(sizeof(T) * 8),
.lanes = 1};
@@ -45,22 +52,31 @@ struct dtype_trait<T> {
template <std::floating_point T>
struct dtype_trait<T> {
inline static constexpr auto value =
DLDataType{.code = DLDataTypeCode::kDLFloat, .bits = static_cast<std::uint8_t>(sizeof(T) * 8), .lanes = 1};
inline static constexpr DLDataType value = {
.code = DLDataTypeCode::kDLFloat, .bits = static_cast<std::uint8_t>(sizeof(T) * 8), .lanes = 1};
};
inline constexpr auto kAnyDeviceID = -1;
inline constexpr auto kAnySize = static_cast<int64_t>(-1);
inline constexpr auto kNullSize = static_cast<int64_t>(-1);
inline constexpr auto kNullDType = static_cast<DLDataTypeCode>(18u);
inline constexpr auto kNullDevice = static_cast<DLDeviceType>(-1);
#ifdef __CUDACC__
template <>
struct dtype_trait<__half> {
inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLFloat, .bits = 16, .lanes = 1};
};
template <>
struct dtype_trait<__nv_bfloat16> {
inline static constexpr DLDataType value = {.code = DLDataTypeCode::kDLBfloat, .bits = 16, .lanes = 1};
};
#endif
template <DLDeviceType Code>
struct device_trait {
inline static constexpr DLDevice value = {.device_type = Code, .device_id = kAnyDeviceID};
};
template <typename... Ts>
inline constexpr auto kDTypeList = std::array<DLDataType, sizeof...(Ts)>{dtype_trait<Ts>::value...};
template <DLDeviceType... Codes>
inline constexpr auto kDeviceList = std::array<DLDevice, sizeof...(Codes)>{
DLDevice{.device_type = static_cast<DLDeviceType>(Codes), .device_id = kAnyDeviceID}...};
inline constexpr auto kDeviceList = std::array<DLDevice, sizeof...(Codes)>{device_trait<Codes>::value...};
template <typename T>
struct PrintAbleSpan {
@@ -103,11 +119,13 @@ struct PrintableDevice {
inline auto& operator<<(std::ostream& os, DLDevice device) {
const auto& mapping = kDeviceStringMap;
const auto entry = static_cast<std::size_t>(device.device_type);
host::RuntimeCheck(entry < mapping.size());
RuntimeCheck(entry < mapping.size());
const auto name = mapping[entry];
host::RuntimeCheck(!name.empty(), "Unknown device: ", int(device.device_type));
RuntimeCheck(!name.empty(), "Unknown device: ", int(device.device_type));
os << name;
if (device.device_id != kAnyDeviceID) os << "[" << device.device_id << "]";
if (device.device_id != kAnyDeviceID && device.device_type != DLDeviceType::kDLCPU) {
os << ":" << device.device_id;
}
return os;
}
@@ -118,7 +136,7 @@ inline auto& operator<<(std::ostream& os, PrintableDevice pd) {
template <typename T>
inline auto& operator<<(std::ostream& os, PrintAbleSpan<T> span) {
os << "[";
for (const auto i : stdv::iota(std::size_t{0}, span.data.size())) {
for (const auto i : irange(span.data.size())) {
if (i > 0) {
os << ", ";
}
@@ -133,37 +151,58 @@ inline auto& operator<<(std::ostream& os, PrintAbleSpan<T> span) {
struct SymbolicSize {
public:
SymbolicSize(std::string_view annotation = {}) : m_value(details::kNullSize), m_annotation(annotation) {}
SymbolicSize(const SymbolicSize&) = delete;
SymbolicSize& operator=(const SymbolicSize&) = delete;
auto get_name() const -> std::string_view {
return m_annotation;
}
auto set_value(int64_t value) -> void {
host::RuntimeCheck(!this->has_value(), "Size value already set");
RuntimeCheck(!this->has_value(), "Size value already set");
m_value = value;
}
auto has_value() const -> bool {
return m_value != details::kNullSize;
}
auto get_value() const -> std::optional<int64_t> {
return this->has_value() ? std::optional{m_value} : std::nullopt;
}
auto unwrap() const -> int64_t {
host::RuntimeCheck(this->has_value(), "Size value is not set");
auto unwrap(DebugInfo info = {}) const -> int64_t {
RuntimeCheck(info, this->has_value(), "Size value is not set");
return m_value;
}
SymbolicSize(const SymbolicSize&) = delete;
SymbolicSize& operator=(const SymbolicSize&) = delete;
auto verify(int64_t dim) -> void {
auto verify(int64_t value, const char* prefix, int64_t dim) -> void {
if (this->has_value()) {
host::RuntimeCheck(m_value == dim, "Size mismatch: expected ", m_value, " but got ", dim);
if (m_value != value) {
[[unlikely]];
Panic("Size mismatch for ", m_name_str(prefix, dim), ": expected ", m_value, " but got ", value);
}
} else {
this->set_value(dim);
this->set_value(value);
}
}
auto value_or_name(const char* prefix, int64_t dim) const -> std::string {
if (const auto value = this->get_value()) {
return std::to_string(*value);
} else {
return m_name_str(prefix, dim);
}
}
private:
auto m_name_str(const char* prefix, int64_t dim) const -> std::string {
std::ostringstream os;
os << prefix << '#' << dim;
if (!m_annotation.empty()) os << "('" << m_annotation << "')";
return std::move(os).str();
}
std::int64_t m_value;
std::string_view m_annotation;
};
@@ -175,27 +214,33 @@ inline auto operator==(DLDevice lhs, DLDevice rhs) -> bool {
struct SymbolicDType {
public:
SymbolicDType() : m_value({details::kNullDType, 0, 0}) {}
SymbolicDType(const SymbolicDType&) = delete;
SymbolicDType& operator=(const SymbolicDType&) = delete;
auto set_value(DLDataType value) -> void {
host::RuntimeCheck(!this->has_value(), "Dtype value already set");
host::RuntimeCheck(
RuntimeCheck(!this->has_value(), "Dtype value already set");
RuntimeCheck(
m_check(value), "Dtype value [", value, "] not in the allowed options: ", details::PrintAbleSpan{m_options});
m_value = value;
}
auto has_value() const -> bool {
return m_value.code != details::kNullDType;
}
auto get_value() const -> std::optional<DLDataType> {
return this->has_value() ? std::optional{m_value} : std::nullopt;
}
auto unwrap() const -> DLDataType {
host::RuntimeCheck(this->has_value(), "Dtype value is not set");
auto unwrap(DebugInfo info = {}) const -> DLDataType {
RuntimeCheck(info, this->has_value(), "Dtype value is not set");
return m_value;
}
auto set_options(std::span<const DLDataType> options) -> void {
m_options = options;
}
template <typename... Ts>
auto set_options() -> void {
m_options = details::kDTypeList<Ts...>;
@@ -203,7 +248,7 @@ struct SymbolicDType {
auto verify(DLDataType dtype) -> void {
if (this->has_value()) {
host::RuntimeCheck(m_value == dtype, "DType mismatch: expected ", m_value, " but got ", dtype);
RuntimeCheck(m_value == dtype, "DType mismatch: expected ", m_value, " but got ", dtype);
} else {
this->set_value(dtype);
}
@@ -221,10 +266,12 @@ struct SymbolicDType {
struct SymbolicDevice {
public:
SymbolicDevice() : m_value({details::kNullDevice, details::kAnyDeviceID}) {}
SymbolicDevice(const SymbolicDevice&) = delete;
SymbolicDevice& operator=(const SymbolicDevice&) = delete;
auto set_value(DLDevice value) -> void {
host::RuntimeCheck(!this->has_value(), "Device value already set");
host::RuntimeCheck(
RuntimeCheck(!this->has_value(), "Device value already set");
RuntimeCheck(
m_check(value),
"Device value [",
details::PrintableDevice{value},
@@ -232,20 +279,24 @@ struct SymbolicDevice {
details::PrintAbleSpan{m_options});
m_value = value;
}
auto has_value() const -> bool {
return m_value.device_type != details::kNullDevice;
}
auto get_value() const -> std::optional<DLDevice> {
return this->has_value() ? std::optional{m_value} : std::nullopt;
}
auto unwrap() const -> DLDevice {
host::RuntimeCheck(this->has_value(), "Device value is not set");
auto unwrap(DebugInfo info = {}) const -> DLDevice {
RuntimeCheck(info, this->has_value(), "Device value is not set");
return m_value;
}
auto set_options(std::span<const DLDevice> options) -> void {
m_options = options;
}
template <DLDeviceType... Codes>
auto set_options() -> void {
m_options = details::kDeviceList<Codes...>;
@@ -253,7 +304,7 @@ struct SymbolicDevice {
auto verify(DLDevice device) -> void {
if (this->has_value()) {
host::RuntimeCheck(
RuntimeCheck(
m_value == device,
"Device mismatch: expected ",
details::PrintableDevice{m_value},
@@ -313,19 +364,6 @@ struct SizeRef : BaseRef<SymbolicSize> {
// otherwise, we can match any size
}
}
auto value_or_name(std::size_t dim) const -> std::string {
if (const auto value = (**this).get_value()) {
return std::to_string(*value);
} else {
const auto annotation = (**this).get_name();
if (annotation.empty()) {
return "dim#" + std::to_string(dim);
} else {
return static_cast<std::string>(annotation);
}
}
}
};
struct DTypeRef : BaseRef<SymbolicDType> {
@@ -361,7 +399,6 @@ struct TensorMatcher {
using SizeRef = details::SizeRef;
using DTypeRef = details::DTypeRef;
using DeviceRef = details::DeviceRef;
using Loc_t = std::source_location;
public:
TensorMatcher(const TensorMatcher&) = delete;
@@ -371,8 +408,8 @@ struct TensorMatcher {
auto with_strides(std::initializer_list<SizeRef> strides) && -> TensorMatcher&& {
// no partial update allowed
host::RuntimeCheck(m_strides.size() == 0, "Strides already specified");
host::RuntimeCheck(m_shape.size() == strides.size(), "Strides size must match shape size");
RuntimeCheck(m_strides.size() == 0, "Strides already specified");
RuntimeCheck(m_shape.size() == strides.size(), "Strides size must match shape size");
m_strides = strides;
return std::move(*this);
}
@@ -381,6 +418,7 @@ struct TensorMatcher {
auto with_dtype(DTypeRef&& dtype) && -> TensorMatcher&& {
m_init_dtype();
m_dtype.rebind(*dtype);
m_dtype->set_options<Ts...>();
return std::move(*this);
}
@@ -396,6 +434,7 @@ struct TensorMatcher {
auto with_device(DeviceRef&& device) && -> TensorMatcher&& {
m_init_device();
m_device.rebind(*device);
m_device->set_options<Codes...>();
return std::move(*this);
}
@@ -408,70 +447,70 @@ struct TensorMatcher {
}
// once we start verification, we cannot modify anymore
auto verify(tvm::ffi::TensorView view, Loc_t loc = Loc_t::current()) const&& -> const TensorMatcher&& {
auto verify(tvm::ffi::TensorView view, DebugInfo info = {}) const&& -> const TensorMatcher&& {
try {
this->m_verify_impl(view);
m_verify_impl(view);
} catch (PanicError& e) {
auto oss = std::ostringstream{};
oss << "Tensor match failed for " << this->debug_str() << " at " << loc.file_name() << ":" << loc.line()
<< "\n- Root cause: " << e.detail();
oss << "Tensor match failed for ";
s_print_tensor(oss, view);
oss << " at " << info.file_name() << ":" << info.line() << "\n- Root cause: " << e.root_cause();
throw PanicError(std::move(oss).str());
}
return std::move(*this);
}
auto debug_str() const -> std::string {
auto oss = std::ostringstream{};
private:
static auto s_print_tensor(std::ostringstream& oss, tvm::ffi::TensorView view) -> void {
oss << "Tensor<";
std::size_t dim = 0;
for (const auto& size_ref : m_shape) {
if (dim > 0) {
int64_t dim = 0;
for (const auto& size : view.shape()) {
if (dim++ > 0) oss << ", ";
oss << size;
}
oss << ">[strides=<";
dim = 0;
for (const auto& stride : view.strides()) {
if (dim++ > 0) {
oss << ", ";
}
oss << size_ref.value_or_name(dim++);
oss << stride;
}
oss << ">";
if (m_strides.size() > 0) {
oss << " [strides=<";
dim = 0;
for (const auto& stride_ref : m_strides) {
if (dim > 0) {
oss << ", ";
}
oss << stride_ref.value_or_name(dim++);
}
oss << ">]";
}
return std::move(oss).str();
oss << ">, dtype=" << view.dtype();
oss << ", device=" << details::PrintableDevice{view.device()} << "]";
}
private:
auto m_verify_impl(tvm::ffi::TensorView view) const -> void {
const auto dim = static_cast<std::size_t>(view.dim());
host::RuntimeCheck(dim == m_shape.size(), "Tensor dimension mismatch: expected ", m_shape.size(), " but got ", dim);
for (const auto i : stdv::iota(std::size_t{0}, dim)) {
m_shape[i]->verify(view.size(i));
RuntimeCheck(dim == m_shape.size(), "Tensor dimension mismatch: expected ", m_shape.size(), " but got ", dim);
for (const auto i : irange(dim)) {
m_shape[i]->verify(view.size(i), "shape", i);
}
if (this->m_has_strides()) {
for (const auto i : stdv::iota(std::size_t{0}, dim)) {
m_strides[i]->verify(view.stride(i));
if (m_has_strides()) {
for (const auto i : irange(dim)) {
if (view.size(i) != 1 || !m_strides[i]->has_value()) {
// skip stride check for size 1 dimension
m_strides[i]->verify(view.stride(i), "stride", i);
}
}
} else {
host::RuntimeCheck(view.is_contiguous(), "Tensor is not contiguous as expected");
RuntimeCheck(view.is_contiguous(), "Tensor is not contiguous as expected");
}
// since we may use the same matcher to verify again, we will force to check
// since we may double verify, we will force to check
m_dtype->verify(view.dtype());
m_device->verify(view.device());
}
auto m_init_dtype() -> void {
host::RuntimeCheck(!m_has_dtype, "DType already specified");
RuntimeCheck(!m_has_dtype, "DType already specified");
m_has_dtype = true;
}
auto m_init_device() -> void {
host::RuntimeCheck(!m_has_device, "Device already specified");
RuntimeCheck(!m_has_device, "Device already specified");
m_has_device = true;
}
auto m_has_strides() const -> bool {
return !m_strides.empty();
}

View File

@@ -7,7 +7,6 @@
#include <concepts>
#include <cstddef>
#include <source_location>
#include <type_traits>
namespace device {
@@ -32,60 +31,63 @@ __always_inline __device__ auto offset(const T* ptr, U... offset) -> const void*
} // namespace pointer
template <typename T, std::size_t N>
struct device_vec {
T data[N];
};
} // namespace device
namespace host {
inline auto
RuntimeDeviceCheck(::cudaError_t error, std::source_location location = std::source_location::current()) -> void {
inline void RuntimeDeviceCheck(::cudaError_t error, DebugInfo location = {}) {
if (error != ::cudaSuccess) {
[[unlikely]];
::host::panic(location, "CUDA error: ", ::cudaGetErrorString(error));
}
}
inline auto RuntimeCudaCheck(std::source_location location = std::source_location::current()) -> void {
inline void RuntimeDeviceCheck(DebugInfo location = {}) {
return RuntimeDeviceCheck(::cudaGetLastError(), location);
}
template <auto F>
inline void set_smem_once(std::size_t smem_size) {
static const auto last_smem_size = [&] {
RuntimeDeviceCheck(::cudaFuncSetAttribute(F, ::cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
return smem_size;
}();
RuntimeCheck(
smem_size <= last_smem_size,
"Dynamic shared memory size exceeds the previously set maximum size: ",
last_smem_size,
" bytes");
}
struct LaunchKernel {
public:
explicit LaunchKernel(
dim3 grid_dim, dim3 block_dim, DLDevice device, std::size_t dynamic_shared_mem_bytes = 0) noexcept
: m_config(s_make_config(grid_dim, block_dim, resolve_device(device), dynamic_shared_mem_bytes)) {}
dim3 grid_dim,
dim3 block_dim,
DLDevice device,
std::size_t dynamic_shared_mem_bytes = 0,
DebugInfo location = {}) noexcept
: m_config(s_make_config(grid_dim, block_dim, resolve_device(device), dynamic_shared_mem_bytes)),
m_location(location) {}
explicit LaunchKernel(
dim3 grid_dim, dim3 block_dim, cudaStream_t stream, std::size_t dynamic_shared_mem_bytes = 0) noexcept
: m_config(s_make_config(grid_dim, block_dim, stream, dynamic_shared_mem_bytes)) {}
dim3 grid_dim,
dim3 block_dim,
cudaStream_t stream,
std::size_t dynamic_shared_mem_bytes = 0,
DebugInfo location = {}) noexcept
: m_config(s_make_config(grid_dim, block_dim, stream, dynamic_shared_mem_bytes)), m_location(location) {}
LaunchKernel(const LaunchKernel&) = delete;
LaunchKernel& operator=(const LaunchKernel&) = delete;
static auto resolve_device(DLDevice device) -> cudaStream_t {
return static_cast<cudaStream_t>(::TVMFFIEnvGetStream(device.device_type, device.device_id));
}
LaunchKernel(const LaunchKernel&) = delete;
LaunchKernel& operator=(const LaunchKernel&) = delete;
template <typename T, typename... Args>
auto operator()(T&& kernel, Args&&... args) const -> void {
host::RuntimeDeviceCheck(::cudaLaunchKernelEx(&m_config, kernel, std::forward<Args>(args)...));
RuntimeDeviceCheck(::cudaLaunchKernelEx(&m_config, kernel, std::forward<Args>(args)...), m_location);
}
private:
static auto
s_make_config(dim3 grid_dim, dim3 block_dim, cudaStream_t stream, std::size_t smem) -> cudaLaunchConfig_t {
static auto s_make_config( // Make a config for kernel launch
dim3 grid_dim,
dim3 block_dim,
cudaStream_t stream,
std::size_t smem) -> cudaLaunchConfig_t {
auto config = ::cudaLaunchConfig_t{};
config.gridDim = grid_dim;
config.blockDim = block_dim;
@@ -94,8 +96,10 @@ struct LaunchKernel {
config.numAttrs = 0;
return config;
}
cudaLaunchConfig_t m_config;
/// TODO: We can add a queue to store the attributes if needed in the future.
const DebugInfo m_location;
/// TODO: We can add a queue to store the attributes (e.g. for PDL) if needed in the future.
};
} // namespace host

View File

@@ -1,23 +1,55 @@
#pragma once
// ref: https://forums.developer.nvidia.com/t/c-20s-source-location-compilation-error-when-using-nvcc-12-1/258026/3
#ifdef __CUDACC__
#pragma push_macro("__cpp_consteval")
#pragma push_macro("_NODISCARD")
#pragma push_macro("__builtin_LINE")
#pragma clang diagnostic push
#pragma clang diagnostic ignored "-Wbuiltin-macro-redefined"
#define __cpp_consteval 201811L
#pragma clang diagnostic pop
#ifdef _NODISCARD
#undef _NODISCARD
#define _NODISCARD
#endif
#define consteval constexpr
#include <source_location>
#undef consteval
#pragma pop_macro("__cpp_consteval")
#pragma pop_macro("_NODISCARD")
#else
#include <source_location>
#endif
#include <dlpack/dlpack.h>
#include <concepts>
#include <cstddef>
#include <ostream>
#include <ranges>
#include <source_location>
#include <sstream>
#include <utility>
namespace host {
struct DebugInfo : public std::source_location {
DebugInfo(std::source_location loc = std::source_location::current()) : std::source_location(loc) {}
};
struct PanicError : public std::runtime_error {
public:
// copy and move constructors
explicit PanicError(std::string msg) : runtime_error(msg), m_message(std::move(msg)) {}
auto detail() const -> std::string_view {
const auto sv = std::string_view{m_message};
const auto pos = sv.find(": ");
return pos == std::string_view::npos ? sv : sv.substr(pos + 2);
auto root_cause() const -> std::string_view {
const auto str = std::string_view{m_message};
const auto pos = str.find(": ");
return pos == std::string_view::npos ? str : str.substr(pos + 2);
}
private:
@@ -26,7 +58,7 @@ struct PanicError : public std::runtime_error {
template <typename... Args>
[[noreturn]]
inline auto panic(std::source_location location, Args&&... args) -> void {
inline auto panic(DebugInfo location, Args&&... args) -> void {
std::ostringstream os;
os << "Runtime check failed at " << location.file_name() << ":" << location.line();
if constexpr (sizeof...(args) > 0) {
@@ -40,32 +72,42 @@ inline auto panic(std::source_location location, Args&&... args) -> void {
template <typename... Args>
struct RuntimeCheck {
using Loc_t = std::source_location;
template <typename Cond>
explicit RuntimeCheck(Cond&& condition, Args&&... args, Loc_t location = Loc_t::current()) {
if (!condition) {
[[unlikely]];
::host::panic(location, std::forward<Args>(args)...);
}
explicit RuntimeCheck(Cond&& condition, Args&&... args, DebugInfo location = {}) {
if (condition) return;
[[unlikely]] ::host::panic(location, std::forward<Args>(args)...);
}
template <typename Cond>
explicit RuntimeCheck(DebugInfo location, Cond&& condition, Args&&... args) {
if (condition) return;
[[unlikely]] ::host::panic(location, std::forward<Args>(args)...);
}
};
template <typename... Args>
struct Panic {
explicit Panic(Args&&... args, DebugInfo location = {}) {
::host::panic(location, std::forward<Args>(args)...);
}
explicit Panic(DebugInfo location, Args&&... args) {
::host::panic(location, std::forward<Args>(args)...);
}
[[noreturn]] ~Panic() {
std::terminate();
}
};
template <typename Cond, typename... Args>
explicit RuntimeCheck(Cond&&, Args&&...) -> RuntimeCheck<Args...>;
template <std::signed_integral T, std::signed_integral U>
inline constexpr auto div_ceil(T a, U b) {
return (a + b - 1) / b;
}
template <typename Cond, typename... Args>
explicit RuntimeCheck(DebugInfo, Cond&&, Args&&...) -> RuntimeCheck<Args...>;
template <std::unsigned_integral T, std::unsigned_integral U>
inline constexpr auto div_ceil(T a, U b) {
return (a + b - 1) / b;
}
template <typename... Args>
explicit Panic(Args&&...) -> Panic<Args...>;
inline auto dtype_bytes(DLDataType dtype) -> std::size_t {
return static_cast<std::size_t>(dtype.bits / 8);
}
template <typename... Args>
explicit Panic(DebugInfo, Args&&...) -> Panic<Args...>;
namespace pointer {
@@ -85,4 +127,26 @@ inline auto offset(const T* ptr, U... offset) -> const void* {
} // namespace pointer
template <std::integral T, std::integral U>
inline constexpr auto div_ceil(T a, U b) {
return (a + b - 1) / b;
}
inline auto dtype_bytes(DLDataType dtype) -> std::size_t {
return static_cast<std::size_t>(dtype.bits / 8);
}
namespace stdr = std::ranges;
namespace stdv = stdr::views;
template <std::integral T>
inline auto irange(T end) {
return stdv::iota(static_cast<T>(0), end);
}
template <std::integral T>
inline auto irange(T start, T end) {
return stdv::iota(start, end);
}
} // namespace host

View File

@@ -1,145 +0,0 @@
#pragma once
#include <sgl_kernel/utils.cuh>
#include <cstddef>
#include <cstdint>
#include <type_traits>
namespace device::warp {
namespace details {
template <std::size_t kUnit>
inline constexpr auto get_mem_package() {
if constexpr (kUnit == 16) {
return uint4{};
} else if constexpr (kUnit == 8) {
return uint2{};
} else if constexpr (kUnit == 4) {
return uint1{};
} else {
static_assert(kUnit == 16 || kUnit == 8 || kUnit == 4, "Unsupported memory package size");
}
}
inline constexpr auto default_unit_size(std::size_t x) -> std::size_t {
if (x % (16 * kWarpThreads) == 0) return 16;
if (x % (8 * kWarpThreads) == 0) return 8;
if (x % (4 * kWarpThreads) == 0) return 4;
return 0; // trigger static assert in _get_mem_package
}
template <std::size_t kBytes, std::size_t kUnit>
using mem_package_t = decltype(get_mem_package<kUnit>());
template <typename T, std::size_t N>
struct storage_vec {
T data[N];
};
__always_inline __device__ auto load_nc(const uint1* __restrict__ src) -> uint1 {
uint32_t tmp;
asm volatile("ld.global.cs.b32 %0,[%1];" : "=r"(tmp) : "l"(src));
return uint1{tmp};
}
__always_inline __device__ auto load_nc(const uint2* __restrict__ src) -> uint2 {
uint32_t tmp0, tmp1;
asm volatile("ld.global.cs.v2.b32 {%0,%1},[%2];" : "=r"(tmp0), "=r"(tmp1) : "l"(src));
return uint2{tmp0, tmp1};
}
__always_inline __device__ auto load_nc(const uint4* __restrict__ src) -> uint4 {
uint32_t tmp0, tmp1, tmp2, tmp3;
asm volatile("ld.global.cs.v4.b32 {%0,%1,%2,%3},[%4];" : "=r"(tmp0), "=r"(tmp1), "=r"(tmp2), "=r"(tmp3) : "l"(src));
return uint4{tmp0, tmp1, tmp2, tmp3};
}
__always_inline __device__ void store_nc(uint1* __restrict__ dst, const uint1& value) {
uint32_t tmp = value.x;
asm volatile("st.global.cs.b32 [%0],%1;" ::"l"(dst), "r"(tmp));
}
__always_inline __device__ void store_nc(uint2* __restrict__ dst, const uint2& value) {
uint32_t tmp0 = value.x;
uint32_t tmp1 = value.y;
asm volatile("st.global.cs.v2.b32 [%0],{%1,%2};" ::"l"(dst), "r"(tmp0), "r"(tmp1));
}
__always_inline __device__ void store_nc(uint4* __restrict__ dst, const uint4& value) {
uint32_t tmp0 = value.x;
uint32_t tmp1 = value.y;
uint32_t tmp2 = value.z;
uint32_t tmp3 = value.w;
asm volatile("st.global.cs.v4.b32 [%0],{%1,%2,%3,%4};" ::"l"(dst), "r"(tmp0), "r"(tmp1), "r"(tmp2), "r"(tmp3));
}
} // namespace details
template <
std::size_t kBytes,
std::size_t kUnit = details::default_unit_size(kBytes),
std::size_t kThreads = ::device::kWarpThreads>
__always_inline __device__ void copy(void* __restrict__ dst, const void* __restrict__ src) {
using Package = details::mem_package_t<kBytes, kUnit>;
constexpr auto kBytesPerLoop = sizeof(Package) * kThreads;
constexpr auto kLoopCount = kBytes / kBytesPerLoop;
static_assert(kBytes % kBytesPerLoop == 0, "kBytes must be multiple of 128 bytes");
const auto dst_packed = static_cast<Package*>(dst);
const auto src_packed = static_cast<const Package*>(src);
const auto lane_id = threadIdx.x % kThreads;
#pragma unroll kLoopCount
for (std::size_t i = 0; i < kLoopCount; ++i) {
const auto j = i * kThreads + lane_id;
dst_packed[j] = src_packed[j];
}
}
template <
std::size_t kBytes,
std::size_t kUnit = details::default_unit_size(kBytes),
std::size_t kThreads = ::device::kWarpThreads>
__always_inline __device__ auto load_vec(const void* __restrict__ src) {
using Package = details::mem_package_t<kBytes, kUnit>;
constexpr auto kBytesPerLoop = sizeof(Package) * kThreads;
constexpr auto kLoopCount = kBytes / kBytesPerLoop;
static_assert(kBytes % kBytesPerLoop == 0, "kBytes must be multiple of 128 bytes");
const auto src_packed = static_cast<const Package*>(src);
const auto lane_id = threadIdx.x % kThreads;
details::storage_vec<Package, kLoopCount> vec;
#pragma unroll kLoopCount
for (std::size_t i = 0; i < kLoopCount; ++i) {
const auto j = i * kThreads + lane_id;
vec.data[i] = details::load_nc(src_packed + j);
}
return vec;
}
template <
std::size_t kBytes,
std::size_t kUnit = details::default_unit_size(kBytes),
std::size_t kThreads = ::device::kWarpThreads,
typename Tp>
__always_inline __device__ void store_vec(void* __restrict__ dst, const Tp& vec) {
using Package = details::mem_package_t<kBytes, kUnit>;
constexpr auto kBytesPerLoop = sizeof(Package) * kThreads;
constexpr auto kLoopCount = kBytes / kBytesPerLoop;
static_assert(kBytes % kBytesPerLoop == 0, "kBytes must be multiple of 128 bytes");
static_assert(std::is_same_v<Tp, details::storage_vec<Package, kLoopCount>>);
const auto dst_packed = static_cast<Package*>(dst);
const auto lane_id = threadIdx.x % kThreads;
#pragma unroll kLoopCount
for (std::size_t i = 0; i < kLoopCount; ++i) {
const auto j = i * kThreads + lane_id;
details::store_nc(dst_packed + j, vec.data[i]);
}
}
} // namespace device::warp