From ab2a5af5049a62e3a9a4805c9ba148cc01284722 Mon Sep 17 00:00:00 2001 From: JoeZhang-0x000 <77204052+JoeZhang-0x000@users.noreply.github.com> Date: Fri, 9 Oct 2026 14:29:43 +0800 Subject: [PATCH] perf(metax): vectorize contiguous embedding lookups --- .../cuda/metax/ops/embedding/kernel.cuh | 37 ++++++++++++++++ src/native/cuda/metax/ops/embedding/kernel.h | 43 +++++++++++++++++++ tests/test_embedding.py | 10 +++++ 3 files changed, 90 insertions(+) create mode 100644 src/native/cuda/metax/ops/embedding/kernel.cuh diff --git a/src/native/cuda/metax/ops/embedding/kernel.cuh b/src/native/cuda/metax/ops/embedding/kernel.cuh new file mode 100644 index 000000000..e490d45f3 --- /dev/null +++ b/src/native/cuda/metax/ops/embedding/kernel.cuh @@ -0,0 +1,37 @@ +#ifndef INFINI_OPS_METAX_EMBEDDING_KERNEL_CUH_ +#define INFINI_OPS_METAX_EMBEDDING_KERNEL_CUH_ + +#include +#include + +namespace infini::ops { + +template +struct alignas(16) MetaxEmbeddingPack { + T values[8]; +}; + +template +__global__ void MetaxEmbeddingVectorizedKernel(T* out, const IndexT* indices, + const T* weight, + int64_t weight_row_stride, + size_t vocab_size) { + using Pack = MetaxEmbeddingPack; + const size_t token = blockIdx.x; + const int64_t index = static_cast(indices[token]); + + // Match the generic kernel: invalid indices leave output rows untouched. + if (index < 0 || static_cast(index) >= vocab_size) { + return; + } + + auto output = reinterpret_cast(out + token * 4096); + auto row = reinterpret_cast(weight + index * weight_row_stride); + for (int pack = threadIdx.x; pack < 512; pack += blockDim.x) { + output[pack] = row[pack]; + } +} + +} // namespace infini::ops + +#endif diff --git a/src/native/cuda/metax/ops/embedding/kernel.h b/src/native/cuda/metax/ops/embedding/kernel.h index 279d9e33f..12eadb2b4 100644 --- a/src/native/cuda/metax/ops/embedding/kernel.h +++ b/src/native/cuda/metax/ops/embedding/kernel.h @@ -4,6 +4,7 @@ #include #include "native/cuda/metax/caster.cuh" +#include "native/cuda/metax/ops/embedding/kernel.cuh" #include "native/cuda/metax/runtime_.h" #include "native/cuda/ops/embedding/kernel.h" @@ -14,6 +15,48 @@ class Operator : public CudaEmbedding> { public: using CudaEmbedding>::CudaEmbedding; + + using CudaEmbedding>::operator(); + + 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) { + return; + } + + if (max_norm.has_value() || !input.IsContiguous() || !out.IsContiguous() || + embedding_dim_ != 4096 || weight.stride(1) != 1 || + weight.stride(0) % 8 != 0 || + reinterpret_cast(weight.data()) % 16 != 0 || + reinterpret_cast(out.data()) % 16 != 0) { + CudaEmbedding>::operator()( + input, weight, padding_idx, max_norm, norm_type, scale_grad_by_freq, + sparse, out); + return; + } + + auto cuda_stream = static_cast::Stream>( + stream_ ? stream_ : 0); + DispatchFunc, + ConcatType, ReducedFloatTypes>>( + {static_cast(input_dtype_), + static_cast(weight_dtype_)}, + [&](auto list_tag) { + using IndexT = + TypeMapType(list_tag)>; + using T = TypeMapType(list_tag)>; + MetaxEmbeddingVectorizedKernel + <<>>( + reinterpret_cast(out.data()), + reinterpret_cast(input.data()), + reinterpret_cast(weight.data()), weight.stride(0), + vocab_size_); + }, + "MetaxEmbedding::operator()"); + } }; } // namespace infini::ops diff --git a/tests/test_embedding.py b/tests/test_embedding.py index 4b0e58e8d..6c7486aea 100644 --- a/tests/test_embedding.py +++ b/tests/test_embedding.py @@ -24,6 +24,11 @@ ((2, 4), (32, 8), None, None, None, torch.int64), ((2, 4), (32, 8), (8, 1), None, (32, 8, 1), torch.int32), ((2, 4), (32, 8), None, (1, 32), None, torch.int64), + ((64,), (32, 4096), None, None, None, torch.int64), + ((2, 4), (32, 4096), None, (4104, 1), None, torch.int32), + ((2, 4), (32, 4096), (8, 1), None, None, torch.int64), + ((2, 4), (32, 4096), None, (8192, 2), None, torch.int32), + ((2, 4), (32, 4096), None, None, (65536, 8192, 1), torch.int64), ) ) + tuple( ((2, 3), (8, 4), None, None, None, torch.int64, options) @@ -40,6 +45,11 @@ (None, 1.0, 1.0, False, False, False), (None, 1.0, 3.0, False, False, False), ) +) + ( + ( + (2, 3), (32, 4096), None, None, None, torch.int64, + (None, 1.0, 2.0, False, False, False), + ), )