Skip to main content

virtio_accel_hexagon/
lower.rs

1//! Strict TOSA 1.0 admission and provider-local QNN graph planning.
2//!
3//! This module deliberately contains no QNN ABI types. It turns a verified TOSA artifact into
4//! owned tensor, binding, and operation metadata that the native boundary can translate while the
5//! source FlatBuffer is no longer borrowed. It compiles and tests on hosts without QAIRT.
6
7#![cfg_attr(not(va_hexagon), allow(dead_code))]
8
9use std::fmt;
10
11use virtio_accel_tosa::{
12    AnalysisError, AnalyzedValueKind, CapabilityDescriptor, DType, DTypeCapability,
13    Error as ParseError, ExtensionSet, GraphCapabilities, Level, NanPropagationMode, Op,
14    OpAttributes, OperatorCapability, OperatorConstraints, ProfileSet, RuntimeCondition,
15    RuntimeConditionSupport, Target, TosaAnalysis, ValueId, ValueRoles, Version, parse,
16};
17
18/// TOSA declaration accepted by the first Hexagon tier.
19///
20/// The target selects the TOSA floating-point profile. Backend admission narrows its tensor types
21/// to FP16 until the selected HTP runtime can prove FP32 computation without precision reduction.
22pub const HEXAGON_TOSA_TARGET: Target = Target::new(
23    Version::TOSA_1_0,
24    ProfileSet::FLOATING_POINT,
25    Level::Level8K,
26    ExtensionSet::NONE,
27);
28
29/// TOSA integer-profile target lowered with exact INT8 storage and INT32 accumulation.
30pub const HEXAGON_TOSA_INTEGER_TARGET: Target = Target::new(
31    Version::TOSA_1_0,
32    ProfileSet::INTEGER,
33    Level::Level8K,
34    ExtensionSet::NONE,
35);
36
37const FLOAT_DTYPES: &[DTypeCapability] = &[
38    DTypeCapability::new(DType::FP16, ValueRoles::ALL),
39    DTypeCapability::new(DType::BOOL, ValueRoles::ALL),
40    DTypeCapability::new(
41        DType::INT32,
42        ValueRoles::OUTPUT
43            .union(ValueRoles::CONSTANT)
44            .union(ValueRoles::INTERMEDIATE),
45    ),
46];
47
48const INTEGER_DTYPES: &[DTypeCapability] = &[
49    DTypeCapability::new(DType::INT8, ValueRoles::ALL),
50    DTypeCapability::new(
51        DType::INT32,
52        ValueRoles::OUTPUT
53            .union(ValueRoles::CONSTANT)
54            .union(ValueRoles::INTERMEDIATE),
55    ),
56];
57
58const FLOAT_OPERATORS: &[OperatorCapability] = &[
59    OperatorCapability::constrained(Op::ARGMAX, OperatorConstraints::PROPAGATING_NAN),
60    OperatorCapability::constrained(Op::MATMUL, OperatorConstraints::ZERO_ZERO_POINTS),
61    OperatorCapability::constrained(
62        Op::MAX_POOL2D,
63        OperatorConstraints::PROPAGATING_NAN.union(OperatorConstraints::ZERO_PADDING),
64    ),
65    OperatorCapability::constrained(Op::CLAMP, OperatorConstraints::PROPAGATING_NAN),
66    OperatorCapability::new(Op::SIGMOID),
67    OperatorCapability::new(Op::TANH),
68    OperatorCapability::new(Op::ADD),
69    OperatorCapability::new(Op::SUB),
70    OperatorCapability::constrained(Op::MUL, OperatorConstraints::ZERO_SHIFT),
71    OperatorCapability::new(Op::POW),
72    OperatorCapability::constrained(Op::MAXIMUM, OperatorConstraints::PROPAGATING_NAN),
73    OperatorCapability::constrained(Op::MINIMUM, OperatorConstraints::PROPAGATING_NAN),
74    OperatorCapability::new(Op::LOGICAL_AND),
75    OperatorCapability::new(Op::LOGICAL_OR),
76    OperatorCapability::new(Op::LOGICAL_XOR),
77    OperatorCapability::new(Op::ABS),
78    OperatorCapability::new(Op::CEIL),
79    OperatorCapability::new(Op::COS),
80    OperatorCapability::new(Op::EXP),
81    OperatorCapability::new(Op::FLOOR),
82    OperatorCapability::new(Op::LOG),
83    OperatorCapability::new(Op::LOGICAL_NOT),
84    OperatorCapability::constrained(Op::NEGATE, OperatorConstraints::ZERO_ZERO_POINTS),
85    OperatorCapability::new(Op::RECIPROCAL),
86    OperatorCapability::new(Op::RSQRT),
87    OperatorCapability::new(Op::SIN),
88    OperatorCapability::new(Op::SELECT),
89    OperatorCapability::new(Op::EQUAL),
90    OperatorCapability::new(Op::GREATER),
91    OperatorCapability::new(Op::GREATER_EQUAL),
92    OperatorCapability::constrained(Op::REDUCE_MAX, OperatorConstraints::PROPAGATING_NAN),
93    OperatorCapability::constrained(Op::REDUCE_MIN, OperatorConstraints::PROPAGATING_NAN),
94    OperatorCapability::new(Op::REDUCE_PRODUCT),
95    OperatorCapability::new(Op::REDUCE_SUM),
96    OperatorCapability::new(Op::CONCAT),
97    OperatorCapability::constrained(Op::RESHAPE, OperatorConstraints::CONSTANT_PARAMETERS),
98    OperatorCapability::new(Op::REVERSE),
99    OperatorCapability::new(Op::TRANSPOSE),
100    OperatorCapability::new(Op::CONST),
101    OperatorCapability::new(Op::IDENTITY),
102    OperatorCapability::new(Op::CONST_SHAPE),
103];
104
105const INTEGER_OPERATORS: &[OperatorCapability] = &[
106    OperatorCapability::new(Op::CONST),
107    OperatorCapability::new(Op::IDENTITY),
108    OperatorCapability::new(Op::MATMUL),
109];
110
111/// Conservative floating-profile capability boundary for the validated QNN HTP tier.
112pub const HEXAGON_TOSA_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
113    target: HEXAGON_TOSA_TARGET,
114    dtypes: FLOAT_DTYPES,
115    operators: FLOAT_OPERATORS,
116    graph: GraphCapabilities {
117        max_regions: 1,
118        max_blocks: 1,
119        dynamic_shapes: false,
120        runtime_conditions: RuntimeConditionSupport::AdvisoryOnly,
121    },
122};
123
124/// Conservative exact integer-profile capability boundary for the validated QNN HTP tier.
125pub const HEXAGON_TOSA_INTEGER_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
126    target: HEXAGON_TOSA_INTEGER_TARGET,
127    dtypes: INTEGER_DTYPES,
128    operators: INTEGER_OPERATORS,
129    graph: GraphCapabilities {
130        max_regions: 1,
131        max_blocks: 1,
132        dynamic_shapes: false,
133        runtime_conditions: RuntimeConditionSupport::None,
134    },
135};
136
137/// Failure while validating and planning a graph for QNN HTP.
138#[derive(Clone, Copy, Debug, PartialEq, Eq)]
139pub enum LoweringError {
140    Parse(ParseError),
141    Analysis(AnalysisError),
142    UnsupportedGraph,
143    UnsupportedType(DType),
144    UnsupportedOperator(Op),
145    InvalidConstant,
146    ResourceLimit,
147}
148
149impl fmt::Display for LoweringError {
150    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
151        write!(formatter, "{self:?}")
152    }
153}
154
155impl std::error::Error for LoweringError {}
156
157/// Whether the first hardware tier has a QNN lowering for `op`.
158pub const fn supports_tosa_operator(op: Op) -> bool {
159    HEXAGON_TOSA_CAPABILITY.supports_operator(op)
160}
161
162fn supports_operator_for_target(op: Op, integer: bool) -> bool {
163    if integer {
164        HEXAGON_TOSA_INTEGER_CAPABILITY.supports_operator(op)
165    } else {
166        supports_tosa_operator(op)
167    }
168}
169
170/// Whether the first hardware tier may expose `dtype` at a model boundary.
171///
172/// FP32 remains deliberately rejected because current HTP floating-point execution may use FP16
173/// math. Integer and packed low-precision tiers require separate targets and evidence.
174pub const fn supports_tosa_dtype(dtype: DType) -> bool {
175    HEXAGON_TOSA_CAPABILITY.supports_dtype(dtype, ValueRoles::INPUT)
176        || HEXAGON_TOSA_CAPABILITY.supports_dtype(dtype, ValueRoles::OUTPUT)
177        || HEXAGON_TOSA_INTEGER_CAPABILITY.supports_dtype(dtype, ValueRoles::INPUT)
178        || HEXAGON_TOSA_INTEGER_CAPABILITY.supports_dtype(dtype, ValueRoles::OUTPUT)
179}
180
181#[derive(Clone, Copy, Debug, PartialEq, Eq)]
182pub(crate) enum Element {
183    Bool,
184    F16,
185    F32,
186    I8,
187    I32,
188}
189
190impl Element {
191    pub(crate) const fn scalar_bytes(self) -> u64 {
192        match self {
193            Self::Bool | Self::I8 => 1,
194            Self::F16 => 2,
195            Self::F32 | Self::I32 => 4,
196        }
197    }
198
199    fn for_dtype(dtype: DType) -> Result<Self, LoweringError> {
200        match dtype {
201            DType::BOOL => Ok(Self::Bool),
202            DType::FP16 => Ok(Self::F16),
203            DType::FP32 => Ok(Self::F32),
204            DType::INT8 => Ok(Self::I8),
205            DType::INT32 => Ok(Self::I32),
206            _ => Err(LoweringError::UnsupportedType(dtype)),
207        }
208    }
209}
210
211#[derive(Clone, Copy, Debug, PartialEq)]
212pub(crate) struct Quantization {
213    pub scale: f32,
214    pub offset: i32,
215}
216
217#[derive(Clone, Copy, Debug, PartialEq, Eq)]
218pub(crate) enum FeatureRole {
219    Input,
220    Output,
221}
222
223/// One exact model-boundary binding.
224#[derive(Clone, Debug, PartialEq, Eq)]
225pub(crate) struct LoweredFeature {
226    pub slot: u32,
227    pub role: FeatureRole,
228    pub io_index: u32,
229    pub value: u32,
230    pub dims: Vec<u32>,
231    pub byte_len: u64,
232}
233
234/// One owned tensor descriptor used while constructing the QNN graph.
235#[derive(Clone, Debug, PartialEq)]
236pub(crate) struct LoweredTensor {
237    pub value: u32,
238    pub element: Element,
239    pub quantization: Option<Quantization>,
240    pub dims: Vec<u32>,
241    pub data: Option<Vec<u8>>,
242}
243
244/// QNN operation selected by portable TOSA lowering.
245#[derive(Clone, Copy, Debug, PartialEq, Eq)]
246pub(crate) enum NodeKind {
247    Identity,
248    Transpose,
249    Reverse,
250    Concat,
251    MatMul,
252    MaxPool2d,
253    Add,
254    Subtract,
255    Multiply,
256    Maximum,
257    Minimum,
258    Power,
259    Abs,
260    Ceil,
261    Cos,
262    Exp,
263    Floor,
264    Log,
265    Negate,
266    Reciprocal,
267    Rsqrt,
268    Sin,
269    Sigmoid,
270    Tanh,
271    Clamp,
272    Equal,
273    Greater,
274    GreaterEqual,
275    Select,
276    LogicalAnd,
277    LogicalOr,
278    LogicalXor,
279    LogicalNot,
280    ArgMax,
281    ReduceMax,
282    ReduceMin,
283    ReduceProduct,
284    ReduceSum,
285}
286
287/// One owned operation descriptor. Parameter meaning is fixed by `kind` and validated natively.
288#[derive(Clone, Debug, PartialEq, Eq)]
289pub(crate) struct LoweredNode {
290    pub kind: NodeKind,
291    pub inputs: Vec<u32>,
292    pub outputs: Vec<u32>,
293    pub parameters: Vec<i32>,
294}
295
296/// Fully owned graph plan produced before entering the native QNN boundary.
297#[derive(Clone, Debug, PartialEq)]
298pub(crate) struct LoweredModel {
299    pub tensors: Vec<LoweredTensor>,
300    pub nodes: Vec<LoweredNode>,
301    pub features: Vec<LoweredFeature>,
302    pub precision: Option<Element>,
303}
304
305impl LoweredModel {
306    pub(crate) fn boundary(&self, value: u32) -> Option<(FeatureRole, u32)> {
307        self.features
308            .iter()
309            .find(|feature| feature.value == value)
310            .map(|feature| (feature.role, feature.io_index))
311    }
312}
313
314pub(crate) fn lower_tosa(bytes: &[u8], target: Target) -> Result<LoweredModel, LoweringError> {
315    let integer = if target == HEXAGON_TOSA_TARGET {
316        false
317    } else if target == HEXAGON_TOSA_INTEGER_TARGET {
318        true
319    } else {
320        return Err(LoweringError::UnsupportedGraph);
321    };
322    let model = parse(bytes).map_err(LoweringError::Parse)?;
323    let analysis = model.analyze_for(target).map_err(LoweringError::Analysis)?;
324    if analysis.regions().len() != 1
325        || analysis.blocks().len() != 1
326        || analysis
327            .conditions()
328            .iter()
329            .any(|condition| !matches!(condition, RuntimeCondition::PowDomain { .. }))
330    {
331        return Err(LoweringError::UnsupportedGraph);
332    }
333
334    validate_types(&analysis, integer)?;
335    let block = analysis.blocks()[0].id();
336    let inputs = analysis.block_inputs(block);
337    let outputs = analysis.block_outputs(block);
338    if inputs.is_empty()
339        || outputs.is_empty()
340        || inputs.iter().any(|input| outputs.contains(input))
341        || inputs
342            .iter()
343            .chain(outputs)
344            .any(|value| !matches!(analysis.value(*value).kind(), AnalyzedValueKind::Tensor(_)))
345        || outputs
346            .iter()
347            .any(|value| analysis.serialized_constant(*value).is_some())
348        || inputs.len().checked_add(outputs.len()).is_none()
349    {
350        return Err(LoweringError::UnsupportedGraph);
351    }
352
353    let mut tensors = Vec::new();
354    tensors
355        .try_reserve_exact(analysis.values().len())
356        .map_err(|_| LoweringError::ResourceLimit)?;
357    for value in analysis.values() {
358        let AnalyzedValueKind::Tensor(tensor) = value.kind() else {
359            if analysis.serialized_constant(value.id()).is_none() {
360                return Err(LoweringError::UnsupportedGraph);
361            }
362            continue;
363        };
364        let element = Element::for_dtype(tensor.dtype())?;
365        tensors.push(LoweredTensor {
366            value: value.id().get(),
367            element,
368            quantization: (integer && matches!(element, Element::I8 | Element::I32)).then_some(
369                Quantization {
370                    scale: 1.0,
371                    offset: 0,
372                },
373            ),
374            dims: static_dims(tensor, true)?,
375            data: analysis.serialized_constant(value.id()).map(<[u8]>::to_vec),
376        });
377    }
378
379    let mut features = Vec::new();
380    features
381        .try_reserve_exact(inputs.len() + outputs.len())
382        .map_err(|_| LoweringError::ResourceLimit)?;
383    for (index, value) in inputs.iter().copied().enumerate() {
384        features.push(lower_feature(
385            &analysis,
386            value,
387            index,
388            index,
389            FeatureRole::Input,
390        )?);
391    }
392    for (index, value) in outputs.iter().copied().enumerate() {
393        features.push(lower_feature(
394            &analysis,
395            value,
396            inputs.len() + index,
397            index,
398            FeatureRole::Output,
399        )?);
400    }
401
402    let mut nodes = Vec::new();
403    let mut quantization_offsets = Vec::new();
404    nodes
405        .try_reserve_exact(analysis.execution_order(block).len())
406        .map_err(|_| LoweringError::ResourceLimit)?;
407    for operator_id in analysis.execution_order(block) {
408        let operator = analysis.operator(*operator_id);
409        let op = operator.op();
410        if !supports_operator_for_target(op, integer) {
411            return Err(LoweringError::UnsupportedOperator(op));
412        }
413        let op_inputs = analysis.operator_inputs(*operator_id);
414        let op_outputs = analysis.operator_outputs(*operator_id);
415        match op {
416            Op::CONST => {
417                if op_inputs.is_empty()
418                    && op_outputs.len() == 1
419                    && analysis.serialized_constant(op_outputs[0]).is_some()
420                {
421                    continue;
422                }
423                return Err(LoweringError::InvalidConstant);
424            }
425            Op::CONST_SHAPE => {
426                if op_inputs.is_empty()
427                    && op_outputs.len() == 1
428                    && analysis.serialized_constant(op_outputs[0]).is_some()
429                {
430                    continue;
431                }
432                return Err(LoweringError::InvalidConstant);
433            }
434            Op::IDENTITY => {
435                require_arity(op_inputs, 1, op_outputs, 1)?;
436                nodes.push(lowered_node(
437                    NodeKind::Identity,
438                    op_inputs,
439                    op_outputs,
440                    Vec::new(),
441                ));
442            }
443            Op::RESHAPE => {
444                require_arity(op_inputs, 2, op_outputs, 1)?;
445                analysis
446                    .serialized_constant(op_inputs[1])
447                    .ok_or(LoweringError::InvalidConstant)?;
448                nodes.push(lowered_node(
449                    NodeKind::Identity,
450                    &op_inputs[..1],
451                    op_outputs,
452                    Vec::new(),
453                ));
454            }
455            Op::TRANSPOSE => {
456                require_arity(op_inputs, 1, op_outputs, 1)?;
457                let OpAttributes::Transpose { perms } = operator.source().attributes() else {
458                    return Err(LoweringError::UnsupportedGraph);
459                };
460                let parameters = perms.iter().collect::<Vec<_>>();
461                if parameters.is_empty() {
462                    return Err(LoweringError::UnsupportedGraph);
463                }
464                nodes.push(lowered_node(
465                    NodeKind::Transpose,
466                    op_inputs,
467                    op_outputs,
468                    parameters,
469                ));
470            }
471            Op::REVERSE => {
472                require_arity(op_inputs, 1, op_outputs, 1)?;
473                let OpAttributes::Reverse { axis } = operator.source().attributes() else {
474                    return Err(LoweringError::UnsupportedGraph);
475                };
476                nodes.push(lowered_node(
477                    NodeKind::Reverse,
478                    op_inputs,
479                    op_outputs,
480                    vec![axis],
481                ));
482            }
483            Op::CONCAT => {
484                if op_inputs.is_empty() || op_outputs.len() != 1 {
485                    return Err(LoweringError::UnsupportedGraph);
486                }
487                let OpAttributes::Concat { axis } = operator.source().attributes() else {
488                    return Err(LoweringError::UnsupportedGraph);
489                };
490                nodes.push(lowered_node(
491                    NodeKind::Concat,
492                    op_inputs,
493                    op_outputs,
494                    vec![axis],
495                ));
496            }
497            Op::MATMUL => {
498                require_arity(op_inputs, 4, op_outputs, 1)?;
499                let left_zero_point = scalar_zero_point(&analysis, op_inputs[2])?;
500                let right_zero_point = scalar_zero_point(&analysis, op_inputs[3])?;
501                if integer {
502                    set_quantization_offset(
503                        &mut tensors,
504                        &mut quantization_offsets,
505                        op_inputs[0],
506                        left_zero_point,
507                    )?;
508                    set_quantization_offset(
509                        &mut tensors,
510                        &mut quantization_offsets,
511                        op_inputs[1],
512                        right_zero_point,
513                    )?;
514                } else if left_zero_point != 0 || right_zero_point != 0 {
515                    return Err(LoweringError::UnsupportedGraph);
516                }
517                nodes.push(LoweredNode {
518                    kind: NodeKind::MatMul,
519                    inputs: vec![op_inputs[0].get(), op_inputs[1].get()],
520                    outputs: vec![op_outputs[0].get()],
521                    parameters: Vec::new(),
522                });
523            }
524            Op::MAX_POOL2D => {
525                require_arity(op_inputs, 1, op_outputs, 1)?;
526                let OpAttributes::MaxPool2d {
527                    kernel,
528                    stride,
529                    pad,
530                    nan_mode,
531                } = operator.source().attributes()
532                else {
533                    return Err(LoweringError::UnsupportedGraph);
534                };
535                if nan_mode != NanPropagationMode::PROPAGATE {
536                    return Err(LoweringError::UnsupportedGraph);
537                }
538                let kernel = fixed_positive_pair(kernel.iter())?;
539                let stride = fixed_positive_pair(stride.iter())?;
540                let pad = pad.iter().collect::<Vec<_>>();
541                if pad.len() != 4 || pad.iter().any(|value| *value != 0) {
542                    return Err(LoweringError::UnsupportedGraph);
543                }
544                let input_dims = tensor_dims(&analysis, op_inputs[0], false)?;
545                let output_dims = tensor_dims(&analysis, op_outputs[0], false)?;
546                if input_dims.len() != 4 || output_dims.len() != 4 {
547                    return Err(LoweringError::UnsupportedGraph);
548                }
549                nodes.push(LoweredNode {
550                    kind: NodeKind::MaxPool2d,
551                    inputs: vec![op_inputs[0].get()],
552                    outputs: vec![op_outputs[0].get()],
553                    parameters: vec![
554                        kernel[0] as i32,
555                        kernel[1] as i32,
556                        stride[0] as i32,
557                        stride[1] as i32,
558                    ],
559                });
560            }
561            Op::ADD | Op::SUB | Op::POW | Op::MAXIMUM | Op::MINIMUM => {
562                require_arity(op_inputs, 2, op_outputs, 1)?;
563                match operator.source().attributes() {
564                    OpAttributes::Maximum { nan_mode } | OpAttributes::Minimum { nan_mode }
565                        if nan_mode != NanPropagationMode::PROPAGATE =>
566                    {
567                        return Err(LoweringError::UnsupportedGraph);
568                    }
569                    _ => {}
570                }
571                let kind = match op {
572                    Op::ADD => NodeKind::Add,
573                    Op::SUB => NodeKind::Subtract,
574                    Op::POW => NodeKind::Power,
575                    Op::MAXIMUM => NodeKind::Maximum,
576                    Op::MINIMUM => NodeKind::Minimum,
577                    _ => unreachable!(),
578                };
579                nodes.push(LoweredNode {
580                    kind,
581                    inputs: op_inputs.iter().map(|value| value.get()).collect(),
582                    outputs: vec![op_outputs[0].get()],
583                    parameters: Vec::new(),
584                });
585            }
586            Op::MUL => {
587                require_arity(op_inputs, 3, op_outputs, 1)?;
588                let shift = analysis
589                    .serialized_constant(op_inputs[2])
590                    .ok_or(LoweringError::InvalidConstant)?;
591                if shift.iter().any(|byte| *byte != 0) {
592                    return Err(LoweringError::UnsupportedGraph);
593                }
594                nodes.push(LoweredNode {
595                    kind: NodeKind::Multiply,
596                    inputs: op_inputs[..2].iter().map(|value| value.get()).collect(),
597                    outputs: vec![op_outputs[0].get()],
598                    parameters: Vec::new(),
599                });
600            }
601            Op::ABS
602            | Op::CEIL
603            | Op::COS
604            | Op::EXP
605            | Op::FLOOR
606            | Op::LOG
607            | Op::RECIPROCAL
608            | Op::RSQRT
609            | Op::SIN
610            | Op::SIGMOID
611            | Op::TANH
612            | Op::LOGICAL_NOT => {
613                require_arity(op_inputs, 1, op_outputs, 1)?;
614                let kind = match op {
615                    Op::ABS => NodeKind::Abs,
616                    Op::CEIL => NodeKind::Ceil,
617                    Op::COS => NodeKind::Cos,
618                    Op::EXP => NodeKind::Exp,
619                    Op::FLOOR => NodeKind::Floor,
620                    Op::LOG => NodeKind::Log,
621                    Op::RECIPROCAL => NodeKind::Reciprocal,
622                    Op::RSQRT => NodeKind::Rsqrt,
623                    Op::SIN => NodeKind::Sin,
624                    Op::SIGMOID => NodeKind::Sigmoid,
625                    Op::TANH => NodeKind::Tanh,
626                    Op::LOGICAL_NOT => NodeKind::LogicalNot,
627                    _ => unreachable!(),
628                };
629                nodes.push(lowered_node(kind, op_inputs, op_outputs, Vec::new()));
630            }
631            Op::NEGATE => {
632                require_arity(op_inputs, 3, op_outputs, 1)?;
633                if scalar_zero_point(&analysis, op_inputs[1])? != 0
634                    || scalar_zero_point(&analysis, op_inputs[2])? != 0
635                {
636                    return Err(LoweringError::UnsupportedGraph);
637                }
638                nodes.push(lowered_node(
639                    NodeKind::Negate,
640                    &op_inputs[..1],
641                    op_outputs,
642                    Vec::new(),
643                ));
644            }
645            Op::CLAMP => {
646                require_arity(op_inputs, 1, op_outputs, 1)?;
647                let OpAttributes::Clamp {
648                    min_val,
649                    max_val,
650                    nan_mode,
651                } = operator.source().attributes()
652                else {
653                    return Err(LoweringError::UnsupportedGraph);
654                };
655                if nan_mode != NanPropagationMode::PROPAGATE {
656                    return Err(LoweringError::UnsupportedGraph);
657                }
658                let dtype = tensor(&analysis, op_inputs[0])?.dtype();
659                let minimum = decode_float(dtype, min_val)?;
660                let maximum = decode_float(dtype, max_val)?;
661                nodes.push(lowered_node(
662                    NodeKind::Clamp,
663                    op_inputs,
664                    op_outputs,
665                    vec![minimum.to_bits() as i32, maximum.to_bits() as i32],
666                ));
667            }
668            Op::EQUAL
669            | Op::GREATER
670            | Op::GREATER_EQUAL
671            | Op::LOGICAL_AND
672            | Op::LOGICAL_OR
673            | Op::LOGICAL_XOR => {
674                require_arity(op_inputs, 2, op_outputs, 1)?;
675                let kind = match op {
676                    Op::EQUAL => NodeKind::Equal,
677                    Op::GREATER => NodeKind::Greater,
678                    Op::GREATER_EQUAL => NodeKind::GreaterEqual,
679                    Op::LOGICAL_AND => NodeKind::LogicalAnd,
680                    Op::LOGICAL_OR => NodeKind::LogicalOr,
681                    Op::LOGICAL_XOR => NodeKind::LogicalXor,
682                    _ => unreachable!(),
683                };
684                nodes.push(lowered_node(kind, op_inputs, op_outputs, Vec::new()));
685            }
686            Op::SELECT => {
687                require_arity(op_inputs, 3, op_outputs, 1)?;
688                nodes.push(lowered_node(
689                    NodeKind::Select,
690                    op_inputs,
691                    op_outputs,
692                    Vec::new(),
693                ));
694            }
695            Op::ARGMAX => {
696                require_arity(op_inputs, 1, op_outputs, 1)?;
697                let OpAttributes::ArgMax { axis, nan_mode } = operator.source().attributes() else {
698                    return Err(LoweringError::UnsupportedGraph);
699                };
700                if nan_mode != NanPropagationMode::PROPAGATE {
701                    return Err(LoweringError::UnsupportedGraph);
702                }
703                nodes.push(lowered_node(
704                    NodeKind::ArgMax,
705                    op_inputs,
706                    op_outputs,
707                    vec![axis],
708                ));
709            }
710            Op::REDUCE_MAX | Op::REDUCE_MIN | Op::REDUCE_PRODUCT | Op::REDUCE_SUM => {
711                require_arity(op_inputs, 1, op_outputs, 1)?;
712                let (kind, axis) = match operator.source().attributes() {
713                    OpAttributes::ReduceMax { axis, nan_mode } => {
714                        if nan_mode != NanPropagationMode::PROPAGATE {
715                            return Err(LoweringError::UnsupportedGraph);
716                        }
717                        (NodeKind::ReduceMax, axis)
718                    }
719                    OpAttributes::ReduceMin { axis, nan_mode } => {
720                        if nan_mode != NanPropagationMode::PROPAGATE {
721                            return Err(LoweringError::UnsupportedGraph);
722                        }
723                        (NodeKind::ReduceMin, axis)
724                    }
725                    OpAttributes::ReduceProduct { axis } => (NodeKind::ReduceProduct, axis),
726                    OpAttributes::ReduceSum { axis } => (NodeKind::ReduceSum, axis),
727                    _ => return Err(LoweringError::UnsupportedGraph),
728                };
729                nodes.push(lowered_node(kind, op_inputs, op_outputs, vec![axis]));
730            }
731            _ => return Err(LoweringError::UnsupportedOperator(op)),
732        }
733    }
734
735    if nodes.is_empty() {
736        return Err(LoweringError::UnsupportedGraph);
737    }
738    Ok(LoweredModel {
739        tensors,
740        nodes,
741        features,
742        precision: if integer {
743            None
744        } else if analysis.values().iter().any(|value| {
745            matches!(value.kind(), AnalyzedValueKind::Tensor(tensor) if tensor.dtype() == DType::FP32)
746        }) {
747            Some(Element::F32)
748        } else {
749            Some(Element::F16)
750        },
751    })
752}
753
754fn validate_types(analysis: &TosaAnalysis<'_>, integer: bool) -> Result<(), LoweringError> {
755    for value in analysis.values() {
756        let AnalyzedValueKind::Tensor(tensor) = value.kind() else {
757            if analysis.serialized_constant(value.id()).is_none() {
758                return Err(LoweringError::UnsupportedGraph);
759            }
760            continue;
761        };
762        let supported = if integer {
763            matches!(tensor.dtype(), DType::INT8 | DType::INT32)
764        } else {
765            matches!(tensor.dtype(), DType::BOOL | DType::FP16 | DType::INT32)
766                || (tensor.dtype() == DType::INT8
767                    && constant_is_parameter_only(analysis, value.id()))
768        };
769        if !supported {
770            return Err(LoweringError::UnsupportedType(tensor.dtype()));
771        }
772    }
773    Ok(())
774}
775
776fn lower_feature(
777    analysis: &TosaAnalysis<'_>,
778    value: ValueId,
779    slot: usize,
780    io_index: usize,
781    role: FeatureRole,
782) -> Result<LoweredFeature, LoweringError> {
783    let dims = tensor_dims(analysis, value, false)?;
784    let element = Element::for_dtype(tensor(analysis, value)?.dtype())?;
785    let byte_len = checked_tensor_byte_len(element, &dims)?;
786    Ok(LoweredFeature {
787        slot: u32::try_from(slot).map_err(|_| LoweringError::ResourceLimit)?,
788        role,
789        io_index: u32::try_from(io_index).map_err(|_| LoweringError::ResourceLimit)?,
790        value: value.get(),
791        dims,
792        byte_len,
793    })
794}
795
796fn checked_tensor_byte_len(element: Element, dims: &[u32]) -> Result<u64, LoweringError> {
797    dims.iter().try_fold(element.scalar_bytes(), |bytes, dim| {
798        bytes
799            .checked_mul(u64::from(*dim))
800            .ok_or(LoweringError::ResourceLimit)
801    })
802}
803
804fn tensor_dims(
805    analysis: &TosaAnalysis<'_>,
806    value: ValueId,
807    allow_scalar: bool,
808) -> Result<Vec<u32>, LoweringError> {
809    let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
810        return Err(LoweringError::UnsupportedGraph);
811    };
812    static_dims(tensor, allow_scalar)
813}
814
815fn static_dims(
816    tensor: virtio_accel_tosa::Tensor<'_>,
817    allow_scalar: bool,
818) -> Result<Vec<u32>, LoweringError> {
819    tensor.rank().ok_or(LoweringError::UnsupportedGraph)?;
820    let dims = tensor
821        .dimensions()
822        .map(|dim| u32::try_from(dim).map_err(|_| LoweringError::UnsupportedGraph))
823        .collect::<Result<Vec<_>, _>>()?;
824    if (!allow_scalar && dims.is_empty()) || dims.contains(&0) {
825        return Err(LoweringError::UnsupportedGraph);
826    }
827    Ok(dims)
828}
829
830fn lowered_node(
831    kind: NodeKind,
832    inputs: &[ValueId],
833    outputs: &[ValueId],
834    parameters: Vec<i32>,
835) -> LoweredNode {
836    LoweredNode {
837        kind,
838        inputs: inputs.iter().map(|value| value.get()).collect(),
839        outputs: outputs.iter().map(|value| value.get()).collect(),
840        parameters,
841    }
842}
843
844fn decode_float(dtype: DType, bytes: &[u8]) -> Result<f32, LoweringError> {
845    match dtype {
846        DType::FP16 if bytes.len() == 2 => Ok(f16_to_f32(u16::from_le_bytes(
847            bytes.try_into().expect("length checked"),
848        ))),
849        DType::FP32 if bytes.len() == 4 => Ok(f32::from_le_bytes(
850            bytes.try_into().expect("length checked"),
851        )),
852        _ => Err(LoweringError::InvalidConstant),
853    }
854}
855
856fn f16_to_f32(bits: u16) -> f32 {
857    let sign = u32::from(bits & 0x8000) << 16;
858    let exponent = (bits >> 10) & 0x1f;
859    let fraction = u32::from(bits & 0x03ff);
860    let output = match exponent {
861        0 if fraction == 0 => sign,
862        0 => {
863            let leading = 31 - fraction.leading_zeros();
864            let normalized_fraction = (fraction << (10 - leading)) & 0x03ff;
865            let exponent32 = 127 - 14 - (10 - leading);
866            sign | (exponent32 << 23) | (normalized_fraction << 13)
867        }
868        0x1f => sign | 0x7f80_0000 | (fraction << 13),
869        _ => sign | (u32::from(exponent + 112) << 23) | (fraction << 13),
870    };
871    f32::from_bits(output)
872}
873
874fn require_arity(
875    inputs: &[ValueId],
876    input_count: usize,
877    outputs: &[ValueId],
878    output_count: usize,
879) -> Result<(), LoweringError> {
880    if inputs.len() == input_count && outputs.len() == output_count {
881        Ok(())
882    } else {
883        Err(LoweringError::UnsupportedGraph)
884    }
885}
886
887fn fixed_positive_pair(values: impl Iterator<Item = i32>) -> Result<[u32; 2], LoweringError> {
888    let values = values.collect::<Vec<_>>();
889    if values.len() != 2 || values.iter().any(|value| *value <= 0) {
890        return Err(LoweringError::UnsupportedGraph);
891    }
892    Ok([
893        u32::try_from(values[0]).map_err(|_| LoweringError::UnsupportedGraph)?,
894        u32::try_from(values[1]).map_err(|_| LoweringError::UnsupportedGraph)?,
895    ])
896}
897
898fn tensor<'a>(
899    analysis: &'a TosaAnalysis<'a>,
900    value: ValueId,
901) -> Result<virtio_accel_tosa::Tensor<'a>, LoweringError> {
902    let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
903        return Err(LoweringError::UnsupportedGraph);
904    };
905    Ok(tensor)
906}
907
908fn scalar_zero_point(analysis: &TosaAnalysis<'_>, value: ValueId) -> Result<i32, LoweringError> {
909    let tensor = tensor(analysis, value)?;
910    if tensor.rank().is_none() || tensor.dimensions().any(|dimension| dimension != 1) {
911        return Err(LoweringError::InvalidConstant);
912    }
913    let bytes = analysis
914        .serialized_constant(value)
915        .ok_or(LoweringError::InvalidConstant)?;
916    match tensor.dtype() {
917        DType::FP16 if bytes.len() == 2 => {
918            let bits = u16::from_le_bytes(bytes.try_into().expect("length checked"));
919            (bits & 0x7fff == 0)
920                .then_some(0)
921                .ok_or(LoweringError::UnsupportedGraph)
922        }
923        DType::FP32 if bytes.len() == 4 => {
924            let value = f32::from_le_bytes(bytes.try_into().expect("length checked"));
925            (value == 0.0)
926                .then_some(0)
927                .ok_or(LoweringError::UnsupportedGraph)
928        }
929        DType::INT8 if bytes.len() == 1 => Ok(i32::from(bytes[0] as i8)),
930        _ => Err(LoweringError::InvalidConstant),
931    }
932}
933
934fn set_quantization_offset(
935    tensors: &mut [LoweredTensor],
936    assigned: &mut Vec<(u32, i32)>,
937    value: ValueId,
938    zero_point: i32,
939) -> Result<(), LoweringError> {
940    let tensor = tensors
941        .iter_mut()
942        .find(|tensor| tensor.value == value.get())
943        .ok_or(LoweringError::UnsupportedGraph)?;
944    let quantization = tensor
945        .quantization
946        .as_mut()
947        .ok_or(LoweringError::UnsupportedGraph)?;
948    let offset = zero_point
949        .checked_neg()
950        .ok_or(LoweringError::UnsupportedGraph)?;
951    if let Some((_, prior)) = assigned
952        .iter()
953        .find(|(assigned_value, _)| *assigned_value == value.get())
954    {
955        if *prior != offset {
956            return Err(LoweringError::UnsupportedGraph);
957        }
958    } else {
959        assigned
960            .try_reserve(1)
961            .map_err(|_| LoweringError::ResourceLimit)?;
962        assigned.push((value.get(), offset));
963    }
964    quantization.offset = offset;
965    Ok(())
966}
967
968fn constant_is_parameter_only(analysis: &TosaAnalysis<'_>, value: ValueId) -> bool {
969    let mut consumed = false;
970    for operator in analysis.operators() {
971        for (index, input) in analysis.operator_inputs(operator.id()).iter().enumerate() {
972            if *input != value {
973                continue;
974            }
975            consumed = true;
976            if !matches!((operator.op(), index), (Op::MATMUL, 2 | 3) | (Op::MUL, 2)) {
977                return false;
978            }
979        }
980    }
981    consumed
982}
983
984#[cfg(test)]
985mod tests {
986    use super::*;
987    use virtio_accel_conformance::numerics::{
988        ADD_FP16, HEXAGON_LOGICAL_CASES, HEXAGON_MOVEMENT_CASES, HEXAGON_REDUCTION_CASES,
989        HEXAGON_UNARY_FP16_CASES, IDENTITY_EDGES_FP16, IDENTITY_EDGES_FP32, IDENTITY_FP8E4M3,
990        IDENTITY_FP8E5M2, IDENTITY_INT4, IDENTITY_INT8, MATMUL_FP16, MATMUL_FP32, MATMUL_INT8,
991        MAX_POOL2D_FP16, MAX_POOL2D_FP32, MAXIMUM_FP16, MINIMUM_FP16, MUL_FP16, POW_FP16, SUB_FP16,
992    };
993
994    #[test]
995    fn plans_the_complete_initial_fp16_corpus_without_qairt() {
996        let cases = [
997            ("identity", IDENTITY_EDGES_FP16.artifact, 1usize, 2usize),
998            ("matmul", MATMUL_FP16.artifact, 1, 3),
999            ("max_pool2d", MAX_POOL2D_FP16.artifact, 1, 2),
1000        ];
1001        for (name, artifact, nodes, features) in cases {
1002            let lowered = lower_tosa(artifact, HEXAGON_TOSA_TARGET)
1003                .unwrap_or_else(|error| panic!("{name} failed to lower: {error}"));
1004            assert_eq!(lowered.nodes.len(), nodes, "{name}");
1005            assert_eq!(lowered.features.len(), features, "{name}");
1006            assert!(lowered.features.iter().all(|feature| feature.byte_len > 0));
1007        }
1008    }
1009
1010    #[test]
1011    fn matmul_discards_only_validated_zero_point_parameters() {
1012        let lowered = lower_tosa(MATMUL_FP16.artifact, HEXAGON_TOSA_TARGET).unwrap();
1013        assert_eq!(lowered.nodes[0].kind, NodeKind::MatMul);
1014        assert_eq!(lowered.features[0].slot, 0);
1015        assert_eq!(lowered.features[1].slot, 1);
1016        assert_eq!(lowered.features[2].slot, 2);
1017        assert_eq!(lowered.features[2].role, FeatureRole::Output);
1018    }
1019
1020    #[test]
1021    fn boundary_indices_do_not_depend_on_tensor_declaration_order() {
1022        let lowered = LoweredModel {
1023            tensors: vec![
1024                LoweredTensor {
1025                    value: 20,
1026                    element: Element::F16,
1027                    quantization: None,
1028                    dims: vec![1],
1029                    data: None,
1030                },
1031                LoweredTensor {
1032                    value: 10,
1033                    element: Element::F16,
1034                    quantization: None,
1035                    dims: vec![1],
1036                    data: None,
1037                },
1038                LoweredTensor {
1039                    value: 30,
1040                    element: Element::F16,
1041                    quantization: None,
1042                    dims: vec![1],
1043                    data: None,
1044                },
1045            ],
1046            nodes: vec![LoweredNode {
1047                kind: NodeKind::MatMul,
1048                inputs: vec![10, 20],
1049                outputs: vec![30],
1050                parameters: Vec::new(),
1051            }],
1052            features: vec![
1053                LoweredFeature {
1054                    slot: 0,
1055                    role: FeatureRole::Input,
1056                    io_index: 0,
1057                    value: 10,
1058                    dims: vec![1],
1059                    byte_len: 2,
1060                },
1061                LoweredFeature {
1062                    slot: 1,
1063                    role: FeatureRole::Input,
1064                    io_index: 1,
1065                    value: 20,
1066                    dims: vec![1],
1067                    byte_len: 2,
1068                },
1069                LoweredFeature {
1070                    slot: 2,
1071                    role: FeatureRole::Output,
1072                    io_index: 0,
1073                    value: 30,
1074                    dims: vec![1],
1075                    byte_len: 2,
1076                },
1077            ],
1078            precision: Some(Element::F16),
1079        };
1080
1081        assert_eq!(lowered.tensors[0].value, 20);
1082        assert_eq!(lowered.boundary(20), Some((FeatureRole::Input, 1)));
1083        assert_eq!(lowered.boundary(10), Some((FeatureRole::Input, 0)));
1084        assert_eq!(lowered.boundary(30), Some((FeatureRole::Output, 0)));
1085    }
1086
1087    #[test]
1088    fn max_pool_keeps_nhwc_shapes_and_attributes() {
1089        let lowered = lower_tosa(MAX_POOL2D_FP16.artifact, HEXAGON_TOSA_TARGET).unwrap();
1090        assert_eq!(lowered.nodes[0].kind, NodeKind::MaxPool2d);
1091        assert_eq!(lowered.nodes[0].parameters, [2, 2, 2, 2]);
1092        assert_eq!(lowered.features[0].dims.len(), 4);
1093        assert_eq!(lowered.features[1].dims.len(), 4);
1094    }
1095
1096    #[test]
1097    fn rejects_fp32_after_htp_precision_probe_detected_fp16_math() {
1098        for case in [IDENTITY_EDGES_FP32, MATMUL_FP32, MAX_POOL2D_FP32] {
1099            assert_eq!(
1100                lower_tosa(case.artifact, HEXAGON_TOSA_TARGET).unwrap_err(),
1101                LoweringError::UnsupportedType(DType::FP32),
1102                "{}",
1103                case.name,
1104            );
1105        }
1106    }
1107
1108    #[test]
1109    fn plans_exact_integer_identity_and_matmul_tier() {
1110        let identity = lower_tosa(IDENTITY_INT8.artifact, HEXAGON_TOSA_INTEGER_TARGET).unwrap();
1111        assert_eq!(identity.precision, None);
1112        assert!(
1113            identity
1114                .features
1115                .iter()
1116                .all(|feature| feature.byte_len == 8)
1117        );
1118
1119        let matmul = lower_tosa(MATMUL_INT8.artifact, HEXAGON_TOSA_INTEGER_TARGET).unwrap();
1120        assert_eq!(matmul.nodes[0].kind, NodeKind::MatMul);
1121        assert_eq!(matmul.features[0].byte_len, 6);
1122        assert_eq!(matmul.features[1].byte_len, 6);
1123        assert_eq!(matmul.features[2].byte_len, 16);
1124        assert!(matmul.tensors.iter().any(|tensor| {
1125            tensor.element == Element::I8
1126                && tensor
1127                    .quantization
1128                    .is_some_and(|quantization| quantization.offset != 0)
1129        }));
1130    }
1131
1132    #[test]
1133    fn integer_target_operator_surface_is_exact() {
1134        for op in [Op::CONST, Op::IDENTITY, Op::MATMUL] {
1135            assert!(supports_operator_for_target(op, true), "{op:?}");
1136        }
1137        for raw in Op::ARGMAX.get()..=Op::CONST_SHAPE.get() {
1138            let op = Op::new(raw);
1139            if !matches!(op, Op::CONST | Op::IDENTITY | Op::MATMUL) {
1140                assert!(!supports_operator_for_target(op, true), "{op:?}");
1141            }
1142        }
1143    }
1144
1145    #[test]
1146    fn descriptor_exposes_hexagon_pool_and_precision_restrictions() {
1147        assert!(!HEXAGON_TOSA_CAPABILITY.supports_dtype(DType::FP32, ValueRoles::INPUT));
1148        assert!(HEXAGON_TOSA_CAPABILITY.supports_dtype(DType::FP16, ValueRoles::INPUT));
1149        assert!(HEXAGON_TOSA_INTEGER_CAPABILITY.supports_dtype(DType::INT8, ValueRoles::INPUT));
1150        let pool = HEXAGON_TOSA_CAPABILITY.operator(Op::MAX_POOL2D).unwrap();
1151        assert!(
1152            pool.constraints
1153                .contains(OperatorConstraints::PROPAGATING_NAN)
1154        );
1155        assert!(pool.constraints.contains(OperatorConstraints::ZERO_PADDING));
1156    }
1157
1158    #[test]
1159    fn rejects_unadvertised_low_precision_profiles_and_extensions() {
1160        for (case, target) in [
1161            (
1162                IDENTITY_INT4,
1163                Target::new(
1164                    Version::TOSA_1_0,
1165                    ProfileSet::INTEGER,
1166                    Level::Level8K,
1167                    ExtensionSet::INT4,
1168                ),
1169            ),
1170            (
1171                IDENTITY_FP8E4M3,
1172                Target::new(
1173                    Version::TOSA_1_0,
1174                    ProfileSet::FLOATING_POINT,
1175                    Level::Level8K,
1176                    ExtensionSet::FP8E4M3,
1177                ),
1178            ),
1179            (
1180                IDENTITY_FP8E5M2,
1181                Target::new(
1182                    Version::TOSA_1_0,
1183                    ProfileSet::FLOATING_POINT,
1184                    Level::Level8K,
1185                    ExtensionSet::FP8E5M2,
1186                ),
1187            ),
1188        ] {
1189            assert_eq!(
1190                lower_tosa(case.artifact, target).unwrap_err(),
1191                LoweringError::UnsupportedGraph,
1192                "{}",
1193                case.name
1194            );
1195        }
1196    }
1197
1198    #[test]
1199    fn rejects_crossed_floating_and_integer_targets() {
1200        assert!(lower_tosa(IDENTITY_INT8.artifact, HEXAGON_TOSA_TARGET).is_err());
1201        assert!(lower_tosa(IDENTITY_EDGES_FP16.artifact, HEXAGON_TOSA_INTEGER_TARGET).is_err());
1202    }
1203
1204    #[test]
1205    fn plans_broadcast_binary_fp16_family() {
1206        for (case, kind) in [
1207            (ADD_FP16, NodeKind::Add),
1208            (SUB_FP16, NodeKind::Subtract),
1209            (MUL_FP16, NodeKind::Multiply),
1210            (POW_FP16, NodeKind::Power),
1211            (MAXIMUM_FP16, NodeKind::Maximum),
1212            (MINIMUM_FP16, NodeKind::Minimum),
1213        ] {
1214            let lowered = lower_tosa(case.artifact, HEXAGON_TOSA_TARGET)
1215                .unwrap_or_else(|error| panic!("{}: {error:?}", case.name));
1216            assert_eq!(lowered.nodes.len(), 1, "{}", case.name);
1217            assert_eq!(lowered.nodes[0].kind, kind, "{}", case.name);
1218            assert_eq!(lowered.nodes[0].inputs.len(), 2, "{}", case.name);
1219            assert_eq!(lowered.nodes[0].outputs.len(), 1, "{}", case.name);
1220        }
1221    }
1222
1223    #[test]
1224    fn plans_every_advertised_operator_family_without_qairt() {
1225        for case in HEXAGON_UNARY_FP16_CASES
1226            .iter()
1227            .chain(HEXAGON_LOGICAL_CASES)
1228            .chain(HEXAGON_REDUCTION_CASES)
1229            .chain(HEXAGON_MOVEMENT_CASES)
1230        {
1231            let lowered = lower_tosa(case.artifact, HEXAGON_TOSA_TARGET)
1232                .unwrap_or_else(|error| panic!("{}: {error:?}", case.name));
1233            assert!(!lowered.nodes.is_empty(), "{}", case.name);
1234            assert_eq!(
1235                lowered.features.len(),
1236                case.inputs.len() + 1,
1237                "{}",
1238                case.name
1239            );
1240        }
1241    }
1242
1243    #[test]
1244    fn advertised_operator_and_dtype_surface_is_exact() {
1245        let mut shared_count = 0;
1246        let mut exceptions = Vec::new();
1247        for raw in Op::ARGMAX.get()..=Op::CONST_SHAPE.get() {
1248            let op = Op::new(raw);
1249            let coreml = virtio_accel_coreml::supports_tosa_operator(op);
1250            let openvino = virtio_accel_openvino::supports_tosa_operator(op);
1251            assert_eq!(coreml, openvino, "shared providers disagree on {op:?}");
1252            if !coreml {
1253                assert!(
1254                    !supports_tosa_operator(op),
1255                    "Hexagon alone advertises {op:?}"
1256                );
1257                continue;
1258            }
1259            shared_count += 1;
1260            if !supports_tosa_operator(op) {
1261                exceptions.push(op);
1262            }
1263        }
1264        assert_eq!(shared_count, 42);
1265        assert_eq!(exceptions, [Op::ERF]);
1266        for dtype in [DType::BOOL, DType::FP16, DType::INT8, DType::INT32] {
1267            assert!(supports_tosa_dtype(dtype), "{dtype:?}");
1268        }
1269        for dtype in [DType::FP32, DType::INT4, DType::FP8E4M3, DType::FP8E5M2] {
1270            assert!(!supports_tosa_dtype(dtype), "{dtype:?}");
1271        }
1272    }
1273
1274    #[test]
1275    fn converts_binary16_attributes_without_losing_special_values() {
1276        assert_eq!(f16_to_f32(0x0000).to_bits(), 0x0000_0000);
1277        assert_eq!(f16_to_f32(0x8000).to_bits(), 0x8000_0000);
1278        assert_eq!(f16_to_f32(0x0001).to_bits(), 0x3380_0000);
1279        assert_eq!(f16_to_f32(0x3c00), 1.0);
1280        assert_eq!(f16_to_f32(0x7c00), f32::INFINITY);
1281        assert_eq!(f16_to_f32(0xfc00), f32::NEG_INFINITY);
1282        assert!(f16_to_f32(0x7e00).is_nan());
1283    }
1284
1285    #[test]
1286    fn rejects_malformed_artifacts_and_storage_overflow_before_native_work() {
1287        let truncated = &IDENTITY_EDGES_FP16.artifact[..IDENTITY_EDGES_FP16.artifact.len() / 2];
1288        assert!(matches!(
1289            lower_tosa(truncated, HEXAGON_TOSA_TARGET),
1290            Err(LoweringError::Parse(_))
1291        ));
1292        assert_eq!(
1293            checked_tensor_byte_len(Element::I32, &[u32::MAX, u32::MAX, u32::MAX]),
1294            Err(LoweringError::ResourceLimit)
1295        );
1296    }
1297}