1use virtio_accel_tosa::{
14 AnalyzedValueKind, CapabilityDescriptor, DType, DTypeCapability, ExtensionSet,
15 GraphCapabilities, Level, NanPropagationMode, Op, OpAttributes, OperatorCapability,
16 OperatorConstraints, ProfileSet, RoundingMode, RuntimeConditionSupport, Target, ValueRoles,
17 Version, parse,
18};
19
20pub const XDNA_TOSA_TARGET: Target = Target::new(
25 Version::TOSA_1_0,
26 ProfileSet::FLOATING_POINT,
27 Level::Level8K,
28 ExtensionSet::BF16,
29);
30
31pub const XDNA_TOSA_FP8_TARGET: Target = Target::new(
37 Version::TOSA_1_0,
38 ProfileSet::FLOATING_POINT,
39 Level::Level8K,
40 ExtensionSet::BF16
41 .union(ExtensionSet::FP8E4M3)
42 .union(ExtensionSet::FP8E5M2),
43);
44
45pub const XDNA_TOSA_INTEGER_TARGET: Target = Target::new(
50 Version::TOSA_1_0,
51 ProfileSet::INTEGER,
52 Level::Level8K,
53 ExtensionSet::NONE,
54);
55
56const BF16_DTYPES: &[DTypeCapability] = &[
57 DTypeCapability::new(DType::BF16, ValueRoles::ALL),
58 DTypeCapability::new(DType::FP32, ValueRoles::OUTPUT),
59];
60
61const BF16_OPERATORS: &[OperatorCapability] = &[
62 OperatorCapability::new(Op::CONST),
63 OperatorCapability::new(Op::IDENTITY),
64 OperatorCapability::constrained(Op::MATMUL, OperatorConstraints::ZERO_ZERO_POINTS),
65 OperatorCapability::constrained(
66 Op::MAX_POOL2D,
67 OperatorConstraints::PROPAGATING_NAN.union(OperatorConstraints::ZERO_PADDING),
68 ),
69];
70
71pub const XDNA_TOSA_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
73 target: XDNA_TOSA_TARGET,
74 dtypes: BF16_DTYPES,
75 operators: BF16_OPERATORS,
76 graph: GraphCapabilities {
77 max_regions: 1,
78 max_blocks: 1,
79 dynamic_shapes: false,
80 runtime_conditions: RuntimeConditionSupport::None,
81 },
82};
83
84const FP8_STORAGE_DTYPES: &[DTypeCapability] = &[
88 DTypeCapability::new(DType::FP8E4M3, ValueRoles::INPUT),
89 DTypeCapability::new(DType::FP8E5M2, ValueRoles::INPUT),
90 DTypeCapability::new(
91 DType::BF16,
92 ValueRoles::OUTPUT.union(ValueRoles::INTERMEDIATE),
93 ),
94 DTypeCapability::new(DType::FP32, ValueRoles::OUTPUT),
95];
96
97const FP8_STORAGE_OPERATORS: &[OperatorCapability] = &[
98 OperatorCapability::new(Op::CAST),
99 OperatorCapability::new(Op::CONST),
100 OperatorCapability::constrained(Op::MATMUL, OperatorConstraints::ZERO_ZERO_POINTS),
101];
102
103pub const XDNA_TOSA_FP8_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
105 target: XDNA_TOSA_FP8_TARGET,
106 dtypes: FP8_STORAGE_DTYPES,
107 operators: FP8_STORAGE_OPERATORS,
108 graph: GraphCapabilities {
109 max_regions: 1,
110 max_blocks: 1,
111 dynamic_shapes: false,
112 runtime_conditions: RuntimeConditionSupport::None,
113 },
114};
115
116const INTEGER_DTYPES: &[DTypeCapability] = &[
122 DTypeCapability::new(DType::INT8, ValueRoles::ALL),
123 DTypeCapability::new(
124 DType::INT32,
125 ValueRoles::INPUT
126 .union(ValueRoles::OUTPUT)
127 .union(ValueRoles::CONSTANT)
128 .union(ValueRoles::INTERMEDIATE),
129 ),
130];
131
132const INTEGER_OPERATORS: &[OperatorCapability] = &[
133 OperatorCapability::new(Op::CONST),
134 OperatorCapability::new(Op::IDENTITY),
135 OperatorCapability::new(Op::MATMUL),
136 OperatorCapability::new(Op::RESCALE),
137];
138
139pub const XDNA_TOSA_INTEGER_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
145 target: XDNA_TOSA_INTEGER_TARGET,
146 dtypes: INTEGER_DTYPES,
147 operators: INTEGER_OPERATORS,
148 graph: GraphCapabilities {
149 max_regions: 1,
150 max_blocks: 1,
151 dynamic_shapes: false,
152 runtime_conditions: RuntimeConditionSupport::None,
153 },
154};
155
156pub(crate) const IDENTITY_LINE_SIZE: usize = 1024;
159
160pub(crate) const FP8_CAST_LINE_SIZE: usize = 1024;
162
163pub(crate) const INT8_IDENTITY_MAX_LINE_SIZE: usize = 1024;
169
170pub(crate) const MATMUL_TILE_M: usize = 32;
179pub(crate) const MATMUL_TILE_K: usize = 64;
180pub(crate) const MATMUL_TILE_N: usize = 32;
181
182pub(crate) const MATMUL_MAX_DIM: usize = 512;
185
186pub(crate) const INT8_MATMUL_MAX_TOTAL_BYTES: usize = 16 * 1024;
192
193pub(crate) const INT8_RESCALE_MAX_TOTAL_BYTES: usize = 16 * 1024;
198
199pub(crate) const MAX_POOL_MAX_KERNEL: usize = 8;
202pub(crate) const MAX_POOL_MAX_STRIDE: usize = 8;
203
204pub(crate) const MAX_POOL_MAX_TOTAL_ELEMENTS: usize = 8 * 1024;
209
210#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
213pub enum CompilerSpec {
214 Identity { elements: usize },
217 Int8Identity { elements: usize, line_size: usize },
220 Fp8ToBf16 { format: Fp8Format, elements: usize },
223 Matmul { m: usize, k: usize, n: usize },
227 Fp8Matmul {
235 format: Fp8Format,
236 m: usize,
237 k: usize,
238 n: usize,
239 },
240 Int8Matmul {
245 m: usize,
246 k: usize,
247 n: usize,
248 left_zero_point: i8,
249 right_zero_point: i8,
250 },
251 Int32ToInt8Rescale {
253 elements: usize,
254 multiplier: i32,
255 shift: i8,
256 output_zero_point: i8,
257 },
258 MaxPool2d {
261 input_h: usize,
262 input_w: usize,
263 channels: usize,
264 output_h: usize,
265 output_w: usize,
266 kernel_h: usize,
267 kernel_w: usize,
268 stride_h: usize,
269 stride_w: usize,
270 },
271}
272
273#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
275pub enum Fp8Format {
276 E4M3,
277 E5M2,
278}
279
280#[derive(Clone, Copy, Debug, PartialEq, Eq)]
282pub enum AdmitError {
283 Parse,
285 Analysis,
287 Unsupported,
289}
290
291impl core::fmt::Display for AdmitError {
292 fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
293 write!(formatter, "{self:?}")
294 }
295}
296
297impl std::error::Error for AdmitError {}
298
299impl From<AdmitError> for virtio_accel_core::BackendError {
305 fn from(error: AdmitError) -> Self {
306 match error {
307 AdmitError::Parse | AdmitError::Analysis => Self::InvalidArgument,
308 AdmitError::Unsupported => Self::Unsupported,
309 }
310 }
311}
312
313pub fn admit(bytes: &[u8], target: Target) -> Result<CompilerSpec, AdmitError> {
330 if target != XDNA_TOSA_TARGET
331 && target != XDNA_TOSA_FP8_TARGET
332 && target != XDNA_TOSA_INTEGER_TARGET
333 {
334 return Err(AdmitError::Unsupported);
335 }
336 let model = parse(bytes).map_err(|_| AdmitError::Parse)?;
337 let analysis = model
338 .analyze_for(target)
339 .map_err(|_| AdmitError::Analysis)?;
340
341 if analysis.regions().len() != 1
342 || analysis.blocks().len() != 1
343 || !analysis.conditions().is_empty()
344 {
345 return Err(AdmitError::Unsupported);
346 }
347 let block = analysis.blocks()[0].id();
348
349 let mut matmul = None;
353 let mut max_pool = None;
354 let mut casts: [Option<virtio_accel_tosa::OperatorId>; 2] = [None, None];
355 let mut cast_count = 0usize;
356 let mut rescale = None;
357 let mut identities = 0usize;
358 let mut constants = 0usize;
359 for operator in analysis.execution_order(block) {
360 match analysis.operator(*operator).op() {
361 Op::IDENTITY => identities += 1,
362 Op::CONST => constants += 1,
363 Op::MATMUL if matmul.is_none() => matmul = Some(*operator),
364 Op::MAX_POOL2D if max_pool.is_none() => max_pool = Some(*operator),
365 Op::CAST if cast_count < casts.len() => {
366 casts[cast_count] = Some(*operator);
367 cast_count += 1;
368 }
369 Op::RESCALE if rescale.is_none() => rescale = Some(*operator),
370 _ => return Err(AdmitError::Unsupported),
371 }
372 }
373 match (
374 target, matmul, max_pool, cast_count, rescale, identities, constants,
375 ) {
376 (XDNA_TOSA_TARGET, None, None, 0, None, _, 0) => admit_identity(&analysis, block),
379 (XDNA_TOSA_TARGET, Some(matmul), None, 0, None, 0, _) => {
380 admit_matmul(&analysis, block, matmul)
381 }
382 (XDNA_TOSA_TARGET, None, Some(max_pool), 0, None, 0, 0) => {
383 admit_max_pool2d(&analysis, block, max_pool)
384 }
385 (XDNA_TOSA_FP8_TARGET, None, None, 1, None, 0, 0) => {
386 admit_fp8_to_bf16(&analysis, block, casts[0].ok_or(AdmitError::Unsupported)?)
387 }
388 (XDNA_TOSA_FP8_TARGET, Some(matmul), None, 2, None, 0, _) => admit_fp8_matmul(
390 &analysis,
391 block,
392 matmul,
393 [
394 casts[0].ok_or(AdmitError::Unsupported)?,
395 casts[1].ok_or(AdmitError::Unsupported)?,
396 ],
397 ),
398 (XDNA_TOSA_INTEGER_TARGET, None, None, 0, None, _, 0) => {
399 admit_int8_identity(&analysis, block)
400 }
401 (XDNA_TOSA_INTEGER_TARGET, Some(matmul), None, 0, None, 0, _) => {
402 admit_int8_matmul(&analysis, block, matmul)
403 }
404 (XDNA_TOSA_INTEGER_TARGET, None, None, 0, Some(rescale), 0, 4) => {
405 admit_int32_to_int8_rescale(&analysis, block, rescale)
406 }
407 _ => Err(AdmitError::Unsupported),
408 }
409}
410
411fn admit_int8_identity(
416 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
417 block: virtio_accel_tosa::BlockId,
418) -> Result<CompilerSpec, AdmitError> {
419 let inputs = analysis.block_inputs(block);
420 let outputs = analysis.block_outputs(block);
421 if inputs.len() != 1 || outputs.len() != 1 {
422 return Err(AdmitError::Unsupported);
423 }
424 for value in analysis.values() {
425 if let AnalyzedValueKind::Tensor(tensor) = value.kind() {
426 if tensor.dtype() != DType::INT8 {
427 return Err(AdmitError::Unsupported);
428 }
429 }
430 }
431
432 let elements = tensor_elements(analysis, outputs[0])?;
433 let line_size = elements.min(INT8_IDENTITY_MAX_LINE_SIZE);
434 if line_size % 4 != 0 || (elements > INT8_IDENTITY_MAX_LINE_SIZE && elements % line_size != 0) {
435 return Err(AdmitError::Unsupported);
436 }
437 Ok(CompilerSpec::Int8Identity {
438 elements,
439 line_size,
440 })
441}
442
443fn admit_int8_matmul(
449 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
450 block: virtio_accel_tosa::BlockId,
451 matmul: virtio_accel_tosa::OperatorId,
452) -> Result<CompilerSpec, AdmitError> {
453 let inputs = analysis.operator_inputs(matmul);
454 let outputs = analysis.operator_outputs(matmul);
455 if inputs.len() != 4
456 || outputs.len() != 1
457 || inputs[0] == inputs[1]
461 || analysis.block_inputs(block) != [inputs[0], inputs[1]]
462 || analysis.block_outputs(block) != [outputs[0]]
463 {
464 return Err(AdmitError::Unsupported);
465 }
466 for operator in analysis.execution_order(block) {
467 if analysis.operator(*operator).op() != Op::CONST {
468 continue;
469 }
470 for produced in analysis.operator_outputs(*operator) {
471 if *produced != inputs[2] && *produced != inputs[3] {
472 return Err(AdmitError::Unsupported);
473 }
474 }
475 }
476
477 let lhs = matmul_dims(analysis, inputs[0], DType::INT8)?;
478 let rhs = matmul_dims(analysis, inputs[1], DType::INT8)?;
479 let out = matmul_dims(analysis, outputs[0], DType::INT32)?;
480 let ([1, m, k], [1, k2, n], [1, m2, n2]) = (lhs, rhs, out) else {
481 return Err(AdmitError::Unsupported);
482 };
483 if k != k2 || m != m2 || n != n2 || [m, k, n].iter().any(|dim| *dim > MATMUL_MAX_DIM) {
484 return Err(AdmitError::Unsupported);
485 }
486
487 let lhs_bytes = align_to_four(m.checked_mul(k).ok_or(AdmitError::Unsupported)?)?;
490 let rhs_bytes = align_to_four(k.checked_mul(n).ok_or(AdmitError::Unsupported)?)?;
491 let output_bytes = m
492 .checked_mul(n)
493 .and_then(|elements| elements.checked_mul(4))
494 .ok_or(AdmitError::Unsupported)?;
495 if lhs_bytes
496 .checked_add(rhs_bytes)
497 .and_then(|bytes| bytes.checked_add(output_bytes))
498 .is_none_or(|bytes| bytes > INT8_MATMUL_MAX_TOTAL_BYTES)
499 {
500 return Err(AdmitError::Unsupported);
501 }
502
503 Ok(CompilerSpec::Int8Matmul {
504 m,
505 k,
506 n,
507 left_zero_point: int8_zero_point(analysis, inputs[2])?,
508 right_zero_point: int8_zero_point(analysis, inputs[3])?,
509 })
510}
511
512fn admit_int32_to_int8_rescale(
519 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
520 block: virtio_accel_tosa::BlockId,
521 rescale: virtio_accel_tosa::OperatorId,
522) -> Result<CompilerSpec, AdmitError> {
523 let inputs = analysis.operator_inputs(rescale);
524 let outputs = analysis.operator_outputs(rescale);
525 if inputs.len() != 5
526 || outputs.len() != 1
527 || analysis.block_inputs(block) != [inputs[0]]
528 || analysis.block_outputs(block) != [outputs[0]]
529 {
530 return Err(AdmitError::Unsupported);
531 }
532 for operator in analysis.execution_order(block) {
533 if analysis.operator(*operator).op() != Op::CONST {
534 continue;
535 }
536 for produced in analysis.operator_outputs(*operator) {
537 if !inputs[1..].contains(produced) {
538 return Err(AdmitError::Unsupported);
539 }
540 }
541 }
542
543 let AnalyzedValueKind::Tensor(input) = analysis.value(inputs[0]).kind() else {
544 return Err(AdmitError::Unsupported);
545 };
546 let AnalyzedValueKind::Tensor(output) = analysis.value(outputs[0]).kind() else {
547 return Err(AdmitError::Unsupported);
548 };
549 if input.dtype() != DType::INT32
550 || output.dtype() != DType::INT8
551 || !input.dimensions().eq(output.dimensions())
552 {
553 return Err(AdmitError::Unsupported);
554 }
555 let OpAttributes::Rescale {
556 scale32,
557 rounding_mode,
558 per_channel,
559 input_unsigned,
560 output_unsigned,
561 } = analysis.operator(rescale).source().attributes()
562 else {
563 return Err(AdmitError::Unsupported);
564 };
565 if !scale32
566 || rounding_mode != RoundingMode::SINGLE_ROUND
567 || per_channel
568 || input_unsigned
569 || output_unsigned
570 {
571 return Err(AdmitError::Unsupported);
572 }
573
574 let multiplier = int32_constant(analysis, inputs[1])?;
575 let shift = int8_zero_point(analysis, inputs[2])?;
576 let input_zero_point = int32_constant(analysis, inputs[3])?;
577 let output_zero_point = int8_zero_point(analysis, inputs[4])?;
578 if multiplier < 0 || !(2..=62).contains(&shift) || input_zero_point != 0 {
579 return Err(AdmitError::Unsupported);
580 }
581
582 let elements = tensor_elements(analysis, inputs[0])?;
583 let input_bytes = elements.checked_mul(4).ok_or(AdmitError::Unsupported)?;
584 let output_bytes = align_to_four(elements)?;
585 if input_bytes
586 .checked_add(output_bytes)
587 .is_none_or(|bytes| bytes > INT8_RESCALE_MAX_TOTAL_BYTES)
588 {
589 return Err(AdmitError::Unsupported);
590 }
591 Ok(CompilerSpec::Int32ToInt8Rescale {
592 elements,
593 multiplier,
594 shift,
595 output_zero_point,
596 })
597}
598
599fn align_to_four(bytes: usize) -> Result<usize, AdmitError> {
600 bytes
601 .checked_add(3)
602 .map(|bytes| bytes & !3)
603 .ok_or(AdmitError::Unsupported)
604}
605
606fn int32_constant(
607 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
608 value: virtio_accel_tosa::ValueId,
609) -> Result<i32, AdmitError> {
610 let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
611 return Err(AdmitError::Unsupported);
612 };
613 if tensor.dtype() != DType::INT32 || tensor.dimensions().ne([1]) {
614 return Err(AdmitError::Unsupported);
615 }
616 let bytes: [u8; 4] = analysis
617 .serialized_constant(value)
618 .ok_or(AdmitError::Unsupported)?
619 .try_into()
620 .map_err(|_| AdmitError::Unsupported)?;
621 Ok(i32::from_le_bytes(bytes))
622}
623
624fn int8_zero_point(
625 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
626 value: virtio_accel_tosa::ValueId,
627) -> Result<i8, AdmitError> {
628 let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
629 return Err(AdmitError::Unsupported);
630 };
631 if tensor.dtype() != DType::INT8 || tensor.dimensions().ne([1]) {
632 return Err(AdmitError::Unsupported);
633 }
634 let [byte] = analysis
635 .serialized_constant(value)
636 .ok_or(AdmitError::Unsupported)?
637 else {
638 return Err(AdmitError::Unsupported);
639 };
640 Ok(*byte as i8)
641}
642
643fn tensor_elements(
644 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
645 value: virtio_accel_tosa::ValueId,
646) -> Result<usize, AdmitError> {
647 let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
648 return Err(AdmitError::Unsupported);
649 };
650 let mut elements = 1usize;
651 for dimension in tensor.dimensions() {
652 let dimension = usize::try_from(dimension).map_err(|_| AdmitError::Unsupported)?;
653 if dimension == 0 {
654 return Err(AdmitError::Unsupported);
655 }
656 elements = elements
657 .checked_mul(dimension)
658 .ok_or(AdmitError::Unsupported)?;
659 }
660 Ok(elements)
661}
662
663fn admit_fp8_to_bf16(
669 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
670 block: virtio_accel_tosa::BlockId,
671 cast: virtio_accel_tosa::OperatorId,
672) -> Result<CompilerSpec, AdmitError> {
673 let inputs = analysis.operator_inputs(cast);
674 let outputs = analysis.operator_outputs(cast);
675 if inputs.len() != 1
676 || outputs.len() != 1
677 || analysis.block_inputs(block) != [inputs[0]]
678 || analysis.block_outputs(block) != [outputs[0]]
679 {
680 return Err(AdmitError::Unsupported);
681 }
682
683 let AnalyzedValueKind::Tensor(input) = analysis.value(inputs[0]).kind() else {
684 return Err(AdmitError::Unsupported);
685 };
686 let AnalyzedValueKind::Tensor(output) = analysis.value(outputs[0]).kind() else {
687 return Err(AdmitError::Unsupported);
688 };
689 let format = match input.dtype() {
690 DType::FP8E4M3 => Fp8Format::E4M3,
691 DType::FP8E5M2 => Fp8Format::E5M2,
692 _ => return Err(AdmitError::Unsupported),
693 };
694 if output.dtype() != DType::BF16 || input.dimensions().ne(output.dimensions()) {
695 return Err(AdmitError::Unsupported);
696 }
697
698 let mut elements = 1usize;
699 for dimension in output.dimensions() {
700 let dimension = usize::try_from(dimension).map_err(|_| AdmitError::Unsupported)?;
701 if dimension == 0 {
702 return Err(AdmitError::Unsupported);
703 }
704 elements = elements
705 .checked_mul(dimension)
706 .ok_or(AdmitError::Unsupported)?;
707 }
708 if elements % FP8_CAST_LINE_SIZE != 0 {
709 return Err(AdmitError::Unsupported);
710 }
711
712 Ok(CompilerSpec::Fp8ToBf16 { format, elements })
713}
714
715fn admit_identity(
720 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
721 block: virtio_accel_tosa::BlockId,
722) -> Result<CompilerSpec, AdmitError> {
723 let inputs = analysis.block_inputs(block);
724 let outputs = analysis.block_outputs(block);
725 if inputs.len() != 1 || outputs.len() != 1 {
726 return Err(AdmitError::Unsupported);
727 }
728 for value in analysis.values() {
729 if let AnalyzedValueKind::Tensor(tensor) = value.kind() {
730 if tensor.dtype() != DType::BF16 {
731 return Err(AdmitError::Unsupported);
732 }
733 }
734 }
735
736 let AnalyzedValueKind::Tensor(output) = analysis.value(outputs[0]).kind() else {
737 return Err(AdmitError::Unsupported);
738 };
739 output.rank().ok_or(AdmitError::Unsupported)?;
740 let mut elements: usize = 1;
741 for dimension in output.dimensions() {
742 let dimension = usize::try_from(dimension).map_err(|_| AdmitError::Unsupported)?;
744 elements = elements
745 .checked_mul(dimension)
746 .ok_or(AdmitError::Unsupported)?;
747 }
748 if elements == 0 || elements % IDENTITY_LINE_SIZE != 0 {
749 return Err(AdmitError::Unsupported);
750 }
751
752 Ok(CompilerSpec::Identity { elements })
753}
754
755fn admit_matmul(
765 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
766 block: virtio_accel_tosa::BlockId,
767 matmul: virtio_accel_tosa::OperatorId,
768) -> Result<CompilerSpec, AdmitError> {
769 let inputs = analysis.operator_inputs(matmul);
770 let outputs = analysis.operator_outputs(matmul);
771 if inputs.len() != 4 || outputs.len() != 1 {
772 return Err(AdmitError::Unsupported);
773 }
774
775 if inputs[0] == inputs[1]
780 || analysis.block_inputs(block) != [inputs[0], inputs[1]]
781 || analysis.block_outputs(block) != [outputs[0]]
782 {
783 return Err(AdmitError::Unsupported);
784 }
785 for operator in analysis.execution_order(block) {
788 if analysis.operator(*operator).op() != Op::CONST {
789 continue;
790 }
791 for produced in analysis.operator_outputs(*operator) {
792 if *produced != inputs[2] && *produced != inputs[3] {
793 return Err(AdmitError::Unsupported);
794 }
795 }
796 }
797
798 let lhs = matmul_dims(analysis, inputs[0], DType::BF16)?;
799 let rhs = matmul_dims(analysis, inputs[1], DType::BF16)?;
800 let out = matmul_dims(analysis, outputs[0], DType::FP32)?;
801
802 let ([1, m, k], [1, k2, n], [1, m2, n2]) = (lhs, rhs, out) else {
805 return Err(AdmitError::Unsupported);
806 };
807 if k != k2 || m != m2 || n != n2 {
808 return Err(AdmitError::Unsupported);
809 }
810 if !tile_admissible(m, MATMUL_TILE_M)
811 || !tile_admissible(k, MATMUL_TILE_K)
812 || !tile_admissible(n, MATMUL_TILE_N)
813 {
814 return Err(AdmitError::Unsupported);
815 }
816
817 Ok(CompilerSpec::Matmul { m, k, n })
818}
819
820fn admit_fp8_matmul(
832 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
833 block: virtio_accel_tosa::BlockId,
834 matmul: virtio_accel_tosa::OperatorId,
835 casts: [virtio_accel_tosa::OperatorId; 2],
836) -> Result<CompilerSpec, AdmitError> {
837 let inputs = analysis.operator_inputs(matmul);
838 let outputs = analysis.operator_outputs(matmul);
839 if inputs.len() != 4 || outputs.len() != 1 {
840 return Err(AdmitError::Unsupported);
841 }
842 if analysis.block_outputs(block) != [outputs[0]] {
843 return Err(AdmitError::Unsupported);
844 }
845
846 let mut promoted: [Option<virtio_accel_tosa::ValueId>; 2] = [None, None];
849 for cast in casts {
850 let cast_inputs = analysis.operator_inputs(cast);
851 let cast_outputs = analysis.operator_outputs(cast);
852 if cast_inputs.len() != 1 || cast_outputs.len() != 1 {
853 return Err(AdmitError::Unsupported);
854 }
855 let operand = if cast_outputs[0] == inputs[0] {
856 0
857 } else if cast_outputs[0] == inputs[1] {
858 1
859 } else {
860 return Err(AdmitError::Unsupported);
862 };
863 if promoted[operand].is_some() {
864 return Err(AdmitError::Unsupported);
865 }
866 if analysis.block_outputs(block).contains(&cast_outputs[0]) {
869 return Err(AdmitError::Unsupported);
870 }
871 promoted[operand] = Some(cast_inputs[0]);
872 }
873 let ([Some(lhs_storage), Some(rhs_storage)], _) = (promoted, ()) else {
874 return Err(AdmitError::Unsupported);
875 };
876
877 if lhs_storage == rhs_storage || analysis.block_inputs(block) != [lhs_storage, rhs_storage] {
880 return Err(AdmitError::Unsupported);
881 }
882 for operator in analysis.execution_order(block) {
884 if analysis.operator(*operator).op() != Op::CONST {
885 continue;
886 }
887 for produced in analysis.operator_outputs(*operator) {
888 if *produced != inputs[2] && *produced != inputs[3] {
889 return Err(AdmitError::Unsupported);
890 }
891 }
892 }
893
894 let lhs_format = fp8_storage_format(analysis, lhs_storage)?;
896 let rhs_format = fp8_storage_format(analysis, rhs_storage)?;
897 if lhs_format != rhs_format {
898 return Err(AdmitError::Unsupported);
899 }
900
901 let lhs = matmul_dims(analysis, lhs_storage, fp8_dtype(lhs_format))?;
902 let rhs = matmul_dims(analysis, rhs_storage, fp8_dtype(rhs_format))?;
903 let lhs_bf16 = matmul_dims(analysis, inputs[0], DType::BF16)?;
904 let rhs_bf16 = matmul_dims(analysis, inputs[1], DType::BF16)?;
905 let out = matmul_dims(analysis, outputs[0], DType::FP32)?;
906
907 if lhs != lhs_bf16 || rhs != rhs_bf16 {
909 return Err(AdmitError::Unsupported);
910 }
911 let ([1, m, k], [1, k2, n], [1, m2, n2]) = (lhs, rhs, out) else {
912 return Err(AdmitError::Unsupported);
913 };
914 if k != k2 || m != m2 || n != n2 {
915 return Err(AdmitError::Unsupported);
916 }
917 if !tile_admissible(m, MATMUL_TILE_M)
918 || !tile_admissible(k, MATMUL_TILE_K)
919 || !tile_admissible(n, MATMUL_TILE_N)
920 {
921 return Err(AdmitError::Unsupported);
922 }
923
924 Ok(CompilerSpec::Fp8Matmul {
925 format: lhs_format,
926 m,
927 k,
928 n,
929 })
930}
931
932fn fp8_storage_format(
934 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
935 value: virtio_accel_tosa::ValueId,
936) -> Result<Fp8Format, AdmitError> {
937 let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
938 return Err(AdmitError::Unsupported);
939 };
940 match tensor.dtype() {
941 DType::FP8E4M3 => Ok(Fp8Format::E4M3),
942 DType::FP8E5M2 => Ok(Fp8Format::E5M2),
943 _ => Err(AdmitError::Unsupported),
944 }
945}
946
947fn fp8_dtype(format: Fp8Format) -> DType {
949 match format {
950 Fp8Format::E4M3 => DType::FP8E4M3,
951 Fp8Format::E5M2 => DType::FP8E5M2,
952 }
953}
954
955fn matmul_dims(
958 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
959 value: virtio_accel_tosa::ValueId,
960 dtype: DType,
961) -> Result<[usize; 3], AdmitError> {
962 let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
963 return Err(AdmitError::Unsupported);
964 };
965 if tensor.dtype() != dtype || tensor.rank() != Some(3) {
966 return Err(AdmitError::Unsupported);
967 }
968 let mut dims = [0usize; 3];
969 for (slot, dimension) in dims.iter_mut().zip(tensor.dimensions()) {
970 *slot = usize::try_from(dimension).map_err(|_| AdmitError::Unsupported)?;
971 if *slot == 0 {
972 return Err(AdmitError::Unsupported);
973 }
974 }
975 Ok(dims)
976}
977
978fn tile_admissible(dim: usize, tile: usize) -> bool {
980 dim > 0 && dim % tile == 0 && dim <= MATMUL_MAX_DIM
981}
982
983fn admit_max_pool2d(
989 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
990 block: virtio_accel_tosa::BlockId,
991 max_pool: virtio_accel_tosa::OperatorId,
992) -> Result<CompilerSpec, AdmitError> {
993 let inputs = analysis.operator_inputs(max_pool);
994 let outputs = analysis.operator_outputs(max_pool);
995 if inputs.len() != 1
996 || outputs.len() != 1
997 || analysis.block_inputs(block) != [inputs[0]]
998 || analysis.block_outputs(block) != [outputs[0]]
999 {
1000 return Err(AdmitError::Unsupported);
1001 }
1002
1003 let OpAttributes::MaxPool2d {
1004 kernel,
1005 stride,
1006 pad,
1007 nan_mode,
1008 } = analysis.operator(max_pool).source().attributes()
1009 else {
1010 return Err(AdmitError::Unsupported);
1011 };
1012 if nan_mode != NanPropagationMode::PROPAGATE {
1013 return Err(AdmitError::Unsupported);
1014 }
1015 let kernel = exact_positive_pair(kernel.iter(), MAX_POOL_MAX_KERNEL)?;
1016 let stride = exact_positive_pair(stride.iter(), MAX_POOL_MAX_STRIDE)?;
1017 let pad: Vec<_> = pad.iter().collect();
1018 if pad != [0, 0, 0, 0] {
1019 return Err(AdmitError::Unsupported);
1020 }
1021
1022 let [batch, input_h, input_w, channels] = pool_dims(analysis, inputs[0])?;
1023 let [output_batch, output_h, output_w, output_channels] = pool_dims(analysis, outputs[0])?;
1024 if batch != 1 || output_batch != 1 || channels != output_channels {
1025 return Err(AdmitError::Unsupported);
1026 }
1027 let input_elements = input_h
1028 .checked_mul(input_w)
1029 .and_then(|elements| elements.checked_mul(channels))
1030 .ok_or(AdmitError::Unsupported)?;
1031 let output_elements = output_h
1032 .checked_mul(output_w)
1033 .and_then(|elements| elements.checked_mul(channels))
1034 .ok_or(AdmitError::Unsupported)?;
1035 if input_elements
1036 .checked_add(output_elements)
1037 .is_none_or(|total| total > MAX_POOL_MAX_TOTAL_ELEMENTS)
1038 {
1039 return Err(AdmitError::Unsupported);
1040 }
1041
1042 Ok(CompilerSpec::MaxPool2d {
1043 input_h,
1044 input_w,
1045 channels,
1046 output_h,
1047 output_w,
1048 kernel_h: kernel[0],
1049 kernel_w: kernel[1],
1050 stride_h: stride[0],
1051 stride_w: stride[1],
1052 })
1053}
1054
1055fn exact_positive_pair(
1056 values: impl Iterator<Item = i32>,
1057 maximum: usize,
1058) -> Result<[usize; 2], AdmitError> {
1059 let values: Vec<_> = values.collect();
1060 let [first, second] = values.as_slice() else {
1061 return Err(AdmitError::Unsupported);
1062 };
1063 let first = usize::try_from(*first).map_err(|_| AdmitError::Unsupported)?;
1064 let second = usize::try_from(*second).map_err(|_| AdmitError::Unsupported)?;
1065 if first == 0 || second == 0 || first > maximum || second > maximum {
1066 return Err(AdmitError::Unsupported);
1067 }
1068 Ok([first, second])
1069}
1070
1071fn pool_dims(
1072 analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
1073 value: virtio_accel_tosa::ValueId,
1074) -> Result<[usize; 4], AdmitError> {
1075 let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
1076 return Err(AdmitError::Unsupported);
1077 };
1078 if tensor.dtype() != DType::BF16 || tensor.rank() != Some(4) {
1079 return Err(AdmitError::Unsupported);
1080 }
1081 let mut dims = [0usize; 4];
1082 for (slot, dimension) in dims.iter_mut().zip(tensor.dimensions()) {
1083 *slot = usize::try_from(dimension).map_err(|_| AdmitError::Unsupported)?;
1084 if *slot == 0 {
1085 return Err(AdmitError::Unsupported);
1086 }
1087 }
1088 Ok(dims)
1089}
1090
1091#[cfg(test)]
1092mod tests {
1093 use super::*;
1094 use virtio_accel_conformance::numerics::RESCALE_INT32_TO_INT8;
1095 use virtio_accel_tosa::Target;
1096 use virtio_accel_tosa_build::{OperatorKind, OwnedGraph, OwnedOperator, OwnedTensor};
1097
1098 fn identity_graph(dtype: DType, shape: Vec<i32>) -> OwnedGraph<'static> {
1099 let mut graph = OwnedGraph::new("main");
1100 graph.push_tensor(OwnedTensor::new("x", shape.clone(), dtype));
1101 graph.push_tensor(OwnedTensor::new("y", shape, dtype));
1102 graph.push_operator(OwnedOperator::new(
1103 OperatorKind::Identity,
1104 vec!["x".into()],
1105 vec!["y".into()],
1106 ));
1107 graph.push_input("x");
1108 graph.push_output("y");
1109 graph
1110 }
1111
1112 fn fp8_cast_graph(input: DType, output: DType, shape: Vec<i32>) -> OwnedGraph<'static> {
1113 let mut graph = OwnedGraph::new("main");
1114 graph
1115 .push_tensor(OwnedTensor::new("x", shape.clone(), input))
1116 .push_tensor(OwnedTensor::new("y", shape, output))
1117 .push_operator(OwnedOperator::new(
1118 OperatorKind::Cast,
1119 vec!["x".into()],
1120 vec!["y".into()],
1121 ))
1122 .push_input("x")
1123 .push_output("y");
1124 graph
1125 }
1126
1127 fn matmul_graph(
1129 m: i32,
1130 k: i32,
1131 n: i32,
1132 in_dtype: DType,
1133 out_dtype: DType,
1134 ) -> OwnedGraph<'static> {
1135 matmul_graph_with_zero_points(m, k, n, in_dtype, out_dtype, 0, 0)
1136 }
1137
1138 fn matmul_graph_with_zero_points(
1139 m: i32,
1140 k: i32,
1141 n: i32,
1142 in_dtype: DType,
1143 out_dtype: DType,
1144 left_zero_point: i8,
1145 right_zero_point: i8,
1146 ) -> OwnedGraph<'static> {
1147 let zero_point = |dtype: DType, value: i8| match dtype {
1148 DType::INT8 => vec![value as u8],
1149 DType::BF16 => vec![0u8; 2],
1150 DType::FP32 => vec![0u8; 4],
1151 _ => Vec::new(),
1152 };
1153 let mut graph = OwnedGraph::new("main");
1154 graph
1155 .push_tensor(OwnedTensor::new("lhs", vec![1, m, k], in_dtype))
1156 .push_tensor(OwnedTensor::new("rhs", vec![1, k, n], in_dtype))
1157 .push_tensor(OwnedTensor::constant(
1158 "lhs_zp",
1159 vec![1],
1160 in_dtype,
1161 zero_point(in_dtype, left_zero_point),
1162 ))
1163 .push_tensor(OwnedTensor::constant(
1164 "rhs_zp",
1165 vec![1],
1166 in_dtype,
1167 zero_point(in_dtype, right_zero_point),
1168 ))
1169 .push_tensor(OwnedTensor::new("output", vec![1, m, n], out_dtype))
1170 .push_operator(OwnedOperator::new(
1171 OperatorKind::Const,
1172 vec![],
1173 vec!["lhs_zp".into()],
1174 ))
1175 .push_operator(OwnedOperator::new(
1176 OperatorKind::Const,
1177 vec![],
1178 vec!["rhs_zp".into()],
1179 ))
1180 .push_operator(OwnedOperator::new(
1181 OperatorKind::MatMul,
1182 vec!["lhs".into(), "rhs".into(), "lhs_zp".into(), "rhs_zp".into()],
1183 vec!["output".into()],
1184 ))
1185 .push_input("lhs")
1186 .push_input("rhs")
1187 .push_output("output");
1188 graph
1189 }
1190
1191 fn rescale_graph(
1192 elements: i32,
1193 multiplier: i32,
1194 shift: i8,
1195 per_channel: bool,
1196 rounding_mode: RoundingMode,
1197 ) -> OwnedGraph<'static> {
1198 let mut graph = OwnedGraph::new("main");
1199 graph
1200 .push_tensor(OwnedTensor::new("input", vec![elements], DType::INT32))
1201 .push_tensor(OwnedTensor::constant(
1202 "multiplier",
1203 vec![1],
1204 DType::INT32,
1205 multiplier.to_le_bytes().to_vec(),
1206 ))
1207 .push_tensor(OwnedTensor::constant(
1208 "shift",
1209 vec![1],
1210 DType::INT8,
1211 vec![shift as u8],
1212 ))
1213 .push_tensor(OwnedTensor::constant(
1214 "input_zp",
1215 vec![1],
1216 DType::INT32,
1217 0_i32.to_le_bytes().to_vec(),
1218 ))
1219 .push_tensor(OwnedTensor::constant(
1220 "output_zp",
1221 vec![1],
1222 DType::INT8,
1223 vec![(-3_i8) as u8],
1224 ))
1225 .push_tensor(OwnedTensor::new("output", vec![elements], DType::INT8));
1226 for parameter in ["multiplier", "shift", "input_zp", "output_zp"] {
1227 graph.push_operator(OwnedOperator::new(
1228 OperatorKind::Const,
1229 vec![],
1230 vec![parameter.into()],
1231 ));
1232 }
1233 graph
1234 .push_operator(OwnedOperator::new(
1235 OperatorKind::Rescale {
1236 scale32: true,
1237 rounding_mode,
1238 per_channel,
1239 input_unsigned: false,
1240 output_unsigned: false,
1241 },
1242 vec![
1243 "input".into(),
1244 "multiplier".into(),
1245 "shift".into(),
1246 "input_zp".into(),
1247 "output_zp".into(),
1248 ],
1249 vec!["output".into()],
1250 ))
1251 .push_input("input")
1252 .push_output("output");
1253 graph
1254 }
1255
1256 struct MaxPoolCase {
1257 input: [i32; 3],
1258 kernel: [i32; 2],
1259 stride: [i32; 2],
1260 pad: [i32; 4],
1261 dtype: DType,
1262 nan_mode: NanPropagationMode,
1263 }
1264
1265 fn max_pool_graph(case: MaxPoolCase) -> OwnedGraph<'static> {
1266 let [input_h, input_w, channels] = case.input;
1267 let output_h = (input_h + case.pad[0] + case.pad[1] - case.kernel[0]) / case.stride[0] + 1;
1268 let output_w = (input_w + case.pad[2] + case.pad[3] - case.kernel[1]) / case.stride[1] + 1;
1269 let mut graph = OwnedGraph::new("main");
1270 graph
1271 .push_tensor(OwnedTensor::new(
1272 "input",
1273 vec![1, input_h, input_w, channels],
1274 case.dtype,
1275 ))
1276 .push_tensor(OwnedTensor::new(
1277 "output",
1278 vec![1, output_h, output_w, channels],
1279 case.dtype,
1280 ))
1281 .push_operator(OwnedOperator::new(
1282 OperatorKind::MaxPool2d {
1283 kernel: case.kernel,
1284 stride: case.stride,
1285 pad: case.pad,
1286 nan_mode: case.nan_mode,
1287 },
1288 vec!["input".into()],
1289 vec!["output".into()],
1290 ))
1291 .push_input("input")
1292 .push_output("output");
1293 graph
1294 }
1295
1296 #[test]
1297 fn both_targets_are_coherent_and_distinct() {
1298 assert_eq!(XDNA_TOSA_TARGET.validate(), Ok(XDNA_TOSA_TARGET));
1299 assert_eq!(
1300 XDNA_TOSA_INTEGER_TARGET.validate(),
1301 Ok(XDNA_TOSA_INTEGER_TARGET)
1302 );
1303 assert_ne!(XDNA_TOSA_TARGET, XDNA_TOSA_INTEGER_TARGET);
1304 for target in [
1305 XDNA_TOSA_TARGET,
1306 XDNA_TOSA_FP8_TARGET,
1307 XDNA_TOSA_INTEGER_TARGET,
1308 ] {
1309 assert_eq!(Target::from_identity(target.to_identity()), Ok(target));
1310 }
1311 }
1312
1313 #[test]
1314 fn admits_bf16_identity() {
1315 let bytes = identity_graph(DType::BF16, vec![1, 4, 1024])
1316 .build(XDNA_TOSA_TARGET)
1317 .expect("build bf16 identity");
1318 let spec = admit(&bytes, XDNA_TOSA_TARGET).expect("admit");
1319 assert_eq!(spec, CompilerSpec::Identity { elements: 4 * 1024 });
1320 }
1321
1322 #[test]
1323 fn integer_capability_preserves_the_openvino_base_and_adds_rescale() {
1324 assert_eq!(
1325 XDNA_TOSA_INTEGER_CAPABILITY.target,
1326 XDNA_TOSA_INTEGER_TARGET
1327 );
1328 assert_eq!(XDNA_TOSA_INTEGER_CAPABILITY.dtypes, INTEGER_DTYPES);
1329 assert_eq!(XDNA_TOSA_INTEGER_CAPABILITY.operators, INTEGER_OPERATORS);
1330 for op in [Op::CONST, Op::IDENTITY, Op::MATMUL, Op::RESCALE] {
1331 assert!(XDNA_TOSA_INTEGER_CAPABILITY.supports_operator(op));
1332 }
1333 assert_eq!(XDNA_TOSA_INTEGER_CAPABILITY.graph.max_regions, 1);
1334 assert_eq!(XDNA_TOSA_INTEGER_CAPABILITY.graph.max_blocks, 1);
1335 assert_eq!(
1336 XDNA_TOSA_INTEGER_CAPABILITY.graph.runtime_conditions,
1337 RuntimeConditionSupport::None
1338 );
1339 }
1340
1341 #[test]
1342 fn admits_shared_int8_identity_shape() {
1343 let bytes = identity_graph(DType::INT8, vec![8])
1344 .build(XDNA_TOSA_INTEGER_TARGET)
1345 .expect("build int8 identity");
1346 assert_eq!(
1347 admit(&bytes, XDNA_TOSA_INTEGER_TARGET),
1348 Ok(CompilerSpec::Int8Identity {
1349 elements: 8,
1350 line_size: 8,
1351 })
1352 );
1353 }
1354
1355 #[test]
1356 fn admits_zero_point_aware_int8_matmul() {
1357 let bytes = matmul_graph_with_zero_points(2, 3, 2, DType::INT8, DType::INT32, -2, 3)
1358 .build(XDNA_TOSA_INTEGER_TARGET)
1359 .expect("build int8 matmul");
1360 assert_eq!(
1361 admit(&bytes, XDNA_TOSA_INTEGER_TARGET),
1362 Ok(CompilerSpec::Int8Matmul {
1363 m: 2,
1364 k: 3,
1365 n: 2,
1366 left_zero_point: -2,
1367 right_zero_point: 3,
1368 })
1369 );
1370 }
1371
1372 #[test]
1373 fn admits_shared_exact_int32_to_int8_rescale() {
1374 assert_eq!(
1375 admit(RESCALE_INT32_TO_INT8.artifact, XDNA_TOSA_INTEGER_TARGET),
1376 Ok(CompilerSpec::Int32ToInt8Rescale {
1377 elements: 16,
1378 multiplier: 1 << 29,
1379 shift: 30,
1380 output_zero_point: -3,
1381 })
1382 );
1383 }
1384
1385 #[test]
1386 fn rescale_rejects_unimplemented_modes_and_invalid_parameters() {
1387 let per_channel = rescale_graph(1, 1 << 29, 30, true, RoundingMode::SINGLE_ROUND)
1388 .build(XDNA_TOSA_INTEGER_TARGET)
1389 .expect("per-channel one-element RESCALE is valid TOSA");
1390 assert_eq!(
1391 admit(&per_channel, XDNA_TOSA_INTEGER_TARGET),
1392 Err(AdmitError::Unsupported)
1393 );
1394
1395 assert!(
1396 rescale_graph(16, 1 << 29, 1, false, RoundingMode::SINGLE_ROUND)
1397 .build(XDNA_TOSA_INTEGER_TARGET)
1398 .is_err(),
1399 "shift 1 violates the released RESCALE range"
1400 );
1401
1402 assert!(
1403 rescale_graph(16, 1 << 29, 30, false, RoundingMode::DOUBLE_ROUND)
1404 .build(XDNA_TOSA_INTEGER_TARGET)
1405 .is_err(),
1406 "DOUBLE_ROUND requires an extension absent from the integer target"
1407 );
1408 }
1409
1410 #[test]
1411 fn rejects_int8_shapes_outside_the_one_core_envelope() {
1412 let non_word_identity = identity_graph(DType::INT8, vec![6])
1413 .build(XDNA_TOSA_INTEGER_TARGET)
1414 .expect("build int8 identity");
1415 assert_eq!(
1416 admit(&non_word_identity, XDNA_TOSA_INTEGER_TARGET),
1417 Err(AdmitError::Unsupported)
1418 );
1419
1420 let non_divisible_identity = identity_graph(DType::INT8, vec![1025])
1421 .build(XDNA_TOSA_INTEGER_TARGET)
1422 .expect("build int8 identity");
1423 assert_eq!(
1424 admit(&non_divisible_identity, XDNA_TOSA_INTEGER_TARGET),
1425 Err(AdmitError::Unsupported)
1426 );
1427
1428 let oversized_matmul =
1429 matmul_graph_with_zero_points(64, 64, 64, DType::INT8, DType::INT32, -2, 3)
1430 .build(XDNA_TOSA_INTEGER_TARGET)
1431 .expect("build int8 matmul");
1432 assert_eq!(
1433 admit(&oversized_matmul, XDNA_TOSA_INTEGER_TARGET),
1434 Err(AdmitError::Unsupported)
1435 );
1436 }
1437
1438 #[test]
1439 fn admits_both_explicit_fp8_to_bf16_casts() {
1440 for (dtype, format) in [
1441 (DType::FP8E4M3, Fp8Format::E4M3),
1442 (DType::FP8E5M2, Fp8Format::E5M2),
1443 ] {
1444 let bytes = fp8_cast_graph(dtype, DType::BF16, vec![1, 1, 4096])
1445 .build(XDNA_TOSA_FP8_TARGET)
1446 .expect("build fp8 cast");
1447 assert_eq!(
1448 admit(&bytes, XDNA_TOSA_FP8_TARGET),
1449 Ok(CompilerSpec::Fp8ToBf16 {
1450 format,
1451 elements: 4096,
1452 })
1453 );
1454 }
1455 }
1456
1457 #[test]
1458 fn fp8_storage_tier_rejects_hidden_or_unsupported_conversion() {
1459 let valid = fp8_cast_graph(DType::FP8E4M3, DType::BF16, vec![1024])
1460 .build(XDNA_TOSA_FP8_TARGET)
1461 .expect("build fp8 cast");
1462 assert_eq!(admit(&valid, XDNA_TOSA_TARGET), Err(AdmitError::Analysis));
1463
1464 for graph in [
1465 fp8_cast_graph(DType::FP8E4M3, DType::FP32, vec![1024]),
1466 fp8_cast_graph(DType::FP8E5M2, DType::BF16, vec![8]),
1467 ] {
1468 let bytes = graph
1469 .build(XDNA_TOSA_FP8_TARGET)
1470 .expect("build semantically valid cast");
1471 assert_eq!(
1472 admit(&bytes, XDNA_TOSA_FP8_TARGET),
1473 Err(AdmitError::Unsupported)
1474 );
1475 }
1476 }
1477
1478 #[test]
1479 fn admits_bf16_matmul_at_tile_multiples() {
1480 for (m, k, n) in [(32, 64, 32), (64, 128, 96)] {
1482 let bytes = matmul_graph(m, k, n, DType::BF16, DType::FP32)
1483 .build(XDNA_TOSA_TARGET)
1484 .expect("build bf16 matmul");
1485 let spec = admit(&bytes, XDNA_TOSA_TARGET).expect("admit");
1486 assert_eq!(
1487 spec,
1488 CompilerSpec::Matmul {
1489 m: m as usize,
1490 k: k as usize,
1491 n: n as usize,
1492 }
1493 );
1494 }
1495 }
1496
1497 #[test]
1498 fn admits_bf16_nhwc_max_pool2d_corpus_shape() {
1499 let bytes = max_pool_graph(MaxPoolCase {
1500 input: [4, 4, 2],
1501 kernel: [2, 2],
1502 stride: [2, 2],
1503 pad: [0; 4],
1504 dtype: DType::BF16,
1505 nan_mode: NanPropagationMode::PROPAGATE,
1506 })
1507 .build(XDNA_TOSA_TARGET)
1508 .expect("build bf16 max pool2d");
1509 assert_eq!(
1510 admit(&bytes, XDNA_TOSA_TARGET),
1511 Ok(CompilerSpec::MaxPool2d {
1512 input_h: 4,
1513 input_w: 4,
1514 channels: 2,
1515 output_h: 2,
1516 output_w: 2,
1517 kernel_h: 2,
1518 kernel_w: 2,
1519 stride_h: 2,
1520 stride_w: 2,
1521 })
1522 );
1523 }
1524
1525 #[test]
1526 fn rejects_max_pool2d_outside_the_proven_envelope() {
1527 let cases = [
1528 max_pool_graph(MaxPoolCase {
1529 input: [4, 4, 2],
1530 kernel: [2, 2],
1531 stride: [2, 2],
1532 pad: [0; 4],
1533 dtype: DType::FP32,
1534 nan_mode: NanPropagationMode::PROPAGATE,
1535 }),
1536 max_pool_graph(MaxPoolCase {
1537 input: [4, 4, 2],
1538 kernel: [2, 2],
1539 stride: [2, 2],
1540 pad: [0; 4],
1541 dtype: DType::BF16,
1542 nan_mode: NanPropagationMode::IGNORE,
1543 }),
1544 max_pool_graph(MaxPoolCase {
1545 input: [4, 4, 2],
1546 kernel: [2, 2],
1547 stride: [2, 2],
1548 pad: [1; 4],
1549 dtype: DType::BF16,
1550 nan_mode: NanPropagationMode::PROPAGATE,
1551 }),
1552 max_pool_graph(MaxPoolCase {
1553 input: [16, 16, 2],
1554 kernel: [9, 2],
1555 stride: [1, 1],
1556 pad: [0; 4],
1557 dtype: DType::BF16,
1558 nan_mode: NanPropagationMode::PROPAGATE,
1559 }),
1560 max_pool_graph(MaxPoolCase {
1561 input: [64, 64, 2],
1562 kernel: [2, 2],
1563 stride: [2, 2],
1564 pad: [0; 4],
1565 dtype: DType::BF16,
1566 nan_mode: NanPropagationMode::PROPAGATE,
1567 }),
1568 ];
1569 for graph in cases {
1570 let bytes = graph
1571 .build(XDNA_TOSA_TARGET)
1572 .expect("build semantically valid max pool2d");
1573 assert_eq!(
1574 admit(&bytes, XDNA_TOSA_TARGET),
1575 Err(AdmitError::Unsupported)
1576 );
1577 }
1578 }
1579
1580 #[test]
1581 fn admits_zero_operator_passthrough() {
1582 let mut graph = OwnedGraph::new("main");
1584 graph.push_tensor(OwnedTensor::new("x", vec![1, 4, 1024], DType::BF16));
1585 graph.push_input("x");
1586 graph.push_output("x");
1587 let bytes = graph.build(XDNA_TOSA_TARGET).expect("build passthrough");
1588 assert_eq!(
1589 admit(&bytes, XDNA_TOSA_TARGET),
1590 Ok(CompilerSpec::Identity { elements: 4 * 1024 })
1591 );
1592 }
1593
1594 #[test]
1595 fn rejects_constant_output_identity() {
1596 let shape = vec![1i32, 4, 1024];
1599 let mut graph = OwnedGraph::new("main");
1600 graph
1601 .push_tensor(OwnedTensor::new("x", shape.clone(), DType::BF16))
1602 .push_tensor(OwnedTensor::constant(
1603 "c",
1604 shape.clone(),
1605 DType::BF16,
1606 vec![0u8; 4 * 1024 * 2],
1607 ))
1608 .push_tensor(OwnedTensor::new("y", shape.clone(), DType::BF16))
1609 .push_tensor(OwnedTensor::new("dead", shape, DType::BF16))
1610 .push_operator(OwnedOperator::new(
1611 OperatorKind::Const,
1612 vec![],
1613 vec!["c".into()],
1614 ))
1615 .push_operator(OwnedOperator::new(
1616 OperatorKind::Identity,
1617 vec!["c".into()],
1618 vec!["y".into()],
1619 ))
1620 .push_operator(OwnedOperator::new(
1621 OperatorKind::Identity,
1622 vec!["x".into()],
1623 vec!["dead".into()],
1624 ))
1625 .push_input("x")
1626 .push_output("y");
1627 let bytes = graph
1628 .build(XDNA_TOSA_TARGET)
1629 .expect("build constant identity");
1630 assert_eq!(
1631 admit(&bytes, XDNA_TOSA_TARGET),
1632 Err(AdmitError::Unsupported)
1633 );
1634 }
1635
1636 #[test]
1637 fn rejects_constant_weights_matmul() {
1638 let (m, k, n) = (32i32, 64i32, 32i32);
1641 let mut graph = OwnedGraph::new("main");
1642 graph
1643 .push_tensor(OwnedTensor::constant(
1644 "lhs",
1645 vec![1, m, k],
1646 DType::BF16,
1647 vec![0u8; (m * k * 2) as usize],
1648 ))
1649 .push_tensor(OwnedTensor::new("rhs", vec![1, k, n], DType::BF16))
1650 .push_tensor(OwnedTensor::constant(
1651 "lhs_zp",
1652 vec![1],
1653 DType::BF16,
1654 vec![0u8; 2],
1655 ))
1656 .push_tensor(OwnedTensor::constant(
1657 "rhs_zp",
1658 vec![1],
1659 DType::BF16,
1660 vec![0u8; 2],
1661 ))
1662 .push_tensor(OwnedTensor::new("output", vec![1, m, n], DType::FP32))
1663 .push_operator(OwnedOperator::new(
1664 OperatorKind::Const,
1665 vec![],
1666 vec!["lhs".into()],
1667 ))
1668 .push_operator(OwnedOperator::new(
1669 OperatorKind::Const,
1670 vec![],
1671 vec!["lhs_zp".into()],
1672 ))
1673 .push_operator(OwnedOperator::new(
1674 OperatorKind::Const,
1675 vec![],
1676 vec!["rhs_zp".into()],
1677 ))
1678 .push_operator(OwnedOperator::new(
1679 OperatorKind::MatMul,
1680 vec!["lhs".into(), "rhs".into(), "lhs_zp".into(), "rhs_zp".into()],
1681 vec!["output".into()],
1682 ))
1683 .push_input("rhs")
1684 .push_output("output");
1685 let bytes = graph
1686 .build(XDNA_TOSA_TARGET)
1687 .expect("build constant-weights matmul");
1688 assert_eq!(
1689 admit(&bytes, XDNA_TOSA_TARGET),
1690 Err(AdmitError::Unsupported)
1691 );
1692 }
1693
1694 #[test]
1695 fn admit_error_maps_to_the_reference_backend_error_codes() {
1696 use virtio_accel_core::BackendError;
1697 assert_eq!(
1698 BackendError::from(AdmitError::Parse),
1699 BackendError::InvalidArgument
1700 );
1701 assert_eq!(
1702 BackendError::from(AdmitError::Analysis),
1703 BackendError::InvalidArgument
1704 );
1705 assert_eq!(
1706 BackendError::from(AdmitError::Unsupported),
1707 BackendError::Unsupported
1708 );
1709 }
1710
1711 #[test]
1715 fn rejects_matmul_with_one_value_feeding_both_operands() {
1716 fn aliased_matmul(dim: i32, in_dtype: DType, out_dtype: DType) -> Vec<u8> {
1717 let zp = match in_dtype {
1718 DType::INT8 => vec![0u8],
1719 _ => vec![0u8; 2],
1720 };
1721 let mut graph = OwnedGraph::new("main");
1722 graph
1723 .push_tensor(OwnedTensor::new("x", vec![1, dim, dim], in_dtype))
1724 .push_tensor(OwnedTensor::constant(
1725 "lhs_zp",
1726 vec![1],
1727 in_dtype,
1728 zp.clone(),
1729 ))
1730 .push_tensor(OwnedTensor::constant("rhs_zp", vec![1], in_dtype, zp))
1731 .push_tensor(OwnedTensor::new("output", vec![1, dim, dim], out_dtype))
1732 .push_operator(OwnedOperator::new(
1733 OperatorKind::Const,
1734 vec![],
1735 vec!["lhs_zp".into()],
1736 ))
1737 .push_operator(OwnedOperator::new(
1738 OperatorKind::Const,
1739 vec![],
1740 vec!["rhs_zp".into()],
1741 ))
1742 .push_operator(OwnedOperator::new(
1743 OperatorKind::MatMul,
1744 vec!["x".into(), "x".into(), "lhs_zp".into(), "rhs_zp".into()],
1745 vec!["output".into()],
1746 ))
1747 .push_input("x")
1748 .push_input("x")
1749 .push_output("output");
1750 let target = if in_dtype == DType::INT8 {
1751 XDNA_TOSA_INTEGER_TARGET
1752 } else {
1753 XDNA_TOSA_TARGET
1754 };
1755 graph.build(target).expect("build aliased matmul")
1756 }
1757
1758 let bf16 = aliased_matmul(64, DType::BF16, DType::FP32);
1759 assert_eq!(admit(&bf16, XDNA_TOSA_TARGET), Err(AdmitError::Unsupported));
1760 let int8 = aliased_matmul(32, DType::INT8, DType::INT32);
1761 assert_eq!(
1762 admit(&int8, XDNA_TOSA_INTEGER_TARGET),
1763 Err(AdmitError::Unsupported)
1764 );
1765 }
1766
1767 fn fp8_matmul_graph(
1770 m: i32,
1771 k: i32,
1772 n: i32,
1773 lhs_dtype: DType,
1774 rhs_dtype: DType,
1775 alias_operands: bool,
1776 escape: bool,
1777 ) -> OwnedGraph<'static> {
1778 let mut graph = OwnedGraph::new("main");
1779 let rhs_name = if alias_operands { "lhs_fp8" } else { "rhs_fp8" };
1780 graph
1781 .push_tensor(OwnedTensor::new("lhs_fp8", vec![1, m, k], lhs_dtype))
1782 .push_tensor(OwnedTensor::new("lhs_bf16", vec![1, m, k], DType::BF16))
1783 .push_tensor(OwnedTensor::new("rhs_bf16", vec![1, k, n], DType::BF16))
1784 .push_tensor(OwnedTensor::constant(
1785 "lhs_zp",
1786 vec![1],
1787 DType::BF16,
1788 vec![0u8; 2],
1789 ))
1790 .push_tensor(OwnedTensor::constant(
1791 "rhs_zp",
1792 vec![1],
1793 DType::BF16,
1794 vec![0u8; 2],
1795 ))
1796 .push_tensor(OwnedTensor::new("output", vec![1, m, n], DType::FP32));
1797 if !alias_operands {
1798 graph.push_tensor(OwnedTensor::new("rhs_fp8", vec![1, k, n], rhs_dtype));
1799 }
1800 graph
1801 .push_operator(OwnedOperator::new(
1802 OperatorKind::Cast,
1803 vec!["lhs_fp8".into()],
1804 vec!["lhs_bf16".into()],
1805 ))
1806 .push_operator(OwnedOperator::new(
1807 OperatorKind::Cast,
1808 vec![rhs_name.into()],
1809 vec!["rhs_bf16".into()],
1810 ))
1811 .push_operator(OwnedOperator::new(
1812 OperatorKind::Const,
1813 vec![],
1814 vec!["lhs_zp".into()],
1815 ))
1816 .push_operator(OwnedOperator::new(
1817 OperatorKind::Const,
1818 vec![],
1819 vec!["rhs_zp".into()],
1820 ))
1821 .push_operator(OwnedOperator::new(
1822 OperatorKind::MatMul,
1823 vec![
1824 "lhs_bf16".into(),
1825 "rhs_bf16".into(),
1826 "lhs_zp".into(),
1827 "rhs_zp".into(),
1828 ],
1829 vec!["output".into()],
1830 ))
1831 .push_input("lhs_fp8");
1832 if !alias_operands {
1833 graph.push_input("rhs_fp8");
1834 } else {
1835 graph.push_input("lhs_fp8");
1836 }
1837 graph.push_output("output");
1838 if escape {
1839 graph.push_output("lhs_bf16");
1840 }
1841 graph
1842 }
1843
1844 #[test]
1845 fn admits_fused_fp8_matmul_for_both_encodings() {
1846 for (dtype, format) in [
1847 (DType::FP8E4M3, Fp8Format::E4M3),
1848 (DType::FP8E5M2, Fp8Format::E5M2),
1849 ] {
1850 let bytes = fp8_matmul_graph(32, 64, 32, dtype, dtype, false, false)
1851 .build(XDNA_TOSA_FP8_TARGET)
1852 .expect("build fused fp8 matmul");
1853 assert_eq!(
1854 admit(&bytes, XDNA_TOSA_FP8_TARGET),
1855 Ok(CompilerSpec::Fp8Matmul {
1856 format,
1857 m: 32,
1858 k: 64,
1859 n: 32
1860 })
1861 );
1862 }
1863 }
1864
1865 #[test]
1868 fn rejects_fused_fp8_matmul_with_mixed_encodings() {
1869 let bytes = fp8_matmul_graph(32, 64, 32, DType::FP8E4M3, DType::FP8E5M2, false, false)
1870 .build(XDNA_TOSA_FP8_TARGET)
1871 .expect("build mixed-encoding fused matmul");
1872 assert_eq!(
1873 admit(&bytes, XDNA_TOSA_FP8_TARGET),
1874 Err(AdmitError::Unsupported)
1875 );
1876 }
1877
1878 #[test]
1881 fn rejects_fused_fp8_matmul_whose_promoted_operand_escapes() {
1882 let bytes = fp8_matmul_graph(32, 64, 32, DType::FP8E4M3, DType::FP8E4M3, false, true)
1883 .build(XDNA_TOSA_FP8_TARGET)
1884 .expect("build escaping fused matmul");
1885 assert_eq!(
1886 admit(&bytes, XDNA_TOSA_FP8_TARGET),
1887 Err(AdmitError::Unsupported)
1888 );
1889 }
1890
1891 #[test]
1894 fn rejects_fused_fp8_matmul_with_one_value_feeding_both_operands() {
1895 let bytes = fp8_matmul_graph(64, 64, 64, DType::FP8E4M3, DType::FP8E4M3, true, false)
1896 .build(XDNA_TOSA_FP8_TARGET)
1897 .expect("build aliased fused matmul");
1898 assert_eq!(
1899 admit(&bytes, XDNA_TOSA_FP8_TARGET),
1900 Err(AdmitError::Unsupported)
1901 );
1902 let distinct = fp8_matmul_graph(64, 64, 64, DType::FP8E4M3, DType::FP8E4M3, false, false)
1905 .build(XDNA_TOSA_FP8_TARGET)
1906 .expect("build distinct fused matmul");
1907 assert!(admit(&distinct, XDNA_TOSA_FP8_TARGET).is_ok());
1908 }
1909
1910 #[test]
1911 fn rejects_fused_fp8_matmul_off_the_tested_tiling() {
1912 for (m, k, n) in [(48, 64, 32), (32, 64, MATMUL_MAX_DIM as i32 + 32)] {
1913 let bytes = fp8_matmul_graph(m, k, n, DType::FP8E4M3, DType::FP8E4M3, false, false)
1914 .build(XDNA_TOSA_FP8_TARGET)
1915 .expect("build fused matmul");
1916 assert_eq!(
1917 admit(&bytes, XDNA_TOSA_FP8_TARGET),
1918 Err(AdmitError::Unsupported)
1919 );
1920 }
1921 }
1922
1923 #[test]
1925 fn fused_fp8_matmul_is_not_admitted_on_the_bf16_target() {
1926 let bytes = fp8_matmul_graph(32, 64, 32, DType::FP8E4M3, DType::FP8E4M3, false, false)
1927 .build(XDNA_TOSA_FP8_TARGET)
1928 .expect("build fused matmul");
1929 assert_eq!(admit(&bytes, XDNA_TOSA_TARGET), Err(AdmitError::Analysis));
1930 }
1931
1932 #[test]
1933 fn rejects_fp32_matmul_inputs() {
1934 let bytes = matmul_graph(32, 64, 32, DType::FP32, DType::FP32)
1936 .build(XDNA_TOSA_TARGET)
1937 .expect("build fp32 matmul");
1938 assert_eq!(
1939 admit(&bytes, XDNA_TOSA_TARGET),
1940 Err(AdmitError::Unsupported)
1941 );
1942 }
1943
1944 #[test]
1945 fn rejects_matmul_shape_off_the_tested_tiling() {
1946 for (m, k, n) in [(48, 64, 32), (32, 64, MATMUL_MAX_DIM as i32 + 32)] {
1948 let bytes = matmul_graph(m, k, n, DType::BF16, DType::FP32)
1949 .build(XDNA_TOSA_TARGET)
1950 .expect("build matmul");
1951 assert_eq!(
1952 admit(&bytes, XDNA_TOSA_TARGET),
1953 Err(AdmitError::Unsupported)
1954 );
1955 }
1956 }
1957
1958 #[test]
1959 fn rejects_fp32_identity() {
1960 let bytes = identity_graph(DType::FP32, vec![1, 4, 1024])
1962 .build(XDNA_TOSA_TARGET)
1963 .expect("build fp32 identity");
1964 assert_eq!(
1965 admit(&bytes, XDNA_TOSA_TARGET),
1966 Err(AdmitError::Unsupported)
1967 );
1968 }
1969
1970 #[test]
1971 fn rejects_non_multiple_of_line_size() {
1972 let bytes = identity_graph(DType::BF16, vec![1, 1, 100])
1973 .build(XDNA_TOSA_TARGET)
1974 .expect("build small identity");
1975 assert_eq!(
1976 admit(&bytes, XDNA_TOSA_TARGET),
1977 Err(AdmitError::Unsupported)
1978 );
1979 }
1980
1981 #[test]
1982 fn rejects_bf16_artifact_under_integer_target() {
1983 let bytes = identity_graph(DType::BF16, vec![1, 4, 1024])
1984 .build(XDNA_TOSA_TARGET)
1985 .expect("build");
1986 assert_eq!(
1987 admit(&bytes, XDNA_TOSA_INTEGER_TARGET),
1988 Err(AdmitError::Analysis)
1989 );
1990 }
1991}