ruprim_host/quantization/
layout.rs1use 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
99fn 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