1pub use cubecl_common::quant::scheme::{
5 BlockScale, BlockSize, QuantMode, QuantScheme, QuantStore, QuantValue, ScaleDtype,
6};
7
8pub const QPARAM_ALIGN: usize = core::mem::align_of::<f32>();
14
15use alloc::vec::Vec;
16use core::any::TypeId;
17use cubecl_common::e4m3;
18use num_traits::PrimInt;
19use serde::{Deserialize, Serialize};
20
21use crate::{DType, Metadata, Shape, bytes::Bytes};
22
23#[derive(new, Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
29pub struct QuantConfig {
30 pub scheme: QuantScheme,
32 pub propagation: QuantPropagation,
34 }
38
39#[derive(
40 Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize, Default,
41)]
42pub enum QuantAcc {
44 #[default]
46 F32,
47 F16,
49 BF16,
51}
52
53pub enum Calibration {
55 MinMax,
57 AbsMean,
62}
63
64#[derive(
67 Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize, Default,
68)]
69pub enum QuantPropagation {
70 Propagate,
72 #[default]
74 Inhibit,
75}
76
77#[derive(Clone, Debug)]
79pub struct QParams<S> {
80 pub scales: S,
82 pub global: Option<S>,
84}
85
86#[derive(Clone, Debug, PartialEq)]
88pub struct DecodedScales {
89 pub block: Vec<f32>,
91 pub global: Option<f32>,
93}
94
95#[derive(Debug, Clone, PartialEq, Eq)]
97pub struct QParamTensor {
98 pub offset_start: usize,
100 pub offset_end: usize,
102 pub metadata: Metadata,
104 pub dtype: DType,
106}
107
108pub fn quantizable(scheme: &QuantScheme) -> bool {
114 if scheme.scale_dtype().round_up(1.0).is_none() {
116 return false;
117 }
118
119 match (scheme.block_scale(), global_scale_dtype(scheme)) {
120 (Some(block), Some(global)) => {
121 block.dtype.max_representable() <= crate::f16::MAX.to_f32() && global == ScaleDtype::F32
125 }
126 _ => true,
127 }
128}
129
130pub fn global_scale_dtype(scheme: &QuantScheme) -> Option<ScaleDtype> {
135 scheme.block_scale().and(scheme.tensor_scale())
136}
137
138pub fn params_shape(data_shape: &Shape, scheme: &QuantScheme) -> Shape {
142 match scheme.block_size() {
143 None => Shape::new([1]),
144 Some(block_size) => Shape::from(block_size.num_blocks(data_shape.as_slice())),
145 }
146}
147
148#[derive(Debug, Clone)]
152pub struct BlockLayout {
153 shape: Shape,
154 block: Vec<u8>,
155 blocks: Shape,
156}
157
158impl BlockLayout {
159 pub fn new(shape: &Shape, block: &BlockSize) -> Self {
161 Self {
162 shape: shape.clone(),
163 block: block.to_dim_vec(shape.num_dims()),
164 blocks: Shape::from(block.num_blocks(shape.as_slice())),
165 }
166 }
167
168 pub fn num_blocks(&self) -> usize {
170 self.blocks.num_elements()
171 }
172
173 pub fn divides(&self) -> bool {
175 self.shape
176 .iter()
177 .zip(&self.block)
178 .all(|(&dim, &extent)| dim.is_multiple_of(extent as usize))
179 }
180
181 pub fn block_of(&self, mut index: usize) -> usize {
183 let mut block = 0;
184 let mut stride = 1;
185 for dim in (0..self.shape.num_dims()).rev() {
186 let coordinate = index % self.shape[dim];
187 index /= self.shape[dim];
188 block += coordinate / self.block[dim] as usize * stride;
189 stride *= self.blocks[dim];
190 }
191 block
192 }
193}
194
195pub struct QuantizedBytes {
205 pub bytes: Bytes,
207 pub scheme: QuantScheme,
209 pub shape: Shape,
212}
213
214impl QuantizedBytes {
215 pub fn new<E: bytemuck::CheckedBitPattern + bytemuck::NoUninit>(
220 value: Vec<E>,
221 shape: impl Into<Shape>,
222 scheme: QuantScheme,
223 scales: &[f32],
224 global: Option<f32>,
225 ) -> Self {
226 let shape = shape.into();
227 assert_eq!(
228 value.len(),
229 shape.num_elements(),
230 "{} quantized values do not fill a tensor of shape {shape:?}",
231 value.len()
232 );
233 if TypeId::of::<E>() != TypeId::of::<i8>() {
235 panic!("Invalid quantized type");
236 }
237
238 let i8s: Vec<i8> = bytemuck::allocation::cast_vec(value);
240 let mut bytes = Bytes::from_elems(i8s);
241
242 let scales = match scheme.block_size() {
243 None => &scales[..1],
244 Some(_) => scales,
245 };
246 let scale_bytes = encode_scales(scales, scheme.scale_dtype());
247 bytes.extend_from_byte_slice_aligned(scale_bytes.as_slice(), QPARAM_ALIGN);
248
249 match (global_scale_dtype(&scheme), global) {
251 (Some(dtype), Some(global)) => {
252 assert_eq!(
255 dtype,
256 ScaleDtype::F32,
257 "a two-level scheme stores its per-tensor scale as f32, got {scheme:?}"
258 );
259 let global_bytes = encode_scales(&[global], dtype);
260 bytes.extend_from_byte_slice_aligned(global_bytes.as_slice(), QPARAM_ALIGN);
261 }
262 (Some(_), None) => panic!("{scheme:?} requires a per-tensor scale"),
263 (None, Some(_)) => panic!("{scheme:?} does not take a per-tensor scale"),
264 (None, None) => {}
265 }
266
267 Self {
268 bytes,
269 scheme,
270 shape,
271 }
272 }
273
274 pub fn num_elements(&self) -> usize {
276 self.shape.num_elements()
277 }
278
279 pub fn into_vec_i8(self) -> (Vec<i8>, DecodedScales) {
281 let scheme = self.scheme;
282 let (values, (qparams, num_params)) = self.split_values_off();
283
284 let global_bytes = global_scale_size(&scheme);
286 let block_end = qparams
287 .len()
288 .checked_sub(global_bytes)
289 .expect("quantized parameter buffer is shorter than the scheme's global scale");
290 let block_start = block_end
291 .checked_sub(scale_size(scheme.scale_dtype()) * num_params)
292 .expect("quantized parameter buffer is shorter than the scheme's block scales");
293
294 let block = decode_scales(&qparams[block_start..block_end], scheme.scale_dtype());
295 let global =
296 global_scale_dtype(&scheme).map(|dtype| decode_scales(&qparams[block_end..], dtype)[0]);
297
298 (values, DecodedScales { block, global })
299 }
300
301 fn split_i8_values(self, scale_bytes: usize) -> (Vec<i8>, Vec<u8>) {
302 let mut values = read_bytes_to_i8(self.bytes);
303
304 let values_end = values
305 .len()
306 .checked_sub(scale_bytes)
307 .expect("quantized tensor data is shorter than its scheme's parameters");
308 let qparams = values.split_off(values_end);
309
310 (values, bytemuck::cast_vec(qparams))
311 }
312
313 fn split_values_off(self) -> (Vec<i8>, (Vec<u8>, usize)) {
318 let num_params = params_shape(&self.shape, &self.scheme).num_elements();
319 let scale_bytes =
320 scale_size(self.scheme.scale_dtype()) * num_params + global_scale_size(&self.scheme);
321
322 if let QuantStore::PackedU32(packed_dim) = self.scheme.store {
323 assert_eq!(
324 packed_dim, 0,
325 "Packing must be on innermost dimension for splitting off values"
326 );
327 }
328
329 let (values, qparams) = match self.scheme.store {
330 QuantStore::Native => self.split_i8_values(scale_bytes),
331 QuantStore::PackedU32(_) => match self.scheme.value {
332 QuantValue::Q8F | QuantValue::Q8S => self.split_i8_values(scale_bytes),
333 QuantValue::Q4F | QuantValue::Q4S | QuantValue::Q2F | QuantValue::Q2S => {
334 let split_at =
335 self.bytes.len().checked_sub(scale_bytes).expect(
336 "quantized tensor data is shorter than its scheme's parameters",
337 );
338 let qparams = self.bytes[split_at..].to_vec();
339 let values = bytemuck::cast_slice::<_, u32>(&self.bytes[..split_at]);
340 let values = unpack_q_to_i8s(values, self.num_elements(), &self.scheme.value);
342 (values, qparams)
343 }
344 QuantValue::E4M3 | QuantValue::E5M2 | QuantValue::E2M1 => {
345 unimplemented!("Not yet supported")
346 }
347 },
348 QuantStore::PackedNative(_) => unimplemented!("Not yet supported"),
349 };
350
351 (values, (qparams, num_params))
352 }
353}
354
355pub fn scale_to_dtype(scale: f32, dtype: ScaleDtype) -> f32 {
365 dtype
366 .round_up(scale)
367 .expect("UE8M0 scales are not yet supported")
368}
369
370pub fn global_scale_size(scheme: &QuantScheme) -> usize {
372 global_scale_dtype(scheme).map_or(0, scale_size)
373}
374
375fn storage_elements(scheme: &QuantScheme, shape: &Shape) -> usize {
383 let num_quants = scheme.num_quants();
384
385 match scheme.store {
386 QuantStore::PackedU32(packed_dim) | QuantStore::PackedNative(packed_dim)
387 if num_quants > 1 && !shape.is_empty() =>
388 {
389 let packed_dim = shape.num_dims() - packed_dim - 1;
390 let mut storage = shape.clone();
391 storage[packed_dim] = storage[packed_dim].div_ceil(num_quants);
392 storage.num_elements()
393 }
394 _ => shape.num_elements().div_ceil(num_quants),
395 }
396}
397
398pub fn quantized_data_len(scheme: &QuantScheme, shape: &Shape) -> usize {
401 let value_bytes = storage_elements(scheme, shape) * scheme.size_bits_stored().div_ceil(8);
402
403 let num_params = params_shape(shape, scheme).num_elements();
404 let scale_bytes = num_params * scale_size(scheme.scale_dtype());
405
406 value_bytes + scale_bytes + global_scale_size(scheme)
407}
408
409pub fn scale_size(dtype: ScaleDtype) -> usize {
411 match dtype {
412 ScaleDtype::F32 => 4,
413 ScaleDtype::F16 | ScaleDtype::BF16 => 2,
414 ScaleDtype::UE8M0 | ScaleDtype::UE4M3 => 1,
415 }
416}
417
418fn decode_scales(bytes: &[u8], dtype: ScaleDtype) -> Vec<f32> {
420 match dtype {
421 ScaleDtype::F32 => bytes
422 .as_chunks::<4>()
423 .0
424 .iter()
425 .map(|c| f32::from_ne_bytes([c[0], c[1], c[2], c[3]]))
426 .collect(),
427 ScaleDtype::F16 => bytes
428 .as_chunks::<2>()
429 .0
430 .iter()
431 .map(|c| crate::f16::from_ne_bytes([c[0], c[1]]).to_f32())
432 .collect(),
433 ScaleDtype::BF16 => bytes
434 .as_chunks::<2>()
435 .0
436 .iter()
437 .map(|c| crate::bf16::from_ne_bytes([c[0], c[1]]).to_f32())
438 .collect(),
439 ScaleDtype::UE4M3 => bytes.iter().map(|b| e4m3::from_bits(*b).to_f32()).collect(),
440 ScaleDtype::UE8M0 => unimplemented!("UE8M0 scales are not yet supported"),
441 }
442}
443
444fn encode_scales(scales: &[f32], dtype: ScaleDtype) -> Vec<u8> {
446 match dtype {
447 ScaleDtype::F32 => scales.iter().flat_map(|s| s.to_ne_bytes()).collect(),
448 ScaleDtype::F16 => scales
449 .iter()
450 .flat_map(|s| crate::f16::from_f32(*s).to_ne_bytes())
451 .collect(),
452 ScaleDtype::BF16 => scales
453 .iter()
454 .flat_map(|s| crate::bf16::from_f32(*s).to_ne_bytes())
455 .collect(),
456 ScaleDtype::UE4M3 => scales
457 .iter()
458 .map(|s| e4m3::from_f32(*s).to_bits())
459 .collect(),
460 ScaleDtype::UE8M0 => unimplemented!("UE8M0 scales are not yet supported"),
461 }
462}
463
464fn read_bytes_to_i8(bytes: Bytes) -> Vec<i8> {
465 match bytes.try_into_vec::<i8>() {
466 Ok(val) => val,
467 Err(bytes) => unsafe { core::mem::transmute::<Vec<u8>, Vec<i8>>(bytes.to_vec()) },
471 }
472}
473
474pub fn pack_i8s_to_u32s(values: Vec<i8>) -> Vec<u32> {
476 #[cfg(target_endian = "big")]
480 {
481 values
482 .chunks(4)
483 .map(|x| {
484 x.iter()
485 .enumerate()
486 .fold(0u32, |acc, (i, x)| acc | (*x as u32 & 0xFF) << (i * 8))
487 })
488 .collect()
489 }
490
491 #[cfg(target_endian = "little")]
494 {
495 let mut values = values;
496 let remainder = values.len() % 4;
497 if remainder != 0 {
498 values.extend(core::iter::repeat_n(0, 4 - remainder));
500 }
501
502 let len = values.len() / 4;
503 let capacity = values.capacity() / 4;
504
505 let mut values = core::mem::ManuallyDrop::new(values);
507 let ptr = values.as_mut_ptr() as *mut u32;
508
509 unsafe { Vec::from_raw_parts(ptr, len, capacity) }
510 }
511}
512
513pub(crate) fn unpack_q_to_i8s<Q: PrimInt>(
515 values: &[Q],
516 numel: usize,
517 value: &QuantValue,
518) -> Vec<i8> {
519 let size_store = size_of::<Q>() * 8;
520 let size_quant = value.size_bits();
521 let num_quants = size_store / size_quant;
522 let mask = Q::from((1 << size_quant) - 1).unwrap();
523 let sign_shift = 8 - size_quant; values
525 .iter()
526 .enumerate()
527 .flat_map(|(i, &packed)| {
528 let n = core::cmp::min(num_quants, numel - i * num_quants);
530 (0..n).map(move |i| {
537 let raw = (packed >> (i * size_quant) & mask).to_u8().unwrap();
538 ((raw << sign_shift) as i8) >> sign_shift
539 })
540 })
541 .collect()
542}
543
544#[cfg(test)]
545mod tests {
546
547 use super::*;
548 use alloc::vec;
549
550 #[test]
551 fn should_pack_i8s_to_u32() {
552 let packed = pack_i8s_to_u32s(vec![-128, 2, -3, 127]);
553
554 assert_eq!(packed, vec![2147287680]);
555 }
556
557 #[test]
558 fn should_pack_i8s_to_u32_padded() {
559 let packed = pack_i8s_to_u32s(vec![-128, 2, -3, 127, 55]);
560 let packed_padded = pack_i8s_to_u32s(vec![-128, 2, -3, 127, 55, 0, 0, 0]);
561
562 assert_eq!(packed, vec![2147287680, 55]);
563 assert_eq!(packed, packed_padded);
564 }
565
566 #[test]
567 fn should_unpack_u32s_to_i8s() {
568 let unpacked = unpack_q_to_i8s(&[2147287680u32], 4, &QuantValue::Q8S);
569
570 assert_eq!(unpacked, vec![-128, 2, -3, 127]);
571 }
572
573 #[test]
574 fn should_unpack_u32s_to_i8s_padded() {
575 let unpacked = unpack_q_to_i8s(&[55u32], 1, &QuantValue::Q8S);
576
577 assert_eq!(unpacked, vec![55]);
578 }
579
580 #[test]
581 fn should_unpack_u32s_to_i8s_arange() {
582 let unpacked = unpack_q_to_i8s(
583 &[
584 0u32, 286331136, 286331153, 572657937, 572662306, 857874978, 858993459, 858993459,
585 1145324612, 1145324612, 1431655748, 1431655765, 1717982549, 1717986918, 2003199590,
586 2004318071,
587 ],
588 128,
589 &QuantValue::Q4S,
590 );
591
592 assert_eq!(
593 unpacked,
594 vec![
595 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,
596 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,
597 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,
598 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,
599 6, 6, 6, 6, 6, 6, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7
600 ]
601 );
602 }
603
604 #[test]
605 fn should_pack_unpack_quantization_parameters_per_tensor_symmetric() {
606 let scale = 0.03937008;
608 let values = vec![0i8, 25, 51, 76, 102, 127];
609
610 let q_bytes = QuantizedBytes::new(
611 values.clone(),
612 [2, 3],
613 QuantScheme::default()
614 .with_value(QuantValue::Q8S)
615 .with_store(QuantStore::Native),
616 &[scale],
617 None,
618 );
619
620 let (q_values, qparams) = q_bytes.into_vec_i8();
621
622 assert_eq!(qparams.block, vec![scale]);
623
624 assert_eq!(q_values, values);
625 }
626
627 #[test]
631 fn scale_to_dtype_survives_the_codec() {
632 let scales = [0.5f32, 0.3, 1.0 / 3.0, 500.0, 1e-3, 7.7e-4];
635
636 for dtype in [
637 ScaleDtype::F32,
638 ScaleDtype::F16,
639 ScaleDtype::BF16,
640 ScaleDtype::UE4M3,
641 ] {
642 let rounded: Vec<f32> = scales.iter().map(|s| scale_to_dtype(*s, dtype)).collect();
643 let via_codec = decode_scales(&encode_scales(&rounded, dtype), dtype);
644
645 assert_eq!(
646 rounded, via_codec,
647 "the codec moves a scale {dtype:?} can already represent"
648 );
649 for (scale, rounded) in scales.iter().zip(&rounded).filter(|(s, _)| **s < 500.0) {
651 assert!(
652 rounded >= scale,
653 "{scale} rounded down to {rounded} for {dtype:?}"
654 );
655 }
656 }
657 }
658
659 #[test]
662 fn should_pack_unpack_two_level_scales() {
663 let block_scales = [0.5f32, 0.125];
665 let global = 3.0f32;
666 let values = vec![0i8, 25, 51, 76, 102, 127, -128, -1];
667
668 let scheme = QuantScheme::default()
669 .with_value(QuantValue::Q8S)
670 .with_store(QuantStore::Native)
671 .per_block([4], ScaleDtype::UE4M3)
672 .per_tensor(ScaleDtype::F32);
673
674 let q_bytes = QuantizedBytes::new(values.clone(), [8], scheme, &block_scales, Some(global));
675
676 assert_eq!(q_bytes.bytes.len(), 8 + 2 + 4);
678
679 let (q_values, scales) = q_bytes.into_vec_i8();
680
681 assert_eq!(q_values, values);
682 assert_eq!(scales.block, block_scales);
683 assert_eq!(scales.global, Some(global));
684 }
685
686 #[test]
687 #[should_panic(expected = "requires a per-tensor scale")]
688 fn two_level_scheme_without_a_global_scale_is_rejected() {
689 let scheme = QuantScheme::default()
690 .with_value(QuantValue::Q8S)
691 .with_store(QuantStore::Native)
692 .per_block([4], ScaleDtype::F32)
693 .per_tensor(ScaleDtype::F32);
694
695 QuantizedBytes::new(vec![0i8; 8], [8], scheme, &[0.5, 0.125], None);
696 }
697
698 #[test]
699 #[should_panic(expected = "stores its per-tensor scale as f32")]
700 fn a_narrower_per_tensor_scale_is_rejected() {
701 let scheme = QuantScheme::default()
702 .with_value(QuantValue::Q8S)
703 .with_store(QuantStore::Native)
704 .per_block([4], ScaleDtype::UE4M3)
705 .per_tensor(ScaleDtype::F16);
706
707 QuantizedBytes::new(vec![0i8; 8], [8], scheme, &[0.5, 0.125], Some(3.0));
708 }
709
710 #[test]
713 fn quantizable_declines_what_no_backend_can_store() {
714 assert!(quantizable(&QuantScheme::default()));
715 assert!(quantizable(
716 &QuantScheme::default().per_block([4], ScaleDtype::F16)
717 ));
718 assert!(quantizable(
719 &QuantScheme::default()
720 .per_block([4], ScaleDtype::UE4M3)
721 .per_tensor(ScaleDtype::F32)
722 ));
723
724 assert!(!quantizable(
726 &QuantScheme::default().per_block([4], ScaleDtype::UE8M0)
727 ));
728 assert!(!quantizable(
729 &QuantScheme::default().per_tensor(ScaleDtype::UE8M0)
730 ));
731
732 assert!(!quantizable(
735 &QuantScheme::default()
736 .per_block([4], ScaleDtype::F32)
737 .per_tensor(ScaleDtype::F32)
738 ));
739 assert!(!quantizable(
740 &QuantScheme::default()
741 .per_block([4], ScaleDtype::UE4M3)
742 .per_tensor(ScaleDtype::BF16)
743 ));
744 }
745
746 #[test]
749 fn encoded_scale_width_matches_scale_size() {
750 let scales = [0.5f32, 0.25, 0.125];
751
752 for dtype in [
753 ScaleDtype::F32,
754 ScaleDtype::F16,
755 ScaleDtype::BF16,
756 ScaleDtype::UE4M3,
757 ] {
758 assert_eq!(
759 encode_scales(&scales, dtype).len(),
760 scale_size(dtype) * scales.len(),
761 "encoded width disagrees with scale_size for {dtype:?}"
762 );
763 }
764 }
765
766 #[test]
767 fn should_pack_unpack_ue4m3_block_scales() {
768 let scales = [0.5f32, 0.125];
771 let values = vec![0i8, 25, 51, 76, 102, 127, -128, -1];
772
773 let q_bytes = QuantizedBytes::new(
774 values.clone(),
775 [8],
776 QuantScheme::default()
777 .with_value(QuantValue::Q8S)
778 .with_store(QuantStore::Native)
779 .per_block([4], ScaleDtype::UE4M3),
780 &scales,
781 None,
782 );
783
784 let (q_values, qparams) = q_bytes.into_vec_i8();
785
786 assert_eq!(qparams.block, scales);
787 assert_eq!(q_values, values);
788 }
789}