InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
div.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_DIV_H_
2#define INFINI_OPS_BASE_DIV_H_
3
4#include <optional>
5#include <string>
6
7#include "operator.h"
8
9namespace infini::ops {
10
11class Div : public Operator<Div> {
12 public:
13 Div(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 Div(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 Div(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
53 [[deprecated("Use the explicit-output rounding-mode overload instead.")]]
54 Div(Tensor input, const Tensor other,
55 const std::optional<std::string> rounding_mode)
56 : input_shape_{input.shape()},
57 input_strides_{input.strides()},
58 input_type_{input.dtype()},
59 other_shape_{other.shape()},
60 other_strides_{other.strides()},
61 other_type_{other.dtype()},
62 rounding_mode_{rounding_mode},
63 device_index_{input.device().index()} {}
64
65 Div(const Tensor input, const double other, Tensor out)
66 : input_shape_{input.shape()},
67 input_strides_{input.strides()},
68 input_type_{input.dtype()},
69 out_shape_{out.shape()},
70 out_strides_{out.strides()},
71 out_type_{out.dtype()},
72 other_{other},
73 device_index_{out.device().index()} {}
74
75 Div(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
101 [[deprecated("Use the explicit-output rounding-mode overload instead.")]]
102 virtual void operator()(
103 Tensor input, const Tensor other,
104 const std::optional<std::string> rounding_mode) const = 0;
105
106 virtual void operator()(const Tensor input, const double other,
107 Tensor out) 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 div.h:11
Tensor::Strides out_strides_
Definition div.h:128
Div(const Tensor input, const Tensor other, const std::optional< std::string > rounding_mode, Tensor out)
Definition div.h:25
Div(Tensor input, const Tensor other)
Definition div.h:42
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
DataType other_type_
Definition div.h:124
Tensor::Shape other_shape_
Definition div.h:120
Tensor::Shape out_shape_
Definition div.h:126
double other_
Definition div.h:134
Tensor::Shape input_shape_
Definition div.h:114
virtual void operator()(const Tensor input, const double other, Tensor out) const =0
DataType out_type_
Definition div.h:130
std::optional< std::string > rounding_mode_
Definition div.h:132
int device_index_
Definition div.h:136
Div(const Tensor input, const double other, const std::optional< std::string > rounding_mode, Tensor out)
Definition div.h:75
virtual void operator()(Tensor input, const Tensor other) const =0
Div(const Tensor input, const Tensor other, Tensor out)
Definition div.h:13
virtual void operator()(Tensor input, const Tensor other, const std::optional< std::string > rounding_mode) const =0
Div(const Tensor input, const double other, Tensor out)
Definition div.h:65
Tensor::Strides other_strides_
Definition div.h:122
virtual void operator()(const Tensor input, const Tensor other, Tensor out) const =0
DataType input_type_
Definition div.h:118
Tensor::Strides input_strides_
Definition div.h:116
Div(Tensor input, const Tensor other, const std::optional< std::string > rounding_mode)
Definition div.h:54
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8