ruda_tensor/ops/tensor.rs
1use super::cat::cat_with_slice_assign;
2use super::grid_sample::float_grid_sample_2d_ref;
3use super::repeat_dim::repeat_with_slice_assign;
4use super::sort::{argsort, sort, sort_with_indices};
5use crate::ops::GridSampleOptions;
6use crate::tensor::{BoolTensor, Device, Float, FloatTensor, IntTensor};
7use crate::{Backend, Distribution, TensorData, get_device_settings};
8use crate::{ExecutionError, Scalar, TensorMetadata, TensorPrimitive};
9use alloc::vec::Vec;
10use ruda_core::tensor::{BoolDType, FloatDType, IntDType, Shape, Slice};
11
12/// Operations on float tensors.
13pub trait FloatTensorOps<B: Backend> {
14 /// Creates a new tensor from the data structure.
15 ///
16 /// # Arguments
17 ///
18 /// * `data` - The data structure.
19 /// * `device` - The device to create the tensor on.
20 ///
21 /// # Returns
22 ///
23 /// The tensor with the given data.
24 fn float_from_data(data: TensorData, device: &Device<B>) -> FloatTensor<B>;
25
26 /// Creates a new tensor with random values.
27 ///
28 /// # Arguments
29 ///
30 /// * `shape` - The shape of the tensor.
31 /// * `distribution` - The distribution to sample from.
32 /// * `device` - The device to create the tensor on.
33 /// * `dtype` - The target data type.
34 ///
35 /// # Returns
36 ///
37 /// The tensor with the given shape and random values.
38 fn float_random(
39 shape: Shape,
40 distribution: Distribution,
41 device: &Device<B>,
42 dtype: FloatDType,
43 ) -> FloatTensor<B>;
44
45 /// Creates a new tensor with zeros.
46 ///
47 /// # Arguments
48 ///
49 /// * `shape` - The shape of the tensor.
50 /// * `device` - The device to create the tensor on.
51 /// * `dtype` - The target data type.
52 ///
53 /// # Returns
54 ///
55 /// The tensor with the given shape and zeros.
56 fn float_zeros(shape: Shape, device: &Device<B>, dtype: FloatDType) -> FloatTensor<B> {
57 Self::float_from_data(TensorData::full_dtype(shape, 0., dtype.into()), device)
58 }
59
60 /// Creates a new tensor with ones.
61 ///
62 /// # Arguments
63 ///
64 /// * `shape` - The shape of the tensor.
65 /// * `device` - The device to create the tensor on.
66 /// * `dtype` - The target data type.
67 ///
68 /// # Returns
69 ///
70 /// The tensor with the given shape and ones.
71 fn float_ones(shape: Shape, device: &Device<B>, dtype: FloatDType) -> FloatTensor<B> {
72 Self::float_from_data(TensorData::full_dtype(shape, 1., dtype.into()), device)
73 }
74
75 /// Creates a tensor filled with given value.
76 ///
77 /// # Arguments
78 ///
79 /// * `shape` - The shape of the tensor.
80 /// * `fill_value` - The value with which to fill the tensor.
81 /// * `device` - The device to create the tensor on.
82 /// * `dtype` - The target data type.
83 ///
84 /// # Returns
85 ///
86 /// The tensor filled with given value
87 fn float_full(
88 shape: Shape,
89 fill_value: Scalar,
90 device: &Device<B>,
91 dtype: FloatDType,
92 ) -> FloatTensor<B> {
93 Self::float_from_data(
94 TensorData::full_dtype(shape, fill_value, dtype.into()),
95 device,
96 )
97 }
98
99 /// Converts the tensor to a data structure.
100 ///
101 /// # Arguments
102 ///
103 /// * `tensor` - The tensor.
104 ///
105 /// # Returns
106 ///
107 /// The data structure with the tensor's data.
108 fn float_into_data(
109 tensor: FloatTensor<B>,
110 ) -> impl Future<Output = Result<TensorData, ExecutionError>> + Send;
111
112 /// Gets the device of the tensor.
113 ///
114 /// # Arguments
115 ///
116 /// * `tensor` - The tensor.
117 ///
118 /// # Returns
119 ///
120 /// The device of the tensor.
121 fn float_device(tensor: &FloatTensor<B>) -> Device<B>;
122
123 /// Moves the tensor to the given device.
124 ///
125 /// # Arguments
126 ///
127 /// * `tensor` - The tensor.
128 /// * `device` - The device to move the tensor to.
129 ///
130 /// # Returns
131 ///
132 /// The tensor on the given device.
133 fn float_to_device(tensor: FloatTensor<B>, device: &Device<B>) -> FloatTensor<B>;
134
135 /// Converts float tensor to int tensor.
136 ///
137 /// # Arguments
138 ///
139 /// * `tensor` - The tensor.
140 /// * `out_dtype` - The output tensor dtype.
141 ///
142 /// # Returns
143 ///
144 /// The int tensor with the same data as the float tensor.
145 fn float_into_int(tensor: FloatTensor<B>, out_dtype: IntDType) -> IntTensor<B>;
146
147 /// Creates an empty tensor with the given shape.
148 ///
149 /// # Arguments
150 ///
151 /// * `shape` - The shape of the tensor.
152 /// * `device` - The device to create the tensor on.
153 /// * `dtype` - The target data type.
154 ///
155 /// # Returns
156 ///
157 /// The empty tensor with the given shape.
158 fn float_empty(shape: Shape, device: &Device<B>, dtype: FloatDType) -> FloatTensor<B>;
159
160 /// Repeat the tensor along the given dimension.
161 ///
162 /// # Arguments
163 ///
164 /// * `tensor` - The tensor.
165 /// * `dim` - The dimension to repeat.
166 /// * `times` - The number of times to repeat the dimension.
167 ///
168 /// # Returns
169 ///
170 /// The tensor with the given dimension repeated.
171 fn float_repeat_dim(tensor: FloatTensor<B>, dim: usize, times: usize) -> FloatTensor<B> {
172 repeat_with_slice_assign::<B, Float>(TensorPrimitive::Float(tensor), dim, times).tensor()
173 }
174
175 /// Adds two tensors together.
176 ///
177 /// # Arguments
178 ///
179 /// * `lhs` - The left-hand side tensor.
180 /// * `rhs` - The right-hand side tensor.
181 ///
182 /// # Returns
183 ///
184 /// The result of adding the two tensors together.
185 fn float_add(lhs: FloatTensor<B>, rhs: FloatTensor<B>) -> FloatTensor<B>;
186
187 /// Adds a scalar to a tensor.
188 ///
189 /// # Arguments
190 ///
191 /// * `lhs` - The left-hand side tensor.
192 /// * `rhs` - The right-hand side scalar.
193 ///
194 /// # Returns
195 ///
196 /// The result of adding the scalar to the tensor.
197 fn float_add_scalar(lhs: FloatTensor<B>, rhs: Scalar) -> FloatTensor<B>;
198
199 /// Clamps a tensor under a minimum value.
200 ///
201 /// # Arguments
202 ///
203 /// * `tensor` - The tensor to clamp.
204 /// * `min` - The minimum value.
205 ///
206 /// # Returns
207 ///
208 /// The clamped tensor.
209 fn float_clamp_min(tensor: FloatTensor<B>, min: Scalar) -> FloatTensor<B> {
210 let dtype = get_device_settings::<B>(&B::float_device(&tensor)).bool_dtype;
211 let mask = Self::float_lower_elem(tensor.clone(), min, dtype);
212 B::float_mask_fill(tensor, mask, min)
213 }
214
215 /// Clamps a tensor over a maximum value.
216 ///
217 /// # Arguments
218 ///
219 /// * `tensor` - The tensor to clamp.
220 /// * `max` - The maximum value.
221 ///
222 /// # Returns
223 ///
224 /// The clamped tensor.
225 fn float_clamp_max(tensor: FloatTensor<B>, max: Scalar) -> FloatTensor<B> {
226 let dtype = get_device_settings::<B>(&B::float_device(&tensor)).bool_dtype;
227 let mask = Self::float_greater_elem(tensor.clone(), max, dtype);
228 B::float_mask_fill(tensor, mask, max)
229 }
230
231 /// Clamps a tensor between a minimum and maximum value.
232 ///
233 /// # Arguments
234 ///
235 /// * `tensor` - The tensor to clamp.
236 /// * `min` - The minimum value.
237 /// * `max` - The maximum value.
238 ///
239 /// # Returns
240 ///
241 /// The clamped tensor.
242 fn float_clamp(tensor: FloatTensor<B>, min: Scalar, max: Scalar) -> FloatTensor<B> {
243 // Default implementation
244 Self::float_clamp_min(Self::float_clamp_max(tensor, max), min)
245 }
246
247 /// Subtracts two tensors.
248 ///
249 /// # Arguments
250 ///
251 /// * `lhs` - The left-hand side tensor.
252 /// * `rhs` - The right-hand side tensor.
253 ///
254 /// # Returns
255 ///
256 /// The result of subtracting the two tensors.
257 fn float_sub(lhs: FloatTensor<B>, rhs: FloatTensor<B>) -> FloatTensor<B>;
258
259 /// Subtracts a scalar from a tensor.
260 ///
261 /// # Arguments
262 ///
263 /// * `lhs` - The left-hand side tensor.
264 /// * `rhs` - The right-hand side scalar.
265 ///
266 /// # Returns
267 ///
268 /// The result of subtracting the scalar from the tensor.
269 fn float_sub_scalar(lhs: FloatTensor<B>, rhs: Scalar) -> FloatTensor<B>;
270
271 /// Multiplies two tensors together element-wise.
272 fn float_mul(lhs: FloatTensor<B>, rhs: FloatTensor<B>) -> FloatTensor<B>;
273
274 /// Multiplies a tensor by a scalar.
275 ///
276 /// # Arguments
277 ///
278 /// * `lhs` - The left-hand side tensor.
279 /// * `rhs` - The right-hand side scalar.
280 ///
281 /// # Returns
282 ///
283 /// The result of multiplying the tensor by the scalar.
284 fn float_mul_scalar(lhs: FloatTensor<B>, rhs: Scalar) -> FloatTensor<B>;
285
286 /// Divides two tensors element-wise.
287 ///
288 /// # Arguments
289 ///
290 /// * `lhs` - The left-hand side tensor.
291 /// * `rhs` - The right-hand side tensor.
292 ///
293 /// # Returns
294 ///
295 /// The result of dividing the two tensors.
296 fn float_div(lhs: FloatTensor<B>, rhs: FloatTensor<B>) -> FloatTensor<B>;
297
298 /// Divides a tensor by a scalar.
299 ///
300 /// # Arguments
301 ///
302 /// * `lhs` - The left-hand side tensor.
303 /// * `rhs` - The right-hand side scalar.
304 ///
305 /// # Returns
306 ///
307 /// The result of dividing the tensor by the scalar.
308 fn float_div_scalar(lhs: FloatTensor<B>, rhs: Scalar) -> FloatTensor<B>;
309
310 /// Computes the remainder of division between two tensors element-wise.
311 ///
312 /// # Arguments
313 ///
314 /// * `lhs` - The left-hand side tensor.
315 /// * `rhs` - The right-hand side tensor.
316 ///
317 /// # Returns
318 ///
319 /// The element-wise remainder when dividing `lhs` by `rhs`.
320 fn float_remainder(lhs: FloatTensor<B>, rhs: FloatTensor<B>) -> FloatTensor<B>;
321
322 /// Computes the modulus of a tensor given a scalar.
323 ///
324 /// # Arguments
325 /// * `lhs` - The left-hand side tensor.
326 /// * `rhs` - The right-hand side scalar.
327 ///
328 /// # Returns
329 ///
330 /// The result of applying the modulus of the scalar to the tensor.
331 fn float_remainder_scalar(lhs: FloatTensor<B>, rhs: Scalar) -> FloatTensor<B>;
332
333 /// Multiplies two tensors together using matrix multiplication.
334 ///
335 /// # Arguments
336 ///
337 /// * `lhs` - The left-hand side tensor.
338 /// * `rhs` - The right-hand side tensor.
339 ///
340 /// # Returns
341 ///
342 /// The result of multiplying the two tensors together using matrix multiplication.
343 fn float_matmul(lhs: FloatTensor<B>, rhs: FloatTensor<B>) -> FloatTensor<B>;
344
345 /// Computes the cross product of two tensors along a given dimension.
346 ///
347 /// # Arguments
348 ///
349 /// * `lhs` - The left-hand side tensor.
350 /// * `rhs` - The right-hand side tensor.
351 /// * `dim` - The dimension to compute the cross product along.
352 ///
353 /// # Returns
354 ///
355 /// The cross product of the two tensors.
356 fn float_cross(lhs: FloatTensor<B>, rhs: FloatTensor<B>, dim: usize) -> FloatTensor<B>;
357
358 /// Negates a tensor element-wise.
359 fn float_neg(tensor: FloatTensor<B>) -> FloatTensor<B> {
360 Self::float_mul_scalar(tensor, (-1f32).into())
361 }
362
363 /// Calculates the reciprocals element-wise
364 fn float_recip(tensor: FloatTensor<B>) -> FloatTensor<B>;
365
366 /// Transposes a tensor.
367 ///
368 /// # Arguments
369 ///
370 /// * `tensor` - The tensor to transpose.
371 ///
372 /// # Returns
373 ///
374 /// The transposed tensor.
375 fn float_transpose(tensor: FloatTensor<B>) -> FloatTensor<B> {
376 let ndims = tensor.shape().num_dims();
377 Self::float_swap_dims(tensor, ndims - 2, ndims - 1)
378 }
379
380 /// Swaps two dimensions of a tensor.
381 ///
382 /// # Arguments
383 ///
384 /// * `tensor` - The tensor to swap the dimensions of.
385 /// * `dim1` - The first dimension to swap.
386 /// * `dim2` - The second dimension to swap.
387 ///
388 /// # Returns
389 ///
390 /// The tensor with the dimensions swapped.
391 fn float_swap_dims(tensor: FloatTensor<B>, dim1: usize, dim2: usize) -> FloatTensor<B>;
392
393 /// Permutes the dimensions of a tensor.
394 ///
395 /// # Arguments
396 ///
397 /// * `tensor` - The tensor to permute the dimensions of.
398 /// * `axes` - The new order of the dimensions.
399 /// # Returns
400 ///
401 /// The tensor with the dimensions permuted.
402 fn float_permute(tensor: FloatTensor<B>, axes: &[usize]) -> FloatTensor<B>;
403
404 /// Reverse the order of elements in a tensor along the given axes.
405 ///
406 /// # Arguments
407 ///
408 /// * `tensor` - The tensor to reverse.
409 /// * `axes` - The axes to reverse.
410 ///
411 /// The tensor with the elements reversed.
412 fn float_flip(tensor: FloatTensor<B>, axes: &[usize]) -> FloatTensor<B>;
413
414 /// Reshapes a tensor.
415 ///
416 /// # Arguments
417 ///
418 /// * `tensor` - The tensor to reshape.
419 /// * `shape` - The new shape of the tensor.
420 ///
421 /// # Returns
422 ///
423 /// The tensor with the new shape.
424 fn float_reshape(tensor: FloatTensor<B>, shape: Shape) -> FloatTensor<B>;
425
426 /// Gather elements from a tensor.
427 ///
428 /// # Arguments
429 ///
430 /// * `dim` - The dimension to gather from.
431 /// * `tensor` - The tensor to gather from.
432 /// * `indices` - The indices to gather.
433 ///
434 /// # Returns
435 ///
436 /// The gathered elements.
437 fn float_gather(dim: usize, tensor: FloatTensor<B>, indices: IntTensor<B>) -> FloatTensor<B>;
438
439 /// Scatter elements into a tensor using sum reduction.
440 ///
441 /// # Arguments
442 ///
443 /// * `dim` - The dimension to scatter into.
444 /// * `tensor` - The tensor to scatter into.
445 /// * `indices` - The indices to scatter into.
446 /// * `value` - The value to scatter.
447 ///
448 /// # Returns
449 ///
450 /// The tensor with the scattered elements.
451 fn float_scatter_add(
452 dim: usize,
453 tensor: FloatTensor<B>,
454 indices: IntTensor<B>,
455 value: FloatTensor<B>,
456 ) -> FloatTensor<B>;
457
458 /// Multi-dimensional scatter: update `data` at locations specified by `indices` with `values`.
459 ///
460 /// # Arguments
461 ///
462 /// * `data` - The tensor to scatter into.
463 /// * `indices` - An M-dimensional integer tensor whose last dimension indexes into `data`.
464 /// * `values` - The values to scatter.
465 /// * `reduction` - How to combine with existing values.
466 ///
467 /// # Returns
468 ///
469 /// The tensor with scattered values.
470 fn float_scatter_nd(
471 _data: FloatTensor<B>,
472 _indices: IntTensor<B>,
473 _values: FloatTensor<B>,
474 _reduction: crate::tensor::IndexingUpdateOp,
475 ) -> FloatTensor<B> {
476 unimplemented!("float_scatter_nd is not implemented for this backend")
477 }
478
479 /// Multi-dimensional gather: collect slices from `data` at locations specified by `indices`.
480 ///
481 /// # Arguments
482 ///
483 /// * `data` - The tensor to gather from.
484 /// * `indices` - An M-dimensional integer tensor whose last dimension indexes into `data`.
485 ///
486 /// # Returns
487 ///
488 /// The gathered tensor.
489 fn float_gather_nd(_data: FloatTensor<B>, _indices: IntTensor<B>) -> FloatTensor<B> {
490 unimplemented!("float_gather_nd is not implemented for this backend")
491 }
492
493 /// Select tensor elements along the given dimension corresponding for the given indices.
494 ///
495 /// # Arguments
496 ///
497 /// * `tensor` - The tensor to select from.
498 /// * `dim` - The dimension to select from.
499 /// * `indices` - The indices to select.
500 ///
501 /// # Returns
502 ///
503 /// The selected elements.
504 fn float_select(tensor: FloatTensor<B>, dim: usize, indices: IntTensor<B>) -> FloatTensor<B>;
505
506 /// Assign the selected elements along the given dimension corresponding for the given indices
507 /// to the given value using sum reduction.
508 ///
509 /// # Arguments
510 ///
511 /// * `tensor` - The tensor to select from.
512 /// * `dim` - The dimension to select from.
513 /// * `indices` - The indices to select.
514 /// * `value` - The value to assign.
515 ///
516 /// # Returns
517 ///
518 /// The tensor with the selected elements assigned to the given value.
519 fn float_select_add(
520 tensor: FloatTensor<B>,
521 dim: usize,
522 indices: IntTensor<B>,
523 value: FloatTensor<B>,
524 ) -> FloatTensor<B>;
525
526 /// Select tensor elements corresponding to the given slices.
527 ///
528 /// # Arguments
529 ///
530 /// * `tensor` - The tensor to select from.
531 /// * `slices` - The slices specifying ranges and steps for each dimension.
532 ///
533 /// # Returns
534 ///
535 /// The selected elements in a new tensor.
536 ///
537 /// # Note
538 ///
539 /// Empty slices (where start >= end) are handled at the high-level tensor API and will not
540 /// be passed to this method. Backend implementations do not need to handle empty slices.
541 fn float_slice(tensor: FloatTensor<B>, slices: &[Slice]) -> FloatTensor<B>;
542
543 /// Assign the selected elements corresponding to the given slices to the given value.
544 ///
545 /// # Arguments
546 ///
547 /// * `tensor` - The tensor to select from.
548 /// * `ranges` - The ranges to select.
549 /// * `value` - The value to assign.
550 ///
551 /// # Returns
552 ///
553 /// The tensor with the selected elements assigned to the given value.
554 ///
555 /// # Note
556 ///
557 /// Empty slice assignments (where any slice range produces 0 elements) are handled at the
558 /// high-level tensor API and will not be passed to this method. Backend implementations do
559 /// not need to handle empty slice assignments.
560 fn float_slice_assign(
561 tensor: FloatTensor<B>,
562 slices: &[Slice],
563 value: FloatTensor<B>,
564 ) -> FloatTensor<B>;
565
566 /// Update the given tensor with the value tensor where the mask is true.
567 ///
568 /// # Arguments
569 ///
570 /// * `tensor` - The tensor to select from.
571 /// * `mask` - The boolean mask to select with.
572 /// * `value` - The value to assign to the selected elements from the value tensor.
573 ///
574 /// # Returns
575 ///
576 /// The tensor with the selected elements assigned to the given value.
577 fn float_mask_where(
578 tensor: FloatTensor<B>,
579 mask: BoolTensor<B>,
580 value: FloatTensor<B>,
581 ) -> FloatTensor<B>;
582
583 /// Update the given tensor with the value where the mask is true.
584 ///
585 /// # Arguments
586 ///
587 /// * `tensor` - The tensor to select from.
588 /// * `mask` - The boolean mask to select with.
589 /// * `value` - The value to assign to the selected elements.
590 ///
591 /// # Returns
592 ///
593 /// The tensor with the selected elements assigned to the given value.
594 fn float_mask_fill(
595 tensor: FloatTensor<B>,
596 mask: BoolTensor<B>,
597 value: Scalar,
598 ) -> FloatTensor<B>;
599
600 /// Equal comparison of two tensors.
601 ///
602 /// # Arguments
603 ///
604 /// * `lhs` - The left-hand side tensor.
605 /// * `rhs` - The right-hand side tensor.
606 /// * `out_dtype` - The output tensor dtype.
607 ///
608 /// # Returns
609 ///
610 /// A boolean tensor with the result of the comparison.
611 fn float_equal(lhs: FloatTensor<B>, rhs: FloatTensor<B>, out_dtype: BoolDType)
612 -> BoolTensor<B>;
613
614 /// Element-wise non-equality comparison.
615 ///
616 /// # Arguments
617 ///
618 /// * `lhs` - The left-hand side tensor.
619 /// * `rhs` - The right-hand side tensor.
620 /// * `out_dtype` - The output tensor dtype.
621 ///
622 /// # Returns
623 ///
624 /// A boolean tensor with the result of the comparison.
625 fn float_not_equal(
626 lhs: FloatTensor<B>,
627 rhs: FloatTensor<B>,
628 out_dtype: BoolDType,
629 ) -> BoolTensor<B> {
630 let equal_tensor = B::float_equal(lhs, rhs, out_dtype);
631 B::bool_not(equal_tensor)
632 }
633
634 /// Equal comparison of a tensor and a scalar.
635 ///
636 /// # Arguments
637 ///
638 /// * `lhs` - The left-hand side tensor.
639 /// * `rhs` - The right-hand side scalar.
640 /// * `out_dtype` - The output tensor dtype.
641 ///
642 /// # Returns
643 ///
644 /// A boolean tensor with the result of the comparison.
645 fn float_equal_elem(lhs: FloatTensor<B>, rhs: Scalar, out_dtype: BoolDType) -> BoolTensor<B>;
646
647 /// Element-wise non-equality comparison with a scalar.
648 ///
649 /// # Arguments
650 ///
651 /// * `lhs` - The left-hand side tensor.
652 /// * `rhs` - The right-hand side scalar.
653 /// * `out_dtype` - The output tensor dtype.
654 ///
655 /// # Returns
656 ///
657 /// A boolean tensor with the result of the comparison.
658 fn float_not_equal_elem(
659 lhs: FloatTensor<B>,
660 rhs: Scalar,
661 out_dtype: BoolDType,
662 ) -> BoolTensor<B> {
663 let equal_tensor = B::float_equal_elem(lhs, rhs, out_dtype);
664 B::bool_not(equal_tensor)
665 }
666
667 /// Greater than comparison of two tensors.
668 ///
669 /// # Arguments
670 ///
671 /// * `lhs` - The left-hand side tensor.
672 /// * `rhs` - The right-hand side tensor.
673 /// * `out_dtype` - The output tensor dtype.
674 ///
675 /// # Returns
676 ///
677 /// A boolean tensor with the result of the comparison.
678 fn float_greater(
679 lhs: FloatTensor<B>,
680 rhs: FloatTensor<B>,
681 out_dtype: BoolDType,
682 ) -> BoolTensor<B>;
683
684 /// Greater than comparison of a tensor and a scalar.
685 ///
686 /// # Arguments
687 ///
688 /// * `lhs` - The left-hand side tensor.
689 /// * `rhs` - The right-hand side scalar.
690 /// * `out_dtype` - The output tensor dtype.
691 ///
692 /// # Returns
693 ///
694 /// A boolean tensor with the result of the comparison.
695 fn float_greater_elem(lhs: FloatTensor<B>, rhs: Scalar, out_dtype: BoolDType) -> BoolTensor<B>;
696
697 /// Greater than or equal comparison of two tensors.
698 ///
699 /// # Arguments
700 ///
701 /// * `lhs` - The left-hand side tensor.
702 /// * `rhs` - The right-hand side tensor.
703 /// * `out_dtype` - The output tensor dtype.
704 ///
705 /// # Returns
706 ///
707 /// A boolean tensor with the result of the comparison.
708 fn float_greater_equal(
709 lhs: FloatTensor<B>,
710 rhs: FloatTensor<B>,
711 out_dtype: BoolDType,
712 ) -> BoolTensor<B>;
713
714 /// Greater than or equal comparison of a tensor and a scalar.
715 ///
716 /// # Arguments
717 ///
718 /// * `lhs` - The left-hand side tensor.
719 /// * `rhs` - The right-hand side scalar.
720 /// * `out_dtype` - The output tensor dtype.
721 ///
722 /// # Returns
723 ///
724 /// A boolean tensor with the result of the comparison.
725 fn float_greater_equal_elem(
726 lhs: FloatTensor<B>,
727 rhs: Scalar,
728 out_dtype: BoolDType,
729 ) -> BoolTensor<B>;
730
731 /// Less than comparison of two tensors.
732 ///
733 /// # Arguments
734 ///
735 /// * `lhs` - The left-hand side tensor.
736 /// * `rhs` - The right-hand side tensor.
737 /// * `out_dtype` - The output tensor dtype.
738 ///
739 /// # Returns
740 ///
741 /// A boolean tensor with the result of the comparison.
742 fn float_lower(lhs: FloatTensor<B>, rhs: FloatTensor<B>, out_dtype: BoolDType)
743 -> BoolTensor<B>;
744
745 /// Less than comparison of a tensor and a scalar.
746 ///
747 /// # Arguments
748 ///
749 /// * `lhs` - The left-hand side tensor.
750 /// * `rhs` - The right-hand side scalar.
751 /// * `out_dtype` - The output tensor dtype.
752 ///
753 /// # Returns
754 ///
755 /// A boolean tensor with the result of the comparison.
756 fn float_lower_elem(lhs: FloatTensor<B>, rhs: Scalar, out_dtype: BoolDType) -> BoolTensor<B>;
757
758 /// Less than or equal comparison of two tensors.
759 ///
760 /// # Arguments
761 ///
762 /// * `lhs` - The left-hand side tensor.
763 /// * `rhs` - The right-hand side tensor.
764 /// * `out_dtype` - The output tensor dtype.
765 ///
766 /// # Returns
767 ///
768 /// A boolean tensor with the result of the comparison.
769 fn float_lower_equal(
770 lhs: FloatTensor<B>,
771 rhs: FloatTensor<B>,
772 out_dtype: BoolDType,
773 ) -> BoolTensor<B>;
774
775 /// Less than or equal comparison of a tensor and a scalar.
776 ///
777 /// # Arguments
778 ///
779 /// * `lhs` - The left-hand side tensor.
780 /// * `rhs` - The right-hand side scalar.
781 /// * `out_dtype` - The output tensor dtype.
782 ///
783 /// # Returns
784 ///
785 /// A boolean tensor with the result of the comparison.
786 fn float_lower_equal_elem(
787 lhs: FloatTensor<B>,
788 rhs: Scalar,
789 out_dtype: BoolDType,
790 ) -> BoolTensor<B>;
791
792 /// Detaches a tensor from the computation graph.
793 fn float_detach(tensor: FloatTensor<B>) -> FloatTensor<B> {
794 // Should only be overridden by autodiff backends.
795 tensor
796 }
797
798 /// Sets the `require_grad` flag of a tensor.
799 fn float_set_require_grad(tensor: FloatTensor<B>, _require_grad: bool) -> FloatTensor<B> {
800 // Should only be overridden by autodiff backends.
801 tensor
802 }
803
804 /// Returns the `require_grad` flag of a tensor.
805 fn float_is_require_grad(_tensor: &FloatTensor<B>) -> bool {
806 // Should only be overridden by autodiff backends.
807 false
808 }
809
810 /// Sum of all elements in a tensor.
811 ///
812 /// # Arguments
813 ///
814 /// * `tensor` - The tensor to sum.
815 ///
816 /// # Returns
817 ///
818 /// A scalar tensor with the sum of all elements in `tensor`.
819 fn float_sum(tensor: FloatTensor<B>) -> FloatTensor<B>;
820
821 /// Sum of all elements in a tensor along a dimension.
822 ///
823 /// # Arguments
824 ///
825 /// * `tensor` - The tensor to sum.
826 /// * `dim` - The dimension along which to sum.
827 ///
828 /// # Returns
829 ///
830 /// A tensor with the sum of all elements in `tensor` along `dim`.
831 fn float_sum_dim(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B>;
832
833 /// Product of all elements in a tensor.
834 ///
835 /// # Arguments
836 ///
837 /// * `tensor` - The tensor to product.
838 ///
839 /// # Returns
840 ///
841 /// A scalar tensor with the product of all elements in `tensor`.
842 fn float_prod(tensor: FloatTensor<B>) -> FloatTensor<B> {
843 let len = tensor.shape().num_elements();
844 let tensor = B::float_reshape(tensor, Shape::new([len]));
845 B::float_prod_dim(tensor, 0)
846 }
847
848 /// Product of all elements in a tensor along a dimension.
849 ///
850 /// # Arguments
851 ///
852 /// * `tensor` - The tensor to product.
853 ///
854 /// # Returns
855 ///
856 /// A tensor with the product of all elements in `tensor` along `dim`.
857 fn float_prod_dim(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B> {
858 let mut shape = tensor.shape();
859 let len = shape[dim];
860 if len == 0 {
861 shape[dim] = 1;
862 return B::float_ones(shape, &B::float_device(&tensor), tensor.dtype().into());
863 }
864 let mut slices = alloc::vec![Slice::full(); shape.num_dims()];
865 slices[dim] = Slice::from(len - 1..len);
866 B::float_slice(B::float_cumprod(tensor, dim), &slices)
867 }
868
869 /// Mean of all elements in a tensor.
870 ///
871 /// # Arguments
872 ///
873 /// * `tensor` - The tensor to mean.
874 ///
875 /// # Returns
876 ///
877 /// A scalar tensor with the mean of all elements in `tensor`.
878 fn float_mean(tensor: FloatTensor<B>) -> FloatTensor<B> {
879 let num_elems = tensor.shape().num_elements() as f32;
880 B::float_div_scalar(B::float_sum(tensor), num_elems.into())
881 }
882
883 /// Mean of all elements in a tensor along a dimension.
884 ///
885 /// # Arguments
886 ///
887 /// * `tensor` - The tensor to mean.
888 /// * `dim` - The dimension along which to mean.
889 ///
890 /// # Returns
891 ///
892 /// A tensor with the mean of all elements in `tensor` along `dim`.
893 fn float_mean_dim(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B>;
894
895 /// Computes the cumulative sum of elements along a dimension.
896 ///
897 /// # Arguments
898 ///
899 /// * `tensor` - The tensor to compute the cumulative sum of.
900 /// * `dim` - The dimension along which to compute the cumulative sum.
901 ///
902 /// # Returns
903 ///
904 /// A tensor with the same shape where each element is the cumulative sum
905 /// of all elements up to and including that position along the dimension.
906 fn float_cumsum(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B>;
907
908 /// Computes the cumulative product of elements along a dimension.
909 ///
910 /// # Arguments
911 ///
912 /// * `tensor` - The tensor to compute the cumulative product of.
913 /// * `dim` - The dimension along which to compute the cumulative product.
914 ///
915 /// # Returns
916 ///
917 /// A tensor with the same shape where each element is the cumulative product
918 /// of all elements up to and including that position along the dimension.
919 fn float_cumprod(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B>;
920
921 /// Computes the cumulative minimum of elements along a dimension.
922 ///
923 /// # Arguments
924 ///
925 /// * `tensor` - The tensor to compute the cumulative minimum of.
926 /// * `dim` - The dimension along which to compute the cumulative minimum.
927 ///
928 /// # Returns
929 ///
930 /// A tensor with the same shape where each element is the minimum
931 /// of all elements up to and including that position along the dimension.
932 fn float_cummin(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B>;
933
934 /// Computes the cumulative maximum of elements along a dimension.
935 ///
936 /// # Arguments
937 ///
938 /// * `tensor` - The tensor to compute the cumulative maximum of.
939 /// * `dim` - The dimension along which to compute the cumulative maximum.
940 ///
941 /// # Returns
942 ///
943 /// A tensor with the same shape where each element is the maximum
944 /// of all elements up to and including that position along the dimension.
945 fn float_cummax(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B>;
946
947 /// Converts a tensor to another floating point data type.
948 ///
949 /// # Arguments
950 ///
951 /// * `tensor` - The tensor to convert.
952 /// * `dtype` - The target data type.
953 ///
954 /// # Returns
955 ///
956 /// A tensor with the same values as `tensor` but in the target floating point data type.
957 fn float_cast(tensor: FloatTensor<B>, dtype: FloatDType) -> FloatTensor<B>;
958
959 /// Returns a new tensor with exponential values.
960 ///
961 /// # Arguments
962 ///
963 /// * `tensor` - The tensor to exponentiate.
964 ///
965 /// # Returns
966 ///
967 /// A tensor with the same shape as `tensor` with exponential values.
968 fn float_exp(tensor: FloatTensor<B>) -> FloatTensor<B>;
969
970 /// Returns a new tensor with natural logarithm values.
971 ///
972 /// # Arguments
973 ///
974 /// * `tensor` - The tensor to take the logarithm of.
975 ///
976 /// # Returns
977 ///
978 /// A tensor with the same shape as `tensor` with natural logarithm values.
979 fn float_log(tensor: FloatTensor<B>) -> FloatTensor<B>;
980
981 /// Returns a new tensor with logarithm values of (1 + Xi).
982 ///
983 /// # Arguments
984 ///
985 /// * `tensor` - The tensor to take the logarithm of.
986 ///
987 /// # Returns
988 ///
989 /// A tensor with the same shape as `tensor` with logarithm values of (1 + Xi).
990 fn float_log1p(tensor: FloatTensor<B>) -> FloatTensor<B>;
991
992 /// Element-wise power with a FloatTensor.
993 ///
994 /// # Arguments
995 ///
996 /// * `lhs` - The left-hand side tensor.
997 /// * `rhs` - The right-hand side tensor.
998 ///
999 /// # Returns
1000 ///
1001 /// The elements of `lhs` raised to the power of the elements of `rhs`.
1002 fn float_powf(lhs: FloatTensor<B>, rhs: FloatTensor<B>) -> FloatTensor<B>;
1003
1004 /// Element-wise power with an IntTensor.
1005 ///
1006 /// # Arguments
1007 ///
1008 /// * `lhs` - The left-hand side tensor.
1009 /// * `rhs` - The right-hand side floatTensor.
1010 ///
1011 /// # Returns
1012 ///
1013 /// The elements of `lhs` raised to the value of `rhs`. Result is an IntTensor.
1014 fn float_powi(lhs: FloatTensor<B>, rhs: IntTensor<B>) -> FloatTensor<B> {
1015 let dtype = lhs.dtype();
1016 Self::float_powf(lhs, B::int_into_float(rhs, dtype.into()))
1017 }
1018
1019 /// Raises a tensor to the power of an int scalar.
1020 ///
1021 /// # Backend Implementors Note
1022 ///
1023 /// A number of common exponent cases can be implemented with operations
1024 /// which are much cheaper than generic exponentiation.
1025 ///
1026 /// This (`Backend` impl overridable) operation handles generic optimizations
1027 /// for several common integer exponent cases; and then dispatches to
1028 /// the (`Backend` impl overridable) [`Self::float_powi_scalar_impl`]
1029 /// operation to handle the generic case.
1030 ///
1031 /// # Arguments
1032 ///
1033 /// * `lhs` - The left-hand side tensor.
1034 /// * `rhs` - The right-hand side scalar.
1035 ///
1036 /// # Returns
1037 ///
1038 /// The elements of `lhs` raised to the value of `rhs`.
1039 fn float_powi_scalar(lhs: FloatTensor<B>, rhs: Scalar) -> FloatTensor<B> {
1040 match rhs.elem::<i64>() {
1041 0 => Self::float_ones(lhs.shape(), &B::float_device(&lhs), lhs.dtype().into()),
1042 1 => lhs,
1043 2 => B::float_mul(lhs.clone(), lhs),
1044 -1 => Self::float_recip(lhs),
1045 -2 => Self::float_recip(B::float_mul(lhs.clone(), lhs)),
1046 _ => Self::float_powi_scalar_impl(lhs, rhs),
1047 }
1048 }
1049
1050 /// Raises a tensor to the power of an int scalar.
1051 ///
1052 /// # Backend Implementors Note
1053 ///
1054 /// This is the generic implementation of integer exponentiation
1055 /// called by [`Self::float_powi_scalar`] in the fallback case.
1056 ///
1057 /// As a general rule, this should not be called directly.
1058 ///
1059 /// # Arguments
1060 ///
1061 /// * `lhs` - The left-hand side tensor.
1062 /// * `rhs` - The right-hand side scalar.
1063 ///
1064 /// # Returns
1065 ///
1066 /// The elements of `lhs` raised to the value of `rhs`.
1067 fn float_powi_scalar_impl(lhs: FloatTensor<B>, rhs: Scalar) -> FloatTensor<B> {
1068 // Avoid a recursive loop by deferring directly to float_powf_scalar_impl.
1069 Self::float_powf_scalar_impl(lhs, rhs)
1070 }
1071
1072 /// Returns a new tensor with values raised to the power of float `value`.
1073 ///
1074 /// # Backend Implementors Note
1075 ///
1076 /// This (`Backend` impl overridable) operation dispatches integer exponentiation
1077 /// to [`Self::float_powi_scalar`], and the remaining non-integer exponent cases to
1078 /// the (`Backend` impl overridable) [`Self::float_powf_scalar_impl`]
1079 /// operation to handle the generic case.
1080 ///
1081 /// # Arguments
1082 ///
1083 /// * `tensor` - The tensor to exponentiate.
1084 /// * `value` - The exponent.
1085 ///
1086 /// # Returns
1087 ///
1088 /// A tensor with the same shape as `tensor` with values raised to the power of `value`.
1089 fn float_powf_scalar(tensor: FloatTensor<B>, value: Scalar) -> FloatTensor<B> {
1090 if let Some(exp) = value.try_as_integer() {
1091 Self::float_powi_scalar(tensor, exp)
1092 } else {
1093 Self::float_powf_scalar_impl(tensor, value)
1094 }
1095 }
1096
1097 /// Returns a new tensor with values raised to the power of float `value`.
1098 ///
1099 /// # Backend Implementors Note
1100 ///
1101 /// This is the generic implementation of integer exponentiation
1102 /// called by [`Self::float_powf_scalar`] in the fallback case.
1103 ///
1104 /// This is the minimal required support a `Backend` must implement
1105 /// for exponentiation.
1106 ///
1107 /// As a general rule, this should not be called directly.
1108 ///
1109 /// # Arguments
1110 ///
1111 /// * `tensor` - The tensor to exponentiate.
1112 /// * `value` - The exponent.
1113 ///
1114 /// # Returns
1115 ///
1116 /// A tensor with the same shape as `tensor` with values raised to the power of `value`.
1117 fn float_powf_scalar_impl(tensor: FloatTensor<B>, value: Scalar) -> FloatTensor<B>;
1118
1119 /// Returns a new tensor with square root values.
1120 ///
1121 /// # Arguments
1122 ///
1123 /// * `tensor` - The tensor to take the square root of.
1124 ///
1125 /// # Returns
1126 ///
1127 /// A tensor with the same shape as `tensor` with square root values.
1128 fn float_sqrt(tensor: FloatTensor<B>) -> FloatTensor<B>;
1129
1130 /// Returns element-wise reciprocal square roots. Backends may provide a native instruction.
1131 fn float_rsqrt(tensor: FloatTensor<B>) -> FloatTensor<B> {
1132 Self::float_recip(Self::float_sqrt(tensor))
1133 }
1134
1135 /// Returns a new tensor with absolute values.
1136 ///
1137 /// # Arguments
1138 ///
1139 /// * `tensor` - The tensor to take absolute value of.
1140 ///
1141 /// # Returns
1142 ///
1143 /// A tensor with the same shape as `tensor` with absolute values.
1144 fn float_abs(tensor: FloatTensor<B>) -> FloatTensor<B>;
1145
1146 /// Returns a new tensor with cosine values.
1147 ///
1148 /// # Arguments
1149 ///
1150 /// * `tensor` - The tensor to take the cosine of.
1151 ///
1152 /// # Returns
1153 ///
1154 /// A tensor with the same shape as `tensor` with cosine values.
1155 fn float_cos(tensor: FloatTensor<B>) -> FloatTensor<B>;
1156
1157 /// Returns a new tensor with sine values.
1158 ///
1159 /// # Arguments
1160 ///
1161 /// * `tensor` - The tensor to take the sine of.
1162 ///
1163 /// # Returns
1164 ///
1165 /// A tensor with the same shape as `tensor` with sine values.
1166 fn float_sin(tensor: FloatTensor<B>) -> FloatTensor<B>;
1167
1168 /// Returns a new tensor with tangent values.
1169 ///
1170 /// # Arguments
1171 ///
1172 /// * `tensor` - The tensor to take the tangent of.
1173 ///
1174 /// # Returns
1175 ///
1176 /// A tensor with the same shape as `tensor` with tangent values.
1177 fn float_tan(tensor: FloatTensor<B>) -> FloatTensor<B>;
1178
1179 /// Returns a new tensor with hyperbolic cosine values.
1180 ///
1181 /// # Arguments
1182 ///
1183 /// * `tensor` - The tensor to take the hyperbolic cosine of.
1184 ///
1185 /// # Returns
1186 ///
1187 /// A tensor with the same shape as `tensor` with hyperbolic cosine values.
1188 fn float_cosh(tensor: FloatTensor<B>) -> FloatTensor<B>;
1189
1190 /// Returns a new tensor with hyperbolic sine values.
1191 ///
1192 /// # Arguments
1193 ///
1194 /// * `tensor` - The tensor to take the hyperbolic sine of.
1195 ///
1196 /// # Returns
1197 ///
1198 /// A tensor with the same shape as `tensor` with hyperbolic sine values.
1199 fn float_sinh(tensor: FloatTensor<B>) -> FloatTensor<B>;
1200
1201 /// Returns a new tensor with hyperbolic tangent values.
1202 ///
1203 /// # Arguments
1204 ///
1205 /// * `tensor` - The tensor to take the hyperbolic tangent of.
1206 ///
1207 /// # Returns
1208 ///
1209 /// A tensor with the same shape as `tensor` with hyperbolic tangent values.
1210 fn float_tanh(tensor: FloatTensor<B>) -> FloatTensor<B>;
1211
1212 /// Returns a new tensor with inverse cosine values.
1213 ///
1214 /// # Arguments
1215 ///
1216 /// * `tensor` - The input tensor.
1217 ///
1218 /// # Returns
1219 ///
1220 /// A tensor with the same shape as `tensor` with inverse cosine values.
1221 fn float_acos(tensor: FloatTensor<B>) -> FloatTensor<B>;
1222
1223 /// Returns a new tensor with inverse hyperbolic cosine values.
1224 ///
1225 /// # Arguments
1226 ///
1227 /// * `tensor` - The input tensor.
1228 ///
1229 /// # Returns
1230 ///
1231 /// A tensor with the same shape as `tensor` with inverse hyperbolic cosine values.
1232 fn float_acosh(tensor: FloatTensor<B>) -> FloatTensor<B>;
1233
1234 /// Returns a new tensor with inverse sine values.
1235 ///
1236 /// # Arguments
1237 ///
1238 /// * `tensor` - The input tensor.
1239 ///
1240 /// # Returns
1241 ///
1242 /// A tensor with the same shape as `tensor` with inverse sine values.
1243 fn float_asin(tensor: FloatTensor<B>) -> FloatTensor<B>;
1244
1245 /// Returns a new tensor with inverse hyperbolic sine values.
1246 ///
1247 /// # Arguments
1248 ///
1249 /// * `tensor` - The input tensor.
1250 ///
1251 /// # Returns
1252 ///
1253 /// A tensor with the same shape as `tensor` with inverse hyperbolic sine values.
1254 fn float_asinh(tensor: FloatTensor<B>) -> FloatTensor<B>;
1255
1256 /// Returns a new tensor with the inverse tangent values.
1257 ///
1258 /// # Arguments
1259 ///
1260 /// * `tensor` - The input tensor.
1261 ///
1262 /// # Returns
1263 ///
1264 /// A tensor with the same shape as `tensor` with the inverse tangent values.
1265 fn float_atan(tensor: FloatTensor<B>) -> FloatTensor<B>;
1266
1267 /// Returns a new tensor with the inverse hyperbolic tangent values.
1268 ///
1269 /// # Arguments
1270 ///
1271 /// * `tensor` - The input tensor.
1272 ///
1273 /// # Returns
1274 ///
1275 /// A tensor with the same shape as `tensor` with the inverse hyperbolic tangent values.
1276 fn float_atanh(tensor: FloatTensor<B>) -> FloatTensor<B>;
1277
1278 /// Returns a tensor with the four-quadrant inverse tangent values of `y` and `x`.
1279 ///
1280 /// # Arguments
1281 ///
1282 /// * `lhs` - The tensor with y coordinates.
1283 /// * `rhs` - The tensor with x coordinates.
1284 ///
1285 /// # Returns
1286 ///
1287 /// A tensor with the four-quadrant inverse tangent values.
1288 fn float_atan2(lhs: FloatTensor<B>, rhs: FloatTensor<B>) -> FloatTensor<B>;
1289
1290 /// Returns a new tensor with rounded values.
1291 ///
1292 /// This function should implement the [round half to even](https://en.wikipedia.org/wiki/Rounding#Rounding_half_to_even)
1293 /// strategy, with halfway cases rounded to the nearest even integer value.
1294 ///
1295 /// # Arguments
1296 ///
1297 /// * `tensor` - The tensor to be rounded.
1298 ///
1299 /// # Returns
1300 ///
1301 /// A tensor with the same shape as `tensor` with rounded values.
1302 fn float_round(tensor: FloatTensor<B>) -> FloatTensor<B>;
1303
1304 /// Returns a new tensor with floored values.
1305 ///
1306 /// # Arguments
1307 ///
1308 /// * `tensor` - The tensor to be floored.
1309 ///
1310 /// # Returns
1311 ///
1312 /// A tensor with the same shape as `tensor` with floored values.
1313 fn float_floor(tensor: FloatTensor<B>) -> FloatTensor<B>;
1314
1315 /// Returns a new tensor with ceiled values.
1316 ///
1317 /// # Arguments
1318 ///
1319 /// * `tensor` - The tensor to be ceiled.
1320 ///
1321 /// # Returns
1322 ///
1323 /// A tensor with the same shape as `tensor` with ceiled values.
1324 fn float_ceil(tensor: FloatTensor<B>) -> FloatTensor<B>;
1325
1326 /// Returns a new tensor with truncated values.
1327 ///
1328 /// # Arguments
1329 ///
1330 /// * `tensor` - The tensor to be truncated.
1331 ///
1332 /// # Returns
1333 ///
1334 /// A tensor with the same shape as `tensor` with truncated values.
1335 fn float_trunc(tensor: FloatTensor<B>) -> FloatTensor<B>;
1336
1337 /// Returns a new tensor with the error function values.
1338 ///
1339 /// # Arguments
1340 ///
1341 /// * `tensor` - The tensor to take the error function of.
1342 ///
1343 /// # Returns
1344 ///
1345 /// A tensor with the same shape as `tensor` with error function values.
1346 fn float_erf(tensor: FloatTensor<B>) -> FloatTensor<B>;
1347
1348 /// Concatenates tensors along a dimension.
1349 ///
1350 /// # Arguments
1351 ///
1352 /// * `tensors` - The tensors to concatenate.
1353 /// * `dim` - The dimension along which to concatenate.
1354 ///
1355 /// # Returns
1356 ///
1357 /// A tensor with the concatenated tensors along `dim`.
1358 ///
1359 /// # Note
1360 ///
1361 /// Empty tensors (where the concatenation dimension has size 0) are filtered out at the
1362 /// high-level tensor API and will not be passed to this method. Backend implementations do
1363 /// not need to handle empty tensors.
1364 fn float_cat(tensors: Vec<FloatTensor<B>>, dim: usize) -> FloatTensor<B> {
1365 cat_with_slice_assign::<B, Float>(
1366 tensors.into_iter().map(TensorPrimitive::Float).collect(),
1367 dim,
1368 )
1369 .tensor()
1370 }
1371
1372 /// Gets the indices of the maximum elements of a tensor along an axis.
1373 ///
1374 /// # Arguments
1375 ///
1376 /// * `tensor` - The tensor to get the maximum elements of.
1377 /// * `dim` - The dimension along which to get the maximum elements.
1378 /// * `out_dtype` - The output tensor dtype.
1379 ///
1380 /// # Returns
1381 ///
1382 /// A tensor with the indices of the maximum elements of `tensor` along `dim`.
1383 fn float_argmax(tensor: FloatTensor<B>, dim: usize, out_dtype: IntDType) -> IntTensor<B>;
1384
1385 /// Gets the indices of the k maximum elements of a tensor along an axis.
1386 /// if two elements are equals, it will be ordered by lowest indices
1387 ///
1388 /// # Arguments
1389 ///
1390 /// * `tensor` - The tensor to get the maximum elements of.
1391 /// * `dim` - The dimension along which to get the maximum elements.
1392 /// * `k` - number of maximum elements
1393 /// * `out_dtype` - The output tensor dtype.
1394 ///
1395 /// # Returns
1396 ///
1397 /// A tensor with the indices of the maximum elements of `tensor` along `dim`.
1398 fn float_argtopk(
1399 tensor: FloatTensor<B>,
1400 dim: usize,
1401 k: usize,
1402 out_dtype: IntDType,
1403 ) -> IntTensor<B>;
1404
1405 /// Gets the values of the k maximum elements of a tensor along an axis.
1406 ///
1407 /// # Arguments
1408 ///
1409 /// * `tensor` - The tensor to get the maximum elements of.
1410 /// * `dim` - The dimension along which to get the maximum elements.
1411 /// * `k` - number of maximum elements
1412 /// * `out_dtype` - The output tensor dtype.
1413 ///
1414 /// # Returns
1415 ///
1416 /// A tensor with the values of the maximum elements of `tensor` along `dim`.
1417 fn float_topk(tensor: FloatTensor<B>, dim: usize, k: usize) -> FloatTensor<B> {
1418 let device = Self::float_device(&tensor);
1419 let dtype = get_device_settings::<B>(&device).int_dtype;
1420 let k_indices = B::int_arange(0..k as i64, &device, dtype);
1421 Self::float_select(Self::float_sort(tensor, dim, true), dim, k_indices)
1422 }
1423
1424 /// Gets the indices of the minimum elements of a tensor along an axis.
1425 ///
1426 /// # Arguments
1427 ///
1428 /// * `tensor` - The tensor to get the minimum elements of.
1429 /// * `dim` - The dimension along which to get the minimum elements.
1430 /// * `out_dtype` - The output tensor dtype.
1431 ///
1432 /// # Returns
1433 ///
1434 /// A tensor with the indices of the minimum elements of `tensor` along `dim`.
1435 fn float_argmin(tensor: FloatTensor<B>, dim: usize, out_dtype: IntDType) -> IntTensor<B>;
1436
1437 /// Gets the maximum element of a tensor.
1438 ///
1439 /// # Arguments
1440 ///
1441 /// * `tensor` - The tensor to get the maximum elements of.
1442 ///
1443 /// # Returns
1444 ///
1445 /// A tensor with the maximum element of `tensor`.
1446 fn float_max(tensor: FloatTensor<B>) -> FloatTensor<B> {
1447 let shape = tensor.shape();
1448 let tensor = B::float_reshape(tensor, Shape::new([shape.num_elements()]));
1449
1450 B::float_max_dim(tensor, 0)
1451 }
1452
1453 /// Gets the maximum elements of a tensor along an axis.
1454 ///
1455 /// # Arguments
1456 ///
1457 /// * `tensor` - The tensor to get the maximum elements of.
1458 /// * `dim` - The dimension along which to get the maximum elements.
1459 ///
1460 /// # Returns
1461 ///
1462 /// A tensor with the maximum elements of `tensor` along `dim`.
1463 fn float_max_dim(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B> {
1464 let dtype = get_device_settings::<B>(&B::float_device(&tensor)).int_dtype;
1465 let index = B::float_argmax(tensor.clone(), dim, dtype);
1466
1467 B::float_gather(dim, tensor, index)
1468 }
1469
1470 /// Gets the maximum elements of a tensor along an axis and their indices.
1471 ///
1472 /// # Arguments
1473 ///
1474 /// * `tensor` - The tensor to get the maximum elements of.
1475 /// * `dim` - The dimension along which to get the maximum elements.
1476 /// * `indices_dtype` - The indices tensor dtype.
1477 ///
1478 /// # Returns
1479 ///
1480 /// A tuple with the maximum elements of `tensor` along `dim` and their indices.
1481 fn float_max_dim_with_indices(
1482 tensor: FloatTensor<B>,
1483 dim: usize,
1484 indices_dtype: IntDType,
1485 ) -> (FloatTensor<B>, IntTensor<B>) {
1486 let index = B::float_argmax(tensor.clone(), dim, indices_dtype);
1487 let values = B::float_gather(dim, tensor, index.clone());
1488
1489 (values, index)
1490 }
1491
1492 /// Gets the minimum element of a tensor.
1493 ///
1494 /// # Arguments
1495 ///
1496 /// * `tensor` - The tensor to get the minimum elements of.
1497 ///
1498 /// # Returns
1499 ///
1500 /// A tensor with the minimum element of `tensor`.
1501 fn float_min(tensor: FloatTensor<B>) -> FloatTensor<B> {
1502 let shape = tensor.shape();
1503 let tensor = B::float_reshape(tensor, Shape::new([shape.num_elements()]));
1504
1505 B::float_min_dim(tensor, 0)
1506 }
1507
1508 /// Gets the minimum elements of a tensor along an axis.
1509 ///
1510 /// # Arguments
1511 ///
1512 /// * `tensor` - The tensor to get the minimum elements of.
1513 /// * `dim` - The dimension along which to get the minimum elements.
1514 ///
1515 /// # Returns
1516 ///
1517 /// A tensor with the minimum elements of `tensor` along `dim`.
1518 fn float_min_dim(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B> {
1519 let dtype = get_device_settings::<B>(&B::float_device(&tensor)).int_dtype;
1520 let index = B::float_argmin(tensor.clone(), dim, dtype);
1521
1522 B::float_gather(dim, tensor, index)
1523 }
1524
1525 /// Gets the minimum elements of a tensor along an axis and their indices.
1526 ///
1527 /// # Arguments
1528 ///
1529 /// * `tensor` - The tensor to get the minimum elements of.
1530 /// * `dim` - The dimension along which to get the minimum elements.
1531 /// * `indices_dtype` - The indices tensor dtype.
1532 ///
1533 /// # Returns
1534 ///
1535 /// A tuple with the minimum elements of `tensor` along `dim` and their indices.
1536 fn float_min_dim_with_indices(
1537 tensor: FloatTensor<B>,
1538 dim: usize,
1539 indices_dtype: IntDType,
1540 ) -> (FloatTensor<B>, IntTensor<B>) {
1541 let index = B::float_argmin(tensor.clone(), dim, indices_dtype);
1542 let values = B::float_gather(dim, tensor, index.clone());
1543
1544 (values, index)
1545 }
1546
1547 /// Gets the maximum absolute element of a tensor.
1548 ///
1549 /// # Arguments
1550 ///
1551 /// * `tensor` - The tensor to get the maximum elements of.
1552 ///
1553 /// # Returns
1554 ///
1555 /// A tensor with the maximum element of `tensor`.
1556 fn float_max_abs(tensor: FloatTensor<B>) -> FloatTensor<B> {
1557 let shape = tensor.shape();
1558 let tensor = B::float_reshape(tensor, Shape::new([shape.num_elements()]));
1559
1560 B::float_max_abs_dim(tensor, 0)
1561 }
1562
1563 /// Gets the maximum absolute elements of a tensor along an axis.
1564 ///
1565 /// # Arguments
1566 ///
1567 /// * `tensor` - The tensor to get the maximum elements of.
1568 /// * `dim` - The dimension along which to get the maximum elements.
1569 ///
1570 /// # Returns
1571 ///
1572 /// A tensor with the maximum elements of `tensor` along `dim`.
1573 fn float_max_abs_dim(tensor: FloatTensor<B>, dim: usize) -> FloatTensor<B> {
1574 B::float_max_dim(B::float_abs(tensor), dim)
1575 }
1576
1577 /// Tests if any element in the float `tensor` evaluates to True.
1578 ///
1579 /// # Arguments
1580 ///
1581 /// * `tensor` - The tensor to test.
1582 /// * `out_dtype` - The output tensor dtype.
1583 ///
1584 /// # Returns
1585 ///
1586 /// A boolean tensor with a single element, True if any element in the tensor is True, False otherwise.
1587 fn float_any(tensor: FloatTensor<B>, out_dtype: BoolDType) -> BoolTensor<B> {
1588 let float_dtype = tensor.dtype();
1589 let bool_tensor = B::float_equal_elem(tensor, 0f32.into(), out_dtype);
1590 let bool_tensor = B::bool_not(bool_tensor);
1591 let sum = B::float_sum(B::bool_into_float(bool_tensor, float_dtype.into()));
1592 B::float_greater_elem(sum, 0f32.into(), out_dtype)
1593 }
1594
1595 /// Tests if any element in the float `tensor` evaluates to True along a given dimension `dim`.
1596 ///
1597 /// # Arguments
1598 ///
1599 /// * `tensor` - The tensor to test.
1600 /// * `dim` - The axis along which to test.
1601 /// * `out_dtype` - The output tensor dtype.
1602 ///
1603 /// # Returns
1604 ///
1605 /// A boolean tensor `Tensor<B, D, Bool>` with the same size as input `tensor`, except in the `dim` axis
1606 /// where the size is 1. The elem in the `dim` axis is True if any element along this dim in the
1607 /// input evaluates to True, False otherwise.
1608 fn float_any_dim(tensor: FloatTensor<B>, dim: usize, out_dtype: BoolDType) -> BoolTensor<B> {
1609 let float_dtype = tensor.dtype();
1610 let bool_tensor = B::float_equal_elem(tensor, 0f32.into(), out_dtype);
1611 let bool_tensor = B::bool_not(bool_tensor);
1612 let sum = B::float_sum_dim(B::bool_into_float(bool_tensor, float_dtype.into()), dim);
1613 B::float_greater_elem(sum, 0f32.into(), out_dtype)
1614 }
1615
1616 /// Tests if all elements in the float `tensor` evaluate to True.
1617 ///
1618 /// # Arguments
1619 ///
1620 /// * `tensor` - The tensor to test.
1621 /// * `out_dtype` - The output tensor dtype.
1622 ///
1623 /// # Returns
1624 ///
1625 /// A boolean tensor `Tensor<B, 1, Bool>` with a single element, True if all elements in the input tensor
1626 /// evaluate to True, False otherwise.
1627 fn float_all(tensor: FloatTensor<B>, out_dtype: BoolDType) -> BoolTensor<B> {
1628 let float_dtype = tensor.dtype();
1629 let num_elems = tensor.shape().num_elements() as f32;
1630 let bool_tensor = B::float_equal_elem(tensor, 0f32.into(), out_dtype);
1631 let bool_tensor = B::bool_not(bool_tensor);
1632 let sum = B::float_sum(B::bool_into_float(bool_tensor, float_dtype.into()));
1633 B::float_equal_elem(sum, num_elems.into(), out_dtype)
1634 }
1635
1636 /// Tests if all elements in the float `tensor` evaluate to True along a given dimension `dim`.
1637 ///
1638 /// # Arguments
1639 ///
1640 /// * `tensor` - The tensor to test.
1641 /// * `dim` - The axis along which to test.
1642 /// * `out_dtype` - The output tensor dtype.
1643 ///
1644 /// # Returns
1645 ///
1646 /// A boolean tensor `Tensor<B, D, Bool>` with the same size as input `tensor`, except in the `dim` axis
1647 /// where the size is 1. The elem in the `dim` axis is True if all elements along this dim in the input
1648 /// evaluates to True, False otherwise.
1649 fn float_all_dim(tensor: FloatTensor<B>, dim: usize, out_dtype: BoolDType) -> BoolTensor<B> {
1650 let float_dtype = tensor.dtype();
1651 let num_elems = tensor.shape()[dim] as f32;
1652 let bool_tensor = B::float_equal_elem(tensor, 0f32.into(), out_dtype);
1653 let bool_tensor = B::bool_not(bool_tensor);
1654 let sum = B::float_sum_dim(B::bool_into_float(bool_tensor, float_dtype.into()), dim);
1655 B::float_equal_elem(sum, num_elems.into(), out_dtype)
1656 }
1657
1658 /// Returns the signs of the float `tensor`.
1659 ///
1660 /// # Arguments
1661 ///
1662 /// * `tensor` - The tensor to extract the signs from.
1663 ///
1664 /// # Returns
1665 ///
1666 /// A tensor with the same shape as `tensor` containing the signs of the elements of `tensor`.
1667 fn float_sign(tensor: FloatTensor<B>) -> FloatTensor<B> {
1668 let device = B::float_device(&tensor);
1669 let bool_dtype = get_device_settings::<B>(&B::float_device(&tensor)).bool_dtype;
1670 let zeros = B::float_zeros(tensor.shape(), &device, tensor.dtype().into());
1671 let less_than_zero = B::float_lower_elem(tensor.clone(), 0f32.into(), bool_dtype);
1672 let greater_than_zero = B::float_greater_elem(tensor, 0f32.into(), bool_dtype);
1673
1674 let mut result = B::float_mask_fill(zeros, less_than_zero, (-1f32).into());
1675 result = B::float_mask_fill(result, greater_than_zero, 1f32.into());
1676 result
1677 }
1678
1679 /// Broadcasts the float `tensor` to the given `shape`.
1680 fn float_expand(tensor: FloatTensor<B>, shape: Shape) -> FloatTensor<B>;
1681
1682 /// Sort the elements of the input `tensor` by value in along a given dimension.
1683 ///
1684 /// This sort is unstable (i.e., may reorder equal elements).
1685 ///
1686 /// # Arguments
1687 ///
1688 /// * `tensor` - The input tensor.
1689 /// * `dim` - The axis along which to sort.
1690 /// * `descending` - The sorting order.
1691 ///
1692 /// # Returns
1693 ///
1694 /// A tensor with the same shape as the input tensor, where the elements are sorted by value.
1695 fn float_sort(tensor: FloatTensor<B>, dim: usize, descending: bool) -> FloatTensor<B> {
1696 sort::<B, Float>(TensorPrimitive::Float(tensor), dim, descending).tensor()
1697 }
1698
1699 /// Sort the elements of the input `tensor` by value in along a given dimension.
1700 ///
1701 /// This sort is unstable (i.e., may reorder equal elements).
1702 ///
1703 /// # Arguments
1704 ///
1705 /// * `tensor` - The input tensor.
1706 /// * `dim` - The axis along which to sort.
1707 /// * `descending` - The sorting order.
1708 /// * `indices_dtype` - The indices tensor dtype.
1709 ///
1710 /// # Returns
1711 ///
1712 /// A tensor with the same shape as the input tensor and corresponding indices, where
1713 /// the elements are sorted by value and the indices map back to the original input tensor.
1714 fn float_sort_with_indices(
1715 tensor: FloatTensor<B>,
1716 dim: usize,
1717 descending: bool,
1718 indices_dtype: IntDType,
1719 ) -> (FloatTensor<B>, IntTensor<B>) {
1720 let (values, indices) = sort_with_indices::<B, Float>(
1721 TensorPrimitive::Float(tensor),
1722 dim,
1723 descending,
1724 indices_dtype,
1725 );
1726 (values.tensor(), indices)
1727 }
1728
1729 /// Returns the indices that sort the elements of the input `tensor` by value along a given dimension.
1730 ///
1731 /// This sort is unstable (i.e., may reorder equal elements).
1732 ///
1733 /// # Arguments
1734 ///
1735 /// * `tensor` - The input tensor.
1736 /// * `dim` - The axis along which to sort.
1737 /// * `descending` - The sorting order.
1738 /// * `out_dtype` - The output tensor dtype.
1739 ///
1740 /// # Returns
1741 ///
1742 /// A tensor with the same shape as the input tensor the indices map back to the original input tensor.
1743 fn float_argsort(
1744 tensor: FloatTensor<B>,
1745 dim: usize,
1746 descending: bool,
1747 out_dtype: IntDType,
1748 ) -> IntTensor<B> {
1749 argsort::<B, Float>(TensorPrimitive::Float(tensor), dim, descending, out_dtype)
1750 }
1751
1752 /// Samples tensor as a two-dimensional spatial grid of (possibly multi-channel) values,
1753 /// using the given locations in [-1, 1].
1754 ///
1755 /// # Arguments
1756 ///
1757 /// * `tensor` - The tensor being sampled from, must be contiguous with shape (N, C, H_in, W_in)
1758 /// * `grid` - A tensor of locations, with shape (N, H_out, W_out, 2). Values are [-1, 1].
1759 /// A [x = -1, y = -1] means top-left, and [x = 1, y = 1] means bottom-right
1760 /// * `options` - Grid sampling options (mode, padding_mode, align_corners)
1761 ///
1762 /// # Returns
1763 ///
1764 /// A tensor with shape (N, C, H_out, W_out)
1765 fn float_grid_sample_2d(
1766 tensor: FloatTensor<B>,
1767 grid: FloatTensor<B>,
1768 options: GridSampleOptions,
1769 ) -> FloatTensor<B> {
1770 // TODO: default impl should get int default dtype
1771 float_grid_sample_2d_ref::<B>(tensor, grid, options)
1772 }
1773
1774 /// Unfold windows along a dimension.
1775 ///
1776 /// Returns a view of the tensor with all complete windows of size `size` in dimension `dim`;
1777 /// where windows are advanced by `step` at each index.
1778 ///
1779 /// The number of windows is `max(0, (shape[dim] - size).ceil_div(step))`.
1780 ///
1781 /// # Arguments
1782 ///
1783 /// * `tensor` - The input tensor to unfold; of shape ``[pre=..., dim shape, post=...]``
1784 /// * `dim` - the selected dim.
1785 /// * `size` - the size of each unfolded window.
1786 /// * `step` - the step between each window.
1787 ///
1788 /// # Returns
1789 ///
1790 /// A tensor view with shape ``[pre=..., windows, size, post=...]``.
1791 fn float_unfold(tensor: FloatTensor<B>, dim: usize, size: usize, step: usize)
1792 -> FloatTensor<B>;
1793
1794 /// Returns a new tensor with boolean elements indicating whether each element of the input is NaN.
1795 ///
1796 /// # Returns
1797 ///
1798 /// A boolean tensor where `true` indicates NaN and `false` indicates a non-NaN value.
1799 fn float_is_nan(tensor: FloatTensor<B>, out_dtype: BoolDType) -> BoolTensor<B> {
1800 // Check if the input tensor is NaN by comparing it to itself
1801 // NaN is the only value that is not equal to itself
1802 B::float_not_equal(tensor.clone(), tensor, out_dtype)
1803 }
1804
1805 /// Returns a new tensor with boolean elements indicating whether each element of the input is infinite (either +INF or -INF).
1806 ///
1807 /// # Returns
1808 ///
1809 /// A boolean tensor where `true` indicates that the value is infinite
1810 fn float_is_inf(tensor: FloatTensor<B>, out_dtype: BoolDType) -> BoolTensor<B> {
1811 B::float_equal_elem(B::float_abs(tensor), f64::INFINITY.into(), out_dtype)
1812 }
1813}