Skip to main content

ruprim_host/sort/
dispatch.rs

1use 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