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