Skip to main content

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}