Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 37 additions & 0 deletions src/native/cuda/metax/ops/embedding/kernel.cuh
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
#ifndef INFINI_OPS_METAX_EMBEDDING_KERNEL_CUH_
#define INFINI_OPS_METAX_EMBEDDING_KERNEL_CUH_

#include <cstddef>
#include <cstdint>

namespace infini::ops {

template <typename T>
struct alignas(16) MetaxEmbeddingPack {
T values[8];
};

template <typename T, typename IndexT>
__global__ void MetaxEmbeddingVectorizedKernel(T* out, const IndexT* indices,
const T* weight,
int64_t weight_row_stride,
size_t vocab_size) {
using Pack = MetaxEmbeddingPack<T>;
const size_t token = blockIdx.x;
const int64_t index = static_cast<int64_t>(indices[token]);

// Match the generic kernel: invalid indices leave output rows untouched.
if (index < 0 || static_cast<size_t>(index) >= vocab_size) {
return;
}

auto output = reinterpret_cast<Pack*>(out + token * 4096);
auto row = reinterpret_cast<const Pack*>(weight + index * weight_row_stride);
for (int pack = threadIdx.x; pack < 512; pack += blockDim.x) {
output[pack] = row[pack];
}
}

} // namespace infini::ops

#endif
43 changes: 43 additions & 0 deletions src/native/cuda/metax/ops/embedding/kernel.h
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
#include <utility>

#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"

Expand All @@ -14,6 +15,48 @@ class Operator<Embedding, Device::Type::kMetax>
: public CudaEmbedding<Runtime<Device::Type::kMetax>> {
public:
using CudaEmbedding<Runtime<Device::Type::kMetax>>::CudaEmbedding;

using CudaEmbedding<Runtime<Device::Type::kMetax>>::operator();

void operator()(const Tensor input, const Tensor weight,
const std::optional<int64_t> padding_idx,
const std::optional<double> 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<uintptr_t>(weight.data()) % 16 != 0 ||
reinterpret_cast<uintptr_t>(out.data()) % 16 != 0) {
CudaEmbedding<Runtime<Device::Type::kMetax>>::operator()(
input, weight, padding_idx, max_norm, norm_type, scale_grad_by_freq,
sparse, out);
return;
}

auto cuda_stream = static_cast<Runtime<Device::Type::kMetax>::Stream>(
stream_ ? stream_ : 0);
DispatchFunc<List<DataType::kInt32, DataType::kInt64>,
ConcatType<List<DataType::kFloat32>, ReducedFloatTypes>>(
{static_cast<int64_t>(input_dtype_),
static_cast<int64_t>(weight_dtype_)},
[&](auto list_tag) {
using IndexT =
TypeMapType<Device::Type::kMetax, ListGet<0>(list_tag)>;
using T = TypeMapType<Device::Type::kMetax, ListGet<1>(list_tag)>;
MetaxEmbeddingVectorizedKernel<T, IndexT>
<<<num_indices_, 256, 0, cuda_stream>>>(
reinterpret_cast<T*>(out.data()),
reinterpret_cast<const IndexT*>(input.data()),
reinterpret_cast<const T*>(weight.data()), weight.stride(0),
vocab_size_);
},
"MetaxEmbedding::operator()");
}
};

} // namespace infini::ops
Expand Down
10 changes: 10 additions & 0 deletions tests/test_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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),
),
)


Expand Down