Skip to main content

ruda_tensor/ops/
qtensor.rs

1use alloc::vec::Vec;
2use ruda_core::tensor::{
3    BoolDType, FloatDType, IntDType, Shape, Slice,
4    quantization::{QuantPropagation, QuantScheme},
5};
6
7use crate::{
8    Backend, ExecutionError, QTensorPrimitive, TensorData, TensorMetadata, TensorPrimitive,
9    get_device_settings,
10};
11use crate::{
12    Scalar,
13    tensor::{
14        BoolTensor, Device, FloatTensor, IntTensor, QuantizedTensor,
15        quantization::{
16            Calibration, QuantizationParametersPrimitive, compute_q_params, compute_range,
17        },
18    },
19};
20
21/// Automatically applies `dequantization -> float operation -> quantization`.
22///
23/// Used for tensor ops that should always return a quantized output.
24#[macro_export]
25macro_rules! dequant_op_quant {
26    // Binary tensor float op w/ lhs & rhs
27    (
28        float_op $float_op:expr, $t1:expr, $t2:expr
29    ) => {{
30        // Heuristic: prioritize lhs scheme
31        let scheme = $t1.scheme().clone();
32
33        let t1_f = Self::dequantize($t1);
34        let t2_f = Self::dequantize($t2);
35        #[allow(clippy::redundant_closure_call)]
36        let out_f = $float_op(t1_f, t2_f);
37
38        Self::quantize_dynamic(out_f, &scheme)
39    }};
40    // Unary tensor float op
41    (
42        float_op $float_op:expr, $tensor:expr
43    ) => {{
44        let scheme = $tensor.scheme().clone();
45        let dtype = get_device_settings::<B>(&Self::q_device(&$tensor)).float_dtype;
46
47        let tensor_f = Self::dequantize($tensor, dtype);
48        #[allow(clippy::redundant_closure_call)]
49        let out_f = $float_op(tensor_f);
50
51        Self::quantize_dynamic(out_f, &scheme)
52    }};
53}
54
55/// Automatically applies `dequantization -> float operation [-> quantization]`.
56///
57/// The output quantization step is optional.
58/// It is only performed when the input quantization scheme is propagated.
59#[macro_export]
60macro_rules! dequant_op_flow {
61    // Binary tensor float op w/ lhs & rhs
62    (
63        float_op $float_op:expr, $t1:expr, $t2:expr
64    ) => {{
65        // Heuristic: prioritize lhs scheme
66        let scheme = $t1.scheme().clone();
67        let propagation = $t1.propagation();
68        let dtype = get_device_settings::<B>(&Self::q_device(&$t1)).float_dtype;
69
70        let t1_f = Self::dequantize($t1, dtype);
71        let t2_f = Self::dequantize($t2, dtype);
72        #[allow(clippy::redundant_closure_call)]
73        let out_f = $float_op(t1_f, t2_f);
74
75        match propagation {
76            QuantPropagation::Propagate => {
77                TensorPrimitive::QFloat(Self::quantize_dynamic(out_f, &scheme))
78            }
79            QuantPropagation::Inhibit => TensorPrimitive::Float(out_f),
80        }
81    }};
82    // Unary tensor float op
83    (
84        float_op $float_op:expr, $tensor:expr
85    ) => {{
86        let scheme = $tensor.scheme().clone();
87        let propagation = $tensor.propagation();
88        let dtype = get_device_settings::<B>(&Self::q_device(&$tensor)).float_dtype;
89
90        let tensor_f = Self::dequantize($tensor, dtype);
91        #[allow(clippy::redundant_closure_call)]
92        let out_f = $float_op(tensor_f);
93
94        match propagation {
95            QuantPropagation::Propagate => {
96                TensorPrimitive::QFloat(Self::quantize_dynamic(out_f, &scheme))
97            }
98            QuantPropagation::Inhibit => TensorPrimitive::Float(out_f),
99        }
100    }};
101}
102
103/// Operations on quantized tensors.
104///
105/// # Return Type Semantics
106///
107/// The return type of each operation indicates how quantization is handled:
108///
109/// ## [`QuantizedTensor<B>`]
110/// If the method returns a `QuantizedTensor<B>`, the operation is expected to preserve the quantized
111/// representation. Implementations should avoid dequantizing when possible to maintain performance.
112/// For example, shape or layout changes such as expand or transpose preserve quantization.
113///
114/// *Note: while this currently doesn't affect the quantized tensor parameters (only per-tensor is
115/// supported at the time of writing), other quantization levels (e.g., per-block) may require re-ordering
116/// the quantization parameters to match the new layout.*
117///
118///
119/// ## [`TensorPrimitive<B>`]
120/// If the method returns a `TensorPrimitive<B>` enum, the return type should align with propagation
121/// strategy specified in the quantization scheme. The output should remain quantized ([`TensorPrimitive::QFloat`])
122/// returned in floating-point form ([`TensorPrimitive::Float`]).
123///
124/// This distinction allows for fine-grained control over mixed-precision flows while still operating
125/// through a unified API.
126pub trait QTensorOps<B: Backend> {
127    /// Creates a new tensor from the data structure.
128    ///
129    /// # Arguments
130    ///
131    /// * `data` - The data structure.
132    /// * `device` - The device to create the tensor on.
133    ///
134    /// # Returns
135    ///
136    /// The tensor with the given data.
137    fn q_from_data(data: TensorData, device: &Device<B>) -> QuantizedTensor<B>;
138
139    /// Convert the tensor to a lower precision data type based on the quantization scheme and parameters.
140    fn quantize(
141        tensor: FloatTensor<B>,
142        scheme: &QuantScheme,
143        qparams: QuantizationParametersPrimitive<B>,
144    ) -> QuantizedTensor<B>;
145
146    /// Dynamically convert the tensor to a lower precision data type based on the quantization scheme.
147    fn quantize_dynamic(tensor: FloatTensor<B>, scheme: &QuantScheme) -> QuantizedTensor<B> {
148        // Dynamically compute min/max tensor range and qparams before quantizing
149        let (min, max) = compute_range::<B>(scheme, tensor.clone(), &Calibration::MinMax);
150        let qparams = compute_q_params(scheme, min, max);
151        Self::quantize(tensor, scheme, qparams)
152    }
153
154    /// Explicit calibration arithmetic independent of the original input and packed parameter storage.
155    /// Quantizes the original input rather than its calibration copy, retaining the selected scheme.
156    fn quantize_dynamic_with_precision(tensor: FloatTensor<B>, scheme: &QuantScheme,
157        calibration_dtype: FloatDType) -> QuantizedTensor<B> {
158        let calibration = B::float_cast(tensor.clone(), calibration_dtype);
159        let (min, max) = compute_range::<B>(scheme, calibration, &Calibration::MinMax);
160        let qparams = compute_q_params::<B>(scheme, min, max);
161        Self::quantize(tensor, scheme, qparams)
162    }
163
164    /// Convert the tensor back to a higher precision data type.
165    fn dequantize(tensor: QuantizedTensor<B>, dtype: FloatDType) -> FloatTensor<B>;
166
167    /// Gets the device of the tensor.
168    ///
169    /// # Arguments
170    ///
171    /// * `tensor` - The tensor.
172    ///
173    /// # Returns
174    ///
175    /// The device of the tensor.
176    fn q_device(tensor: &QuantizedTensor<B>) -> Device<B>;
177
178    /// Moves the tensor to the given device.
179    ///
180    /// # Arguments
181    ///
182    /// * `tensor` - The tensor.
183    /// * `device` - The device to move the tensor to.
184    ///
185    /// # Returns
186    ///
187    /// The tensor on the given device.
188    fn q_to_device(tensor: QuantizedTensor<B>, device: &Device<B>) -> QuantizedTensor<B>;
189
190    /// Reshapes a tensor.
191    ///
192    /// # Arguments
193    ///
194    /// * `tensor` - The tensor to reshape.
195    /// * `shape` - The new shape of the tensor.
196    ///
197    /// # Returns
198    ///
199    /// The tensor with the new shape.
200    fn q_reshape(tensor: QuantizedTensor<B>, shape: Shape) -> QuantizedTensor<B>;
201
202    /// Converts the tensor to a data structure.
203    ///
204    /// # Arguments
205    ///
206    /// * `tensor` - The tensor.
207    ///
208    /// # Returns
209    ///
210    /// The data structure with the tensor's data.
211    fn q_into_data(
212        tensor: QuantizedTensor<B>,
213    ) -> impl Future<Output = Result<TensorData, ExecutionError>> + Send;
214
215    /// Detaches a tensor from the computation graph.
216    fn q_detach(tensor: QuantizedTensor<B>) -> QuantizedTensor<B> {
217        // Should only be overridden by autodiff backends.
218        tensor
219    }
220
221    /// Sets the `require_grad` flag of a tensor.
222    fn q_set_require_grad(tensor: QuantizedTensor<B>, _require_grad: bool) -> QuantizedTensor<B> {
223        // Should only be overridden by autodiff backends.
224        tensor
225    }
226
227    /// Returns the `require_grad` flag of a tensor.
228    fn q_is_require_grad(_tensor: &QuantizedTensor<B>) -> bool {
229        // Should only be overridden by autodiff backends.
230        false
231    }
232
233    /// Broadcasts the `tensor` to the given `shape`.
234    fn q_expand(tensor: QuantizedTensor<B>, shape: Shape) -> QuantizedTensor<B>;
235
236    /// Transposes a tensor.
237    ///
238    /// # Arguments
239    ///
240    /// * `tensor` - The tensor to transpose.
241    ///
242    /// # Returns
243    ///
244    /// The transposed tensor.
245    fn q_transpose(tensor: QuantizedTensor<B>) -> QuantizedTensor<B> {
246        let ndims = tensor.shape().num_dims();
247        Self::q_swap_dims(tensor, ndims - 2, ndims - 1)
248    }
249
250    /// Swaps two dimensions of a tensor.
251    ///
252    /// # Arguments
253    ///
254    /// * `tensor` - The tensor to swap the dimensions of.
255    /// * `dim1` - The first dimension to swap.
256    /// * `dim2` - The second dimension to swap.
257    ///
258    /// # Returns
259    ///
260    /// The tensor with the dimensions swapped.
261    fn q_swap_dims(tensor: QuantizedTensor<B>, dim1: usize, dim2: usize) -> QuantizedTensor<B>;
262
263    /// Permutes the dimensions of a tensor.
264    ///
265    /// # Arguments
266    ///
267    /// * `tensor` - The tensor to permute the dimensions of.
268    /// * `axes` - The new order of the dimensions.
269    /// # Returns
270    ///
271    /// The tensor with the dimensions permuted.
272    fn q_permute(tensor: QuantizedTensor<B>, axes: &[usize]) -> QuantizedTensor<B>;
273
274    /// Reverse the order of elements in a tensor along the given axes.
275    ///
276    /// # Arguments
277    ///
278    /// * `tensor` - The tensor to reverse.
279    /// * `axes` - The axes to reverse.
280    ///
281    /// The tensor with the elements reversed.
282    fn q_flip(tensor: QuantizedTensor<B>, axes: &[usize]) -> QuantizedTensor<B>;
283
284    /// Select tensor elements along the given dimension corresponding for the given indices.
285    ///
286    /// # Arguments
287    ///
288    /// * `tensor` - The tensor to select from.
289    /// * `dim` - The dimension to select from.
290    /// * `indices` - The indices to select.
291    ///
292    /// # Returns
293    ///
294    /// The selected elements.
295    fn q_select(
296        tensor: QuantizedTensor<B>,
297        dim: usize,
298        indices: IntTensor<B>,
299    ) -> QuantizedTensor<B>;
300
301    /// Select tensor elements corresponding to the given slices.
302    ///
303    /// # Arguments
304    ///
305    /// * `tensor` - The tensor to select from.
306    /// * `slices` - The slices specifying ranges and steps for each dimension.
307    ///
308    /// # Returns
309    ///
310    /// The selected elements in a new tensor.
311    fn q_slice(tensor: QuantizedTensor<B>, slices: &[Slice]) -> QuantizedTensor<B>;
312
313    /// Gather elements from a tensor.
314    ///
315    /// # Arguments
316    ///
317    /// * `dim` - The dimension to gather from.
318    /// * `tensor` - The tensor to gather from.
319    /// * `indices` - The indices to gather.
320    ///
321    /// # Returns
322    ///
323    /// The gathered elements.
324    fn q_gather(
325        dim: usize,
326        tensor: QuantizedTensor<B>,
327        indices: IntTensor<B>,
328    ) -> QuantizedTensor<B> {
329        // Default implementation. Backends can gather on the quantized values when supported.
330        dequant_op_quant!(
331            float_op | tensor | B::float_gather(dim, tensor, indices),
332            tensor
333        )
334    }
335
336    /// Repeat the tensor along the given dimension.
337    ///
338    /// # Arguments
339    ///
340    /// * `tensor` - The tensor.
341    /// * `dim` - The dimension to repeat.
342    /// * `times` - The number of times to repeat the dimension.
343    ///
344    /// # Returns
345    ///
346    /// The tensor with the given dimension repeated.
347    fn q_repeat_dim(tensor: QuantizedTensor<B>, dim: usize, times: usize) -> QuantizedTensor<B> {
348        dequant_op_quant!(
349            float_op | tensor | B::float_repeat_dim(tensor, dim, times),
350            tensor
351        )
352    }
353
354    /// Adds two tensors together.
355    ///
356    /// # Arguments
357    ///
358    /// * `lhs` - The left hand side tensor.
359    /// * `rhs` - The right hand side tensor.
360    ///
361    /// # Returns
362    ///
363    /// The result of adding the two tensors together.
364    fn q_add(lhs: QuantizedTensor<B>, rhs: QuantizedTensor<B>) -> TensorPrimitive<B> {
365        dequant_op_flow!(float_op | lhs, rhs | B::float_add(lhs, rhs), lhs, rhs)
366    }
367
368    /// Adds a scalar to a tensor.
369    ///
370    /// # Arguments
371    ///
372    /// * `lhs` - The left hand side tensor.
373    /// * `rhs` - The right hand side scalar.
374    ///
375    /// # Returns
376    ///
377    /// The result of adding the scalar to the tensor.
378    fn q_add_scalar(lhs: QuantizedTensor<B>, rhs: Scalar) -> TensorPrimitive<B> {
379        dequant_op_flow!(float_op | tensor | B::float_add_scalar(tensor, rhs), lhs)
380    }
381
382    /// Clamps a tensor under a minimum value.
383    ///
384    /// # Arguments
385    ///
386    /// * `tensor` - The tensor to clamp.
387    /// * `min` - The minimum value.
388    ///
389    /// # Returns
390    ///
391    /// The clamped tensor.
392    fn q_clamp_min(tensor: QuantizedTensor<B>, min: Scalar) -> TensorPrimitive<B> {
393        dequant_op_flow!(float_op | tensor | B::float_clamp_min(tensor, min), tensor)
394    }
395
396    /// Clamps a tensor over a maximum value.
397    ///
398    /// # Arguments
399    ///
400    /// * `tensor` - The tensor to clamp.
401    /// * `max` - The maximum value.
402    ///
403    /// # Returns
404    ///
405    /// The clamped tensor.
406    fn q_clamp_max(tensor: QuantizedTensor<B>, max: Scalar) -> TensorPrimitive<B> {
407        dequant_op_flow!(float_op | tensor | B::float_clamp_max(tensor, max), tensor)
408    }
409
410    /// Clamps a tensor between a minimum and maximum value.
411    ///
412    /// # Arguments
413    ///
414    /// * `tensor` - The tensor to clamp.
415    /// * `min` - The minimum value.
416    /// * `max` - The maximum value.
417    ///
418    /// # Returns
419    ///
420    /// The clamped tensor.
421    fn q_clamp(tensor: QuantizedTensor<B>, min: Scalar, max: Scalar) -> TensorPrimitive<B> {
422        dequant_op_flow!(float_op | tensor | B::float_clamp(tensor, min, max), tensor)
423    }
424
425    /// Subtracts two tensors.
426    ///
427    /// # Arguments
428    ///
429    /// * `lhs` - The left hand side tensor.
430    /// * `rhs` - The right hand side tensor.
431    ///
432    /// # Returns
433    ///
434    /// The result of subtracting the two tensors.
435    fn q_sub(lhs: QuantizedTensor<B>, rhs: QuantizedTensor<B>) -> TensorPrimitive<B> {
436        dequant_op_flow!(float_op | lhs, rhs | B::float_sub(lhs, rhs), lhs, rhs)
437    }
438
439    /// Subtracts a scalar from a tensor.
440    ///
441    /// # Arguments
442    ///
443    /// * `lhs` - The left hand side tensor.
444    /// * `rhs` - The right hand side scalar.
445    ///
446    /// # Returns
447    ///
448    /// The result of subtracting the scalar from the tensor.
449    fn q_sub_scalar(lhs: QuantizedTensor<B>, rhs: Scalar) -> TensorPrimitive<B> {
450        dequant_op_flow!(float_op | tensor | B::float_sub_scalar(tensor, rhs), lhs)
451    }
452
453    /// Multiplies two tensors together element-wise.
454    fn q_mul(lhs: QuantizedTensor<B>, rhs: QuantizedTensor<B>) -> TensorPrimitive<B> {
455        dequant_op_flow!(float_op | lhs, rhs | B::float_mul(lhs, rhs), lhs, rhs)
456    }
457
458    /// Multiplies a tensor by a scalar.
459    ///
460    /// # Arguments
461    ///
462    /// * `lhs` - The left hand side tensor.
463    /// * `rhs` - The right hand side scalar.
464    ///
465    /// # Returns
466    ///
467    /// The result of multiplying the tensor by the scalar.
468    fn q_mul_scalar(lhs: QuantizedTensor<B>, rhs: Scalar) -> TensorPrimitive<B> {
469        dequant_op_flow!(float_op | tensor | B::float_mul_scalar(tensor, rhs), lhs)
470    }
471
472    /// Divides two tensors element-wise.
473    ///
474    /// # Arguments
475    ///
476    /// * `lhs` - The left hand side tensor.
477    /// * `rhs` - The right hand side tensor.
478    ///
479    /// # Returns
480    ///
481    /// The result of dividing the two tensors.
482    fn q_div(lhs: QuantizedTensor<B>, rhs: QuantizedTensor<B>) -> TensorPrimitive<B> {
483        dequant_op_flow!(float_op | lhs, rhs | B::float_div(lhs, rhs), lhs, rhs)
484    }
485
486    /// Divides a tensor by a scalar.
487    ///
488    /// # Arguments
489    ///
490    /// * `lhs` - The left hand side tensor.
491    /// * `rhs` - The right hand side scalar.
492    ///
493    /// # Returns
494    ///
495    /// The result of dividing the tensor by the scalar.
496    fn q_div_scalar(lhs: QuantizedTensor<B>, rhs: Scalar) -> TensorPrimitive<B> {
497        dequant_op_flow!(float_op | tensor | B::float_div_scalar(tensor, rhs), lhs)
498    }
499
500    /// Multiplies two tensors together using matrix multiplication.
501    ///
502    /// # Arguments
503    ///
504    /// * `lhs` - The left hand side tensor.
505    /// * `rhs` - The right hand side tensor.
506    ///
507    /// # Returns
508    ///
509    /// The result of multiplying the two tensors together using matrix multiplication.
510    fn q_matmul(lhs: TensorPrimitive<B>, rhs: TensorPrimitive<B>) -> TensorPrimitive<B> {
511        Self::q_matmul_default(lhs, rhs)
512    }
513
514    /// Existing dequantized arithmetic and propagation contract for unsupported native combinations.
515    fn q_matmul_default(lhs: TensorPrimitive<B>, rhs: TensorPrimitive<B>) -> TensorPrimitive<B> {
516        let mut propagation = QuantPropagation::Inhibit;
517        let mut scheme = QuantScheme::default();
518
519        // Pick a target dtype for any dequantization. If either operand is already
520        // a Float tensor, take its dtype so a Float-QFloat (or QFloat-Float) pair
521        // ends up matching after dequantize and `float_matmul` doesn't see a
522        // dtype mismatch. Only when both operands are QFloat do we fall back to
523        // the device default.
524        let target_dtype: Option<FloatDType> = match (&lhs, &rhs) {
525            (TensorPrimitive::Float(t), _) | (_, TensorPrimitive::Float(t)) => {
526                Some(t.dtype().into())
527            }
528            _ => None,
529        };
530
531        let lhs = match lhs {
532            TensorPrimitive::Float(lhs) => lhs,
533            TensorPrimitive::QFloat(lhs) => {
534                propagation = lhs.propagation();
535                scheme = *lhs.scheme();
536                let float_dtype = target_dtype
537                    .unwrap_or_else(|| get_device_settings::<B>(&Self::q_device(&lhs)).float_dtype);
538
539                Self::dequantize(lhs, float_dtype)
540            }
541        };
542        let rhs = match rhs {
543            TensorPrimitive::Float(rhs) => rhs,
544            TensorPrimitive::QFloat(rhs) => {
545                propagation = rhs.propagation();
546                scheme = *rhs.scheme();
547                let float_dtype = target_dtype
548                    .unwrap_or_else(|| get_device_settings::<B>(&Self::q_device(&rhs)).float_dtype);
549
550                Self::dequantize(rhs, float_dtype)
551            }
552        };
553
554        let out_f = B::float_matmul(lhs, rhs);
555        match propagation {
556            QuantPropagation::Propagate => {
557                TensorPrimitive::QFloat(<Self>::quantize_dynamic(out_f, &scheme))
558            }
559            QuantPropagation::Inhibit => TensorPrimitive::Float(out_f),
560        }
561    }
562
563    /// Negates a tensor element-wise.
564    fn q_neg(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
565        dequant_op_flow!(float_op | tensor | B::float_neg(tensor), tensor)
566    }
567
568    /// Calculates the reciprocals element-wise
569    fn q_recip(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
570        dequant_op_flow!(float_op | tensor | B::float_recip(tensor), tensor)
571    }
572
573    /// Sum of all elements in a tensor.
574    ///
575    /// # Arguments
576    ///
577    /// * `tensor` - The tensor to sum.
578    ///
579    /// # Returns
580    ///
581    /// A scalar tensor with the sum of all elements in `tensor`.
582    fn q_sum(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
583        dequant_op_flow!(float_op | tensor | B::float_sum(tensor), tensor)
584    }
585
586    /// Sum of all elements in a tensor along a dimension.
587    ///
588    /// # Arguments
589    ///
590    /// * `tensor` - The tensor to sum.
591    /// * `dim` - The dimension along which to sum.
592    ///
593    /// # Returns
594    ///
595    /// A tensor with the sum of all elements in `tensor` along `dim`.
596    fn q_sum_dim(tensor: QuantizedTensor<B>, dim: usize) -> TensorPrimitive<B> {
597        dequant_op_flow!(float_op | tensor | B::float_sum_dim(tensor, dim), tensor)
598    }
599
600    /// Product of all elements in a tensor.
601    ///
602    /// # Arguments
603    ///
604    /// * `tensor` - The tensor to product.
605    ///
606    /// # Returns
607    ///
608    /// A scalar tensor with the product of all elements in `tensor`.
609    fn q_prod(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
610        dequant_op_flow!(float_op | tensor | B::float_prod(tensor), tensor)
611    }
612
613    /// Product of all elements in a tensor along a dimension.
614    ///
615    /// # Arguments
616    ///
617    /// * `tensor` - The tensor to product.
618    ///
619    /// # Returns
620    ///
621    /// A tensor with the product of all elements in `tensor` along `dim`.
622    fn q_prod_dim(tensor: QuantizedTensor<B>, dim: usize) -> TensorPrimitive<B> {
623        dequant_op_flow!(float_op | tensor | B::float_prod_dim(tensor, dim), tensor)
624    }
625
626    /// Mean of all elements in a tensor.
627    ///
628    /// # Arguments
629    ///
630    /// * `tensor` - The tensor to mean.
631    ///
632    /// # Returns
633    ///
634    /// A scalar tensor with the mean of all elements in `tensor`.
635    fn q_mean(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
636        dequant_op_flow!(float_op | tensor | B::float_mean(tensor), tensor)
637    }
638
639    /// Mean of all elements in a tensor along a dimension.
640    ///
641    /// # Arguments
642    ///
643    /// * `tensor` - The tensor to mean.
644    /// * `dim` - The dimension along which to mean.
645    ///
646    /// # Returns
647    ///
648    /// A tensor with the mean of all elements in `tensor` along `dim`.
649    fn q_mean_dim(tensor: QuantizedTensor<B>, dim: usize) -> TensorPrimitive<B> {
650        dequant_op_flow!(float_op | tensor | B::float_mean_dim(tensor, dim), tensor)
651    }
652
653    /// Computes the cumulative sum of elements along a dimension.
654    ///
655    /// # Arguments
656    ///
657    /// * `tensor` - The tensor to compute the cumulative sum of.
658    /// * `dim` - The dimension along which to compute the cumulative sum.
659    ///
660    /// # Returns
661    ///
662    /// A tensor with the same shape where each element is the cumulative sum
663    /// of all elements up to and including that position along the dimension.
664    fn q_cumsum(tensor: QuantizedTensor<B>, dim: usize) -> TensorPrimitive<B> {
665        dequant_op_flow!(float_op | tensor | B::float_cumsum(tensor, dim), tensor)
666    }
667
668    /// Computes the cumulative product of elements along a dimension.
669    ///
670    /// # Arguments
671    ///
672    /// * `tensor` - The tensor to compute the cumulative product of.
673    /// * `dim` - The dimension along which to compute the cumulative product.
674    ///
675    /// # Returns
676    ///
677    /// A tensor with the same shape where each element is the cumulative product
678    /// of all elements up to and including that position along the dimension.
679    fn q_cumprod(tensor: QuantizedTensor<B>, dim: usize) -> TensorPrimitive<B> {
680        dequant_op_flow!(float_op | tensor | B::float_cumprod(tensor, dim), tensor)
681    }
682
683    /// Computes the cumulative minimum of elements along a dimension.
684    ///
685    /// # Arguments
686    ///
687    /// * `tensor` - The tensor to compute the cumulative minimum of.
688    /// * `dim` - The dimension along which to compute the cumulative minimum.
689    ///
690    /// # Returns
691    ///
692    /// A tensor with the same shape where each element is the minimum
693    /// of all elements up to and including that position along the dimension.
694    fn q_cummin(tensor: QuantizedTensor<B>, dim: usize) -> TensorPrimitive<B> {
695        dequant_op_flow!(float_op | tensor | B::float_cummin(tensor, dim), tensor)
696    }
697
698    /// Computes the cumulative maximum of elements along a dimension.
699    ///
700    /// # Arguments
701    ///
702    /// * `tensor` - The tensor to compute the cumulative maximum of.
703    /// * `dim` - The dimension along which to compute the cumulative maximum.
704    ///
705    /// # Returns
706    ///
707    /// A tensor with the same shape where each element is the maximum
708    /// of all elements up to and including that position along the dimension.
709    fn q_cummax(tensor: QuantizedTensor<B>, dim: usize) -> TensorPrimitive<B> {
710        dequant_op_flow!(float_op | tensor | B::float_cummax(tensor, dim), tensor)
711    }
712
713    /// Returns a new tensor with exponential values.
714    ///
715    /// # Arguments
716    ///
717    /// * `tensor` - The tensor to exponentiate.
718    ///
719    /// # Returns
720    ///
721    /// A tensor with the same shape as `tensor` with exponential values.
722    fn q_exp(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
723        dequant_op_flow!(float_op | tensor | B::float_exp(tensor), tensor)
724    }
725
726    /// Returns a new tensor with natural logarithm values.
727    ///
728    /// # Arguments
729    ///
730    /// * `tensor` - The tensor to take the logarithm of.
731    ///
732    /// # Returns
733    ///
734    /// A tensor with the same shape as `tensor` with natural logarithm values.
735    fn q_log(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
736        dequant_op_flow!(float_op | tensor | B::float_log(tensor), tensor)
737    }
738
739    /// Returns a new tensor with logarithm values of (1 + Xi).
740    ///
741    /// # Arguments
742    ///
743    /// * `tensor` - The tensor to take the logarithm of.
744    ///
745    /// # Returns
746    ///
747    /// A tensor with the same shape as `tensor` with logarithm values of (1 + Xi).
748    fn q_log1p(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
749        dequant_op_flow!(float_op | tensor | B::float_log1p(tensor), tensor)
750    }
751
752    /// Element-wise power with another tensor.
753    ///
754    /// # Arguments
755    ///
756    /// * `lhs` - The left hand side tensor.
757    /// * `rhs` - The right hand side tensor.
758    ///
759    /// # Returns
760    ///
761    /// The elements of `lhs` raised to the power of the elements of `rhs`.
762    fn q_powf(lhs: QuantizedTensor<B>, rhs: QuantizedTensor<B>) -> TensorPrimitive<B> {
763        dequant_op_flow!(float_op | lhs, rhs | B::float_powf(lhs, rhs), lhs, rhs)
764    }
765
766    /// Element-wise power with an IntTensor.
767    ///
768    /// # Arguments
769    ///
770    /// * `lhs` - The left hand side tensor.
771    /// * `rhs` - The right hand side floatTensor.
772    ///
773    /// # Returns
774    ///
775    /// The elements of `lhs` raised to the value of `rhs`. Result is an IntTensor.
776    fn q_powi(lhs: QuantizedTensor<B>, rhs: IntTensor<B>) -> TensorPrimitive<B> {
777        dequant_op_flow!(float_op | tensor | B::float_powi(tensor, rhs), lhs)
778    }
779
780    /// Element-wise power with an int scalar.
781    ///
782    /// # Arguments
783    ///
784    /// * `lhs` - The left hand side tensor.
785    /// * `rhs` - The right hand side scalar.
786    ///
787    /// # Returns
788    ///
789    /// The elements of `lhs` raised to the value of `rhs`.
790    fn q_powi_scalar(lhs: QuantizedTensor<B>, rhs: Scalar) -> TensorPrimitive<B> {
791        dequant_op_flow!(float_op | tensor | B::float_powi_scalar(tensor, rhs), lhs)
792    }
793
794    /// Element-wise power with a float scalar.
795    ///
796    /// # Arguments
797    ///
798    /// * `tensor` - The tensor to exponentiate.
799    /// * `value` - The exponent.
800    ///
801    /// # Returns
802    ///
803    /// A tensor with the same shape as `tensor` with values raised to the power of `value`.
804    fn q_powf_scalar(tensor: QuantizedTensor<B>, value: Scalar) -> TensorPrimitive<B> {
805        dequant_op_flow!(
806            float_op | tensor | B::float_powf_scalar(tensor, value),
807            tensor
808        )
809    }
810
811    /// Returns a new tensor with square root values.
812    ///
813    /// # Arguments
814    ///
815    /// * `tensor` - The tensor to take the square root of.
816    ///
817    /// # Returns
818    ///
819    /// A tensor with the same shape as `tensor` with square root values.
820    fn q_sqrt(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
821        dequant_op_flow!(float_op | tensor | B::float_sqrt(tensor), tensor)
822    }
823
824    /// Returns a new tensor with absolute values.
825    ///
826    /// # Arguments
827    ///
828    /// * `tensor` - The tensor to take absolute value of.
829    ///
830    /// # Returns
831    ///
832    /// A tensor with the same shape as `tensor` with absolute values.
833    fn q_abs(tensor: QuantizedTensor<B>) -> QuantizedTensor<B> {
834        dequant_op_quant!(float_op | tensor | B::float_abs(tensor), tensor)
835    }
836
837    /// Returns a new tensor with cosine values.
838    ///
839    /// # Arguments
840    ///
841    /// * `tensor` - The tensor to take the cosine of.
842    ///
843    /// # Returns
844    ///
845    /// A tensor with the same shape as `tensor` with cosine values.
846    fn q_cos(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
847        dequant_op_flow!(float_op | tensor | B::float_cos(tensor), tensor)
848    }
849
850    /// Returns a new tensor with sine values.
851    ///
852    /// # Arguments
853    ///
854    /// * `tensor` - The tensor to take the sine of.
855    ///
856    /// # Returns
857    ///
858    /// A tensor with the same shape as `tensor` with sine values.
859    fn q_sin(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
860        dequant_op_flow!(float_op | tensor | B::float_sin(tensor), tensor)
861    }
862
863    /// Returns a new tensor with tangent values.
864    ///
865    /// # Arguments
866    ///
867    /// * `tensor` - The tensor to take the tangent of.
868    ///
869    /// # Returns
870    ///
871    /// A tensor with the same shape as `tensor` with tangent values.
872    fn q_tan(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
873        dequant_op_flow!(float_op | tensor | B::float_tan(tensor), tensor)
874    }
875
876    /// Returns a new tensor with hyperbolic cosine values.
877    ///
878    /// # Arguments
879    ///
880    /// * `tensor` - The tensor to take the hyperbolic cosine of.
881    ///
882    /// # Returns
883    ///
884    /// A tensor with the same shape as `tensor` with hyperbolic cosine values.
885    fn q_cosh(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
886        dequant_op_flow!(float_op | tensor | B::float_cosh(tensor), tensor)
887    }
888
889    /// Returns a new tensor with hyperbolic sine values.
890    ///
891    /// # Arguments
892    ///
893    /// * `tensor` - The tensor to take the hyperbolic sine of.
894    ///
895    /// # Returns
896    ///
897    /// A tensor with the same shape as `tensor` with hyperbolic sine values.
898    fn q_sinh(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
899        dequant_op_flow!(float_op | tensor | B::float_sinh(tensor), tensor)
900    }
901
902    /// Returns a new tensor with hyperbolic tangent values.
903    ///
904    /// # Arguments
905    ///
906    /// * `tensor` - The tensor to take the hyperbolic tangent of.
907    ///
908    /// # Returns
909    ///
910    /// A tensor with the same shape as `tensor` with hyperbolic tangent values.
911    fn q_tanh(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
912        dequant_op_flow!(float_op | tensor | B::float_tanh(tensor), tensor)
913    }
914
915    /// Returns a new tensor with the error function values.
916    ///
917    /// # Arguments
918    ///
919    /// * `tensor` - The tensor to take the error function of.
920    ///
921    /// # Returns
922    ///
923    /// A tensor with the same shape as `tensor` with error function values.
924    fn q_erf(tensor: QuantizedTensor<B>) -> TensorPrimitive<B> {
925        dequant_op_flow!(float_op | tensor | B::float_erf(tensor), tensor)
926    }
927
928    /// Concatenates tensors along a dimension.
929    ///
930    /// # Arguments
931    ///
932    /// * `tensors` - The tensors to concatenate.
933    /// * `dim` - The dimension along which to concatenate.
934    ///
935    /// # Returns
936    ///
937    /// A tensor with the concatenated tensors along `dim`.
938    fn q_cat(tensors: Vec<QuantizedTensor<B>>, dim: usize) -> QuantizedTensor<B> {
939        // Heuristic: prioritize first tensor scheme
940        let first = tensors.first().unwrap();
941        let scheme = *first.scheme();
942        let dtype = get_device_settings::<B>(&Self::q_device(first)).float_dtype;
943
944        let tensor_f = tensors
945            .into_iter()
946            .map(|tensor| Self::dequantize(tensor, dtype))
947            .collect();
948
949        let out_f = B::float_cat(tensor_f, dim);
950
951        Self::quantize_dynamic(out_f, &scheme)
952    }
953
954    /// Gets the indices of the maximum elements of a tensor along an axis.
955    ///
956    /// # Arguments
957    ///
958    /// * `tensor` - The tensor to get the maximum elements of.
959    /// * `dim` - The dimension along which to get the maximum elements.
960    /// * `out_dtype` - The output tensor dtype.
961    ///
962    /// # Returns
963    ///
964    /// A tensor with the indices of the maximum elements of `tensor` along `dim`.
965    fn q_argmax(tensor: QuantizedTensor<B>, dim: usize, out_dtype: IntDType) -> IntTensor<B> {
966        let dtype = get_device_settings::<B>(&Self::q_device(&tensor)).float_dtype;
967        let tensor_f = Self::dequantize(tensor, dtype);
968        B::float_argmax(tensor_f, dim, out_dtype)
969    }
970
971    /// Gets the indices of the k maximum elements of a tensor along an axis.
972    /// If two elements are equals, order them by the lowest indices
973    ///
974    /// # Arguments
975    ///
976    /// * `tensor` - The tensor to get the k maximum elements of.
977    /// * `dim` - The dimension along which to get the maximum elements.
978    /// * `k` - number of k maximums to keep
979    /// * `out_dtype` - The output tensor dtype.
980    ///
981    /// # Returns
982    ///
983    /// A tensor with the indices of the `k` maximum elements of `tensor` along `dim`.
984    fn q_argtopk(
985        tensor: QuantizedTensor<B>,
986        dim: usize,
987        k: usize,
988        out_dtype: IntDType,
989    ) -> IntTensor<B> {
990        let dtype = get_device_settings::<B>(&Self::q_device(&tensor)).float_dtype;
991        let tensor_f = Self::dequantize(tensor, dtype);
992        B::float_argtopk(tensor_f, dim, k, out_dtype)
993    }
994
995    /// Gets the values of the k maximum elements of a tensor along an axis.
996    ///
997    /// # Arguments
998    ///
999    /// * `tensor` - The tensor to get the k maximum elements of.
1000    /// * `dim` - The dimension along which to get the maximum elements.
1001    /// * `k` - number of k maximums to keep
1002    /// * `out_dtype` - The output tensor dtype.
1003    ///
1004    /// # Returns
1005    ///
1006    /// A tensor with the values of the `k` maximum elements of `tensor` along `dim`.
1007    fn q_topk(tensor: QuantizedTensor<B>, dim: usize, k: usize) -> QuantizedTensor<B> {
1008        dequant_op_quant!(float_op | tensor | B::float_topk(tensor, dim, k), tensor)
1009    }
1010
1011    /// Gets the indices of the minimum elements of a tensor along an axis.
1012    ///
1013    /// # Arguments
1014    ///
1015    /// * `tensor` - The tensor to get the minimum elements of.
1016    /// * `dim` - The dimension along which to get the minimum elements.
1017    /// * `out_dtype` - The output tensor dtype.
1018    ///
1019    /// # Returns
1020    ///
1021    /// A tensor with the indices of the minimum elements of `tensor` along `dim`.
1022    fn q_argmin(tensor: QuantizedTensor<B>, dim: usize, out_dtype: IntDType) -> IntTensor<B> {
1023        let dtype = get_device_settings::<B>(&Self::q_device(&tensor)).float_dtype;
1024        let tensor_f = Self::dequantize(tensor, dtype);
1025        B::float_argmin(tensor_f, dim, out_dtype)
1026    }
1027
1028    /// Gets the maximum element of a tensor.
1029    ///
1030    /// # Arguments
1031    ///
1032    /// * `tensor` - The tensor to get the maximum elements of.
1033    ///
1034    /// # Returns
1035    ///
1036    /// A tensor with the maximum element of `tensor`.
1037    fn q_max(tensor: QuantizedTensor<B>) -> QuantizedTensor<B> {
1038        let shape = tensor.shape();
1039        let tensor = B::q_reshape(tensor, Shape::new([shape.num_elements()]));
1040
1041        B::q_max_dim(tensor, 0)
1042    }
1043
1044    /// Gets the maximum elements of a tensor along an axis.
1045    ///
1046    /// # Arguments
1047    ///
1048    /// * `tensor` - The tensor to get the maximum elements of.
1049    /// * `dim` - The dimension along which to get the maximum elements.
1050    ///
1051    /// # Returns
1052    ///
1053    /// A tensor with the maximum elements of `tensor` along `dim`.
1054    fn q_max_dim(tensor: QuantizedTensor<B>, dim: usize) -> QuantizedTensor<B> {
1055        let int_dtype = get_device_settings::<B>(&B::q_device(&tensor)).int_dtype;
1056        let index = B::q_argmax(tensor.clone(), dim, int_dtype);
1057
1058        B::q_gather(dim, tensor, index)
1059    }
1060
1061    /// Gets the maximum elements of a tensor along an axis and their indices.
1062    ///
1063    /// # Arguments
1064    ///
1065    /// * `tensor` - The tensor to get the maximum elements of.
1066    /// * `dim` - The dimension along which to get the maximum elements.
1067    ///
1068    /// # Returns
1069    ///
1070    /// A tuple with the maximum elements of `tensor` along `dim` and their indices.
1071    fn q_max_dim_with_indices(
1072        tensor: QuantizedTensor<B>,
1073        dim: usize,
1074        out_dtype: IntDType,
1075    ) -> (QuantizedTensor<B>, IntTensor<B>) {
1076        let index = B::q_argmax(tensor.clone(), dim, out_dtype);
1077        let values = B::q_gather(dim, tensor, index.clone());
1078
1079        (values, index)
1080    }
1081
1082    /// Gets the minimum element of a tensor.
1083    ///
1084    /// # Arguments
1085    ///
1086    /// * `tensor` - The tensor to get the minimum elements of.
1087    ///
1088    /// # Returns
1089    ///
1090    /// A tensor with the minimum element of `tensor`.
1091    fn q_min(tensor: QuantizedTensor<B>) -> QuantizedTensor<B> {
1092        let shape = tensor.shape();
1093        let tensor = B::q_reshape(tensor, Shape::new([shape.num_elements()]));
1094
1095        B::q_min_dim(tensor, 0)
1096    }
1097
1098    /// Gets the minimum elements of a tensor along an axis.
1099    ///
1100    /// # Arguments
1101    ///
1102    /// * `tensor` - The tensor to get the minimum elements of.
1103    /// * `dim` - The dimension along which to get the minimum elements.
1104    ///
1105    /// # Returns
1106    ///
1107    /// A tensor with the minimum elements of `tensor` along `dim`.
1108    fn q_min_dim(tensor: QuantizedTensor<B>, dim: usize) -> QuantizedTensor<B> {
1109        let int_dtype = get_device_settings::<B>(&B::q_device(&tensor)).int_dtype;
1110        let index = B::q_argmin(tensor.clone(), dim, int_dtype);
1111
1112        B::q_gather(dim, tensor, index)
1113    }
1114
1115    /// Gets the minimum elements of a tensor along an axis and their indices.
1116    ///
1117    /// # Arguments
1118    ///
1119    /// * `tensor` - The tensor to get the minimum elements of.
1120    /// * `dim` - The dimension along which to get the minimum elements.
1121    ///
1122    /// # Returns
1123    ///
1124    /// A tuple with the minimum elements of `tensor` along `dim` and their indices.
1125    fn q_min_dim_with_indices(
1126        tensor: QuantizedTensor<B>,
1127        dim: usize,
1128        out_dtype: IntDType,
1129    ) -> (QuantizedTensor<B>, IntTensor<B>) {
1130        let index = B::q_argmin(tensor.clone(), dim, out_dtype);
1131        let values = B::q_gather(dim, tensor, index.clone());
1132
1133        (values, index)
1134    }
1135
1136    /// Gets the maximum element of a tensor.
1137    ///
1138    /// # Arguments
1139    ///
1140    /// * `tensor` - The tensor to get the maximum elements of.
1141    ///
1142    /// # Returns
1143    ///
1144    /// A tensor with the maximum element of `tensor`.
1145    fn q_max_abs(tensor: QuantizedTensor<B>) -> QuantizedTensor<B> {
1146        let shape = tensor.shape();
1147        let tensor = B::q_reshape(tensor, Shape::new([shape.num_elements()]));
1148
1149        B::q_max_abs_dim(tensor, 0)
1150    }
1151
1152    /// Gets the maximum elements of a tensor along an axis.
1153    ///
1154    /// # Arguments
1155    ///
1156    /// * `tensor` - The tensor to get the maximum elements of.
1157    /// * `dim` - The dimension along which to get the maximum elements.
1158    ///
1159    /// # Returns
1160    ///
1161    /// A tensor with the maximum elements of `tensor` along `dim`.
1162    fn q_max_abs_dim(tensor: QuantizedTensor<B>, dim: usize) -> QuantizedTensor<B> {
1163        let int_dtype = get_device_settings::<B>(&B::q_device(&tensor)).int_dtype;
1164        let index = B::q_argmax(B::q_abs(tensor.clone()), dim, int_dtype);
1165
1166        B::q_gather(dim, tensor, index)
1167    }
1168
1169    /// Tests if any element in the `tensor` evaluates to True.
1170    ///
1171    /// # Arguments
1172    ///
1173    /// * `tensor` - The tensor to test.
1174    ///
1175    /// # Returns
1176    ///
1177    /// A boolean tensor with a single element, True if any element in the tensor is True, False otherwise.
1178    fn q_any(tensor: QuantizedTensor<B>, out_dtype: BoolDType) -> BoolTensor<B> {
1179        let dtype = get_device_settings::<B>(&Self::q_device(&tensor)).float_dtype;
1180        let tensor_f = Self::dequantize(tensor, dtype);
1181        B::float_any(tensor_f, out_dtype)
1182    }
1183
1184    /// Tests if any element in the float `tensor` evaluates to True along a given dimension `dim`.
1185    ///
1186    /// # Arguments
1187    ///
1188    /// * `tensor` - The tensor to test.
1189    /// * `dim` - The axis along which to test.
1190    ///
1191    /// # Returns
1192    ///
1193    /// A boolean tensor `Tensor<B, D, Bool>` with the same size as input `tensor`, except in the `dim` axis
1194    /// where the size is 1. The elem in the `dim` axis is True if any element along this dim in the
1195    /// input evaluates to True, False otherwise.
1196    fn q_any_dim(tensor: QuantizedTensor<B>, dim: usize, out_dtype: BoolDType) -> BoolTensor<B> {
1197        let dtype = get_device_settings::<B>(&Self::q_device(&tensor)).float_dtype;
1198        let tensor_f = Self::dequantize(tensor, dtype);
1199        B::float_any_dim(tensor_f, dim, out_dtype)
1200    }
1201
1202    /// Tests if all elements in the `tensor` evaluate to True.
1203    ///
1204    /// # Arguments
1205    ///
1206    /// * `tensor` - The tensor to test.
1207    ///
1208    /// # Returns
1209    ///
1210    /// A boolean tensor `Tensor<B, 1, Bool>` with a single element, True if all elements in the input tensor
1211    /// evaluate to True, False otherwise.
1212    fn q_all(tensor: QuantizedTensor<B>, out_dtype: BoolDType) -> BoolTensor<B> {
1213        let dtype = get_device_settings::<B>(&Self::q_device(&tensor)).float_dtype;
1214        let tensor_f = Self::dequantize(tensor, dtype);
1215        B::float_all(tensor_f, out_dtype)
1216    }
1217
1218    /// Tests if all elements in the `tensor` evaluate to True along a given dimension `dim`.
1219    ///
1220    /// # Arguments
1221    ///
1222    /// * `tensor` - The tensor to test.
1223    /// * `dim` - The axis along which to test.
1224    ///
1225    /// # Returns
1226    ///
1227    /// A boolean tensor `Tensor<B, D, Bool>` with the same size as input `tensor`, except in the `dim` axis
1228    /// where the size is 1. The elem in the `dim` axis is True if all elements along this dim in the input
1229    /// evaluates to True, False otherwise.
1230    fn q_all_dim(tensor: QuantizedTensor<B>, dim: usize, out_dtype: BoolDType) -> BoolTensor<B> {
1231        let dtype = get_device_settings::<B>(&Self::q_device(&tensor)).float_dtype;
1232        let tensor_f = Self::dequantize(tensor, dtype);
1233        B::float_all_dim(tensor_f, dim, out_dtype)
1234    }
1235
1236    /// Sort the elements of the input `tensor` by value in along a given dimension.
1237    ///
1238    /// This sort is unstable (i.e., may reorder equal elements).
1239    ///
1240    /// # Arguments
1241    ///
1242    /// * `tensor` - The input tensor.
1243    /// * `dim` - The axis along which to sort.
1244    /// * `descending` - The sorting order.
1245    ///
1246    /// # Returns
1247    ///
1248    /// A tensor with the same shape as the input tensor, where the elements are sorted by value.
1249    fn q_sort(tensor: QuantizedTensor<B>, dim: usize, descending: bool) -> QuantizedTensor<B> {
1250        // Default implementation. Backends can sort on the int values since qparams remain the same.
1251        dequant_op_quant!(
1252            float_op | tensor | B::float_sort(tensor, dim, descending),
1253            tensor
1254        )
1255    }
1256
1257    /// Sort the elements of the input `tensor` by value in along a given dimension.
1258    ///
1259    /// This sort is unstable (i.e., may reorder equal elements).
1260    ///
1261    /// # Arguments
1262    ///
1263    /// * `tensor` - The input tensor.
1264    /// * `dim` - The axis along which to sort.
1265    /// * `descending` - The sorting order.
1266    ///
1267    /// # Returns
1268    ///
1269    /// A tensor with the same shape as the input tensor and corresponding indices, where
1270    /// the elements are sorted by value and the indices map back to the original input tensor.
1271    fn q_sort_with_indices(
1272        tensor: QuantizedTensor<B>,
1273        dim: usize,
1274        descending: bool,
1275        out_dtype: IntDType,
1276    ) -> (QuantizedTensor<B>, IntTensor<B>) {
1277        let scheme = *tensor.scheme();
1278        let dtype = get_device_settings::<B>(&Self::q_device(&tensor)).float_dtype;
1279
1280        let tensor_f = Self::dequantize(tensor, dtype);
1281        let (out_f, indices) = B::float_sort_with_indices(tensor_f, dim, descending, out_dtype);
1282
1283        (Self::quantize_dynamic(out_f, &scheme), indices)
1284    }
1285
1286    /// Returns the indices that sort the elements of the input `tensor` by value along a given dimension.
1287    ///
1288    /// This sort is unstable (i.e., may reorder equal elements).
1289    ///
1290    /// # Arguments
1291    ///
1292    /// * `tensor` - The input tensor.
1293    /// * `dim` - The axis along which to sort.
1294    /// * `descending` - The sorting order.
1295    ///
1296    /// # Returns
1297    ///
1298    /// A tensor with the same shape as the input tensor the indices map back to the original input tensor.
1299    fn q_argsort(
1300        tensor: QuantizedTensor<B>,
1301        dim: usize,
1302        descending: bool,
1303        out_dtype: IntDType,
1304    ) -> IntTensor<B> {
1305        let dtype = get_device_settings::<B>(&Self::q_device(&tensor)).float_dtype;
1306        let tensor_f = Self::dequantize(tensor, dtype);
1307        B::float_argsort(tensor_f, dim, descending, out_dtype)
1308    }
1309}