InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
lu_unpack.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_LU_UNPACK_H_
2#define INFINI_OPS_BASE_LU_UNPACK_H_
3
4#include "operator.h"
5
6namespace infini::ops {
7
8class LuUnpack : public Operator<LuUnpack> {
9 public:
10 LuUnpack(const Tensor LU_data, const Tensor LU_pivots, const bool unpack_data,
11 const bool unpack_pivots, Tensor P, Tensor L, Tensor U)
12 : LU_data_shape_{LU_data.shape()},
13 LU_data_strides_{LU_data.strides()},
14 LU_data_type_{LU_data.dtype()},
15 LU_pivots_shape_{LU_pivots.shape()},
16 LU_pivots_strides_{LU_pivots.strides()},
17 LU_pivots_type_{LU_pivots.dtype()},
18 P_shape_{P.shape()},
19 P_strides_{P.strides()},
20 P_type_{P.dtype()},
21 L_shape_{L.shape()},
22 L_strides_{L.strides()},
23 L_type_{L.dtype()},
24 U_shape_{U.shape()},
25 U_strides_{U.strides()},
26 U_type_{U.dtype()},
27 unpack_data_{unpack_data},
28 unpack_pivots_{unpack_pivots},
29 device_index_{P.device().index()} {}
30
31 virtual void operator()(const Tensor LU_data, const Tensor LU_pivots,
32 const bool unpack_data, const bool unpack_pivots,
33 Tensor P, Tensor L, Tensor U) const = 0;
34
35 protected:
36 Tensor::Shape LU_data_shape_;
37
38 Tensor::Strides LU_data_strides_;
39
40 DataType LU_data_type_;
41
42 Tensor::Shape LU_pivots_shape_;
43
44 Tensor::Strides LU_pivots_strides_;
45
47
48 Tensor::Shape P_shape_;
49
50 Tensor::Strides P_strides_;
51
52 DataType P_type_;
53
54 Tensor::Shape L_shape_;
55
56 Tensor::Strides L_strides_;
57
58 DataType L_type_;
59
60 Tensor::Shape U_shape_;
61
62 Tensor::Strides U_strides_;
63
64 DataType U_type_;
65
67
69
71};
72
73} // namespace infini::ops
74
75#endif
Definition lu_unpack.h:8
Tensor::Strides L_strides_
Definition lu_unpack.h:56
Tensor::Shape L_shape_
Definition lu_unpack.h:54
Tensor::Strides LU_data_strides_
Definition lu_unpack.h:38
Tensor::Shape LU_pivots_shape_
Definition lu_unpack.h:42
Tensor::Shape P_shape_
Definition lu_unpack.h:48
Tensor::Strides P_strides_
Definition lu_unpack.h:50
virtual void operator()(const Tensor LU_data, const Tensor LU_pivots, const bool unpack_data, const bool unpack_pivots, Tensor P, Tensor L, Tensor U) const =0
DataType P_type_
Definition lu_unpack.h:52
DataType LU_data_type_
Definition lu_unpack.h:40
bool unpack_data_
Definition lu_unpack.h:66
bool unpack_pivots_
Definition lu_unpack.h:68
Tensor::Strides U_strides_
Definition lu_unpack.h:62
DataType LU_pivots_type_
Definition lu_unpack.h:46
DataType L_type_
Definition lu_unpack.h:58
Tensor::Strides LU_pivots_strides_
Definition lu_unpack.h:44
LuUnpack(const Tensor LU_data, const Tensor LU_pivots, const bool unpack_data, const bool unpack_pivots, Tensor P, Tensor L, Tensor U)
Definition lu_unpack.h:10
Tensor::Shape LU_data_shape_
Definition lu_unpack.h:36
int device_index_
Definition lu_unpack.h:70
DataType U_type_
Definition lu_unpack.h:64
Tensor::Shape U_shape_
Definition lu_unpack.h:60
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8