1pub use cubecl_common::quant::scheme::{
5 BlockSize, QuantLevel, QuantMode, QuantParam, QuantScheme, QuantStore, QuantValue,
6};
7
8pub const QPARAM_ALIGN: usize = core::mem::align_of::<f32>();
14
15use alloc::vec::Vec;
16use core::any::TypeId;
17use num_traits::PrimInt;
18use serde::{Deserialize, Serialize};
19
20use crate::{DType, Metadata, Shape, bytes::Bytes};
21
22#[derive(new, Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
28pub struct QuantConfig {
29 pub scheme: QuantScheme,
31 pub propagation: QuantPropagation,
33 }
37
38#[derive(
39 Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize, Default,
40)]
41pub enum QuantAcc {
43 #[default]
45 F32,
46 F16,
48 BF16,
50}
51
52pub enum Calibration {
54 MinMax,
56 AbsMean,
61}
62
63#[derive(
66 Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize, Default,
67)]
68pub enum QuantPropagation {
69 Propagate,
71 #[default]
73 Inhibit,
74}
75
76#[derive(Clone, Debug)]
78pub struct QParams<S> {
79 pub scales: S,
81}
82
83#[derive(Debug, Clone, PartialEq, Eq)]
85pub struct QParamTensor {
86 pub offset_start: usize,
88 pub offset_end: usize,
90 pub metadata: Metadata,
92 pub dtype: DType,
94}
95
96pub fn params_shape(data_shape: &Shape, level: QuantLevel) -> Shape {
98 match level {
99 QuantLevel::Tensor => Shape::new([1]),
100 QuantLevel::Block(block_size) => {
101 let mut params_shape = data_shape.clone();
102 let block_size = block_size.to_dim_vec(data_shape.num_dims());
103
104 for (shape, block_size) in params_shape.iter_mut().zip(block_size) {
105 *shape = (*shape).div_ceil(block_size as usize);
106 }
107
108 params_shape
109 }
110 }
111}
112
113pub struct QuantizedBytes {
123 pub bytes: Bytes,
125 pub scheme: QuantScheme,
127 pub num_elements: usize,
129}
130
131impl QuantizedBytes {
132 pub fn new<E: bytemuck::CheckedBitPattern + bytemuck::NoUninit>(
134 value: Vec<E>,
135 scheme: QuantScheme,
136 scales: &[f32],
137 ) -> Self {
138 let num_elements = value.len();
139 if TypeId::of::<E>() != TypeId::of::<i8>() {
141 panic!("Invalid quantized type");
142 }
143
144 let i8s: Vec<i8> = bytemuck::allocation::cast_vec(value);
146 let mut bytes = Bytes::from_elems(i8s);
147
148 let scales = match scheme.level {
149 QuantLevel::Tensor => &scales[..1],
150 QuantLevel::Block(_block_size) => scales,
151 };
152 let scale_bytes = encode_scales(scales, scheme.param);
153 bytes.extend_from_byte_slice_aligned(scale_bytes.as_slice(), QPARAM_ALIGN);
154
155 Self {
156 bytes,
157 scheme,
158 num_elements,
159 }
160 }
161
162 pub fn into_vec_i8(self) -> (Vec<i8>, QParams<Vec<f32>>) {
164 let param = self.scheme.param;
165 let (values, (qparams, num_params)) = self.split_values_off();
166
167 let scales_size = scale_size(param) * num_params;
173 let scales = decode_scales(&qparams[qparams.len() - scales_size..], param);
174
175 (values, QParams { scales })
176 }
177
178 fn split_i8_values(self, scale_bytes: usize) -> (Vec<i8>, Vec<u8>) {
179 let mut values = read_bytes_to_i8(self.bytes);
180
181 let values_end = values.len() - scale_bytes;
182 let qparams = values.split_off(values_end);
183
184 (values, bytemuck::cast_vec(qparams))
185 }
186
187 fn split_values_off(self) -> (Vec<i8>, (Vec<u8>, usize)) {
192 let num_params = match self.scheme.level {
193 QuantLevel::Tensor => 1,
194 QuantLevel::Block(block_size) => self.num_elements / block_size.num_elements(),
195 };
196 let scale_bytes = scale_size(self.scheme.param) * num_params;
197
198 if let QuantStore::PackedU32(packed_dim) = self.scheme.store {
199 assert_eq!(
200 packed_dim, 0,
201 "Packing must be on innermost dimension for splitting off values"
202 );
203 }
204
205 let (values, qparams) = match self.scheme.store {
206 QuantStore::Native => self.split_i8_values(scale_bytes),
207 QuantStore::PackedU32(_) => match self.scheme.value {
208 QuantValue::Q8F | QuantValue::Q8S => self.split_i8_values(scale_bytes),
209 QuantValue::Q4F | QuantValue::Q4S | QuantValue::Q2F | QuantValue::Q2S => {
210 let split_at = self.bytes.len() - scale_bytes;
211 let qparams = self.bytes[split_at..].to_vec();
212 let values = bytemuck::cast_slice::<_, u32>(&self.bytes[..split_at]);
213 let values = unpack_q_to_i8s(values, self.num_elements, &self.scheme.value);
215 (values, qparams)
216 }
217 QuantValue::E4M3 | QuantValue::E5M2 | QuantValue::E2M1 => {
218 unimplemented!("Not yet supported")
219 }
220 },
221 QuantStore::PackedNative(_) => unimplemented!("Not yet supported"),
222 };
223
224 (values, (qparams, num_params))
225 }
226}
227
228fn scale_size(param: QuantParam) -> usize {
230 match param {
231 QuantParam::F32 => 4,
232 QuantParam::F16 | QuantParam::BF16 => 2,
233 QuantParam::UE8M0 | QuantParam::UE4M3 => 1,
234 }
235}
236
237fn decode_scales(bytes: &[u8], param: QuantParam) -> Vec<f32> {
239 match param {
240 QuantParam::F32 => bytes
241 .chunks_exact(4)
242 .map(|c| f32::from_ne_bytes([c[0], c[1], c[2], c[3]]))
243 .collect(),
244 QuantParam::F16 => bytes
245 .chunks_exact(2)
246 .map(|c| crate::f16::from_ne_bytes([c[0], c[1]]).to_f32())
247 .collect(),
248 QuantParam::BF16 => bytes
249 .chunks_exact(2)
250 .map(|c| crate::bf16::from_ne_bytes([c[0], c[1]]).to_f32())
251 .collect(),
252 QuantParam::UE8M0 | QuantParam::UE4M3 => unimplemented!("Not yet supported"),
253 }
254}
255
256fn encode_scales(scales: &[f32], param: QuantParam) -> Vec<u8> {
258 match param {
259 QuantParam::F32 => scales.iter().flat_map(|s| s.to_ne_bytes()).collect(),
260 QuantParam::F16 => scales
261 .iter()
262 .flat_map(|s| crate::f16::from_f32(*s).to_ne_bytes())
263 .collect(),
264 QuantParam::BF16 => scales
265 .iter()
266 .flat_map(|s| crate::bf16::from_f32(*s).to_ne_bytes())
267 .collect(),
268 QuantParam::UE8M0 | QuantParam::UE4M3 => unimplemented!("Not yet supported"),
269 }
270}
271
272fn read_bytes_to_i8(bytes: Bytes) -> Vec<i8> {
273 match bytes.try_into_vec::<i8>() {
274 Ok(val) => val,
275 Err(bytes) => unsafe { core::mem::transmute::<Vec<u8>, Vec<i8>>(bytes.to_vec()) },
279 }
280}
281
282pub fn pack_i8s_to_u32s(values: Vec<i8>) -> Vec<u32> {
284 #[cfg(target_endian = "big")]
288 {
289 values
290 .chunks(4)
291 .map(|x| {
292 x.iter()
293 .enumerate()
294 .fold(0u32, |acc, (i, x)| acc | (*x as u32 & 0xFF) << (i * 8))
295 })
296 .collect()
297 }
298
299 #[cfg(target_endian = "little")]
302 {
303 let mut values = values;
304 let remainder = values.len() % 4;
305 if remainder != 0 {
306 values.extend(core::iter::repeat_n(0, 4 - remainder));
308 }
309
310 let len = values.len() / 4;
311 let capacity = values.capacity() / 4;
312
313 let mut values = core::mem::ManuallyDrop::new(values);
315 let ptr = values.as_mut_ptr() as *mut u32;
316
317 unsafe { Vec::from_raw_parts(ptr, len, capacity) }
318 }
319}
320
321pub(crate) fn unpack_q_to_i8s<Q: PrimInt>(
323 values: &[Q],
324 numel: usize,
325 value: &QuantValue,
326) -> Vec<i8> {
327 let size_store = size_of::<Q>() * 8;
328 let size_quant = value.size_bits();
329 let num_quants = size_store / size_quant;
330 let mask = Q::from((1 << size_quant) - 1).unwrap();
331 let sign_shift = 8 - size_quant; values
333 .iter()
334 .enumerate()
335 .flat_map(|(i, &packed)| {
336 let n = core::cmp::min(num_quants, numel - i * num_quants);
338 (0..n).map(move |i| {
345 let raw = (packed >> (i * size_quant) & mask).to_u8().unwrap();
346 ((raw << sign_shift) as i8) >> sign_shift
347 })
348 })
349 .collect()
350}
351
352#[cfg(test)]
353mod tests {
354
355 use super::*;
356 use alloc::vec;
357
358 #[test]
359 fn should_pack_i8s_to_u32() {
360 let packed = pack_i8s_to_u32s(vec![-128, 2, -3, 127]);
361
362 assert_eq!(packed, vec![2147287680]);
363 }
364
365 #[test]
366 fn should_pack_i8s_to_u32_padded() {
367 let packed = pack_i8s_to_u32s(vec![-128, 2, -3, 127, 55]);
368 let packed_padded = pack_i8s_to_u32s(vec![-128, 2, -3, 127, 55, 0, 0, 0]);
369
370 assert_eq!(packed, vec![2147287680, 55]);
371 assert_eq!(packed, packed_padded);
372 }
373
374 #[test]
375 fn should_unpack_u32s_to_i8s() {
376 let unpacked = unpack_q_to_i8s(&[2147287680u32], 4, &QuantValue::Q8S);
377
378 assert_eq!(unpacked, vec![-128, 2, -3, 127]);
379 }
380
381 #[test]
382 fn should_unpack_u32s_to_i8s_padded() {
383 let unpacked = unpack_q_to_i8s(&[55u32], 1, &QuantValue::Q8S);
384
385 assert_eq!(unpacked, vec![55]);
386 }
387
388 #[test]
389 fn should_unpack_u32s_to_i8s_arange() {
390 let unpacked = unpack_q_to_i8s(
391 &[
392 0u32, 286331136, 286331153, 572657937, 572662306, 857874978, 858993459, 858993459,
393 1145324612, 1145324612, 1431655748, 1431655765, 1717982549, 1717986918, 2003199590,
394 2004318071,
395 ],
396 128,
397 &QuantValue::Q4S,
398 );
399
400 assert_eq!(
401 unpacked,
402 vec![
403 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
404 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3,
405 3, 3, 3, 3, 3, 3, 3, 3, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 5, 5,
406 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6,
407 6, 6, 6, 6, 6, 6, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7
408 ]
409 );
410 }
411
412 #[test]
413 fn should_pack_unpack_quantization_parameters_per_tensor_symmetric() {
414 let scale = 0.03937008;
416 let values = vec![0i8, 25, 51, 76, 102, 127];
417
418 let q_bytes = QuantizedBytes::new(
419 values.clone(),
420 QuantScheme::default()
421 .with_value(QuantValue::Q8S)
422 .with_store(QuantStore::Native),
423 &[scale],
424 );
425
426 let (q_values, qparams) = q_bytes.into_vec_i8();
427
428 assert_eq!(qparams.scales, vec![scale]);
429
430 assert_eq!(q_values, values);
431 }
432}