ruprim_host/sort/
dispatch.rs1use ruda_core::tensor::{DType, host::HostTensor};
2
3pub fn float_argtopk(
4 tensor: HostTensor,
5 dim: usize,
6 k: usize,
7 out_dtype: ruda_core::tensor::IntDType,
8) -> HostTensor {
9 let indices = super::argtopk(tensor, dim, k);
10 if indices.dtype() != DType::from(out_dtype) {
11 crate::cast::int_cast(indices, out_dtype)
12 } else {
13 indices
14 }
15}
16
17pub fn int_argtopk(tensor: HostTensor, dim: usize, k: usize) -> HostTensor {
18 let dtype = tensor.dtype();
19 let indices = super::argtopk(tensor, dim, k);
20 if indices.dtype() != dtype {
21 crate::cast::int_cast(indices, dtype.into())
22 } else {
23 indices
24 }
25}
26
27pub fn float_sort_with_indices(
28 tensor: HostTensor,
29 dim: usize,
30 descending: bool,
31 indices_dtype: ruda_core::tensor::IntDType,
32) -> (HostTensor, HostTensor) {
33 let (values, indices) = crate::sort::sort_with_indices(tensor, dim, descending);
34 let indices = if indices.dtype() != DType::from(indices_dtype) {
35 crate::cast::int_cast(indices, indices_dtype)
36 } else {
37 indices
38 };
39 (values, indices)
40}
41
42pub fn float_argsort(
43 tensor: HostTensor,
44 dim: usize,
45 descending: bool,
46 out_dtype: ruda_core::tensor::IntDType,
47) -> HostTensor {
48 let indices = crate::sort::argsort(tensor, dim, descending);
49 if indices.dtype() != DType::from(out_dtype) {
50 crate::cast::int_cast(indices, out_dtype)
51 } else {
52 indices
53 }
54}
55