InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
linalg_lstsq.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_LINALG_LSTSQ_H_
2#define INFINI_OPS_BASE_LINALG_LSTSQ_H_
3
4#include <optional>
5#include <string>
6
7#include "operator.h"
8
9namespace infini::ops::linalg {
10
11class Lstsq : public Operator<Lstsq> {
12 public:
13 Lstsq(const Tensor input, const Tensor b, const std::optional<double> rcond,
14 const std::optional<std::string> driver, Tensor solution,
15 Tensor residuals, Tensor rank, Tensor singular_values)
16 : input_shape_{input.shape()},
17 input_strides_{input.strides()},
18 input_type_{input.dtype()},
19 b_shape_{b.shape()},
20 b_strides_{b.strides()},
21 b_type_{b.dtype()},
22 solution_shape_{solution.shape()},
23 solution_strides_{solution.strides()},
24 solution_type_{solution.dtype()},
25 residuals_shape_{residuals.shape()},
26 residuals_strides_{residuals.strides()},
27 residuals_type_{residuals.dtype()},
28 rank_shape_{rank.shape()},
29 rank_strides_{rank.strides()},
30 rank_type_{rank.dtype()},
31 singular_values_shape_{singular_values.shape()},
32 singular_values_strides_{singular_values.strides()},
33 singular_values_type_{singular_values.dtype()},
34 rcond_{rcond},
35 driver_{driver},
36 device_index_{solution.device().index()} {}
37
38 virtual void operator()(const Tensor input, const Tensor b,
39 const std::optional<double> rcond,
40 const std::optional<std::string> driver,
41 Tensor solution, Tensor residuals, Tensor rank,
42 Tensor singular_values) const = 0;
43
44 protected:
45 Tensor::Shape input_shape_;
46
47 Tensor::Strides input_strides_;
48
49 DataType input_type_;
50
51 Tensor::Shape b_shape_;
52
53 Tensor::Strides b_strides_;
54
55 DataType b_type_;
56
57 Tensor::Shape solution_shape_;
58
59 Tensor::Strides solution_strides_;
60
62
63 Tensor::Shape residuals_shape_;
64
65 Tensor::Strides residuals_strides_;
66
68
69 Tensor::Shape rank_shape_;
70
71 Tensor::Strides rank_strides_;
72
73 DataType rank_type_;
74
76
77 Tensor::Strides singular_values_strides_;
78
80
81 std::optional<double> rcond_{};
82
83 std::optional<std::string> driver_{};
84
86};
87
88} // namespace infini::ops::linalg
89
90#endif
Definition generated/include/operator.h:282
Definition linalg_lstsq.h:11
std::optional< double > rcond_
Definition linalg_lstsq.h:81
Tensor::Shape solution_shape_
Definition linalg_lstsq.h:57
Tensor::Strides residuals_strides_
Definition linalg_lstsq.h:65
DataType rank_type_
Definition linalg_lstsq.h:73
Tensor::Strides input_strides_
Definition linalg_lstsq.h:47
Tensor::Strides rank_strides_
Definition linalg_lstsq.h:71
DataType b_type_
Definition linalg_lstsq.h:55
Tensor::Shape input_shape_
Definition linalg_lstsq.h:45
Lstsq(const Tensor input, const Tensor b, const std::optional< double > rcond, const std::optional< std::string > driver, Tensor solution, Tensor residuals, Tensor rank, Tensor singular_values)
Definition linalg_lstsq.h:13
DataType solution_type_
Definition linalg_lstsq.h:61
Tensor::Shape rank_shape_
Definition linalg_lstsq.h:69
Tensor::Shape singular_values_shape_
Definition linalg_lstsq.h:75
DataType residuals_type_
Definition linalg_lstsq.h:67
Tensor::Strides singular_values_strides_
Definition linalg_lstsq.h:77
std::optional< std::string > driver_
Definition linalg_lstsq.h:83
DataType singular_values_type_
Definition linalg_lstsq.h:79
Tensor::Shape residuals_shape_
Definition linalg_lstsq.h:63
int device_index_
Definition linalg_lstsq.h:85
DataType input_type_
Definition linalg_lstsq.h:49
Tensor::Shape b_shape_
Definition linalg_lstsq.h:51
Tensor::Strides solution_strides_
Definition linalg_lstsq.h:59
Tensor::Strides b_strides_
Definition linalg_lstsq.h:53
virtual void operator()(const Tensor input, const Tensor b, const std::optional< double > rcond, const std::optional< std::string > driver, Tensor solution, Tensor residuals, Tensor rank, Tensor singular_values) const =0
Definition linalg_cholesky.h:6
infini::rt::TensorView Tensor
Definition tensor.h:8