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}