InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
var.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_VAR_H_
2#define INFINI_OPS_BASE_VAR_H_
3
4#include <optional>
5#include <vector>
6
7#include "operator.h"
8
9namespace infini::ops {
10
11class Var : public Operator<Var> {
12 public:
13 Var(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 Var(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 var.h:11
std::optional< double > correction_
Definition var.h:68
Var(const Tensor input, const std::optional< std::vector< int64_t > > dim, const std::optional< double > correction, const bool keepdim, Tensor out)
Definition var.h:26
bool keepdim_
Definition var.h:66
Var(const Tensor input, const std::optional< std::vector< int64_t > > dim, const bool unbiased, const bool keepdim, Tensor out)
Definition var.h:13
DataType input_type_
Definition var.h:54
DataType out_type_
Definition var.h:60
std::optional< std::vector< int64_t > > dim_
Definition var.h:62
Tensor::Strides input_strides_
Definition var.h:52
virtual void operator()(const Tensor input, const std::optional< std::vector< int64_t > > dim, const bool unbiased, const bool keepdim, Tensor out) const =0
Tensor::Strides out_strides_
Definition var.h:58
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 var.h:64
Tensor::Shape out_shape_
Definition var.h:56
int device_index_
Definition var.h:70
Tensor::Shape input_shape_
Definition var.h:50
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8