1#![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
18pub const HEXAGON_TOSA_TARGET: Target = Target::new(
23 Version::TOSA_1_0,
24 ProfileSet::FLOATING_POINT,
25 Level::Level8K,
26 ExtensionSet::NONE,
27);
28
29pub 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
111pub 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
124pub 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#[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
157pub 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
170pub 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#[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#[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#[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#[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#[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}