ruprim_host/quantization/
conversion.rs1use super::*;
2use ruda_core::tensor::quantization::params_shape;
3
4pub fn quantize_dynamic(tensor: HostTensor, scheme: &QuantScheme) -> HostQTensor {
5 let shape = tensor.shape();
6 let tensor = tensor.to_contiguous();
7 let float_data = float_storage_as_f32(&tensor);
8 let (a, b) = scheme.value.range();
9 let range = b - a;
10
11 let (quantized, scales) = match scheme.level {
12 QuantLevel::Tensor => {
13 let mut alpha: f32 = 0.0;
15 for &x in &*float_data {
16 let abs = x.abs();
17 if abs > alpha {
18 alpha = abs;
19 }
20 }
21 let scale = validated_scale(2.0 * alpha / range);
22 let inv_scale = 1.0 / scale;
23
24 let quantized = float_data
26 .iter()
27 .map(|&x| (x * inv_scale).round().clamp(a, b) as i8)
28 .collect::<Vec<i8>>();
29
30 (quantized, alloc::vec![scale])
31 }
32 QuantLevel::Block(block_size) => {
33 let block_dims = block_size.to_dim_vec(shape.rank());
34 let params_shape = params_shape(&shape, scheme.level);
35 let mut alphas = alloc::vec![0.0f32; params_shape.num_elements()];
36 for (index, &x) in float_data.iter().enumerate() {
37 let block = block_param_index(index, &shape, &block_dims, ¶ms_shape);
38 let abs = x.abs();
39 if abs > alphas[block] {
40 alphas[block] = abs;
41 }
42 }
43 let scales = alphas.into_iter()
44 .map(|alpha| validated_scale(2.0 * alpha / range))
45 .collect::<Vec<_>>();
46 let inv_scales = scales.iter().map(|scale| 1.0 / scale).collect::<Vec<_>>();
47 let quantized = float_data.iter().enumerate()
48 .map(|(index, &x)| {
49 let block = block_param_index(index, &shape, &block_dims, ¶ms_shape);
50 (x * inv_scales[block]).round().clamp(a, b) as i8
51 })
52 .collect();
53 (quantized, scales)
54 }
55 };
56
57 let bytes = Bytes::from_elems(quantized);
58 let layout = Layout::contiguous(shape);
59 let qt = HostTensor::new(bytes, layout, DType::I8);
60
61 HostQTensor::new(qt, scheme.with_store(QuantStore::Native), scales)
62}
63
64pub fn quantize(
65 tensor: HostTensor,
66 scheme: &QuantScheme,
67 qparams: QParams<HostTensor>,
68) -> HostQTensor {
69 let shape = tensor.shape();
70 let tensor = tensor.to_contiguous();
71 let float_data = float_storage_as_f32(&tensor);
72
73 let scales_tensor = qparams.scales.to_contiguous();
78 let scales_data = float_storage_as_f32(&scales_tensor);
79 let scales: Vec<f32> = scales_data.iter().copied().map(validated_scale).collect();
80 assert_eq!(
81 scales.len(), params_shape(&shape, scheme.level).num_elements(),
82 "quantized scale count must match the parameter shape"
83 );
84
85 let (a, b) = scheme.value.range();
86
87 let quantized = match scheme.level {
88 QuantLevel::Tensor => {
89 let inv_scale = 1.0 / scales[0];
90 float_data
91 .iter()
92 .map(|&x| (x * inv_scale).round().clamp(a, b) as i8)
93 .collect::<Vec<i8>>()
94 }
95 QuantLevel::Block(block_size) => {
96 let block_dims = block_size.to_dim_vec(shape.rank());
97 let params_shape = params_shape(&shape, scheme.level);
98 let inv_scales = scales.iter().map(|scale| 1.0 / scale).collect::<Vec<_>>();
99 float_data.iter().enumerate()
100 .map(|(index, &x)| {
101 let block = block_param_index(index, &shape, &block_dims, ¶ms_shape);
102 (x * inv_scales[block]).round().clamp(a, b) as i8
103 })
104 .collect::<Vec<_>>()
105 }
106 };
107
108 let bytes = Bytes::from_elems(quantized);
109 let layout = Layout::contiguous(shape);
110 let qt = HostTensor::new(bytes, layout, DType::I8);
111
112 HostQTensor::new(qt, scheme.with_store(QuantStore::Native), scales)
113}
114
115pub fn dequantize(tensor: HostQTensor, dtype: FloatDType) -> HostTensor {
116 let shape = tensor.tensor.shape();
117 let qt = tensor.tensor.to_contiguous();
118 let q_data: &[i8] = qt.storage();
119
120 let dequantized = match tensor.scheme.level {
121 QuantLevel::Tensor => {
122 let scale = tensor.scales[0];
123 q_data
124 .iter()
125 .map(|&x_q| scale * x_q as f32)
126 .collect::<Vec<f32>>()
127 }
128 QuantLevel::Block(block_size) => {
129 let block_dims = block_size.to_dim_vec(shape.rank());
130 let params_shape = params_shape(&shape, tensor.scheme.level);
131 q_data
132 .iter().enumerate()
133 .map(|(index, &x_q)| {
134 let block = block_param_index(index, &shape, &block_dims, ¶ms_shape);
135 tensor.scales[block] * x_q as f32
136 })
137 .collect::<Vec<f32>>()
138 }
139 };
140
141 let layout = Layout::contiguous(shape);
142 match dtype {
143 FloatDType::F32 | FloatDType::Flex32 => {
144 HostTensor::new(Bytes::from_elems(dequantized), layout, DType::F32)
145 }
146 FloatDType::F64 => {
147 let data: Vec<f64> = dequantized.iter().map(|&v| v as f64).collect();
148 HostTensor::new(Bytes::from_elems(data), layout, DType::F64)
149 }
150 FloatDType::F16 => {
151 let data: Vec<f16> = dequantized.iter().map(|&v| f16::from_f32(v)).collect();
152 HostTensor::new(Bytes::from_elems(data), layout, DType::F16)
153 }
154 FloatDType::BF16 => {
155 let data: Vec<bf16> = dequantized.iter().map(|&v| bf16::from_f32(v)).collect();
156 HostTensor::new(Bytes::from_elems(data), layout, DType::BF16)
157 }
158 }
159}
160
161fn block_param_index(mut index: usize, shape: &Shape, block_dims: &[u8], params_shape: &Shape) -> usize {
162 let mut parameter = 0;
163 let mut stride = 1;
164 for axis in (0..shape.rank()).rev() {
165 let coordinate = index % shape[axis];
166 index /= shape[axis];
167 parameter += coordinate / block_dims[axis] as usize * stride;
168 stride *= params_shape[axis];
169 }
170 parameter
171}
172
173fn validated_scale(scale: f32) -> f32 {
175 if scale.is_normal() {
176 scale
177 } else {
178 f32::MIN_POSITIVE
179 }
180}
181