InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
triangular_solve.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_TRIANGULAR_SOLVE_H_
2#define INFINI_OPS_BASE_TRIANGULAR_SOLVE_H_
3
4#include "operator.h"
5
6namespace infini::ops {
7
8class TriangularSolve : public Operator<TriangularSolve> {
9 public:
10 TriangularSolve(const Tensor input, const Tensor A, const bool upper,
11 const bool transpose, const bool unitriangular, Tensor X,
12 Tensor M)
13 : input_shape_{input.shape()},
14 input_strides_{input.strides()},
15 input_type_{input.dtype()},
16 A_shape_{A.shape()},
17 A_strides_{A.strides()},
18 A_type_{A.dtype()},
19 X_shape_{X.shape()},
20 X_strides_{X.strides()},
21 X_type_{X.dtype()},
22 M_shape_{M.shape()},
23 M_strides_{M.strides()},
24 M_type_{M.dtype()},
25 upper_{upper},
26 transpose_{transpose},
27 unitriangular_{unitriangular},
28 device_index_{X.device().index()} {}
29
30 virtual void operator()(const Tensor input, const Tensor A, const bool upper,
31 const bool transpose, const bool unitriangular,
32 Tensor X, Tensor M) const = 0;
33
34 protected:
35 Tensor::Shape input_shape_;
36
37 Tensor::Strides input_strides_;
38
39 DataType input_type_;
40
41 Tensor::Shape A_shape_;
42
43 Tensor::Strides A_strides_;
44
45 DataType A_type_;
46
47 Tensor::Shape X_shape_;
48
49 Tensor::Strides X_strides_;
50
51 DataType X_type_;
52
53 Tensor::Shape M_shape_;
54
55 Tensor::Strides M_strides_;
56
57 DataType M_type_;
58
59 bool upper_{};
60
61 bool transpose_{};
62
64
66};
67
68} // namespace infini::ops
69
70#endif
Definition generated/include/operator.h:282
Definition triangular_solve.h:8
DataType M_type_
Definition triangular_solve.h:57
Tensor::Strides X_strides_
Definition triangular_solve.h:49
Tensor::Shape A_shape_
Definition triangular_solve.h:41
int device_index_
Definition triangular_solve.h:65
DataType X_type_
Definition triangular_solve.h:51
Tensor::Strides A_strides_
Definition triangular_solve.h:43
Tensor::Shape X_shape_
Definition triangular_solve.h:47
Tensor::Shape input_shape_
Definition triangular_solve.h:35
DataType input_type_
Definition triangular_solve.h:39
Tensor::Shape M_shape_
Definition triangular_solve.h:53
Tensor::Strides M_strides_
Definition triangular_solve.h:55
Tensor::Strides input_strides_
Definition triangular_solve.h:37
DataType A_type_
Definition triangular_solve.h:45
virtual void operator()(const Tensor input, const Tensor A, const bool upper, const bool transpose, const bool unitriangular, Tensor X, Tensor M) const =0
TriangularSolve(const Tensor input, const Tensor A, const bool upper, const bool transpose, const bool unitriangular, Tensor X, Tensor M)
Definition triangular_solve.h:10
bool transpose_
Definition triangular_solve.h:61
bool unitriangular_
Definition triangular_solve.h:63
bool upper_
Definition triangular_solve.h:59
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8