1#ifndef INFINI_OPS_BASE_ADD_H_
2#define INFINI_OPS_BASE_ADD_H_
27 input.IsContiguous()},
29 other.IsContiguous()},
31 assert(!out.HasBroadcastDim() &&
32 "the output of `Add` should NOT have broadcasted dim!");
36 "operator `Add` requires all input and output Tensors to have the "
42 assert(alpha ==
static_cast<int64_t
>(alpha) &&
43 "operator `Add` requires integral `alpha` for integer tensors");
49 :
Add{input, other, 1.0, out} {}
52 const double alpha,
Tensor out)
const = 0;
55 (*this)(input, other, 1.0, out);
58 template <
typename TensorLike>
60 const TensorLike& other) {
64 template <
typename TensorLike>
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);
72 for (std::size_t i = 0; i < ndim; ++i) {
73 auto input_dim = i < ndim - input_shape.size()
75 : input_shape[i + input_shape.size() - ndim];
76 auto other_dim = i < ndim - other_shape.size()
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);
84 return TensorLike::Empty(out_shape, input.dtype(), input.device());
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();
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 "
100 strides[out_dim] = input.size(i) == 1 ? 0 : input.stride(i);
108 for (Tensor::Size i = 0; i < out.ndim(); ++i) {
109 auto input_dim = i < out.ndim() - input.ndim()
111 : input.size(i + input.ndim() - out.ndim());
112 auto other_dim = i < out.ndim() - other.ndim()
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 "
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