1use cubecl::prelude::*;
2use cubecl_common::quant::scheme::*;
3use cubecl_common::{e2m1x2, e4m3, e5m2};
4use cubecl_core as cubecl;
5
6#[cube]
16pub fn dequantize_aligned<Q: Scalar, S: CubePrimitive, F: Numeric, NQ: Size, NF: Size>(
17 value: Vector<Q, NQ>,
18 scale: S,
19 table: ComptimeOption<Box<[f32]>>,
20 #[comptime] scheme: QuantScheme,
21) -> Vector<F, NF> {
22 comptime!(crate::quant::check_table_bindings(&scheme, table.is_some()));
23
24 let q_values = match scheme.store {
25 QuantStore::Native | QuantStore::PackedNative(_) => Vector::<F, NF>::cast_from(value),
26 QuantStore::PackedU32(_) => {
27 unpack_cast_u32::<F, NQ, NF>(Vector::cast_from(value), table.clone(), scheme)
28 }
29 };
30
31 match scheme.mode {
32 QuantMode::Symmetric | QuantMode::Lookup => q_values * Vector::<F, NF>::cast_from(scale),
35 }
36}
37
38#[cube]
46pub fn dequantize_aligned_wide<Q: Scalar, F: Numeric, NQ: Size, NF: Size>(
47 value: Vector<Q, NQ>,
48 scale: f32,
49 table: ComptimeOption<Box<[f32]>>,
50 #[comptime] scheme: QuantScheme,
51) -> Vector<F, NF> {
52 Vector::<F, NF>::cast_from(dequantize_aligned::<Q, f32, f32, NQ, NF>(
53 value, scale, table, scheme,
54 ))
55}
56
57#[cube]
63pub fn multiply_global_scale<S: CubePrimitive>(global_scale: f32, scale: S) -> f32 {
64 global_scale * f32::cast_from(scale)
65}
66
67#[cube]
71pub fn unpack_cast_u32<F: Numeric, NQ: Size, NF: Size>(
72 value: Vector<u32, NQ>,
73 table: ComptimeOption<Box<[f32]>>,
74 #[comptime] scheme: QuantScheme,
75) -> Vector<F, NF> {
76 let num_quants = scheme.num_quants();
77 let native_packing = scheme.native_packing();
78 let size_bits = scheme.size_bits_value();
79 let mask = comptime![packing_mask(scheme)];
80 let size!(NP) = native_packing;
81
82 let mut out = Vector::<F, NF>::empty();
83
84 #[unroll]
85 for vector_idx in 0..value.vector_size() {
86 let packed_val = value.extract(vector_idx);
87 let out_offset = vector_idx * num_quants;
88 #[unroll]
89 for packed_idx in range_stepped(0, num_quants, native_packing) {
90 let shift = packed_idx * size_bits;
91 let value = (packed_val >> shift as u32) & mask;
92
93 let float_value = cast_masked::<F, NP>(value, table.clone(), scheme);
94
95 #[unroll]
96 for native_idx in 0..native_packing {
97 let out_offset = out_offset + packed_idx + native_idx;
98 out.insert(out_offset, float_value.extract(native_idx));
99 }
100 }
101 }
102
103 out
104}
105
106#[cube]
116pub fn unpack_fields<F: Numeric, NF: Size>(
117 word: u32,
118 first: u32,
119 table: ComptimeOption<Box<[f32]>>,
120 #[comptime] scheme: QuantScheme,
121) -> Vector<F, NF> {
122 comptime!(assert!(
123 !matches!(scheme.value, QuantValue::E2M1),
124 "unpack_fields: e2m1 decodes in native pairs, which a sub-word line would split"
125 ));
126 let size_bits = scheme.size_bits_value();
127 let mask = comptime![packing_mask(scheme)];
128 let size!(N1) = 1usize;
129
130 let mut out = Vector::<F, NF>::empty();
131 #[unroll]
132 for j in 0..NF::value() {
133 let shift = (first + j as u32) * size_bits as u32;
134 let field = (word >> shift) & mask;
135 let value = cast_masked::<F, N1>(field, table.clone(), scheme);
136 out.insert(j, value.extract(0usize));
137 }
138 out
139}
140
141fn packing_mask(scheme: QuantScheme) -> u32 {
144 let bits = match scheme.value {
145 QuantValue::E2M1 => 8, other => other.size_bits(),
147 };
148 (1u32 << bits) - 1
149}
150
151#[cube]
160fn cast_masked<F: Numeric, N: Size>(
161 value: u32,
162 table: ComptimeOption<Box<[f32]>>,
163 #[comptime] scheme: QuantScheme,
164) -> Vector<F, N> {
165 #[comptime]
166 match table {
167 ComptimeOption::Some(t) => Vector::<F, N>::cast_from(t[value as usize]),
170 ComptimeOption::None => cast_masked_plain::<F, N>(value, scheme),
171 }
172}
173
174#[cube]
177fn cast_masked_plain<F: Numeric, N: Size>(
178 value: u32,
179 #[comptime] scheme: QuantScheme,
180) -> Vector<F, N> {
181 match scheme.value {
182 QuantValue::E5M2 => Vector::<F, N>::cast_from(e5m2::from_bits(value as u8)),
184 QuantValue::E4M3 => Vector::<F, N>::cast_from(e4m3::from_bits(value as u8)),
185 QuantValue::E2M1 => Vector::<F, N>::cast_from(e2m1x2::from_bits(value as u8)),
186 QuantValue::Q8F
187 | QuantValue::Q4F
188 | QuantValue::Q2F
189 | QuantValue::Q8S
190 | QuantValue::Q4S
191 | QuantValue::Q2S => {
192 let size_quant = scheme.size_bits_value() as u32;
193 let sign_bit = 1u32 << (size_quant - 1);
194
195 let signed_value = (value ^ sign_bit) as i32 - sign_bit as i32;
199 Vector::<F, N>::cast_from(signed_value)
200 }
201 }
202}
203
204#[cfg(test)]
205mod tests {
206 use super::*;
207 use cubecl_core::ir::{ElemType, Scope, UIntKind};
208 use cubecl_core::{define_size, ir::settings::Dim3};
209
210 define_size!(N1);
211
212 fn test_scope() -> Scope {
214 let scope = Scope::root(KernelSettings::new(
215 Dim3::new_single(),
216 ExecutionMode::Checked,
217 AddressType::U32,
218 ));
219 scope.register_size::<N1>(1);
220 scope.register_type::<usize>(ElemType::UInt(UIntKind::U32));
221 scope
222 }
223
224 #[test]
227 fn expanding_takes_one_scale_whatever_the_levels() {
228 let scope = test_scope();
229 let one = f32::__expand_new(&scope, 1.0);
230 let value = Vector::<f32, N1>::__expand_new(&scope, one);
231
232 for scheme in [
233 QuantScheme::default(),
234 QuantScheme::default().per_block([32], ScaleDtype::F32),
235 QuantScheme::default()
236 .per_block([32], ScaleDtype::F32)
237 .per_tensor(ScaleDtype::F32),
238 ] {
239 dequantize_aligned::expand::<f32, f32, f32, N1, N1>(
240 &scope,
241 value,
242 one,
243 ComptimeOptionExpand::None,
244 scheme,
245 );
246 }
247 }
248}