1use half::bf16;
24
25#[path = "kv_dequant/luts.rs"]
26mod luts;
27pub use luts::{NVFP4_E2M1_LUT, NVFP4_GROUP_SIZE, TURBO3_LUT, TURBO4_LUT, e4m3_lut};
28
29pub fn dequant_fp8_to_bf16(fp8_bytes: &[u8], scale: f32, out: &mut [bf16]) {
31 debug_assert_eq!(fp8_bytes.len(), out.len());
32 let lut = e4m3_lut();
33 for (i, b) in fp8_bytes.iter().enumerate() {
34 out[i] = bf16::from_f32(lut[*b as usize] * scale);
35 }
36}
37
38pub fn dequant_4bit_block_to_bf16(
44 raw: &[u8],
45 bs: usize,
46 nkv: usize,
47 hd: usize,
48 lut: &[f32; 16],
49 out: &mut [bf16],
50) {
51 debug_assert!(hd.is_multiple_of(NVFP4_GROUP_SIZE));
52 debug_assert!(hd.is_multiple_of(2));
53 debug_assert_eq!(out.len(), bs * nkv * hd);
54 let head_data_bytes = hd / 2;
55 let head_scale_bytes = hd / NVFP4_GROUP_SIZE;
56 let token_data_stride = nkv * head_data_bytes;
57 let token_scale_stride = nkv * head_scale_bytes;
58 let data_section_bytes = bs * token_data_stride;
59 debug_assert!(raw.len() >= data_section_bytes + bs * token_scale_stride);
60 let (data, scales) = raw.split_at(data_section_bytes);
61 let e4m3 = e4m3_lut();
62 for tok in 0..bs {
63 for kv_h in 0..nkv {
64 let d_off = tok * token_data_stride + kv_h * head_data_bytes;
65 let s_off = tok * token_scale_stride + kv_h * head_scale_bytes;
66 for byte_idx in 0..head_data_bytes {
67 let byte = data[d_off + byte_idx];
68 let n0 = (byte & 0x0F) as usize;
69 let n1 = ((byte >> 4) & 0x0F) as usize;
70 let elem_pair_idx = byte_idx * 2;
71 let group_idx = elem_pair_idx / NVFP4_GROUP_SIZE;
72 let scale = e4m3[scales[s_off + group_idx] as usize];
73 let v0 = lut[n0] * scale;
74 let v1 = lut[n1] * scale;
75 let out_base = (tok * nkv + kv_h) * hd + elem_pair_idx;
76 out[out_base] = bf16::from_f32(v0);
77 out[out_base + 1] = bf16::from_f32(v1);
78 }
79 }
80 }
81}
82
83pub fn dequant_turbo3_block_to_bf16(
99 raw: &[u8],
100 bs: usize,
101 nkv: usize,
102 hd: usize,
103 out: &mut [bf16],
104) {
105 debug_assert!(hd.is_multiple_of(8));
106 debug_assert!(hd.is_multiple_of(NVFP4_GROUP_SIZE));
107 debug_assert_eq!(out.len(), bs * nkv * hd);
108 let head_data_bytes = hd * 3 / 8;
109 let head_scale_bytes = hd / NVFP4_GROUP_SIZE;
110 let token_data_stride = nkv * head_data_bytes;
111 let token_scale_stride = nkv * head_scale_bytes;
112 let data_section_bytes = bs * token_data_stride;
113 debug_assert!(raw.len() >= data_section_bytes + bs * token_scale_stride);
114 let (data, scales) = raw.split_at(data_section_bytes);
115 let e4m3 = e4m3_lut();
116 for tok in 0..bs {
117 for kv_h in 0..nkv {
118 let d_off = tok * token_data_stride + kv_h * head_data_bytes;
119 let s_off = tok * token_scale_stride + kv_h * head_scale_bytes;
120 for triplet_idx in 0..hd / 8 {
121 let b0 = data[d_off + triplet_idx * 3] as u32;
122 let b1 = data[d_off + triplet_idx * 3 + 1] as u32;
123 let b2 = data[d_off + triplet_idx * 3 + 2] as u32;
124 let nibbles = [
125 (b0) & 0x7,
126 (b0 >> 3) & 0x7,
127 ((b0 >> 6) | (b1 << 2)) & 0x7,
128 (b1 >> 1) & 0x7,
129 (b1 >> 4) & 0x7,
130 ((b1 >> 7) | (b2 << 1)) & 0x7,
131 (b2 >> 2) & 0x7,
132 (b2 >> 5) & 0x7,
133 ];
134 let elem_base_in_head = triplet_idx * 8;
135 for k in 0..8 {
136 let elem = elem_base_in_head + k;
137 let group_idx = elem / NVFP4_GROUP_SIZE;
138 let scale = e4m3[scales[s_off + group_idx] as usize];
139 let v = TURBO3_LUT[nibbles[k] as usize] * scale;
140 let out_idx = (tok * nkv + kv_h) * hd + elem;
141 out[out_idx] = bf16::from_f32(v);
142 }
143 }
144 }
145 }
146}
147
148pub fn dequant_turbo8_block_to_bf16(
155 raw: &[u8],
156 bs: usize,
157 nkv: usize,
158 hd: usize,
159 out: &mut [bf16],
160) {
161 debug_assert!(hd.is_multiple_of(NVFP4_GROUP_SIZE));
162 debug_assert_eq!(out.len(), bs * nkv * hd);
163 let head_data_bytes = hd;
164 let head_scale_bytes = (hd / NVFP4_GROUP_SIZE) * 2; let token_data_stride = nkv * head_data_bytes;
166 let token_scale_stride = nkv * head_scale_bytes;
167 let data_section_bytes = bs * token_data_stride;
168 debug_assert!(raw.len() >= data_section_bytes + bs * token_scale_stride);
169 let (data, scales) = raw.split_at(data_section_bytes);
170 let e4m3 = e4m3_lut();
171 for tok in 0..bs {
172 for kv_h in 0..nkv {
173 let d_off = tok * token_data_stride + kv_h * head_data_bytes;
174 let s_off = tok * token_scale_stride + kv_h * head_scale_bytes;
175 for i in 0..hd {
176 let byte = data[d_off + i] as usize;
177 let group_idx = i / NVFP4_GROUP_SIZE;
178 let s_byte_off = s_off + group_idx * 2;
180 let scale_bf16 = bf16::from_le_bytes([scales[s_byte_off], scales[s_byte_off + 1]]);
181 let scale = scale_bf16.to_f32();
182 let v = e4m3[byte] * scale;
183 let out_idx = (tok * nkv + kv_h) * hd + i;
184 out[out_idx] = bf16::from_f32(v);
185 }
186 }
187 }
188}
189
190#[cfg(test)]
191mod tests {
192 use super::*;
193
194 const FP8_ONE: u8 = 0x38;
196
197 #[test]
198 fn e4m3_lut_basics() {
199 let lut = e4m3_lut();
200 assert_eq!(lut[0x00], 0.0);
201 assert_eq!(lut[0x80], -0.0);
202 assert!((lut[0x38] - 1.0).abs() < 1e-6); assert!((lut[0xB8] + 1.0).abs() < 1e-6); assert!((lut[0x3F] - 1.875).abs() < 1e-6); assert!((lut[0x40] - 2.0).abs() < 1e-6); assert!((lut[0x78] - 256.0).abs() < 1e-6); assert!((lut[0x7E] - 448.0).abs() < 1e-6); assert!(lut[0x7F].is_nan());
209 assert!(lut[0xFF].is_nan());
210 assert!((lut[0x01] - (1.0 / 512.0)).abs() < 1e-9); }
212
213 #[test]
214 fn fp8_dequant_with_scale() {
215 let bytes = [0x38u8, 0x3F, 0x40, 0xB8];
216 let mut out = vec![bf16::ZERO; 4];
217 dequant_fp8_to_bf16(&bytes, 0.5, &mut out);
218 let f: Vec<f32> = out.iter().map(|x| x.to_f32()).collect();
219 assert!((f[0] - 0.5).abs() < 1e-2);
220 assert!((f[1] - 0.9375).abs() < 1e-2);
221 assert!((f[2] - 1.0).abs() < 1e-2);
222 assert!((f[3] + 0.5).abs() < 1e-2);
223 }
224
225 #[test]
226 #[allow(clippy::needless_range_loop)]
227 fn nvfp4_dequant_layout() {
228 let bs = 1;
229 let nkv = 1;
230 let hd = 16;
231 let mut raw = vec![0u8; 8 + 1];
232 for i in 0..8 {
233 let lo = (2 * i) & 0xF;
234 let hi = (2 * i + 1) & 0xF;
235 raw[i] = (lo as u8) | ((hi as u8) << 4);
236 }
237 raw[8] = FP8_ONE;
238 let mut out = vec![bf16::ZERO; bs * nkv * hd];
239 dequant_4bit_block_to_bf16(&raw, bs, nkv, hd, &NVFP4_E2M1_LUT, &mut out);
240 for i in 0..hd {
241 let expected = NVFP4_E2M1_LUT[i];
242 assert!(
243 (out[i].to_f32() - expected).abs() < 1e-2,
244 "elem {i}: expected {expected}, got {}",
245 out[i].to_f32(),
246 );
247 }
248 }
249
250 #[test]
251 #[allow(clippy::needless_range_loop)]
252 fn turbo4_dequant_layout() {
253 let bs = 1;
255 let nkv = 1;
256 let hd = 16;
257 let mut raw = vec![0u8; 8 + 1];
258 for i in 0..8 {
259 raw[i] = ((2 * i) as u8 & 0xF) | (((2 * i + 1) as u8 & 0xF) << 4);
260 }
261 raw[8] = FP8_ONE;
262 let mut out = vec![bf16::ZERO; hd];
263 dequant_4bit_block_to_bf16(&raw, bs, nkv, hd, &TURBO4_LUT, &mut out);
264 for i in 0..hd {
265 assert!(
266 (out[i].to_f32() - TURBO4_LUT[i]).abs() < 1e-2,
267 "elem {i}: expected {}, got {}",
268 TURBO4_LUT[i],
269 out[i].to_f32(),
270 );
271 }
272 }
273
274 #[test]
275 fn turbo3_unpack_round_trip() {
276 let bs = 1;
277 let nkv = 1;
278 let hd = 16;
279 let head_data_bytes = hd * 3 / 8; let mut raw = vec![0u8; head_data_bytes + 1];
281 let pack8 = |vals: [u8; 8]| -> [u8; 3] {
282 let b0 = vals[0] | (vals[1] << 3) | (vals[2] << 6);
283 let b1 = (vals[2] >> 2) | (vals[3] << 1) | (vals[4] << 4) | (vals[5] << 7);
284 let b2 = (vals[5] >> 1) | (vals[6] << 2) | (vals[7] << 5);
285 [b0, b1, b2]
286 };
287 let t0 = pack8([0, 1, 2, 3, 4, 5, 6, 7]);
288 let t1 = pack8([7, 6, 5, 4, 3, 2, 1, 0]);
289 raw[..3].copy_from_slice(&t0);
290 raw[3..6].copy_from_slice(&t1);
291 raw[head_data_bytes] = FP8_ONE;
292 let mut out = vec![bf16::ZERO; bs * nkv * hd];
293 dequant_turbo3_block_to_bf16(&raw, bs, nkv, hd, &mut out);
294 let expect: Vec<f32> = (0..8u32)
295 .map(|i| TURBO3_LUT[i as usize])
296 .chain((0..8u32).rev().map(|i| TURBO3_LUT[i as usize]))
297 .collect();
298 for (i, e) in expect.iter().enumerate() {
299 assert!(
300 (out[i].to_f32() - e).abs() < 1e-2,
301 "elem {i}: expected {e}, got {}",
302 out[i].to_f32(),
303 );
304 }
305 }
306
307 #[test]
308 #[allow(clippy::needless_range_loop)]
309 fn turbo8_dequant_layout() {
310 let bs = 1;
313 let nkv = 1;
314 let hd = 16;
315 let mut raw = vec![FP8_ONE; hd + 2];
316 let scale_bytes = bf16::from_f32(1.0).to_le_bytes();
318 raw[hd] = scale_bytes[0];
319 raw[hd + 1] = scale_bytes[1];
320 let mut out = vec![bf16::ZERO; bs * nkv * hd];
321 dequant_turbo8_block_to_bf16(&raw, bs, nkv, hd, &mut out);
322 for i in 0..hd {
323 assert!(
324 (out[i].to_f32() - 1.0).abs() < 1e-2,
325 "elem {i}: expected 1.0, got {}",
326 out[i].to_f32(),
327 );
328 }
329 }
330
331 #[test]
332 fn multi_head_multi_token_consistency() {
333 let bs = 2;
336 let nkv = 2;
337 let hd = 16;
338 let head_data_bytes = hd / 2;
339 let head_scale_bytes = hd / NVFP4_GROUP_SIZE;
340 let token_data_stride = nkv * head_data_bytes;
341 let token_scale_stride = nkv * head_scale_bytes;
342 let data_section_bytes = bs * token_data_stride;
343 let scale_section_bytes = bs * token_scale_stride;
344 let mut raw = vec![0u8; data_section_bytes + scale_section_bytes];
345 for tok in 0..bs {
347 for kv_h in 0..nkv {
348 let d_off = tok * token_data_stride + kv_h * head_data_bytes;
349 let nibble_lo = ((tok * 2 + kv_h) as u8) & 0xF;
350 let nibble_hi = (((tok * 2 + kv_h) + 8) as u8) & 0xF;
351 for byte_idx in 0..head_data_bytes {
352 raw[d_off + byte_idx] = nibble_lo | (nibble_hi << 4);
353 }
354 let s_off = data_section_bytes + tok * token_scale_stride + kv_h * head_scale_bytes;
355 raw[s_off] = FP8_ONE;
356 }
357 }
358 let mut out = vec![bf16::ZERO; bs * nkv * hd];
359 dequant_4bit_block_to_bf16(&raw, bs, nkv, hd, &NVFP4_E2M1_LUT, &mut out);
360 for tok in 0..bs {
361 for kv_h in 0..nkv {
362 let nibble_lo = (tok * 2 + kv_h) & 0xF;
363 let nibble_hi = ((tok * 2 + kv_h) + 8) & 0xF;
364 let exp_lo = NVFP4_E2M1_LUT[nibble_lo];
365 let exp_hi = NVFP4_E2M1_LUT[nibble_hi];
366 let base = (tok * nkv + kv_h) * hd;
367 for elem in 0..hd {
368 let expected = if elem % 2 == 0 { exp_lo } else { exp_hi };
369 assert!(
370 (out[base + elem].to_f32() - expected).abs() < 1e-2,
371 "tok={tok} kv_h={kv_h} elem={elem}: expected {expected}, got {}",
372 out[base + elem].to_f32(),
373 );
374 }
375 }
376 }
377 }
378}