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