InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
divide.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_DIVIDE_H_
2#define INFINI_OPS_BASE_DIVIDE_H_
3
4#include <optional>
5#include <string>
6
7#include "operator.h"
8
9namespace infini::ops {
10
11class Divide : public Operator<Divide> {
12 public:
13 Divide(const Tensor input, const Tensor other, Tensor out)
14 : input_shape_{input.shape()},
15 input_strides_{input.strides()},
16 input_type_{input.dtype()},
17 other_shape_{other.shape()},
18 other_strides_{other.strides()},
19 other_type_{other.dtype()},
20 out_shape_{out.shape()},
21 out_strides_{out.strides()},
22 out_type_{out.dtype()},
23 device_index_{out.device().index()} {}
24
25 Divide(const Tensor input, const Tensor other,
26 const std::optional<std::string> rounding_mode, Tensor out)
27 : input_shape_{input.shape()},
28 input_strides_{input.strides()},
29 input_type_{input.dtype()},
30 other_shape_{other.shape()},
31 other_strides_{other.strides()},
32 other_type_{other.dtype()},
33 out_shape_{out.shape()},
34 out_strides_{out.strides()},
35 out_type_{out.dtype()},
36 rounding_mode_{rounding_mode},
37 device_index_{out.device().index()} {}
38
41 [[deprecated("Use the `(input, other, out)` overload instead.")]]
42 Divide(Tensor input, const Tensor other)
43 : input_shape_{input.shape()},
44 input_strides_{input.strides()},
45 input_type_{input.dtype()},
46 other_shape_{other.shape()},
47 other_strides_{other.strides()},
48 other_type_{other.dtype()},
49 device_index_{input.device().index()} {}
50
51 Divide(const Tensor input, const double other, Tensor out)
52 : input_shape_{input.shape()},
53 input_strides_{input.strides()},
54 input_type_{input.dtype()},
55 out_shape_{out.shape()},
56 out_strides_{out.strides()},
57 out_type_{out.dtype()},
58 other_{other},
59 device_index_{out.device().index()} {}
60
63 [[deprecated("Use the explicit-output rounding-mode overload instead.")]]
64 Divide(Tensor input, const Tensor other,
65 const std::optional<std::string> rounding_mode)
66 : input_shape_{input.shape()},
67 input_strides_{input.strides()},
68 input_type_{input.dtype()},
69 other_shape_{other.shape()},
70 other_strides_{other.strides()},
71 other_type_{other.dtype()},
72 rounding_mode_{rounding_mode},
73 device_index_{input.device().index()} {}
74
75 Divide(const Tensor input, const double other,
76 const std::optional<std::string> rounding_mode, Tensor out)
77 : input_shape_{input.shape()},
78 input_strides_{input.strides()},
79 input_type_{input.dtype()},
80 out_shape_{out.shape()},
81 out_strides_{out.strides()},
82 out_type_{out.dtype()},
83 rounding_mode_{rounding_mode},
84 other_{other},
85 device_index_{out.device().index()} {}
86
87 virtual void operator()(const Tensor input, const Tensor other,
88 Tensor out) const = 0;
89
90 virtual void operator()(const Tensor input, const Tensor other,
91 const std::optional<std::string> rounding_mode,
92 Tensor out) const = 0;
93
96 [[deprecated("Use `operator()(input, other, out)` instead.")]]
97 virtual void operator()(Tensor input, const Tensor other) const = 0;
98
99 virtual void operator()(const Tensor input, const double other,
100 Tensor out) const = 0;
101
104 [[deprecated("Use the explicit-output rounding-mode overload instead.")]]
105 virtual void operator()(
106 Tensor input, const Tensor other,
107 const std::optional<std::string> rounding_mode) const = 0;
108
109 virtual void operator()(const Tensor input, const double other,
110 const std::optional<std::string> rounding_mode,
111 Tensor out) const = 0;
112
113 protected:
114 Tensor::Shape input_shape_;
115
116 Tensor::Strides input_strides_;
117
118 DataType input_type_;
119
120 Tensor::Shape other_shape_;
121
122 Tensor::Strides other_strides_;
123
124 DataType other_type_;
125
126 Tensor::Shape out_shape_;
127
128 Tensor::Strides out_strides_;
129
130 DataType out_type_;
131
132 std::optional<std::string> rounding_mode_{};
133
134 double other_{};
135
137};
138
139} // namespace infini::ops
140
141#endif
Definition divide.h:11
Tensor::Strides input_strides_
Definition divide.h:116
virtual void operator()(Tensor input, const Tensor other) const =0
virtual void operator()(const Tensor input, const double other, const std::optional< std::string > rounding_mode, Tensor out) const =0
virtual void operator()(const Tensor input, const Tensor other, const std::optional< std::string > rounding_mode, Tensor out) const =0
std::optional< std::string > rounding_mode_
Definition divide.h:132
DataType other_type_
Definition divide.h:124
Divide(const Tensor input, const double other, const std::optional< std::string > rounding_mode, Tensor out)
Definition divide.h:75
DataType out_type_
Definition divide.h:130
virtual void operator()(Tensor input, const Tensor other, const std::optional< std::string > rounding_mode) const =0
Divide(const Tensor input, const Tensor other, Tensor out)
Definition divide.h:13
Divide(const Tensor input, const double other, Tensor out)
Definition divide.h:51
Tensor::Shape input_shape_
Definition divide.h:114
virtual void operator()(const Tensor input, const double other, Tensor out) const =0
Tensor::Shape out_shape_
Definition divide.h:126
Divide(Tensor input, const Tensor other, const std::optional< std::string > rounding_mode)
Definition divide.h:64
Tensor::Strides other_strides_
Definition divide.h:122
DataType input_type_
Definition divide.h:118
Divide(const Tensor input, const Tensor other, const std::optional< std::string > rounding_mode, Tensor out)
Definition divide.h:25
int device_index_
Definition divide.h:136
Tensor::Shape other_shape_
Definition divide.h:120
Divide(Tensor input, const Tensor other)
Definition divide.h:42
virtual void operator()(const Tensor input, const Tensor other, Tensor out) const =0
Tensor::Strides out_strides_
Definition divide.h:128
double other_
Definition divide.h:134
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8