Skip to main content

ruprim_host/quantization/
layout.rs

1use super::*;
2
3pub fn q_reshape(tensor: HostQTensor, shape: Shape) -> HostQTensor {
4    block_safe_layout_op(tensor, |t| t.reshape(shape))
5}
6
7pub fn q_swap_dims(
8    tensor: HostQTensor,
9    dim1: usize,
10    dim2: usize,
11) -> HostQTensor {
12    block_safe_layout_op(tensor, |t| t.transpose(dim1, dim2))
13}
14
15pub fn q_permute(tensor: HostQTensor, axes: &[usize]) -> HostQTensor {
16    block_safe_layout_op(tensor, |t| t.permute(axes))
17}
18
19pub fn q_flip(tensor: HostQTensor, axes: &[usize]) -> HostQTensor {
20    block_safe_layout_op(tensor, |t| crate::flip::flip(t, axes))
21}
22
23pub fn q_expand(tensor: HostQTensor, shape: Shape) -> HostQTensor {
24    block_safe_layout_op(tensor, |t| crate::expand::expand(t, shape))
25}
26
27pub fn q_select(
28    tensor: HostQTensor,
29    dim: usize,
30    indices: HostTensor,
31) -> HostQTensor {
32    match tensor.scheme.level {
33        QuantLevel::Tensor => HostQTensor::new(
34            crate::gather_scatter::select::<i8>(tensor.tensor, dim, indices),
35            tensor.scheme,
36            tensor.scales,
37        ),
38        QuantLevel::Block(_) => {
39            let scheme = tensor.scheme;
40            let float_tensor = crate::quantization::dequantize(tensor, FloatDType::F32);
41            let result = crate::gather_scatter::select::<f32>(float_tensor, dim, indices);
42            crate::quantization::quantize_dynamic(result, &scheme)
43        }
44    }
45}
46
47pub fn q_slice(tensor: HostQTensor, slices: &[Slice]) -> HostQTensor {
48    block_safe_layout_op(tensor, |t| crate::slice::slice(t, slices))
49}
50
51pub fn q_argmax(
52    tensor: HostQTensor,
53    dim: usize,
54    out_dtype: ruda_core::tensor::IntDType,
55) -> HostTensor {
56    let tensor = crate::quantization::dequantize(tensor, FloatDType::F32);
57    let result = crate::reduce::argmax(tensor, dim);
58    if result.dtype() != DType::from(out_dtype) {
59        crate::cast::int_cast(result, out_dtype)
60    } else {
61        result
62    }
63}
64
65pub fn q_argmin(
66    tensor: HostQTensor,
67    dim: usize,
68    out_dtype: ruda_core::tensor::IntDType,
69) -> HostTensor {
70    let tensor = crate::quantization::dequantize(tensor, FloatDType::F32);
71    let result = crate::reduce::argmin(tensor, dim);
72    if result.dtype() != DType::from(out_dtype) {
73        crate::cast::int_cast(result, out_dtype)
74    } else {
75        result
76    }
77}
78
79pub fn q_gather(
80    dim: usize,
81    tensor: HostQTensor,
82    indices: HostTensor,
83) -> HostQTensor {
84    match tensor.scheme.level {
85        QuantLevel::Tensor => HostQTensor::new(
86            crate::gather_scatter::gather::<i8>(tensor.tensor, dim, indices),
87            tensor.scheme,
88            tensor.scales,
89        ),
90        QuantLevel::Block(_) => {
91            let scheme = tensor.scheme;
92            let float_tensor = crate::quantization::dequantize(tensor, FloatDType::F32);
93            let result = crate::gather_scatter::gather::<f32>(float_tensor, dim, indices);
94            crate::quantization::quantize_dynamic(result, &scheme)
95        }
96    }
97}
98
99/// Apply a layout operation to a quantized tensor.
100/// For block-quantized tensors, dequantizes and requantizes to preserve
101/// correct scale-to-block mapping.
102fn block_safe_layout_op(
103    qtensor: HostQTensor,
104    op: impl FnOnce(HostTensor) -> HostTensor,
105) -> HostQTensor {
106    match qtensor.scheme.level {
107        QuantLevel::Tensor => HostQTensor::new(op(qtensor.tensor), qtensor.scheme, qtensor.scales),
108        QuantLevel::Block(_) => {
109            let scheme = qtensor.scheme;
110            let float_tensor = crate::quantization::dequantize(qtensor, FloatDType::F32);
111            let result = op(float_tensor);
112            crate::quantization::quantize_dynamic(result, &scheme)
113        }
114    }
115}
116