InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
add.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_ADD_H_
2#define INFINI_OPS_BASE_ADD_H_
3
4#include <algorithm>
5#include <cstddef>
6#include <cstdint>
7
8#include "operator.h"
9
10namespace infini::ops {
11
12class Add : public Operator<Add> {
13 public:
14 Add(const Tensor input, const Tensor other, const double alpha, Tensor out)
15 : ndim_{out.ndim()},
16 output_size_{out.numel()},
17 input_type_{input.dtype()},
18 other_type_{other.dtype()},
19 out_type_{out.dtype()},
20 input_shape_{out.shape()},
21 other_shape_{out.shape()},
22 out_shape_{out.shape()},
25 out_strides_{out.strides()},
26 is_input_contiguous_{input.shape() == out.shape() &&
27 input.IsContiguous()},
28 is_other_contiguous_{other.shape() == out.shape() &&
29 other.IsContiguous()},
30 is_out_contiguous_{out.IsContiguous()} {
31 assert(!out.HasBroadcastDim() &&
32 "the output of `Add` should NOT have broadcasted dim!");
33 // TODO(lzm): support mix-precision later using the generic elementwise
34 // framework.
36 "operator `Add` requires all input and output Tensors to have the "
37 "same dtype");
38 if (input_type_ != DataType::kFloat16 &&
39 input_type_ != DataType::kBFloat16 &&
40 input_type_ != DataType::kFloat32 &&
41 input_type_ != DataType::kFloat64) {
42 assert(alpha == static_cast<int64_t>(alpha) &&
43 "operator `Add` requires integral `alpha` for integer tensors");
44 }
45 ValidateBroadcast(input, other, out);
46 }
47
48 Add(const Tensor input, const Tensor other, Tensor out)
49 : Add{input, other, 1.0, out} {}
50
51 virtual void operator()(const Tensor input, const Tensor other,
52 const double alpha, Tensor out) const = 0;
53
54 void operator()(const Tensor input, const Tensor other, Tensor out) const {
55 (*this)(input, other, 1.0, out);
56 }
57
58 template <typename TensorLike>
59 static auto MakeReturnValue(const TensorLike& input,
60 const TensorLike& other) {
61 return MakeReturnValue(input, other, 1.0);
62 }
63
64 template <typename TensorLike>
65 static auto MakeReturnValue(const TensorLike& input, const TensorLike& other,
66 const double /*alpha*/) {
67 auto input_shape = input.shape();
68 auto other_shape = other.shape();
69 auto ndim = std::max(input_shape.size(), other_shape.size());
70 typename TensorLike::Shape out_shape(ndim, 1);
71
72 for (std::size_t i = 0; i < ndim; ++i) {
73 auto input_dim = i < ndim - input_shape.size()
74 ? 1
75 : input_shape[i + input_shape.size() - ndim];
76 auto other_dim = i < ndim - other_shape.size()
77 ? 1
78 : other_shape[i + other_shape.size() - ndim];
79 assert((input_dim == other_dim || input_dim == 1 || other_dim == 1) &&
80 "operator `Add` requires broadcast-compatible input shapes");
81 out_shape[i] = std::max(input_dim, other_dim);
82 }
83
84 return TensorLike::Empty(out_shape, input.dtype(), input.device());
85 }
86
87 protected:
88 static Tensor::Strides BroadcastStrides(const Tensor input,
89 const Tensor out) {
90 assert(input.ndim() <= out.ndim() &&
91 "operator `Add` input rank must not exceed output rank");
92 Tensor::Strides strides(out.ndim(), 0);
93 auto offset = out.ndim() - input.ndim();
94
95 for (Tensor::Size i = 0; i < input.ndim(); ++i) {
96 auto out_dim = i + offset;
97 assert((input.size(i) == 1 || input.size(i) == out.size(out_dim)) &&
98 "operator `Add` input shape is not broadcast-compatible with "
99 "output shape");
100 strides[out_dim] = input.size(i) == 1 ? 0 : input.stride(i);
101 }
102
103 return strides;
104 }
105
106 static void ValidateBroadcast(const Tensor input, const Tensor other,
107 const Tensor out) {
108 for (Tensor::Size i = 0; i < out.ndim(); ++i) {
109 auto input_dim = i < out.ndim() - input.ndim()
110 ? 1
111 : input.size(i + input.ndim() - out.ndim());
112 auto other_dim = i < out.ndim() - other.ndim()
113 ? 1
114 : other.size(i + other.ndim() - out.ndim());
115 assert(out.size(i) == std::max(input_dim, other_dim) &&
116 "operator `Add` output shape must equal the broadcasted input "
117 "shape");
118 }
119 }
120
121 Tensor::Size ndim_{0};
122
123 Tensor::Size output_size_{0};
124
125 const DataType input_type_;
126
127 const DataType other_type_;
128
129 const DataType out_type_;
130
131 Tensor::Shape input_shape_;
132
133 Tensor::Shape other_shape_;
134
135 Tensor::Shape out_shape_;
136
137 Tensor::Strides input_strides_;
138
139 Tensor::Strides other_strides_;
140
141 Tensor::Strides out_strides_;
142
144
146
148};
149
150} // namespace infini::ops
151
152#endif
Definition add.h:12
const DataType other_type_
Definition add.h:127
const DataType input_type_
Definition add.h:125
bool is_out_contiguous_
Definition add.h:147
bool is_other_contiguous_
Definition add.h:145
Tensor::Size output_size_
Definition add.h:123
Tensor::Shape out_shape_
Definition add.h:135
Tensor::Shape input_shape_
Definition add.h:131
Tensor::Strides out_strides_
Definition add.h:141
Tensor::Shape other_shape_
Definition add.h:133
Tensor::Strides other_strides_
Definition add.h:139
bool is_input_contiguous_
Definition add.h:143
const DataType out_type_
Definition add.h:129
virtual void operator()(const Tensor input, const Tensor other, const double alpha, Tensor out) const =0
Tensor::Size ndim_
Definition add.h:121
static auto MakeReturnValue(const TensorLike &input, const TensorLike &other, const double)
Definition add.h:65
static Tensor::Strides BroadcastStrides(const Tensor input, const Tensor out)
Definition add.h:88
Tensor::Strides input_strides_
Definition add.h:137
Add(const Tensor input, const Tensor other, Tensor out)
Definition add.h:48
static void ValidateBroadcast(const Tensor input, const Tensor other, const Tensor out)
Definition add.h:106
void operator()(const Tensor input, const Tensor other, Tensor out) const
Definition add.h:54
static auto MakeReturnValue(const TensorLike &input, const TensorLike &other)
Definition add.h:59
Add(const Tensor input, const Tensor other, const double alpha, Tensor out)
Definition add.h:14
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8