1#ifndef INFINI_OPS_BASE_EMBEDDING_H_
2#define INFINI_OPS_BASE_EMBEDDING_H_
16 const std::optional<int64_t> padding_idx,
17 const std::optional<double> max_norm,
const double norm_type,
18 const bool scale_grad_by_freq,
const bool sparse,
Tensor out)
36 assert(weight.ndim() == 2 &&
"`Embedding` requires 2D `weight`");
37 assert(out.ndim() == input.ndim() + 1 &&
38 "`Embedding` output rank must be input rank + 1");
40 for (Tensor::Size i = 0; i < input.ndim(); ++i) {
41 assert(out.size(i) == input.size(i) &&
42 "`Embedding` output shape must match `input` on non-last "
47 "`Embedding` output last dim must equal `weight` embedding dim");
50 "`Embedding` supports int32 and int64 indices only");
54 "`Embedding` supports float32, float16, and bfloat16 weights only");
56 "`Embedding` output dtype must match `weight` dtype");
60 "`Embedding` padding_idx must be within the weight rows");
64 :
Embedding{input, weight, std::nullopt, std::nullopt,
65 2.0, false, false, out} {}
69 [[deprecated(
"Use the PyTorch-compatible overload instead.")]]
71 const bool scale_grad_by_freq,
const bool sparse,
Tensor out)
73 std::nullopt, 2.0, scale_grad_by_freq,
77 const std::optional<int64_t> padding_idx,
78 const std::optional<double> max_norm,
79 const double norm_type,
const bool scale_grad_by_freq,
80 const bool sparse,
Tensor out)
const = 0;
83 (*this)(input, weight, std::nullopt, std::nullopt, 2.0,
false,
false, out);
88 [[deprecated(
"Use the PyTorch-compatible overload instead.")]]
90 const int64_t padding_idx,
const bool scale_grad_by_freq,
91 const bool sparse,
Tensor out)
const {
92 (*this)(input, weight, std::optional<int64_t>{padding_idx}, std::nullopt,
93 2.0, scale_grad_by_freq, sparse, out);
96 template <
typename TensorLike>
98 const TensorLike& input,
const TensorLike& weight,
99 const std::optional<int64_t> = std::nullopt,
100 const std::optional<double> = std::nullopt,
103 const bool =
false) {
104 typename TensorLike::Shape out_shape{input.shape()};
105 out_shape.push_back(weight.size(1));
107 return TensorLike::Empty(out_shape, weight.dtype(), weight.device());
111 static Tensor::Size
NumIndices(
const Tensor::Shape& input_shape) {
112 Tensor::Size num_indices = 1;
114 for (Tensor::Size dim : input_shape) {
Definition embedding.h:13
virtual 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 =0
static auto MakeReturnValue(const TensorLike &input, const TensorLike &weight, const std::optional< int64_t >=std::nullopt, const std::optional< double >=std::nullopt, const double=2.0, const bool=false, const bool=false)
Definition embedding.h:97
Tensor::Shape input_shape_
Definition embedding.h:121
double norm_type_
Definition embedding.h:149
Tensor::Strides out_strides_
Definition embedding.h:131
Tensor::Size num_indices_
Definition embedding.h:139
Embedding(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)
Definition embedding.h:15
std::optional< int64_t > padding_idx_
Definition embedding.h:145
DataType weight_dtype_
Definition embedding.h:135
bool sparse_
Definition embedding.h:153
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
Definition embedding.h:89
Tensor::Shape out_shape_
Definition embedding.h:125
Tensor::Shape weight_shape_
Definition embedding.h:123
Tensor::Size embedding_dim_
Definition embedding.h:143
Embedding(const Tensor input, const Tensor weight, const int64_t padding_idx, const bool scale_grad_by_freq, const bool sparse, Tensor out)
Definition embedding.h:70
Tensor::Strides weight_strides_
Definition embedding.h:129
bool scale_grad_by_freq_
Definition embedding.h:151
DataType input_dtype_
Definition embedding.h:133
Tensor::Size vocab_size_
Definition embedding.h:141
Tensor::Strides input_strides_
Definition embedding.h:127
Embedding(const Tensor input, const Tensor weight, Tensor out)
Definition embedding.h:63
static Tensor::Size NumIndices(const Tensor::Shape &input_shape)
Definition embedding.h:111
std::optional< double > max_norm_
Definition embedding.h:147
DataType out_dtype_
Definition embedding.h:137
void operator()(const Tensor input, const Tensor weight, Tensor out) const
Definition embedding.h:82
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8