InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
searchsorted.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SEARCHSORTED_H_
2#define INFINI_OPS_BASE_SEARCHSORTED_H_
3
4#include <optional>
5#include <string>
6
7#include "operator.h"
8
9namespace infini::ops {
10
11class Searchsorted : public Operator<Searchsorted> {
12 public:
13 Searchsorted(const Tensor sorted_sequence, const Tensor input,
14 const std::optional<Tensor> sorter, const bool out_int32,
15 const bool right, const std::optional<std::string> side,
16 Tensor out)
17 : sorted_sequence_shape_{sorted_sequence.shape()},
18 sorted_sequence_strides_{sorted_sequence.strides()},
19 sorted_sequence_type_{sorted_sequence.dtype()},
20 input_shape_{input.shape()},
21 input_strides_{input.strides()},
22 input_type_{input.dtype()},
23 out_shape_{out.shape()},
24 out_strides_{out.strides()},
25 out_type_{out.dtype()},
26 has_sorter_{sorter.has_value()},
27 sorter_shape_{sorter ? Tensor::Shape{sorter->shape()}
28 : Tensor::Shape{}},
29 sorter_strides_{sorter ? Tensor::Strides{sorter->strides()}
30 : Tensor::Strides{}},
31 sorter_type_{sorter ? sorter->dtype() : DataType::kFloat32},
32 out_int32_{out_int32},
33 right_{right},
34 side_{side},
35 device_index_{out.device().index()} {}
36
39 [[deprecated("Place the optional `sorter` Tensor before attributes.")]]
40 Searchsorted(const Tensor sorted_sequence, const Tensor input,
41 const bool out_int32, const bool right,
42 const std::optional<std::string> side,
43 const std::optional<Tensor> sorter, Tensor out)
44 : Searchsorted{sorted_sequence, input, sorter, out_int32,
45 right, side, out} {}
46
47 Searchsorted(const Tensor sorted_sequence, const double input,
48 const std::optional<Tensor> sorter, const bool out_int32,
49 const bool right, const std::optional<std::string> side,
50 Tensor out)
51 : sorted_sequence_shape_{sorted_sequence.shape()},
52 sorted_sequence_strides_{sorted_sequence.strides()},
53 sorted_sequence_type_{sorted_sequence.dtype()},
54 out_shape_{out.shape()},
55 out_strides_{out.strides()},
56 out_type_{out.dtype()},
57 has_sorter_{sorter.has_value()},
58 sorter_shape_{sorter ? Tensor::Shape{sorter->shape()}
59 : Tensor::Shape{}},
60 sorter_strides_{sorter ? Tensor::Strides{sorter->strides()}
61 : Tensor::Strides{}},
62 sorter_type_{sorter ? sorter->dtype() : DataType::kFloat32},
63 out_int32_{out_int32},
64 right_{right},
65 side_{side},
66 input_{input},
67 device_index_{out.device().index()} {}
68
71 [[deprecated("Place the optional `sorter` Tensor before attributes.")]]
72 Searchsorted(const Tensor sorted_sequence, const double input,
73 const bool out_int32, const bool right,
74 const std::optional<std::string> side,
75 const std::optional<Tensor> sorter, Tensor out)
76 : Searchsorted{sorted_sequence, input, sorter, out_int32,
77 right, side, out} {}
78
79 void operator()(const Tensor sorted_sequence, const Tensor input,
80 const std::optional<Tensor> sorter, const bool out_int32,
81 const bool right, const std::optional<std::string> side,
82 Tensor out) const {
83 (*this)(sorted_sequence, input, out_int32, right, side, sorter, out);
84 }
85
86 void operator()(const Tensor sorted_sequence, const double input,
87 const std::optional<Tensor> sorter, const bool out_int32,
88 const bool right, const std::optional<std::string> side,
89 Tensor out) const {
90 (*this)(sorted_sequence, input, out_int32, right, side, sorter, out);
91 }
92
95 [[deprecated("Place the optional `sorter` Tensor before attributes.")]]
96 virtual void operator()(const Tensor sorted_sequence, const Tensor input,
97 const bool out_int32, const bool right,
98 const std::optional<std::string> side,
99 const std::optional<Tensor> sorter,
100 Tensor out) const = 0;
101
104 [[deprecated("Place the optional `sorter` Tensor before attributes.")]]
105 virtual void operator()(const Tensor sorted_sequence, const double input,
106 const bool out_int32, const bool right,
107 const std::optional<std::string> side,
108 const std::optional<Tensor> sorter,
109 Tensor out) const = 0;
110
111 protected:
113
115
117
118 Tensor::Shape input_shape_;
119
120 Tensor::Strides input_strides_;
121
122 DataType input_type_;
123
124 Tensor::Shape out_shape_;
125
126 Tensor::Strides out_strides_;
127
128 DataType out_type_;
129
130 bool has_sorter_{false};
131
132 Tensor::Shape sorter_shape_;
133
134 Tensor::Strides sorter_strides_;
135
136 DataType sorter_type_{DataType::kFloat32};
137
139
140 bool right_{};
141
142 std::optional<std::string> side_{};
143
144 double input_{};
145
147};
148
149} // namespace infini::ops
150
151#endif
Definition generated/include/operator.h:282
Definition searchsorted.h:11
Tensor::Shape sorted_sequence_shape_
Definition searchsorted.h:112
Searchsorted(const Tensor sorted_sequence, const double input, const bool out_int32, const bool right, const std::optional< std::string > side, const std::optional< Tensor > sorter, Tensor out)
Definition searchsorted.h:72
DataType sorted_sequence_type_
Definition searchsorted.h:116
Tensor::Shape sorter_shape_
Definition searchsorted.h:132
void operator()(const Tensor sorted_sequence, const Tensor input, const std::optional< Tensor > sorter, const bool out_int32, const bool right, const std::optional< std::string > side, Tensor out) const
Definition searchsorted.h:79
DataType out_type_
Definition searchsorted.h:128
DataType sorter_type_
Definition searchsorted.h:136
Tensor::Strides out_strides_
Definition searchsorted.h:126
int device_index_
Definition searchsorted.h:146
Searchsorted(const Tensor sorted_sequence, const Tensor input, const bool out_int32, const bool right, const std::optional< std::string > side, const std::optional< Tensor > sorter, Tensor out)
Definition searchsorted.h:40
virtual void operator()(const Tensor sorted_sequence, const double input, const bool out_int32, const bool right, const std::optional< std::string > side, const std::optional< Tensor > sorter, Tensor out) const =0
std::optional< std::string > side_
Definition searchsorted.h:142
Tensor::Shape input_shape_
Definition searchsorted.h:118
Tensor::Strides sorted_sequence_strides_
Definition searchsorted.h:114
Tensor::Strides sorter_strides_
Definition searchsorted.h:134
Searchsorted(const Tensor sorted_sequence, const Tensor input, const std::optional< Tensor > sorter, const bool out_int32, const bool right, const std::optional< std::string > side, Tensor out)
Definition searchsorted.h:13
Tensor::Strides input_strides_
Definition searchsorted.h:120
void operator()(const Tensor sorted_sequence, const double input, const std::optional< Tensor > sorter, const bool out_int32, const bool right, const std::optional< std::string > side, Tensor out) const
Definition searchsorted.h:86
bool out_int32_
Definition searchsorted.h:138
virtual void operator()(const Tensor sorted_sequence, const Tensor input, const bool out_int32, const bool right, const std::optional< std::string > side, const std::optional< Tensor > sorter, Tensor out) const =0
bool has_sorter_
Definition searchsorted.h:130
DataType input_type_
Definition searchsorted.h:122
double input_
Definition searchsorted.h:144
Searchsorted(const Tensor sorted_sequence, const double input, const std::optional< Tensor > sorter, const bool out_int32, const bool right, const std::optional< std::string > side, Tensor out)
Definition searchsorted.h:47
Tensor::Shape out_shape_
Definition searchsorted.h:124
bool right_
Definition searchsorted.h:140
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8