InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
sub.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SUB_H_
2#define INFINI_OPS_BASE_SUB_H_
3
4#include "operator.h"
5
6namespace infini::ops {
7
8class Sub : public Operator<Sub> {
9 public:
10 Sub(const Tensor input, const Tensor other, const double alpha, Tensor out)
11 : input_shape_{input.shape()},
12 input_strides_{input.strides()},
13 input_type_{input.dtype()},
14 other_shape_{other.shape()},
15 other_strides_{other.strides()},
16 other_type_{other.dtype()},
17 out_shape_{out.shape()},
18 out_strides_{out.strides()},
19 out_type_{out.dtype()},
20 alpha_{alpha},
21 device_index_{out.device().index()} {}
22
25 [[deprecated("Use the `(input, other, alpha, out)` overload instead.")]]
26 Sub(Tensor input, const Tensor other, const double alpha)
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 alpha_{alpha},
34 device_index_{input.device().index()} {}
35
36 Sub(const Tensor input, const double other, const double alpha, Tensor out)
37 : input_shape_{input.shape()},
38 input_strides_{input.strides()},
39 input_type_{input.dtype()},
40 out_shape_{out.shape()},
41 out_strides_{out.strides()},
42 out_type_{out.dtype()},
43 alpha_{alpha},
44 other_{other},
45 device_index_{out.device().index()} {}
46
47 virtual void operator()(const Tensor input, const Tensor other,
48 const double alpha, Tensor out) const = 0;
49
52 [[deprecated("Use `operator()(input, other, alpha, out)` instead.")]]
53 virtual void operator()(Tensor input, const Tensor other,
54 const double alpha) const = 0;
55
56 virtual void operator()(const Tensor input, const double other,
57 const double alpha, Tensor out) const = 0;
58
59 protected:
60 Tensor::Shape input_shape_;
61
62 Tensor::Strides input_strides_;
63
64 DataType input_type_;
65
66 Tensor::Shape other_shape_;
67
68 Tensor::Strides other_strides_;
69
70 DataType other_type_;
71
72 Tensor::Shape out_shape_;
73
74 Tensor::Strides out_strides_;
75
76 DataType out_type_;
77
78 double alpha_{};
79
80 double other_{};
81
83};
84
85} // namespace infini::ops
86
87#endif
Definition generated/include/operator.h:282
Definition sub.h:8
Tensor::Strides out_strides_
Definition sub.h:74
virtual void operator()(const Tensor input, const Tensor other, const double alpha, Tensor out) const =0
virtual void operator()(Tensor input, const Tensor other, const double alpha) const =0
Tensor::Shape out_shape_
Definition sub.h:72
double alpha_
Definition sub.h:78
Sub(const Tensor input, const Tensor other, const double alpha, Tensor out)
Definition sub.h:10
DataType input_type_
Definition sub.h:64
Sub(Tensor input, const Tensor other, const double alpha)
Definition sub.h:26
Tensor::Shape other_shape_
Definition sub.h:66
Sub(const Tensor input, const double other, const double alpha, Tensor out)
Definition sub.h:36
int device_index_
Definition sub.h:82
virtual void operator()(const Tensor input, const double other, const double alpha, Tensor out) const =0
double other_
Definition sub.h:80
DataType out_type_
Definition sub.h:76
Tensor::Shape input_shape_
Definition sub.h:60
DataType other_type_
Definition sub.h:70
Tensor::Strides input_strides_
Definition sub.h:62
Tensor::Strides other_strides_
Definition sub.h:68
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8