Skip to main content

virtio_accel_tosa/
runtime.rs

1use core::fmt;
2
3use crate::{
4    AnalyzedValueKind, DType, LevelLimits, Op, OpAttributes, OperatorId, RuntimeCondition,
5    TosaAnalysis, ValueId,
6};
7
8/// One host-readable dynamic CTC value supplied for specialization.
9#[derive(Clone, Copy, Debug)]
10pub struct RuntimeValue<'a> {
11    pub value: ValueId,
12    pub bytes: &'a [u8],
13}
14
15/// Failure while resolving or checking dynamic CTC data.
16#[derive(Clone, Copy, Debug, PartialEq, Eq)]
17pub enum RuntimeErrorKind {
18    ValuesOutOfOrder,
19    MissingValue,
20    InvalidEncoding,
21    RequiredConditionFailed,
22}
23
24/// Located runtime-specialization failure.
25#[derive(Clone, Copy, Debug, PartialEq, Eq)]
26pub struct RuntimeError {
27    pub operator: Option<OperatorId>,
28    pub value: Option<ValueId>,
29    pub kind: RuntimeErrorKind,
30}
31
32impl fmt::Display for RuntimeError {
33    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
34        write!(formatter, "{self:?}")
35    }
36}
37
38/// Result of checking every dynamic CTC value needed by one specialization.
39#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
40pub struct RuntimeValidation {
41    /// At least one advisory TOSA `REQUIRE` condition failed.
42    pub unpredictable: bool,
43}
44
45/// Validate dynamic CTC encodings and every associated mandatory `ERROR_IF` condition.
46///
47/// `values` must be strictly sorted by [`ValueId`]. Only dynamic CTC values are inspected; normal
48/// tensor inputs stay on the provider's direct execution path. Per-element advisory conditions are
49/// represented in [`TosaAnalysis::conditions`] but deliberately are not scanned here.
50pub fn validate_runtime_values(
51    analysis: &TosaAnalysis<'_>,
52    values: &[RuntimeValue<'_>],
53) -> Result<RuntimeValidation, RuntimeError> {
54    if values.windows(2).any(|pair| pair[0].value >= pair[1].value) {
55        return Err(RuntimeError {
56            operator: None,
57            value: None,
58            kind: RuntimeErrorKind::ValuesOutOfOrder,
59        });
60    }
61
62    let mut validation = RuntimeValidation::default();
63    for operator in analysis.operators() {
64        let dynamic = analysis
65            .operator_conditions(operator.id())
66            .iter()
67            .any(|condition| matches!(condition, RuntimeCondition::DynamicCompileTimeInput { .. }));
68        if !dynamic {
69            continue;
70        }
71        for condition in analysis.operator_conditions(operator.id()) {
72            if let RuntimeCondition::DynamicCompileTimeInput { value, .. } = *condition {
73                let bytes = lookup(values, value).ok_or(RuntimeError {
74                    operator: Some(operator.id()),
75                    value: Some(value),
76                    kind: RuntimeErrorKind::MissingValue,
77                })?;
78                if !valid_encoding(analysis, value, bytes) {
79                    return Err(RuntimeError {
80                        operator: Some(operator.id()),
81                        value: Some(value),
82                        kind: RuntimeErrorKind::InvalidEncoding,
83                    });
84                }
85            }
86        }
87        match check_operator(analysis, values, operator.id())? {
88            ConditionResult::Valid => {}
89            ConditionResult::Unpredictable => validation.unpredictable = true,
90        }
91    }
92    Ok(validation)
93}
94
95#[derive(Clone, Copy, Debug, PartialEq, Eq)]
96enum ConditionResult {
97    Valid,
98    Unpredictable,
99}
100
101fn lookup<'a>(values: &'a [RuntimeValue<'_>], value: ValueId) -> Option<&'a [u8]> {
102    values
103        .binary_search_by_key(&value, |candidate| candidate.value)
104        .ok()
105        .map(|index| values[index].bytes)
106}
107
108fn data_for<'a>(
109    analysis: &'a TosaAnalysis<'_>,
110    values: &'a [RuntimeValue<'_>],
111    value: ValueId,
112) -> Option<&'a [u8]> {
113    analysis
114        .serialized_constant(value)
115        .or_else(|| lookup(values, value))
116}
117
118fn valid_encoding(analysis: &TosaAnalysis<'_>, value: ValueId, bytes: &[u8]) -> bool {
119    let analyzed = analysis.value(value);
120    let (dtype, elements) = match analyzed.kind() {
121        AnalyzedValueKind::Tensor(tensor) => {
122            let Some(_) = tensor.rank() else {
123                return false;
124            };
125            let Some(elements) = tensor.dimensions().try_fold(1_usize, |count, dimension| {
126                count.checked_mul(usize::try_from(dimension).ok()?)
127            }) else {
128                return false;
129            };
130            (tensor.dtype(), elements)
131        }
132        AnalyzedValueKind::Shape(shape) => (DType::SHAPE, shape.rank() as usize),
133    };
134    let expected = match dtype {
135        DType::INT4 => elements.div_ceil(2),
136        DType::BOOL | DType::INT8 | DType::FP8E4M3 | DType::FP8E5M2 => elements,
137        DType::INT16 | DType::FP16 | DType::BF16 => elements.saturating_mul(2),
138        DType::INT32 | DType::FP32 => elements.saturating_mul(4),
139        DType::INT48 | DType::SHAPE => elements.saturating_mul(8),
140        _ => return false,
141    };
142    if bytes.len() != expected {
143        return false;
144    }
145    match dtype {
146        DType::BOOL => bytes.iter().all(|value| *value <= 1),
147        DType::INT4 => (0..elements).all(|index| crate::unpack_int4(bytes, index) != Some(-8)),
148        DType::INT48 => (0..elements).all(|index| {
149            integer_at(dtype, bytes, index)
150                .is_some_and(|value| (-(1_i64 << 47)..(1_i64 << 47)).contains(&value))
151        }),
152        DType::SHAPE => {
153            let magnitude = 1_i128 << analysis.target().level.limits().max_log2_size;
154            (0..elements).all(|index| {
155                shape_at(bytes, index).is_some_and(|value| {
156                    i128::from(value) >= -magnitude && i128::from(value) < magnitude
157                })
158            })
159        }
160        _ => true,
161    }
162}
163
164fn check_operator(
165    analysis: &TosaAnalysis<'_>,
166    runtime: &[RuntimeValue<'_>],
167    operator: OperatorId,
168) -> Result<ConditionResult, RuntimeError> {
169    let plan = analysis.operator(operator);
170    let inputs = analysis.operator_inputs(operator);
171    let fail = |value| RuntimeError {
172        operator: Some(operator),
173        value: Some(value),
174        kind: RuntimeErrorKind::RequiredConditionFailed,
175    };
176    let data = |index: usize| {
177        data_for(analysis, runtime, inputs[index]).ok_or(RuntimeError {
178            operator: Some(operator),
179            value: Some(inputs[index]),
180            kind: RuntimeErrorKind::MissingValue,
181        })
182    };
183
184    match plan.op() {
185        Op::AVG_POOL2D => {
186            for index in [1, 2] {
187                if !zero_point_valid(analysis, inputs[index], data(index)?, false) {
188                    return Err(fail(inputs[index]));
189                }
190            }
191            Ok(ConditionResult::Valid)
192        }
193        Op::CONV2D | Op::CONV3D | Op::DEPTHWISE_CONV2D | Op::TRANSPOSE_CONV2D => {
194            for index in [3, 4] {
195                if !zero_point_valid(analysis, inputs[index], data(index)?, false) {
196                    return Err(fail(inputs[index]));
197                }
198            }
199            Ok(ConditionResult::Valid)
200        }
201        Op::MATMUL => {
202            for index in [2, 3] {
203                if !zero_point_valid(analysis, inputs[index], data(index)?, false) {
204                    return Err(fail(inputs[index]));
205                }
206            }
207            Ok(ConditionResult::Valid)
208        }
209        Op::MUL => {
210            let shift = integer_value(analysis, inputs[2], data(2)?, 0);
211            let dtype = tensor_dtype(analysis, inputs[0]);
212            if shift.is_some_and(|shift| {
213                (0..=63).contains(&shift) && (dtype == DType::INT32 || shift == 0)
214            }) {
215                Ok(ConditionResult::Valid)
216            } else {
217                Ok(ConditionResult::Unpredictable)
218            }
219        }
220        Op::TABLE => Ok(ConditionResult::Valid),
221        Op::NEGATE => {
222            for index in [1, 2] {
223                if !zero_point_valid(analysis, inputs[index], data(index)?, false) {
224                    return Err(fail(inputs[index]));
225                }
226            }
227            Ok(ConditionResult::Valid)
228        }
229        Op::PAD => {
230            if !zero_point_valid(analysis, inputs[2], data(2)?, false) {
231                return Err(fail(inputs[2]));
232            }
233            if !pad_valid(analysis, operator, inputs, data(1)?) {
234                return Err(fail(inputs[1]));
235            }
236            Ok(ConditionResult::Valid)
237        }
238        Op::RESHAPE => {
239            if reshape_valid(analysis, operator, inputs, data(1)?) {
240                Ok(ConditionResult::Valid)
241            } else {
242                Err(fail(inputs[1]))
243            }
244        }
245        Op::SLICE => {
246            if slice_valid(analysis, operator, inputs, data(1)?, data(2)?) {
247                Ok(ConditionResult::Valid)
248            } else {
249                Err(fail(inputs[1]))
250            }
251        }
252        Op::TILE => {
253            if tile_valid(analysis, operator, inputs, data(1)?) {
254                Ok(ConditionResult::Valid)
255            } else {
256                Err(fail(inputs[1]))
257            }
258        }
259        Op::RESIZE => {
260            if resize_valid(
261                analysis,
262                operator,
263                inputs,
264                data(1)?,
265                data(2)?,
266                data(3)?,
267                analysis.target().level.limits(),
268            ) {
269                Ok(ConditionResult::Valid)
270            } else {
271                Err(fail(inputs[1]))
272            }
273        }
274        Op::RESCALE => {
275            let OpAttributes::Rescale {
276                input_unsigned,
277                output_unsigned,
278                ..
279            } = plan.source().attributes()
280            else {
281                unreachable!()
282            };
283            if !zero_point_valid(analysis, inputs[3], data(3)?, input_unsigned) {
284                return Err(fail(inputs[3]));
285            }
286            if !zero_point_valid(analysis, inputs[4], data(4)?, output_unsigned) {
287                return Err(fail(inputs[4]));
288            }
289            let multipliers = data(1)?;
290            let shifts = data(2)?;
291            let count = tensor_elements(analysis, inputs[1]).unwrap_or(0);
292            if (0..count).all(|index| {
293                integer_value(analysis, inputs[1], multipliers, index)
294                    .is_some_and(|value| value >= 0)
295                    && integer_value(analysis, inputs[2], shifts, index)
296                        .is_some_and(|value| (2..=62).contains(&value))
297            }) {
298                Ok(ConditionResult::Valid)
299            } else {
300                Ok(ConditionResult::Unpredictable)
301            }
302        }
303        _ => Ok(ConditionResult::Valid),
304    }
305}
306
307fn zero_point_valid(
308    analysis: &TosaAnalysis<'_>,
309    value: ValueId,
310    bytes: &[u8],
311    unsigned: bool,
312) -> bool {
313    let dtype = tensor_dtype(analysis, value);
314    dtype == DType::INT8
315        || (dtype == DType::INT16
316            && unsigned
317            && bytes
318                .get(..2)
319                .and_then(|bytes| <[u8; 2]>::try_from(bytes).ok())
320                .is_some_and(|bytes| matches!(u16::from_le_bytes(bytes), 0 | 32_768)))
321        || value_is_zero(dtype, bytes, 0)
322}
323
324fn value_is_zero(dtype: DType, bytes: &[u8], index: usize) -> bool {
325    match dtype {
326        DType::INT4 | DType::INT8 | DType::INT16 | DType::INT32 | DType::INT48 => {
327            integer_at(dtype, bytes, index) == Some(0)
328        }
329        DType::FP8E4M3 | DType::FP8E5M2 => bytes.get(index).is_some_and(|bits| bits & 0x7f == 0),
330        DType::FP16 | DType::BF16 => bytes
331            .get(index * 2..index * 2 + 2)
332            .and_then(|bytes| <[u8; 2]>::try_from(bytes).ok())
333            .is_some_and(|bytes| u16::from_le_bytes(bytes) & 0x7fff == 0),
334        DType::FP32 => bytes
335            .get(index * 4..index * 4 + 4)
336            .and_then(|bytes| <[u8; 4]>::try_from(bytes).ok())
337            .is_some_and(|bytes| u32::from_le_bytes(bytes) & 0x7fff_ffff == 0),
338        _ => false,
339    }
340}
341
342fn integer_value(
343    analysis: &TosaAnalysis<'_>,
344    value: ValueId,
345    bytes: &[u8],
346    index: usize,
347) -> Option<i64> {
348    integer_at(tensor_dtype(analysis, value), bytes, index)
349}
350
351fn integer_at(dtype: DType, bytes: &[u8], index: usize) -> Option<i64> {
352    let width = match dtype {
353        DType::INT4 | DType::INT8 => 1,
354        DType::INT16 => 2,
355        DType::INT32 => 4,
356        DType::INT48 => 8,
357        _ => return None,
358    };
359    let start = index.checked_mul(width)?;
360    match dtype {
361        DType::INT4 => Some(i64::from(crate::unpack_int4(bytes, index)?)),
362        DType::INT8 => Some(i64::from(*bytes.get(start)? as i8)),
363        DType::INT16 => Some(i64::from(i16::from_le_bytes(
364            bytes.get(start..start + 2)?.try_into().ok()?,
365        ))),
366        DType::INT32 => Some(i64::from(i32::from_le_bytes(
367            bytes.get(start..start + 4)?.try_into().ok()?,
368        ))),
369        DType::INT48 => Some(i64::from_le_bytes(
370            bytes.get(start..start + 8)?.try_into().ok()?,
371        )),
372        _ => None,
373    }
374}
375
376fn shape_at(bytes: &[u8], index: usize) -> Option<i64> {
377    let start = index.checked_mul(8)?;
378    Some(i64::from_le_bytes(
379        bytes.get(start..start + 8)?.try_into().ok()?,
380    ))
381}
382
383fn tensor_dtype(analysis: &TosaAnalysis<'_>, value: ValueId) -> DType {
384    match analysis.value(value).kind() {
385        AnalyzedValueKind::Tensor(tensor) => tensor.dtype(),
386        AnalyzedValueKind::Shape(_) => DType::SHAPE,
387    }
388}
389
390fn tensor_elements(analysis: &TosaAnalysis<'_>, value: ValueId) -> Option<usize> {
391    match analysis.value(value).kind() {
392        AnalyzedValueKind::Tensor(tensor) => {
393            let _ = tensor.rank()?;
394            tensor.dimensions().try_fold(1_usize, |count, dimension| {
395                count.checked_mul(usize::try_from(dimension).ok()?)
396            })
397        }
398        AnalyzedValueKind::Shape(shape) => Some(shape.rank() as usize),
399    }
400}
401
402fn pad_valid(
403    analysis: &TosaAnalysis<'_>,
404    operator: OperatorId,
405    inputs: &[ValueId],
406    padding: &[u8],
407) -> bool {
408    let output = analysis.operator_outputs(operator)[0];
409    let Some(rank) = tensor_rank(analysis, inputs[0]) else {
410        return false;
411    };
412    tensor_rank(analysis, output) == Some(rank)
413        && (0..rank).all(|index| {
414            let (Some(input), Some(output)) = (
415                tensor_dimension(analysis, inputs[0], index),
416                tensor_dimension(analysis, output, index),
417            ) else {
418                return false;
419            };
420            let before = shape_at(padding, index * 2);
421            let after = shape_at(padding, index * 2 + 1);
422            before.is_some_and(|before| before >= 0)
423                && after.is_some_and(|after| after >= 0)
424                && before.zip(after).is_some_and(|(before, after)| {
425                    i128::from(input) + i128::from(before) + i128::from(after) == i128::from(output)
426                })
427        })
428}
429
430fn reshape_valid(
431    analysis: &TosaAnalysis<'_>,
432    operator: OperatorId,
433    inputs: &[ValueId],
434    shape: &[u8],
435) -> bool {
436    let output = analysis.operator_outputs(operator)[0];
437    let dimensions_match = tensor_rank(analysis, output).is_some_and(|rank| {
438        (0..rank).all(|index| {
439            tensor_dimension(analysis, output, index)
440                .is_some_and(|dimension| shape_at(shape, index) == Some(i64::from(dimension)))
441        })
442    });
443    dimensions_match && tensor_elements(analysis, inputs[0]) == tensor_elements(analysis, output)
444}
445
446fn slice_valid(
447    analysis: &TosaAnalysis<'_>,
448    operator: OperatorId,
449    inputs: &[ValueId],
450    starts: &[u8],
451    sizes: &[u8],
452) -> bool {
453    let output = analysis.operator_outputs(operator)[0];
454    let Some(rank) = tensor_rank(analysis, inputs[0]) else {
455        return false;
456    };
457    tensor_rank(analysis, output) == Some(rank)
458        && (0..rank).all(|index| {
459            let (Some(input), Some(output)) = (
460                tensor_dimension(analysis, inputs[0], index),
461                tensor_dimension(analysis, output, index),
462            ) else {
463                return false;
464            };
465            let start = shape_at(starts, index);
466            let size = shape_at(sizes, index);
467            start.is_some_and(|start| start >= 0)
468                && size.is_some_and(|size| size > 0)
469                && start.zip(size).is_some_and(|(start, size)| {
470                    i128::from(start) + i128::from(size) <= i128::from(input)
471                        && size == i64::from(output)
472                })
473        })
474}
475
476fn tile_valid(
477    analysis: &TosaAnalysis<'_>,
478    operator: OperatorId,
479    inputs: &[ValueId],
480    multiples: &[u8],
481) -> bool {
482    let output = analysis.operator_outputs(operator)[0];
483    let Some(rank) = tensor_rank(analysis, inputs[0]) else {
484        return false;
485    };
486    tensor_rank(analysis, output) == Some(rank)
487        && (0..rank).all(|index| {
488            let (Some(input), Some(output)) = (
489                tensor_dimension(analysis, inputs[0], index),
490                tensor_dimension(analysis, output, index),
491            ) else {
492                return false;
493            };
494            shape_at(multiples, index).is_some_and(|multiple| {
495                multiple >= 1 && i128::from(input) * i128::from(multiple) == i128::from(output)
496            })
497        })
498}
499
500fn resize_valid(
501    analysis: &TosaAnalysis<'_>,
502    operator: OperatorId,
503    inputs: &[ValueId],
504    scale: &[u8],
505    offset: &[u8],
506    border: &[u8],
507    limits: LevelLimits,
508) -> bool {
509    let Some(input) = dimensions4(analysis, inputs[0]) else {
510        return false;
511    };
512    let Some(output) = dimensions4(analysis, analysis.operator_outputs(operator)[0]) else {
513        return false;
514    };
515    let (Some(yn_), Some(yd), Some(xn), Some(xd), Some(oy), Some(ox), Some(by), Some(bx)) = (
516        shape_at(scale, 0).map(i128::from),
517        shape_at(scale, 1).map(i128::from),
518        shape_at(scale, 2).map(i128::from),
519        shape_at(scale, 3).map(i128::from),
520        shape_at(offset, 0).map(i128::from),
521        shape_at(offset, 1).map(i128::from),
522        shape_at(border, 0).map(i128::from),
523        shape_at(border, 1).map(i128::from),
524    ) else {
525        return false;
526    };
527    if yn_ <= 0
528        || yd <= 0
529        || xn <= 0
530        || xd <= 0
531        || yn_ > 2_048
532        || xn > 2_048
533        || yn_ > i128::from(limits.max_scale) * yd
534        || xn > i128::from(limits.max_scale) * xd
535        || yd >= 16 * yn_
536        || xd >= 16 * xn
537        || !(-yn_..16 * yn_).contains(&oy)
538        || !(-xn..16 * xn).contains(&ox)
539        || !(-16 * yn_..yn_).contains(&by)
540        || !(-16 * xn..xn).contains(&bx)
541        || [input[1], input[2], output[1], output[2]]
542            .into_iter()
543            .any(|value| value >= 16_384)
544    {
545        return false;
546    }
547    let height_numerator = (i128::from(input[1]) - 1) * yn_ - oy + by;
548    let width_numerator = (i128::from(input[2]) - 1) * xn - ox + bx;
549    height_numerator % yd == 0
550        && width_numerator % xd == 0
551        && output
552            == [
553                input[0],
554                i64::try_from(height_numerator / yd + 1).unwrap_or(i64::MIN),
555                i64::try_from(width_numerator / xd + 1).unwrap_or(i64::MIN),
556                input[3],
557            ]
558}
559
560fn dimensions4(analysis: &TosaAnalysis<'_>, value: ValueId) -> Option<[i64; 4]> {
561    (tensor_rank(analysis, value) == Some(4)).then(|| {
562        [
563            i64::from(tensor_dimension(analysis, value, 0).unwrap()),
564            i64::from(tensor_dimension(analysis, value, 1).unwrap()),
565            i64::from(tensor_dimension(analysis, value, 2).unwrap()),
566            i64::from(tensor_dimension(analysis, value, 3).unwrap()),
567        ]
568    })
569}
570
571fn tensor_rank(analysis: &TosaAnalysis<'_>, value: ValueId) -> Option<usize> {
572    match analysis.value(value).kind() {
573        AnalyzedValueKind::Tensor(tensor) => tensor.rank(),
574        AnalyzedValueKind::Shape(_) => None,
575    }
576}
577
578fn tensor_dimension(analysis: &TosaAnalysis<'_>, value: ValueId, index: usize) -> Option<i32> {
579    match analysis.value(value).kind() {
580        AnalyzedValueKind::Tensor(tensor) => tensor.dimension(index),
581        AnalyzedValueKind::Shape(_) => None,
582    }
583}