Skip to main content

ruda_tensor/api/
float.rs

1use crate::api::AsIndex;
2use crate::api::Cast;
3use crate::api::Tensor;
4use crate::api::cast::ToElement;
5use crate::api::check;
6use crate::api::check::TensorCheck;
7use crate::api::ops::GridSampleOptions;
8use crate::api::quantization::{QuantScheme, QuantizationParameters};
9use crate::api::backend::Backend;
10use crate::api::stats;
11use crate::api::{Distribution, TensorData};
12use crate::api::{Bool, Float, Int, TensorPrimitive};
13#[cfg(feature = "api-distributed")]
14use crate::AutodiffBackend;
15use crate::ElementConversion;
16use crate::Scalar;
17use crate::TensorMetadata;
18#[cfg(feature = "api-distributed")]
19use crate::distributed::DistributedParamId;
20use crate::get_device_settings;
21use crate::tensor::FloatMathOps;
22use crate::tensor::quantization::QuantizationParametersPrimitive;
23use core::f32;
24
25/// Default RTOL value for `is_close` and `all_close`.
26pub const DEFAULT_RTOL: f64 = 1e-5;
27
28/// Default ATOL value for `is_close` and `all_close`.
29pub const DEFAULT_ATOL: f64 = 1e-8;
30
31impl<const D: usize, B> Tensor<B, D>
32where
33    B: Backend,
34{
35    /// Applies the [error function](https://en.wikipedia.org/wiki/Error_function) element wise.
36    ///
37    #[cfg_attr(
38        doc,
39        doc = r#"
40$y_i = \text{erf}\(x_i\)$
41
42The error function is defined as:
43
44$$\text{erf}\(x\) = \frac{2}{\sqrt{\pi}} \int_0^x e^{-t^2} dt$$
45"#
46    )]
47    #[cfg_attr(not(doc), doc = "`y_i = erf(x_i)`")]
48    pub fn erf(self) -> Self {
49        Self::new(TensorPrimitive::Float(B::float_erf(
50            self.primitive.tensor(),
51        )))
52    }
53
54    /// Applies [reciprocal operation](https://en.wikipedia.org/wiki/Multiplicative_inverse)
55    /// (or multiplicative inverse) element wise.
56    ///
57    #[cfg_attr(doc, doc = r#"$y_i = \frac{1}{x_i}$"#)]
58    #[cfg_attr(not(doc), doc = "`y_i = 1/x_i`")]
59    pub fn recip(self) -> Self {
60        Self::new(TensorPrimitive::Float(B::float_recip(
61            self.primitive.tensor(),
62        )))
63    }
64
65    /// Applies the reciprocal square root element-wise, preserving shape and dtype.
66    pub fn rsqrt(self) -> Self {
67        Self::new(TensorPrimitive::Float(B::float_rsqrt(self.primitive.tensor())))
68    }
69
70    /// Converts each of the elements of the input tensor from angles in degrees to radians.
71    ///
72    /// # Example
73    /// ```ignore
74    /// let tensor_in_radians = tensor.deg2rad();
75    /// ```
76    pub fn deg2rad(self) -> Self {
77        self.mul_scalar(f32::consts::PI / 180.0)
78    }
79
80    /// Converts each of the elements of the input tensor from angles in radians to degrees.
81    ///
82    /// # Example
83    /// ```ignore
84    /// let tensor_in_degrees = tensor.rad2deg();
85    /// ```
86    pub fn rad2deg(self) -> Self {
87        self.mul_scalar(180.0 / f32::consts::PI)
88    }
89
90    /// Applies element wise round operation.
91    ///
92    /// This function implements the [round half to even](https://en.wikipedia.org/wiki/Rounding#Rounding_half_to_even)
93    /// strategy, with halfway cases rounded to the nearest even integer value.
94    pub fn round(self) -> Self {
95        Self::new(TensorPrimitive::Float(B::float_round(
96            self.primitive.tensor(),
97        )))
98    }
99
100    /// Applies element wise floor operation.
101    pub fn floor(self) -> Self {
102        Self::new(TensorPrimitive::Float(B::float_floor(
103            self.primitive.tensor(),
104        )))
105    }
106
107    /// Applies element wise ceil operation.
108    pub fn ceil(self) -> Self {
109        Self::new(TensorPrimitive::Float(B::float_ceil(
110            self.primitive.tensor(),
111        )))
112    }
113
114    /// Create a tensor from floats (f32) on a given device.
115    ///
116    /// # Example
117    ///
118    /// ```rust
119    /// use ruda_tensor::api::backend::Backend;
120    /// use ruda_tensor::api::Tensor;
121    ///
122    /// fn example<B: Backend>() {
123    ///     let device = B::Device::default();
124    ///     let _ = Tensor::<B, 1>::from_floats([1.0, 2.0], &device);
125    ///     let _ = Tensor::<B, 2>::from_floats([[1.0, 2.0], [3.0, 4.0]], &device);
126    /// }
127    /// ```
128    pub fn from_floats<A: Into<TensorData>>(floats: A, device: &B::Device) -> Self {
129        Self::from_data(floats.into().convert::<f32>(), device)
130    }
131
132    /// Returns a new tensor with the same shape and device as the current tensor and the data
133    /// cast to Integer.
134    ///
135    /// # Example
136    ///
137    /// ```rust
138    /// use ruda_tensor::api::backend::Backend;
139    /// use ruda_tensor::api::Tensor;
140    ///
141    /// fn example<B: Backend>() {
142    ///     let device = Default::default();
143    ///     let float_tensor = Tensor::<B, 1>::from_floats([1.0, 2.0], &device);
144    ///     let int_tensor = float_tensor.int();
145    /// }
146    /// ```
147    pub fn int(self) -> Tensor<B, D, Int> {
148        let out_dtype = get_device_settings::<B>(&self.device()).int_dtype;
149        Tensor::new(B::float_into_int(self.primitive.tensor(), out_dtype))
150    }
151
152    /// Returns a new tensor with the same shape, dtype, and device as the current tensor filled random
153    /// values sampled from the given distribution.
154    pub fn random_like(&self, distribution: Distribution) -> Self {
155        Self::new(TensorPrimitive::Float(B::float_random(
156            self.shape(),
157            distribution,
158            &self.device(),
159            self.dtype().into(),
160        )))
161    }
162
163    /// Calculate the variance along the given dimension.
164    pub fn var(self, dim: usize) -> Self {
165        stats::var(self, dim)
166    }
167
168    /// Calculate the variance along the given dimension without applying the Bessel’s correction.
169    pub fn var_bias(self, dim: usize) -> Self {
170        stats::var_bias(self, dim)
171    }
172
173    /// Calculate the variance along the given dimension and also returns the mean.
174    pub fn var_mean(self, dim: usize) -> (Self, Self) {
175        let mean = self.clone().mean_dim(dim);
176        let var = stats::var_with_mean(self, mean.clone(), dim);
177        (var, mean)
178    }
179
180    /// Calculate the variance along the given dimension without applying the Bessel’s correction and also returns the mean.
181    pub fn var_mean_bias(self, dim: usize) -> (Self, Self) {
182        let mean = self.clone().mean_dim(dim);
183        let var = stats::var_with_mean_bias(self, mean.clone(), dim);
184        (var, mean)
185    }
186
187    /// Returns the median value along the specified dimension.
188    ///
189    /// The median is not unique for input tensors with an even number of elements
190    /// in the reduced dimension. In this case, the lower of the two medians is returned,
191    /// following PyTorch's behavior.
192    ///
193    /// # Note
194    ///
195    /// The current implementation performs a full sort along the specified dimension,
196    /// which has O(nlog(n)) complexity. Additionally, most backends currently fall back
197    /// to CPU for the sort operation, which may result in slower performance compared
198    /// to native GPU operations.
199    ///
200    /// # Arguments
201    ///
202    /// - `dim` - The dimension along which to compute the median.
203    ///
204    /// # Returns
205    ///
206    /// - A tensor containing the median values along the specified dimension.
207    ///
208    /// # Example 1
209    ///
210    /// ```ignore
211    /// // Assuming backend B
212    /// let device = B::Device::default();
213    /// let tensor = Tensor::<B, 2>::from_data(
214    ///     [[1.0, 5.0, 3.0, 2.0], [8.0, 4.0, 6.0, 7.0]],
215    ///     &device,
216    /// );
217    ///
218    /// // Median along dimension 0:
219    /// // sorted columns are [1.0, 8.0], [4.0, 5.0], [3.0, 6.0], [2.0, 7.0]
220    /// let median = tensor.median(0);
221    /// // Result: [[1.0, 4.0, 3.0, 2.0]]
222    ///
223    /// // Median along dimension 1:
224    /// // sorted rows are [1.0, 2.0, 3.0, 5.0] and [4.0, 6.0, 7.0, 8.0]
225    /// let median = tensor.median(1);
226    /// // Result: [[2.0], [6.0]]
227    /// ```
228    ///
229    /// # Example 2
230    ///
231    /// The median across all elements can be calculated as follows:
232    ///
233    /// ```ignore
234    /// // D is the number of dimensions of the tensor
235    /// let flattened_tensor: Tensor<B, 1> = tensor.flatten(0, D - 1);
236    ///
237    /// // Calculate median for dim 0 since the tensor has become 1 dimensional
238    /// let median = flattened_tensor.median(0);
239    /// // Result: [4.0]
240    /// ```
241    pub fn median(self, dim: usize) -> Self {
242        // TODO: Allow backend specialization. Optimally, implement a median kernel for ruda
243        // instead of leveraging a full sort to get the median.
244        stats::median(self, dim)
245    }
246
247    /// Returns the median value along the specified dimension and its index.
248    ///
249    /// The median is not unique for input tensors with an even number of elements
250    /// in the reduced dimension. In this case, the lower of the two medians is returned,
251    /// following PyTorch's behavior.
252    ///
253    /// # Note
254    ///
255    /// The current implementation performs a full sort along the specified dimension,
256    /// which has O(nlog(n)) complexity. Additionally, most backends currently fall back
257    /// to CPU for the sort operation, which may result in slower performance compared
258    /// to native GPU operations.
259    ///
260    /// # Arguments
261    ///
262    /// - `dim` - The dimension along which to compute the median.
263    ///
264    /// # Returns
265    ///
266    /// A tuple containing:
267    /// - A tensor with the median values.
268    /// - A tensor with the indices of the median values in the original tensor.
269    ///
270    /// # Example
271    ///
272    /// ```ignore
273    /// // Assuming backend B
274    /// let device = B::Device::default();
275    /// let tensor = Tensor::<B, 2>::from_data(
276    ///     [[1.0, 5.0, 3.0, 2.0], [8.0, 4.0, 6.0, 7.0]],
277    ///     &device,
278    /// );
279    ///
280    /// // Median along dimension 1:
281    /// // sorted rows are [1.0, 2.0, 3.0, 5.0] and [4.0, 6.0, 7.0, 8.0]
282    /// let (values, indices) = tensor.median_with_indices(1);
283    /// // values: [[2.0], [6.0]], indices: [[3], [2]] (position in the original tensor)
284    /// ```
285    pub fn median_with_indices(self, dim: usize) -> (Self, Tensor<B, D, Int>) {
286        // TODO: Allow backend specialization. Optimally, implement a median kernel for ruda
287        // instead of leveraging a full sort to get the median.
288        stats::median_with_indices(self, dim)
289    }
290
291    /// Converts a tensor to the specified data type.
292    ///
293    /// Supports both within-kind casting (e.g., `FloatDType::F64`) and cross-kind casting
294    /// (e.g., `IntDType::I64` to produce an int tensor).
295    ///
296    /// This is a no-op when casting to the current dtype within the same kind.
297    ///
298    /// # Example
299    ///
300    /// ```rust
301    /// use ruda_tensor::api::backend::Backend;
302    /// use ruda_tensor::api::{Tensor, FloatDType, IntDType};
303    ///
304    /// fn example<B: Backend>() {
305    ///     let device = Default::default();
306    ///     let float_tensor = Tensor::<B, 1>::from_floats([1.0, 2.5], &device);
307    ///
308    ///     // Within-kind cast (float to float)
309    ///     let f64_tensor = float_tensor.clone().cast(FloatDType::F64);
310    ///
311    ///     // Cross-kind cast (float to int)
312    ///     let int_tensor = float_tensor.cast(IntDType::I64);
313    /// }
314    /// ```
315    #[must_use]
316    pub fn cast<T: Cast<B, Float>>(self, dtype: T) -> Tensor<B, D, T::OutputKind> {
317        Tensor::new(T::cast(self.primitive, dtype))
318    }
319
320    /// Detach the current tensor from the autodiff graph.
321    ///
322    /// This function does nothing when autodiff is not enabled.
323    /// This can be used in batchers or elsewhere to ensure that previous operations are not
324    /// considered in the autodiff graph.
325    pub fn detach(self) -> Self {
326        Self::new(TensorPrimitive::Float(B::float_detach(
327            self.primitive.tensor(),
328        )))
329    }
330
331    /// Mark the tensor to keep gradients during the backward pass.
332    ///
333    /// This function does nothing when autodiff is not enabled.
334    pub fn require_grad(self) -> Self {
335        self.set_require_grad(true)
336    }
337
338    /// Returns true if the tensor requires gradients during the backward pass.
339    pub fn is_require_grad(&self) -> bool {
340        match &self.primitive {
341            TensorPrimitive::Float(tensor) => B::float_is_require_grad(tensor),
342            TensorPrimitive::QFloat(tensor) => B::q_is_require_grad(tensor),
343        }
344    }
345
346    /// Mark the tensor as tracked or untracked depending on the require_grad argument.
347    /// When tracked, the gradients will be available after the backward pass.
348    ///
349    /// This function does nothing when autodiff is not enabled.
350    pub fn set_require_grad(self, require_grad: bool) -> Self {
351        let primitive = match self.primitive {
352            TensorPrimitive::Float(tensor) => {
353                TensorPrimitive::Float(B::float_set_require_grad(tensor, require_grad))
354            }
355            TensorPrimitive::QFloat(tensor) => {
356                TensorPrimitive::QFloat(B::q_set_require_grad(tensor, require_grad))
357            }
358        };
359        Self::new(primitive)
360    }
361
362    /// Applies the relu function to the tensor.
363    pub(crate) fn relu(self) -> Self {
364        Self::new(TensorPrimitive::Float(B::relu(self.primitive.tensor())))
365    }
366
367    /// Calculate covaraince matrix between different entries alongside a given dimension.
368    ///
369    /// # Arguments
370    ///
371    /// * `size` - The size of the square matrix.
372    /// * `correction_factor` - Is usually 1 for samples and 0 for population.
373    pub fn cov(self, dim: usize, correction_factor: usize) -> Tensor<B, D> {
374        let n = self.dims()[dim];
375        let centered = (self.clone() - self.mean_dim(dim)).swap_dims(dim, 0);
376        centered
377            .clone()
378            .transpose()
379            .matmul(centered)
380            .div_scalar(n as f32 - correction_factor as f32)
381    }
382
383    /// Convert the tensor to a lower precision data type based on the quantization scheme.
384    ///
385    /// # Arguments
386    ///
387    /// * `scheme` - The quantization scheme.
388    /// * `qparams` - The pre-computed quantization parameters.
389    ///
390    /// # Returns
391    ///
392    /// The quantized tensor.
393    pub fn quantize(
394        self,
395        scheme: &QuantScheme,
396        qparams: QuantizationParameters<B>,
397    ) -> Tensor<B, D> {
398        let tensor = self.primitive.tensor();
399        let scales = qparams.scales.primitive.tensor();
400        let scales_shape = crate::quantization::params_shape(&tensor.shape(), scheme.level);
401        assert_eq!(
402            scales.shape().num_elements(),
403            scales_shape.num_elements(),
404            "Quantization scale count must match the parameter shape"
405        );
406        let scales = B::float_reshape(scales, scales_shape);
407        Tensor::new(TensorPrimitive::QFloat(B::quantize(
408            tensor,
409            scheme,
410            QuantizationParametersPrimitive { scales },
411        )))
412    }
413
414    /// Dynamically convert the tensor to a lower precision data type based on the quantization scheme.
415    ///
416    /// # Arguments
417    ///
418    /// * `scheme` - The quantization scheme.
419    ///
420    /// # Returns
421    ///
422    /// The quantized tensor.
423    ///
424    /// # Notes
425    /// This uses [min-max calibration](crate::api::quantization::Calibration::MinMax).
426    pub fn quantize_dynamic(self, scheme: &QuantScheme) -> Tensor<B, D> {
427        Tensor::new(TensorPrimitive::QFloat(B::quantize_dynamic(
428            self.primitive.tensor(),
429            scheme,
430        )))
431    }
432
433    /// Quantize using explicitly selected calibration arithmetic without changing original input storage.
434    /// For example, FP16/BF16 inputs can retain FP32 range/scale arithmetic for any supported packed scheme.
435    /// The scheme still determines INT2/4/8 or FP4/FP8 values, scale storage, blocks and packing.
436    pub fn quantize_dynamic_with_precision(self, scheme: &QuantScheme, calibration_dtype: crate::FloatDType) -> Tensor<B, D> {
437        Tensor::new(TensorPrimitive::QFloat(B::quantize_dynamic_with_precision(self.primitive.tensor(), scheme, calibration_dtype)))
438    }
439
440    /// Dequantize directly into the requested floating storage, independent of device defaults.
441    /// Ordinary floating tensors are explicitly cast rather than returned with an unrelated dtype.
442    pub fn dequantize_with_dtype(self, dtype: crate::FloatDType) -> Tensor<B, D> {
443        let tensor = match self.primitive {
444            TensorPrimitive::QFloat(tensor) => B::dequantize(tensor, dtype),
445            TensorPrimitive::Float(tensor) => B::float_cast(tensor, dtype),
446        };
447        Tensor::new(TensorPrimitive::Float(tensor))
448    }
449
450    /// Convert the tensor back to a higher precision data type.
451    ///
452    /// If the tensor is not quantized, its value is simply returned.
453    ///
454    /// # Returns
455    ///
456    /// The dequantized tensor.
457    pub fn dequantize(self) -> Tensor<B, D> {
458        Tensor::new(TensorPrimitive::Float(self.primitive.tensor()))
459    }
460
461    /// Checks element wise if the tensor is close to another tensor.
462    ///
463    /// The tolerance is defined by the following equation:
464    ///
465    /// ```text
466    /// abs(a - b) <= (atol + rtol * abs(b))
467    ///
468    /// where `a` is the first tensor, `b` is the second tensor, `rtol` is the relative tolerance,
469    /// and `atol` is the absolute tolerance.
470    /// ```
471    ///
472    /// # Arguments
473    ///
474    /// * `other` - The tensor to compare with.
475    /// * `rtol` - Optional relative tolerance. Default is 1e-5; see `DEFAULT_RTOL`.
476    /// * `atol` - Optional absolute tolerance. Default is 1e-8; see `DEFAULT_ATOL`.
477    ///
478    /// # Returns
479    ///
480    /// A boolean tensor with the same shape as the input tensors.
481    ///
482    /// # Example
483    ///
484    /// ```rust
485    /// use ruda_tensor::api::backend::Backend;
486    /// use ruda_tensor::api::{Tensor, Shape};
487    ///
488    /// fn example<B: Backend>() {
489    ///    let device = B::Device::default();
490    ///    let tensor1 = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
491    ///    let tensor2 = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
492    ///    let tensor = tensor1.is_close(tensor2, None, None);
493    ///    println!("{tensor}");
494    ///    // [[true, true, true], [true, true, true]]
495    /// }
496    /// ```
497    pub fn is_close(self, other: Self, rtol: Option<f64>, atol: Option<f64>) -> Tensor<B, D, Bool> {
498        let rtol = rtol.unwrap_or(DEFAULT_RTOL);
499        let atol = atol.unwrap_or(DEFAULT_ATOL);
500
501        // check finite difference is close
502        let is_close_finite_val = self
503            .clone()
504            .sub(other.clone())
505            .abs()
506            .lower_equal(other.clone().abs().mul_scalar(rtol).add_scalar(atol))
507            .bool_and(self.clone().is_finite())
508            .bool_and(other.clone().is_finite());
509
510        // check if both are infinite and have same sign
511        let inf_same_sign = self
512            .clone()
513            .is_finite()
514            .bool_not()
515            .bool_and(other.clone().is_finite().bool_not())
516            .bool_and(self.equal(other));
517
518        is_close_finite_val.bool_or(inf_same_sign)
519    }
520
521    /// Checks if all elements are close to another tensor.
522    ///
523    /// The tolerance is defined by the following equation:
524    ///
525    /// ```text
526    ///
527    /// abs(a - b) <= (atol + rtol * abs(b))
528    ///
529    /// where `a` is the first tensor, `b` is the second tensor, `rtol` is the relative tolerance,
530    /// and `atol` is the absolute tolerance.
531    ///
532    /// ```
533    ///
534    /// # Arguments
535    ///
536    /// * `other` - The tensor to compare with.
537    /// * `rtol` - Optional relative tolerance. Default is 1e-5; see `DEFAULT_RTOL`.
538    /// * `atol` - Optional absolute tolerance. Default is 1e-8; see `DEFAULT_ATOL`.
539    ///
540    /// # Returns
541    ///
542    /// A boolean scalar.
543    ///
544    /// # Remarks
545    ///
546    /// # Example
547    ///
548    /// ```rust
549    /// use ruda_tensor::api::backend::Backend;
550    /// use ruda_tensor::api::{Tensor, Shape};
551    ///
552    /// fn example<B: Backend>() {
553    ///    let device = B::Device::default();
554    ///    let tensor1 = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
555    ///    let tensor2 = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
556    ///    let result = tensor1.all_close(tensor2, None, None);
557    ///    println!("{}", result);
558    ///    // true
559    /// }
560    /// ```
561    pub fn all_close(self, other: Self, rtol: Option<f64>, atol: Option<f64>) -> bool {
562        self.is_close(other, rtol, atol)
563            .all()
564            .into_scalar()
565            .to_bool()
566    }
567
568    /// Returns a new tensor with boolean elements indicating whether each element of the input is NaN.
569    ///
570    /// # Returns
571    ///
572    /// A boolean tensor where `true` indicates NaN and `false` indicates a non-NaN value.
573    ///
574    /// # Example
575    ///
576    /// ```rust
577    /// use ruda_tensor::api::backend::Backend;
578    /// use ruda_tensor::api::{Tensor, Bool, Shape};
579    ///
580    /// fn example<B: Backend>() {
581    ///    let device = B::Device::default();
582    ///    let tensor = Tensor::<B, 2>::from_data([[1.0, f64::NAN, 3.0], [5.0, 9.0, 6.0]], &device);
583    ///    let tensor = tensor.is_nan();
584    ///    println!("{tensor}");
585    ///    // [[false, true, false], [false, false, false]]
586    /// }
587    /// ```
588    pub fn is_nan(self) -> Tensor<B, D, Bool> {
589        let out_dtype = get_device_settings::<B>(&self.device()).bool_dtype;
590        Tensor::new(B::float_is_nan(self.primitive.tensor(), out_dtype))
591    }
592
593    /// Checks if the tensor contains any NaN values.
594    ///
595    /// # Returns
596    ///
597    /// A boolean tensor with a single element indicating whether the tensor contains any NaN values.
598    ///
599    /// # Example
600    ///
601    /// ```rust
602    /// use ruda_tensor::api::backend::Backend;
603    /// use ruda_tensor::api::{Tensor, Bool, Shape};
604    ///
605    /// fn example<B: Backend>() {
606    ///   let device = B::Device::default();
607    ///   let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [f64::NAN, 9.0, 6.0]], &device);
608    ///   let tensor = tensor.contains_nan();
609    ///   println!("{tensor}");
610    ///   // [true]
611    ///   let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
612    ///   let tensor = tensor.contains_nan();
613    ///   println!("{tensor}");
614    ///   // [false]
615    /// }
616    /// ```
617    pub fn contains_nan(self) -> Tensor<B, 1, Bool> {
618        // Summing the tensor will result in NaN if the tensor contains any NaN values
619        // This is faster than checking each element individually
620        // because it rolls up the NaN values into a single value
621        let sum = self.sum();
622
623        sum.is_nan()
624    }
625
626    /// Returns a new tensor with boolean elements indicating whether each element of the input is infinite (either +INF or -INF).
627    ///
628    /// # Returns
629    ///
630    /// A boolean tensor where `true` indicates that the value is infinite
631    ///
632    /// # Example
633    ///
634    /// ```rust
635    /// use ruda_tensor::api::backend::Backend;
636    /// use ruda_tensor::api::{Tensor, Bool, Shape};
637    ///
638    /// fn example<B: Backend>() {
639    ///    let device = B::Device::default();
640    ///    let tensor = Tensor::<B, 2>::from_data([[1.0, f64::INFINITY, 3.0], [f64::NAN, 9.0, 6.0]], &device);
641    ///    let tensor = tensor.is_finite();
642    ///    println!("{tensor}");
643    ///    // [[false, true, false], [false, false, false]]
644    /// }
645    /// ```
646    pub fn is_inf(self) -> Tensor<B, D, Bool> {
647        let out_dtype = get_device_settings::<B>(&self.device()).bool_dtype;
648        Tensor::new(B::float_is_inf(self.primitive.tensor(), out_dtype))
649    }
650
651    /// Returns a new tensor with boolean elements indicating whether each element of the input is finite
652    ///
653    /// # Returns
654    ///
655    /// A boolean tensor where `true` indicates that the value is finite and `false` indicates
656    /// either INF, -INF or NAN
657    ///
658    /// # Example
659    ///
660    /// ```rust
661    /// use ruda_tensor::api::backend::Backend;
662    /// use ruda_tensor::api::{Tensor, Bool, Shape};
663    ///
664    /// fn example<B: Backend>() {
665    ///    let device = B::Device::default();
666    ///    let tensor = Tensor::<B, 2>::from_data([[1.0, f64::INFINITY, 3.0], [f64::NAN, 9.0, 6.0]], &device);
667    ///    let tensor = tensor.is_finite();
668    ///    println!("{tensor}");
669    ///    // [[true, false, true], [false, true, true]]
670    /// }
671    /// ```
672    pub fn is_finite(self) -> Tensor<B, D, Bool> {
673        self.clone()
674            .is_nan()
675            .bool_not()
676            .bool_and(self.is_inf().bool_not())
677    }
678
679    /// Samples tensor as a two-dimensional spatial grid of (possibly multi-channel) values,
680    /// using the given locations in [-1, 1].
681    ///
682    /// # Arguments
683    ///
684    /// * `grid` - A tensor of locations, with shape (N, H_out, W_out, 2). Values are [-1, 1].
685    ///   A [x = -1, y = -1] means top-left, and [x = 1, y = 1] means bottom-right
686    /// * `options` - Grid sampling options (mode, padding_mode, align_corners)
687    ///
688    /// # Returns
689    ///
690    /// A tensor with shape (N, C, H_out, W_out)
691    ///
692    /// # Example
693    ///
694    /// ```ignore
695    /// use ruda_tensor::api::ops::{GridSampleOptions, GridSamplePaddingMode, InterpolateMode};
696    ///
697    /// // Default options (bilinear, zeros padding, align_corners=false)
698    /// let output = tensor.grid_sample_2d(grid, GridSampleOptions::default());
699    ///
700    /// // Custom options
701    /// let options = GridSampleOptions::new(InterpolateMode::Bilinear)
702    ///     .with_padding_mode(GridSamplePaddingMode::Border)
703    ///     .with_align_corners(true);
704    /// let output = tensor.grid_sample_2d(grid, options);
705    /// ```
706    pub fn grid_sample_2d(
707        self,
708        grid: Tensor<B, D>,
709        options: impl Into<GridSampleOptions>,
710    ) -> Tensor<B, D> {
711        Tensor::new(TensorPrimitive::Float(B::float_grid_sample_2d(
712            self.primitive.tensor(),
713            grid.primitive.tensor(),
714            options.into(),
715        )))
716    }
717
718    /// Computes the cross product of `self` and another tensor along a given dimension.
719    ///
720    /// Both `self` and `other` **must have size 3** along the specified `dim`,
721    /// because the cross product is only defined in three-dimensional space.
722    ///
723    /// # Arguments
724    ///
725    /// * `other` - The other tensor to take the cross product with.
726    /// * `dim`   - The dimension along which to compute the cross product.
727    ///
728    /// # Returns
729    ///
730    /// A tensor containing the cross product of `self` and `other` along `dim`.
731    pub fn cross<Dim: AsIndex>(self, other: Tensor<B, D>, dim: Dim) -> Tensor<B, D> {
732        let dim = dim.expect_dim_index(D);
733        check!(TensorCheck::cross(&self, &other, dim));
734        Tensor::new(TensorPrimitive::Float(B::float_cross(
735            self.primitive.tensor(),
736            other.primitive.tensor(),
737            dim,
738        )))
739    }
740
741    /// Applies element wise power operation with a float Tensor
742    ///
743    /// # Arguments
744    ///
745    /// * `other` - The tensor to apply the power operation with.
746    ///
747    /// # Example
748    ///
749    /// ```rust
750    /// use ruda_tensor::api::backend::Backend;
751    /// use ruda_tensor::api::{Tensor, Shape};
752    ///
753    /// fn example<B: Backend>() {
754    ///    let device = B::Device::default();
755    ///    let tensor1 = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
756    ///    let tensor2 = Tensor::<B, 2>::from_data([[2.0, 3.0, 4.0], [1.0, 2.0, 3.0]], &device);
757    ///    let tensor = tensor1.powf(tensor2);
758    ///    println!("{tensor}");
759    ///    // [[1.0, 8.0, 81.0], [5.0, 81.0, 216.0]]
760    /// }
761    /// ```
762    pub fn powf(self, other: Self) -> Self {
763        let primitive = match (self.primitive, other.primitive) {
764            (TensorPrimitive::Float(lhs), TensorPrimitive::Float(rhs)) => {
765                TensorPrimitive::Float(B::float_powf(lhs, rhs))
766            }
767            (TensorPrimitive::QFloat(lhs), TensorPrimitive::QFloat(rhs)) => B::q_powf(lhs, rhs),
768            (TensorPrimitive::QFloat(lhs), TensorPrimitive::Float(rhs)) => {
769                let dtype = rhs.dtype();
770                TensorPrimitive::Float(B::float_powf(B::dequantize(lhs, dtype.into()), rhs))
771            }
772            (TensorPrimitive::Float(lhs), TensorPrimitive::QFloat(rhs)) => {
773                let dtype = lhs.dtype();
774                TensorPrimitive::Float(B::float_powf(lhs, B::dequantize(rhs, dtype.into())))
775            }
776        };
777
778        Tensor::new(primitive)
779    }
780
781    /// Applies element wise power operation with a float scalar
782    ///
783    /// # Arguments
784    ///
785    /// * `other` - The scalar to apply the power operation with.
786    ///
787    /// # Example
788    ///
789    /// ```rust
790    /// use ruda_tensor::api::backend::Backend;
791    /// use ruda_tensor::api::{Tensor, Shape};
792    ///
793    /// fn example<B: Backend>() {
794    ///    let device = B::Device::default();
795    ///    let tensor = Tensor::<B, 2>::from_data([[1.0, -2.0, 3.0], [5.0, 9.0, 6.0]], &device);
796    ///    let tensor = tensor.powf_scalar(2.0);
797    ///    println!("{tensor}");
798    ///    // [[1.0, 4.0, 9.0], [25.0, 81.0, 36.0]]
799    /// }
800    /// ```
801    pub fn powf_scalar<E: ElementConversion>(self, other: E) -> Self {
802        let rhs = Scalar::new(other, &self.dtype());
803
804        let primitive = match self.primitive {
805            TensorPrimitive::Float(lhs) => TensorPrimitive::Float(B::float_powf_scalar(lhs, rhs)),
806            TensorPrimitive::QFloat(lhs) => B::q_powf_scalar(lhs, rhs),
807        };
808
809        Tensor::new(primitive)
810    }
811}
812
813impl<const D: usize, B: Backend> Tensor<B, D> {
814    /// Draws samples from a categorical distribution defined by the last dimension
815    /// of the input tensor.
816    ///
817    /// The last dimension is treated as a (possibly unnormalized) set of weights
818    /// defining a categorical distribution over categories. All leading dimensions
819    /// are treated as batch dimensions. The method returns integer indices of the
820    /// sampled categories.
821    ///
822    /// # Arguments
823    ///
824    /// * `num_samples` - Number of samples to draw per distribution. Must be >= 1.
825    ///
826    /// # Panics
827    ///
828    /// Panics if `num_samples` is 0.
829    ///
830    /// # Note
831    ///
832    /// Distributions with all-zero weights produce undefined (NaN-based) sampling
833    /// results. Callers should ensure each distribution has at least one positive
834    /// weight.
835    ///
836    /// # Returns
837    ///
838    /// An integer tensor with the same shape as the input, except the last dimension
839    /// is replaced by `num_samples`, containing sampled category indices in
840    /// `[0, num_categories)`.
841    ///
842    /// # Example
843    ///
844    /// ```rust
845    /// use ruda_tensor::api::backend::Backend;
846    /// use ruda_tensor::api::Tensor;
847    ///
848    /// fn example<B: Backend>() {
849    ///     let device = B::Device::default();
850    ///     let probs = Tensor::<B, 2>::from_floats(
851    ///         [[0.0, 1.0, 0.0], [0.0, 0.0, 1.0]],
852    ///         &device,
853    ///     );
854    ///     let samples = probs.categorical(4);
855    ///     // First row always samples index 1, second row always samples index 2
856    ///     println!("{samples}");
857    /// }
858    /// ```
859    pub fn categorical(self, num_samples: usize) -> Tensor<B, D, Int> {
860        assert!(num_samples > 0, "categorical: num_samples must be >= 1");
861
862        let shape = self.shape();
863        let num_categories = shape[D - 1];
864        let batch_size = (shape.num_elements() / num_categories).max(1);
865        let device = self.device();
866
867        // Flatten leading dimensions into a single batch dimension: [batch, categories]
868        let flat: Tensor<B, 2> = self.reshape([batch_size, num_categories]);
869
870        // Normalize weights to probabilities
871        let sum = flat.clone().sum_dim(1); // [batch, 1]
872        let probs = flat / sum;
873
874        // Cumulative sum along categories dimension
875        let cumsum = probs.cumsum(1); // [batch, categories]
876
877        // Uniform random values for each sample
878        let uniform = Tensor::<B, 2>::random(
879            [batch_size, num_samples],
880            Distribution::Uniform(0.0, 1.0),
881            &device,
882        ); // [batch, num_samples]
883
884        // Expand dimensions for broadcasting:
885        //   cumsum: [batch, categories, 1]
886        //   uniform: [batch, 1, num_samples]
887        let cumsum_3d: Tensor<B, 3> = cumsum.unsqueeze_dim(2);
888        let uniform_3d: Tensor<B, 3> = uniform.unsqueeze_dim(1);
889
890        // Count categories where cumsum < uniform (inverse CDF)
891        let mask: Tensor<B, 3, Bool> = cumsum_3d.lower(uniform_3d);
892        let indices: Tensor<B, 2, Int> = mask.int().sum_dim(1).squeeze_dim::<2>(1);
893
894        // Clamp to valid range to guard against floating-point imprecision in cumsum
895        let indices = indices.clamp(0, num_categories as i64 - 1);
896
897        // Reshape back to [...leading_dims, num_samples]
898        let mut out_shape = shape;
899        out_shape[D - 1] = num_samples;
900        indices.reshape(out_shape)
901    }
902}
903
904#[cfg(feature = "api-distributed")]
905impl<const D: usize, B> Tensor<B, D>
906where
907    B: AutodiffBackend,
908{
909    /// Returns true if the tensor is marked as distributed.
910    pub fn is_distributed(&self) -> bool {
911        match &self.primitive {
912            TensorPrimitive::Float(tensor) => B::is_distributed(tensor),
913            TensorPrimitive::QFloat(_) => unimplemented!(),
914        }
915    }
916
917    /// Mark the tensor as distributed.
918    ///
919    /// This function does nothing when autodiff or distributed is not enabled.
920    pub fn set_distributed(self, param_id: DistributedParamId) -> Self {
921        let primitive = match self.primitive {
922            TensorPrimitive::Float(tensor) => {
923                TensorPrimitive::Float(B::set_distributed_params(tensor, param_id))
924            }
925            TensorPrimitive::QFloat(_) => unimplemented!(),
926        };
927        Self::new(primitive)
928    }
929}
930
931impl<B, const D: usize, K> Tensor<B, D, K>
932where
933    B: Backend,
934    K: FloatMathOps<B>,
935{
936    /// Applies element wise square operation.
937    ///
938    #[cfg_attr(doc, doc = r#"$y_i = x_i * x_i$"#)]
939    #[cfg_attr(not(doc), doc = "`y_i = x_i * x_i`")]
940    pub fn square(self) -> Self {
941        Self::new(K::square(self.primitive))
942    }
943
944    /// Applies element wise exponential operation.
945    ///
946    #[cfg_attr(doc, doc = r#"$y_i = e^{x_i}$"#)]
947    #[cfg_attr(not(doc), doc = "`y = e^x`")]
948    pub fn exp(self) -> Self {
949        Self::new(K::exp(self.primitive))
950    }
951
952    /// Applies element wise natural logarithm of one plus the input tensor.
953    ///
954    #[cfg_attr(doc, doc = r#"$y_i = \log_e\(x_i + 1\)$"#)]
955    #[cfg_attr(not(doc), doc = "`y_i = log1p(x_i)`")]
956    pub fn log1p(self) -> Self {
957        Self::new(K::log1p(self.primitive))
958    }
959
960    /// Applies element wise natural log operation *ln*.
961    ///
962    #[cfg_attr(doc, doc = r#"$y_i = \log_e\(x_i\)$"#)]
963    #[cfg_attr(not(doc), doc = "`y_i = log(x_i)`")]
964    pub fn log(self) -> Self {
965        Self::new(K::log(self.primitive))
966    }
967
968    /// Applies element wise square root operation.
969    ///
970    pub fn sqrt(self) -> Self {
971        Tensor::new(K::sqrt(self.primitive))
972    }
973    /// Applies element wise cosine operation.
974    ///
975    #[cfg_attr(doc, doc = r#"$y_i = \cos\(x_i\)$"#)]
976    #[cfg_attr(not(doc), doc = "`y_i = cos(x_i)`")]
977    pub fn cos(self) -> Self {
978        Tensor::new(K::cos(self.primitive))
979    }
980
981    /// Applies element wise sine operation.
982    ///
983    #[cfg_attr(doc, doc = r#"$y_i = \sin\(x_i\)$"#)]
984    #[cfg_attr(not(doc), doc = "`y_i = sin(x_i)`")]
985    pub fn sin(self) -> Self {
986        Tensor::new(K::sin(self.primitive))
987    }
988
989    /// Applies element wise tangent operation.
990    ///
991    #[cfg_attr(doc, doc = r#"$y_i = \tan\(x_i\)$"#)]
992    #[cfg_attr(not(doc), doc = "`y_i = tan(x_i)`")]
993    pub fn tan(self) -> Self {
994        Tensor::new(K::tan(self.primitive))
995    }
996
997    /// Applies element wise hyperbolic cosine operation.
998    ///
999    #[cfg_attr(doc, doc = r#"$y_i = \cosh\(x_i\)$"#)]
1000    #[cfg_attr(not(doc), doc = "`y_i = cosh(x_i)`")]
1001    ///
1002    /// # Example
1003    ///
1004    /// ```rust
1005    /// use ruda_tensor::api::backend::Backend;
1006    /// use ruda_tensor::api::Tensor;
1007    ///
1008    /// fn example<B: Backend>() {
1009    ///     let device = Default::default();
1010    ///
1011    ///     let tensor = Tensor::<B, 1>::from_data([0.0, -1.0, 2.0], &device);
1012    ///     println!("{}", tensor.cosh()); // [1.0, 1.5430, 3.7621]
1013    /// }
1014    /// ```
1015    pub fn cosh(self) -> Self {
1016        Tensor::new(K::cosh(self.primitive))
1017    }
1018
1019    /// Applies element wise hyperbolic sine operation.
1020    ///
1021    #[cfg_attr(doc, doc = r#"$y_i = \sinh\(x_i\)$"#)]
1022    #[cfg_attr(not(doc), doc = "`y_i = sinh(x_i)`")]
1023    ///
1024    /// # Example
1025    ///
1026    /// ```rust
1027    /// use ruda_tensor::api::backend::Backend;
1028    /// use ruda_tensor::api::Tensor;
1029    ///
1030    /// fn example<B: Backend>() {
1031    ///     let device = Default::default();
1032    ///
1033    ///     let tensor = Tensor::<B, 1>::from_data([0.0, -1.0, 2.0], &device);
1034    ///     println!("{}", tensor.sinh()); // [0.0, -1.1752, 3.6269]
1035    /// }
1036    /// ```
1037    pub fn sinh(self) -> Self {
1038        Tensor::new(K::sinh(self.primitive))
1039    }
1040
1041    /// Applies element wise hyperbolic tangent operation.
1042    ///
1043    #[cfg_attr(doc, doc = r#"$y_i = \tanh\(x_i\)$"#)]
1044    #[cfg_attr(not(doc), doc = "`y_i = tanh(x_i)`")]
1045    ///
1046    /// # Example
1047    ///
1048    /// ```rust
1049    /// use ruda_tensor::api::backend::Backend;
1050    /// use ruda_tensor::api::Tensor;
1051    ///
1052    /// fn example<B: Backend>() {
1053    ///     let device = Default::default();
1054    ///
1055    ///     let tensor = Tensor::<B, 1>::from_data([0.0, -1.0, 2.0], &device);
1056    ///     println!("{}", tensor.tanh()); // [0.0, -0.7616, 0.9640]
1057    /// }
1058    /// ```
1059    pub fn tanh(self) -> Self {
1060        Tensor::new(K::tanh(self.primitive))
1061    }
1062
1063    /// Applies element wise inverse cosine operation.
1064    ///
1065    #[cfg_attr(doc, doc = r#"$y_i = \acos\(x_i\)$"#)]
1066    #[cfg_attr(not(doc), doc = "`y_i = acos(x_i)`")]
1067    ///
1068    /// # Example
1069    ///
1070    /// ```rust
1071    /// use ruda_tensor::api::backend::Backend;
1072    /// use ruda_tensor::api::Tensor;
1073    ///
1074    /// fn example<B: Backend>() {
1075    ///     let device = Default::default();
1076    ///
1077    ///     let tensor = Tensor::<B, 1>::from_data([0.0, -1.0, 1.0], &device);
1078    ///     println!("{}", tensor.acos()); // [1.5708, 3.1416, 0.0]
1079    /// }
1080    /// ```
1081    pub fn acos(self) -> Self {
1082        Tensor::new(K::acos(self.primitive))
1083    }
1084
1085    /// Applies element wise inverse hyperbolic cosine operation.
1086    ///
1087    #[cfg_attr(doc, doc = r#"$y_i = \acosh\(x_i\)$"#)]
1088    #[cfg_attr(not(doc), doc = "`y_i = acosh(x_i)`")]
1089    ///
1090    /// # Example
1091    ///
1092    /// ```rust
1093    /// use ruda_tensor::api::backend::Backend;
1094    /// use ruda_tensor::api::Tensor;
1095    ///
1096    /// fn example<B: Backend>() {
1097    ///     let device = Default::default();
1098    ///
1099    ///     let tensor = Tensor::<B, 1>::from_data([1.0, 2.0, 3.0], &device);
1100    ///     println!("{}", tensor.acosh()); // [0.0000, 1.3170, 1.7627]
1101    /// }
1102    /// ```
1103    pub fn acosh(self) -> Self {
1104        Tensor::new(K::acosh(self.primitive))
1105    }
1106
1107    /// Applies element wise inverse sine operation.
1108    ///
1109    #[cfg_attr(doc, doc = r#"$y_i = \asin\(x_i\)$"#)]
1110    #[cfg_attr(not(doc), doc = "`y_i = asin(x_i)`")]
1111    ///
1112    /// # Example
1113    ///
1114    /// ```rust
1115    /// use ruda_tensor::api::backend::Backend;
1116    /// use ruda_tensor::api::Tensor;
1117    ///
1118    /// fn example<B: Backend>() {
1119    ///     let device = Default::default();
1120    ///
1121    ///     let tensor = Tensor::<B, 1>::from_data([0.0, -1.0, 1.0], &device);
1122    ///     println!("{}", tensor.asin()); // [ 0.0000, -1.5708,  1.5708]
1123    /// }
1124    /// ```
1125    pub fn asin(self) -> Self {
1126        Tensor::new(K::asin(self.primitive))
1127    }
1128
1129    /// Applies element wise inverse hyperbolic sine operation.
1130    ///
1131    #[cfg_attr(doc, doc = r#"$y_i = \asinh\(x_i\)$"#)]
1132    #[cfg_attr(not(doc), doc = "`y_i = asinh(x_i)`")]
1133    ///
1134    /// # Example
1135    ///
1136    /// ```rust
1137    /// use ruda_tensor::api::backend::Backend;
1138    /// use ruda_tensor::api::Tensor;
1139    ///
1140    /// fn example<B: Backend>() {
1141    ///     let device = Default::default();
1142    ///
1143    ///     let tensor = Tensor::<B, 1>::from_data([0.0, -1.0, 1.0], &device);
1144    ///     println!("{}", tensor.asinh()); // [ 0.0000, -0.8814,  0.8814]
1145    /// }
1146    /// ```
1147    pub fn asinh(self) -> Self {
1148        Tensor::new(K::asinh(self.primitive))
1149    }
1150
1151    /// Applies element wise inverse tangent operation.
1152    ///
1153    #[cfg_attr(doc, doc = r#"$y_i = \atan\(x_i\)$"#)]
1154    #[cfg_attr(not(doc), doc = "`y_i = atan(x_i)`")]
1155    ///
1156    /// # Example
1157    ///
1158    /// ```rust
1159    /// use ruda_tensor::api::backend::Backend;
1160    /// use ruda_tensor::api::Tensor;
1161    ///
1162    /// fn example<B: Backend>() {
1163    ///     let device = Default::default();
1164    ///
1165    ///     let tensor = Tensor::<B, 1>::from_data([0.0, -1.0, 2.0], &device);
1166    ///     println!("{}", tensor.atan()); // [ 0.0, -0.7854,  1.1071]
1167    /// }
1168    /// ```
1169    pub fn atan(self) -> Self {
1170        Tensor::new(K::atan(self.primitive))
1171    }
1172
1173    /// Applies element wise inverse hyperbolic tangent operation.
1174    ///
1175    #[cfg_attr(doc, doc = r#"$y_i = \atanh\(x_i\)$"#)]
1176    #[cfg_attr(not(doc), doc = "`y_i = atanh(x_i)`")]
1177    ///
1178    /// # Example
1179    ///
1180    /// ```rust
1181    /// use ruda_tensor::api::backend::Backend;
1182    /// use ruda_tensor::api::Tensor;
1183    ///
1184    /// fn example<B: Backend>() {
1185    ///     let device = Default::default();
1186    ///
1187    ///     let tensor = Tensor::<B, 1>::from_data([0.0, -0.5, 0.5], &device);
1188    ///     println!("{}", tensor.atanh()); // [ 0.0, -0.5493,  0.5493]
1189    /// }
1190    /// ```
1191    pub fn atanh(self) -> Self {
1192        Tensor::new(K::atanh(self.primitive))
1193    }
1194
1195    /// Applies element wise inverse tangent operation using the signs of arguments to determine the correct quadrant.
1196    ///
1197    #[cfg_attr(doc, doc = r#"$z_i = \atan2\(y_i, x_i\)$"#)]
1198    #[cfg_attr(not(doc), doc = "`z_i = atan2(y_i, x_i)`")]
1199    ///
1200    /// # Example
1201    ///
1202    /// ```rust
1203    /// use ruda_tensor::api::backend::Backend;
1204    /// use ruda_tensor::api::Tensor;
1205    ///
1206    /// fn example<B: Backend>() {
1207    ///     let device = Default::default();
1208    ///
1209    ///     let lhs = Tensor::<B, 1>::from_data([-2.0, 2.0, -2.0], &device);
1210    ///     let rhs = Tensor::<B, 1>::from_data([1.0, -1.0, -1.0], &device);
1211    ///     println!("{}", lhs.atan2(rhs)); // [-1.1071,  2.0344, -2.0344]
1212    /// }
1213    /// ```
1214    pub fn atan2(self, other: Self) -> Self {
1215        Tensor::new(K::atan2(self.primitive, other.primitive))
1216    }
1217}