Skip to main content

ruda_model/
tensor.rs

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}