InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
nanquantile.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_NANQUANTILE_H_
2#define INFINI_OPS_BASE_NANQUANTILE_H_
3
4#include <optional>
5#include <string>
6
7#include "operator.h"
8
9namespace infini::ops {
10
11class Nanquantile : public Operator<Nanquantile> {
12 public:
13 Nanquantile(const Tensor input, const Tensor q,
14 const std::optional<int64_t> dim, const bool keepdim,
15 const std::string interpolation, Tensor out)
16 : input_shape_{input.shape()},
17 input_strides_{input.strides()},
18 input_type_{input.dtype()},
19 q_shape_{q.shape()},
20 q_strides_{q.strides()},
21 q_type_{q.dtype()},
22 out_shape_{out.shape()},
23 out_strides_{out.strides()},
24 out_type_{out.dtype()},
25 dim_{dim},
26 keepdim_{keepdim},
27 interpolation_{interpolation},
28 device_index_{out.device().index()} {}
29
30 Nanquantile(const Tensor input, const double q,
31 const std::optional<int64_t> dim, const bool keepdim,
32 const std::string interpolation, Tensor out)
33 : input_shape_{input.shape()},
34 input_strides_{input.strides()},
35 input_type_{input.dtype()},
36 out_shape_{out.shape()},
37 out_strides_{out.strides()},
38 out_type_{out.dtype()},
39 dim_{dim},
40 keepdim_{keepdim},
41 interpolation_{interpolation},
42 q_{q},
43 device_index_{out.device().index()} {}
44
45 virtual void operator()(const Tensor input, const Tensor q,
46 const std::optional<int64_t> dim, const bool keepdim,
47 const std::string interpolation,
48 Tensor out) const = 0;
49
50 virtual void operator()(const Tensor input, const double q,
51 const std::optional<int64_t> dim, const bool keepdim,
52 const std::string interpolation,
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 q_shape_;
63
64 Tensor::Strides q_strides_;
65
66 DataType q_type_;
67
68 Tensor::Shape out_shape_;
69
70 Tensor::Strides out_strides_;
71
72 DataType out_type_;
73
74 std::optional<int64_t> dim_{};
75
76 bool keepdim_{};
77
78 std::string interpolation_{};
79
80 double q_{};
81
83};
84
85} // namespace infini::ops
86
87#endif
Definition nanquantile.h:11
virtual void operator()(const Tensor input, const Tensor q, const std::optional< int64_t > dim, const bool keepdim, const std::string interpolation, Tensor out) const =0
int device_index_
Definition nanquantile.h:82
Nanquantile(const Tensor input, const double q, const std::optional< int64_t > dim, const bool keepdim, const std::string interpolation, Tensor out)
Definition nanquantile.h:30
Tensor::Strides input_strides_
Definition nanquantile.h:58
double q_
Definition nanquantile.h:80
bool keepdim_
Definition nanquantile.h:76
std::optional< int64_t > dim_
Definition nanquantile.h:74
virtual void operator()(const Tensor input, const double q, const std::optional< int64_t > dim, const bool keepdim, const std::string interpolation, Tensor out) const =0
Tensor::Strides out_strides_
Definition nanquantile.h:70
std::string interpolation_
Definition nanquantile.h:78
Tensor::Strides q_strides_
Definition nanquantile.h:64
Nanquantile(const Tensor input, const Tensor q, const std::optional< int64_t > dim, const bool keepdim, const std::string interpolation, Tensor out)
Definition nanquantile.h:13
DataType input_type_
Definition nanquantile.h:60
Tensor::Shape input_shape_
Definition nanquantile.h:56
DataType out_type_
Definition nanquantile.h:72
Tensor::Shape out_shape_
Definition nanquantile.h:68
Tensor::Shape q_shape_
Definition nanquantile.h:62
DataType q_type_
Definition nanquantile.h:66
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8