spark_runtime/
kv_dequant.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Host-side dequantization of paged KV cache blocks → BF16.
4//!
5//! Used by `--high-speed-swap` Phase 6.2.c to produce BF16 source data for
6//! the orchestrator's tile-streaming attention kernel from any of Atlas's
7//! quantized KV layouts. The kernel-side packing/scale layouts mirrored
8//! here:
9//!
10//! | Quant   | data bytes/elem | scale | LUT                             | source kernel                      |
11//! |---------|-----------------|-------|---------------------------------|------------------------------------|
12//! | BF16    | 2               | none  | identity                        | (direct copy, not in this module)  |
13//! | FP8     | 1 (E4M3)        | tensor| `e4m3_lut`                      | reshape_and_cache_fp8.cu           |
14//! | NVFP4   | 0.5 (4-bit)     | group | `NVFP4_E2M1_LUT`                | paged_decode_attn_nvfp4.cu         |
15//! | Turbo4  | 0.5 (4-bit)     | group | `TURBO4_LUT` (Lloyd-Max 16)     | paged_decode_attn_turbo4.cu        |
16//! | Turbo3  | 0.375 (3-bit)   | group | `TURBO3_LUT` (Lloyd-Max 8)      | paged_decode_attn_turbo3.cu        |
17//! | Turbo8  | 1 (E4M3)        | group | `e4m3_lut`                      | paged_decode_attn_turbo8.cu        |
18//!
19//! "Group" scales cover 16 elements per scale byte (`NVFP4_GROUP_SIZE`) and
20//! are stored in a separate section after the data section within each
21//! block. All LUTs match their CUDA-side counterparts byte-for-byte.
22
23use 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
29/// Dequant FP8 (E4M3) bytes to BF16, applying a per-tensor scale.
30pub 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
38/// Dequant a 4-bit packed (NVFP4 or Turbo4) KV block to BF16.
39///
40/// Layout per block:
41///   data:   `bs * nkv * (hd / 2)` bytes (2 nibbles/byte).
42///   scales: `bs * nkv * (hd / NVFP4_GROUP_SIZE)` bytes (1 FP8 scale per group).
43pub 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
83/// Dequant a Turbo3 (3-bit packed) KV block to BF16.
84///
85/// Layout per block:
86///   data:   `bs * nkv * (hd * 3 / 8)` bytes — 8 values packed in 3 bytes.
87///   scales: `bs * nkv * (hd / NVFP4_GROUP_SIZE)` bytes.
88/// Bit packing for 8 vals v0..v7 in 3 bytes b0,b1,b2 (mirrors
89/// `kernels/gb10/common/paged_decode_attn_turbo3.cu:67-75`):
90///   v0 = b0 & 0x7
91///   v1 = (b0 >> 3) & 0x7
92///   v2 = ((b0 >> 6) | (b1 << 2)) & 0x7
93///   v3 = (b1 >> 1) & 0x7
94///   v4 = (b1 >> 4) & 0x7
95///   v5 = ((b1 >> 7) | (b2 << 1)) & 0x7
96///   v6 = (b2 >> 2) & 0x7
97///   v7 = (b2 >> 5) & 0x7
98pub 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
148/// Dequant a Turbo8 (FP8 E4M3 data + per-group **BF16** scales) KV block to BF16.
149///
150/// Layout per block (post 2026-04-28 BF16-scale upgrade):
151///   data:   `bs * nkv * hd` bytes (1 FP8 byte per element).
152///   scales: `bs * nkv * (hd / NVFP4_GROUP_SIZE) * 2` bytes (BF16 = 2 bytes/scale).
153/// Mirrors `kernels/gb10/common/paged_decode_attn_turbo8*.cu` post-upgrade.
154pub 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; // BF16 scales: 2 bytes per group
165    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                // Read BF16 scale (2 bytes, little-endian).
179                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    /// FP8 E4M3 byte for +1.0: sign=0, exp=7, mantissa=0 → 0b00111000 = 0x38.
195    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); // +1.0
203        assert!((lut[0xB8] + 1.0).abs() < 1e-6); // -1.0
204        assert!((lut[0x3F] - 1.875).abs() < 1e-6); // mantissa max
205        assert!((lut[0x40] - 2.0).abs() < 1e-6); // exp +1
206        assert!((lut[0x78] - 256.0).abs() < 1e-6); // exp +8
207        assert!((lut[0x7E] - 448.0).abs() < 1e-6); // max finite
208        assert!(lut[0x7F].is_nan());
209        assert!(lut[0xFF].is_nan());
210        assert!((lut[0x01] - (1.0 / 512.0)).abs() < 1e-9); // smallest subnormal
211    }
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        // Same nibble pattern, different LUT — verifies the LUT is honored.
254        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; // 6
280        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        // 2026-04-28: Turbo8 scales are BF16 (2 bytes), not FP8 (1 byte).
311        // 1 token × 1 kv_head × hd=16, 1 group of 16. Total bytes = 16 (data) + 2 (scale).
312        let bs = 1;
313        let nkv = 1;
314        let hd = 16;
315        let mut raw = vec![FP8_ONE; hd + 2];
316        // Scale = 1.0 in BF16 (= 0x3F80 little-endian bytes [0x80, 0x3F]).
317        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        // 2 tokens × 2 kv_heads × hd=16 → 4 (token,head) groups, each 8 data
334        // bytes + 1 scale byte. Deterministically build & verify the indexing.
335        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        // Each (tok, kv_head) gets its own marker pattern.
346        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}