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