1use core::fmt;
2
3use crate::{
4 AnalyzedValueKind, DType, LevelLimits, Op, OpAttributes, OperatorId, RuntimeCondition,
5 TosaAnalysis, ValueId,
6};
7
8#[derive(Clone, Copy, Debug)]
10pub struct RuntimeValue<'a> {
11 pub value: ValueId,
12 pub bytes: &'a [u8],
13}
14
15#[derive(Clone, Copy, Debug, PartialEq, Eq)]
17pub enum RuntimeErrorKind {
18 ValuesOutOfOrder,
19 MissingValue,
20 InvalidEncoding,
21 RequiredConditionFailed,
22}
23
24#[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#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
40pub struct RuntimeValidation {
41 pub unpredictable: bool,
43}
44
45pub 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}