InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
linalg_norm.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_LINALG_NORM_H_
2#define INFINI_OPS_BASE_LINALG_NORM_H_
3
4#include <optional>
5#include <string>
6#include <vector>
7
8#include "operator.h"
9
10namespace infini::ops::linalg {
11
12class Norm : public Operator<Norm> {
13 public:
14 Norm(const Tensor input, const std::optional<double> ord,
15 const std::optional<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 Norm(const Tensor input, const std::string ord,
30 const std::optional<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 std::optional<double> ord,
44 const std::optional<std::vector<int64_t>> dim,
45 const bool keepdim,
46 const std::optional<DataType> dtype,
47 Tensor out) const = 0;
48
49 virtual void operator()(const Tensor input, const std::string ord,
50 const std::optional<std::vector<int64_t>> dim,
51 const bool keepdim,
52 const std::optional<DataType> dtype,
53 Tensor out) const = 0;
54
55 protected:
56 Tensor::Shape input_shape_;
57
58 Tensor::Strides input_strides_;
59
60 DataType input_type_;
61
62 Tensor::Shape out_shape_;
63
64 Tensor::Strides out_strides_;
65
66 DataType out_type_;
67
68 std::optional<double> ord_{};
69
70 std::optional<std::vector<int64_t>> dim_{};
71
72 bool keepdim_{};
73
74 std::optional<DataType> dtype_{};
75
77};
78
79} // namespace infini::ops::linalg
80
81#endif
Definition generated/include/operator.h:282
Definition linalg_norm.h:12
bool keepdim_
Definition linalg_norm.h:72
int device_index_
Definition linalg_norm.h:76
DataType out_type_
Definition linalg_norm.h:66
Tensor::Strides input_strides_
Definition linalg_norm.h:58
std::optional< double > ord_
Definition linalg_norm.h:68
Tensor::Shape input_shape_
Definition linalg_norm.h:56
virtual void operator()(const Tensor input, const std::string ord, const std::optional< std::vector< int64_t > > dim, const bool keepdim, const std::optional< DataType > dtype, Tensor out) const =0
Norm(const Tensor input, const std::optional< double > ord, const std::optional< std::vector< int64_t > > dim, const bool keepdim, const std::optional< DataType > dtype, Tensor out)
Definition linalg_norm.h:14
std::optional< std::vector< int64_t > > dim_
Definition linalg_norm.h:70
std::optional< DataType > dtype_
Definition linalg_norm.h:74
Norm(const Tensor input, const std::string ord, const std::optional< std::vector< int64_t > > dim, const bool keepdim, const std::optional< DataType > dtype, Tensor out)
Definition linalg_norm.h:29
Tensor::Shape out_shape_
Definition linalg_norm.h:62
Tensor::Strides out_strides_
Definition linalg_norm.h:64
DataType input_type_
Definition linalg_norm.h:60
virtual void operator()(const Tensor input, const std::optional< double > ord, const std::optional< std::vector< int64_t > > dim, const bool keepdim, const std::optional< DataType > dtype, Tensor out) const =0
Definition linalg_cholesky.h:6
infini::rt::TensorView Tensor
Definition tensor.h:8