Skip to main content

virtio_accel_openvino/
lower.rs

1//! TOSA 1.0 to OpenVINO IR lowering.
2//!
3//! This module intentionally owns the OpenVINO IR (version 11) XML and weights encoding.
4//! Portable crates expose only the verified TOSA model and provider-neutral analysis; no
5//! OpenVINO type, header, or dependency crosses the backend boundary, so this encoder compiles
6//! and unit-tests on every platform.
7//!
8//! Emission order is load-bearing: `Parameter` layers are written first in input-slot order and
9//! `Result` layers last in output-slot order, because the IR frontend builds its parameter and
10//! result vectors in encounter order and those indices drive the runtime's
11//! `set_input/output_tensor_by_index` calls. Every layer, tensor, and attribute string written
12//! into the document is generated from fixed tables and integer IDs — TOSA-declared names never
13//! reach the XML, so no escaping surface exists.
14
15// Builds without a detected OpenVINO runtime type-check and unit-test this backend-local
16// encoder, but only the native runtime modules call it from `load_program`.
17#![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
29/// TOSA target currently lowered by the OpenVINO backend.
30pub const OPENVINO_TOSA_TARGET: Target = Target::new(
31    Version::TOSA_1_0,
32    ProfileSet::FLOATING_POINT,
33    Level::Level8K,
34    ExtensionSet::NONE,
35);
36
37/// TOSA integer-profile target lowered with exact INT8 storage and INT32 arithmetic.
38pub const OPENVINO_TOSA_INTEGER_TARGET: Target = Target::new(
39    Version::TOSA_1_0,
40    ProfileSet::INTEGER,
41    Level::Level8K,
42    ExtensionSet::NONE,
43);
44
45/// TOSA FP8 target (ADR 0009), mirroring `VULKAN_TOSA_FP8_TARGET` exactly.
46///
47/// TOSA gates FP8 legality on the `FP8E4M3` / `FP8E5M2` extensions rather than on the base
48/// floating-point profile, so the FP8 envelope is a distinct *target* rather than a dtype
49/// narrowing of [`OPENVINO_TOSA_TARGET`] -- unlike FP16, which shares the float target identity.
50pub 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
127/// `CAST` for the FP8 tier only.
128///
129/// It is deliberately absent from [`FLOAT_OPERATORS`]: that list is the operator surface shared
130/// with the Core ML provider, and `advertised_operator_and_dtype_surface_is_exact` in the Hexagon
131/// crate asserts the two agree operator for operator through `supports_tosa_operator`. Adding
132/// `CAST` there would have OpenVINO advertise an operator Core ML does not. The FP8 tier extends
133/// the shared list instead of mutating it.
134const CAST_CAPABILITY: OperatorCapability = OperatorCapability::new(Op::CAST);
135
136/// The FP8 tier's operators: every shared float operator, plus `CAST` as the narrowing path the
137/// Vulkan FP8 tier also carries. Built from [`FLOAT_OPERATORS`] rather than restated, so the two
138/// envelopes cannot drift.
139const 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
149/// The FP8 tier's dtypes: both encodings in every role, plus FP16 because TOSA's FP8 `MATMUL`
150/// accumulates into it and INT32 because `ARGMAX` indexes with it. Mirrors the Vulkan tier's
151/// `FLOAT8_DTYPES` so a graph admitted by one backend is admitted by the other.
152const 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
165/// Conservative floating-profile capability boundary for OpenVINO lowering.
166pub 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
178/// Conservative exact integer-profile capability boundary for OpenVINO lowering.
179pub 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
191/// FP8 capability boundary (ADR 0009): `(FP8, FP8) -> FP16` MATMUL and exact FP8 data movement.
192pub const OPENVINO_TOSA_FP8_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
193    target: OPENVINO_TOSA_FP8_TARGET,
194    dtypes: FLOAT8_DTYPES,
195    // The float tier's envelope plus CAST, narrowed by dtype rather than by a restated list --
196    // close to how the FP16 tier relates to FP32. TOSA's own per-operator dtype rules do the real
197    // constraining: it admits no FP8 elementwise operator, so an FP8-typed ADD is rejected by
198    // `analyze_for` even though ADD appears here. What that leaves reachable is FP8 MATMUL, FP8
199    // data movement, FP8 CAST, and FP16 arithmetic over FP8-sourced operands -- a superset of the
200    // Vulkan FP8 tier, so any graph Vulkan admits is admitted here too.
201    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
210/// The descriptor that governs `target`, or `None` if this backend does not lower it.
211///
212/// Admission must consult the descriptor for the target actually being lowered rather than a
213/// fixed one: the FP8 tier carries operators (`CAST`) the float tier does not, and rejects most
214/// of the operators the float tier admits, so one hardcoded descriptor cannot gate both.
215fn 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
227/// Weights-blob entries are aligned generously so every element type loads aligned.
228const 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/// Element kinds this lowering writes into IR documents.
250#[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    /// The `element_type` attribute spelling.
264    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    /// The port `precision` attribute spelling.
278    const fn precision(self) -> &'static str {
279        match self {
280            Self::F32 => "FP32",
281            Self::F16 => "FP16",
282            // Not "FP8E4M3": OpenVINO writes the FP8 port precisions without the FP prefix.
283            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    /// Storage bytes per scalar.
293    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
317/// A model-boundary element type; I64 stays internal to the graph.
318fn 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/// One model-boundary tensor and the binding slot that carries it.
338#[derive(Clone, Debug, PartialEq, Eq)]
339pub(crate) struct LoweredFeature {
340    pub slot: u32,
341    pub role: LoweredFeatureRole,
342    /// Index within the model's inputs or outputs, for `*_tensor_by_index` access.
343    pub io_index: u32,
344    pub element: OvElement,
345    pub dims: Vec<i64>,
346    /// Exact tensor bytes: the required length of a binding over this slot.
347    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
357/// A one-element FP8 `IDENTITY` document, for asking a device's compiler whether it accepts FP8
358/// at all.
359///
360/// Built through the same [`IrBuilder`] the real lowering uses, so it cannot drift from the
361/// document format that is known to work -- hand-written probe XML would be a second format to
362/// keep correct. FP8E4M3 stands in for both encodings: a compiler with no FP8 support at all
363/// rejects the type, and one that supports the extension supports both halves of it.
364pub(crate) fn fp8_probe_document() -> LoweredModel {
365    const ELEMENT: OvElement = OvElement::F8E4M3;
366    const DIMS: &[i64] = &[1];
367    // Three layers and two values: inside every bound `IrBuilder::new` checks.
368    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
396/// Whether the initial OpenVINO lowering tier can lower `op` for supported types and attributes.
397pub const fn supports_tosa_operator(op: Op) -> bool {
398    OPENVINO_TOSA_CAPABILITY.supports_operator(op)
399}
400
401/// Whether this lowering can expose `dtype` at an OpenVINO model boundary.
402///
403/// Operator-specific and target-specific validation still applies. INT8 is currently limited to
404/// exact identity and zero-point-aware MATMUL; other integer operators are rejected instead of
405/// silently dequantized.
406pub 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/// A produced tensor inside the document: one output port of one layer.
416#[derive(Clone, Copy, Debug, PartialEq, Eq)]
417struct PortRef {
418    layer: u32,
419    port: u32,
420}
421
422/// Incremental IR v11 document builder.
423struct IrBuilder {
424    layers: String,
425    edges: String,
426    weights: Vec<u8>,
427    next_layer: u32,
428    /// The producing port of each analyzed value, indexed by dense `ValueId`.
429    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    /// Emit one layer; input ports take `0..inputs.len()` and output ports follow.
457    ///
458    /// `data` is a preformatted attribute list (`key="value" ...`) built exclusively from fixed
459    /// tables and integer formatting.
460    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    /// Append raw constant bytes to the weights blob and emit its `Const` layer.
520    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    /// Emit an `i64` vector (or rank-0 scalar) constant used as an operator parameter input.
555    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
706/// Keep the artifact's declared TOSA profile load-bearing after individual operators gain support
707/// for more than one dtype. TOSA profile analysis validates operator legality, but a simple op such
708/// as IDENTITY can be legal for several profiles; the provider target still selects exactly one
709/// lowering tier and must not be used to smuggle a differently typed graph into it.
710fn 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            // The FP8 tier's own dtypes and no others. FP16 is legal here because TOSA's FP8
722            // `MATMUL` accumulates into it and INT32 because `ARGMAX` indexes with it; a graph
723            // carrying FP32 or INT8 belongs to another tier and must not be relabeled into this
724            // one. The float branch above cannot serve: it admits FP32 and FP16 alike.
725            !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
749/// Static, positive dimensions of a tensor; empty (rank-0) is allowed only inside the graph.
750fn 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
759/// Model-boundary tensors additionally reject rank-0.
760fn 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    // Compile-time constants consumed only by a layer parameter are deliberately absent from the
837    // OpenVINO graph. They have already been validated by TOSA analysis.
838    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    // (layer type, extra `<data .../>` attributes)
873    const NUMPY: &str = "auto_broadcast=\"numpy\"";
874    let (kind, data) = match op {
875        // A same-type Convert materializes the copy. Identity must not become a bare
876        // parameter-to-result edge: the runtime accepts that document but completes without
877        // writing a caller-bound output tensor (observed with the 2026.3 CPU plugin), which the
878        // output-pointer honesty check cannot detect.
879        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    // No plugin has an FP8 pooling primitive -- the NPU compiler's IE dialect admits FP8 on
1087    // MaxPool only as an optional scale, and the CPU plugin reports an empty primitive list --
1088    // so the window itself runs in binary16 while the transposes around it stay FP8. That keeps
1089    // the bytes that move narrow, and the round trip is exact for every finite value: each
1090    // widened tap is an exact FP8 value, the maximum is therefore one of them, and narrowing an
1091    // exactly-representable value returns it. A NaN that propagates through the window may come
1092    // back canonicalized, losing its sign.
1093    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
1169/// Lower TOSA INT8 MATMUL without relying on provider-specific implicit quantization.
1170///
1171/// TOSA requires `(a - a_zp) * (b - b_zp)` with exact INT32 accumulation. OpenVINO's ordinary
1172/// MatMul does not carry TOSA zero points, so both operands are widened and adjusted explicitly.
1173/// This preserves the integer-profile contract on every plugin; a plugin may fuse the pattern when
1174/// it can prove an equivalent optimized kernel.
1175fn 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
1247/// TOSA `ARGMAX` lowers to `TopK` (k = 1, stable lowest-index ties) plus a `Squeeze` that drops
1248/// the kept axis; only the TopK indices output is consumed and the values port dangles, which
1249/// the runtime accepts (pinned by a native unit test).
1250fn 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
1306/// `RECIPROCAL` and `RSQRT` lower through `Power` with a `-1` exponent.
1307///
1308/// `RSQRT` deliberately becomes `Sqrt` followed by `Power(x, -1)` rather than `Power(x, -0.5)`:
1309/// IEEE `pow(-0.0, -0.5)` is `+inf`, while TOSA `rsqrt(-0.0)` is `-inf`. `sqrt(-0.0) = -0.0`
1310/// and `pow(-0.0, -1) = -inf` preserve the signed-zero edge, and negative inputs still produce
1311/// NaN through `Sqrt`.
1312fn 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        // Normalize TOSA bool serialization to strict 0/1 storage bytes.
1373        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
1405/// The element type an FP8 operand of `op` must be widened to before emission, if any.
1406///
1407/// The NPU compiler's IE dialect declares `MatMul` operands as
1408/// `RankedTensorOf<[F16, F32, F64, SI32, quant_QuantizedType]>`, so MLIR's own verifier rejects a
1409/// MatMul over raw FP8 before any hardware question is asked. Widening is exact -- every FP8
1410/// value is representable in binary16 -- and costs nothing semantically, because TOSA's FP8
1411/// `MATMUL` already accumulates into FP16: the widened operands and the declared output type
1412/// agree, so nothing narrows back. Data movement is deliberately absent here; it must stay FP8
1413/// to remain bit-exact, since a widen/narrow round trip may canonicalize a NaN payload.
1414fn 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        // CAST is the one operator whose input and output types are *meant* to differ, so each
1456        // side is checked independently rather than against a shared predicate.
1457        Op::CAST => {
1458            for value in inputs.iter().chain(outputs) {
1459                require(*value, is_float_or_fp8)?;
1460            }
1461        }
1462        // FP8 operands are legal and widened at emission; the result is not FP8. TOSA's FP8
1463        // MATMUL accumulates into FP16, so requiring a float output states that rule locally.
1464        Op::MATMUL => {
1465            for value in inputs {
1466                require(*value, is_float_or_fp8)?;
1467            }
1468            require(outputs[0], is_float)?;
1469        }
1470        // Data movement and pooling carry FP8 through unchanged: these are the operators whose
1471        // FP8 form must stay bit-exact, so nothing widens and nothing rounds.
1472        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
1515/// Format a float attribute value; IR attribute text has no NaN spelling this crate relies on.
1516fn 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        // Identity materializes as a same-type Convert; a bare parameter-to-result edge would
1664        // complete without writing a caller-bound output tensor.
1665        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    /// FP8 MATMUL must reach the device as an FP16 MatMul fed by two explicit widening Converts.
1704    /// The NPU compiler's IE dialect declares MatMul operands without the FP8 types, so emitting
1705    /// a raw FP8 MatMul is rejected by MLIR's verifier -- this asserts the structure that avoids
1706    /// it, rather than numerics the plugin is free to accumulate its own way.
1707    #[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            // The FP8 boundary survives: the parameters are still FP8 at the model edge, which
1721            // is where the bandwidth saving lives.
1722            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    /// Pooling widens around the window only. No plugin has an FP8 pooling primitive, but the
1735    /// transposes either side stay FP8, so the bytes that move stay narrow.
1736    #[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    /// Data movement must not widen: a round trip through binary16 could canonicalize a NaN
1753    /// payload, and the tier's claim for movement is that it is bit-exact.
1754    #[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        // INT4 has no tier here and must stay out of the reported boundary. FP8 left this list
1772        // when the FP8 tier landed (ADR 0009); it is now reported like any other admitted dtype.
1773        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        // A graph-visible INT8 boundary remains an integer-tier program.
1881        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        // The zero-point operands are parameter-only constants and never become Const layers.
1918        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        // Two i64 permutation constants, each 64-byte aligned in the weights blob.
1942        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        // A quiet NaN with a payload, negative zero, and a subnormal: raw storage must survive.
1982        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        // A second constant lands at the next 64-byte boundary.
1989        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}