1#![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
21pub 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
45pub 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
62const 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 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 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}