InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
sort.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_SORT_H_
2#define INFINI_OPS_BASE_SORT_H_
3
4#include <optional>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class Sort : public Operator<Sort> {
11 public:
12 Sort(const Tensor input, const int64_t dim, const bool descending,
13 Tensor values, Tensor indices)
14 : input_shape_{input.shape()},
15 input_strides_{input.strides()},
16 input_type_{input.dtype()},
17 values_shape_{values.shape()},
18 values_strides_{values.strides()},
19 values_type_{values.dtype()},
20 indices_shape_{indices.shape()},
21 indices_strides_{indices.strides()},
22 indices_type_{indices.dtype()},
23 dim_{dim},
24 descending_{descending},
25 device_index_{values.device().index()} {}
26
27 Sort(const Tensor input, const std::optional<bool> stable, const int64_t dim,
28 const bool descending, Tensor values, Tensor indices)
29 : input_shape_{input.shape()},
30 input_strides_{input.strides()},
31 input_type_{input.dtype()},
32 values_shape_{values.shape()},
33 values_strides_{values.strides()},
34 values_type_{values.dtype()},
35 indices_shape_{indices.shape()},
36 indices_strides_{indices.strides()},
37 indices_type_{indices.dtype()},
38 dim_{dim},
39 descending_{descending},
40 stable_{stable},
41 device_index_{values.device().index()} {}
42
43 virtual void operator()(const Tensor input, const int64_t dim,
44 const bool descending, Tensor values,
45 Tensor indices) const = 0;
46
47 virtual void operator()(const Tensor input, const std::optional<bool> stable,
48 const int64_t dim, const bool descending,
49 Tensor values, Tensor indices) const = 0;
50
51 protected:
52 Tensor::Shape input_shape_;
53
54 Tensor::Strides input_strides_;
55
56 DataType input_type_;
57
58 Tensor::Shape values_shape_;
59
60 Tensor::Strides values_strides_;
61
62 DataType values_type_;
63
64 Tensor::Shape indices_shape_;
65
66 Tensor::Strides indices_strides_;
67
68 DataType indices_type_;
69
70 int64_t dim_{};
71
73
74 std::optional<bool> stable_{};
75
77};
78
79} // namespace infini::ops
80
81#endif
Definition generated/include/operator.h:282
Definition sort.h:10
virtual void operator()(const Tensor input, const int64_t dim, const bool descending, Tensor values, Tensor indices) const =0
int device_index_
Definition sort.h:76
Sort(const Tensor input, const int64_t dim, const bool descending, Tensor values, Tensor indices)
Definition sort.h:12
DataType indices_type_
Definition sort.h:68
Tensor::Strides input_strides_
Definition sort.h:54
virtual void operator()(const Tensor input, const std::optional< bool > stable, const int64_t dim, const bool descending, Tensor values, Tensor indices) const =0
Tensor::Strides indices_strides_
Definition sort.h:66
DataType values_type_
Definition sort.h:62
Tensor::Strides values_strides_
Definition sort.h:60
Tensor::Shape values_shape_
Definition sort.h:58
int64_t dim_
Definition sort.h:70
Tensor::Shape indices_shape_
Definition sort.h:64
std::optional< bool > stable_
Definition sort.h:74
bool descending_
Definition sort.h:72
Sort(const Tensor input, const std::optional< bool > stable, const int64_t dim, const bool descending, Tensor values, Tensor indices)
Definition sort.h:27
Tensor::Shape input_shape_
Definition sort.h:52
DataType input_type_
Definition sort.h:56
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8