1#![cfg_attr(not(target_os = "macos"), allow(dead_code))]
10
11use std::fmt;
12
13use virtio_accel_tosa::{
14 AnalysisError, AnalyzedValueKind, DType, Error as ParseError, ExtensionSet, Level,
15 NanPropagationMode, Op, OpAttributes, ProfileSet, Target, TosaAnalysis, ValueId, Version,
16 parse,
17};
18
19pub const COREML_TOSA_TARGET: Target = Target::new(
21 Version::TOSA_1_0,
22 ProfileSet::FLOATING_POINT,
23 Level::Level8K,
24 ExtensionSet::NONE,
25);
26
27const COREML_SPECIFICATION_VERSION: u64 = 7;
30const COREML_FLOAT16: u64 = 65_552;
31const COREML_FLOAT32: u64 = 65_568;
32const COREML_INT32: u64 = 131_104;
33
34#[derive(Clone, Copy, Debug, PartialEq, Eq)]
35pub enum LoweringError {
36 Parse(ParseError),
37 Analysis(AnalysisError),
38 UnsupportedGraph,
39 UnsupportedType(DType),
40 UnsupportedOperator(Op),
41 InvalidConstant,
42 ResourceLimit,
43}
44
45impl fmt::Display for LoweringError {
46 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
47 write!(formatter, "{self:?}")
48 }
49}
50
51impl std::error::Error for LoweringError {}
52
53#[derive(Clone, Copy, Debug, PartialEq, Eq)]
54pub(crate) enum LoweredFeatureRole {
55 Input,
56 Output,
57}
58
59#[derive(Clone, Debug, PartialEq, Eq)]
60pub(crate) struct LoweredFeature {
61 pub slot: u32,
62 pub role: LoweredFeatureRole,
63 pub name: String,
64}
65
66#[derive(Clone, Debug)]
67pub(crate) struct LoweredModel {
68 pub bytes: Vec<u8>,
69 pub features: Vec<LoweredFeature>,
70}
71
72pub const fn supports_tosa_operator(op: Op) -> bool {
74 matches!(
75 op,
76 Op::ARGMAX
77 | Op::MATMUL
78 | Op::MAX_POOL2D
79 | Op::CLAMP
80 | Op::ERF
81 | Op::SIGMOID
82 | Op::TANH
83 | Op::ADD
84 | Op::LOGICAL_AND
85 | Op::LOGICAL_OR
86 | Op::LOGICAL_XOR
87 | Op::MAXIMUM
88 | Op::MINIMUM
89 | Op::MUL
90 | Op::POW
91 | Op::SUB
92 | Op::ABS
93 | Op::CEIL
94 | Op::COS
95 | Op::EXP
96 | Op::FLOOR
97 | Op::LOG
98 | Op::LOGICAL_NOT
99 | Op::NEGATE
100 | Op::RECIPROCAL
101 | Op::RSQRT
102 | Op::SIN
103 | Op::SELECT
104 | Op::EQUAL
105 | Op::GREATER
106 | Op::GREATER_EQUAL
107 | Op::REDUCE_MAX
108 | Op::REDUCE_MIN
109 | Op::REDUCE_PRODUCT
110 | Op::REDUCE_SUM
111 | Op::CONCAT
112 | Op::RESHAPE
113 | Op::REVERSE
114 | Op::TRANSPOSE
115 | Op::CONST
116 | Op::CONST_SHAPE
117 | Op::IDENTITY
118 )
119}
120
121pub const fn supports_tosa_dtype(dtype: DType) -> bool {
127 matches!(dtype, DType::FP16 | DType::FP32 | DType::INT32)
128}
129
130pub(crate) fn lower_tosa(bytes: &[u8], target: Target) -> Result<LoweredModel, LoweringError> {
131 if target != COREML_TOSA_TARGET {
132 return Err(LoweringError::UnsupportedGraph);
133 }
134 let model = parse(bytes).map_err(LoweringError::Parse)?;
135 let analysis = model.analyze_for(target).map_err(LoweringError::Analysis)?;
136 if analysis.regions().len() != 1
137 || analysis.blocks().len() != 1
138 || !analysis.conditions().is_empty()
139 {
140 return Err(LoweringError::UnsupportedGraph);
141 }
142 let block = analysis.blocks()[0].id();
143 let inputs = analysis.block_inputs(block);
144 let outputs = analysis.block_outputs(block);
145 if inputs.is_empty()
146 || outputs.is_empty()
147 || inputs.iter().any(|input| outputs.contains(input))
148 || inputs.len().checked_add(outputs.len()).is_none()
149 {
150 return Err(LoweringError::UnsupportedGraph);
151 }
152
153 let mut names = analysis
154 .values()
155 .iter()
156 .map(|value| format!("v{}", value.id().get()))
157 .collect::<Vec<_>>();
158 let mut features = Vec::new();
159 features
160 .try_reserve_exact(inputs.len() + outputs.len())
161 .map_err(|_| LoweringError::ResourceLimit)?;
162 let mut description = Vec::new();
163
164 for (index, value) in inputs.iter().copied().enumerate() {
165 let name = format!("input_{index}");
166 names[value.get() as usize] = name.clone();
167 let tensor = tensor(&analysis, value)?;
168 encode_feature(&mut description, 1, &name, tensor)?;
169 features.push(LoweredFeature {
170 slot: u32::try_from(index).map_err(|_| LoweringError::ResourceLimit)?,
171 role: LoweredFeatureRole::Input,
172 name,
173 });
174 }
175 for (index, value) in outputs.iter().copied().enumerate() {
176 let name = format!("output_{index}");
177 names[value.get() as usize] = name.clone();
178 let tensor = tensor(&analysis, value)?;
179 encode_feature(&mut description, 10, &name, tensor)?;
180 features.push(LoweredFeature {
181 slot: u32::try_from(inputs.len() + index).map_err(|_| LoweringError::ResourceLimit)?,
182 role: LoweredFeatureRole::Output,
183 name,
184 });
185 }
186
187 let mut network = Vec::new();
188 for operator in analysis.execution_order(block) {
189 encode_operator(&mut network, &analysis, *operator, &names)?;
190 }
191 field_varint(&mut network, 5, 1);
193
194 let mut encoded = Vec::new();
195 field_varint(&mut encoded, 1, COREML_SPECIFICATION_VERSION);
196 field_message(&mut encoded, 2, &description);
197 field_message(&mut encoded, 500, &network);
198 Ok(LoweredModel {
199 bytes: encoded,
200 features,
201 })
202}
203
204fn tensor<'a>(
205 analysis: &'a TosaAnalysis<'a>,
206 value: ValueId,
207) -> Result<virtio_accel_tosa::Tensor<'a>, LoweringError> {
208 match analysis.value(value).kind() {
209 AnalyzedValueKind::Tensor(tensor) => Ok(tensor),
210 AnalyzedValueKind::Shape(_) => Err(LoweringError::UnsupportedGraph),
211 }
212}
213
214fn encode_feature(
215 description: &mut Vec<u8>,
216 field: u32,
217 name: &str,
218 tensor: virtio_accel_tosa::Tensor<'_>,
219) -> Result<(), LoweringError> {
220 let shape = static_shape(tensor)?;
221 if shape.is_empty() {
222 return Err(LoweringError::UnsupportedGraph);
223 }
224 let data_type = coreml_data_type(tensor.dtype())?;
225 let mut array = Vec::new();
226 field_packed_varints(
227 &mut array,
228 1,
229 shape.iter().copied().map(|value| value as u64),
230 );
231 field_varint(&mut array, 2, data_type);
232 let mut feature_type = Vec::new();
233 field_message(&mut feature_type, 5, &array);
234 let mut feature = Vec::new();
235 field_string(&mut feature, 1, name);
236 field_message(&mut feature, 3, &feature_type);
237 field_message(description, field, &feature);
238 Ok(())
239}
240
241fn encode_operator(
242 network: &mut Vec<u8>,
243 analysis: &TosaAnalysis<'_>,
244 operator_id: virtio_accel_tosa::OperatorId,
245 names: &[String],
246) -> Result<(), LoweringError> {
247 let operator = analysis.operator(operator_id);
248 let op = operator.op();
249 if !supports_tosa_operator(op) {
250 return Err(LoweringError::UnsupportedOperator(op));
251 }
252 let all_inputs = analysis.operator_inputs(operator_id);
253 let outputs = analysis.operator_outputs(operator_id);
254 let inputs = match op {
255 Op::MATMUL => {
256 for zero_point in &all_inputs[2..4] {
257 let bytes = analysis
258 .serialized_constant(*zero_point)
259 .ok_or(LoweringError::UnsupportedGraph)?;
260 if !serialized_float_is_zero(tensor(analysis, *zero_point)?.dtype(), bytes) {
261 return Err(LoweringError::UnsupportedGraph);
262 }
263 }
264 &all_inputs[..2]
265 }
266 Op::MUL => {
267 let shift = analysis
268 .serialized_constant(all_inputs[2])
269 .ok_or(LoweringError::UnsupportedGraph)?;
270 if shift.iter().any(|byte| *byte != 0) {
271 return Err(LoweringError::UnsupportedGraph);
272 }
273 &all_inputs[..2]
274 }
275 Op::NEGATE => {
276 for zero_point in &all_inputs[1..3] {
277 let bytes = analysis
278 .serialized_constant(*zero_point)
279 .ok_or(LoweringError::UnsupportedGraph)?;
280 if bytes.iter().any(|byte| *byte != 0) {
281 return Err(LoweringError::UnsupportedGraph);
282 }
283 }
284 &all_inputs[..1]
285 }
286 Op::RESHAPE => {
287 analysis
288 .serialized_constant(all_inputs[1])
289 .ok_or(LoweringError::UnsupportedGraph)?;
290 &all_inputs[..1]
291 }
292 _ => all_inputs,
293 };
294
295 if op == Op::CONST_SHAPE {
298 return Ok(());
299 }
300 if op == Op::CONST {
301 let output = outputs[0];
302 if constant_is_parameter_only(analysis, output) {
303 return Ok(());
304 }
305 let dtype = tensor(analysis, output)?.dtype();
306 if !matches!(dtype, DType::FP16 | DType::FP32 | DType::BOOL) {
307 return Err(LoweringError::UnsupportedType(dtype));
308 }
309 }
310
311 validate_operator_types(analysis, op, inputs, outputs)?;
312 match operator.source().attributes() {
313 OpAttributes::Maximum { nan_mode } | OpAttributes::Minimum { nan_mode } => {
314 require_propagating_nan(nan_mode)?;
315 }
316 _ => {}
317 }
318
319 if op == Op::MAX_POOL2D {
320 return encode_max_pool2d(network, analysis, operator_id, inputs, outputs, names);
321 }
322
323 let mut layer = Vec::new();
324 field_string(
325 &mut layer,
326 1,
327 &format!("tosa_{}_{}", operator_id.get(), op.name().unwrap_or("op")),
328 );
329 for value in inputs {
330 field_string(&mut layer, 2, &names[value.get() as usize]);
331 }
332 for value in outputs {
333 field_string(&mut layer, 3, &names[value.get() as usize]);
334 }
335
336 match op {
337 Op::IDENTITY => field_message(&mut layer, 600, &[]),
338 Op::MATMUL => field_message(&mut layer, 1045, &[]),
339 Op::ADD => field_message(&mut layer, 880, &[]),
340 Op::SUB => field_message(&mut layer, 905, &[]),
341 Op::MUL => field_message(&mut layer, 900, &[]),
342 Op::POW => field_message(&mut layer, 885, &[]),
343 Op::MAXIMUM => field_message(&mut layer, 875, &[]),
344 Op::MINIMUM => field_message(&mut layer, 870, &[]),
345 Op::EQUAL => field_message(&mut layer, 815, &[]),
346 Op::GREATER => field_message(&mut layer, 830, &[]),
347 Op::GREATER_EQUAL => field_message(&mut layer, 832, &[]),
348 Op::LOGICAL_OR => field_message(&mut layer, 840, &[]),
349 Op::LOGICAL_XOR => field_message(&mut layer, 845, &[]),
350 Op::LOGICAL_NOT => field_message(&mut layer, 850, &[]),
351 Op::LOGICAL_AND => field_message(&mut layer, 855, &[]),
352 Op::SELECT => field_message(&mut layer, 1330, &[]),
353 Op::CEIL => field_message(&mut layer, 665, &[]),
354 Op::FLOOR => field_message(&mut layer, 670, &[]),
355 Op::SIN => field_message(&mut layer, 710, &[]),
356 Op::COS => field_message(&mut layer, 715, &[]),
357 Op::TANH => field_message(&mut layer, 760, &[]),
358 Op::ERF => field_message(&mut layer, 790, &[]),
359 Op::SIGMOID => {
360 let mut activation = Vec::new();
361 field_message(&mut activation, 40, &[]);
362 field_message(&mut layer, 130, &activation);
363 }
364 Op::ABS => encode_unary(&mut layer, 6, None),
365 Op::EXP => encode_unary(&mut layer, 4, None),
366 Op::LOG => encode_unary(&mut layer, 5, None),
367 Op::RECIPROCAL => encode_unary(&mut layer, 3, Some(-1.0)),
370 Op::RSQRT => encode_unary(&mut layer, 3, Some(-0.5)),
371 Op::NEGATE => {
372 let mut multiply = Vec::new();
373 field_float(&mut multiply, 1, -1.0);
374 field_message(&mut layer, 231, &multiply);
375 }
376 Op::CLAMP => {
377 let OpAttributes::Clamp {
378 min_val,
379 max_val,
380 nan_mode,
381 } = operator.source().attributes()
382 else {
383 return Err(LoweringError::UnsupportedGraph);
384 };
385 require_propagating_nan(nan_mode)?;
386 let dtype = tensor(analysis, inputs[0])?.dtype();
387 let mut clip = Vec::new();
388 field_float(&mut clip, 1, decode_float(dtype, min_val)?);
389 field_float(&mut clip, 2, decode_float(dtype, max_val)?);
390 field_message(&mut layer, 660, &clip);
391 }
392 Op::ARGMAX => {
393 let OpAttributes::ArgMax { axis, nan_mode } = operator.source().attributes() else {
394 return Err(LoweringError::UnsupportedGraph);
395 };
396 require_propagating_nan(nan_mode)?;
397 let mut params = Vec::new();
398 field_signed(&mut params, 1, i64::from(axis));
399 field_varint(&mut params, 2, 1);
400 field_message(&mut layer, 1025, ¶ms);
401 }
402 Op::REDUCE_MAX | Op::REDUCE_MIN | Op::REDUCE_PRODUCT | Op::REDUCE_SUM => {
403 let axis = match operator.source().attributes() {
404 OpAttributes::ReduceMax { axis, nan_mode }
405 | OpAttributes::ReduceMin { axis, nan_mode } => {
406 require_propagating_nan(nan_mode)?;
407 axis
408 }
409 OpAttributes::ReduceProduct { axis } | OpAttributes::ReduceSum { axis } => axis,
410 _ => return Err(LoweringError::UnsupportedGraph),
411 };
412 let mut params = Vec::new();
413 field_packed_varints(&mut params, 1, [axis as i64 as u64]);
414 field_varint(&mut params, 2, 1);
415 let field = match op {
416 Op::REDUCE_MAX => 1260,
417 Op::REDUCE_MIN => 1265,
418 Op::REDUCE_SUM => 1270,
419 _ => 1275,
420 };
421 field_message(&mut layer, field, ¶ms);
422 }
423 Op::CONCAT => {
424 let OpAttributes::Concat { axis } = operator.source().attributes() else {
425 return Err(LoweringError::UnsupportedGraph);
426 };
427 let mut params = Vec::new();
428 field_signed(&mut params, 1, i64::from(axis));
429 field_message(&mut layer, 980, ¶ms);
430 }
431 Op::RESHAPE => {
432 let shape = static_shape(tensor(analysis, outputs[0])?)?;
433 let mut params = Vec::new();
434 field_packed_varints(
435 &mut params,
436 1,
437 shape.iter().copied().map(|value| value as u64),
438 );
439 field_message(&mut layer, 1140, ¶ms);
440 }
441 Op::REVERSE => {
442 let OpAttributes::Reverse { axis } = operator.source().attributes() else {
443 return Err(LoweringError::UnsupportedGraph);
444 };
445 let rank = tensor(analysis, inputs[0])?
446 .rank()
447 .ok_or(LoweringError::UnsupportedGraph)?;
448 let axis = usize::try_from(axis).map_err(|_| LoweringError::UnsupportedGraph)?;
449 let mut params = Vec::new();
450 field_packed_varints(
451 &mut params,
452 1,
453 (0..rank).map(|index| u64::from(index == axis)),
454 );
455 field_message(&mut layer, 960, ¶ms);
456 }
457 Op::TRANSPOSE => {
458 let OpAttributes::Transpose { perms } = operator.source().attributes() else {
459 return Err(LoweringError::UnsupportedGraph);
460 };
461 let mut params = Vec::new();
462 field_packed_varints(&mut params, 1, perms.iter().map(|axis| axis as u64));
463 field_message(&mut layer, 985, ¶ms);
464 }
465 Op::CONST => encode_constant(&mut layer, analysis, outputs[0])?,
466 _ => return Err(LoweringError::UnsupportedOperator(op)),
467 }
468 field_message(network, 1, &layer);
469 Ok(())
470}
471
472fn encode_max_pool2d(
473 network: &mut Vec<u8>,
474 analysis: &TosaAnalysis<'_>,
475 operator_id: virtio_accel_tosa::OperatorId,
476 inputs: &[ValueId],
477 outputs: &[ValueId],
478 names: &[String],
479) -> Result<(), LoweringError> {
480 let OpAttributes::MaxPool2d {
481 kernel,
482 stride,
483 pad,
484 nan_mode,
485 } = analysis.operator(operator_id).source().attributes()
486 else {
487 return Err(LoweringError::UnsupportedGraph);
488 };
489 require_propagating_nan(nan_mode)?;
490 let kernel = kernel.iter().collect::<Vec<_>>();
491 let stride = stride.iter().collect::<Vec<_>>();
492 let pad = pad.iter().collect::<Vec<_>>();
493 if kernel.len() != 2
494 || stride.len() != 2
495 || pad.len() != 4
496 || kernel.iter().chain(&stride).any(|value| *value <= 0)
497 || pad.iter().any(|value| *value != 0)
498 {
499 return Err(LoweringError::UnsupportedGraph);
500 }
501
502 let stem = format!("tosa_{}_max_pool2d", operator_id.get());
503 let nchw_input = format!("{stem}_nchw_input");
504 let nchw_output = format!("{stem}_nchw_output");
505 encode_transpose_layer(
506 network,
507 &format!("{stem}_to_nchw"),
508 &names[inputs[0].get() as usize],
509 &nchw_input,
510 [0, 3, 1, 2],
511 );
512
513 let mut params = Vec::new();
514 field_packed_varints(
515 &mut params,
516 10,
517 kernel.into_iter().map(|value| value as u64),
518 );
519 field_packed_varints(
520 &mut params,
521 20,
522 stride.into_iter().map(|value| value as u64),
523 );
524 field_message(&mut params, 30, &[]);
525 let mut pooling = Vec::new();
526 field_string(&mut pooling, 1, &stem);
527 field_string(&mut pooling, 2, &nchw_input);
528 field_string(&mut pooling, 3, &nchw_output);
529 field_message(&mut pooling, 120, ¶ms);
530 field_message(network, 1, &pooling);
531
532 encode_transpose_layer(
533 network,
534 &format!("{stem}_to_nhwc"),
535 &nchw_output,
536 &names[outputs[0].get() as usize],
537 [0, 2, 3, 1],
538 );
539 Ok(())
540}
541
542fn encode_transpose_layer(
543 network: &mut Vec<u8>,
544 name: &str,
545 input: &str,
546 output: &str,
547 axes: impl IntoIterator<Item = u64>,
548) {
549 let mut params = Vec::new();
550 field_packed_varints(&mut params, 1, axes);
551 let mut layer = Vec::new();
552 field_string(&mut layer, 1, name);
553 field_string(&mut layer, 2, input);
554 field_string(&mut layer, 3, output);
555 field_message(&mut layer, 985, ¶ms);
556 field_message(network, 1, &layer);
557}
558
559fn constant_is_parameter_only(analysis: &TosaAnalysis<'_>, value: ValueId) -> bool {
560 let mut consumed = false;
561 for operator in analysis.operators() {
562 for (index, input) in analysis.operator_inputs(operator.id()).iter().enumerate() {
563 if *input != value {
564 continue;
565 }
566 consumed = true;
567 if !matches!(
568 (operator.op(), index),
569 (Op::MATMUL, 2 | 3) | (Op::MUL, 2) | (Op::NEGATE, 1 | 2) | (Op::RESHAPE, 1)
570 ) {
571 return false;
572 }
573 }
574 }
575 consumed
576}
577
578fn validate_operator_types(
579 analysis: &TosaAnalysis<'_>,
580 op: Op,
581 inputs: &[ValueId],
582 outputs: &[ValueId],
583) -> Result<(), LoweringError> {
584 let require = |value, predicate: fn(DType) -> bool| {
585 let dtype = tensor(analysis, value)?.dtype();
586 if predicate(dtype) {
587 Ok(())
588 } else {
589 Err(LoweringError::UnsupportedType(dtype))
590 }
591 };
592 let is_float = |dtype| matches!(dtype, DType::FP16 | DType::FP32);
593 let is_bool = |dtype| dtype == DType::BOOL;
594 let is_int32 = |dtype| dtype == DType::INT32;
595
596 match op {
597 Op::CONST => {
598 require(outputs[0], |dtype| {
599 matches!(dtype, DType::FP16 | DType::FP32 | DType::BOOL)
600 })?;
601 }
602 Op::LOGICAL_AND | Op::LOGICAL_OR | Op::LOGICAL_XOR | Op::LOGICAL_NOT => {
603 for value in inputs.iter().chain(outputs) {
604 require(*value, is_bool)?;
605 }
606 }
607 Op::EQUAL | Op::GREATER | Op::GREATER_EQUAL => {
608 for value in inputs {
609 require(*value, is_float)?;
610 }
611 require(outputs[0], is_bool)?;
612 }
613 Op::SELECT => {
614 require(inputs[0], is_bool)?;
615 for value in inputs[1..].iter().chain(outputs) {
616 require(*value, is_float)?;
617 }
618 }
619 Op::ARGMAX => {
620 require(inputs[0], is_float)?;
621 require(outputs[0], is_int32)?;
622 }
623 _ => {
624 for value in inputs.iter().chain(outputs) {
625 require(*value, is_float)?;
626 }
627 }
628 }
629 Ok(())
630}
631
632fn require_propagating_nan(nan_mode: NanPropagationMode) -> Result<(), LoweringError> {
633 if nan_mode == NanPropagationMode::PROPAGATE {
634 Ok(())
635 } else {
636 Err(LoweringError::UnsupportedGraph)
637 }
638}
639
640fn encode_unary(layer: &mut Vec<u8>, operation: u64, alpha: Option<f32>) {
641 let mut params = Vec::new();
642 field_varint(&mut params, 1, operation);
643 if let Some(alpha) = alpha {
644 field_float(&mut params, 2, alpha);
645 }
646 field_message(layer, 220, ¶ms);
647}
648
649fn encode_constant(
650 layer: &mut Vec<u8>,
651 analysis: &TosaAnalysis<'_>,
652 output: ValueId,
653) -> Result<(), LoweringError> {
654 let tensor = tensor(analysis, output)?;
655 let data = analysis
656 .serialized_constant(output)
657 .ok_or(LoweringError::InvalidConstant)?;
658 let mut shape = static_shape(tensor)?;
659 if shape.is_empty() {
660 shape.push(1);
661 }
662 let mut weights = Vec::new();
663 match tensor.dtype() {
664 DType::FP32 => {
665 if data.len() % 4 != 0 {
666 return Err(LoweringError::InvalidConstant);
667 }
668 field_bytes(&mut weights, 1, data);
669 }
670 DType::FP16 => field_bytes(&mut weights, 2, data),
671 DType::BOOL => {
672 let mut floats = Vec::new();
673 floats
674 .try_reserve_exact(data.len() * 4)
675 .map_err(|_| LoweringError::ResourceLimit)?;
676 for value in data {
677 floats.extend_from_slice(&f32::from(*value != 0).to_le_bytes());
678 }
679 field_bytes(&mut weights, 1, &floats);
680 }
681 dtype => return Err(LoweringError::UnsupportedType(dtype)),
682 }
683 let mut params = Vec::new();
684 field_packed_varints(
685 &mut params,
686 1,
687 shape.iter().copied().map(|value| value as u64),
688 );
689 field_message(&mut params, 2, &weights);
690 field_message(layer, 1070, ¶ms);
691 Ok(())
692}
693
694fn static_shape(tensor: virtio_accel_tosa::Tensor<'_>) -> Result<Vec<i32>, LoweringError> {
695 tensor.rank().ok_or(LoweringError::UnsupportedGraph)?;
696 let shape = tensor.dimensions().collect::<Vec<_>>();
697 if shape.iter().any(|dimension| *dimension <= 0) {
698 return Err(LoweringError::UnsupportedGraph);
699 }
700 Ok(shape)
701}
702
703fn coreml_data_type(dtype: DType) -> Result<u64, LoweringError> {
704 match dtype {
705 DType::FP16 => Ok(COREML_FLOAT16),
706 DType::FP32 => Ok(COREML_FLOAT32),
707 DType::INT32 => Ok(COREML_INT32),
708 _ => Err(LoweringError::UnsupportedType(dtype)),
709 }
710}
711
712fn decode_float(dtype: DType, bytes: &[u8]) -> Result<f32, LoweringError> {
713 match dtype {
714 DType::FP16 if bytes.len() == 2 => Ok(f16_to_f32(u16::from_le_bytes(
715 bytes.try_into().expect("length checked"),
716 ))),
717 DType::FP32 if bytes.len() == 4 => Ok(f32::from_le_bytes(bytes.try_into().unwrap())),
718 _ => Err(LoweringError::UnsupportedType(dtype)),
719 }
720}
721
722fn f16_to_f32(bits: u16) -> f32 {
723 let sign = u32::from(bits & 0x8000) << 16;
724 let exponent = (bits >> 10) & 0x1f;
725 let fraction = u32::from(bits & 0x03ff);
726 let converted = match exponent {
727 0 if fraction == 0 => sign,
728 0 => {
729 let shift = fraction.leading_zeros() - 21;
730 let normalized = fraction << shift;
731 sign | ((127 - 15 - shift + 1) << 23) | ((normalized & 0x03ff) << 13)
732 }
733 0x1f => sign | 0x7f80_0000 | (fraction << 13),
734 _ => sign | ((u32::from(exponent) + 127 - 15) << 23) | (fraction << 13),
735 };
736 f32::from_bits(converted)
737}
738
739fn serialized_float_is_zero(dtype: DType, bytes: &[u8]) -> bool {
740 match dtype {
741 DType::FP16 if bytes.len() == 2 => {
742 u16::from_le_bytes(bytes.try_into().expect("length checked")) & 0x7fff == 0
743 }
744 DType::FP32 if bytes.len() == 4 => {
745 u32::from_le_bytes(bytes.try_into().expect("length checked")) & 0x7fff_ffff == 0
746 }
747 _ => false,
748 }
749}
750
751fn field_varint(target: &mut Vec<u8>, field: u32, value: u64) {
752 varint(target, u64::from(field) << 3);
753 varint(target, value);
754}
755
756fn field_signed(target: &mut Vec<u8>, field: u32, value: i64) {
757 field_varint(target, field, value as u64);
758}
759
760fn field_float(target: &mut Vec<u8>, field: u32, value: f32) {
761 varint(target, (u64::from(field) << 3) | 5);
762 target.extend_from_slice(&value.to_le_bytes());
763}
764
765fn field_string(target: &mut Vec<u8>, field: u32, value: &str) {
766 field_bytes(target, field, value.as_bytes());
767}
768
769fn field_message(target: &mut Vec<u8>, field: u32, message: &[u8]) {
770 field_bytes(target, field, message);
771}
772
773fn field_bytes(target: &mut Vec<u8>, field: u32, bytes: &[u8]) {
774 varint(target, (u64::from(field) << 3) | 2);
775 varint(target, bytes.len() as u64);
776 target.extend_from_slice(bytes);
777}
778
779fn field_packed_varints(target: &mut Vec<u8>, field: u32, values: impl IntoIterator<Item = u64>) {
780 let mut packed = Vec::new();
781 for value in values {
782 varint(&mut packed, value);
783 }
784 field_bytes(target, field, &packed);
785}
786
787fn varint(target: &mut Vec<u8>, mut value: u64) {
788 while value >= 0x80 {
789 target.push((value as u8) | 0x80);
790 value >>= 7;
791 }
792 target.push(value as u8);
793}
794
795#[cfg(test)]
796mod tests {
797 use super::*;
798
799 const IDENTITY_FP32: &[u8] = include_bytes!("../tests/data/identity-fp32-v1.0.0.tosa");
800 #[test]
801 fn lowers_a_verified_tosa_model_without_host_dependencies() {
802 let lowered = lower_tosa(IDENTITY_FP32, COREML_TOSA_TARGET).unwrap();
803
804 assert!(!lowered.bytes.is_empty());
805 assert_eq!(lowered.features.len(), 2);
806 assert_eq!(lowered.features[0].slot, 0);
807 assert_eq!(lowered.features[0].role, LoweredFeatureRole::Input);
808 assert_eq!(lowered.features[1].slot, 1);
809 assert_eq!(lowered.features[1].role, LoweredFeatureRole::Output);
810 }
811
812 #[test]
813 fn rejects_a_different_tosa_target_before_parsing() {
814 let target = Target::new(
815 Version::TOSA_1_0,
816 ProfileSet::INTEGER,
817 Level::Level8K,
818 ExtensionSet::NONE,
819 );
820
821 assert!(matches!(
822 lower_tosa(IDENTITY_FP32, target),
823 Err(LoweringError::UnsupportedGraph)
824 ));
825 }
826
827 #[test]
828 fn reports_low_precision_types_as_unsupported_in_the_neural_network_lowering() {
829 assert!(supports_tosa_dtype(DType::FP16));
830 assert!(supports_tosa_dtype(DType::FP32));
831 assert!(supports_tosa_dtype(DType::INT32));
832 assert!(!supports_tosa_dtype(DType::INT8));
833 assert!(!supports_tosa_dtype(DType::INT4));
834 assert!(!supports_tosa_dtype(DType::FP8E4M3));
835 assert!(!supports_tosa_dtype(DType::FP8E5M2));
836 }
837
838 #[test]
839 fn rejects_shared_quantized_artifacts_at_the_declared_target_boundary() {
840 use virtio_accel_conformance::numerics::{
841 IDENTITY_FP8E4M3, IDENTITY_FP8E5M2, IDENTITY_INT4, IDENTITY_INT8,
842 };
843
844 for (case, target) in [
845 (
846 IDENTITY_INT8,
847 Target::new(
848 Version::TOSA_1_0,
849 ProfileSet::INTEGER,
850 Level::Level8K,
851 ExtensionSet::NONE,
852 ),
853 ),
854 (
855 IDENTITY_INT4,
856 Target::new(
857 Version::TOSA_1_0,
858 ProfileSet::INTEGER,
859 Level::Level8K,
860 ExtensionSet::INT4,
861 ),
862 ),
863 (
864 IDENTITY_FP8E4M3,
865 Target::new(
866 Version::TOSA_1_0,
867 ProfileSet::FLOATING_POINT,
868 Level::Level8K,
869 ExtensionSet::FP8E4M3,
870 ),
871 ),
872 (
873 IDENTITY_FP8E5M2,
874 Target::new(
875 Version::TOSA_1_0,
876 ProfileSet::FLOATING_POINT,
877 Level::Level8K,
878 ExtensionSet::FP8E5M2,
879 ),
880 ),
881 ] {
882 assert!(matches!(
883 lower_tosa(case.artifact, target),
884 Err(LoweringError::UnsupportedGraph)
885 ));
886 }
887 }
888
889 #[test]
890 fn lowers_batched_matmul_without_encoding_parameter_constants() {
891 let lowered = lower_tosa(
892 virtio_accel_conformance::numerics::MATMUL_FP32.artifact,
893 COREML_TOSA_TARGET,
894 )
895 .unwrap();
896
897 assert!(!lowered.bytes.is_empty());
898 assert_eq!(lowered.features.len(), 3);
899 assert_eq!(lowered.features[0].slot, 0);
900 assert_eq!(lowered.features[1].slot, 1);
901 assert_eq!(lowered.features[2].slot, 2);
902 assert!(lowered.bytes.windows(2).any(|bytes| bytes == [0xaa, 0x41]));
904 }
905
906 #[test]
907 fn lowers_the_shared_fp32_edge_identity_artifact() {
908 let lowered = lower_tosa(
909 virtio_accel_conformance::numerics::IDENTITY_EDGES_FP32.artifact,
910 COREML_TOSA_TARGET,
911 )
912 .unwrap();
913
914 assert_eq!(lowered.features.len(), 2);
915 assert!(!lowered.bytes.is_empty());
916 }
917
918 #[test]
919 fn lowers_nhwc_max_pool_through_explicit_layout_transposes() {
920 let lowered = lower_tosa(
921 virtio_accel_conformance::numerics::MAX_POOL2D_FP32.artifact,
922 COREML_TOSA_TARGET,
923 )
924 .unwrap();
925
926 assert_eq!(
929 lowered
930 .bytes
931 .windows(2)
932 .filter(|bytes| *bytes == [0xca, 0x3d])
933 .count(),
934 2
935 );
936 assert!(lowered.bytes.windows(2).any(|bytes| bytes == [0xc2, 0x07]));
937 }
938
939 #[test]
940 fn lowers_every_shared_fp16_numerical_artifact() {
941 use virtio_accel_conformance::numerics::{
942 IDENTITY_EDGES_FP16, MATMUL_FP16, MAX_POOL2D_FP16,
943 };
944
945 for case in [IDENTITY_EDGES_FP16, MATMUL_FP16, MAX_POOL2D_FP16] {
946 let lowered = lower_tosa(case.artifact, COREML_TOSA_TARGET).unwrap();
947 assert!(!lowered.bytes.is_empty(), "{}", case.name);
948 }
949 }
950
951 #[test]
952 fn greater_equal_uses_the_distinct_core_ml_field() {
953 assert!(supports_tosa_operator(Op::GREATER_EQUAL));
954 let mut layer = Vec::new();
955 field_message(&mut layer, 832, &[]);
956 assert_eq!(layer, [0x82, 0x34, 0x00]);
957 }
958
959 #[test]
960 fn fp16_parameters_preserve_zero_finite_and_nan_classes() {
961 assert_eq!(decode_float(DType::FP16, &0_u16.to_le_bytes()), Ok(0.0));
962 assert_eq!(
963 decode_float(DType::FP16, &0x8000_u16.to_le_bytes())
964 .unwrap()
965 .to_bits(),
966 (-0.0_f32).to_bits()
967 );
968 assert_eq!(
969 decode_float(DType::FP16, &0x3c00_u16.to_le_bytes()),
970 Ok(1.0)
971 );
972 assert!(
973 decode_float(DType::FP16, &0x7e00_u16.to_le_bytes())
974 .unwrap()
975 .is_nan()
976 );
977 assert_eq!(
978 decode_float(DType::FP16, &0x0001_u16.to_le_bytes())
979 .unwrap()
980 .to_bits(),
981 (2.0_f32.powi(-24)).to_bits()
982 );
983 }
984}