1pub use ruda_tensor::api::*;
2
3#[cfg(test)]
4mod tests {
5 use super::*;
6 use crate::{TestAutodiffBackend as B, TestBackend};
7
8 #[test]
9 fn full_axis_topk_preserves_values_indices_and_all_input_gradients() {
10 let device = Default::default();
11 let input =
12 Tensor::<B, 2>::from_floats([[1., 3., 2.], [-2., 0., -1.]], &device).require_grad();
13 let expected = TensorData::from([[3., 2., 1.], [0., -1., -2.]]);
14 let indices = TensorData::from([[1i32, 2, 0], [1, 2, 0]]);
15 input
16 .clone()
17 .argtopk(3, 1)
18 .to_data()
19 .assert_eq(&indices, false);
20 let output = input.clone().topk(3, 1);
21 output.to_data().assert_eq(&expected, false);
22 let gradients = output.sum().backward();
23 input
24 .grad(&gradients)
25 .unwrap()
26 .to_data()
27 .assert_eq(&TensorData::from([[1., 1., 1.], [1., 1., 1.]]), false);
28 let (values, actual_indices) = input.clone().topk_with_indices(3, 1);
29 values.to_data().assert_eq(&expected, false);
30 actual_indices.to_data().assert_eq(&indices, false);
31 assert_eq!(input.clone().topk(0, 1).dims(), [2, 0]);
32 assert_eq!(input.argtopk(0, 1).dims(), [2, 0]);
33
34 let input = Tensor::<TestBackend, 2, Int>::from_ints([[1, 3, 2], [-2, 0, -1]], &device);
35 input
36 .clone()
37 .topk(3, 1)
38 .to_data()
39 .assert_eq(&TensorData::from([[3i32, 2, 1], [0, -1, -2]]), false);
40 input
41 .clone()
42 .argtopk(3, 1)
43 .to_data()
44 .assert_eq(&indices, false);
45 assert_eq!(input.topk(0, 1).dims(), [2, 0]);
46 let empty = Tensor::<TestBackend, 2>::zeros([2, 0], &device);
47 assert_eq!(empty.clone().topk(0, 1).dims(), [2, 0]);
48 assert_eq!(empty.argtopk(0, 1).dims(), [2, 0]);
49 }
50}