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}