InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
embedding.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_EMBEDDING_H_
2#define INFINI_OPS_BASE_EMBEDDING_H_
3
4#include <cstddef>
5#include <optional>
6
7#include "data_type.h"
8#include "operator.h"
9#include "tensor.h"
10
11namespace infini::ops {
12
13class Embedding : public Operator<Embedding> {
14 public:
15 Embedding(const Tensor input, const Tensor weight,
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)
19 : input_shape_{input.shape()},
20 weight_shape_{weight.shape()},
21 out_shape_{out.shape()},
22 input_strides_{input.strides()},
23 weight_strides_{weight.strides()},
24 out_strides_{out.strides()},
25 input_dtype_{input.dtype()},
26 weight_dtype_{weight.dtype()},
27 out_dtype_{out.dtype()},
29 vocab_size_{weight.size(0)},
30 embedding_dim_{weight.size(1)},
31 padding_idx_{padding_idx},
32 max_norm_{max_norm},
33 norm_type_{norm_type},
34 scale_grad_by_freq_{scale_grad_by_freq},
35 sparse_{sparse} {
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");
39
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 "
43 "dims");
44 }
45
46 assert(out.size(-1) == embedding_dim_ &&
47 "`Embedding` output last dim must equal `weight` embedding dim");
48 assert((input_dtype_ == DataType::kInt32 ||
49 input_dtype_ == DataType::kInt64) &&
50 "`Embedding` supports int32 and int64 indices only");
51 assert((weight_dtype_ == DataType::kFloat32 ||
52 weight_dtype_ == DataType::kFloat16 ||
53 weight_dtype_ == DataType::kBFloat16) &&
54 "`Embedding` supports float32, float16, and bfloat16 weights only");
55 assert(out_dtype_ == weight_dtype_ &&
56 "`Embedding` output dtype must match `weight` dtype");
57 assert((!padding_idx_.has_value() ||
58 (*padding_idx_ >= -static_cast<int64_t>(vocab_size_) &&
59 *padding_idx_ < static_cast<int64_t>(vocab_size_))) &&
60 "`Embedding` padding_idx must be within the weight rows");
61 }
62
63 Embedding(const Tensor input, const Tensor weight, Tensor out)
64 : Embedding{input, weight, std::nullopt, std::nullopt,
65 2.0, false, false, out} {}
66
69 [[deprecated("Use the PyTorch-compatible overload instead.")]]
70 Embedding(const Tensor input, const Tensor weight, const int64_t padding_idx,
71 const bool scale_grad_by_freq, const bool sparse, Tensor out)
72 : Embedding{input, weight, padding_idx,
73 std::nullopt, 2.0, scale_grad_by_freq,
74 sparse, out} {}
75
76 virtual void operator()(const Tensor input, const Tensor weight,
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;
81
82 void operator()(const Tensor input, const Tensor weight, Tensor out) const {
83 (*this)(input, weight, std::nullopt, std::nullopt, 2.0, false, false, out);
84 }
85
88 [[deprecated("Use the PyTorch-compatible overload instead.")]]
89 void operator()(const Tensor input, const Tensor weight,
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);
94 }
95
96 template <typename TensorLike>
97 static auto MakeReturnValue(
98 const TensorLike& input, const TensorLike& weight,
99 const std::optional<int64_t> /*padding_idx*/ = std::nullopt,
100 const std::optional<double> /*max_norm*/ = std::nullopt,
101 const double /*norm_type*/ = 2.0,
102 const bool /*scale_grad_by_freq*/ = false,
103 const bool /*sparse*/ = false) {
104 typename TensorLike::Shape out_shape{input.shape()};
105 out_shape.push_back(weight.size(1));
106
107 return TensorLike::Empty(out_shape, weight.dtype(), weight.device());
108 }
109
110 protected:
111 static Tensor::Size NumIndices(const Tensor::Shape& input_shape) {
112 Tensor::Size num_indices = 1;
113
114 for (Tensor::Size dim : input_shape) {
115 num_indices *= dim;
116 }
117
118 return num_indices;
119 }
120
121 Tensor::Shape input_shape_;
122
123 Tensor::Shape weight_shape_;
124
125 Tensor::Shape out_shape_;
126
127 Tensor::Strides input_strides_;
128
129 Tensor::Strides weight_strides_;
130
131 Tensor::Strides out_strides_;
132
133 DataType input_dtype_;
134
136
137 DataType out_dtype_;
138
139 Tensor::Size num_indices_{0};
140
141 Tensor::Size vocab_size_{0};
142
143 Tensor::Size embedding_dim_{0};
144
145 std::optional<int64_t> padding_idx_{};
146
147 std::optional<double> max_norm_{};
148
149 double norm_type_{2.0};
150
152
153 bool sparse_{false};
154};
155
156} // namespace infini::ops
157
158#endif
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