diff --git a/src/base/embedding.h b/src/base/embedding.h index ea73c1a91..c3dd16b0e 100644 --- a/src/base/embedding.h +++ b/src/base/embedding.h @@ -2,6 +2,7 @@ #define INFINI_OPS_BASE_EMBEDDING_H_ #include +#include #include "data_type.h" #include "operator.h" @@ -11,7 +12,9 @@ namespace infini::ops { class Embedding : public Operator { public: - Embedding(const Tensor input, const Tensor weight, const int64_t padding_idx, + Embedding(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) : input_shape_{input.shape()}, weight_shape_{weight.shape()}, @@ -26,6 +29,8 @@ class Embedding : public Operator { vocab_size_{weight.size(0)}, embedding_dim_{weight.size(1)}, padding_idx_{padding_idx}, + max_norm_{max_norm}, + norm_type_{norm_type}, scale_grad_by_freq_{scale_grad_by_freq}, sparse_{sparse} { assert(weight.ndim() == 2 && "`Embedding` requires 2D `weight`"); @@ -49,28 +54,53 @@ class Embedding : public Operator { "`Embedding` supports float32, float16, and bfloat16 weights only"); assert(out_dtype_ == weight_dtype_ && "`Embedding` output dtype must match `weight` dtype"); - assert(padding_idx_ >= -static_cast(vocab_size_) && - padding_idx_ < static_cast(vocab_size_) && + assert((!padding_idx_.has_value() || + (*padding_idx_ >= -static_cast(vocab_size_) && + *padding_idx_ < static_cast(vocab_size_))) && "`Embedding` padding_idx must be within the weight rows"); } Embedding(const Tensor input, const Tensor weight, Tensor out) - : Embedding{input, weight, -1, false, false, out} {} + : Embedding{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.")]] + Embedding(const Tensor input, const Tensor weight, const int64_t padding_idx, + const bool scale_grad_by_freq, const bool sparse, Tensor out) + : Embedding{input, weight, padding_idx, + std::nullopt, 2.0, scale_grad_by_freq, + sparse, out} {} virtual void operator()(const Tensor input, const Tensor weight, - const int64_t padding_idx, - const bool scale_grad_by_freq, const bool sparse, - Tensor out) const = 0; + 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 = 0; void operator()(const Tensor input, const Tensor weight, Tensor out) const { - (*this)(input, weight, -1, false, false, out); + (*this)(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.")]] + void operator()(const Tensor input, const Tensor weight, + const int64_t padding_idx, const bool scale_grad_by_freq, + const bool sparse, Tensor out) const { + (*this)(input, weight, std::optional{padding_idx}, std::nullopt, + 2.0, scale_grad_by_freq, sparse, out); } template - static auto MakeReturnValue(const TensorLike& input, const TensorLike& weight, - const int64_t /*padding_idx*/ = -1, - const bool /*scale_grad_by_freq*/ = false, - const bool /*sparse*/ = false) { + static auto MakeReturnValue( + const TensorLike& input, const TensorLike& weight, + const std::optional /*padding_idx*/ = std::nullopt, + const std::optional /*max_norm*/ = std::nullopt, + const double /*norm_type*/ = 2.0, + const bool /*scale_grad_by_freq*/ = false, + const bool /*sparse*/ = false) { auto out_shape = input.shape(); out_shape.push_back(weight.size(1)); @@ -112,7 +142,11 @@ class Embedding : public Operator { Tensor::Size embedding_dim_{0}; - int64_t padding_idx_{0}; + std::optional padding_idx_{}; + + std::optional max_norm_{}; + + double norm_type_{2.0}; bool scale_grad_by_freq_{false}; diff --git a/src/native/ascend/ops/embedding/kernel.h b/src/native/ascend/ops/embedding/kernel.h index b0a08f4ef..0b6e0ae09 100644 --- a/src/native/ascend/ops/embedding/kernel.h +++ b/src/native/ascend/ops/embedding/kernel.h @@ -1,11 +1,13 @@ #ifndef INFINI_OPS_ASCEND_EMBEDDING_KERNEL_H_ #define INFINI_OPS_ASCEND_EMBEDDING_KERNEL_H_ +#include #include #include "acl/acl.h" #include "aclnn/aclnn_base.h" #include "aclnnop/aclnn_embedding.h" +#include "aclnnop/aclnn_embedding_renorm.h" #include "base/embedding.h" #include "native/ascend/common.h" #include "native/ascend/workspace_pool_.h" @@ -16,9 +18,12 @@ namespace infini::ops { template <> class Operator : public Embedding { public: - Operator(const Tensor input, const Tensor weight, const int64_t padding_idx, + 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, scale_grad_by_freq, sparse, out), + : Embedding(input, weight, padding_idx, max_norm, norm_type, + scale_grad_by_freq, sparse, out), input_cache_(input), weight_cache_(weight), out_cache_(out) { @@ -30,7 +35,16 @@ class Operator : public Embedding { } Operator(const Tensor input, const Tensor weight, Tensor out) - : Operator(input, weight, -1, false, false, 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) {} ~Operator() { if (!ascend::IsAclRuntimeAlive()) return; @@ -41,7 +55,8 @@ class Operator : public Embedding { } void operator()(const Tensor input, const Tensor weight, - const int64_t /*padding_idx*/, + 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 { auto stream = static_cast(stream_); @@ -50,21 +65,44 @@ class Operator : public Embedding { auto t_input = input_cache_.get(const_cast(input.data())); auto t_out = out_cache_.get(out.data()); - if (!executor_) { - auto ret = aclnnEmbeddingGetWorkspaceSize(t_weight, t_input, t_out, - &ws_size_, &executor_); + if (max_norm.has_value() && !renorm_executor_) { + auto ret = aclnnEmbeddingRenormGetWorkspaceSize( + t_weight, t_input, *max_norm, norm_type, &renorm_ws_size_, + &renorm_executor_); + assert(ret == ACL_SUCCESS && + "`aclnnEmbeddingRenormGetWorkspaceSize` failed"); + aclSetAclOpExecutorRepeatable(renorm_executor_); + } else if (max_norm.has_value()) { + aclSetInputTensorAddr(renorm_executor_, 0, t_weight, + const_cast(weight.data())); + aclSetInputTensorAddr(renorm_executor_, 1, t_input, + const_cast(input.data())); + } + + if (!embedding_executor_) { + auto ret = aclnnEmbeddingGetWorkspaceSize( + t_weight, t_input, t_out, &embedding_ws_size_, &embedding_executor_); assert(ret == ACL_SUCCESS && "`aclnnEmbeddingGetWorkspaceSize` failed"); - aclSetAclOpExecutorRepeatable(executor_); + aclSetAclOpExecutorRepeatable(embedding_executor_); } else { - aclSetInputTensorAddr(executor_, 0, t_weight, + aclSetInputTensorAddr(embedding_executor_, 0, t_weight, const_cast(weight.data())); - aclSetInputTensorAddr(executor_, 1, t_input, + aclSetInputTensorAddr(embedding_executor_, 1, t_input, const_cast(input.data())); - aclSetOutputTensorAddr(executor_, 0, t_out, out.data()); + aclSetOutputTensorAddr(embedding_executor_, 0, t_out, out.data()); + } + + const auto workspace_size = std::max(renorm_ws_size_, embedding_ws_size_); + auto& arena = ascend::GetWorkspacePool().Ensure(stream, workspace_size); + + if (max_norm.has_value()) { + auto ret = aclnnEmbeddingRenorm(arena.buf, renorm_ws_size_, + renorm_executor_, stream); + assert(ret == ACL_SUCCESS && "`aclnnEmbeddingRenorm` failed"); } - auto& arena = ascend::GetWorkspacePool().Ensure(stream, ws_size_); - auto ret = aclnnEmbedding(arena.buf, ws_size_, executor_, stream); + auto ret = aclnnEmbedding(arena.buf, embedding_ws_size_, + embedding_executor_, stream); assert(ret == ACL_SUCCESS && "`aclnnEmbedding` failed"); } @@ -75,9 +113,13 @@ class Operator : public Embedding { mutable ascend::AclTensorCache out_cache_; - mutable aclOpExecutor* executor_ = nullptr; + mutable aclOpExecutor* renorm_executor_ = nullptr; + + mutable aclOpExecutor* embedding_executor_ = nullptr; + + mutable uint64_t renorm_ws_size_ = 0; - mutable uint64_t ws_size_ = 0; + mutable uint64_t embedding_ws_size_ = 0; }; } // namespace infini::ops diff --git a/src/native/cuda/ops/embedding/kernel.cuh b/src/native/cuda/ops/embedding/kernel.cuh index c2aa5009f..d80b66564 100644 --- a/src/native/cuda/ops/embedding/kernel.cuh +++ b/src/native/cuda/ops/embedding/kernel.cuh @@ -1,15 +1,51 @@ #ifndef INFINI_OPS_CUDA_EMBEDDING_KERNEL_CUH_ #define INFINI_OPS_CUDA_EMBEDDING_KERNEL_CUH_ +#include #include #include +#include #include +#include "native/cuda/caster.cuh" #include "native/cuda/kernel_commons.cuh" namespace infini::ops { namespace embedding_detail { +struct MaxOp { + __device__ float operator()(float lhs, float rhs) const { + return lhs > rhs ? lhs : rhs; + } +}; + +struct MinOp { + __device__ float operator()(float lhs, float rhs) const { + return lhs < rhs ? lhs : rhs; + } +}; + +template +__device__ float ReduceSum(float value) { + using BlockReduce = cub::BlockReduce; + __shared__ typename BlockReduce::TempStorage temp_storage; + return BlockReduce(temp_storage).Sum(value); +} + +template +__device__ float ReduceMax(float value) { + using BlockReduce = cub::BlockReduce; + __shared__ typename BlockReduce::TempStorage temp_storage; + return BlockReduce(temp_storage).Reduce(value, MaxOp{}); +} + +template +__device__ float ReduceMin(float value) { + using BlockReduce = cub::BlockReduce; + __shared__ typename BlockReduce::TempStorage temp_storage; + return BlockReduce(temp_storage).Reduce(value, MinOp{}); +} + __forceinline__ __device__ bool IsAligned(const void* ptr, size_t alignment) { return (reinterpret_cast(ptr) % alignment) == 0; } @@ -134,6 +170,116 @@ __forceinline__ __device__ void CopyRow(T* __restrict__ dst, } // namespace embedding_detail +template +__global__ void ClearEmbeddingVisitedKernel(int* visited, size_t vocab_size) { + const size_t index = blockIdx.x * block_size + threadIdx.x; + + if (index < vocab_size) { + visited[index] = 0; + } +} + +template +__global__ void EmbeddingRenormKernel( + T* weight, const IndexT* indices, int* visited, size_t num_indices, + size_t input_ndim, const size_t* input_shape, + const ptrdiff_t* input_strides, ptrdiff_t weight_row_stride, + ptrdiff_t weight_col_stride, size_t embedding_dim, size_t vocab_size, + bool input_contiguous, double max_norm, double norm_type) { + const size_t flat_index = blockIdx.x; + + if (flat_index >= num_indices) { + return; + } + + __shared__ size_t row; + __shared__ bool claimed; + + if (threadIdx.x == 0) { + const size_t input_offset = + input_contiguous + ? flat_index + : IndexToOffset(flat_index, input_ndim, input_shape, input_strides); + const IndexT index = indices[input_offset]; + claimed = index >= 0 && static_cast(index) < vocab_size; + + if (claimed) { + row = static_cast(index); + claimed = atomicCAS(visited + row, 0, 1) == 0; + } + } + + __syncthreads(); + + if (!claimed) { + return; + } + + const float p = static_cast(norm_type); + float norm; + + if (isinf(p) && p > 0.0f) { + float local_max = 0.0f; + for (size_t column = threadIdx.x; column < embedding_dim; + column += block_size) { + const float value = Caster::template Cast( + weight[row * weight_row_stride + column * weight_col_stride]); + local_max = fmaxf(local_max, fabsf(value)); + } + norm = embedding_detail::ReduceMax(local_max); + } else if (isinf(p)) { + float local_min = INFINITY; + for (size_t column = threadIdx.x; column < embedding_dim; + column += block_size) { + const float value = Caster::template Cast( + weight[row * weight_row_stride + column * weight_col_stride]); + local_min = fminf(local_min, fabsf(value)); + } + norm = embedding_detail::ReduceMin(local_min); + } else { + float local_sum = 0.0f; + for (size_t column = threadIdx.x; column < embedding_dim; + column += block_size) { + const float value = Caster::template Cast( + weight[row * weight_row_stride + column * weight_col_stride]); + const float abs_value = fabsf(value); + + if (p == 0.0f) { + local_sum += abs_value == 0.0f ? 0.0f : 1.0f; + } else if (p == 1.0f) { + local_sum += abs_value; + } else if (p == 2.0f) { + local_sum += value * value; + } else { + local_sum += powf(abs_value, p); + } + } + + const float sum = embedding_detail::ReduceSum(local_sum); + norm = p == 0.0f ? sum : powf(sum, 1.0f / p); + } + + __shared__ float factor; + if (threadIdx.x == 0) { + const float limit = static_cast(max_norm); + factor = norm > limit ? limit / (norm + 1e-7f) : 1.0f; + } + + __syncthreads(); + + if (factor == 1.0f) { + return; + } + + for (size_t column = threadIdx.x; column < embedding_dim; + column += block_size) { + const auto offset = row * weight_row_stride + column * weight_col_stride; + const float current = Caster::template Cast(weight[offset]); + weight[offset] = Caster::template Cast(current * factor); + } +} + template __global__ void EmbeddingKernel( T* __restrict__ output, const IndexT* __restrict__ indices, diff --git a/src/native/cuda/ops/embedding/kernel.h b/src/native/cuda/ops/embedding/kernel.h index 418ead416..d5d8f7e28 100644 --- a/src/native/cuda/ops/embedding/kernel.h +++ b/src/native/cuda/ops/embedding/kernel.h @@ -19,9 +19,12 @@ template class CudaEmbedding : public Embedding { public: CudaEmbedding(const Tensor input, const Tensor weight, - const int64_t padding_idx, const bool scale_grad_by_freq, - const bool sparse, Tensor out) - : Embedding{input, weight, padding_idx, scale_grad_by_freq, sparse, out}, + 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()}, out_ndim_{out.ndim()}, is_input_contiguous_{input.IsContiguous()}, @@ -60,15 +63,37 @@ class CudaEmbedding : public Embedding { Backend::Memcpy(d_metadata_, metadata.data(), metadata_size, Backend::kMemcpyHostToDevice); + + if (max_norm.has_value() && vocab_size_ > 0) { + Backend::Malloc(reinterpret_cast(&d_visited_), + vocab_size_ * sizeof(*d_visited_)); + } } CudaEmbedding(const Tensor input, const Tensor weight, Tensor out) - : CudaEmbedding(input, weight, -1, false, false, out) {} + : CudaEmbedding(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.")]] + CudaEmbedding(const Tensor input, const Tensor weight, + const int64_t padding_idx, const bool scale_grad_by_freq, + const bool sparse, Tensor out) + : CudaEmbedding(input, weight, padding_idx, std::nullopt, 2.0, + scale_grad_by_freq, sparse, out) {} - ~CudaEmbedding() { Backend::Free(d_metadata_); } + ~CudaEmbedding() { + Backend::Free(d_metadata_); + + if (d_visited_) { + Backend::Free(d_visited_); + } + } void operator()(const Tensor input, const Tensor weight, - const int64_t /*padding_idx*/, + 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) { @@ -78,6 +103,16 @@ class CudaEmbedding : public Embedding { auto cuda_stream = static_cast(stream_ ? stream_ : 0); + constexpr size_t kRenormBlockSize = 256; + if (max_norm.has_value()) { + assert(d_visited_ && "`CudaEmbedding` renorm state is not initialized"); + const size_t clear_grid_size = + utils::CeilDiv(vocab_size_, kRenormBlockSize); + ClearEmbeddingVisitedKernel + <<>>(d_visited_, + vocab_size_); + } + size_t block_size = 256; if (embedding_dim_ <= 64) { block_size = 512; @@ -96,6 +131,17 @@ class CudaEmbedding : public Embedding { TypeMapType(list_tag)>; using T = TypeMapType(list_tag)>; + if (max_norm.has_value()) { + EmbeddingRenormKernel + <<>>( + reinterpret_cast(const_cast(weight.data())), + reinterpret_cast(input.data()), d_visited_, + num_indices_, input_ndim_, d_input_shape_, d_input_strides_, + weight_row_stride_, weight_col_stride_, embedding_dim_, + vocab_size_, is_input_contiguous_, *max_norm, norm_type); + } + EmbeddingKernel <<>>( reinterpret_cast(out.data()), @@ -122,6 +168,8 @@ class CudaEmbedding : public Embedding { Tensor::Stride weight_col_stride_{0}; + int* d_visited_{nullptr}; + std::byte* d_metadata_{nullptr}; Tensor::Size* d_input_shape_{nullptr}; diff --git a/tests/test_embedding.py b/tests/test_embedding.py index f58c2aeab..b40ba0a7f 100644 --- a/tests/test_embedding.py +++ b/tests/test_embedding.py @@ -27,14 +27,17 @@ ) + tuple( ((2, 3), (8, 4), None, None, None, torch.int64, options) for options in ( - (-1, False, False), - (-1, False, True), - (-1, True, False), - (-1, True, True), - (0, False, False), - (0, False, True), - (0, True, False), - (0, True, True), + *( + (padding_idx, None, 2.0, scale_grad_by_freq, sparse, True) + for padding_idx in (-1, 0) + for scale_grad_by_freq in (False, True) + for sparse in (False, True) + ), + (None, None, 2.0, False, False, False), + (-1, None, 2.0, True, True, False), + (None, 0.5, 2.0, False, False, False), + (None, 1.0, 1.0, False, False, False), + (None, 1.0, 3.0, False, False, False), ) ) @@ -77,7 +80,24 @@ def test_embedding( vocab_size = weight_shape[0] embedding_dim = weight_shape[1] output_shape = (*input_shape, embedding_dim) - padding_idx, scale_grad_by_freq, sparse = options or (None, False, False) + if options is None: + padding_idx, max_norm, norm_type, scale_grad_by_freq, sparse = ( + None, + None, + 2.0, + False, + False, + ) + use_legacy_overload = False + else: + ( + padding_idx, + max_norm, + norm_type, + scale_grad_by_freq, + sparse, + use_legacy_overload, + ) = options input = randint_strided( 0 if padding_idx is not None else 1, @@ -94,14 +114,20 @@ def test_embedding( lambda *args, **kwargs: _embedding( *args, padding_idx=padding_idx, + max_norm=max_norm, + norm_type=norm_type, scale_grad_by_freq=scale_grad_by_freq, sparse=sparse, + use_default_overload=options is None, + use_legacy_overload=use_legacy_overload, implementation_index=implementation_index, **kwargs, ), lambda *args, **kwargs: _torch_embedding( *args, padding_idx=padding_idx, + max_norm=max_norm, + norm_type=norm_type, scale_grad_by_freq=scale_grad_by_freq, sparse=sparse, **kwargs, @@ -119,8 +145,12 @@ def _embedding( *, out, padding_idx, + max_norm, + norm_type, scale_grad_by_freq, sparse, + use_default_overload, + use_legacy_overload, implementation_index, ): kwargs = { @@ -128,13 +158,25 @@ def _embedding( "stream": get_stream(input.device), } - if padding_idx is None: + if use_default_overload: infini.ops.embedding(input, weight, out, **kwargs) + elif use_legacy_overload: + infini.ops.embedding( + input, + weight, + padding_idx, + scale_grad_by_freq, + sparse, + out, + **kwargs, + ) else: infini.ops.embedding( input, weight, padding_idx, + max_norm, + norm_type, scale_grad_by_freq, sparse, out, @@ -150,16 +192,25 @@ def _torch_embedding( *, out, padding_idx, + max_norm, + norm_type, scale_grad_by_freq, sparse, ): + if max_norm is not None: + # Use PyTorch's CPU path as the backend-independent renorm reference. + input = input.cpu() + weight = weight.cpu() + result = torch.nn.functional.embedding( input, weight, padding_idx=padding_idx, + max_norm=max_norm, + norm_type=norm_type, scale_grad_by_freq=scale_grad_by_freq, sparse=sparse, ) - out.copy_(result) + out.copy_(result.to(out.device)) return out