1#![cfg_attr(not(va_openvino), allow(dead_code))]
18
19use std::fmt;
20use std::fmt::Write as _;
21
22use virtio_accel_tosa::{
23 AnalysisError, AnalyzedValueKind, CapabilityDescriptor, DType, DTypeCapability,
24 DTypeConstraints, Error as ParseError, ExtensionSet, GraphCapabilities, Level,
25 NanPropagationMode, Op, OpAttributes, OperatorCapability, OperatorConstraints, ProfileSet,
26 RuntimeConditionSupport, Target, TosaAnalysis, ValueId, ValueRoles, Version, parse,
27};
28
29pub const OPENVINO_TOSA_TARGET: Target = Target::new(
31 Version::TOSA_1_0,
32 ProfileSet::FLOATING_POINT,
33 Level::Level8K,
34 ExtensionSet::NONE,
35);
36
37pub const OPENVINO_TOSA_INTEGER_TARGET: Target = Target::new(
39 Version::TOSA_1_0,
40 ProfileSet::INTEGER,
41 Level::Level8K,
42 ExtensionSet::NONE,
43);
44
45pub const OPENVINO_TOSA_FP8_TARGET: Target = Target::new(
51 Version::TOSA_1_0,
52 ProfileSet::FLOATING_POINT,
53 Level::Level8K,
54 ExtensionSet::FP8E4M3.union(ExtensionSet::FP8E5M2),
55);
56
57const FLOAT_DTYPES: &[DTypeCapability] = &[
58 DTypeCapability::new(DType::FP16, ValueRoles::ALL),
59 DTypeCapability::new(DType::FP32, ValueRoles::ALL),
60 DTypeCapability::new(DType::BOOL, ValueRoles::ALL),
61 DTypeCapability::new(DType::INT32, ValueRoles::ALL),
62 DTypeCapability::constrained(
63 DType::INT8,
64 ValueRoles::CONSTANT,
65 DTypeConstraints::PARAMETER_ONLY,
66 ),
67];
68
69const INTEGER_DTYPES: &[DTypeCapability] = &[
70 DTypeCapability::new(DType::INT8, ValueRoles::ALL),
71 DTypeCapability::new(
72 DType::INT32,
73 ValueRoles::OUTPUT
74 .union(ValueRoles::CONSTANT)
75 .union(ValueRoles::INTERMEDIATE),
76 ),
77];
78
79const FLOAT_OPERATORS: &[OperatorCapability] = &[
80 OperatorCapability::constrained(Op::ARGMAX, OperatorConstraints::PROPAGATING_NAN),
81 OperatorCapability::constrained(Op::MATMUL, OperatorConstraints::ZERO_ZERO_POINTS),
82 OperatorCapability::constrained(
83 Op::MAX_POOL2D,
84 OperatorConstraints::PROPAGATING_NAN.union(OperatorConstraints::ZERO_PADDING),
85 ),
86 OperatorCapability::constrained(Op::CLAMP, OperatorConstraints::PROPAGATING_NAN),
87 OperatorCapability::new(Op::ERF),
88 OperatorCapability::new(Op::SIGMOID),
89 OperatorCapability::new(Op::TANH),
90 OperatorCapability::new(Op::ADD),
91 OperatorCapability::new(Op::LOGICAL_AND),
92 OperatorCapability::new(Op::LOGICAL_OR),
93 OperatorCapability::new(Op::LOGICAL_XOR),
94 OperatorCapability::constrained(Op::MAXIMUM, OperatorConstraints::PROPAGATING_NAN),
95 OperatorCapability::constrained(Op::MINIMUM, OperatorConstraints::PROPAGATING_NAN),
96 OperatorCapability::constrained(Op::MUL, OperatorConstraints::ZERO_SHIFT),
97 OperatorCapability::new(Op::POW),
98 OperatorCapability::new(Op::SUB),
99 OperatorCapability::new(Op::ABS),
100 OperatorCapability::new(Op::CEIL),
101 OperatorCapability::new(Op::COS),
102 OperatorCapability::new(Op::EXP),
103 OperatorCapability::new(Op::FLOOR),
104 OperatorCapability::new(Op::LOG),
105 OperatorCapability::new(Op::LOGICAL_NOT),
106 OperatorCapability::constrained(Op::NEGATE, OperatorConstraints::ZERO_ZERO_POINTS),
107 OperatorCapability::new(Op::RECIPROCAL),
108 OperatorCapability::new(Op::RSQRT),
109 OperatorCapability::new(Op::SIN),
110 OperatorCapability::new(Op::SELECT),
111 OperatorCapability::new(Op::EQUAL),
112 OperatorCapability::new(Op::GREATER),
113 OperatorCapability::new(Op::GREATER_EQUAL),
114 OperatorCapability::constrained(Op::REDUCE_MAX, OperatorConstraints::PROPAGATING_NAN),
115 OperatorCapability::constrained(Op::REDUCE_MIN, OperatorConstraints::PROPAGATING_NAN),
116 OperatorCapability::new(Op::REDUCE_PRODUCT),
117 OperatorCapability::new(Op::REDUCE_SUM),
118 OperatorCapability::new(Op::CONCAT),
119 OperatorCapability::constrained(Op::RESHAPE, OperatorConstraints::CONSTANT_PARAMETERS),
120 OperatorCapability::new(Op::REVERSE),
121 OperatorCapability::new(Op::TRANSPOSE),
122 OperatorCapability::new(Op::CONST),
123 OperatorCapability::new(Op::CONST_SHAPE),
124 OperatorCapability::new(Op::IDENTITY),
125];
126
127const CAST_CAPABILITY: OperatorCapability = OperatorCapability::new(Op::CAST);
135
136const FLOAT8_OPERATORS: &[OperatorCapability] = &{
140 let mut extended = [CAST_CAPABILITY; FLOAT_OPERATORS.len() + 1];
141 let mut index = 0;
142 while index < FLOAT_OPERATORS.len() {
143 extended[index] = FLOAT_OPERATORS[index];
144 index += 1;
145 }
146 extended
147};
148
149const FLOAT8_DTYPES: &[DTypeCapability] = &[
153 DTypeCapability::new(DType::FP8E4M3, ValueRoles::ALL),
154 DTypeCapability::new(DType::FP8E5M2, ValueRoles::ALL),
155 DTypeCapability::new(DType::FP16, ValueRoles::ALL),
156 DTypeCapability::new(DType::INT32, ValueRoles::ALL),
157];
158
159const INTEGER_OPERATORS: &[OperatorCapability] = &[
160 OperatorCapability::new(Op::CONST),
161 OperatorCapability::new(Op::IDENTITY),
162 OperatorCapability::new(Op::MATMUL),
163];
164
165pub const OPENVINO_TOSA_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
167 target: OPENVINO_TOSA_TARGET,
168 dtypes: FLOAT_DTYPES,
169 operators: FLOAT_OPERATORS,
170 graph: GraphCapabilities {
171 max_regions: 1,
172 max_blocks: 1,
173 dynamic_shapes: false,
174 runtime_conditions: RuntimeConditionSupport::None,
175 },
176};
177
178pub const OPENVINO_TOSA_INTEGER_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
180 target: OPENVINO_TOSA_INTEGER_TARGET,
181 dtypes: INTEGER_DTYPES,
182 operators: INTEGER_OPERATORS,
183 graph: GraphCapabilities {
184 max_regions: 1,
185 max_blocks: 1,
186 dynamic_shapes: false,
187 runtime_conditions: RuntimeConditionSupport::None,
188 },
189};
190
191pub const OPENVINO_TOSA_FP8_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
193 target: OPENVINO_TOSA_FP8_TARGET,
194 dtypes: FLOAT8_DTYPES,
195 operators: FLOAT8_OPERATORS,
202 graph: GraphCapabilities {
203 max_regions: 1,
204 max_blocks: 1,
205 dynamic_shapes: false,
206 runtime_conditions: RuntimeConditionSupport::None,
207 },
208};
209
210fn capability_for(target: Target) -> Option<&'static CapabilityDescriptor> {
216 if target == OPENVINO_TOSA_TARGET {
217 Some(&OPENVINO_TOSA_CAPABILITY)
218 } else if target == OPENVINO_TOSA_INTEGER_TARGET {
219 Some(&OPENVINO_TOSA_INTEGER_CAPABILITY)
220 } else if target == OPENVINO_TOSA_FP8_TARGET {
221 Some(&OPENVINO_TOSA_FP8_CAPABILITY)
222 } else {
223 None
224 }
225}
226
227const WEIGHTS_ALIGNMENT: usize = 64;
229
230#[derive(Clone, Copy, Debug, PartialEq, Eq)]
231pub enum LoweringError {
232 Parse(ParseError),
233 Analysis(AnalysisError),
234 UnsupportedGraph,
235 UnsupportedType(DType),
236 UnsupportedOperator(Op),
237 InvalidConstant,
238 ResourceLimit,
239}
240
241impl fmt::Display for LoweringError {
242 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
243 write!(formatter, "{self:?}")
244 }
245}
246
247impl std::error::Error for LoweringError {}
248
249#[derive(Clone, Copy, Debug, PartialEq, Eq)]
251pub(crate) enum OvElement {
252 F32,
253 F16,
254 F8E4M3,
255 F8E5M2,
256 I8,
257 I32,
258 I64,
259 Bool,
260}
261
262impl OvElement {
263 pub(crate) const fn element_type(self) -> &'static str {
265 match self {
266 Self::F32 => "f32",
267 Self::F16 => "f16",
268 Self::F8E4M3 => "f8e4m3",
269 Self::F8E5M2 => "f8e5m2",
270 Self::I8 => "i8",
271 Self::I32 => "i32",
272 Self::I64 => "i64",
273 Self::Bool => "boolean",
274 }
275 }
276
277 const fn precision(self) -> &'static str {
279 match self {
280 Self::F32 => "FP32",
281 Self::F16 => "FP16",
282 Self::F8E4M3 => "F8E4M3",
284 Self::F8E5M2 => "F8E5M2",
285 Self::I8 => "I8",
286 Self::I32 => "I32",
287 Self::I64 => "I64",
288 Self::Bool => "BOOL",
289 }
290 }
291
292 pub(crate) const fn scalar_bytes(self) -> u32 {
294 match self {
295 Self::F32 | Self::I32 => 4,
296 Self::F16 => 2,
297 Self::F8E4M3 | Self::F8E5M2 | Self::I8 => 1,
298 Self::I64 => 8,
299 Self::Bool => 1,
300 }
301 }
302
303 fn for_dtype(dtype: DType) -> Result<Self, LoweringError> {
304 match dtype {
305 DType::FP32 => Ok(Self::F32),
306 DType::FP16 => Ok(Self::F16),
307 DType::FP8E4M3 => Ok(Self::F8E4M3),
308 DType::FP8E5M2 => Ok(Self::F8E5M2),
309 DType::INT8 => Ok(Self::I8),
310 DType::INT32 => Ok(Self::I32),
311 DType::BOOL => Ok(Self::Bool),
312 _ => Err(LoweringError::UnsupportedType(dtype)),
313 }
314 }
315}
316
317fn boundary_element(dtype: DType) -> Result<OvElement, LoweringError> {
319 match dtype {
320 DType::FP16 => Ok(OvElement::F16),
321 DType::FP32 => Ok(OvElement::F32),
322 DType::FP8E4M3 => Ok(OvElement::F8E4M3),
323 DType::FP8E5M2 => Ok(OvElement::F8E5M2),
324 DType::INT8 => Ok(OvElement::I8),
325 DType::INT32 => Ok(OvElement::I32),
326 DType::BOOL => Ok(OvElement::Bool),
327 _ => Err(LoweringError::UnsupportedType(dtype)),
328 }
329}
330
331#[derive(Clone, Copy, Debug, PartialEq, Eq)]
332pub(crate) enum LoweredFeatureRole {
333 Input,
334 Output,
335}
336
337#[derive(Clone, Debug, PartialEq, Eq)]
339pub(crate) struct LoweredFeature {
340 pub slot: u32,
341 pub role: LoweredFeatureRole,
342 pub io_index: u32,
344 pub element: OvElement,
345 pub dims: Vec<i64>,
346 pub byte_len: u64,
348}
349
350#[derive(Clone, Debug)]
351pub(crate) struct LoweredModel {
352 pub xml: Vec<u8>,
353 pub weights: Vec<u8>,
354 pub features: Vec<LoweredFeature>,
355}
356
357pub(crate) fn fp8_probe_document() -> LoweredModel {
365 const ELEMENT: OvElement = OvElement::F8E4M3;
366 const DIMS: &[i64] = &[1];
367 let mut builder = IrBuilder::new(2).expect("a three-layer document is always within bounds");
369 let parameter = builder.emit_layer(
370 "Parameter",
371 "opset1",
372 "fp8_probe_input",
373 &format!("shape=\"1\" element_type=\"{}\"", ELEMENT.element_type()),
374 &[],
375 &[(ELEMENT, DIMS)],
376 )[0];
377 let identity = builder.emit_layer(
378 "Convert",
379 "opset1",
380 "fp8_probe_identity",
381 &format!("destination_type=\"{}\"", ELEMENT.element_type()),
382 &[(parameter, DIMS)],
383 &[(ELEMENT, DIMS)],
384 )[0];
385 builder.emit_layer(
386 "Result",
387 "opset1",
388 "fp8_probe_output",
389 "",
390 &[(identity, DIMS)],
391 &[],
392 );
393 builder.finish()
394}
395
396pub const fn supports_tosa_operator(op: Op) -> bool {
398 OPENVINO_TOSA_CAPABILITY.supports_operator(op)
399}
400
401pub const fn supports_tosa_dtype(dtype: DType) -> bool {
407 OPENVINO_TOSA_CAPABILITY.supports_dtype(dtype, ValueRoles::INPUT)
408 || OPENVINO_TOSA_CAPABILITY.supports_dtype(dtype, ValueRoles::OUTPUT)
409 || OPENVINO_TOSA_INTEGER_CAPABILITY.supports_dtype(dtype, ValueRoles::INPUT)
410 || OPENVINO_TOSA_INTEGER_CAPABILITY.supports_dtype(dtype, ValueRoles::OUTPUT)
411 || OPENVINO_TOSA_FP8_CAPABILITY.supports_dtype(dtype, ValueRoles::INPUT)
412 || OPENVINO_TOSA_FP8_CAPABILITY.supports_dtype(dtype, ValueRoles::OUTPUT)
413}
414
415#[derive(Clone, Copy, Debug, PartialEq, Eq)]
417struct PortRef {
418 layer: u32,
419 port: u32,
420}
421
422struct IrBuilder {
424 layers: String,
425 edges: String,
426 weights: Vec<u8>,
427 next_layer: u32,
428 sources: Vec<Option<PortRef>>,
430}
431
432impl IrBuilder {
433 fn new(values: usize) -> Result<Self, LoweringError> {
434 let mut sources = Vec::new();
435 sources
436 .try_reserve_exact(values)
437 .map_err(|_| LoweringError::ResourceLimit)?;
438 sources.resize(values, None);
439 Ok(Self {
440 layers: String::new(),
441 edges: String::new(),
442 weights: Vec::new(),
443 next_layer: 0,
444 sources,
445 })
446 }
447
448 fn source(&self, value: ValueId) -> Result<PortRef, LoweringError> {
449 self.sources[value.get() as usize].ok_or(LoweringError::UnsupportedGraph)
450 }
451
452 fn set_source(&mut self, value: ValueId, port: PortRef) {
453 self.sources[value.get() as usize] = Some(port);
454 }
455
456 fn emit_layer(
461 &mut self,
462 kind: &str,
463 version: &str,
464 name: &str,
465 data: &str,
466 inputs: &[(PortRef, &[i64])],
467 outputs: &[(OvElement, &[i64])],
468 ) -> Vec<PortRef> {
469 let layer = self.next_layer;
470 self.next_layer += 1;
471 let _ = write!(
472 self.layers,
473 "<layer id=\"{layer}\" name=\"{name}\" type=\"{kind}\" version=\"{version}\">"
474 );
475 if !data.is_empty() {
476 let _ = write!(self.layers, "<data {data}/>");
477 }
478 if !inputs.is_empty() {
479 self.layers.push_str("<input>");
480 for (port, (source, dims)) in inputs.iter().enumerate() {
481 let port = port as u32;
482 let _ = write!(self.layers, "<port id=\"{port}\">");
483 Self::write_dims(&mut self.layers, dims);
484 self.layers.push_str("</port>");
485 let _ = write!(
486 self.edges,
487 "<edge from-layer=\"{}\" from-port=\"{}\" to-layer=\"{layer}\" to-port=\"{port}\"/>",
488 source.layer, source.port
489 );
490 }
491 self.layers.push_str("</input>");
492 }
493 let mut ports = Vec::with_capacity(outputs.len());
494 if !outputs.is_empty() {
495 self.layers.push_str("<output>");
496 for (index, (element, dims)) in outputs.iter().enumerate() {
497 let port = (inputs.len() + index) as u32;
498 let _ = write!(
499 self.layers,
500 "<port id=\"{port}\" precision=\"{}\">",
501 element.precision()
502 );
503 Self::write_dims(&mut self.layers, dims);
504 self.layers.push_str("</port>");
505 ports.push(PortRef { layer, port });
506 }
507 self.layers.push_str("</output>");
508 }
509 self.layers.push_str("</layer>");
510 ports
511 }
512
513 fn write_dims(target: &mut String, dims: &[i64]) {
514 for dim in dims {
515 let _ = write!(target, "<dim>{dim}</dim>");
516 }
517 }
518
519 fn emit_const(
521 &mut self,
522 name: &str,
523 element: OvElement,
524 dims: &[i64],
525 bytes: &[u8],
526 ) -> Result<PortRef, LoweringError> {
527 let expected = element_byte_len(element, dims)?;
528 if expected != bytes.len() as u64 {
529 return Err(LoweringError::InvalidConstant);
530 }
531 let padding = self.weights.len().next_multiple_of(WEIGHTS_ALIGNMENT) - self.weights.len();
532 self.weights
533 .try_reserve_exact(padding + bytes.len())
534 .map_err(|_| LoweringError::ResourceLimit)?;
535 self.weights.resize(self.weights.len() + padding, 0);
536 let offset = self.weights.len();
537 self.weights.extend_from_slice(bytes);
538 let mut shape = String::new();
539 for (index, dim) in dims.iter().enumerate() {
540 if index > 0 {
541 shape.push(',');
542 }
543 let _ = write!(shape, "{dim}");
544 }
545 let data = format!(
546 "element_type=\"{}\" shape=\"{shape}\" offset=\"{offset}\" size=\"{}\"",
547 element.element_type(),
548 bytes.len()
549 );
550 let ports = self.emit_layer("Const", "opset1", name, &data, &[], &[(element, dims)]);
551 Ok(ports[0])
552 }
553
554 fn emit_i64_const(
556 &mut self,
557 name: &str,
558 values: &[i64],
559 rank0: bool,
560 ) -> Result<PortRef, LoweringError> {
561 let mut bytes = Vec::new();
562 bytes
563 .try_reserve_exact(values.len() * 8)
564 .map_err(|_| LoweringError::ResourceLimit)?;
565 for value in values {
566 bytes.extend_from_slice(&value.to_le_bytes());
567 }
568 let dims = [values.len() as i64];
569 let dims: &[i64] = if rank0 { &[] } else { &dims };
570 self.emit_const(name, OvElement::I64, dims, &bytes)
571 }
572
573 fn finish(self) -> LoweredModel {
574 let mut document = String::with_capacity(
575 self.layers.len()
576 + self.edges.len()
577 + "<net name=\"tosa\" version=\"11\"><layers></layers><edges></edges></net>".len()
578 + 24,
579 );
580 document.push_str("<?xml version=\"1.0\"?>");
581 document.push_str("<net name=\"tosa\" version=\"11\">");
582 document.push_str("<layers>");
583 document.push_str(&self.layers);
584 document.push_str("</layers>");
585 document.push_str("<edges>");
586 document.push_str(&self.edges);
587 document.push_str("</edges>");
588 document.push_str("</net>");
589 LoweredModel {
590 xml: document.into_bytes(),
591 weights: self.weights,
592 features: Vec::new(),
593 }
594 }
595}
596
597fn element_byte_len(element: OvElement, dims: &[i64]) -> Result<u64, LoweringError> {
598 let mut total = u64::from(element.scalar_bytes());
599 for dim in dims {
600 let dim = u64::try_from(*dim).map_err(|_| LoweringError::UnsupportedGraph)?;
601 total = total
602 .checked_mul(dim)
603 .ok_or(LoweringError::UnsupportedGraph)?;
604 }
605 Ok(total)
606}
607
608pub(crate) fn lower_tosa(bytes: &[u8], target: Target) -> Result<LoweredModel, LoweringError> {
609 let capability = capability_for(target).ok_or(LoweringError::UnsupportedGraph)?;
610 let model = parse(bytes).map_err(LoweringError::Parse)?;
611 let analysis = model.analyze_for(target).map_err(LoweringError::Analysis)?;
612 validate_target_types(&analysis, target)?;
613 if analysis.regions().len() != 1
614 || analysis.blocks().len() != 1
615 || !analysis.conditions().is_empty()
616 {
617 return Err(LoweringError::UnsupportedGraph);
618 }
619 let block = analysis.blocks()[0].id();
620 let inputs = analysis.block_inputs(block);
621 let outputs = analysis.block_outputs(block);
622 if inputs.is_empty()
623 || outputs.is_empty()
624 || inputs.iter().any(|input| outputs.contains(input))
625 || inputs.len().checked_add(outputs.len()).is_none()
626 {
627 return Err(LoweringError::UnsupportedGraph);
628 }
629
630 let mut builder = IrBuilder::new(analysis.values().len())?;
631 let mut features = Vec::new();
632 features
633 .try_reserve_exact(inputs.len() + outputs.len())
634 .map_err(|_| LoweringError::ResourceLimit)?;
635
636 for (index, value) in inputs.iter().copied().enumerate() {
637 let tensor = tensor(&analysis, value)?;
638 let element = boundary_element(tensor.dtype())?;
639 let dims = feature_dims(tensor)?;
640 let byte_len = element_byte_len(element, &dims)?;
641 let name = format!("input_{index}");
642 let mut shape = String::new();
643 for (position, dim) in dims.iter().enumerate() {
644 if position > 0 {
645 shape.push(',');
646 }
647 let _ = write!(shape, "{dim}");
648 }
649 let data = format!(
650 "shape=\"{shape}\" element_type=\"{}\"",
651 element.element_type()
652 );
653 let ports = builder.emit_layer(
654 "Parameter",
655 "opset1",
656 &name,
657 &data,
658 &[],
659 &[(element, dims.as_slice())],
660 );
661 builder.set_source(value, ports[0]);
662 features.push(LoweredFeature {
663 slot: u32::try_from(index).map_err(|_| LoweringError::ResourceLimit)?,
664 role: LoweredFeatureRole::Input,
665 io_index: u32::try_from(index).map_err(|_| LoweringError::ResourceLimit)?,
666 element,
667 dims,
668 byte_len,
669 });
670 }
671
672 for operator in analysis.execution_order(block) {
673 encode_operator(&mut builder, &analysis, *operator, capability)?;
674 }
675
676 for (index, value) in outputs.iter().copied().enumerate() {
677 let tensor = tensor(&analysis, value)?;
678 let element = boundary_element(tensor.dtype())?;
679 let dims = feature_dims(tensor)?;
680 let byte_len = element_byte_len(element, &dims)?;
681 let source = builder.source(value)?;
682 let name = format!("output_{index}");
683 builder.emit_layer(
684 "Result",
685 "opset1",
686 &name,
687 "",
688 &[(source, dims.as_slice())],
689 &[],
690 );
691 features.push(LoweredFeature {
692 slot: u32::try_from(inputs.len() + index).map_err(|_| LoweringError::ResourceLimit)?,
693 role: LoweredFeatureRole::Output,
694 io_index: u32::try_from(index).map_err(|_| LoweringError::ResourceLimit)?,
695 element,
696 dims,
697 byte_len,
698 });
699 }
700
701 let mut lowered = builder.finish();
702 lowered.features = features;
703 Ok(lowered)
704}
705
706fn validate_target_types(analysis: &TosaAnalysis<'_>, target: Target) -> Result<(), LoweringError> {
711 for value in analysis.values() {
712 let AnalyzedValueKind::Tensor(tensor) = value.kind() else {
713 continue;
714 };
715 let dtype = tensor.dtype();
716 let mismatched = if target == OPENVINO_TOSA_TARGET {
717 dtype == DType::INT8
718 && !(analysis.serialized_constant(value.id()).is_some()
719 && constant_is_parameter_only(analysis, value.id()))
720 } else if target == OPENVINO_TOSA_FP8_TARGET {
721 !matches!(
726 dtype,
727 DType::FP8E4M3 | DType::FP8E5M2 | DType::FP16 | DType::INT32
728 )
729 } else {
730 matches!(dtype, DType::FP16 | DType::FP32)
731 };
732 if mismatched {
733 return Err(LoweringError::UnsupportedType(dtype));
734 }
735 }
736 Ok(())
737}
738
739fn tensor<'a>(
740 analysis: &'a TosaAnalysis<'a>,
741 value: ValueId,
742) -> Result<virtio_accel_tosa::Tensor<'a>, LoweringError> {
743 match analysis.value(value).kind() {
744 AnalyzedValueKind::Tensor(tensor) => Ok(tensor),
745 AnalyzedValueKind::Shape(_) => Err(LoweringError::UnsupportedGraph),
746 }
747}
748
749fn static_dims(tensor: virtio_accel_tosa::Tensor<'_>) -> Result<Vec<i64>, LoweringError> {
751 tensor.rank().ok_or(LoweringError::UnsupportedGraph)?;
752 let dims = tensor.dimensions().map(i64::from).collect::<Vec<_>>();
753 if dims.iter().any(|dimension| *dimension <= 0) {
754 return Err(LoweringError::UnsupportedGraph);
755 }
756 Ok(dims)
757}
758
759fn feature_dims(tensor: virtio_accel_tosa::Tensor<'_>) -> Result<Vec<i64>, LoweringError> {
761 let dims = static_dims(tensor)?;
762 if dims.is_empty() {
763 return Err(LoweringError::UnsupportedGraph);
764 }
765 Ok(dims)
766}
767
768fn value_port_dims(
769 builder: &IrBuilder,
770 analysis: &TosaAnalysis<'_>,
771 value: ValueId,
772) -> Result<(PortRef, Vec<i64>), LoweringError> {
773 Ok((
774 builder.source(value)?,
775 static_dims(tensor(analysis, value)?)?,
776 ))
777}
778
779fn encode_operator(
780 builder: &mut IrBuilder,
781 analysis: &TosaAnalysis<'_>,
782 operator_id: virtio_accel_tosa::OperatorId,
783 capability: &CapabilityDescriptor,
784) -> Result<(), LoweringError> {
785 let operator = analysis.operator(operator_id);
786 let op = operator.op();
787 if !capability.supports_operator(op) {
788 return Err(LoweringError::UnsupportedOperator(op));
789 }
790 let all_inputs = analysis.operator_inputs(operator_id);
791 let outputs = analysis.operator_outputs(operator_id);
792 if op == Op::MATMUL && tensor(analysis, all_inputs[0])?.dtype() == DType::INT8 {
793 return encode_int8_matmul(builder, analysis, operator_id, all_inputs, outputs);
794 }
795 let inputs = match op {
796 Op::MATMUL => {
797 for zero_point in &all_inputs[2..4] {
798 let bytes = analysis
799 .serialized_constant(*zero_point)
800 .ok_or(LoweringError::UnsupportedGraph)?;
801 if !serialized_float_is_zero(tensor(analysis, *zero_point)?.dtype(), bytes) {
802 return Err(LoweringError::UnsupportedGraph);
803 }
804 }
805 &all_inputs[..2]
806 }
807 Op::MUL => {
808 let shift = analysis
809 .serialized_constant(all_inputs[2])
810 .ok_or(LoweringError::UnsupportedGraph)?;
811 if shift.iter().any(|byte| *byte != 0) {
812 return Err(LoweringError::UnsupportedGraph);
813 }
814 &all_inputs[..2]
815 }
816 Op::NEGATE => {
817 for zero_point in &all_inputs[1..3] {
818 let bytes = analysis
819 .serialized_constant(*zero_point)
820 .ok_or(LoweringError::UnsupportedGraph)?;
821 if bytes.iter().any(|byte| *byte != 0) {
822 return Err(LoweringError::UnsupportedGraph);
823 }
824 }
825 &all_inputs[..1]
826 }
827 Op::RESHAPE => {
828 analysis
829 .serialized_constant(all_inputs[1])
830 .ok_or(LoweringError::UnsupportedGraph)?;
831 &all_inputs[..1]
832 }
833 _ => all_inputs,
834 };
835
836 if op == Op::CONST_SHAPE {
839 return Ok(());
840 }
841 if op == Op::CONST {
842 let output = outputs[0];
843 if constant_is_parameter_only(analysis, output) {
844 return Ok(());
845 }
846 return encode_tosa_constant(builder, analysis, operator_id, output);
847 }
848
849 validate_operator_types(analysis, op, inputs, outputs)?;
850 match operator.source().attributes() {
851 OpAttributes::Maximum { nan_mode } | OpAttributes::Minimum { nan_mode } => {
852 require_propagating_nan(nan_mode)?;
853 }
854 _ => {}
855 }
856
857 let stem = format!("tosa_{}_{}", operator_id.get(), op.name().unwrap_or("op"));
858 if op == Op::MAX_POOL2D {
859 return encode_max_pool2d(builder, analysis, operator_id, inputs, outputs, &stem);
860 }
861 if op == Op::ARGMAX {
862 return encode_argmax(builder, analysis, operator_id, inputs, outputs, &stem);
863 }
864 if matches!(op, Op::RECIPROCAL | Op::RSQRT) {
865 return encode_negative_power(builder, analysis, op, inputs, outputs, &stem);
866 }
867
868 let output_tensor = tensor(analysis, outputs[0])?;
869 let output_element = OvElement::for_dtype(output_tensor.dtype())?;
870 let output_dims = static_dims(output_tensor)?;
871
872 const NUMPY: &str = "auto_broadcast=\"numpy\"";
874 let (kind, data) = match op {
875 Op::IDENTITY => (
880 "Convert",
881 format!("destination_type=\"{}\"", output_element.element_type()),
882 ),
883 Op::CAST => (
884 "Convert",
885 format!("destination_type=\"{}\"", output_element.element_type()),
886 ),
887 Op::ADD => ("Add", NUMPY.to_owned()),
888 Op::SUB => ("Subtract", NUMPY.to_owned()),
889 Op::MUL => ("Multiply", NUMPY.to_owned()),
890 Op::POW => ("Power", NUMPY.to_owned()),
891 Op::MAXIMUM => ("Maximum", NUMPY.to_owned()),
892 Op::MINIMUM => ("Minimum", NUMPY.to_owned()),
893 Op::EQUAL => ("Equal", NUMPY.to_owned()),
894 Op::GREATER => ("Greater", NUMPY.to_owned()),
895 Op::GREATER_EQUAL => ("GreaterEqual", NUMPY.to_owned()),
896 Op::LOGICAL_AND => ("LogicalAnd", NUMPY.to_owned()),
897 Op::LOGICAL_OR => ("LogicalOr", NUMPY.to_owned()),
898 Op::LOGICAL_XOR => ("LogicalXor", NUMPY.to_owned()),
899 Op::SELECT => ("Select", NUMPY.to_owned()),
900 Op::LOGICAL_NOT => ("LogicalNot", String::new()),
901 Op::ABS => ("Abs", String::new()),
902 Op::CEIL => ("Ceiling", String::new()),
903 Op::COS => ("Cos", String::new()),
904 Op::ERF => ("Erf", String::new()),
905 Op::EXP => ("Exp", String::new()),
906 Op::FLOOR => ("Floor", String::new()),
907 Op::LOG => ("Log", String::new()),
908 Op::NEGATE => ("Negative", String::new()),
909 Op::SIN => ("Sin", String::new()),
910 Op::SIGMOID => ("Sigmoid", String::new()),
911 Op::TANH => ("Tanh", String::new()),
912 Op::MATMUL => (
913 "MatMul",
914 "transpose_a=\"false\" transpose_b=\"false\"".to_owned(),
915 ),
916 Op::CLAMP => {
917 let OpAttributes::Clamp {
918 min_val,
919 max_val,
920 nan_mode,
921 } = operator.source().attributes()
922 else {
923 return Err(LoweringError::UnsupportedGraph);
924 };
925 require_propagating_nan(nan_mode)?;
926 let dtype = tensor(analysis, inputs[0])?.dtype();
927 let min = attr_float(decode_float(dtype, min_val)?)?;
928 let max = attr_float(decode_float(dtype, max_val)?)?;
929 ("Clamp", format!("min=\"{min}\" max=\"{max}\""))
930 }
931 Op::CONCAT => {
932 let OpAttributes::Concat { axis } = operator.source().attributes() else {
933 return Err(LoweringError::UnsupportedGraph);
934 };
935 ("Concat", format!("axis=\"{axis}\""))
936 }
937 Op::REDUCE_MAX | Op::REDUCE_MIN | Op::REDUCE_PRODUCT | Op::REDUCE_SUM => {
938 let axis = match operator.source().attributes() {
939 OpAttributes::ReduceMax { axis, nan_mode }
940 | OpAttributes::ReduceMin { axis, nan_mode } => {
941 require_propagating_nan(nan_mode)?;
942 axis
943 }
944 OpAttributes::ReduceProduct { axis } | OpAttributes::ReduceSum { axis } => axis,
945 _ => return Err(LoweringError::UnsupportedGraph),
946 };
947 let axes =
948 builder.emit_i64_const(&format!("{stem}_axes"), &[i64::from(axis)], false)?;
949 let (source, dims) = value_port_dims(builder, analysis, inputs[0])?;
950 let kind = match op {
951 Op::REDUCE_MAX => "ReduceMax",
952 Op::REDUCE_MIN => "ReduceMin",
953 Op::REDUCE_SUM => "ReduceSum",
954 _ => "ReduceProd",
955 };
956 let ports = builder.emit_layer(
957 kind,
958 "opset1",
959 &stem,
960 "keep_dims=\"true\"",
961 &[(source, dims.as_slice()), (axes, &[1])],
962 &[(output_element, output_dims.as_slice())],
963 );
964 builder.set_source(outputs[0], ports[0]);
965 return Ok(());
966 }
967 Op::RESHAPE => {
968 let shape = builder.emit_i64_const(&format!("{stem}_shape"), &output_dims, false)?;
969 let (source, dims) = value_port_dims(builder, analysis, inputs[0])?;
970 let ports = builder.emit_layer(
971 "Reshape",
972 "opset1",
973 &stem,
974 "special_zero=\"false\"",
975 &[
976 (source, dims.as_slice()),
977 (shape, &[output_dims.len() as i64]),
978 ],
979 &[(output_element, output_dims.as_slice())],
980 );
981 builder.set_source(outputs[0], ports[0]);
982 return Ok(());
983 }
984 Op::REVERSE => {
985 let OpAttributes::Reverse { axis } = operator.source().attributes() else {
986 return Err(LoweringError::UnsupportedGraph);
987 };
988 let axes =
989 builder.emit_i64_const(&format!("{stem}_axes"), &[i64::from(axis)], false)?;
990 let (source, dims) = value_port_dims(builder, analysis, inputs[0])?;
991 let ports = builder.emit_layer(
992 "Reverse",
993 "opset1",
994 &stem,
995 "mode=\"index\"",
996 &[(source, dims.as_slice()), (axes, &[1])],
997 &[(output_element, output_dims.as_slice())],
998 );
999 builder.set_source(outputs[0], ports[0]);
1000 return Ok(());
1001 }
1002 Op::TRANSPOSE => {
1003 let OpAttributes::Transpose { perms } = operator.source().attributes() else {
1004 return Err(LoweringError::UnsupportedGraph);
1005 };
1006 let perms = perms.iter().map(i64::from).collect::<Vec<_>>();
1007 let order = builder.emit_i64_const(&format!("{stem}_perms"), &perms, false)?;
1008 let (source, dims) = value_port_dims(builder, analysis, inputs[0])?;
1009 let ports = builder.emit_layer(
1010 "Transpose",
1011 "opset1",
1012 &stem,
1013 "",
1014 &[(source, dims.as_slice()), (order, &[perms.len() as i64])],
1015 &[(output_element, output_dims.as_slice())],
1016 );
1017 builder.set_source(outputs[0], ports[0]);
1018 return Ok(());
1019 }
1020 _ => return Err(LoweringError::UnsupportedOperator(op)),
1021 };
1022
1023 let mut connected = Vec::with_capacity(inputs.len());
1024 for (index, value) in inputs.iter().enumerate() {
1025 let (mut source, dims) = value_port_dims(builder, analysis, *value)?;
1026 if let Some(widened) = fp8_widening(op, tensor(analysis, *value)?.dtype()) {
1027 source = builder.emit_layer(
1028 "Convert",
1029 "opset1",
1030 &format!("{stem}_widen_{index}"),
1031 &format!("destination_type=\"{}\"", widened.element_type()),
1032 &[(source, dims.as_slice())],
1033 &[(widened, dims.as_slice())],
1034 )[0];
1035 }
1036 connected.push((source, dims));
1037 }
1038 let connected = connected
1039 .iter()
1040 .map(|(source, dims)| (*source, dims.as_slice()))
1041 .collect::<Vec<_>>();
1042 let ports = builder.emit_layer(
1043 kind,
1044 "opset1",
1045 &stem,
1046 &data,
1047 &connected,
1048 &[(output_element, output_dims.as_slice())],
1049 );
1050 builder.set_source(outputs[0], ports[0]);
1051 Ok(())
1052}
1053
1054fn encode_max_pool2d(
1055 builder: &mut IrBuilder,
1056 analysis: &TosaAnalysis<'_>,
1057 operator_id: virtio_accel_tosa::OperatorId,
1058 inputs: &[ValueId],
1059 outputs: &[ValueId],
1060 stem: &str,
1061) -> Result<(), LoweringError> {
1062 let OpAttributes::MaxPool2d {
1063 kernel,
1064 stride,
1065 pad,
1066 nan_mode,
1067 } = analysis.operator(operator_id).source().attributes()
1068 else {
1069 return Err(LoweringError::UnsupportedGraph);
1070 };
1071 require_propagating_nan(nan_mode)?;
1072 let kernel = kernel.iter().collect::<Vec<_>>();
1073 let stride = stride.iter().collect::<Vec<_>>();
1074 let pad = pad.iter().collect::<Vec<_>>();
1075 if kernel.len() != 2
1076 || stride.len() != 2
1077 || pad.len() != 4
1078 || kernel.iter().chain(&stride).any(|value| *value <= 0)
1079 || pad.iter().any(|value| *value != 0)
1080 {
1081 return Err(LoweringError::UnsupportedGraph);
1082 }
1083
1084 let input_tensor = tensor(analysis, inputs[0])?;
1085 let element = OvElement::for_dtype(input_tensor.dtype())?;
1086 let pool_element = match element {
1094 OvElement::F8E4M3 | OvElement::F8E5M2 => OvElement::F16,
1095 other => other,
1096 };
1097 let widens = pool_element != element;
1098 let nhwc_in = static_dims(input_tensor)?;
1099 let nhwc_out = static_dims(tensor(analysis, outputs[0])?)?;
1100 if nhwc_in.len() != 4 || nhwc_out.len() != 4 {
1101 return Err(LoweringError::UnsupportedGraph);
1102 }
1103 let permute = |dims: &[i64], perms: [usize; 4]| perms.map(|axis| dims[axis]);
1104 let nchw_in = permute(&nhwc_in, [0, 3, 1, 2]);
1105 let nchw_out = permute(&nhwc_out, [0, 3, 1, 2]);
1106
1107 let to_nchw = builder.emit_i64_const(&format!("{stem}_nchw_perms"), &[0, 3, 1, 2], false)?;
1108 let source = builder.source(inputs[0])?;
1109 let nchw_input = builder.emit_layer(
1110 "Transpose",
1111 "opset1",
1112 &format!("{stem}_to_nchw"),
1113 "",
1114 &[(source, nhwc_in.as_slice()), (to_nchw, &[4])],
1115 &[(element, nchw_in.as_slice())],
1116 )[0];
1117
1118 let data = format!(
1119 "strides=\"{},{}\" kernel=\"{},{}\" pads_begin=\"0,0\" pads_end=\"0,0\" \
1120 rounding_type=\"floor\" auto_pad=\"explicit\"",
1121 stride[0], stride[1], kernel[0], kernel[1]
1122 );
1123 let widened_input = if widens {
1124 builder.emit_layer(
1125 "Convert",
1126 "opset1",
1127 &format!("{stem}_widen"),
1128 &format!("destination_type=\"{}\"", pool_element.element_type()),
1129 &[(nchw_input, nchw_in.as_slice())],
1130 &[(pool_element, nchw_in.as_slice())],
1131 )[0]
1132 } else {
1133 nchw_input
1134 };
1135 let pooled = builder.emit_layer(
1136 "MaxPool",
1137 "opset1",
1138 stem,
1139 &data,
1140 &[(widened_input, nchw_in.as_slice())],
1141 &[(pool_element, nchw_out.as_slice())],
1142 )[0];
1143 let pooled = if widens {
1144 builder.emit_layer(
1145 "Convert",
1146 "opset1",
1147 &format!("{stem}_narrow"),
1148 &format!("destination_type=\"{}\"", element.element_type()),
1149 &[(pooled, nchw_out.as_slice())],
1150 &[(element, nchw_out.as_slice())],
1151 )[0]
1152 } else {
1153 pooled
1154 };
1155
1156 let to_nhwc = builder.emit_i64_const(&format!("{stem}_nhwc_perms"), &[0, 2, 3, 1], false)?;
1157 let restored = builder.emit_layer(
1158 "Transpose",
1159 "opset1",
1160 &format!("{stem}_to_nhwc"),
1161 "",
1162 &[(pooled, nchw_out.as_slice()), (to_nhwc, &[4])],
1163 &[(element, nhwc_out.as_slice())],
1164 )[0];
1165 builder.set_source(outputs[0], restored);
1166 Ok(())
1167}
1168
1169fn encode_int8_matmul(
1176 builder: &mut IrBuilder,
1177 analysis: &TosaAnalysis<'_>,
1178 operator_id: virtio_accel_tosa::OperatorId,
1179 inputs: &[ValueId],
1180 outputs: &[ValueId],
1181) -> Result<(), LoweringError> {
1182 if inputs.len() != 4 || outputs.len() != 1 {
1183 return Err(LoweringError::UnsupportedGraph);
1184 }
1185 let output = tensor(analysis, outputs[0])?;
1186 if tensor(analysis, inputs[0])?.dtype() != DType::INT8
1187 || tensor(analysis, inputs[1])?.dtype() != DType::INT8
1188 || output.dtype() != DType::INT32
1189 {
1190 return Err(LoweringError::UnsupportedGraph);
1191 }
1192 let read_zero_point = |value| {
1193 let bytes = analysis
1194 .serialized_constant(value)
1195 .ok_or(LoweringError::UnsupportedGraph)?;
1196 if bytes.len() != 1 || tensor(analysis, value)?.dtype() != DType::INT8 {
1197 return Err(LoweringError::UnsupportedGraph);
1198 }
1199 Ok(i32::from(bytes[0] as i8))
1200 };
1201 let zero_points = [read_zero_point(inputs[2])?, read_zero_point(inputs[3])?];
1202 let stem = format!("tosa_{}_matmul", operator_id.get());
1203 let mut adjusted = Vec::with_capacity(2);
1204 for index in 0..2 {
1205 let dims = static_dims(tensor(analysis, inputs[index])?)?;
1206 let source = builder.source(inputs[index])?;
1207 let widened = builder.emit_layer(
1208 "Convert",
1209 "opset1",
1210 &format!("{stem}_widen_{index}"),
1211 "destination_type=\"i32\"",
1212 &[(source, dims.as_slice())],
1213 &[(OvElement::I32, dims.as_slice())],
1214 )[0];
1215 let zero_point = builder.emit_const(
1216 &format!("{stem}_zero_point_{index}"),
1217 OvElement::I32,
1218 &[],
1219 &zero_points[index].to_le_bytes(),
1220 )?;
1221 let shifted = builder.emit_layer(
1222 "Subtract",
1223 "opset1",
1224 &format!("{stem}_shift_{index}"),
1225 "auto_broadcast=\"numpy\"",
1226 &[(widened, dims.as_slice()), (zero_point, &[])],
1227 &[(OvElement::I32, dims.as_slice())],
1228 )[0];
1229 adjusted.push((shifted, dims));
1230 }
1231 let output_dims = static_dims(output)?;
1232 let result = builder.emit_layer(
1233 "MatMul",
1234 "opset1",
1235 &stem,
1236 "transpose_a=\"false\" transpose_b=\"false\"",
1237 &[
1238 (adjusted[0].0, adjusted[0].1.as_slice()),
1239 (adjusted[1].0, adjusted[1].1.as_slice()),
1240 ],
1241 &[(OvElement::I32, output_dims.as_slice())],
1242 )[0];
1243 builder.set_source(outputs[0], result);
1244 Ok(())
1245}
1246
1247fn encode_argmax(
1251 builder: &mut IrBuilder,
1252 analysis: &TosaAnalysis<'_>,
1253 operator_id: virtio_accel_tosa::OperatorId,
1254 inputs: &[ValueId],
1255 outputs: &[ValueId],
1256 stem: &str,
1257) -> Result<(), LoweringError> {
1258 let OpAttributes::ArgMax { axis, nan_mode } =
1259 analysis.operator(operator_id).source().attributes()
1260 else {
1261 return Err(LoweringError::UnsupportedGraph);
1262 };
1263 require_propagating_nan(nan_mode)?;
1264 let input_tensor = tensor(analysis, inputs[0])?;
1265 let element = OvElement::for_dtype(input_tensor.dtype())?;
1266 let input_dims = static_dims(input_tensor)?;
1267 let axis_index = usize::try_from(axis).map_err(|_| LoweringError::UnsupportedGraph)?;
1268 if axis_index >= input_dims.len() {
1269 return Err(LoweringError::UnsupportedGraph);
1270 }
1271 let mut kept_dims = input_dims.clone();
1272 kept_dims[axis_index] = 1;
1273
1274 let k = builder.emit_i64_const(&format!("{stem}_k"), &[1], true)?;
1275 let source = builder.source(inputs[0])?;
1276 let data = format!(
1277 "axis=\"{axis}\" mode=\"max\" sort=\"value\" stable=\"true\" index_element_type=\"i32\""
1278 );
1279 let topk = builder.emit_layer(
1280 "TopK",
1281 "opset11",
1282 &format!("{stem}_topk"),
1283 &data,
1284 &[(source, input_dims.as_slice()), (k, &[])],
1285 &[
1286 (element, kept_dims.as_slice()),
1287 (OvElement::I32, kept_dims.as_slice()),
1288 ],
1289 );
1290 let indices = topk[1];
1291
1292 let output_dims = static_dims(tensor(analysis, outputs[0])?)?;
1293 let axes = builder.emit_i64_const(&format!("{stem}_axes"), &[i64::from(axis)], false)?;
1294 let squeezed = builder.emit_layer(
1295 "Squeeze",
1296 "opset1",
1297 &format!("{stem}_squeeze"),
1298 "",
1299 &[(indices, kept_dims.as_slice()), (axes, &[1])],
1300 &[(OvElement::I32, output_dims.as_slice())],
1301 )[0];
1302 builder.set_source(outputs[0], squeezed);
1303 Ok(())
1304}
1305
1306fn encode_negative_power(
1313 builder: &mut IrBuilder,
1314 analysis: &TosaAnalysis<'_>,
1315 op: Op,
1316 inputs: &[ValueId],
1317 outputs: &[ValueId],
1318 stem: &str,
1319) -> Result<(), LoweringError> {
1320 let input_tensor = tensor(analysis, inputs[0])?;
1321 let element = OvElement::for_dtype(input_tensor.dtype())?;
1322 let dims = static_dims(input_tensor)?;
1323 let exponent_bytes: &[u8] = match element {
1324 OvElement::F32 => &(-1.0f32).to_le_bytes(),
1325 OvElement::F16 => &0xbc00u16.to_le_bytes(),
1326 _ => return Err(LoweringError::UnsupportedType(input_tensor.dtype())),
1327 };
1328 let exponent = builder.emit_const(&format!("{stem}_exponent"), element, &[], exponent_bytes)?;
1329
1330 let mut source = builder.source(inputs[0])?;
1331 if op == Op::RSQRT {
1332 source = builder.emit_layer(
1333 "Sqrt",
1334 "opset1",
1335 &format!("{stem}_sqrt"),
1336 "",
1337 &[(source, dims.as_slice())],
1338 &[(element, dims.as_slice())],
1339 )[0];
1340 }
1341 let output_dims = static_dims(tensor(analysis, outputs[0])?)?;
1342 let powered = builder.emit_layer(
1343 "Power",
1344 "opset1",
1345 stem,
1346 "auto_broadcast=\"numpy\"",
1347 &[(source, dims.as_slice()), (exponent, &[])],
1348 &[(element, output_dims.as_slice())],
1349 )[0];
1350 builder.set_source(outputs[0], powered);
1351 Ok(())
1352}
1353
1354fn encode_tosa_constant(
1355 builder: &mut IrBuilder,
1356 analysis: &TosaAnalysis<'_>,
1357 operator_id: virtio_accel_tosa::OperatorId,
1358 output: ValueId,
1359) -> Result<(), LoweringError> {
1360 let tensor = tensor(analysis, output)?;
1361 let dtype = tensor.dtype();
1362 if !matches!(dtype, DType::FP16 | DType::FP32 | DType::INT8 | DType::BOOL) {
1363 return Err(LoweringError::UnsupportedType(dtype));
1364 }
1365 let element = OvElement::for_dtype(dtype)?;
1366 let data = analysis
1367 .serialized_constant(output)
1368 .ok_or(LoweringError::InvalidConstant)?;
1369 let dims = static_dims(tensor)?;
1370 let name = format!("tosa_{}_const", operator_id.get());
1371 let port = if element == OvElement::Bool {
1372 let mut normalized = Vec::new();
1374 normalized
1375 .try_reserve_exact(data.len())
1376 .map_err(|_| LoweringError::ResourceLimit)?;
1377 normalized.extend(data.iter().map(|byte| u8::from(*byte != 0)));
1378 builder.emit_const(&name, element, &dims, &normalized)?
1379 } else {
1380 builder.emit_const(&name, element, &dims, data)?
1381 };
1382 builder.set_source(output, port);
1383 Ok(())
1384}
1385
1386fn constant_is_parameter_only(analysis: &TosaAnalysis<'_>, value: ValueId) -> bool {
1387 let mut consumed = false;
1388 for operator in analysis.operators() {
1389 for (index, input) in analysis.operator_inputs(operator.id()).iter().enumerate() {
1390 if *input != value {
1391 continue;
1392 }
1393 consumed = true;
1394 if !matches!(
1395 (operator.op(), index),
1396 (Op::MATMUL, 2 | 3) | (Op::MUL, 2) | (Op::NEGATE, 1 | 2) | (Op::RESHAPE, 1)
1397 ) {
1398 return false;
1399 }
1400 }
1401 }
1402 consumed
1403}
1404
1405fn fp8_widening(op: Op, dtype: DType) -> Option<OvElement> {
1415 let is_fp8 = matches!(dtype, DType::FP8E4M3 | DType::FP8E5M2);
1416 (is_fp8 && op == Op::MATMUL).then_some(OvElement::F16)
1417}
1418
1419fn validate_operator_types(
1420 analysis: &TosaAnalysis<'_>,
1421 op: Op,
1422 inputs: &[ValueId],
1423 outputs: &[ValueId],
1424) -> Result<(), LoweringError> {
1425 let require = |value, predicate: fn(DType) -> bool| {
1426 let dtype = tensor(analysis, value)?.dtype();
1427 if predicate(dtype) {
1428 Ok(())
1429 } else {
1430 Err(LoweringError::UnsupportedType(dtype))
1431 }
1432 };
1433 let is_float = |dtype| matches!(dtype, DType::FP16 | DType::FP32);
1434 let is_float_or_fp8 = |dtype| {
1435 matches!(
1436 dtype,
1437 DType::FP16 | DType::FP32 | DType::FP8E4M3 | DType::FP8E5M2
1438 )
1439 };
1440 let is_movable = |dtype| {
1441 matches!(
1442 dtype,
1443 DType::FP16 | DType::FP32 | DType::INT8 | DType::FP8E4M3 | DType::FP8E5M2
1444 )
1445 };
1446 let is_bool = |dtype| dtype == DType::BOOL;
1447 let is_int32 = |dtype| dtype == DType::INT32;
1448
1449 match op {
1450 Op::IDENTITY => {
1451 for value in inputs.iter().chain(outputs) {
1452 require(*value, is_movable)?;
1453 }
1454 }
1455 Op::CAST => {
1458 for value in inputs.iter().chain(outputs) {
1459 require(*value, is_float_or_fp8)?;
1460 }
1461 }
1462 Op::MATMUL => {
1465 for value in inputs {
1466 require(*value, is_float_or_fp8)?;
1467 }
1468 require(outputs[0], is_float)?;
1469 }
1470 Op::CONCAT | Op::RESHAPE | Op::REVERSE | Op::TRANSPOSE | Op::MAX_POOL2D => {
1473 for value in inputs.iter().chain(outputs) {
1474 require(*value, is_float_or_fp8)?;
1475 }
1476 }
1477 Op::LOGICAL_AND | Op::LOGICAL_OR | Op::LOGICAL_XOR | Op::LOGICAL_NOT => {
1478 for value in inputs.iter().chain(outputs) {
1479 require(*value, is_bool)?;
1480 }
1481 }
1482 Op::EQUAL | Op::GREATER | Op::GREATER_EQUAL => {
1483 for value in inputs {
1484 require(*value, is_float)?;
1485 }
1486 require(outputs[0], is_bool)?;
1487 }
1488 Op::SELECT => {
1489 require(inputs[0], is_bool)?;
1490 for value in inputs[1..].iter().chain(outputs) {
1491 require(*value, is_float)?;
1492 }
1493 }
1494 Op::ARGMAX => {
1495 require(inputs[0], is_float_or_fp8)?;
1496 require(outputs[0], is_int32)?;
1497 }
1498 _ => {
1499 for value in inputs.iter().chain(outputs) {
1500 require(*value, is_float)?;
1501 }
1502 }
1503 }
1504 Ok(())
1505}
1506
1507fn require_propagating_nan(nan_mode: NanPropagationMode) -> Result<(), LoweringError> {
1508 if nan_mode == NanPropagationMode::PROPAGATE {
1509 Ok(())
1510 } else {
1511 Err(LoweringError::UnsupportedGraph)
1512 }
1513}
1514
1515fn attr_float(value: f32) -> Result<String, LoweringError> {
1517 if value.is_nan() {
1518 return Err(LoweringError::UnsupportedGraph);
1519 }
1520 Ok(format!("{value}"))
1521}
1522
1523fn decode_float(dtype: DType, bytes: &[u8]) -> Result<f32, LoweringError> {
1524 match dtype {
1525 DType::FP16 if bytes.len() == 2 => Ok(f16_to_f32(u16::from_le_bytes(
1526 bytes.try_into().expect("length checked"),
1527 ))),
1528 DType::FP32 if bytes.len() == 4 => Ok(f32::from_le_bytes(bytes.try_into().unwrap())),
1529 _ => Err(LoweringError::UnsupportedType(dtype)),
1530 }
1531}
1532
1533fn f16_to_f32(bits: u16) -> f32 {
1534 let sign = u32::from(bits & 0x8000) << 16;
1535 let exponent = (bits >> 10) & 0x1f;
1536 let fraction = u32::from(bits & 0x03ff);
1537 let converted = match exponent {
1538 0 if fraction == 0 => sign,
1539 0 => {
1540 let shift = fraction.leading_zeros() - 21;
1541 let normalized = fraction << shift;
1542 sign | ((127 - 15 - shift + 1) << 23) | ((normalized & 0x03ff) << 13)
1543 }
1544 0x1f => sign | 0x7f80_0000 | (fraction << 13),
1545 _ => sign | ((u32::from(exponent) + 127 - 15) << 23) | (fraction << 13),
1546 };
1547 f32::from_bits(converted)
1548}
1549
1550fn serialized_float_is_zero(dtype: DType, bytes: &[u8]) -> bool {
1551 match dtype {
1552 DType::FP16 if bytes.len() == 2 => {
1553 u16::from_le_bytes(bytes.try_into().expect("length checked")) & 0x7fff == 0
1554 }
1555 DType::FP32 if bytes.len() == 4 => {
1556 u32::from_le_bytes(bytes.try_into().expect("length checked")) & 0x7fff_ffff == 0
1557 }
1558 DType::FP8E4M3 | DType::FP8E5M2 if bytes.len() == 1 => bytes[0] & 0x7f == 0,
1559 _ => false,
1560 }
1561}
1562
1563#[cfg(test)]
1564mod fp8_graphs_impl {
1565 use super::*;
1566 use virtio_accel_tosa_build::{OperatorKind, OwnedGraph, OwnedOperator, OwnedTensor};
1567 pub(crate) fn fp8_matmul_graph(dtype: DType) -> Vec<u8> {
1568 let (m, k, n) = (4, 6, 2);
1569 let mut graph = OwnedGraph::new("main");
1570 graph
1571 .push_tensor(OwnedTensor::new("a", vec![1, m, k], dtype))
1572 .push_tensor(OwnedTensor::new("b", vec![1, k, n], dtype))
1573 .push_tensor(OwnedTensor::constant("a_zp", vec![1], dtype, vec![0]))
1574 .push_tensor(OwnedTensor::constant("b_zp", vec![1], dtype, vec![0]))
1575 .push_tensor(OwnedTensor::new("y", vec![1, m, n], DType::FP16))
1576 .push_operator(OwnedOperator::new(
1577 OperatorKind::Const,
1578 vec![],
1579 vec!["a_zp".into()],
1580 ))
1581 .push_operator(OwnedOperator::new(
1582 OperatorKind::Const,
1583 vec![],
1584 vec!["b_zp".into()],
1585 ))
1586 .push_operator(OwnedOperator::new(
1587 OperatorKind::MatMul,
1588 vec!["a".into(), "b".into(), "a_zp".into(), "b_zp".into()],
1589 vec!["y".into()],
1590 ))
1591 .push_input("a")
1592 .push_input("b")
1593 .push_output("y");
1594 graph.build(OPENVINO_TOSA_FP8_TARGET).unwrap()
1595 }
1596
1597 pub(crate) fn fp8_pool_graph(dtype: DType) -> Vec<u8> {
1598 let mut graph = OwnedGraph::new("main");
1599 graph
1600 .push_tensor(OwnedTensor::new("x", vec![1, 4, 4, 1], dtype))
1601 .push_tensor(OwnedTensor::new("y", vec![1, 2, 2, 1], dtype))
1602 .push_operator(OwnedOperator::new(
1603 OperatorKind::MaxPool2d {
1604 kernel: [2, 2],
1605 stride: [2, 2],
1606 pad: [0; 4],
1607 nan_mode: NanPropagationMode::PROPAGATE,
1608 },
1609 vec!["x".into()],
1610 vec!["y".into()],
1611 ))
1612 .push_input("x")
1613 .push_output("y");
1614 graph.build(OPENVINO_TOSA_FP8_TARGET).unwrap()
1615 }
1616
1617 pub(crate) fn fp8_transpose_graph(dtype: DType) -> Vec<u8> {
1618 let mut graph = OwnedGraph::new("main");
1619 graph
1620 .push_tensor(OwnedTensor::new("x", vec![2, 3], dtype))
1621 .push_tensor(OwnedTensor::new("y", vec![3, 2], dtype))
1622 .push_operator(OwnedOperator::new(
1623 OperatorKind::Transpose {
1624 perms: [1, 0, 0, 0, 0, 0],
1625 rank: 2,
1626 },
1627 vec!["x".into()],
1628 vec!["y".into()],
1629 ))
1630 .push_input("x")
1631 .push_output("y");
1632 graph.build(OPENVINO_TOSA_FP8_TARGET).unwrap()
1633 }
1634}
1635
1636#[cfg(test)]
1637pub(crate) use fp8_graphs_impl::{fp8_matmul_graph, fp8_pool_graph, fp8_transpose_graph};
1638
1639#[cfg(test)]
1640mod tests {
1641 use super::*;
1642
1643 use virtio_accel_conformance::numerics::{
1644 HEXAGON_LOGICAL_CASES, IDENTITY_EDGES_FP16, IDENTITY_EDGES_FP32, IDENTITY_FP8E4M3,
1645 IDENTITY_FP8E5M2, IDENTITY_INT4, IDENTITY_INT8, MATMUL_FP16, MATMUL_FP32, MATMUL_INT8,
1646 MAX_POOL2D_FP16, MAX_POOL2D_FP32, MUL_FP16,
1647 };
1648
1649 const IDENTITY_FP32_LOCAL: &[u8] = include_bytes!("../tests/data/identity-fp32-v1.0.0.tosa");
1650
1651 fn xml_str(lowered: &LoweredModel) -> &str {
1652 core::str::from_utf8(&lowered.xml).expect("lowered documents are UTF-8")
1653 }
1654
1655 #[test]
1656 fn lowers_a_verified_tosa_model_without_host_dependencies() {
1657 let lowered = lower_tosa(IDENTITY_FP32_LOCAL, OPENVINO_TOSA_TARGET).unwrap();
1658 let xml = xml_str(&lowered);
1659 assert!(xml.starts_with("<?xml version=\"1.0\"?><net name=\"tosa\" version=\"11\">"));
1660 assert_eq!(xml.matches("type=\"Parameter\"").count(), 1);
1661 assert_eq!(xml.matches("type=\"Result\"").count(), 1);
1662 assert_eq!(xml.matches("type=\"Const\"").count(), 0);
1663 assert_eq!(xml.matches("type=\"Convert\"").count(), 1);
1666 assert!(xml.contains("destination_type=\"f32\""));
1667 assert!(
1668 xml.contains("<edge from-layer=\"0\" from-port=\"0\" to-layer=\"1\" to-port=\"0\"/>")
1669 );
1670 assert!(
1671 xml.contains("<edge from-layer=\"1\" from-port=\"1\" to-layer=\"2\" to-port=\"0\"/>")
1672 );
1673 assert!(lowered.weights.is_empty());
1674 assert_eq!(lowered.features.len(), 2);
1675 assert_eq!(
1676 (lowered.features[0].slot, lowered.features[0].role),
1677 (0, LoweredFeatureRole::Input)
1678 );
1679 assert_eq!(
1680 (lowered.features[1].slot, lowered.features[1].role),
1681 (1, LoweredFeatureRole::Output)
1682 );
1683 assert_eq!(lowered.features[0].io_index, 0);
1684 assert_eq!(lowered.features[1].io_index, 0);
1685 assert_eq!(lowered.features[0].element, OvElement::F32);
1686 assert_eq!(lowered.features[0].byte_len, lowered.features[1].byte_len);
1687 }
1688
1689 #[test]
1690 fn rejects_a_different_tosa_target_before_parsing() {
1691 let integer_target = Target::new(
1692 Version::TOSA_1_0,
1693 ProfileSet::INTEGER,
1694 Level::Level8K,
1695 ExtensionSet::INT4,
1696 );
1697 assert_eq!(
1698 lower_tosa(IDENTITY_FP32_LOCAL, integer_target).unwrap_err(),
1699 LoweringError::UnsupportedGraph
1700 );
1701 }
1702
1703 #[test]
1708 fn fp8_matmul_reaches_the_device_widened() {
1709 for dtype in [DType::FP8E4M3, DType::FP8E5M2] {
1710 let lowered = lower_tosa(&fp8_matmul_graph(dtype), OPENVINO_TOSA_FP8_TARGET)
1711 .unwrap_or_else(|error| panic!("{dtype:?}: {error:?}"));
1712 let xml = xml_str(&lowered);
1713 assert_eq!(
1714 xml.matches("type=\"Convert\"").count(),
1715 2,
1716 "{dtype:?}: expected one widening Convert per operand\n{xml}"
1717 );
1718 assert_eq!(xml.matches("type=\"MatMul\"").count(), 1, "{dtype:?}");
1719 assert!(xml.contains("destination_type=\"f16\""), "{dtype:?}");
1720 let spelling = if dtype == DType::FP8E4M3 {
1723 "F8E4M3"
1724 } else {
1725 "F8E5M2"
1726 };
1727 assert!(
1728 xml.contains(spelling),
1729 "{dtype:?}: boundary lost its FP8 type"
1730 );
1731 }
1732 }
1733
1734 #[test]
1737 fn fp8_pooling_widens_only_around_the_window() {
1738 for dtype in [DType::FP8E4M3, DType::FP8E5M2] {
1739 let lowered = lower_tosa(&fp8_pool_graph(dtype), OPENVINO_TOSA_FP8_TARGET)
1740 .unwrap_or_else(|error| panic!("{dtype:?}: {error:?}"));
1741 let xml = xml_str(&lowered);
1742 assert_eq!(
1743 xml.matches("type=\"Convert\"").count(),
1744 2,
1745 "{dtype:?}: expected a widen and a narrow\n{xml}"
1746 );
1747 assert_eq!(xml.matches("type=\"MaxPool\"").count(), 1, "{dtype:?}");
1748 assert_eq!(xml.matches("type=\"Transpose\"").count(), 2, "{dtype:?}");
1749 }
1750 }
1751
1752 #[test]
1755 fn fp8_data_movement_never_widens() {
1756 for dtype in [DType::FP8E4M3, DType::FP8E5M2] {
1757 let lowered = lower_tosa(&fp8_transpose_graph(dtype), OPENVINO_TOSA_FP8_TARGET)
1758 .unwrap_or_else(|error| panic!("{dtype:?}: {error:?}"));
1759 let xml = xml_str(&lowered);
1760 assert_eq!(
1761 xml.matches("type=\"Convert\"").count(),
1762 0,
1763 "{dtype:?}: FP8 movement was widened\n{xml}"
1764 );
1765 assert_eq!(xml.matches("type=\"Transpose\"").count(), 1, "{dtype:?}");
1766 }
1767 }
1768
1769 #[test]
1770 fn reports_the_exact_integer_boundary_independently_of_other_low_precision_types() {
1771 assert!(!supports_tosa_dtype(DType::INT4), "INT4");
1774 for dtype in [
1775 DType::FP16,
1776 DType::FP32,
1777 DType::FP8E4M3,
1778 DType::FP8E5M2,
1779 DType::INT8,
1780 DType::INT32,
1781 DType::BOOL,
1782 ] {
1783 assert!(supports_tosa_dtype(dtype), "{dtype:?}");
1784 }
1785 }
1786
1787 #[test]
1788 fn descriptors_separate_parameter_constants_and_integer_boundaries() {
1789 let parameter = OPENVINO_TOSA_CAPABILITY.dtype(DType::INT8).unwrap();
1790 assert_eq!(parameter.roles, ValueRoles::CONSTANT);
1791 assert!(
1792 parameter
1793 .constraints
1794 .contains(DTypeConstraints::PARAMETER_ONLY)
1795 );
1796 assert!(OPENVINO_TOSA_INTEGER_CAPABILITY.supports_dtype(DType::INT8, ValueRoles::INPUT));
1797 assert!(!OPENVINO_TOSA_INTEGER_CAPABILITY.supports_operator(Op::ADD));
1798 let pool = OPENVINO_TOSA_CAPABILITY.operator(Op::MAX_POOL2D).unwrap();
1799 assert!(
1800 pool.constraints
1801 .contains(OperatorConstraints::PROPAGATING_NAN)
1802 );
1803 assert!(pool.constraints.contains(OperatorConstraints::ZERO_PADDING));
1804 }
1805
1806 #[test]
1807 fn lowers_boolean_model_boundaries_as_direct_bytes() {
1808 let case = HEXAGON_LOGICAL_CASES
1809 .iter()
1810 .find(|case| case.name == "logical-or")
1811 .unwrap();
1812 let lowered = lower_tosa(case.artifact, OPENVINO_TOSA_TARGET).unwrap();
1813 let xml = xml_str(&lowered);
1814 assert_eq!(xml.matches("type=\"Parameter\"").count(), 2);
1815 assert_eq!(xml.matches("type=\"LogicalOr\"").count(), 1);
1816 assert_eq!(xml.matches("type=\"Result\"").count(), 1);
1817 assert!(xml.contains("element_type=\"boolean\""));
1818 assert!(xml.contains("precision=\"BOOL\""));
1819 assert_eq!(lowered.features.len(), 3);
1820 for feature in &lowered.features {
1821 assert_eq!(feature.element, OvElement::Bool);
1822 assert_eq!(feature.byte_len, 4);
1823 }
1824 }
1825
1826 #[test]
1827 fn rejects_unimplemented_low_precision_extensions_at_the_declared_target_boundary() {
1828 for (case, extensions, profiles) in [
1829 (IDENTITY_INT4, ExtensionSet::INT4, ProfileSet::INTEGER),
1830 (
1831 IDENTITY_FP8E4M3,
1832 ExtensionSet::FP8E4M3,
1833 ProfileSet::FLOATING_POINT,
1834 ),
1835 (
1836 IDENTITY_FP8E5M2,
1837 ExtensionSet::FP8E5M2,
1838 ProfileSet::FLOATING_POINT,
1839 ),
1840 ] {
1841 let target = Target::new(Version::TOSA_1_0, profiles, Level::Level8K, extensions);
1842 assert_eq!(
1843 lower_tosa(case.artifact, target).unwrap_err(),
1844 LoweringError::UnsupportedGraph,
1845 "{}",
1846 case.name
1847 );
1848 }
1849 }
1850
1851 #[test]
1852 fn lowers_int8_identity_with_direct_byte_boundaries() {
1853 let lowered = lower_tosa(IDENTITY_INT8.artifact, OPENVINO_TOSA_INTEGER_TARGET).unwrap();
1854 let xml = xml_str(&lowered);
1855 assert!(xml.contains("element_type=\"i8\""));
1856 assert!(xml.contains("precision=\"I8\""));
1857 assert!(xml.contains("destination_type=\"i8\""));
1858 assert_eq!(lowered.features[0].element, OvElement::I8);
1859 assert_eq!(lowered.features[0].byte_len, 8);
1860 assert_eq!(lowered.features[1].byte_len, 8);
1861 }
1862
1863 #[test]
1864 fn target_profiles_cannot_admit_the_other_tiers_tensor_types() {
1865 assert_eq!(
1866 lower_tosa(IDENTITY_INT8.artifact, OPENVINO_TOSA_TARGET).unwrap_err(),
1867 LoweringError::UnsupportedType(DType::INT8)
1868 );
1869 assert!(matches!(
1870 lower_tosa(IDENTITY_FP32_LOCAL, OPENVINO_TOSA_INTEGER_TARGET),
1871 Err(LoweringError::Analysis(_))
1872 ));
1873 }
1874
1875 #[test]
1876 fn floating_target_admits_only_parameter_only_int8_constants() {
1877 let lowered = lower_tosa(MUL_FP16.artifact, OPENVINO_TOSA_TARGET).unwrap();
1878 assert!(xml_str(&lowered).contains("type=\"Multiply\""));
1879
1880 assert_eq!(
1882 lower_tosa(IDENTITY_INT8.artifact, OPENVINO_TOSA_TARGET).unwrap_err(),
1883 LoweringError::UnsupportedType(DType::INT8)
1884 );
1885 }
1886
1887 #[test]
1888 fn lowers_int8_matmul_with_explicit_zero_point_legalization() {
1889 let lowered = lower_tosa(MATMUL_INT8.artifact, OPENVINO_TOSA_INTEGER_TARGET).unwrap();
1890 let xml = xml_str(&lowered);
1891 assert_eq!(xml.matches("type=\"Parameter\"").count(), 2);
1892 assert_eq!(xml.matches("type=\"Convert\"").count(), 2);
1893 assert_eq!(xml.matches("type=\"Subtract\"").count(), 2);
1894 assert_eq!(xml.matches("type=\"MatMul\"").count(), 1);
1895 assert!(xml.contains("destination_type=\"i32\""));
1896 assert!(xml.contains("element_type=\"i8\""));
1897 assert!(xml.contains("element_type=\"i32\""));
1898 assert_eq!(lowered.features[0].element, OvElement::I8);
1899 assert_eq!(lowered.features[1].element, OvElement::I8);
1900 assert_eq!(lowered.features[2].element, OvElement::I32);
1901 assert_eq!(lowered.features[2].byte_len, 16);
1902 let zero_points = [
1903 i32::from_le_bytes(lowered.weights[0..4].try_into().unwrap()),
1904 i32::from_le_bytes(lowered.weights[64..68].try_into().unwrap()),
1905 ];
1906 assert_eq!(zero_points, MATMUL_INT8.zero_points.map(i32::from));
1907 }
1908
1909 #[test]
1910 fn lowers_batched_matmul_without_encoding_parameter_constants() {
1911 let lowered = lower_tosa(MATMUL_FP32.artifact, OPENVINO_TOSA_TARGET).unwrap();
1912 let xml = xml_str(&lowered);
1913 assert_eq!(xml.matches("type=\"Parameter\"").count(), 2);
1914 assert_eq!(xml.matches("type=\"Result\"").count(), 1);
1915 assert!(xml.contains("type=\"MatMul\""));
1916 assert!(xml.contains("transpose_a=\"false\" transpose_b=\"false\""));
1917 assert_eq!(xml.matches("type=\"Const\"").count(), 0);
1919 assert!(lowered.weights.is_empty());
1920 assert_eq!(lowered.features.len(), 3);
1921 assert_eq!(lowered.features[2].io_index, 0);
1922 }
1923
1924 #[test]
1925 fn lowers_the_shared_fp32_edge_identity_artifact() {
1926 let lowered = lower_tosa(IDENTITY_EDGES_FP32.artifact, OPENVINO_TOSA_TARGET).unwrap();
1927 assert_eq!(lowered.features.len(), 2);
1928 assert_eq!(lowered.features[0].byte_len, 32);
1929 }
1930
1931 #[test]
1932 fn lowers_nhwc_max_pool_through_explicit_layout_transposes() {
1933 let lowered = lower_tosa(MAX_POOL2D_FP32.artifact, OPENVINO_TOSA_TARGET).unwrap();
1934 let xml = xml_str(&lowered);
1935 assert_eq!(xml.matches("type=\"Transpose\"").count(), 2);
1936 assert_eq!(xml.matches("type=\"MaxPool\"").count(), 1);
1937 assert!(xml.contains(
1938 "strides=\"2,2\" kernel=\"2,2\" pads_begin=\"0,0\" pads_end=\"0,0\" \
1939 rounding_type=\"floor\" auto_pad=\"explicit\""
1940 ));
1941 assert_eq!(xml.matches("type=\"Const\"").count(), 2);
1943 assert!(xml.contains("offset=\"0\" size=\"32\""));
1944 assert!(xml.contains("offset=\"64\" size=\"32\""));
1945 assert_eq!(lowered.weights.len(), 96);
1946 let perms = |offset: usize| {
1947 (0..4)
1948 .map(|index| {
1949 i64::from_le_bytes(
1950 lowered.weights[offset + index * 8..offset + index * 8 + 8]
1951 .try_into()
1952 .unwrap(),
1953 )
1954 })
1955 .collect::<Vec<_>>()
1956 };
1957 assert_eq!(perms(0), [0, 3, 1, 2]);
1958 assert_eq!(perms(64), [0, 2, 3, 1]);
1959 }
1960
1961 #[test]
1962 fn lowers_every_shared_fp16_numerical_artifact() {
1963 for (name, artifact) in [
1964 ("identity", IDENTITY_EDGES_FP16.artifact),
1965 ("matmul", MATMUL_FP16.artifact),
1966 ("max_pool2d", MAX_POOL2D_FP16.artifact),
1967 ] {
1968 let lowered = lower_tosa(artifact, OPENVINO_TOSA_TARGET).unwrap_or_else(|error| {
1969 panic!("fp16 {name} failed to lower: {error}");
1970 });
1971 let xml = xml_str(&lowered);
1972 assert!(xml.contains("element_type=\"f16\""), "{name}");
1973 assert!(xml.contains("precision=\"FP16\""), "{name}");
1974 assert_eq!(lowered.features[0].element, OvElement::F16, "{name}");
1975 }
1976 }
1977
1978 #[test]
1979 fn fp16_constant_bytes_are_preserved_bit_exactly() {
1980 let mut builder = IrBuilder::new(0).unwrap();
1981 let bytes: [u8; 6] = [0x01, 0x7e, 0x00, 0x80, 0x01, 0x00];
1983 let port = builder
1984 .emit_const("payload", OvElement::F16, &[3], &bytes)
1985 .unwrap();
1986 assert_eq!(port, PortRef { layer: 0, port: 0 });
1987 assert_eq!(&builder.weights[..6], &bytes);
1988 builder
1990 .emit_const("aligned", OvElement::F32, &[1], &1.0f32.to_le_bytes())
1991 .unwrap();
1992 assert_eq!(builder.weights.len(), 68);
1993 assert_eq!(&builder.weights[64..68], &1.0f32.to_le_bytes());
1994 assert!(builder.layers.contains("offset=\"64\" size=\"4\""));
1995 }
1996
1997 #[test]
1998 fn scalar_constants_are_rank_zero() {
1999 let mut builder = IrBuilder::new(0).unwrap();
2000 builder.emit_i64_const("k", &[1], true).unwrap();
2001 assert!(builder.layers.contains("shape=\"\""));
2002 assert!(builder.layers.contains("size=\"8\""));
2003 assert!(!builder.layers.contains("<dim>"));
2004 assert_eq!(
2005 builder
2006 .emit_const("wrong", OvElement::F32, &[2], &[0u8; 4])
2007 .unwrap_err(),
2008 LoweringError::InvalidConstant
2009 );
2010 }
2011
2012 #[test]
2013 fn f16_to_f32_preserves_zero_finite_and_nan_classes() {
2014 assert_eq!(f16_to_f32(0x0000), 0.0);
2015 assert!(f16_to_f32(0x8000).is_sign_negative());
2016 assert_eq!(f16_to_f32(0x3c00), 1.0);
2017 assert_eq!(f16_to_f32(0xc000), -2.0);
2018 assert_eq!(f16_to_f32(0x7bff), 65504.0);
2019 assert_eq!(f16_to_f32(0x0001), 5.960_464_5e-8);
2020 assert_eq!(f16_to_f32(0x7c00), f32::INFINITY);
2021 assert_eq!(f16_to_f32(0xfc00), f32::NEG_INFINITY);
2022 assert!(f16_to_f32(0x7e01).is_nan());
2023 }
2024
2025 #[test]
2026 fn attribute_floats_reject_nan_and_render_infinities() {
2027 assert_eq!(attr_float(1.5).unwrap(), "1.5");
2028 assert_eq!(attr_float(f32::INFINITY).unwrap(), "inf");
2029 assert_eq!(attr_float(f32::NEG_INFINITY).unwrap(), "-inf");
2030 assert_eq!(
2031 attr_float(f32::NAN).unwrap_err(),
2032 LoweringError::UnsupportedGraph
2033 );
2034 }
2035
2036 #[test]
2037 fn generated_documents_never_contain_escapable_text() {
2038 for artifact in [
2039 IDENTITY_FP32_LOCAL,
2040 MATMUL_FP32.artifact,
2041 MAX_POOL2D_FP32.artifact,
2042 ] {
2043 let lowered = lower_tosa(artifact, OPENVINO_TOSA_TARGET).unwrap();
2044 let xml = xml_str(&lowered);
2045 assert!(!xml.contains('&'));
2046 assert!(!xml.contains('\''));
2047 assert_eq!(xml.matches('<').count(), xml.matches('>').count());
2048 }
2049 }
2050}