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