InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
linalg_matrix_norm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_LINALG_MATRIX_NORM_H_
2#define INFINI_OPS_BASE_LINALG_MATRIX_NORM_H_
3
4#include <optional>
5#include <string>
6#include <vector>
7
8#include "operator.h"
9
10namespace infini::ops::linalg {
11
12class MatrixNorm : public Operator<MatrixNorm> {
13 public:
14 MatrixNorm(const Tensor input, const double ord,
15 const std::vector<int64_t> dim, const bool keepdim,
16 const std::optional<DataType> dtype, Tensor out)
17 : input_shape_{input.shape()},
18 input_strides_{input.strides()},
19 input_type_{input.dtype()},
20 out_shape_{out.shape()},
21 out_strides_{out.strides()},
22 out_type_{out.dtype()},
23 ord_{ord},
24 dim_{dim},
25 keepdim_{keepdim},
26 dtype_{dtype},
27 device_index_{out.device().index()} {}
28
29 MatrixNorm(const Tensor input, const std::string ord,
30 const std::vector<int64_t> dim, const bool keepdim,
31 const std::optional<DataType> dtype, Tensor out)
32 : input_shape_{input.shape()},
33 input_strides_{input.strides()},
34 input_type_{input.dtype()},
35 out_shape_{out.shape()},
36 out_strides_{out.strides()},
37 out_type_{out.dtype()},
38 dim_{dim},
39 keepdim_{keepdim},
40 dtype_{dtype},
41 device_index_{out.device().index()} {}
42
43 virtual void operator()(const Tensor input, const double ord,
44 const std::vector<int64_t> dim, const bool keepdim,
45 const std::optional<DataType> dtype,
46 Tensor out) const = 0;
47
48 virtual void operator()(const Tensor input, const std::string ord,
49 const std::vector<int64_t> dim, const bool keepdim,
50 const std::optional<DataType> dtype,
51 Tensor out) const = 0;
52
53 protected:
54 Tensor::Shape input_shape_;
55
56 Tensor::Strides input_strides_;
57
58 DataType input_type_;
59
60 Tensor::Shape out_shape_;
61
62 Tensor::Strides out_strides_;
63
64 DataType out_type_;
65
66 double ord_{};
67
68 std::vector<int64_t> dim_{};
69
70 bool keepdim_{};
71
72 std::optional<DataType> dtype_{};
73
75};
76
77} // namespace infini::ops::linalg
78
79#endif
Definition generated/include/operator.h:282
Definition linalg_matrix_norm.h:12
std::vector< int64_t > dim_
Definition linalg_matrix_norm.h:68
Tensor::Strides out_strides_
Definition linalg_matrix_norm.h:62
std::optional< DataType > dtype_
Definition linalg_matrix_norm.h:72
Tensor::Strides input_strides_
Definition linalg_matrix_norm.h:56
double ord_
Definition linalg_matrix_norm.h:66
MatrixNorm(const Tensor input, const double ord, const std::vector< int64_t > dim, const bool keepdim, const std::optional< DataType > dtype, Tensor out)
Definition linalg_matrix_norm.h:14
virtual void operator()(const Tensor input, const std::string ord, const std::vector< int64_t > dim, const bool keepdim, const std::optional< DataType > dtype, Tensor out) const =0
int device_index_
Definition linalg_matrix_norm.h:74
virtual void operator()(const Tensor input, const double ord, const std::vector< int64_t > dim, const bool keepdim, const std::optional< DataType > dtype, Tensor out) const =0
Tensor::Shape input_shape_
Definition linalg_matrix_norm.h:54
Tensor::Shape out_shape_
Definition linalg_matrix_norm.h:60
bool keepdim_
Definition linalg_matrix_norm.h:70
MatrixNorm(const Tensor input, const std::string ord, const std::vector< int64_t > dim, const bool keepdim, const std::optional< DataType > dtype, Tensor out)
Definition linalg_matrix_norm.h:29
DataType out_type_
Definition linalg_matrix_norm.h:64
DataType input_type_
Definition linalg_matrix_norm.h:58
Definition linalg_cholesky.h:6
infini::rt::TensorView Tensor
Definition tensor.h:8