Skip to main content

ruda_tensor/api/
orderable.rs

1use crate::{
2    Backend, ElementConversion, Scalar,
3    tensor::{Bool, IndexingUpdateOp, Int, Ordered},
4};
5use ruda_core::tensor::indexing::AsIndex;
6
7use crate::api::check;
8use crate::api::{Tensor, check::TensorCheck};
9
10impl<B, const D: usize, K> Tensor<B, D, K>
11where
12    B: Backend,
13    K: Ordered<B>,
14{
15    /// Sort the elements by value in ascending order along a given dimension.
16    ///
17    /// This sort is unstable (i.e., may reorder equal elements).
18    ///
19    /// # Arguments
20    ///
21    /// * `dim` - The dimension to sort along.
22    ///
23    /// # Returns
24    ///
25    /// A new tensor with the elements sorted in ascending order along the given dimension.
26    ///
27    /// # Example
28    ///
29    /// ```rust
30    /// use ruda_tensor::api::backend::Backend;
31    /// use ruda_tensor::api::{Tensor, Shape};
32    ///
33    /// fn example<B: Backend>() {
34    ///   let device = B::Device::default();
35    ///   let tensor = Tensor::<B, 2>::from_data([[12.0, -2.0, 3.0], [5.0, 3.0, 6.0]], &device);
36    ///   let tensor = tensor.sort(0);
37    ///   println!("{tensor}");
38    ///   // [[5.0, -2.0, 3.0], [12.0, 3.0, 6.0]]
39    ///   let tensor = tensor.sort(1);
40    ///   println!("{tensor}");
41    ///   // [[-2.0, 3.0, 12.0], [3.0, 5.0, 6.0]]
42    /// }
43    /// ```
44    pub fn sort(self, dim: usize) -> Self {
45        check!(TensorCheck::sort_dim::<D>("Sort", dim));
46        Tensor::new(K::sort(self.primitive, dim, /*descending*/ false))
47    }
48
49    /// Sort the elements by value in descending order along a given dimension.
50    ///
51    /// This sort is unstable (i.e., may reorder equal elements).
52    ///
53    /// # Arguments
54    ///
55    /// * `dim` - The dimension to sort along.
56    ///
57    /// # Returns
58    ///
59    /// A new tensor with the elements sorted in descending order along the given dimension.
60    ///
61    /// # Example
62    ///
63    /// ```rust
64    /// use ruda_tensor::api::backend::Backend;
65    /// use ruda_tensor::api::{Tensor, Shape};
66    ///
67    /// fn example<B: Backend>() {
68    ///    let device = B::Device::default();
69    ///    let tensor = Tensor::<B, 2>::from_data([[12.0, -2.0, 3.0], [5.0, 3.0, 6.0]], &device);
70    ///    let tensor = tensor.sort_descending(0);
71    ///    println!("{tensor}");
72    ///    // [[12.0, 3.0, 6.0], [5.0, -2.0, 3.0]]
73    ///    let tensor = tensor.sort_descending(1);
74    ///    println!("{tensor}");
75    ///    // [[12.0, 3.0, -2.0], [6.0, 5.0, 3.0]]
76    /// }
77    /// ```
78    pub fn sort_descending(self, dim: usize) -> Self {
79        check!(TensorCheck::sort_dim::<D>("Sort", dim));
80        Tensor::new(K::sort(self.primitive, dim, /*descending*/ true))
81    }
82
83    /// Sort the elements by value in ascending order along a given dimension.
84    /// Also returns the indices.
85    ///
86    /// This sort is unstable (i.e., may reorder equal elements).
87    ///
88    /// # Arguments
89    ///
90    /// * `dim` - The dimension to sort along.
91    ///
92    /// # Returns
93    ///
94    /// A tuple containing the sorted tensor and the indices tensor.
95    ///
96    /// # Example
97    ///
98    /// ```rust
99    /// use ruda_tensor::api::backend::Backend;
100    /// use ruda_tensor::api::{Tensor, Shape};
101    ///
102    /// fn example<B: Backend>() {
103    ///   let device = B::Device::default();
104    ///   let tensor = Tensor::<B, 2>::from_data([[12.0, -2.0, 3.0], [5.0, 3.0, 6.0]], &device);
105    ///   let (tensor, indices) = tensor.sort_with_indices(0);
106    ///   println!("{tensor}");
107    ///   // [[5.0, -2.0, 3.0], [12.0, 3.0, 6.0]]
108    ///   println!("{}", indices);
109    ///   // [[1, 0, 0], [0, 1, 1]]
110    /// }
111    /// ```
112    pub fn sort_with_indices(self, dim: usize) -> (Self, Tensor<B, D, Int>) {
113        check!(TensorCheck::sort_dim::<D>("Sort_with_indices", dim));
114        let (values, indices) =
115            K::sort_with_indices(self.primitive, dim, /*descending*/ false);
116        (Tensor::new(values), Tensor::new(indices))
117    }
118
119    /// Sort the elements by value in descending order along a given dimension.
120    /// Also returns the indices.
121    ///
122    /// This sort is unstable (i.e., may reorder equal elements).
123    ///
124    /// # Arguments
125    ///
126    /// * `dim` - The dimension to sort along.
127    ///
128    /// # Example
129    ///
130    /// ```rust
131    /// use ruda_tensor::api::backend::Backend;
132    /// use ruda_tensor::api::{Tensor, Shape};
133    ///
134    /// fn example<B: Backend>() {
135    ///    let device = B::Device::default();
136    ///    let tensor = Tensor::<B, 2>::from_data([[12.0, -2.0, 3.0], [5.0, 3.0, 6.0]], &device);
137    ///    let (tensor, indices) = tensor.sort_descending_with_indices(0);
138    ///    println!("{tensor}");
139    ///    // [[12.0, 3.0, 6.0], [5.0, -2.0, 3.0]]
140    ///    println!("{}", indices);
141    ///    // [[0, 1, 1], [1, 0, 0]]
142    /// }
143    /// ```
144    pub fn sort_descending_with_indices(self, dim: usize) -> (Self, Tensor<B, D, Int>) {
145        check!(TensorCheck::sort_dim::<D>("Sort_with_indices", dim));
146        let (values, indices) = K::sort_with_indices(self.primitive, dim, /*descending*/ true);
147        (Tensor::new(values), Tensor::new(indices))
148    }
149
150    /// Returns the indices that sort the elements by value in ascending order along a given dimension.
151    ///
152    /// This sort is unstable (i.e., may reorder equal elements).
153    ///
154    /// # Arguments
155    ///
156    /// * `dim` - The dimension to sort along.
157    ///
158    /// # Example
159    ///
160    /// ```rust
161    /// use ruda_tensor::api::backend::Backend;
162    /// use ruda_tensor::api::{Tensor, Shape};
163    ///
164    /// fn example<B: Backend>() {
165    ///    let device = B::Device::default();
166    ///    let tensor = Tensor::<B, 2>::from_data([[12.0, -2.0, 3.0], [5.0, 3.0, 6.0]], &device);
167    ///    let tensor = tensor.argsort(0);
168    ///    println!("{tensor}");
169    ///    // [[1, 0, 0], [0, 1, 1]]
170    /// }
171    /// ```
172    pub fn argsort(self, dim: usize) -> Tensor<B, D, Int> {
173        check!(TensorCheck::sort_dim::<D>("Argsort", dim));
174        Tensor::new(K::argsort(self.primitive, dim, /*descending*/ false))
175    }
176
177    /// Returns the indices that sort the elements by value in descending order along a given dimension.
178    ///
179    /// This sort is unstable (i.e., may reorder equal elements).
180    ///
181    /// # Arguments
182    ///
183    /// * `dim` - The dimension to sort along.
184    ///
185    /// # Example
186    ///
187    /// ```rust
188    /// use ruda_tensor::api::backend::Backend;
189    /// use ruda_tensor::api::{Tensor, Shape};
190    ///
191    /// fn example<B: Backend>() {
192    ///    let device = B::Device::default();
193    ///    let tensor = Tensor::<B, 2>::from_data([[12.0, -2.0, 3.0], [5.0, 3.0, 6.0]], &device);
194    ///    let tensor = tensor.argsort_descending(0);
195    ///    println!("{tensor}");
196    ///    // [[0, 1, 1], [1, 0, 0]]
197    ///    let tensor = tensor.argsort_descending(1);
198    ///    println!("{tensor}");
199    ///    // [[0, 2, 1], [2, 0, 1]]
200    /// }
201    /// ```
202    pub fn argsort_descending(self, dim: usize) -> Tensor<B, D, Int> {
203        check!(TensorCheck::sort_dim::<D>("Argsort", dim));
204        Tensor::new(K::argsort(self.primitive, dim, /*descending*/ true))
205    }
206
207    /// Returns the `k` largest elements of the given input tensor along a given dimension.
208    ///
209    /// # Arguments
210    ///
211    /// * `k` - The number of elements to return.
212    ///
213    /// # Returns
214    ///
215    /// A new tensor with the `k` largest elements along the given dimension.
216    ///
217    /// # Example
218    ///
219    /// ```rust
220    /// use ruda_tensor::api::backend::Backend;
221    /// use ruda_tensor::api::{Tensor, Shape};
222    ///
223    /// fn example<B: Backend>() {
224    ///   let device = B::Device::default();
225    ///   let tensor = Tensor::<B, 2>::from_data([[12.0, -2.0, 3.0], [5.0, 3.0, 6.0]], &device);
226    ///   let tensor = tensor.topk(2, 0);
227    ///   println!("{tensor}");
228    ///   // [[12.0, 3.0, 6.0], [5.0, -2.0, 3.0]]
229    ///   let tensor = tensor.topk(1, 1);
230    ///   println!("{tensor}");
231    ///   // [[12.0], [6.0]]
232    /// }
233    /// ```
234    pub fn topk(self, k: usize, dim: usize) -> Self {
235        assert!(self.shape()[dim] >= k);
236        Tensor::new(K::topk(self.primitive, dim, k))
237    }
238
239    /// Returns the `k` largest elements of the given input tensor along a given dimension.
240    /// Also returns the indices.
241    ///
242    /// # Arguments
243    ///
244    /// * `k` - The number of elements to return.
245    /// * `dim` - The dimension to sort along.
246    ///
247    /// # Example
248    ///
249    /// ```rust
250    /// use ruda_tensor::api::backend::Backend;
251    /// use ruda_tensor::api::{Tensor, Shape};
252    ///
253    /// fn example<B: Backend>() {
254    ///    let device = B::Device::default();
255    ///    let tensor = Tensor::<B, 2>::from_data([[12.0, -2.0, 3.0], [5.0, 3.0, 6.0]], &device);
256    ///    let (tensor, indices) = tensor.topk_with_indices(2, 0);
257    ///    println!("{tensor}");
258    ///    // [[12.0, 3.0, 6.0], [5.0, -2.0, 3.0]]
259    ///    println!("{}", indices);
260    ///    // [[0, 1, 1], [1, 0, 0]]
261    ///    let (tensor, indices) = tensor.topk_with_indices(1, 1);
262    ///    println!("{tensor}");
263    ///    // [[12.0], [6.0]]
264    ///    println!("{indices}");
265    ///    // [[0], [2]]
266    /// }
267    /// ```
268    pub fn topk_with_indices(self, k: usize, dim: usize) -> (Self, Tensor<B, D, Int>) {
269        assert!(self.shape()[dim] >= k);
270        let k_indices = Tensor::arange(0..k as i64, &self.device());
271        let (values, indices) = self.sort_descending_with_indices(dim);
272        (
273            values.select(dim, k_indices.clone()),
274            indices.select(dim, k_indices),
275        )
276    }
277
278    /// Create a one hot tensor.
279    ///
280    /// # Example
281    ///
282    /// ```rust
283    /// use ruda_tensor::api::backend::Backend;
284    /// use ruda_tensor::api::Tensor;
285    ///
286    /// fn example<B: Backend>(){
287    ///     let device = Default::default();
288    ///     let indices: Tensor<B, 1> = Tensor::from_floats([0.0, 1.0, 2.0, 3.0], &device);
289    ///     let one_hot: Tensor<B, 2> = indices.one_hot(4);
290    ///     println!("{}", one_hot.to_data());
291    ///     // [[1.0, 0.0, 0.0, 0.0], [0.0, 1.0, 0.0, 0.0], [0.0, 0.0, 1.0, 0.0], [0.0, 0.0, 0.0, 1.0]]
292    /// }
293    /// ```
294    pub fn one_hot<const D2: usize>(self, num_classes: usize) -> Tensor<B, D2, K> {
295        check!(TensorCheck::one_hot_tensor(self.clone(), num_classes));
296        self.one_hot_fill(num_classes, 1.0, 0.0, -1)
297    }
298
299    /// Create a one-hot encoded tensor with configurable `num_classes`, `on_value`, `off_value`, and `axis` including high-ranked tensors.
300    ///
301    /// # Arguments
302    ///
303    /// * `num_classes`: The number of classes for the one-hot encoding, which defines the size of the one-hot dimension.
304    /// * `on_value`: The value to assign for active positions (corresponding to indices).
305    /// * `off_value`: The value to assign for inactive positions.
306    /// * `axis`: The axis along which the one-hot dimension is added. Supports negative indexing.
307    ///
308    /// # Returns
309    ///
310    /// A tensor with one additional dimension for the one-hot encoding, where active positions are filled with `on_value` and others with `off_value`.
311    ///
312    /// # Example
313    /// ```rust
314    /// use ruda_tensor::api::backend::Backend;
315    /// use ruda_tensor::api::{Tensor, Float};
316    /// fn example<B: Backend<FloatElem: From<f32>>>() {
317    ///     let device = B::Device::default();
318    ///     let indices: Tensor<B, 2, Float> = Tensor::from_floats([[0., 2.], [1., -1.]], &device);
319    ///     // One-hot encoding
320    ///     let tensor:Tensor<B, 3, Float> = indices.one_hot_fill(3, 5.0.into(), 0.0.into(), -1);
321    ///     println!("{tensor}");
322    ///     // [[[5.0, 0.0, 0.0],
323    ///     // [0.0, 0.0, 5.0]],
324    ///     // [[0.0, 5.0, 0.0],
325    ///     // [0.0, 0.0, 5.0]]]
326    /// }
327    /// ```
328    pub fn one_hot_fill<const D2: usize>(
329        self,
330        num_classes: usize,
331        on_value: f32,
332        off_value: f32,
333        axis: i64,
334    ) -> Tensor<B, D2, K> {
335        check!(TensorCheck::one_hot_tensor_rank::<D, D2>());
336        // Initialize shape from the current tensor dimensions and prepare for modification
337        let mut shape = self.shape();
338        let device = self.device();
339        let rank = self.dims().len();
340
341        // Adjust negative axis to a positive index
342        let axis = if axis < 0 {
343            axis + rank as i64 + 1
344        } else {
345            axis
346        };
347
348        // Ensure axis is within valid range
349        if axis < 0 || axis > rank as i64 {
350            panic!("Axis out of range. Accepted range is [-r-1, r] where r = rank(indices).");
351        }
352        // Convert the input tensor to integer indices
353        let indices: Tensor<B, D, Int> =
354            Tensor::from_data(self.to_data().convert::<i64>(), &device);
355        // Insert the new dimension for the one-hot representation
356        shape.insert(axis as usize, num_classes);
357        // Adjust indices to valid range and handle invalid indices
358        let adjusted_indices = indices
359            .clone()
360            .mask_fill(self.clone().lower_elem(0), num_classes as i64) // Handle negative indices
361            .add(indices.clone().mask_fill(self.clone().greater_elem(0), 0)); // Handle positive indices
362        // Unsqueeze the indices tensor along the specified axis
363        let indices_unsqueezed: Tensor<B, D2, Int> = adjusted_indices.unsqueeze_dim(axis as usize);
364
365        // Initialize the output tensor with the off_value
366        let output = Tensor::full(shape.clone(), off_value, &device);
367
368        // Prepare scatter tensor for on_value and off_value adjustments
369        let scatter_on_values = Tensor::full(indices_unsqueezed.shape(), on_value, &device)
370            - Tensor::full(indices_unsqueezed.shape(), off_value, &self.device());
371
372        // Scatter on_value at the appropriate indices to create the one-hot representation
373        output.scatter(
374            axis as usize,
375            indices_unsqueezed,
376            scatter_on_values,
377            IndexingUpdateOp::Add,
378        )
379    }
380
381    /// Applies element wise greater comparison and returns a boolean tensor.
382    ///
383    /// # Panics
384    ///
385    /// If the two tensors don't have the same shape.
386    ///
387    /// # Example
388    ///
389    /// ```rust
390    /// use ruda_tensor::api::backend::Backend;
391    /// use ruda_tensor::api::{Tensor, Shape};
392    ///
393    /// fn example<B: Backend>() {
394    ///   let device = B::Device::default();
395    ///   let tensor1 = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
396    ///   let tensor2 = Tensor::<B, 2>::from_data([[1.0, 3.0, 4.0], [1.0, 2.0, 3.0]], &device);
397    ///   let tensor = tensor1.greater(tensor2);
398    ///   println!("{tensor}");
399    ///   // [[false, false, false], [true, true, true]]
400    /// }
401    /// ```
402    pub fn greater(self, other: Self) -> Tensor<B, D, Bool> {
403        check!(TensorCheck::binary_ops_ew("Greater", &self, &other));
404        Tensor::new(K::greater(self.primitive, other.primitive))
405    }
406
407    /// Applies element wise greater-equal comparison and returns a boolean tensor.
408    ///
409    /// # Panics
410    ///
411    /// If the two tensors don't have the same shape.
412    ///
413    /// # Example
414    ///
415    /// ```rust
416    /// use ruda_tensor::api::backend::Backend;
417    /// use ruda_tensor::api::{Tensor, Shape};
418    ///
419    /// fn example<B: Backend>() {
420    ///    let device = B::Device::default();
421    ///    let tensor1 = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
422    ///    let tensor2 = Tensor::<B, 2>::from_data([[1.0, 3.0, 4.0], [1.0, 2.0, 3.0]], &device);
423    ///    let tensor = tensor1.greater_equal(tensor2);
424    ///    println!("{tensor}");
425    ///    // [[true, false, false], [true, true, true]]
426    /// }
427    /// ```
428    pub fn greater_equal(self, other: Self) -> Tensor<B, D, Bool> {
429        check!(TensorCheck::binary_ops_ew("Greater_equal", &self, &other));
430        Tensor::new(K::greater_equal(self.primitive, other.primitive))
431    }
432
433    /// Applies element wise lower comparison and returns a boolean tensor.
434    ///
435    /// # Panics
436    ///
437    /// If the two tensors don't have the same shape.
438    ///
439    /// # Example
440    ///
441    /// ```rust
442    /// use ruda_tensor::api::backend::Backend;
443    /// use ruda_tensor::api::{Tensor, Shape};
444    ///
445    /// fn example<B: Backend>() {
446    ///    let device = B::Device::default();
447    ///    let tensor1 = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
448    ///    let tensor2 = Tensor::<B, 2>::from_data([[1.0, 3.0, 4.0], [1.0, 2.0, 3.0]], &device);
449    ///    let tensor = tensor1.lower(tensor2);
450    ///    println!("{tensor}");
451    ///    // [[false, true, true], [false, false, false]]
452    /// }
453    /// ```
454    pub fn lower(self, other: Self) -> Tensor<B, D, Bool> {
455        check!(TensorCheck::binary_ops_ew("Lower", &self, &other));
456        Tensor::new(K::lower(self.primitive, other.primitive))
457    }
458
459    /// Applies element wise lower-equal comparison and returns a boolean tensor.
460    ///
461    /// # Panics
462    ///
463    /// If the two tensors don't have the same shape.
464    ///
465    /// # Example
466    ///
467    /// ```rust
468    /// use ruda_tensor::api::backend::Backend;
469    /// use ruda_tensor::api::{Tensor, Shape};
470    ///
471    /// fn example<B: Backend>() {
472    ///    let device = B::Device::default();
473    ///    let tensor1 = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
474    ///    let tensor2 = Tensor::<B, 2>::from_data([[1.0, 3.0, 4.0], [1.0, 2.0, 3.0]], &device);
475    ///    let tensor = tensor1.lower_equal(tensor2);
476    ///    println!("{tensor}");
477    ///    // [[true, true, true], [false, false, false]]
478    /// }
479    /// ```
480    pub fn lower_equal(self, other: Self) -> Tensor<B, D, Bool> {
481        check!(TensorCheck::binary_ops_ew("Lower_equal", &self, &other));
482        Tensor::new(K::lower_equal(self.primitive, other.primitive))
483    }
484
485    /// Applies greater than `other` comparison and returns a boolean tensor.
486    ///
487    /// # Arguments
488    ///
489    /// * `other` - The element to compare.
490    ///
491    /// # Example
492    ///
493    /// ```rust
494    /// use ruda_tensor::api::backend::Backend;
495    /// use ruda_tensor::api::{Tensor, Shape};
496    ///
497    /// fn example<B: Backend>() {
498    ///    let device = B::Device::default();
499    ///    let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
500    ///    let tensor = tensor.greater_elem(3.0);
501    ///    println!("{tensor}");
502    ///    // [[false, false, true], [true, true, true]]
503    /// }
504    /// ```
505    pub fn greater_elem<E: ElementConversion>(self, other: E) -> Tensor<B, D, Bool> {
506        let other = Scalar::new(other, &self.dtype());
507        Tensor::new(K::greater_elem(self.primitive, other))
508    }
509
510    /// Applies greater-equal than `other` comparison and returns a boolean tensor.
511    ///
512    /// # Arguments
513    ///
514    /// * `other` - The element to compare.
515    ///
516    /// # Example
517    ///
518    /// ```rust
519    /// use ruda_tensor::api::backend::Backend;
520    /// use ruda_tensor::api::{Tensor, Shape};
521    ///
522    /// fn example<B: Backend>() {
523    ///    let device = B::Device::default();
524    ///    let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
525    ///    let tensor = tensor.greater_equal_elem(3.0);
526    ///    println!("{tensor}");
527    ///    // [[false, false, true], [true, true, true]]
528    /// }
529    /// ```
530    pub fn greater_equal_elem<E: ElementConversion>(self, other: E) -> Tensor<B, D, Bool> {
531        let other = Scalar::new(other, &self.dtype());
532        Tensor::new(K::greater_equal_elem(self.primitive, other))
533    }
534
535    /// Applies lower than `other` comparison and returns a boolean tensor.
536    ///
537    /// # Arguments
538    ///
539    /// * `other` - The element to compare.
540    ///
541    /// # Example
542    ///
543    /// ```rust
544    /// use ruda_tensor::api::backend::Backend;
545    /// use ruda_tensor::api::{Tensor, Shape};
546    ///
547    /// fn example<B: Backend>() {
548    ///     let device = B::Device::default();
549    ///     let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
550    ///     let tensor = tensor.lower_elem(3.0);
551    ///     println!("{tensor}");
552    ///     // [[true, true, false], [false, false, false]]
553    /// }
554    /// ```
555    pub fn lower_elem<E: ElementConversion>(self, other: E) -> Tensor<B, D, Bool> {
556        let other = Scalar::new(other, &self.dtype());
557        Tensor::new(K::lower_elem(self.primitive, other))
558    }
559
560    /// Applies lower-equal than `other` comparison and returns a boolean tensor.
561    ///
562    /// # Arguments
563    ///
564    /// * `other` - The element to compare.
565    ///
566    /// # Example
567    ///
568    /// ```rust
569    /// use ruda_tensor::api::backend::Backend;
570    /// use ruda_tensor::api::{Tensor, Shape};
571    ///
572    /// fn example<B: Backend>() {
573    ///    let device = B::Device::default();
574    ///    let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
575    ///    let tensor = tensor.lower_equal_elem(3.0);
576    ///    println!("{tensor}");
577    ///    // [[true, true, true], [false, false, false]]
578    /// }
579    /// ```
580    pub fn lower_equal_elem<E: ElementConversion>(self, other: E) -> Tensor<B, D, Bool> {
581        let other = Scalar::new(other, &self.dtype());
582        Tensor::new(K::lower_equal_elem(self.primitive, other))
583    }
584
585    /// Applies the argmax function along the given dimension and returns an integer tensor.
586    ///
587    /// # Example
588    ///
589    /// ```rust
590    /// use ruda_tensor::api::backend::Backend;
591    /// use ruda_tensor::api::{Tensor, Shape};
592    ///
593    /// fn example<B: Backend>() {
594    ///     let device = B::Device::default();
595    ///     let tensor = Tensor::<B, 3>::ones(Shape::new([2, 3, 3]), &device);
596    ///     let tensor = tensor.argmax(1);
597    ///     println!("{:?}", tensor.shape());
598    ///     // Shape { dims: [2, 1, 3] }
599    /// }
600    /// ```
601    pub fn argmax(self, dim: usize) -> Tensor<B, D, Int> {
602        Tensor::new(K::argmax(self.primitive, dim))
603    }
604
605    /// Applies the argtopk function along the given dimension and returns an integer tensor.
606    ///
607    /// # Example
608    ///
609    /// ```rust
610    /// use ruda_tensor::api::backend::Backend;
611    /// use ruda_tensor::api::{Tensor, Shape};
612    ///
613    /// fn example<B: Backend>() {
614    ///     let device = B::Device::default();
615    ///     let tensor = Tensor::<B, 3>::ones(Shape::new([2, 3, 3]), &device);
616    ///     let tensor = tensor.argtopk(1, 2);
617    ///     println!("{:?}", tensor.shape());
618    /// }
619    /// ```
620    pub fn argtopk(self, k: usize, dim: usize) -> Tensor<B, D, Int> {
621        assert!(self.shape()[dim] >= k);
622        Tensor::new(K::argtopk(self.primitive, dim, k))
623    }
624
625    /// Find the maximum value.
626    ///
627    /// # Example
628    ///
629    /// ```rust
630    /// use ruda_tensor::api::backend::Backend;
631    /// use ruda_tensor::api::{Tensor, Shape};
632    ///
633    /// fn example<B: Backend>() {
634    ///   let device = B::Device::default();
635    ///   let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
636    ///   let tensor = tensor.max();
637    ///   println!("{tensor}");
638    ///   // [9.0]
639    /// }
640    /// ```
641    pub fn max(self) -> Tensor<B, 1, K> {
642        Tensor::new(K::max(self.primitive))
643    }
644
645    /// Find the maximum value along the given dimension.
646    ///
647    /// Also returns the indices.
648    ///
649    /// # Example
650    ///
651    /// ```rust
652    /// use ruda_tensor::api::backend::Backend;
653    /// use ruda_tensor::api::{Tensor, Shape};
654    ///
655    /// fn example<B: Backend>() {
656    ///    let device = B::Device::default();
657    ///    let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
658    ///    let (tensor, index) = tensor.max_dim_with_indices(0);
659    ///    // [[5.0, 9.0, 6.0]]
660    ///    println!("{tensor}");
661    ///    // [[1, 1, 1]]
662    ///    println!("{index}");
663    /// }
664    /// ```
665    pub fn max_dim_with_indices<I: AsIndex>(self, dim: I) -> (Self, Tensor<B, D, Int>) {
666        let dim = dim.expect_dim_index(D);
667        check!(TensorCheck::aggregate_dim::<D>("Max", dim));
668
669        let (tensor, index) = K::max_dim_with_indices(self.primitive, dim);
670
671        let tensor = Tensor::new(tensor);
672        let index = Tensor::new(index);
673
674        (tensor, index)
675    }
676
677    /// Find the maximum absolute value.
678    ///
679    /// # Example
680    ///
681    /// ```rust
682    /// use ruda_tensor::api::backend::Backend;
683    /// use ruda_tensor::api::{Tensor, Shape};
684    ///
685    /// fn example<B: Backend>() {
686    ///   let device = B::Device::default();
687    ///   let tensor = Tensor::<B, 2>::from_data([[1.0, -7.0, 3.0], [5.0, -1.0, 6.0]], &device);
688    ///   let tensor = tensor.max_abs();
689    ///   println!("{tensor}");
690    ///   // [7.0]
691    /// }
692    /// ```
693    pub fn max_abs(self) -> Tensor<B, 1, K> {
694        Tensor::new(K::max_abs(self.primitive))
695    }
696
697    /// Finds the maximum pair wise values with another tensor.
698    ///
699    /// # Arguments
700    ///
701    /// * `other` - Other tensor to find maximum elements with
702    ///
703    /// # Returns
704    ///
705    /// A tensor with the same shape as the input tensors containing the maximum value found
706    /// in the input tensors.
707    ///
708    /// # Example
709    ///
710    /// ```rust
711    /// use ruda_tensor::api::backend::Backend;
712    /// use ruda_tensor::api::{Tensor, Shape};
713    ///
714    /// fn example<B: Backend>() {
715    ///    let device = B::Device::default();
716    ///    let tensor1 = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
717    ///    let tensor2 = Tensor::<B, 2>::from_data([[2.0, 3.0, 4.0], [1.0, 2.0, 3.0]], &device);
718    ///    let tensor = tensor1.max_pair(tensor2);
719    ///    println!("{tensor}");
720    ///    // [[2.0, 3.0, 4.0], [5.0, 9.0, 6.0]]
721    /// }
722    /// ```
723    pub fn max_pair(self, other: Self) -> Self {
724        let mask = self.clone().lower(other.clone());
725        self.mask_where(mask, other)
726    }
727
728    /// Find the maximum absolute value along the given dimension.
729    ///
730    /// # Arguments
731    ///
732    /// * `dim` - The dimension or axis along which to aggregate the elements,
733    ///   supports negative indexing.
734    ///
735    /// # Returns
736    ///
737    /// The returned tensor will have the same rank,
738    /// but the aggregated dimension will have size 1.
739    ///
740    /// # Example
741    ///
742    /// ```rust
743    /// use ruda_tensor::api::backend::Backend;
744    /// use ruda_tensor::api::{Tensor, Shape};
745    ///
746    /// fn example<B: Backend>() {
747    ///   let device = B::Device::default();
748    ///   let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
749    ///   let tensor = tensor.max_dim(0);
750    ///   println!("{tensor}");
751    ///   // [[5.0, 9.0, 6.0]]
752    /// }
753    /// ```
754    pub fn max_abs_dim<I: AsIndex>(self, dim: I) -> Self {
755        let dim = dim.expect_dim_index(D);
756        check!(TensorCheck::aggregate_dim::<D>("MaxAbs", dim));
757
758        Tensor::new(K::max_abs_dim(self.primitive, dim))
759    }
760
761    /// Find the maximum absolute value along the given dimensions.
762    ///
763    /// # Arguments
764    ///
765    /// * `dims` - The dimensions or axes along which to aggregate the elements,
766    ///   supports negative indexing.
767    ///
768    /// # Returns
769    ///
770    /// The returned tensor will have the same rank,
771    /// but the aggregated dimensions will have size 1.
772    ///
773    /// # Example
774    ///
775    /// ```rust
776    /// use ruda_tensor::api::backend::Backend;
777    /// use ruda_tensor::api::{Tensor, Shape};
778    ///
779    /// fn example<B: Backend>() {
780    ///   let device = B::Device::default();
781    ///   let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
782    ///   let tensor = tensor.max_abs_dims(&[0, 1]);
783    ///   println!("{tensor}");
784    ///   // [[9.0]]
785    /// }
786    /// ```
787    pub fn max_abs_dims<I: AsIndex>(self, dims: &[I]) -> Self {
788        dims.iter()
789            .fold(self, |tensor, &dim| tensor.max_abs_dim(dim))
790    }
791
792    /// Applies the argmin function along the given dimension and returns an integer tensor.
793    ///
794    /// # Example
795    ///
796    /// ```rust
797    /// use ruda_tensor::api::backend::Backend;
798    /// use ruda_tensor::api::{Tensor, Shape};
799    ///
800    /// fn example<B: Backend>() {
801    ///     let device = Default::default();
802    ///     let tensor = Tensor::<B, 3>::ones(Shape::new([2, 3, 3]), &device);
803    ///     let tensor = tensor.argmin(1);
804    ///     println!("{:?}", tensor.shape());
805    ///     // Shape { dims: [2, 1, 3] }
806    /// }
807    /// ```
808    pub fn argmin(self, dim: usize) -> Tensor<B, D, Int> {
809        Tensor::new(K::argmin(self.primitive, dim))
810    }
811
812    /// Find the minimum value.
813    ///
814    /// # Example
815    ///
816    /// ```rust
817    /// use ruda_tensor::api::backend::Backend;
818    /// use ruda_tensor::api::{Tensor, Shape};
819    ///
820    /// fn example<B: Backend>() {
821    ///    let device = B::Device::default();
822    ///    let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
823    ///    let tensor = tensor.min();
824    ///    println!("{tensor}");
825    ///    // [-2.0]
826    /// }
827    /// ```
828    pub fn min(self) -> Tensor<B, 1, K> {
829        Tensor::new(K::min(self.primitive))
830    }
831
832    /// Find the minimum value along the given dimension.
833    ///
834    /// # Arguments
835    ///
836    /// * `dim` - The dimension or axis along which to aggregate the elements;
837    ///   supports negative indexing.
838    ///
839    /// # Returns
840    ///
841    /// The returned tensor will have the same rank,
842    /// but the aggregated dimension will have size 1.
843    ///
844    /// # Example
845    ///
846    /// ```rust
847    /// use ruda_tensor::api::backend::Backend;
848    /// use ruda_tensor::api::{Tensor, Shape};
849    ///
850    /// fn example<B: Backend>() {
851    ///    let device = B::Device::default();
852    ///    let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
853    ///    let tensor = tensor.min_dim(0);
854    ///    println!("{tensor}");
855    ///    // [[1.0, -2.0, 3.0]]
856    /// }
857    /// ```
858    pub fn min_dim<I: AsIndex>(self, dim: I) -> Self {
859        let dim = dim.expect_dim_index(D);
860        check!(TensorCheck::aggregate_dim::<D>("Min", dim));
861        Tensor::new(K::min_dim(self.primitive, dim))
862    }
863
864    /// Find the minimum value along the given dimensions.
865    ///
866    /// # Arguments
867    ///
868    /// * `dims` - The dimensions or axes along which to aggregate the elements;
869    ///   supports negative indexing.
870    ///
871    /// # Returns
872    ///
873    /// The returned tensor will have the same rank,
874    /// but the aggregated dimensions will have size 1.
875    ///
876    /// # Example
877    ///
878    /// ```rust
879    /// use ruda_tensor::api::backend::Backend;
880    /// use ruda_tensor::api::{Tensor, Shape};
881    ///
882    /// fn example<B: Backend>() {
883    ///   let device = B::Device::default();
884    ///   let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
885    ///   let tensor = tensor.min_dims(&[0, 1]);
886    ///   println!("{tensor}");
887    ///   // [[-2.0]]
888    /// }
889    /// ```
890    pub fn min_dims<I: AsIndex>(self, dims: &[I]) -> Self {
891        dims.iter().fold(self, |tensor, &dim| tensor.min_dim(dim))
892    }
893
894    /// Find the minimum value along the given dimension.
895    ///
896    /// Also returns the indices.
897    ///
898    /// # Example
899    ///
900    /// ```rust
901    /// use ruda_tensor::api::backend::Backend;
902    /// use ruda_tensor::api::{Tensor, Shape};
903    ///
904    /// fn example<B: Backend>() {
905    ///    let device = B::Device::default();
906    ///    let tensor = Tensor::<B, 2>::from_data([[7.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
907    ///    let (tensor, index) = tensor.min_dim_with_indices(0);
908    ///    println!("{tensor}");
909    ///    // [[5.0, -2.0, 3.0]]
910    ///    println!("{}", index);
911    ///    // [[1, 0, 0]]
912    /// }
913    /// ```
914    pub fn min_dim_with_indices<I: AsIndex>(self, dim: I) -> (Self, Tensor<B, D, Int>) {
915        let dim = dim.expect_dim_index(D);
916        check!(TensorCheck::aggregate_dim::<D>("Min", dim));
917
918        let (tensor, index) = K::min_dim_with_indices(self.primitive, dim);
919
920        let tensor = Tensor::new(tensor);
921        let index = Tensor::new(index);
922
923        (tensor, index)
924    }
925
926    /// Finds the minimum pair wise values with another tensor.
927    ///
928    /// # Arguments
929    ///
930    /// * `other` - Other tensor to find minimum elements with
931    ///
932    /// # Returns
933    ///
934    /// A tensor with the same shape as the input tensors containing the minimum value found
935    /// between each element of the two source tensors.
936    ///
937    /// # Example
938    ///
939    /// ```rust
940    /// use ruda_tensor::api::backend::Backend;
941    /// use ruda_tensor::api::{Tensor, Shape};
942    ///
943    /// fn example<B: Backend>() {
944    ///    let device = B::Device::default();
945    ///    let tensor1 = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
946    ///    let tensor2 = Tensor::<B, 2>::from_data([[2.0, 3.0, 4.0], [1.0, 2.0, 3.0]], &device);
947    ///    let tensor = tensor1.min_pair(tensor2);
948    ///    println!("{tensor}");
949    ///    // [[1.0, -2.0, 3.0], [1.0, 2.0, 3.0]]
950    /// }
951    pub fn min_pair(self, other: Self) -> Self {
952        let mask = other.clone().lower(self.clone());
953        self.mask_where(mask, other)
954    }
955
956    /// Clamp element wise between the given min and max values.
957    ///
958    /// # Arguments
959    ///
960    /// * `min` - The minimum value.
961    /// * `max` - The maximum value.
962    ///
963    /// # Returns
964    ///
965    /// A new tensor with the values clamped between the given min and max values.
966    ///
967    /// # Example
968    ///
969    /// ```rust
970    /// use ruda_tensor::api::backend::Backend;
971    /// use ruda_tensor::api::{Int, Tensor};
972    ///
973    /// fn example<B: Backend>() {
974    ///   let device = Default::default();
975    ///   let tensor = Tensor::<B, 2, Int>::from_ints(
976    ///    [
977    ///     [1, 2, 3],
978    ///     [4, 5, 6],
979    ///     [7, 8, 9]
980    ///    ],
981    ///    &device);
982    ///    let tensor = tensor.clamp(2, 6);
983    ///    println!("{tensor}");
984    ///    // [[2, 2, 3], [4, 5, 6], [6, 6, 6]]
985    /// }
986    /// ```
987    pub fn clamp<E: ElementConversion>(self, min: E, max: E) -> Self {
988        let dtype = self.dtype();
989        Self::new(K::clamp(
990            self.primitive,
991            Scalar::new(min, &dtype),
992            Scalar::new(max, &dtype),
993        ))
994    }
995
996    /// Clamp element wise under a minimum value.
997    ///
998    /// # Arguments
999    ///
1000    /// * `tensor` - The tensor to clamp.
1001    /// * `min` - The minimum value.
1002    ///
1003    /// # Returns
1004    ///
1005    /// A new tensor with the values clamped under the given min value.
1006    ///
1007    /// # Example
1008    ///
1009    /// ```rust
1010    /// use ruda_tensor::api::backend::Backend;
1011    /// use ruda_tensor::api::{Int, Tensor};
1012    ///
1013    /// fn example<B: Backend>() {
1014    ///    let device = Default::default();
1015    ///    let tensor = Tensor::<B, 2, Int>::from_ints(
1016    ///    [[1, 2, 3], [4, 5, 6], [7, 8, 9]],
1017    ///    &device);
1018    ///    let tensor = tensor.clamp_min(4);
1019    ///    println!("{tensor}");
1020    ///    // [[4, 4, 4], [4, 5, 6], [7, 8, 9]]
1021    /// }
1022    /// ```
1023    pub fn clamp_min<E: ElementConversion>(self, min: E) -> Self {
1024        let min = Scalar::new(min, &self.dtype());
1025        Self::new(K::clamp_min(self.primitive, min))
1026    }
1027
1028    /// Clamp element wise over a maximum value.
1029    ///
1030    /// # Arguments
1031    ///
1032    /// * `tensor` - The tensor to clamp.
1033    /// * `max` - The maximum value.
1034    ///
1035    /// # Returns
1036    ///
1037    /// A new tensor with the values clamped over the given max value.
1038    ///
1039    /// # Example
1040    ///
1041    /// ```rust
1042    /// use ruda_tensor::api::backend::Backend;
1043    /// use ruda_tensor::api::{Int, Tensor};
1044    ///
1045    /// fn example<B: Backend>() {
1046    ///    let device = Default::default();
1047    ///    let tensor = Tensor::<B, 2, Int>::from_ints(
1048    ///    [[1, 2, 3], [4, 5, 6], [7, 8, 9]],
1049    ///    &device);
1050    ///    let tensor = tensor.clamp_max(5);
1051    ///    println!("{tensor}");
1052    ///    // [[1, 2, 3], [4, 5, 5], [5, 5, 5]]
1053    /// }
1054    /// ```
1055    pub fn clamp_max<E: ElementConversion>(self, max: E) -> Self {
1056        let max = Scalar::new(max, &self.dtype());
1057        Self::new(K::clamp_max(self.primitive, max))
1058    }
1059
1060    /// Computes the cumulative minimum of elements along the given *dimension* or *axis*.
1061    ///
1062    /// # Arguments
1063    ///
1064    /// * `dim` - The dimension or axis along which to compute the cumulative minimum.
1065    ///
1066    /// # Example
1067    ///
1068    /// ```rust
1069    /// use ruda_tensor::api::backend::Backend;
1070    /// use ruda_tensor::api::{Tensor, Shape};
1071    ///
1072    /// fn example<B: Backend>() {
1073    ///    let device = B::Device::default();
1074    ///    let tensor = Tensor::<B, 2>::from_data([[3.0, 5.0, 2.0], [4.0, 1.0, 6.0]], &device);
1075    ///    let result = tensor.clone().cummin(0);
1076    ///    println!("{result}");
1077    ///    // [[3.0, 5.0, 2.0], [3.0, 1.0, 2.0]]
1078    ///    let result = tensor.cummin(1);
1079    ///    println!("{result}");
1080    ///    // [[3.0, 3.0, 2.0], [4.0, 1.0, 1.0]]
1081    /// }
1082    /// ```
1083    pub fn cummin(self, dim: usize) -> Self {
1084        check!(TensorCheck::aggregate_dim::<D>("CumMin", dim));
1085        Self::new(K::cummin(self.primitive, dim))
1086    }
1087
1088    /// Computes the cumulative maximum of elements along the given *dimension* or *axis*.
1089    ///
1090    /// # Arguments
1091    ///
1092    /// * `dim` - The dimension or axis along which to compute the cumulative maximum.
1093    ///
1094    /// # Example
1095    ///
1096    /// ```rust
1097    /// use ruda_tensor::api::backend::Backend;
1098    /// use ruda_tensor::api::{Tensor, Shape};
1099    ///
1100    /// fn example<B: Backend>() {
1101    ///    let device = B::Device::default();
1102    ///    let tensor = Tensor::<B, 2>::from_data([[3.0, 1.0, 2.0], [4.0, 5.0, 2.0]], &device);
1103    ///    let result = tensor.clone().cummax(0);
1104    ///    println!("{result}");
1105    ///    // [[3.0, 1.0, 2.0], [4.0, 5.0, 2.0]]
1106    ///    let result = tensor.cummax(1);
1107    ///    println!("{result}");
1108    ///    // [[3.0, 3.0, 3.0], [4.0, 5.0, 5.0]]
1109    /// }
1110    /// ```
1111    pub fn cummax(self, dim: usize) -> Self {
1112        check!(TensorCheck::aggregate_dim::<D>("CumMax", dim));
1113        Self::new(K::cummax(self.primitive, dim))
1114    }
1115    /// Find the maximum value along the given dimension.
1116    ///
1117    /// # Arguments
1118    ///
1119    /// * `dim` - The dimension or axis along which to aggregate the elements;
1120    ///   supports negative indexing.
1121    ///
1122    /// # Returns
1123    ///
1124    /// The returned tensor will have the same rank,
1125    /// but the aggregated dimension will have size 1.
1126    ///
1127    /// # Example
1128    ///
1129    /// ```rust
1130    /// use ruda_tensor::api::backend::Backend;
1131    /// use ruda_tensor::api::{Tensor, Shape};
1132    ///
1133    /// fn example<B: Backend>() {
1134    ///   let device = B::Device::default();
1135    ///   let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
1136    ///   let tensor = tensor.max_dim(0);
1137    ///   println!("{tensor}");
1138    ///   // [[5.0, 9.0, 6.0]]
1139    /// }
1140    /// ```
1141    pub fn max_dim<I: AsIndex>(self, dim: I) -> Self {
1142        let dim = dim.expect_dim_index(D);
1143        check!(TensorCheck::aggregate_dim::<D>("Max", dim));
1144        Tensor::new(K::max_dim(self.primitive, dim))
1145    }
1146
1147    /// Find the maximum value along the given dimensions.
1148    ///
1149    /// # Arguments
1150    ///
1151    /// * `dims` - The dimensions or axis along which to aggregate the elements;
1152    ///   supports negative indexing.
1153    ///
1154    /// # Returns
1155    ///
1156    /// The returned tensor will have the same rank,
1157    /// but the aggregated dimensions will have size 1.
1158    ///
1159    /// # Example
1160    ///
1161    /// ```rust
1162    /// use ruda_tensor::api::backend::Backend;
1163    /// use ruda_tensor::api::{Tensor, Shape};
1164    ///
1165    /// fn example<B: Backend>() {
1166    ///   let device = B::Device::default();
1167    ///   let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
1168    ///   let tensor = tensor.max_dims(&[0, 1]);
1169    ///   println!("{tensor}");
1170    ///   // [[9.0]]
1171    /// }
1172    /// ```
1173    pub fn max_dims<I: AsIndex>(self, dims: &[I]) -> Self {
1174        dims.iter().fold(self, |tensor, &dim| tensor.max_dim(dim))
1175    }
1176}