InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
linalg_pinv.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_LINALG_PINV_H_
2#define INFINI_OPS_BASE_LINALG_PINV_H_
3
4#include <optional>
5
6#include "operator.h"
7
8namespace infini::ops::linalg {
9
10class Pinv : public Operator<Pinv> {
11 public:
12 Pinv(const Tensor input, const double rcond, const bool hermitian, Tensor out)
13 : input_shape_{input.shape()},
14 input_strides_{input.strides()},
15 input_type_{input.dtype()},
16 out_shape_{out.shape()},
17 out_strides_{out.strides()},
18 out_type_{out.dtype()},
19 rcond_{rcond},
20 hermitian_{hermitian},
21 device_index_{out.device().index()} {}
22
23 Pinv(const Tensor input, const std::optional<Tensor> atol,
24 const std::optional<Tensor> rtol, const bool hermitian, Tensor out)
25 : input_shape_{input.shape()},
26 input_strides_{input.strides()},
27 input_type_{input.dtype()},
28 out_shape_{out.shape()},
29 out_strides_{out.strides()},
30 out_type_{out.dtype()},
31 has_atol_{atol.has_value()},
32 atol_shape_{atol ? Tensor::Shape{atol->shape()} : Tensor::Shape{}},
33 atol_strides_{atol ? Tensor::Strides{atol->strides()}
34 : Tensor::Strides{}},
35 atol_type_{atol ? atol->dtype() : DataType::kFloat32},
36 has_rtol_{rtol.has_value()},
37 rtol_shape_{rtol ? Tensor::Shape{rtol->shape()} : Tensor::Shape{}},
38 rtol_strides_{rtol ? Tensor::Strides{rtol->strides()}
39 : Tensor::Strides{}},
40 rtol_type_{rtol ? rtol->dtype() : DataType::kFloat32},
41 hermitian_{hermitian},
42 device_index_{out.device().index()} {}
43
44 Pinv(const Tensor input, const std::optional<double> atol,
45 const std::optional<double> rtol, const bool hermitian, Tensor out)
46 : input_shape_{input.shape()},
47 input_strides_{input.strides()},
48 input_type_{input.dtype()},
49 out_shape_{out.shape()},
50 out_strides_{out.strides()},
51 out_type_{out.dtype()},
52 hermitian_{hermitian},
53 atol_{atol},
54 rtol_{rtol},
55 device_index_{out.device().index()} {}
56
57 Pinv(const Tensor input, const Tensor rcond, const bool hermitian, Tensor out)
58 : input_shape_{input.shape()},
59 input_strides_{input.strides()},
60 input_type_{input.dtype()},
61 out_shape_{out.shape()},
62 out_strides_{out.strides()},
63 out_type_{out.dtype()},
64 rcond_shape_{rcond.shape()},
65 rcond_strides_{rcond.strides()},
66 rcond_type_{rcond.dtype()},
67 hermitian_{hermitian},
68 device_index_{out.device().index()} {}
69
70 virtual void operator()(const Tensor input, const double rcond,
71 const bool hermitian, Tensor out) const = 0;
72
73 virtual void operator()(const Tensor input, const std::optional<Tensor> atol,
74 const std::optional<Tensor> rtol,
75 const bool hermitian, Tensor out) const = 0;
76
77 virtual void operator()(const Tensor input, const std::optional<double> atol,
78 const std::optional<double> rtol,
79 const bool hermitian, Tensor out) const = 0;
80
81 virtual void operator()(const Tensor input, const Tensor rcond,
82 const bool hermitian, Tensor out) const = 0;
83
84 protected:
85 Tensor::Shape input_shape_;
86
87 Tensor::Strides input_strides_;
88
89 DataType input_type_;
90
91 Tensor::Shape out_shape_;
92
93 Tensor::Strides out_strides_;
94
95 DataType out_type_;
96
97 double rcond_{};
98
99 bool hermitian_{};
100
101 bool has_atol_{false};
102
103 Tensor::Shape atol_shape_;
104
105 Tensor::Strides atol_strides_;
106
107 DataType atol_type_{DataType::kFloat32};
108
109 bool has_rtol_{false};
110
111 Tensor::Shape rtol_shape_;
112
113 Tensor::Strides rtol_strides_;
114
115 DataType rtol_type_{DataType::kFloat32};
116
117 std::optional<double> atol_{};
118
119 std::optional<double> rtol_{};
120
121 Tensor::Shape rcond_shape_;
122
123 Tensor::Strides rcond_strides_;
124
125 DataType rcond_type_;
126
128};
129
130} // namespace infini::ops::linalg
131
132#endif
Definition generated/include/operator.h:282
Definition linalg_pinv.h:10
Pinv(const Tensor input, const double rcond, const bool hermitian, Tensor out)
Definition linalg_pinv.h:12
DataType atol_type_
Definition linalg_pinv.h:107
virtual void operator()(const Tensor input, const std::optional< Tensor > atol, const std::optional< Tensor > rtol, const bool hermitian, Tensor out) const =0
Tensor::Strides atol_strides_
Definition linalg_pinv.h:105
Tensor::Shape atol_shape_
Definition linalg_pinv.h:103
DataType out_type_
Definition linalg_pinv.h:95
bool has_atol_
Definition linalg_pinv.h:101
virtual void operator()(const Tensor input, const Tensor rcond, const bool hermitian, Tensor out) const =0
DataType input_type_
Definition linalg_pinv.h:89
std::optional< double > atol_
Definition linalg_pinv.h:117
Tensor::Shape rtol_shape_
Definition linalg_pinv.h:111
Tensor::Strides rtol_strides_
Definition linalg_pinv.h:113
double rcond_
Definition linalg_pinv.h:97
Tensor::Strides out_strides_
Definition linalg_pinv.h:93
Pinv(const Tensor input, const std::optional< double > atol, const std::optional< double > rtol, const bool hermitian, Tensor out)
Definition linalg_pinv.h:44
virtual void operator()(const Tensor input, const double rcond, const bool hermitian, Tensor out) const =0
bool hermitian_
Definition linalg_pinv.h:99
Tensor::Strides input_strides_
Definition linalg_pinv.h:87
int device_index_
Definition linalg_pinv.h:127
std::optional< double > rtol_
Definition linalg_pinv.h:119
Tensor::Shape input_shape_
Definition linalg_pinv.h:85
virtual void operator()(const Tensor input, const std::optional< double > atol, const std::optional< double > rtol, const bool hermitian, Tensor out) const =0
bool has_rtol_
Definition linalg_pinv.h:109
Tensor::Shape out_shape_
Definition linalg_pinv.h:91
Pinv(const Tensor input, const std::optional< Tensor > atol, const std::optional< Tensor > rtol, const bool hermitian, Tensor out)
Definition linalg_pinv.h:23
Pinv(const Tensor input, const Tensor rcond, const bool hermitian, Tensor out)
Definition linalg_pinv.h:57
Tensor::Shape rcond_shape_
Definition linalg_pinv.h:121
DataType rtol_type_
Definition linalg_pinv.h:115
Tensor::Strides rcond_strides_
Definition linalg_pinv.h:123
DataType rcond_type_
Definition linalg_pinv.h:125
Definition linalg_cholesky.h:6
infini::rt::TensorView Tensor
Definition tensor.h:8