InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
histogram.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_HISTOGRAM_H_
2#define INFINI_OPS_BASE_HISTOGRAM_H_
3
4#include <optional>
5#include <vector>
6
7#include "operator.h"
8
9namespace infini::ops {
10
11class Histogram : public Operator<Histogram> {
12 public:
13 Histogram(const Tensor input, const Tensor bins,
14 const std::optional<Tensor> weight, const bool density, Tensor hist,
15 Tensor bin_edges)
16 : input_shape_{input.shape()},
17 input_strides_{input.strides()},
18 input_type_{input.dtype()},
19 bins_shape_{bins.shape()},
20 bins_strides_{bins.strides()},
21 bins_type_{bins.dtype()},
22 hist_shape_{hist.shape()},
23 hist_strides_{hist.strides()},
24 hist_type_{hist.dtype()},
25 bin_edges_shape_{bin_edges.shape()},
26 bin_edges_strides_{bin_edges.strides()},
27 bin_edges_type_{bin_edges.dtype()},
28 has_weight_{weight.has_value()},
29 weight_shape_{weight ? Tensor::Shape{weight->shape()}
30 : Tensor::Shape{}},
31 weight_strides_{weight ? Tensor::Strides{weight->strides()}
32 : Tensor::Strides{}},
33 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
34 density_{density},
35 device_index_{hist.device().index()} {}
36
37 Histogram(const Tensor input, const int64_t bins,
38 const std::optional<Tensor> weight,
39 const std::optional<std::vector<double>> range, const bool density,
40 Tensor hist, Tensor bin_edges)
41 : input_shape_{input.shape()},
42 input_strides_{input.strides()},
43 input_type_{input.dtype()},
44 hist_shape_{hist.shape()},
45 hist_strides_{hist.strides()},
46 hist_type_{hist.dtype()},
47 bin_edges_shape_{bin_edges.shape()},
48 bin_edges_strides_{bin_edges.strides()},
49 bin_edges_type_{bin_edges.dtype()},
50 has_weight_{weight.has_value()},
51 weight_shape_{weight ? Tensor::Shape{weight->shape()}
52 : Tensor::Shape{}},
53 weight_strides_{weight ? Tensor::Strides{weight->strides()}
54 : Tensor::Strides{}},
55 weight_type_{weight ? weight->dtype() : DataType::kFloat32},
56 density_{density},
57 bins_{bins},
58 range_{range},
59 device_index_{hist.device().index()} {}
60
61 virtual void operator()(const Tensor input, const Tensor bins,
62 const std::optional<Tensor> weight,
63 const bool density, Tensor hist,
64 Tensor bin_edges) const = 0;
65
66 virtual void operator()(const Tensor input, const int64_t bins,
67 const std::optional<Tensor> weight,
68 const std::optional<std::vector<double>> range,
69 const bool density, Tensor hist,
70 Tensor bin_edges) const = 0;
71
72 protected:
73 Tensor::Shape input_shape_;
74
75 Tensor::Strides input_strides_;
76
77 DataType input_type_;
78
79 Tensor::Shape bins_shape_;
80
81 Tensor::Strides bins_strides_;
82
83 DataType bins_type_;
84
85 Tensor::Shape hist_shape_;
86
87 Tensor::Strides hist_strides_;
88
89 DataType hist_type_;
90
91 Tensor::Shape bin_edges_shape_;
92
93 Tensor::Strides bin_edges_strides_;
94
96
97 bool has_weight_{false};
98
99 Tensor::Shape weight_shape_;
100
101 Tensor::Strides weight_strides_;
102
103 DataType weight_type_{DataType::kFloat32};
104
105 bool density_{};
106
107 int64_t bins_{};
108
109 std::optional<std::vector<double>> range_{};
110
112};
113
114} // namespace infini::ops
115
116#endif
Definition histogram.h:11
Tensor::Strides bins_strides_
Definition histogram.h:81
bool has_weight_
Definition histogram.h:97
DataType weight_type_
Definition histogram.h:103
bool density_
Definition histogram.h:105
Tensor::Strides weight_strides_
Definition histogram.h:101
DataType input_type_
Definition histogram.h:77
Tensor::Strides input_strides_
Definition histogram.h:75
virtual void operator()(const Tensor input, const Tensor bins, const std::optional< Tensor > weight, const bool density, Tensor hist, Tensor bin_edges) const =0
int64_t bins_
Definition histogram.h:107
virtual void operator()(const Tensor input, const int64_t bins, const std::optional< Tensor > weight, const std::optional< std::vector< double > > range, const bool density, Tensor hist, Tensor bin_edges) const =0
Tensor::Strides hist_strides_
Definition histogram.h:87
Tensor::Shape weight_shape_
Definition histogram.h:99
int device_index_
Definition histogram.h:111
DataType hist_type_
Definition histogram.h:89
Tensor::Strides bin_edges_strides_
Definition histogram.h:93
Histogram(const Tensor input, const int64_t bins, const std::optional< Tensor > weight, const std::optional< std::vector< double > > range, const bool density, Tensor hist, Tensor bin_edges)
Definition histogram.h:37
std::optional< std::vector< double > > range_
Definition histogram.h:109
Histogram(const Tensor input, const Tensor bins, const std::optional< Tensor > weight, const bool density, Tensor hist, Tensor bin_edges)
Definition histogram.h:13
DataType bins_type_
Definition histogram.h:83
Tensor::Shape hist_shape_
Definition histogram.h:85
Tensor::Shape input_shape_
Definition histogram.h:73
Tensor::Shape bins_shape_
Definition histogram.h:79
Tensor::Shape bin_edges_shape_
Definition histogram.h:91
DataType bin_edges_type_
Definition histogram.h:95
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8