1use 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#[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 SquaredRelu = 3,
41}
42
43const MOE_MAGIC: [u8; 8] = *b"VKNVMOE\0";
44const MOE_HEADER_BYTES: usize = 160;
45
46#[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#[derive(Clone, Copy, Debug, PartialEq)]
113pub struct Nvfp4Artifact {
114 bytes: [u8; HEADER_BYTES],
115}
116
117impl Nvfp4Artifact {
118 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 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 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 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 assert_eq!(plan.dispatches[0].spec[2..4], [1, 0]);
590 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}