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