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