Skip to main content

virtio_accel_coreml/
mlprogram.rs

1//! Dependency-free Core ML ML Program encoding for the exact TOSA integer tier.
2//!
3//! Core ML's older `NeuralNetwork` model family cannot expose INT8 multi-array boundaries. The
4//! ML Program format added that boundary on macOS 26. This module deliberately implements only
5//! operations whose integer semantics have been executed against the shared Rust oracle:
6//! same-type identity and INT8 MATMUL with explicit zero-point subtraction and INT32 accumulation.
7//! No path converts integer tensors through floating point.
8
9#![cfg_attr(not(target_os = "macos"), allow(dead_code))]
10
11use crate::lower::{
12    LoweredExecution, LoweredFeature, LoweredFeatureRole, LoweredModel, LoweringError,
13    encode_feature, static_shape,
14};
15use virtio_accel_tosa::{
16    AnalyzedValueKind, CapabilityDescriptor, DType, DTypeCapability, ExtensionSet,
17    GraphCapabilities, Level, Op, OperatorCapability, ProfileSet, RuntimeConditionSupport, Target,
18    TosaAnalysis, ValueId, ValueRoles, Version, parse,
19};
20
21/// TOSA integer-profile target lowered to Core ML ML Program on macOS 26 or newer.
22pub const COREML_TOSA_INTEGER_TARGET: Target = Target::new(
23    Version::TOSA_1_0,
24    ProfileSet::INTEGER,
25    Level::Level8K,
26    ExtensionSet::NONE,
27);
28
29const INTEGER_DTYPES: &[DTypeCapability] = &[
30    DTypeCapability::new(DType::INT8, ValueRoles::ALL),
31    DTypeCapability::new(
32        DType::INT32,
33        ValueRoles::OUTPUT
34            .union(ValueRoles::CONSTANT)
35            .union(ValueRoles::INTERMEDIATE),
36    ),
37];
38
39const INTEGER_OPERATORS: &[OperatorCapability] = &[
40    OperatorCapability::new(Op::CONST),
41    OperatorCapability::new(Op::IDENTITY),
42    OperatorCapability::new(Op::MATMUL),
43];
44
45/// Conservative integer-profile boundary for the macOS 26+ ML Program tier.
46pub const COREML_TOSA_INTEGER_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
47    target: COREML_TOSA_INTEGER_TARGET,
48    dtypes: INTEGER_DTYPES,
49    operators: INTEGER_OPERATORS,
50    graph: GraphCapabilities {
51        max_regions: 1,
52        max_blocks: 1,
53        dynamic_shapes: false,
54        runtime_conditions: RuntimeConditionSupport::None,
55    },
56};
57
58const COREML_SPECIFICATION_VERSION: u64 = 10;
59const MLPROGRAM_VERSION: u64 = 1;
60const OPSET: &str = "CoreML9";
61
62// MILSpec.DataType values. These are distinct from MLMultiArrayDataType values in the model
63// description encoded by `encode_feature`.
64const MIL_BOOL: u64 = 1;
65const MIL_STRING: u64 = 2;
66const MIL_INT8: u64 = 21;
67const MIL_INT32: u64 = 23;
68
69pub(crate) fn lower_integer_tosa(
70    bytes: &[u8],
71    target: Target,
72) -> Result<LoweredModel, LoweringError> {
73    if target != COREML_TOSA_INTEGER_TARGET {
74        return Err(LoweringError::UnsupportedGraph);
75    }
76    let model = parse(bytes).map_err(LoweringError::Parse)?;
77    let analysis = model.analyze_for(target).map_err(LoweringError::Analysis)?;
78    if analysis.regions().len() != 1
79        || analysis.blocks().len() != 1
80        || !analysis.conditions().is_empty()
81    {
82        return Err(LoweringError::UnsupportedGraph);
83    }
84    let block = analysis.blocks()[0].id();
85    let inputs = analysis.block_inputs(block);
86    let outputs = analysis.block_outputs(block);
87    if inputs.is_empty() || outputs.is_empty() || inputs.iter().any(|id| outputs.contains(id)) {
88        return Err(LoweringError::UnsupportedGraph);
89    }
90
91    let mut description = Vec::new();
92    let mut features = Vec::new();
93    features
94        .try_reserve_exact(
95            inputs
96                .len()
97                .checked_add(outputs.len())
98                .ok_or(LoweringError::ResourceLimit)?,
99        )
100        .map_err(|_| LoweringError::ResourceLimit)?;
101    for (index, value) in inputs.iter().copied().enumerate() {
102        let tensor = tensor(&analysis, value)?;
103        if tensor.dtype() != DType::INT8 {
104            return Err(LoweringError::UnsupportedType(tensor.dtype()));
105        }
106        let name = format!("input_{index}");
107        encode_feature(&mut description, 1, &name, tensor)?;
108        features.push(LoweredFeature {
109            slot: u32::try_from(index).map_err(|_| LoweringError::ResourceLimit)?,
110            role: LoweredFeatureRole::Input,
111            name,
112        });
113    }
114    for (index, value) in outputs.iter().copied().enumerate() {
115        let tensor = tensor(&analysis, value)?;
116        if !matches!(tensor.dtype(), DType::INT8 | DType::INT32) {
117            return Err(LoweringError::UnsupportedType(tensor.dtype()));
118        }
119        let name = format!("output_{index}");
120        encode_feature(&mut description, 10, &name, tensor)?;
121        features.push(LoweredFeature {
122            slot: u32::try_from(inputs.len() + index).map_err(|_| LoweringError::ResourceLimit)?,
123            role: LoweredFeatureRole::Output,
124            name,
125        });
126    }
127
128    let executable = analysis
129        .execution_order(block)
130        .iter()
131        .copied()
132        .filter(|operator| {
133            !matches!(
134                analysis.operator(*operator).op(),
135                Op::CONST | Op::CONST_SHAPE
136            )
137        })
138        .collect::<Vec<_>>();
139    if executable.len() != 1 {
140        return Err(LoweringError::UnsupportedGraph);
141    }
142    let operator = executable[0];
143    let operations = match analysis.operator(operator).op() {
144        Op::IDENTITY => encode_identity(&analysis, operator, inputs, outputs)?,
145        Op::MATMUL => encode_matmul(&analysis, operator, inputs, outputs)?,
146        op => return Err(LoweringError::UnsupportedOperator(op)),
147    };
148
149    let mut block_body = Vec::new();
150    for index in 0..outputs.len() {
151        field_string(&mut block_body, 2, &format!("output_{index}"));
152    }
153    for operation in operations {
154        field_message(&mut block_body, 3, &operation);
155    }
156
157    let mut function = Vec::new();
158    for (index, value) in inputs.iter().copied().enumerate() {
159        let shape = static_shape(tensor(&analysis, value)?)?;
160        field_message(
161            &mut function,
162            1,
163            &named_value_type(&format!("input_{index}"), MIL_INT8, &shape),
164        );
165    }
166    field_string(&mut function, 2, OPSET);
167    field_message(&mut function, 3, &map_entry(OPSET, &block_body));
168
169    let mut program = Vec::new();
170    field_varint(&mut program, 1, MLPROGRAM_VERSION);
171    field_message(&mut program, 2, &map_entry("main", &function));
172
173    let mut encoded = Vec::new();
174    field_varint(&mut encoded, 1, COREML_SPECIFICATION_VERSION);
175    field_message(&mut encoded, 2, &description);
176    field_message(&mut encoded, 502, &program);
177    Ok(LoweredModel {
178        bytes: encoded,
179        features,
180        execution: LoweredExecution::CoreMl,
181    })
182}
183
184fn encode_identity(
185    analysis: &TosaAnalysis<'_>,
186    operator: virtio_accel_tosa::OperatorId,
187    block_inputs: &[ValueId],
188    block_outputs: &[ValueId],
189) -> Result<Vec<Vec<u8>>, LoweringError> {
190    let inputs = analysis.operator_inputs(operator);
191    let outputs = analysis.operator_outputs(operator);
192    if block_inputs.len() != 1
193        || block_outputs.len() != 1
194        || inputs != block_inputs
195        || outputs != block_outputs
196    {
197        return Err(LoweringError::UnsupportedGraph);
198    }
199    let input = tensor(analysis, inputs[0])?;
200    let output = tensor(analysis, outputs[0])?;
201    let input_shape = static_shape(input)?;
202    let output_shape = static_shape(output)?;
203    if input.dtype() != DType::INT8 || output.dtype() != DType::INT8 || input_shape != output_shape
204    {
205        return Err(LoweringError::UnsupportedGraph);
206    }
207
208    // Core ML rejects an INT8 `identity` MIL op. A widening cast, exact add-zero, and narrowing
209    // cast is accepted and preserves every signed byte value exactly.
210    Ok(vec![
211        const_string("identity_wide_dtype", "int32"),
212        unary_operation(
213            "cast",
214            &[("x", "input_0"), ("dtype", "identity_wide_dtype")],
215            "identity_wide",
216            MIL_INT32,
217            &input_shape,
218        ),
219        const_int32("identity_zero", 0),
220        unary_operation(
221            "add",
222            &[("x", "identity_wide"), ("y", "identity_zero")],
223            "identity_exact",
224            MIL_INT32,
225            &input_shape,
226        ),
227        const_string("identity_narrow_dtype", "int8"),
228        unary_operation(
229            "cast",
230            &[("x", "identity_exact"), ("dtype", "identity_narrow_dtype")],
231            "output_0",
232            MIL_INT8,
233            &output_shape,
234        ),
235    ])
236}
237
238fn encode_matmul(
239    analysis: &TosaAnalysis<'_>,
240    operator: virtio_accel_tosa::OperatorId,
241    block_inputs: &[ValueId],
242    block_outputs: &[ValueId],
243) -> Result<Vec<Vec<u8>>, LoweringError> {
244    let inputs = analysis.operator_inputs(operator);
245    let outputs = analysis.operator_outputs(operator);
246    if block_inputs.len() != 2
247        || block_outputs.len() != 1
248        || inputs.len() != 4
249        || outputs.len() != 1
250        || inputs[..2] != *block_inputs
251        || outputs != block_outputs
252    {
253        return Err(LoweringError::UnsupportedGraph);
254    }
255    let lhs = tensor(analysis, inputs[0])?;
256    let rhs = tensor(analysis, inputs[1])?;
257    let output = tensor(analysis, outputs[0])?;
258    let lhs_shape = static_shape(lhs)?;
259    let rhs_shape = static_shape(rhs)?;
260    let output_shape = static_shape(output)?;
261    if lhs.dtype() != DType::INT8
262        || rhs.dtype() != DType::INT8
263        || output.dtype() != DType::INT32
264        || lhs_shape.len() != 3
265        || rhs_shape.len() != 3
266        || output_shape.len() != 3
267    {
268        return Err(LoweringError::UnsupportedGraph);
269    }
270    let zero_point = |value: ValueId| {
271        let bytes = analysis
272            .serialized_constant(value)
273            .ok_or(LoweringError::UnsupportedGraph)?;
274        if tensor(analysis, value)?.dtype() != DType::INT8 || bytes.len() != 1 {
275            return Err(LoweringError::UnsupportedGraph);
276        }
277        Ok(i32::from(bytes[0] as i8))
278    };
279    let lhs_zero_point = zero_point(inputs[2])?;
280    let rhs_zero_point = zero_point(inputs[3])?;
281
282    let mut operations = Vec::new();
283    operations
284        .try_reserve_exact(11)
285        .map_err(|_| LoweringError::ResourceLimit)?;
286    operations.push(const_string("lhs_wide_dtype", "int32"));
287    operations.push(unary_operation(
288        "cast",
289        &[("x", "input_0"), ("dtype", "lhs_wide_dtype")],
290        "lhs_wide",
291        MIL_INT32,
292        &lhs_shape,
293    ));
294    operations.push(const_string("rhs_wide_dtype", "int32"));
295    operations.push(unary_operation(
296        "cast",
297        &[("x", "input_1"), ("dtype", "rhs_wide_dtype")],
298        "rhs_wide",
299        MIL_INT32,
300        &rhs_shape,
301    ));
302    operations.push(const_int32("lhs_zero_point", lhs_zero_point));
303    operations.push(unary_operation(
304        "sub",
305        &[("x", "lhs_wide"), ("y", "lhs_zero_point")],
306        "lhs_centered",
307        MIL_INT32,
308        &lhs_shape,
309    ));
310    operations.push(const_int32("rhs_zero_point", rhs_zero_point));
311    operations.push(unary_operation(
312        "sub",
313        &[("x", "rhs_wide"), ("y", "rhs_zero_point")],
314        "rhs_centered",
315        MIL_INT32,
316        &rhs_shape,
317    ));
318    operations.push(const_bool("matmul_transpose_x", false));
319    operations.push(const_bool("matmul_transpose_y", false));
320    operations.push(unary_operation(
321        "matmul",
322        &[
323            ("x", "lhs_centered"),
324            ("y", "rhs_centered"),
325            ("transpose_x", "matmul_transpose_x"),
326            ("transpose_y", "matmul_transpose_y"),
327        ],
328        "output_0",
329        MIL_INT32,
330        &output_shape,
331    ));
332    Ok(operations)
333}
334
335fn tensor<'a>(
336    analysis: &'a TosaAnalysis<'a>,
337    value: ValueId,
338) -> Result<virtio_accel_tosa::Tensor<'a>, LoweringError> {
339    match analysis.value(value).kind() {
340        AnalyzedValueKind::Tensor(tensor) => Ok(tensor),
341        AnalyzedValueKind::Shape(_) => Err(LoweringError::UnsupportedGraph),
342    }
343}
344
345fn unary_operation(
346    kind: &str,
347    inputs: &[(&str, &str)],
348    output: &str,
349    dtype: u64,
350    shape: &[i32],
351) -> Vec<u8> {
352    let mut operation = Vec::new();
353    field_string(&mut operation, 1, kind);
354    for (key, name) in inputs {
355        field_message(&mut operation, 2, &argument_entry(key, name));
356    }
357    field_message(&mut operation, 3, &named_value_type(output, dtype, shape));
358    field_message(
359        &mut operation,
360        5,
361        &attribute_entry("name", &string_value(output)),
362    );
363    operation
364}
365
366fn const_string(name: &str, value: &str) -> Vec<u8> {
367    const_operation(name, MIL_STRING, &string_value(value))
368}
369
370fn const_int32(name: &str, value: i32) -> Vec<u8> {
371    const_operation(name, MIL_INT32, &int32_value(value))
372}
373
374fn const_bool(name: &str, value: bool) -> Vec<u8> {
375    const_operation(name, MIL_BOOL, &bool_value(value))
376}
377
378fn const_operation(name: &str, dtype: u64, value: &[u8]) -> Vec<u8> {
379    let mut operation = Vec::new();
380    field_string(&mut operation, 1, "const");
381    field_message(&mut operation, 3, &named_value_type(name, dtype, &[]));
382    field_message(&mut operation, 5, &attribute_entry("val", value));
383    field_message(
384        &mut operation,
385        5,
386        &attribute_entry("name", &string_value(name)),
387    );
388    operation
389}
390
391fn argument_entry(key: &str, name: &str) -> Vec<u8> {
392    let mut binding = Vec::new();
393    field_string(&mut binding, 1, name);
394    let mut argument = Vec::new();
395    field_message(&mut argument, 1, &binding);
396    map_entry(key, &argument)
397}
398
399fn attribute_entry(key: &str, value: &[u8]) -> Vec<u8> {
400    map_entry(key, value)
401}
402
403fn map_entry(key: &str, value: &[u8]) -> Vec<u8> {
404    let mut entry = Vec::new();
405    field_string(&mut entry, 1, key);
406    field_message(&mut entry, 2, value);
407    entry
408}
409
410fn named_value_type(name: &str, dtype: u64, shape: &[i32]) -> Vec<u8> {
411    let mut named = Vec::new();
412    field_string(&mut named, 1, name);
413    field_message(&mut named, 2, &value_type(dtype, shape));
414    named
415}
416
417fn value_type(dtype: u64, shape: &[i32]) -> Vec<u8> {
418    let mut tensor = Vec::new();
419    field_varint(&mut tensor, 1, dtype);
420    if !shape.is_empty() {
421        field_varint(&mut tensor, 2, shape.len() as u64);
422        for dimension in shape {
423            let mut constant = Vec::new();
424            field_varint(&mut constant, 1, *dimension as u64);
425            let mut dimension_message = Vec::new();
426            field_message(&mut dimension_message, 1, &constant);
427            field_message(&mut tensor, 3, &dimension_message);
428        }
429    }
430    let mut value_type = Vec::new();
431    field_message(&mut value_type, 1, &tensor);
432    value_type
433}
434
435fn string_value(value: &str) -> Vec<u8> {
436    let mut repeated = Vec::new();
437    field_string(&mut repeated, 1, value);
438    immediate_tensor_value(MIL_STRING, 4, &repeated)
439}
440
441fn int32_value(value: i32) -> Vec<u8> {
442    let mut packed = Vec::new();
443    varint(&mut packed, value as i64 as u64);
444    let mut repeated = Vec::new();
445    field_bytes(&mut repeated, 1, &packed);
446    immediate_tensor_value(MIL_INT32, 2, &repeated)
447}
448
449fn bool_value(value: bool) -> Vec<u8> {
450    let mut packed = Vec::new();
451    varint(&mut packed, value as u64);
452    let mut repeated = Vec::new();
453    field_bytes(&mut repeated, 1, &packed);
454    immediate_tensor_value(MIL_BOOL, 3, &repeated)
455}
456
457fn immediate_tensor_value(dtype: u64, tensor_field: u32, repeated: &[u8]) -> Vec<u8> {
458    let mut tensor_value = Vec::new();
459    field_message(&mut tensor_value, tensor_field, repeated);
460    let mut immediate = Vec::new();
461    field_message(&mut immediate, 1, &tensor_value);
462    let mut value = Vec::new();
463    field_message(&mut value, 2, &value_type(dtype, &[]));
464    field_message(&mut value, 3, &immediate);
465    value
466}
467
468fn field_varint(target: &mut Vec<u8>, field: u32, value: u64) {
469    varint(target, u64::from(field) << 3);
470    varint(target, value);
471}
472
473fn field_string(target: &mut Vec<u8>, field: u32, value: &str) {
474    field_bytes(target, field, value.as_bytes());
475}
476
477fn field_message(target: &mut Vec<u8>, field: u32, message: &[u8]) {
478    field_bytes(target, field, message);
479}
480
481fn field_bytes(target: &mut Vec<u8>, field: u32, bytes: &[u8]) {
482    varint(target, (u64::from(field) << 3) | 2);
483    varint(target, bytes.len() as u64);
484    target.extend_from_slice(bytes);
485}
486
487fn varint(target: &mut Vec<u8>, mut value: u64) {
488    while value >= 0x80 {
489        target.push((value as u8) | 0x80);
490        value >>= 7;
491    }
492    target.push(value as u8);
493}
494
495#[cfg(test)]
496mod tests {
497    use super::*;
498    use virtio_accel_conformance::numerics::{IDENTITY_INT8, MATMUL_INT8};
499
500    #[test]
501    fn lowers_shared_int8_identity_to_ml_program() {
502        let lowered =
503            lower_integer_tosa(IDENTITY_INT8.artifact, COREML_TOSA_INTEGER_TARGET).unwrap();
504        assert_eq!(lowered.features.len(), 2);
505        assert!(lowered.bytes.windows(7).any(|bytes| bytes == b"CoreML9"));
506        assert!(lowered.bytes.windows(4).any(|bytes| bytes == b"cast"));
507        assert!(lowered.bytes.windows(3).any(|bytes| bytes == b"add"));
508    }
509
510    #[test]
511    fn lowers_shared_int8_matmul_with_explicit_zero_points() {
512        let lowered = lower_integer_tosa(MATMUL_INT8.artifact, COREML_TOSA_INTEGER_TARGET).unwrap();
513        assert_eq!(lowered.features.len(), 3);
514        assert!(lowered.bytes.windows(6).any(|bytes| bytes == b"matmul"));
515        assert!(lowered.bytes.windows(3).any(|bytes| bytes == b"sub"));
516        // Negative lhs zero-point is encoded as a sign-extended protobuf int32 varint.
517        assert!(
518            lowered
519                .bytes
520                .windows(10)
521                .any(|bytes| bytes == [0xfe, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0x01])
522        );
523    }
524
525    #[test]
526    fn rejects_float_artifacts_at_the_integer_target() {
527        let error = lower_integer_tosa(
528            virtio_accel_conformance::numerics::MATMUL_FP32.artifact,
529            COREML_TOSA_INTEGER_TARGET,
530        )
531        .unwrap_err();
532        assert!(matches!(error, LoweringError::Analysis(_)));
533    }
534}