[Minor] Enhance JIT kernel and add dev docs (#14570)
This commit is contained in:
@@ -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();
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user