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