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}