InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
nanmean.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_NANMEAN_H_
2#define INFINI_OPS_BASE_NANMEAN_H_
3
4#include <optional>
5#include <vector>
6
7#include "operator.h"
8
9namespace infini::ops {
10
11class Nanmean : public Operator<Nanmean> {
12 public:
13 Nanmean(const Tensor input, const std::optional<std::vector<int64_t>> dim,
14 const bool keepdim, const std::optional<DataType> dtype, 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 keepdim_{keepdim},
23 dtype_{dtype},
24 device_index_{out.device().index()} {}
25
26 virtual void operator()(const Tensor input,
27 const std::optional<std::vector<int64_t>> dim,
28 const bool keepdim,
29 const std::optional<DataType> dtype,
30 Tensor out) const = 0;
31
32 protected:
33 Tensor::Shape input_shape_;
34
35 Tensor::Strides input_strides_;
36
37 DataType input_type_;
38
39 Tensor::Shape out_shape_;
40
41 Tensor::Strides out_strides_;
42
43 DataType out_type_;
44
45 std::optional<std::vector<int64_t>> dim_{};
46
47 bool keepdim_{};
48
49 std::optional<DataType> dtype_{};
50
52};
53
54} // namespace infini::ops
55
56#endif
Definition nanmean.h:11
Tensor::Strides out_strides_
Definition nanmean.h:41
virtual void operator()(const Tensor input, const std::optional< std::vector< int64_t > > dim, const bool keepdim, const std::optional< DataType > dtype, Tensor out) const =0
std::optional< std::vector< int64_t > > dim_
Definition nanmean.h:45
Tensor::Shape out_shape_
Definition nanmean.h:39
Nanmean(const Tensor input, const std::optional< std::vector< int64_t > > dim, const bool keepdim, const std::optional< DataType > dtype, Tensor out)
Definition nanmean.h:13
DataType out_type_
Definition nanmean.h:43
Tensor::Shape input_shape_
Definition nanmean.h:33
DataType input_type_
Definition nanmean.h:37
std::optional< DataType > dtype_
Definition nanmean.h:49
bool keepdim_
Definition nanmean.h:47
int device_index_
Definition nanmean.h:51
Tensor::Strides input_strides_
Definition nanmean.h:35
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8