diff --git a/src/native/cambricon/cnnl_utils.h b/src/native/cambricon/cnnl_utils.h new file mode 100644 index 000000000..a32c65f10 --- /dev/null +++ b/src/native/cambricon/cnnl_utils.h @@ -0,0 +1,127 @@ +#ifndef INFINI_OPS_CAMBRICON_CNNL_UTILS_H_ +#define INFINI_OPS_CAMBRICON_CNNL_UTILS_H_ + +#include +#include +#include +#include +#include +#include + +#include "native/cambricon/common.h" +#include "tensor.h" + +namespace infini::ops::cnnl_utils { + +struct HandleDeleter { + using pointer = cnnlHandle_t; + + void operator()(pointer handle) const noexcept { + if (handle) { + (void)cnnlDestroy(handle); + } + } +}; + +using Handle = + std::unique_ptr, HandleDeleter>; + +inline Handle CreateHandle() { + cnnlHandle_t handle{nullptr}; + [[maybe_unused]] const auto status = cnnlCreate(&handle); + assert(status == CNNL_STATUS_SUCCESS && "`cnnlCreate` failed."); + + return Handle{handle}; +} + +struct TensorDescriptorDeleter { + using pointer = cnnlTensorDescriptor_t; + + void operator()(pointer desc) const noexcept { + if (desc) { + (void)cnnlDestroyTensorDescriptor(desc); + } + } +}; + +using TensorDescriptor = + std::unique_ptr, + TensorDescriptorDeleter>; + +inline TensorDescriptor CreateTensorDescriptor() { + cnnlTensorDescriptor_t desc{nullptr}; + [[maybe_unused]] const auto status = cnnlCreateTensorDescriptor(&desc); + assert(status == CNNL_STATUS_SUCCESS && + "`cnnlCreateTensorDescriptor` failed."); + + return TensorDescriptor{desc}; +} + +namespace detail { + +template +int CheckedInt(Integer value) { + static_assert(std::is_integral_v); + + [[maybe_unused]] bool out_of_range{false}; + if constexpr (std::is_signed_v) { + const auto wide = static_cast(value); + out_of_range = wide < std::numeric_limits::min() || + wide > std::numeric_limits::max(); + } else { + const auto wide = static_cast(value); + out_of_range = + wide > static_cast(std::numeric_limits::max()); + } + + assert(!out_of_range && + "`CNNL tensor descriptor` value does not fit in `int`."); + + return static_cast(value); +} + +template +std::vector CheckedIntVector(const Values& values) { + std::vector result; + result.reserve(values.size()); + for (const auto value : values) { + result.push_back(CheckedInt(value)); + } + return result; +} + +} // namespace detail + +inline void SetTensorDescriptor(cnnlTensorDescriptor_t desc, DataType dtype, + const Tensor::Shape& shape, + const Tensor::Strides& strides) { + assert(!shape.empty() && shape.size() == strides.size() && + "`CNNL tensor descriptor` requires matching non-empty shape and " + "strides."); + + const auto cnnl_dtype = GetDataType(dtype); + assert(cnnl_dtype != CNNL_DTYPE_INVALID && + "`CNNL tensor descriptor` does not support this data type."); + + const auto cnnl_shape = detail::CheckedIntVector(shape); + const auto cnnl_strides = detail::CheckedIntVector(strides); + const auto ndim = detail::CheckedInt(shape.size()); + + [[maybe_unused]] const auto status = + cnnlSetTensorDescriptorEx(desc, CNNL_LAYOUT_ARRAY, cnnl_dtype, ndim, + cnnl_shape.data(), cnnl_strides.data()); + assert(status == CNNL_STATUS_SUCCESS && + "`cnnlSetTensorDescriptorEx` failed."); +} + +inline TensorDescriptor MakeTensorDescriptor(DataType dtype, + const Tensor::Shape& shape, + const Tensor::Strides& strides) { + auto desc = CreateTensorDescriptor(); + SetTensorDescriptor(desc.get(), dtype, shape, strides); + return desc; +} + +} // namespace infini::ops::cnnl_utils + +#endif diff --git a/src/native/cambricon/cnrt_utils.h b/src/native/cambricon/cnrt_utils.h new file mode 100644 index 000000000..3cd18f6ce --- /dev/null +++ b/src/native/cambricon/cnrt_utils.h @@ -0,0 +1,36 @@ +#ifndef INFINI_OPS_CAMBRICON_CNRT_UTILS_H_ +#define INFINI_OPS_CAMBRICON_CNRT_UTILS_H_ + +#include + +#include +#include + +namespace infini::ops::cnrt_utils { + +struct DeviceBufferDeleter { + using pointer = void*; + + void operator()(pointer buffer) const noexcept { + if (buffer) { + (void)cnrtFree(buffer); + } + } +}; + +using DeviceBuffer = std::unique_ptr; + +inline DeviceBuffer AllocateDeviceBuffer(std::size_t size) { + if (size == 0) { + return {}; + } + + void* buffer{nullptr}; + CNRT_CHECK(cnrtMalloc(&buffer, size)); + + return DeviceBuffer{buffer}; +} + +} // namespace infini::ops::cnrt_utils + +#endif diff --git a/src/native/cambricon/ops/embedding/kernel.h b/src/native/cambricon/ops/embedding/kernel.h new file mode 100644 index 000000000..26f613e6b --- /dev/null +++ b/src/native/cambricon/ops/embedding/kernel.h @@ -0,0 +1,190 @@ +#ifndef INFINI_OPS_CAMBRICON_EMBEDDING_KERNEL_H_ +#define INFINI_OPS_CAMBRICON_EMBEDDING_KERNEL_H_ + +#include +#include +#include + +// clang-format off +#include +#include +// clang-format on + +#include "base/embedding.h" +#include "native/cambricon/cnnl_utils.h" +#include "native/cambricon/cnrt_utils.h" +#include "native/cambricon/common.h" + +namespace infini::ops { + +namespace embedding_detail { + +constexpr std::size_t AlignUp(std::size_t value, std::size_t alignment) { + return (value + alignment - 1) / alignment * alignment; +} + +constexpr std::size_t MetadataSize(std::size_t input_ndim) { + return input_ndim * sizeof(std::size_t) + + input_ndim * sizeof(std::ptrdiff_t) + + (input_ndim + 1) * sizeof(std::ptrdiff_t); +} + +constexpr std::size_t VisitedOffset(std::size_t input_ndim) { + return AlignUp(MetadataSize(input_ndim), alignof(std::int32_t)); +} + +} // namespace embedding_detail + +inline std::size_t EmbeddingWorkspaceSize(std::size_t input_ndim, + std::size_t vocab_size, + bool apply_max_norm, + bool launch_custom_forward) { + if (!apply_max_norm && !launch_custom_forward) { + return 0; + } + + const auto visited_size = + apply_max_norm ? vocab_size * sizeof(std::int32_t) : 0; + + return embedding_detail::VisitedOffset(input_ndim) + visited_size; +} + +void EmbeddingKernelLaunch( + void* workspace, DataType input_dtype, DataType weight_dtype, + int core_per_cluster, int cluster_count, cnrtQueue_t queue, void* output, + const void* input, void* weight, std::size_t num_indices, + std::size_t input_ndim, const std::size_t* input_shape, + const std::ptrdiff_t* input_strides, const std::ptrdiff_t* output_strides, + std::ptrdiff_t weight_row_stride, std::ptrdiff_t weight_col_stride, + std::size_t embedding_dim, std::size_t vocab_size, bool apply_max_norm, + float max_norm, float norm_type, bool launch_custom_forward); + +template <> +class Operator : public Embedding { + public: + Operator(const Tensor input, const Tensor weight, + const std::optional padding_idx, + const std::optional max_norm, const double norm_type, + const bool scale_grad_by_freq, const bool sparse, Tensor out) + : Embedding{input, weight, padding_idx, + max_norm, norm_type, scale_grad_by_freq, + sparse, out}, + input_ndim_{input.ndim()}, + weight_row_stride_{weight.stride(0)}, + weight_col_stride_{weight.stride(1)}, + use_cnnl_forward_{input.ndim() > 0 && weight.size(0) > 0 && + input.IsContiguous() && weight.IsContiguous() && + out.IsContiguous()} { + cnrt_utils::GetLaunchConfig(input.device(), &core_per_cluster_, + &cluster_count_); + + if (use_cnnl_forward_) { + cnnl_handle_ = cnnl_utils::CreateHandle(); + input_desc_ = cnnl_utils::MakeTensorDescriptor(input_dtype_, input_shape_, + input_strides_); + weight_desc_ = cnnl_utils::MakeTensorDescriptor( + weight_dtype_, weight_shape_, weight_strides_); + out_desc_ = cnnl_utils::MakeTensorDescriptor(out_dtype_, out_shape_, + out_strides_); + } + + workspace_size_ = EmbeddingWorkspaceSize( + input_ndim_, vocab_size_, max_norm.has_value(), !use_cnnl_forward_); + default_workspace_ = cnrt_utils::AllocateDeviceBuffer(workspace_size_); + } + + Operator(const Tensor input, const Tensor weight, Tensor out) + : Operator(input, weight, std::nullopt, std::nullopt, 2.0, false, false, + out) {} + + /// \deprecated Use the overload that also accepts `max_norm` and + /// `norm_type` instead. + [[deprecated("Use the PyTorch-compatible overload instead.")]] + Operator(const Tensor input, const Tensor weight, const int64_t padding_idx, + const bool scale_grad_by_freq, const bool sparse, Tensor out) + : Operator(input, weight, padding_idx, std::nullopt, 2.0, + scale_grad_by_freq, sparse, out) {} + + std::size_t workspace_size_in_bytes() const override { + return workspace_size_; + } + + void operator()(const Tensor input, const Tensor weight, + const std::optional /*padding_idx*/, + const std::optional max_norm, const double norm_type, + const bool /*scale_grad_by_freq*/, const bool /*sparse*/, + Tensor out) const override { + if (num_indices_ == 0 || embedding_dim_ == 0) { + return; + } + + assert(max_norm.has_value() == max_norm_.has_value() && + "`CambriconEmbedding` max_norm presence changed after creation"); + + auto queue = static_cast(stream_ ? stream_ : 0); + const bool launch_custom_forward = !use_cnnl_forward_; + + if (max_norm.has_value() || launch_custom_forward) { + void* workspace = workspace_ ? workspace_ : default_workspace_.get(); + [[maybe_unused]] const auto workspace_size = + workspace_ ? workspace_size_in_bytes_ : workspace_size_; + assert(workspace && workspace_size >= workspace_size_ && + "`CambriconEmbedding` requires a sufficiently large workspace."); + + EmbeddingKernelLaunch( + workspace, input.dtype(), weight.dtype(), core_per_cluster_, + cluster_count_, queue, out.data(), input.data(), + const_cast(weight.data()), num_indices_, input_ndim_, + input_shape_.data(), input_strides_.data(), out_strides_.data(), + weight_row_stride_, weight_col_stride_, embedding_dim_, vocab_size_, + max_norm.has_value(), static_cast(max_norm.value_or(0.0)), + static_cast(norm_type), launch_custom_forward); + } + + if (launch_custom_forward) { + return; + } + + [[maybe_unused]] const auto set_queue_status = + cnnlSetQueue(cnnl_handle_.get(), queue); + assert(set_queue_status == CNNL_STATUS_SUCCESS && "`cnnlSetQueue` failed."); + + // A non-negative CNNL padding_idx zeroes the corresponding weight row. + // InfiniOps forward semantics only use padding_idx for backward behavior. + [[maybe_unused]] const auto embedding_status = cnnlEmbeddingForward_v2( + cnnl_handle_.get(), weight_desc_.get(), weight.data(), + input_desc_.get(), input.data(), -1, nullptr, nullptr, out_desc_.get(), + out.data()); + assert(embedding_status == CNNL_STATUS_SUCCESS && + "`cnnlEmbeddingForward_v2` failed."); + } + + private: + std::size_t input_ndim_{0}; + + std::ptrdiff_t weight_row_stride_{0}; + + std::ptrdiff_t weight_col_stride_{0}; + + int core_per_cluster_{0}; + + int cluster_count_{0}; + + std::size_t workspace_size_{0}; + + cnrt_utils::DeviceBuffer default_workspace_{}; + + bool use_cnnl_forward_{false}; + + cnnl_utils::Handle cnnl_handle_{}; + + cnnl_utils::TensorDescriptor input_desc_{}; + + cnnl_utils::TensorDescriptor weight_desc_{}; + + cnnl_utils::TensorDescriptor out_desc_{}; +}; + +} // namespace infini::ops + +#endif diff --git a/src/native/cambricon/ops/embedding/kernel.mlu b/src/native/cambricon/ops/embedding/kernel.mlu new file mode 100644 index 000000000..3d1c45fdb --- /dev/null +++ b/src/native/cambricon/ops/embedding/kernel.mlu @@ -0,0 +1,296 @@ +#include +#include + +#include "dispatcher.h" +#include "kernel.h" +#include "native/cambricon/data_type_.h" + +namespace infini::ops { +namespace embedding_detail { + +struct WorkspaceLayout { + std::size_t* input_shape; + std::ptrdiff_t* input_strides; + std::ptrdiff_t* output_strides; + std::int32_t* visited; +}; + +static WorkspaceLayout GetWorkspaceLayout(void* workspace, + std::size_t input_ndim) { + auto* base = static_cast(workspace); + auto* input_shape = reinterpret_cast(base); + auto* input_strides = reinterpret_cast( + base + input_ndim * sizeof(std::size_t)); + auto* output_strides = reinterpret_cast( + reinterpret_cast(input_strides) + + input_ndim * sizeof(std::ptrdiff_t)); + auto* visited = + reinterpret_cast(base + VisitedOffset(input_ndim)); + + return {input_shape, input_strides, output_strides, visited}; +} + +static void CopyMetadataAsync(const WorkspaceLayout& layout, cnrtQueue_t queue, + std::size_t input_ndim, + const std::size_t* input_shape, + const std::ptrdiff_t* input_strides, + const std::ptrdiff_t* output_strides, + bool copy_output_strides) { + if (input_ndim > 0) { + CNRT_CHECK(cnrtMemcpyAsync( + layout.input_shape, const_cast(input_shape), + input_ndim * sizeof(std::size_t), queue, cnrtMemcpyHostToDev)); + + CNRT_CHECK(cnrtMemcpyAsync( + layout.input_strides, const_cast(input_strides), + input_ndim * sizeof(std::ptrdiff_t), queue, cnrtMemcpyHostToDev)); + } + + if (copy_output_strides) { + CNRT_CHECK(cnrtMemcpyAsync( + layout.output_strides, const_cast(output_strides), + (input_ndim + 1) * sizeof(std::ptrdiff_t), queue, cnrtMemcpyHostToDev)); + } +} + +__mlu_func__ ptrdiff_t LogicalOffset(size_t logical_index, size_t ndim, + const size_t* shape, + const ptrdiff_t* strides) { + ptrdiff_t offset = 0; + + for (size_t dim = ndim; dim > 0; --dim) { + const size_t axis = dim - 1; + const size_t coordinate = logical_index % shape[axis]; + logical_index /= shape[axis]; + offset += static_cast(coordinate) * strides[axis]; + } + + return offset; +} + +template +__mlu_func__ float LoadFloat(const T* source) { + if constexpr (std::is_same::value) { + return __half2float(*source); + } else if constexpr (std::is_same::value) { + return __bfloat162float(*source); + } else { + return *source; + } +} + +template +__mlu_func__ void StoreFloat(T* destination, float value) { + *destination = value; +} + +__mlu_global__ void EmbeddingClearVisitedKernel(int32_t* visited, + size_t vocab_size) { + if (__is_mpu()) { + return; + } + + for (size_t row = taskId; row < vocab_size; row += taskDim) { + visited[row] = 0; + } +} + +template +__mlu_global__ void EmbeddingMarkVisitedKernel( + const IndexT* indices, int32_t* visited, size_t num_indices, + size_t input_ndim, const size_t* input_shape, + const ptrdiff_t* input_strides, size_t vocab_size) { + if (__is_mpu()) { + return; + } + + for (size_t logical_index = taskId; logical_index < num_indices; + logical_index += taskDim) { + const ptrdiff_t input_offset = + LogicalOffset(logical_index, input_ndim, input_shape, input_strides); + const IndexT index = indices[input_offset]; + + if (index < static_cast(0) || + index >= static_cast(vocab_size)) { + continue; + } + + __bang_atomic_reduce_or(visited + static_cast(index), 1); + } +} + +template +__mlu_global__ void EmbeddingRenormKernel(T* weight, const int32_t* visited, + ptrdiff_t weight_row_stride, + ptrdiff_t weight_col_stride, + size_t embedding_dim, + size_t vocab_size, float max_norm, + float norm_type) { + if (__is_mpu()) { + return; + } + + for (size_t row = taskId; row < vocab_size; row += taskDim) { + if (visited[row] == 0) { + continue; + } + + float accumulator = 0.0f; + + if (isinf(norm_type) && norm_type < 0.0f) { + accumulator = INFINITY; + } + + for (size_t column = 0; column < embedding_dim; ++column) { + const ptrdiff_t weight_offset = + static_cast(row) * weight_row_stride + + static_cast(column) * weight_col_stride; + const float value = LoadFloat(weight + weight_offset); + const float absolute = fabsf(value); + + if (isinf(norm_type) && norm_type > 0.0f) { + accumulator = fmaxf(accumulator, absolute); + } else if (isinf(norm_type)) { + accumulator = fminf(accumulator, absolute); + } else if (norm_type == 0.0f) { + accumulator += absolute == 0.0f ? 0.0f : 1.0f; + } else if (norm_type == 1.0f) { + accumulator += absolute; + } else if (norm_type == 2.0f) { + accumulator += value * value; + } else { + accumulator += powf(absolute, norm_type); + } + } + + const float norm = norm_type == 0.0f || isinf(norm_type) + ? accumulator + : powf(accumulator, 1.0f / norm_type); + + if (norm <= max_norm) { + continue; + } + + const float factor = max_norm / (norm + 1e-7f); + for (size_t column = 0; column < embedding_dim; ++column) { + const ptrdiff_t weight_offset = + static_cast(row) * weight_row_stride + + static_cast(column) * weight_col_stride; + const float value = LoadFloat(weight + weight_offset); + StoreFloat(weight + weight_offset, value * factor); + } + } +} + +template +__mlu_global__ void EmbeddingForwardKernel( + T* output, const IndexT* indices, const T* weight, size_t num_indices, + size_t input_ndim, const size_t* input_shape, + const ptrdiff_t* input_strides, const ptrdiff_t* output_strides, + ptrdiff_t weight_row_stride, ptrdiff_t weight_col_stride, + size_t embedding_dim, size_t vocab_size) { + if (__is_mpu()) { + return; + } + + const size_t element_count = num_indices * embedding_dim; + + for (size_t logical_element = taskId; logical_element < element_count; + logical_element += taskDim) { + const size_t logical_index = logical_element / embedding_dim; + const size_t column = logical_element % embedding_dim; + const ptrdiff_t input_offset = + LogicalOffset(logical_index, input_ndim, input_shape, input_strides); + const IndexT index = indices[input_offset]; + + const ptrdiff_t output_offset = + LogicalOffset(logical_index, input_ndim, input_shape, output_strides) + + static_cast(column) * output_strides[input_ndim]; + + if (index < static_cast(0) || + index >= static_cast(vocab_size)) { + output[output_offset] = static_cast(0.0f); + } else { + const ptrdiff_t weight_offset = + static_cast(index) * weight_row_stride + + static_cast(column) * weight_col_stride; + output[output_offset] = weight[weight_offset]; + } + } +} + +} // namespace embedding_detail + +void EmbeddingKernelLaunch( + void* workspace, DataType input_dtype, DataType weight_dtype, + int core_per_cluster, int cluster_count, cnrtQueue_t queue, void* output, + const void* input, void* weight, std::size_t num_indices, + std::size_t input_ndim, const std::size_t* input_shape, + const std::ptrdiff_t* input_strides, const std::ptrdiff_t* output_strides, + std::ptrdiff_t weight_row_stride, std::ptrdiff_t weight_col_stride, + std::size_t embedding_dim, std::size_t vocab_size, bool apply_max_norm, + float max_norm, float norm_type, bool launch_custom_forward) { + if (!apply_max_norm && !launch_custom_forward) { + return; + } + + const auto layout = + embedding_detail::GetWorkspaceLayout(workspace, input_ndim); + embedding_detail::CopyMetadataAsync(layout, queue, input_ndim, input_shape, + input_strides, output_strides, + launch_custom_forward); + + const cnrtDim3_t kernel_dim = {static_cast(core_per_cluster), + static_cast(cluster_count), 1}; + + if (apply_max_norm && vocab_size > 0) { + (void)cnrtGetLastError(); + embedding_detail:: + EmbeddingClearVisitedKernel<<>>( + layout.visited, vocab_size); + CNRT_CHECK(cnrtGetLastError()); + } + + DispatchFunc< + Device::Type::kCambricon, List, + List>( + {input_dtype, weight_dtype}, + [&](auto input_tag, auto weight_tag) { + using IndexT = typename decltype(input_tag)::type; + using T = typename decltype(weight_tag)::type; + + if (apply_max_norm && vocab_size > 0) { + (void)cnrtGetLastError(); + embedding_detail::EmbeddingMarkVisitedKernel + <<>>( + reinterpret_cast(input), layout.visited, + num_indices, input_ndim, layout.input_shape, + layout.input_strides, vocab_size); + CNRT_CHECK(cnrtGetLastError()); + + (void)cnrtGetLastError(); + embedding_detail::EmbeddingRenormKernel + <<>>( + reinterpret_cast(weight), layout.visited, + weight_row_stride, weight_col_stride, embedding_dim, + vocab_size, max_norm, norm_type); + CNRT_CHECK(cnrtGetLastError()); + } + + if (launch_custom_forward) { + (void)cnrtGetLastError(); + embedding_detail::EmbeddingForwardKernel + <<>>( + reinterpret_cast(output), + reinterpret_cast(input), + reinterpret_cast(weight), num_indices, input_ndim, + layout.input_shape, layout.input_strides, + layout.output_strides, weight_row_stride, weight_col_stride, + embedding_dim, vocab_size); + CNRT_CHECK(cnrtGetLastError()); + } + }, + "CambriconEmbeddingKernelLaunch"); +} + +} // namespace infini::ops diff --git a/tests/test_embedding.py b/tests/test_embedding.py index b40ba0a7f..4b0e58e8d 100644 --- a/tests/test_embedding.py +++ b/tests/test_embedding.py @@ -17,6 +17,7 @@ _TEST_CASES = tuple( (*case, None) for case in ( + ((), (8, 4), None, None, None, torch.int64), ((1, 5), (32000, 4), None, None, None, torch.int64), ((2, 10), (32000, 2048), None, None, None, torch.int32), ((1, 5), (10, 10), None, None, None, torch.int64),