Skip to main content

virtio_accel_vulkan/
nvfp4.rs

1//! Provider-native NVFP4 matrix products.
2//!
3//! TOSA has no FP4 tensor type. Encoding packed weights as an integer TOSA
4//! graph would hide the operation from provider admission and explode it into
5//! scalar nodes, so Vulkan owns a small artifact format for the exact storage
6//! contract used by NVFP4 checkpoints. This is an artifact-format extension,
7//! not a virtio-accel wire change.
8
9use virtio_accel_core::{ArtifactFormat, ArtifactRef, TargetIdentity};
10
11use crate::lower::{
12    DispatchPlan, KernelSpec, LoweringError, ProgramPlan, SlotPlan, SlotRole, Work,
13};
14use crate::shader::{Nvfp4MatmulSpec, Operand, Storage};
15
16pub const VULKAN_NVFP4_FORMAT: ArtifactFormat = match ArtifactFormat::new(0x564e_4634) {
17    Some(format) => format,
18    None => panic!("nonzero artifact format"),
19};
20
21pub const VULKAN_NVFP4_TARGET: TargetIdentity =
22    TargetIdentity([0x564b_4e46, 0x5034_0001, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]);
23
24const MAGIC: [u8; 8] = *b"VKNVFP4\0";
25const VERSION: u32 = 1;
26const HEADER_BYTES: usize = 32;
27
28/// Fused projection epilogue. Its integer representation is part of artifact v1.
29#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
30#[repr(u32)]
31pub enum Nvfp4Activation {
32    #[default]
33    None = 0,
34    Silu = 1,
35    Sigmoid = 2,
36    /// `max(x, 0)^2` after the projection has applied `tensor_scale` once.
37    /// A fused routed-expert caller folds the omitted second scale squared
38    /// into the following projection's tensor scale. This keeps the private
39    /// intermediate normal without changing the composed result.
40    SquaredRelu = 3,
41}
42
43const MOE_MAGIC: [u8; 8] = *b"VKNVMOE\0";
44const MOE_HEADER_BYTES: usize = 160;
45
46/// Two native NVFP4 expert projections with a device-local squared-ReLU
47/// intermediate. Each expert occupies one externally bound cache-slot span;
48/// the four offsets identify its up/down packed and block-scale planes.
49#[derive(Clone, Copy, Debug, PartialEq)]
50pub struct Nvfp4MoeArtifact {
51    bytes: [u8; MOE_HEADER_BYTES],
52}
53
54impl Nvfp4MoeArtifact {
55    pub fn new(
56        batch: u32,
57        width: u32,
58        inner: u32,
59        expert_bytes: u32,
60        offsets: &[[u32; 4]],
61    ) -> Result<Self, LoweringError> {
62        if batch == 0
63            || batch > 6
64            || offsets.len() != batch as usize
65            || width % 16 != 0
66            || inner % 16 != 0
67        {
68            return Err(LoweringError::UnsupportedGraph);
69        }
70        let fits = |offset: u32, bytes: u64| {
71            u64::from(offset)
72                .checked_add(bytes)
73                .is_some_and(|last| last <= u64::from(expert_bytes))
74        };
75        for &[up_packed, up_scales, down_packed, down_scales] in offsets {
76            if [up_packed, up_scales, down_packed, down_scales]
77                .into_iter()
78                .any(|offset| offset % 4 != 0)
79                || !fits(up_packed, u64::from(inner) * u64::from(width) / 2)
80                || !fits(up_scales, u64::from(inner) * u64::from(width) / 16)
81                || !fits(down_packed, u64::from(width) * u64::from(inner) / 2)
82                || !fits(down_scales, u64::from(width) * u64::from(inner) / 16)
83            {
84                return Err(LoweringError::ResourceLimit);
85            }
86        }
87        let mut bytes = [0; MOE_HEADER_BYTES];
88        bytes[..8].copy_from_slice(&MOE_MAGIC);
89        for (index, value) in [VERSION, batch, width, inner, expert_bytes]
90            .into_iter()
91            .enumerate()
92        {
93            bytes[8 + index * 4..12 + index * 4].copy_from_slice(&value.to_le_bytes());
94        }
95        for (index, value) in offsets.iter().flatten().copied().enumerate() {
96            bytes[28 + index * 4..32 + index * 4].copy_from_slice(&value.to_le_bytes());
97        }
98        Ok(Self { bytes })
99    }
100
101    pub fn as_ref(&self) -> ArtifactRef<'_> {
102        ArtifactRef {
103            format: VULKAN_NVFP4_FORMAT,
104            target: VULKAN_NVFP4_TARGET,
105            payload: &self.bytes,
106            resident_bytes: crate::REQUIRED_RESIDENT_BYTES,
107        }
108    }
109}
110
111/// A validated F32 × NVFP4 projection artifact.
112#[derive(Clone, Copy, Debug, PartialEq)]
113pub struct Nvfp4Artifact {
114    bytes: [u8; HEADER_BYTES],
115}
116
117impl Nvfp4Artifact {
118    /// Describe `[m,k] F32 × [n,k] NVFP4 -> [m,n] F32`.
119    ///
120    /// The tensor scale is a one-element F32 input rather than artifact metadata, so one loaded
121    /// program serves every checkpoint tensor with this geometry.
122    pub fn new(m: u32, n: u32, k: u32, activation: Nvfp4Activation) -> Result<Self, LoweringError> {
123        Self::build(m, n, k, activation, 0)
124    }
125
126    /// Describe `batch` independent F32 vectors of length `k` multiplied by
127    /// NVFP4 weight matrices with `n` rows and `k` columns.
128    ///
129    /// Activations, packed weights, block scales and tensor scales all carry
130    /// the leading batch dimension. This is the natural execution unit for a
131    /// routed MoE layer and amortizes submission over all selected experts.
132    pub fn new_batched(
133        batch: u32,
134        n: u32,
135        k: u32,
136        activation: Nvfp4Activation,
137    ) -> Result<Self, LoweringError> {
138        Self::build(batch, n, k, activation, 1)
139    }
140
141    /// Describe an expert batch whose packed weights and block scales are
142    /// separate bindings. This retains zero-copy imports for independently
143    /// cached experts while still issuing one dispatch. The Vulkan descriptor
144    /// budget permits up to six experts.
145    pub fn new_expert_batch(
146        batch: u32,
147        n: u32,
148        k: u32,
149        activation: Nvfp4Activation,
150    ) -> Result<Self, LoweringError> {
151        Self::build(batch, n, k, activation, 2)
152    }
153
154    fn build(
155        m: u32,
156        n: u32,
157        k: u32,
158        activation: Nvfp4Activation,
159        weight_mode: u32,
160    ) -> Result<Self, LoweringError> {
161        if m == 0
162            || n == 0
163            || k == 0
164            || k % 16 != 0
165            || weight_mode > 2
166            || (weight_mode == 2 && m > 6)
167        {
168            return Err(LoweringError::UnsupportedGraph);
169        }
170        checked_lengths(m, n, k, weight_mode != 0)?;
171        let mut bytes = [0; HEADER_BYTES];
172        bytes[..8].copy_from_slice(&MAGIC);
173        for (offset, value) in [VERSION, m, n, k, activation as u32, weight_mode]
174            .into_iter()
175            .enumerate()
176        {
177            bytes[8 + offset * 4..12 + offset * 4].copy_from_slice(&value.to_le_bytes());
178        }
179        Ok(Self { bytes })
180    }
181
182    pub fn as_bytes(&self) -> &[u8] {
183        &self.bytes
184    }
185
186    pub fn as_ref(&self) -> ArtifactRef<'_> {
187        ArtifactRef {
188            format: VULKAN_NVFP4_FORMAT,
189            target: VULKAN_NVFP4_TARGET,
190            payload: &self.bytes,
191            resident_bytes: crate::REQUIRED_RESIDENT_BYTES,
192        }
193    }
194}
195
196fn checked_lengths(
197    m: u32,
198    n: u32,
199    k: u32,
200    batched_weights: bool,
201) -> Result<[u64; 5], LoweringError> {
202    let product = |a: u32, b: u32, bytes: u64| {
203        u64::from(a)
204            .checked_mul(u64::from(b))
205            .and_then(|v| v.checked_mul(bytes))
206            .ok_or(LoweringError::ResourceLimit)
207    };
208    let batches = if batched_weights { m } else { 1 };
209    Ok([
210        product(m, k, 4)?,
211        product(n, k, u64::from(batches))? / 2,
212        product(n, k, u64::from(batches))? / 16,
213        u64::from(batches) * 4,
214        product(m, n, 4)?,
215    ])
216}
217
218fn word(bytes: &[u8], offset: usize) -> Result<u32, LoweringError> {
219    bytes
220        .get(offset..offset + 4)
221        .and_then(|v| v.try_into().ok())
222        .map(u32::from_le_bytes)
223        .ok_or(LoweringError::UnsupportedGraph)
224}
225
226pub(crate) fn lower_nvfp4(bytes: &[u8]) -> Result<ProgramPlan, LoweringError> {
227    if bytes.len() == MOE_HEADER_BYTES && bytes[..8] == MOE_MAGIC {
228        return lower_nvfp4_moe(bytes);
229    }
230    if bytes.len() != HEADER_BYTES || bytes[..8] != MAGIC {
231        return Err(LoweringError::UnsupportedGraph);
232    }
233    let version = word(bytes, 8)?;
234    let m = word(bytes, 12)?;
235    let n = word(bytes, 16)?;
236    let k = word(bytes, 20)?;
237    let activation = word(bytes, 24)?;
238    let weight_mode = word(bytes, 28)?;
239    if version != VERSION
240        || activation > Nvfp4Activation::SquaredRelu as u32
241        || weight_mode > 2
242        || (weight_mode == 2 && m > 6)
243        || m == 0
244        || n == 0
245        || k == 0
246        || k % 16 != 0
247    {
248        return Err(LoweringError::UnsupportedGraph);
249    }
250    let lengths = checked_lengths(m, n, k, weight_mode != 0)?;
251    let operand = |slot| Operand {
252        buffer: slot,
253        base: 0,
254    };
255    if weight_mode == 2 {
256        let packed_bytes = u64::from(n) * u64::from(k) / 2;
257        let scale_bytes = u64::from(n) * u64::from(k) / 16;
258        let mut slots = vec![SlotPlan {
259            slot: 0,
260            role: SlotRole::Input,
261            byte_len: lengths[0],
262            storage: Storage::Word,
263        }];
264        let mut packed = Vec::with_capacity(m as usize);
265        let mut scales = Vec::with_capacity(m as usize);
266        for expert in 0..m {
267            let packed_slot = 1 + expert * 2;
268            let scale_slot = packed_slot + 1;
269            slots.push(SlotPlan {
270                slot: packed_slot,
271                role: SlotRole::Input,
272                byte_len: packed_bytes,
273                storage: Storage::Byte,
274            });
275            slots.push(SlotPlan {
276                slot: scale_slot,
277                role: SlotRole::Input,
278                byte_len: scale_bytes,
279                storage: Storage::Byte,
280            });
281            packed.push(operand(packed_slot));
282            scales.push(operand(scale_slot));
283        }
284        let tensor_scale_slot = 1 + m * 2;
285        let output_slot = tensor_scale_slot + 1;
286        slots.push(SlotPlan {
287            slot: tensor_scale_slot,
288            role: SlotRole::Input,
289            byte_len: u64::from(m) * 4,
290            storage: Storage::Word,
291        });
292        slots.push(SlotPlan {
293            slot: output_slot,
294            role: SlotRole::Output,
295            byte_len: lengths[4],
296            storage: Storage::Word,
297        });
298        return Ok(ProgramPlan {
299            slots,
300            arena_bytes: 0,
301            constants: Vec::new(),
302            dispatches: vec![DispatchPlan {
303                kernel: KernelSpec::Nvfp4Matmul { cooperative: false },
304                spec: Nvfp4MatmulSpec {
305                    activation: operand(0),
306                    packed: &packed,
307                    block_scales: &scales,
308                    tensor_scale: operand(tensor_scale_slot),
309                    output: operand(output_slot),
310                    m,
311                    n,
312                    k,
313                    epilogue: activation,
314                    weight_mode,
315                }
316                .words(),
317                work: Work::Nvfp4Matmul { m, n },
318                barrier_before: false,
319            }],
320        });
321    }
322    Ok(ProgramPlan {
323        slots: vec![
324            SlotPlan {
325                slot: 0,
326                role: SlotRole::Input,
327                byte_len: lengths[0],
328                storage: Storage::Word,
329            },
330            SlotPlan {
331                slot: 1,
332                role: SlotRole::Input,
333                byte_len: lengths[1],
334                storage: Storage::Byte,
335            },
336            SlotPlan {
337                slot: 2,
338                role: SlotRole::Input,
339                byte_len: lengths[2],
340                storage: Storage::Byte,
341            },
342            SlotPlan {
343                slot: 3,
344                role: SlotRole::Input,
345                byte_len: lengths[3],
346                storage: Storage::Word,
347            },
348            SlotPlan {
349                slot: 4,
350                role: SlotRole::Output,
351                byte_len: lengths[4],
352                storage: Storage::Word,
353            },
354        ],
355        arena_bytes: 0,
356        constants: Vec::new(),
357        dispatches: vec![DispatchPlan {
358            kernel: KernelSpec::Nvfp4Matmul {
359                cooperative: weight_mode == 0 && m >= 8,
360            },
361            spec: Nvfp4MatmulSpec {
362                activation: operand(0),
363                packed: &[operand(1)],
364                block_scales: &[operand(2)],
365                tensor_scale: operand(3),
366                output: operand(4),
367                m,
368                n,
369                k,
370                epilogue: activation,
371                weight_mode,
372            }
373            .words(),
374            work: Work::Nvfp4Matmul { m, n },
375            barrier_before: false,
376        }],
377    })
378}
379
380fn lower_nvfp4_moe(bytes: &[u8]) -> Result<ProgramPlan, LoweringError> {
381    let values: Vec<u32> = (0..5)
382        .map(|i| word(bytes, 8 + i * 4))
383        .collect::<Result<_, _>>()?;
384    let [version, batch, width, inner, expert_bytes] = values.as_slice() else {
385        return Err(LoweringError::UnsupportedGraph);
386    };
387    if *version != VERSION || *batch == 0 || *batch > 6 || *expert_bytes == 0 {
388        return Err(LoweringError::UnsupportedGraph);
389    }
390    let offsets = (0..*batch as usize)
391        .map(|expert| {
392            Ok([
393                word(bytes, 28 + (expert * 4) * 4)?,
394                word(bytes, 28 + (expert * 4 + 1) * 4)?,
395                word(bytes, 28 + (expert * 4 + 2) * 4)?,
396                word(bytes, 28 + (expert * 4 + 3) * 4)?,
397            ])
398        })
399        .collect::<Result<Vec<[u32; 4]>, LoweringError>>()?;
400    let artifact = Nvfp4MoeArtifact::new(*batch, *width, *inner, *expert_bytes, &offsets)?;
401    let _ = artifact;
402    let activation_bytes = u64::from(*batch) * u64::from(*width) * 4;
403    let intermediate_bytes = u64::from(*batch) * u64::from(*inner) * 4;
404    let output_bytes = activation_bytes;
405    let mut slots = vec![SlotPlan {
406        slot: 0,
407        role: SlotRole::Input,
408        byte_len: activation_bytes,
409        storage: Storage::Word,
410    }];
411    for expert in 0..*batch {
412        slots.push(SlotPlan {
413            slot: 1 + expert,
414            role: SlotRole::Input,
415            byte_len: u64::from(*expert_bytes),
416            storage: Storage::Byte,
417        });
418    }
419    let up_tensor_slot = 1 + *batch;
420    let down_tensor_slot = up_tensor_slot + 1;
421    let output_slot = down_tensor_slot + 1;
422    for slot in [up_tensor_slot, down_tensor_slot] {
423        slots.push(SlotPlan {
424            slot,
425            role: SlotRole::Input,
426            byte_len: u64::from(*batch) * 4,
427            storage: Storage::Word,
428        });
429    }
430    slots.push(SlotPlan {
431        slot: output_slot,
432        role: SlotRole::Output,
433        byte_len: output_bytes,
434        storage: Storage::Word,
435    });
436    let arena_slot = slots.len() as u32;
437    let operands = |plane: usize| {
438        (0..*batch)
439            .map(|expert| Operand {
440                buffer: 1 + expert,
441                // SPIR-V storage operands index 32-bit words; the artifact
442                // exposes byte offsets to match host tensor views.
443                base: offsets[expert as usize][plane] / 4,
444            })
445            .collect::<Vec<_>>()
446    };
447    Ok(ProgramPlan {
448        slots,
449        arena_bytes: intermediate_bytes,
450        constants: Vec::new(),
451        dispatches: vec![
452            DispatchPlan {
453                kernel: KernelSpec::Nvfp4Matmul { cooperative: false },
454                spec: Nvfp4MatmulSpec {
455                    activation: Operand { buffer: 0, base: 0 },
456                    packed: &operands(0),
457                    block_scales: &operands(1),
458                    tensor_scale: Operand {
459                        buffer: up_tensor_slot,
460                        base: 0,
461                    },
462                    output: Operand {
463                        buffer: arena_slot,
464                        base: 0,
465                    },
466                    m: *batch,
467                    n: *inner,
468                    k: *width,
469                    epilogue: Nvfp4Activation::SquaredRelu as u32,
470                    weight_mode: 2,
471                }
472                .words(),
473                work: Work::Nvfp4Matmul {
474                    m: *batch,
475                    n: *inner,
476                },
477                barrier_before: false,
478            },
479            DispatchPlan {
480                kernel: KernelSpec::Nvfp4Matmul { cooperative: false },
481                spec: Nvfp4MatmulSpec {
482                    activation: Operand {
483                        buffer: arena_slot,
484                        base: 0,
485                    },
486                    packed: &operands(2),
487                    block_scales: &operands(3),
488                    tensor_scale: Operand {
489                        buffer: down_tensor_slot,
490                        base: 0,
491                    },
492                    output: Operand {
493                        buffer: output_slot,
494                        base: 0,
495                    },
496                    m: *batch,
497                    n: *width,
498                    k: *inner,
499                    epilogue: Nvfp4Activation::None as u32,
500                    weight_mode: 2,
501                }
502                .words(),
503                work: Work::Nvfp4Matmul {
504                    m: *batch,
505                    n: *width,
506                },
507                barrier_before: true,
508            },
509        ],
510    })
511}
512
513#[cfg(test)]
514mod tests {
515    use super::*;
516
517    #[test]
518    fn artifact_round_trips_into_exact_buffer_contract() {
519        let artifact = Nvfp4Artifact::new(3, 17, 2688, Nvfp4Activation::Silu).unwrap();
520        let plan = lower_nvfp4(artifact.as_bytes()).unwrap();
521        assert_eq!(
522            plan.slots
523                .iter()
524                .map(|slot| slot.byte_len)
525                .collect::<Vec<_>>(),
526            vec![3 * 2688 * 4, 17 * 2688 / 2, 17 * 2688 / 16, 4, 3 * 17 * 4,]
527        );
528        assert_eq!(plan.dispatches[0].work, Work::Nvfp4Matmul { m: 3, n: 17 });
529    }
530
531    #[test]
532    fn batched_artifact_batches_weights_scales_and_tensor_scales() {
533        let artifact = Nvfp4Artifact::new_batched(6, 1856, 2688, Nvfp4Activation::None).unwrap();
534        let plan = lower_nvfp4(artifact.as_bytes()).unwrap();
535        assert_eq!(
536            plan.slots
537                .iter()
538                .map(|slot| slot.byte_len)
539                .collect::<Vec<_>>(),
540            vec![
541                6 * 2688 * 4,
542                6 * 1856 * 2688 / 2,
543                6 * 1856 * 2688 / 16,
544                6 * 4,
545                6 * 1856 * 4,
546            ]
547        );
548    }
549
550    #[test]
551    fn expert_batch_keeps_each_weight_pair_in_its_own_binding() {
552        let artifact =
553            Nvfp4Artifact::new_expert_batch(6, 1856, 2688, Nvfp4Activation::None).unwrap();
554        let plan = lower_nvfp4(artifact.as_bytes()).unwrap();
555        assert_eq!(plan.slots.len(), 15);
556        assert_eq!(plan.slots[0].byte_len, 6 * 2688 * 4);
557        for expert in 0..6 {
558            assert_eq!(plan.slots[1 + expert * 2].byte_len, 1856 * 2688 / 2);
559            assert_eq!(plan.slots[2 + expert * 2].byte_len, 1856 * 2688 / 16);
560        }
561        assert_eq!(plan.slots[13].byte_len, 6 * 4);
562        assert_eq!(plan.slots[14].byte_len, 6 * 1856 * 4);
563    }
564
565    #[test]
566    fn fused_moe_uses_one_expert_binding_and_a_private_intermediate() {
567        let width = 2688;
568        let inner = 1856;
569        let up_bytes = inner * width / 2;
570        let down_bytes = width * inner / 2;
571        let scale_bytes = inner * width / 16;
572        let artifact = Nvfp4MoeArtifact::new(
573            6,
574            width,
575            inner,
576            6 << 20,
577            &[[0, up_bytes, 3 << 20, (3 << 20) + down_bytes]; 6],
578        )
579        .unwrap();
580        assert!(scale_bytes < 1 << 20);
581        let plan = lower_nvfp4(&artifact.bytes).unwrap();
582        assert_eq!(plan.slots.len(), 10);
583        assert_eq!(plan.arena_bytes, 6 * u64::from(inner) * 4);
584        assert_eq!(plan.dispatches.len(), 2);
585        assert!(!plan.dispatches[0].barrier_before);
586        assert!(plan.dispatches[1].barrier_before);
587        // First packed operand follows the activation operand in the
588        // specialization payload and names byte offset zero in words.
589        assert_eq!(plan.dispatches[0].spec[2..4], [1, 0]);
590        // The first down packed operand begins at 3 MiB, encoded as words.
591        assert_eq!(plan.dispatches[1].spec[2..4], [1, (3 << 20) / 4]);
592    }
593
594    #[test]
595    fn malformed_geometry_is_rejected() {
596        assert!(Nvfp4Artifact::new(1, 1, 15, Nvfp4Activation::None).is_err());
597    }
598}