Skip to main content

virtio_accel_coreml/
lower.rs

1//! TOSA 1.0 to Core ML neural-network lowering.
2//!
3//! This module intentionally owns Core ML's protobuf encoding. Portable crates expose only the
4//! verified TOSA model and provider-neutral analysis; no Core ML type, path, or dependency crosses
5//! the backend boundary.
6
7// Non-macOS builds type-check and unit-test this backend-local encoder, but only the macOS runtime
8// calls it from `load_program`.
9#![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
20/// TOSA target currently lowered by the Core ML backend.
21pub 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
91/// Conservative static floating-profile capability boundary for Core ML lowering.
92pub 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
104// Float16 MLMultiArray model boundaries require the iOS 16 / macOS 13 format revision. The
105// backend itself requires macOS 14, so all production TOSA models can use this version uniformly.
106const 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/// Backend-local execution choice for a verified TOSA program.
152///
153/// A whole-program identity is data movement rather than arithmetic. Keeping it out of Core ML
154/// preserves the exact byte representation required by TOSA for NaNs, signed zero, and subnormals.
155#[derive(Clone, Copy, Debug, PartialEq, Eq)]
156pub(crate) enum LoweredExecution {
157    CoreMl,
158    ExactCopy { input_slot: u32, output_slot: u32 },
159}
160
161/// Whether the initial Core ML lowering tier can lower `op` for supported types and attributes.
162pub const fn supports_tosa_operator(op: Op) -> bool {
163    COREML_TOSA_CAPABILITY.supports_operator(op)
164}
165
166/// Whether this lowering can expose `dtype` at a Core ML model boundary.
167///
168/// Operator-specific and target-specific validation still applies. INT8 is admitted only through
169/// the integer-profile ML Program tier on macOS 26 or newer; it is never silently dequantized into
170/// the floating-point NeuralNetwork tier.
171pub 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    // Exact rank mapping is mandatory for the general-ND layers used by this lowering.
245    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    // CTC constants consumed only by a layer parameter are deliberately absent from the Core ML
378    // graph. They have already been validated by TOSA analysis.
379    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        // Core ML's INVERSE and RSQRT unary modes force a nonzero default epsilon when zero is
450        // encoded. POWER avoids that numerical mismatch while retaining one fused unary layer.
451        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, &params);
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, &params);
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, &params);
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, &params);
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, &params);
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, &params);
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, &params);
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, &params);
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, &params);
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, &params);
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        // NeuralNetworkLayer.batchedMatmul is field 1045 (wire key 8362 = 0xaa 0x41).
1009        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        // The lowering emits transpose -> pooling -> transpose. Pooling is field 120
1033        // (wire key 962 = 0xc2 0x07); transpose is field 985 (0xca 0x3d).
1034        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}