InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
argsort.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_ARGSORT_H_
2#define INFINI_OPS_BASE_ARGSORT_H_
3
4#include "operator.h"
5
6namespace infini::ops {
7
8class Argsort : public Operator<Argsort> {
9 public:
10 Argsort(const Tensor input, const int64_t dim, const bool descending,
11 const bool stable, Tensor out)
12 : input_shape_{input.shape()},
13 input_strides_{input.strides()},
14 input_type_{input.dtype()},
15 out_shape_{out.shape()},
16 out_strides_{out.strides()},
17 out_type_{out.dtype()},
18 stable_{stable},
19 dim_{dim},
20 descending_{descending},
21 device_index_{out.device().index()} {}
22
25 [[deprecated("Use the PyTorch-compatible parameter order instead.")]]
26 Argsort(const Tensor input, const bool stable, const int64_t dim,
27 const bool descending, Tensor out)
28 : Argsort{input, dim, descending, stable, out} {}
29
30 void operator()(const Tensor input, const int64_t dim, const bool descending,
31 const bool stable, Tensor out) const {
32 (*this)(input, stable, dim, descending, out);
33 }
34
37 [[deprecated("Use the PyTorch-compatible parameter order instead.")]]
38 virtual void operator()(const Tensor input, const bool stable,
39 const int64_t dim, const bool descending,
40 Tensor out) const = 0;
41
42 protected:
43 Tensor::Shape input_shape_;
44
45 Tensor::Strides input_strides_;
46
47 DataType input_type_;
48
49 Tensor::Shape out_shape_;
50
51 Tensor::Strides out_strides_;
52
53 DataType out_type_;
54
55 bool stable_{};
56
57 int64_t dim_{};
58
60
62};
63
64} // namespace infini::ops
65
66#endif
Definition argsort.h:8
void operator()(const Tensor input, const int64_t dim, const bool descending, const bool stable, Tensor out) const
Definition argsort.h:30
int64_t dim_
Definition argsort.h:57
DataType input_type_
Definition argsort.h:47
Tensor::Shape input_shape_
Definition argsort.h:43
bool descending_
Definition argsort.h:59
Argsort(const Tensor input, const int64_t dim, const bool descending, const bool stable, Tensor out)
Definition argsort.h:10
Tensor::Strides input_strides_
Definition argsort.h:45
int device_index_
Definition argsort.h:61
bool stable_
Definition argsort.h:55
virtual void operator()(const Tensor input, const bool stable, const int64_t dim, const bool descending, Tensor out) const =0
DataType out_type_
Definition argsort.h:53
Argsort(const Tensor input, const bool stable, const int64_t dim, const bool descending, Tensor out)
Definition argsort.h:26
Tensor::Strides out_strides_
Definition argsort.h:51
Tensor::Shape out_shape_
Definition argsort.h:49
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8