1#[derive(Clone, Copy, Debug, PartialEq, Eq)]
12pub struct Float16Tensor {
13 pub shape: &'static [usize],
15 pub bits: &'static [u16],
17}
18
19#[derive(Clone, Copy, Debug, PartialEq, Eq)]
21pub struct TosaFloat16Case {
22 pub name: &'static str,
24 pub artifact: &'static [u8],
26 pub inputs: &'static [Float16Tensor],
28 pub outputs: &'static [Float16Tensor],
30}
31
32impl TosaFloat16Case {
33 pub fn output_matches(&self, output: usize, actual: &[u16]) -> bool {
38 let Some(expected) = self.outputs.get(output).map(|tensor| tensor.bits) else {
39 return false;
40 };
41 expected.len() == actual.len()
42 && expected.iter().zip(actual).all(|(expected, actual)| {
43 if is_binary16_nan(*expected) {
44 is_binary16_nan(*actual)
45 } else {
46 expected == actual
47 }
48 })
49 }
50}
51
52#[derive(Clone, Copy, Debug, PartialEq, Eq)]
54pub struct Bfloat16Tensor {
55 pub shape: &'static [usize],
57 pub bits: &'static [u16],
59}
60
61#[derive(Clone, Copy, Debug, PartialEq, Eq)]
63pub struct TosaBfloat16Case {
64 pub name: &'static str,
66 pub artifact: &'static [u8],
68 pub inputs: &'static [Bfloat16Tensor],
70 pub outputs: &'static [Bfloat16Tensor],
72}
73
74impl TosaBfloat16Case {
75 pub fn output_matches(&self, output: usize, actual: &[u16]) -> bool {
77 let Some(expected) = self.outputs.get(output).map(|tensor| tensor.bits) else {
78 return false;
79 };
80 expected.len() == actual.len()
81 && expected.iter().zip(actual).all(|(expected, actual)| {
82 if is_bfloat16_nan(*expected) {
83 is_bfloat16_nan(*actual)
84 } else {
85 expected == actual
86 }
87 })
88 }
89}
90
91#[derive(Clone, Copy, Debug, PartialEq, Eq)]
93pub enum RawTensor {
94 Fp16(&'static [u16]),
96 Bool(&'static [u8]),
98 Int32(&'static [i32]),
100}
101
102impl RawTensor {
103 pub fn bytes(self) -> Vec<u8> {
105 match self {
106 Self::Fp16(values) => values
107 .iter()
108 .flat_map(|value| value.to_le_bytes())
109 .collect(),
110 Self::Bool(values) => values.to_vec(),
111 Self::Int32(values) => values
112 .iter()
113 .flat_map(|value| value.to_le_bytes())
114 .collect(),
115 }
116 }
117
118 pub fn byte_len(self) -> usize {
120 match self {
121 Self::Fp16(values) => values.len() * 2,
122 Self::Bool(values) => values.len(),
123 Self::Int32(values) => values.len() * 4,
124 }
125 }
126}
127
128#[derive(Clone, Copy, Debug, PartialEq, Eq)]
130pub struct TosaRawCase {
131 pub name: &'static str,
133 pub artifact: &'static [u8],
135 pub inputs: &'static [RawTensor],
137 pub output: RawTensor,
139 pub fp16_max_ulps: u16,
141}
142
143impl TosaRawCase {
144 pub fn output_matches(self, actual: &[u8]) -> bool {
146 if actual.len() != self.output.byte_len() {
147 return false;
148 }
149 match self.output {
150 RawTensor::Bool(expected) => actual == expected,
151 RawTensor::Int32(expected) => actual
152 .chunks_exact(4)
153 .map(|bytes| i32::from_le_bytes(bytes.try_into().expect("four-byte chunk")))
154 .eq(expected.iter().copied()),
155 RawTensor::Fp16(expected) => {
156 actual
157 .chunks_exact(2)
158 .zip(expected)
159 .all(|(bytes, expected)| {
160 let actual = u16::from_le_bytes([bytes[0], bytes[1]]);
161 if is_binary16_nan(*expected) {
162 is_binary16_nan(actual)
163 } else if self.fp16_max_ulps == 0
164 || (*expected & 0x8000) != (actual & 0x8000)
165 || (*expected & 0x7fff) == 0
166 {
167 actual == *expected
168 } else {
169 actual.abs_diff(*expected) <= self.fp16_max_ulps
170 }
171 })
172 }
173 }
174 }
175}
176
177const fn is_binary16_nan(bits: u16) -> bool {
178 bits & 0x7c00 == 0x7c00 && bits & 0x03ff != 0
179}
180
181const fn is_bfloat16_nan(bits: u16) -> bool {
182 bits & 0x7f80 == 0x7f80 && bits & 0x007f != 0
183}
184
185#[derive(Clone, Copy, Debug, PartialEq, Eq)]
187pub enum PackedDType {
188 Int4,
190 Int8,
192 Fp8E4M3,
194 Fp8E5M2,
196}
197
198impl PackedDType {
199 pub const fn storage_bytes(self, elements: usize) -> Option<usize> {
201 match self {
202 Self::Int4 => Some(elements / 2 + elements % 2),
203 Self::Int8 | Self::Fp8E4M3 | Self::Fp8E5M2 => Some(elements),
204 }
205 }
206}
207
208#[derive(Clone, Copy, Debug, PartialEq, Eq)]
210pub struct PackedTensor {
211 pub shape: &'static [usize],
213 pub bytes: &'static [u8],
215}
216
217#[derive(Clone, Copy, Debug, PartialEq, Eq)]
219pub struct TosaFp8ToBfloat16Case {
220 pub name: &'static str,
222 pub input_dtype: PackedDType,
224 pub artifact: &'static [u8],
226 pub input: PackedTensor,
228 pub output: Bfloat16Tensor,
230}
231
232impl TosaFp8ToBfloat16Case {
233 pub fn output_matches(self, actual: &[u16]) -> bool {
235 self.output.bits.len() == actual.len()
236 && self
237 .output
238 .bits
239 .iter()
240 .zip(actual)
241 .all(|(expected, actual)| {
242 if is_bfloat16_nan(*expected) {
243 is_bfloat16_nan(*actual)
244 } else {
245 expected == actual
246 }
247 })
248 }
249}
250
251#[derive(Clone, Copy, Debug, PartialEq, Eq)]
253pub struct Int32Tensor {
254 pub shape: &'static [usize],
256 pub values: &'static [i32],
258}
259
260#[derive(Clone, Copy, Debug, PartialEq, Eq)]
262pub struct TosaInt8MatmulCase {
263 pub name: &'static str,
265 pub artifact: &'static [u8],
267 pub inputs: &'static [PackedTensor],
269 pub zero_points: [i8; 2],
271 pub outputs: &'static [Int32Tensor],
273}
274
275impl TosaInt8MatmulCase {
276 pub fn output_matches(&self, output: usize, actual: &[i32]) -> bool {
278 self.outputs
279 .get(output)
280 .is_some_and(|expected| expected.values == actual)
281 }
282}
283
284#[derive(Clone, Copy, Debug, PartialEq, Eq)]
286pub struct TosaInt32ToInt8RescaleCase {
287 pub name: &'static str,
289 pub artifact: &'static [u8],
291 pub input: Int32Tensor,
293 pub multiplier: i32,
295 pub shift: i8,
297 pub output_zero_point: i8,
299 pub output: PackedTensor,
301}
302
303impl TosaInt32ToInt8RescaleCase {
304 pub fn output_matches(self, actual: &[u8]) -> bool {
306 self.output.bytes == actual
307 }
308}
309
310#[derive(Clone, Copy, Debug, PartialEq, Eq)]
312pub struct TosaPackedCase {
313 pub name: &'static str,
315 pub dtype: PackedDType,
317 pub artifact: &'static [u8],
319 pub inputs: &'static [PackedTensor],
321 pub outputs: &'static [PackedTensor],
323}
324
325impl TosaPackedCase {
326 pub fn output_matches(&self, output: usize, actual: &[u8]) -> bool {
331 let Some(expected) = self.outputs.get(output).map(|tensor| tensor.bytes) else {
332 return false;
333 };
334 expected.len() == actual.len()
335 && expected.iter().zip(actual).all(|(expected, actual)| {
336 if packed_is_nan(self.dtype, *expected) {
337 packed_is_nan(self.dtype, *actual)
338 } else {
339 expected == actual
340 }
341 })
342 }
343}
344
345const fn packed_is_nan(dtype: PackedDType, bits: u8) -> bool {
346 match dtype {
347 PackedDType::Fp8E4M3 => bits & 0x7f == 0x7f,
348 PackedDType::Fp8E5M2 => bits & 0x7c == 0x7c && bits & 0x03 != 0,
349 PackedDType::Int4 | PackedDType::Int8 => false,
350 }
351}
352
353#[derive(Clone, Copy, Debug, PartialEq)]
355pub struct Float32Tensor {
356 pub shape: &'static [usize],
358 pub values: &'static [f32],
360}
361
362#[derive(Clone, Copy, Debug, PartialEq)]
364pub struct TosaFloat32Case {
365 pub name: &'static str,
367 pub artifact: &'static [u8],
369 pub inputs: &'static [Float32Tensor],
371 pub outputs: &'static [Float32Tensor],
373 pub absolute_tolerance: f32,
375 pub relative_tolerance: f32,
377}
378
379impl TosaFloat32Case {
380 pub fn output_matches(&self, output: usize, actual: &[f32]) -> bool {
382 let Some(expected) = self.outputs.get(output).map(|tensor| tensor.values) else {
383 return false;
384 };
385 expected.len() == actual.len()
386 && expected.iter().zip(actual).all(|(expected, actual)| {
387 if expected.is_nan() {
388 actual.is_nan()
389 } else if expected.is_infinite() || *expected == 0.0 {
390 expected.to_bits() == actual.to_bits()
391 } else {
392 let difference = (expected - actual).abs();
393 difference <= self.absolute_tolerance
394 || difference <= self.relative_tolerance * expected.abs()
395 }
396 })
397 }
398}
399
400const MATMUL_LHS: &[f32] = &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
401const MATMUL_RHS: &[f32] = &[7.0, 8.0, 9.0, 10.0, 11.0, 12.0];
402const MATMUL_OUTPUT: &[f32] = &[58.0, 64.0, 139.0, 154.0];
403const MATMUL_INPUTS: &[Float32Tensor] = &[
404 Float32Tensor {
405 shape: &[1, 2, 3],
406 values: MATMUL_LHS,
407 },
408 Float32Tensor {
409 shape: &[1, 3, 2],
410 values: MATMUL_RHS,
411 },
412];
413const MATMUL_OUTPUTS: &[Float32Tensor] = &[Float32Tensor {
414 shape: &[1, 2, 2],
415 values: MATMUL_OUTPUT,
416}];
417
418pub const MATMUL_FP32: TosaFloat32Case = TosaFloat32Case {
420 name: "matmul-fp32",
421 artifact: include_bytes!("data/matmul-fp32-v1.0.0.tosa"),
422 inputs: MATMUL_INPUTS,
423 outputs: MATMUL_OUTPUTS,
424 absolute_tolerance: 1.0e-5,
425 relative_tolerance: 1.0e-5,
426};
427
428const MAX_POOL2D_INPUT: &[f32] = &[
429 1.0, 101.0, 2.0, 102.0, 3.0, 103.0, 4.0, 104.0, 5.0, 105.0, 6.0, 106.0, 7.0, 107.0, 8.0, 108.0,
430 9.0, 109.0, 10.0, 110.0, 11.0, 111.0, 12.0, 112.0, 13.0, 113.0, 14.0, 114.0, 15.0, 115.0, 16.0,
431 116.0,
432];
433const MAX_POOL2D_OUTPUT: &[f32] = &[6.0, 106.0, 8.0, 108.0, 14.0, 114.0, 16.0, 116.0];
434const MAX_POOL2D_INPUTS: &[Float32Tensor] = &[Float32Tensor {
435 shape: &[1, 4, 4, 2],
436 values: MAX_POOL2D_INPUT,
437}];
438const MAX_POOL2D_OUTPUTS: &[Float32Tensor] = &[Float32Tensor {
439 shape: &[1, 2, 2, 2],
440 values: MAX_POOL2D_OUTPUT,
441}];
442
443pub const MAX_POOL2D_FP32: TosaFloat32Case = TosaFloat32Case {
445 name: "max-pool2d-fp32",
446 artifact: include_bytes!("data/max-pool2d-fp32-v1.0.0.tosa"),
447 inputs: MAX_POOL2D_INPUTS,
448 outputs: MAX_POOL2D_OUTPUTS,
449 absolute_tolerance: 0.0,
450 relative_tolerance: 0.0,
451};
452
453const MAX_POOL2D_INPUT_BF16_BITS: &[u16] = &[
454 0x3f80, 0x42ca, 0x4000, 0x42cc, 0x4040, 0x42ce, 0x4080, 0x42d0, 0x40a0, 0x42d2, 0x40c0, 0x42d4,
455 0x40e0, 0x42d6, 0x4100, 0x42d8, 0x4110, 0x42da, 0x4120, 0x42dc, 0x4130, 0x42de, 0x4140, 0x42e0,
456 0x4150, 0x42e2, 0x4160, 0x42e4, 0x4170, 0x42e6, 0x4180, 0x42e8,
457];
458const MAX_POOL2D_OUTPUT_BF16_BITS: &[u16] = &[
459 0x40c0, 0x42d4, 0x4100, 0x42d8, 0x4160, 0x42e4, 0x4180, 0x42e8,
460];
461const MAX_POOL2D_INPUTS_BF16: &[Bfloat16Tensor] = &[Bfloat16Tensor {
462 shape: &[1, 4, 4, 2],
463 bits: MAX_POOL2D_INPUT_BF16_BITS,
464}];
465const MAX_POOL2D_OUTPUTS_BF16: &[Bfloat16Tensor] = &[Bfloat16Tensor {
466 shape: &[1, 2, 2, 2],
467 bits: MAX_POOL2D_OUTPUT_BF16_BITS,
468}];
469
470pub const MAX_POOL2D_BF16: TosaBfloat16Case = TosaBfloat16Case {
473 name: "max-pool2d-bf16",
474 artifact: include_bytes!("data/max-pool2d-bf16-v1.0.0.tosa"),
475 inputs: MAX_POOL2D_INPUTS_BF16,
476 outputs: MAX_POOL2D_OUTPUTS_BF16,
477};
478
479const IDENTITY_EDGE_VALUES: &[f32] = &[
480 f32::NAN,
481 f32::NEG_INFINITY,
482 -0.0,
483 0.0,
484 f32::from_bits(1),
485 f32::MIN_POSITIVE,
486 1.0,
487 f32::INFINITY,
488];
489const IDENTITY_EDGE_INPUTS: &[Float32Tensor] = &[Float32Tensor {
490 shape: &[8],
491 values: IDENTITY_EDGE_VALUES,
492}];
493const IDENTITY_EDGE_OUTPUTS: &[Float32Tensor] = IDENTITY_EDGE_INPUTS;
494
495pub const IDENTITY_EDGES_FP32: TosaFloat32Case = TosaFloat32Case {
497 name: "identity-edges-fp32",
498 artifact: include_bytes!("data/identity-edges-fp32-v1.0.0.tosa"),
499 inputs: IDENTITY_EDGE_INPUTS,
500 outputs: IDENTITY_EDGE_OUTPUTS,
501 absolute_tolerance: 0.0,
502 relative_tolerance: 0.0,
503};
504
505const MATMUL_LHS_FP16_BITS: &[u16] = &[0x3c00, 0x4000, 0x4200, 0x4400, 0x4500, 0x4600];
506const MATMUL_RHS_FP16_BITS: &[u16] = &[0x4700, 0x4800, 0x4880, 0x4900, 0x4980, 0x4a00];
507const MATMUL_OUTPUT_FP16_BITS: &[u16] = &[0x5340, 0x5400, 0x5858, 0x58d0];
508const MATMUL_INPUTS_FP16: &[Float16Tensor] = &[
509 Float16Tensor {
510 shape: &[1, 2, 3],
511 bits: MATMUL_LHS_FP16_BITS,
512 },
513 Float16Tensor {
514 shape: &[1, 3, 2],
515 bits: MATMUL_RHS_FP16_BITS,
516 },
517];
518const MATMUL_OUTPUTS_FP16: &[Float16Tensor] = &[Float16Tensor {
519 shape: &[1, 2, 2],
520 bits: MATMUL_OUTPUT_FP16_BITS,
521}];
522
523pub const MATMUL_FP16: TosaFloat16Case = TosaFloat16Case {
525 name: "matmul-fp16",
526 artifact: include_bytes!("data/matmul-fp16-v1.0.0.tosa"),
527 inputs: MATMUL_INPUTS_FP16,
528 outputs: MATMUL_OUTPUTS_FP16,
529};
530
531const BINARY_LEFT_FP16_BITS: &[u16] = &[0x4000, 0x4400];
532const BINARY_RIGHT_FP16_BITS: &[u16] = &[0x3c00, 0x4000, 0x4200];
533const BINARY_INPUTS_FP16: &[Float16Tensor] = &[
534 Float16Tensor {
535 shape: &[2, 1],
536 bits: BINARY_LEFT_FP16_BITS,
537 },
538 Float16Tensor {
539 shape: &[1, 3],
540 bits: BINARY_RIGHT_FP16_BITS,
541 },
542];
543const ADD_OUTPUT_FP16_BITS: &[u16] = &[0x4200, 0x4400, 0x4500, 0x4500, 0x4600, 0x4700];
544const SUB_OUTPUT_FP16_BITS: &[u16] = &[0x3c00, 0x0000, 0xbc00, 0x4200, 0x4000, 0x3c00];
545const MUL_OUTPUT_FP16_BITS: &[u16] = &[0x4000, 0x4400, 0x4600, 0x4400, 0x4800, 0x4a00];
546const POW_OUTPUT_FP16_BITS: &[u16] = &[0x4000, 0x4400, 0x4800, 0x4400, 0x4c00, 0x5400];
547const MAXIMUM_OUTPUT_FP16_BITS: &[u16] = &[0x4000, 0x4000, 0x4200, 0x4400, 0x4400, 0x4400];
548const MINIMUM_OUTPUT_FP16_BITS: &[u16] = &[0x3c00, 0x4000, 0x4000, 0x3c00, 0x4000, 0x4200];
549
550const fn binary_output(bits: &'static [u16]) -> [Float16Tensor; 1] {
551 [Float16Tensor {
552 shape: &[2, 3],
553 bits,
554 }]
555}
556
557pub const ADD_FP16: TosaFloat16Case = TosaFloat16Case {
559 name: "add-fp16",
560 artifact: include_bytes!("data/add-fp16-v1.0.0.tosa"),
561 inputs: BINARY_INPUTS_FP16,
562 outputs: &binary_output(ADD_OUTPUT_FP16_BITS),
563};
564
565pub const SUB_FP16: TosaFloat16Case = TosaFloat16Case {
567 name: "sub-fp16",
568 artifact: include_bytes!("data/sub-fp16-v1.0.0.tosa"),
569 inputs: BINARY_INPUTS_FP16,
570 outputs: &binary_output(SUB_OUTPUT_FP16_BITS),
571};
572
573pub const MUL_FP16: TosaFloat16Case = TosaFloat16Case {
575 name: "mul-fp16",
576 artifact: include_bytes!("data/mul-fp16-v1.0.0.tosa"),
577 inputs: BINARY_INPUTS_FP16,
578 outputs: &binary_output(MUL_OUTPUT_FP16_BITS),
579};
580
581pub const POW_FP16: TosaFloat16Case = TosaFloat16Case {
583 name: "pow-fp16",
584 artifact: include_bytes!("data/pow-fp16-v1.0.0.tosa"),
585 inputs: BINARY_INPUTS_FP16,
586 outputs: &binary_output(POW_OUTPUT_FP16_BITS),
587};
588
589pub const MAXIMUM_FP16: TosaFloat16Case = TosaFloat16Case {
591 name: "maximum-fp16",
592 artifact: include_bytes!("data/maximum-fp16-v1.0.0.tosa"),
593 inputs: BINARY_INPUTS_FP16,
594 outputs: &binary_output(MAXIMUM_OUTPUT_FP16_BITS),
595};
596
597pub const MINIMUM_FP16: TosaFloat16Case = TosaFloat16Case {
599 name: "minimum-fp16",
600 artifact: include_bytes!("data/minimum-fp16-v1.0.0.tosa"),
601 inputs: BINARY_INPUTS_FP16,
602 outputs: &binary_output(MINIMUM_OUTPUT_FP16_BITS),
603};
604
605const UNARY_INPUTS_RAW: &[RawTensor] = &[RawTensor::Fp16(&[0x3800, 0x3c00, 0x4000, 0x4400])];
606
607const fn unary_raw_case(
608 name: &'static str,
609 artifact: &'static [u8],
610 output: &'static [u16],
611 fp16_max_ulps: u16,
612) -> TosaRawCase {
613 TosaRawCase {
614 name,
615 artifact,
616 inputs: UNARY_INPUTS_RAW,
617 output: RawTensor::Fp16(output),
618 fp16_max_ulps,
619 }
620}
621
622pub const HEXAGON_UNARY_FP16_CASES: &[TosaRawCase] = &[
624 unary_raw_case(
625 "abs-fp16",
626 include_bytes!("data/abs-fp16-v1.0.0.tosa"),
627 &[0x3800, 0x3c00, 0x4000, 0x4400],
628 0,
629 ),
630 unary_raw_case(
631 "ceil-fp16",
632 include_bytes!("data/ceil-fp16-v1.0.0.tosa"),
633 &[0x3c00, 0x3c00, 0x4000, 0x4400],
634 0,
635 ),
636 unary_raw_case(
637 "cos-fp16",
638 include_bytes!("data/cos-fp16-v1.0.0.tosa"),
639 &[0x3b05, 0x3853, 0xb6a9, 0xb93b],
640 8,
641 ),
642 unary_raw_case(
643 "exp-fp16",
644 include_bytes!("data/exp-fp16-v1.0.0.tosa"),
645 &[0x3e98, 0x4170, 0x4764, 0x52d3],
646 4,
647 ),
648 unary_raw_case(
649 "floor-fp16",
650 include_bytes!("data/floor-fp16-v1.0.0.tosa"),
651 &[0x0000, 0x3c00, 0x4000, 0x4400],
652 0,
653 ),
654 unary_raw_case(
655 "log-fp16",
656 include_bytes!("data/log-fp16-v1.0.0.tosa"),
657 &[0xb98c, 0x0000, 0x398c, 0x3d8c],
658 4,
659 ),
660 unary_raw_case(
661 "negate-fp16",
662 include_bytes!("data/negate-fp16-v1.0.0.tosa"),
663 &[0xb800, 0xbc00, 0xc000, 0xc400],
664 0,
665 ),
666 unary_raw_case(
667 "reciprocal-fp16",
668 include_bytes!("data/reciprocal-fp16-v1.0.0.tosa"),
669 &[0x4000, 0x3c00, 0x3800, 0x3400],
670 2,
671 ),
672 unary_raw_case(
673 "rsqrt-fp16",
674 include_bytes!("data/rsqrt-fp16-v1.0.0.tosa"),
675 &[0x3da8, 0x3c00, 0x39a8, 0x3800],
676 4,
677 ),
678 unary_raw_case(
679 "sin-fp16",
680 include_bytes!("data/sin-fp16-v1.0.0.tosa"),
681 &[0x37ac, 0x3abb, 0x3b46, 0xba0e],
682 8,
683 ),
684 unary_raw_case(
685 "sigmoid-fp16",
686 include_bytes!("data/sigmoid-fp16-v1.0.0.tosa"),
687 &[0x38fb, 0x39d9, 0x3b0c, 0x3bdb],
688 4,
689 ),
690 unary_raw_case(
691 "tanh-fp16",
692 include_bytes!("data/tanh-fp16-v1.0.0.tosa"),
693 &[0x3765, 0x3a18, 0x3bb6, 0x3bff],
694 4,
695 ),
696 unary_raw_case(
697 "clamp-fp16",
698 include_bytes!("data/clamp-fp16-v1.0.0.tosa"),
699 &[0x3800, 0x3c00, 0x3c00, 0x3c00],
700 0,
701 ),
702];
703
704const COMPARISON_INPUTS_RAW: &[RawTensor] = &[
705 RawTensor::Fp16(&[0x3c00, 0x4000, 0x4200, 0x4400]),
706 RawTensor::Fp16(&[0x3c00, 0x4200, 0x4000, 0x4400]),
707];
708const LOGICAL_INPUTS_RAW: &[RawTensor] = &[
709 RawTensor::Bool(&[0, 0, 1, 1]),
710 RawTensor::Bool(&[0, 1, 0, 1]),
711];
712const LOGICAL_NOT_INPUT_RAW: &[RawTensor] = &[RawTensor::Bool(&[0, 0, 1, 1])];
713const SELECT_INPUTS_RAW: &[RawTensor] = &[
714 RawTensor::Bool(&[0, 1, 0, 1]),
715 RawTensor::Fp16(&[0x3c00, 0x4000, 0x4200, 0x4400]),
716 RawTensor::Fp16(&[0x4500, 0x4600, 0x4700, 0x4800]),
717];
718
719pub const HEXAGON_LOGICAL_CASES: &[TosaRawCase] = &[
721 TosaRawCase {
722 name: "equal-fp16",
723 artifact: include_bytes!("data/equal-fp16-v1.0.0.tosa"),
724 inputs: COMPARISON_INPUTS_RAW,
725 output: RawTensor::Bool(&[1, 0, 0, 1]),
726 fp16_max_ulps: 0,
727 },
728 TosaRawCase {
729 name: "greater-fp16",
730 artifact: include_bytes!("data/greater-fp16-v1.0.0.tosa"),
731 inputs: COMPARISON_INPUTS_RAW,
732 output: RawTensor::Bool(&[0, 0, 1, 0]),
733 fp16_max_ulps: 0,
734 },
735 TosaRawCase {
736 name: "greater-equal-fp16",
737 artifact: include_bytes!("data/greater-equal-fp16-v1.0.0.tosa"),
738 inputs: COMPARISON_INPUTS_RAW,
739 output: RawTensor::Bool(&[1, 0, 1, 1]),
740 fp16_max_ulps: 0,
741 },
742 TosaRawCase {
743 name: "logical-and",
744 artifact: include_bytes!("data/logical-and-fp16-v1.0.0.tosa"),
745 inputs: LOGICAL_INPUTS_RAW,
746 output: RawTensor::Bool(&[0, 0, 0, 1]),
747 fp16_max_ulps: 0,
748 },
749 TosaRawCase {
750 name: "logical-or",
751 artifact: include_bytes!("data/logical-or-fp16-v1.0.0.tosa"),
752 inputs: LOGICAL_INPUTS_RAW,
753 output: RawTensor::Bool(&[0, 1, 1, 1]),
754 fp16_max_ulps: 0,
755 },
756 TosaRawCase {
757 name: "logical-xor",
758 artifact: include_bytes!("data/logical-xor-fp16-v1.0.0.tosa"),
759 inputs: LOGICAL_INPUTS_RAW,
760 output: RawTensor::Bool(&[0, 1, 1, 0]),
761 fp16_max_ulps: 0,
762 },
763 TosaRawCase {
764 name: "logical-not",
765 artifact: include_bytes!("data/logical-not-fp16-v1.0.0.tosa"),
766 inputs: LOGICAL_NOT_INPUT_RAW,
767 output: RawTensor::Bool(&[1, 1, 0, 0]),
768 fp16_max_ulps: 0,
769 },
770 TosaRawCase {
771 name: "select-fp16",
772 artifact: include_bytes!("data/select-fp16-v1.0.0.tosa"),
773 inputs: SELECT_INPUTS_RAW,
774 output: RawTensor::Fp16(&[0x4500, 0x4000, 0x4700, 0x4400]),
775 fp16_max_ulps: 0,
776 },
777];
778
779const REDUCTION_INPUTS_RAW: &[RawTensor] = &[RawTensor::Fp16(&[
780 0x3c00, 0x4200, 0x4000, 0xbc00, 0x4400, 0x4000,
781])];
782
783pub const HEXAGON_REDUCTION_CASES: &[TosaRawCase] = &[
785 TosaRawCase {
786 name: "argmax-fp16",
787 artifact: include_bytes!("data/argmax-fp16-v1.0.0.tosa"),
788 inputs: REDUCTION_INPUTS_RAW,
789 output: RawTensor::Int32(&[1, 1]),
790 fp16_max_ulps: 0,
791 },
792 TosaRawCase {
793 name: "reduce-max-fp16",
794 artifact: include_bytes!("data/reduce-max-fp16-v1.0.0.tosa"),
795 inputs: REDUCTION_INPUTS_RAW,
796 output: RawTensor::Fp16(&[0x4200, 0x4400]),
797 fp16_max_ulps: 0,
798 },
799 TosaRawCase {
800 name: "reduce-min-fp16",
801 artifact: include_bytes!("data/reduce-min-fp16-v1.0.0.tosa"),
802 inputs: REDUCTION_INPUTS_RAW,
803 output: RawTensor::Fp16(&[0x3c00, 0xbc00]),
804 fp16_max_ulps: 0,
805 },
806 TosaRawCase {
807 name: "reduce-product-fp16",
808 artifact: include_bytes!("data/reduce-product-fp16-v1.0.0.tosa"),
809 inputs: REDUCTION_INPUTS_RAW,
810 output: RawTensor::Fp16(&[0x4600, 0xc800]),
811 fp16_max_ulps: 1,
812 },
813 TosaRawCase {
814 name: "reduce-sum-fp16",
815 artifact: include_bytes!("data/reduce-sum-fp16-v1.0.0.tosa"),
816 inputs: REDUCTION_INPUTS_RAW,
817 output: RawTensor::Fp16(&[0x4600, 0x4500]),
818 fp16_max_ulps: 1,
819 },
820];
821
822pub const HEXAGON_MOVEMENT_CASES: &[TosaRawCase] = &[
824 TosaRawCase {
825 name: "const-add-fp16",
826 artifact: include_bytes!("data/const-fp16-v1.0.0.tosa"),
827 inputs: &[RawTensor::Fp16(&[0x4900, 0x4d00, 0x4f80, 0x5100])],
828 output: RawTensor::Fp16(&[0x4980, 0x4d80, 0x5020, 0x5180]),
829 fp16_max_ulps: 0,
830 },
831 TosaRawCase {
832 name: "reshape-const-shape-fp16",
833 artifact: include_bytes!("data/reshape-fp16-v1.0.0.tosa"),
834 inputs: &[RawTensor::Fp16(&[0x3c00, 0x4000, 0x4200, 0x4400])],
835 output: RawTensor::Fp16(&[0x3c00, 0x4000, 0x4200, 0x4400]),
836 fp16_max_ulps: 0,
837 },
838 TosaRawCase {
839 name: "transpose-fp16",
840 artifact: include_bytes!("data/transpose-fp16-v1.0.0.tosa"),
841 inputs: &[RawTensor::Fp16(&[
842 0x3c00, 0x4000, 0x4200, 0x4400, 0x4500, 0x4600,
843 ])],
844 output: RawTensor::Fp16(&[0x3c00, 0x4400, 0x4000, 0x4500, 0x4200, 0x4600]),
845 fp16_max_ulps: 0,
846 },
847 TosaRawCase {
848 name: "reverse-fp16",
849 artifact: include_bytes!("data/reverse-fp16-v1.0.0.tosa"),
850 inputs: &[RawTensor::Fp16(&[
851 0x3c00, 0x4000, 0x4200, 0x4400, 0x4500, 0x4600,
852 ])],
853 output: RawTensor::Fp16(&[0x4200, 0x4000, 0x3c00, 0x4600, 0x4500, 0x4400]),
854 fp16_max_ulps: 0,
855 },
856 TosaRawCase {
857 name: "concat-fp16",
858 artifact: include_bytes!("data/concat-fp16-v1.0.0.tosa"),
859 inputs: &[
860 RawTensor::Fp16(&[0x3c00, 0x4000]),
861 RawTensor::Fp16(&[0x4200, 0x4400]),
862 ],
863 output: RawTensor::Fp16(&[0x3c00, 0x4200, 0x4000, 0x4400]),
864 fp16_max_ulps: 0,
865 },
866];
867
868const MOCK_CLASSIFIER_FEATURES_FP16_BITS: &[u16] = &[
869 0x3c00, 0x4000, 0x4200, 0xbc00, 0x3800, 0x4000, ];
872const MOCK_CLASSIFIER_WEIGHTS_FP16_BITS: &[u16] = &[
873 0x3c00, 0x0000, 0x0000, 0x3c00, 0x3c00, 0xbc00, ];
877const MOCK_CLASSIFIER_LOGITS_FP16_BITS: &[u16] = &[
878 0x4400, 0xbc00, 0x3c00, 0xbe00, ];
881const MOCK_CLASSIFIER_INPUTS_FP16: &[Float16Tensor] = &[
882 Float16Tensor {
883 shape: &[1, 2, 3],
884 bits: MOCK_CLASSIFIER_FEATURES_FP16_BITS,
885 },
886 Float16Tensor {
887 shape: &[1, 3, 2],
888 bits: MOCK_CLASSIFIER_WEIGHTS_FP16_BITS,
889 },
890];
891const MOCK_CLASSIFIER_OUTPUTS_FP16: &[Float16Tensor] = &[Float16Tensor {
892 shape: &[1, 2, 2],
893 bits: MOCK_CLASSIFIER_LOGITS_FP16_BITS,
894}];
895
896pub const MOCK_LINEAR_CLASSIFIER_FP16: TosaFloat16Case = TosaFloat16Case {
898 name: "mock-linear-classifier-fp16",
899 artifact: MATMUL_FP16.artifact,
900 inputs: MOCK_CLASSIFIER_INPUTS_FP16,
901 outputs: MOCK_CLASSIFIER_OUTPUTS_FP16,
902};
903
904const MAX_POOL2D_INPUT_FP16_BITS: &[u16] = &[
905 0x3c00, 0x5650, 0x4000, 0x5660, 0x4200, 0x5670, 0x4400, 0x5680, 0x4500, 0x5690, 0x4600, 0x56a0,
906 0x4700, 0x56b0, 0x4800, 0x56c0, 0x4880, 0x56d0, 0x4900, 0x56e0, 0x4980, 0x56f0, 0x4a00, 0x5700,
907 0x4a80, 0x5710, 0x4b00, 0x5720, 0x4b80, 0x5730, 0x4c00, 0x5740,
908];
909const MAX_POOL2D_OUTPUT_FP16_BITS: &[u16] = &[
910 0x4600, 0x56a0, 0x4800, 0x56c0, 0x4b00, 0x5720, 0x4c00, 0x5740,
911];
912const MAX_POOL2D_INPUTS_FP16: &[Float16Tensor] = &[Float16Tensor {
913 shape: &[1, 4, 4, 2],
914 bits: MAX_POOL2D_INPUT_FP16_BITS,
915}];
916const MAX_POOL2D_OUTPUTS_FP16: &[Float16Tensor] = &[Float16Tensor {
917 shape: &[1, 2, 2, 2],
918 bits: MAX_POOL2D_OUTPUT_FP16_BITS,
919}];
920
921pub const MAX_POOL2D_FP16: TosaFloat16Case = TosaFloat16Case {
923 name: "max-pool2d-fp16",
924 artifact: include_bytes!("data/max-pool2d-fp16-v1.0.0.tosa"),
925 inputs: MAX_POOL2D_INPUTS_FP16,
926 outputs: MAX_POOL2D_OUTPUTS_FP16,
927};
928
929const IDENTITY_EDGE_FP16_BITS: &[u16] = &[
930 0x7e00, 0xfc00, 0x8000, 0x0000, 0x0001, 0x0400, 0x3c00, 0x7c00,
931];
932const IDENTITY_EDGE_INPUTS_FP16: &[Float16Tensor] = &[Float16Tensor {
933 shape: &[8],
934 bits: IDENTITY_EDGE_FP16_BITS,
935}];
936const IDENTITY_EDGE_OUTPUTS_FP16: &[Float16Tensor] = IDENTITY_EDGE_INPUTS_FP16;
937
938pub const IDENTITY_EDGES_FP16: TosaFloat16Case = TosaFloat16Case {
940 name: "identity-edges-fp16",
941 artifact: include_bytes!("data/identity-edges-fp16-v1.0.0.tosa"),
942 inputs: IDENTITY_EDGE_INPUTS_FP16,
943 outputs: IDENTITY_EDGE_OUTPUTS_FP16,
944};
945
946const IDENTITY_INT8_BYTES: &[u8] = &[0x80, 0x81, 0xff, 0x00, 0x01, 0x7e, 0x7f, 0x2a];
947const IDENTITY_INT8_TENSORS: &[PackedTensor] = &[PackedTensor {
948 shape: &[8],
949 bytes: IDENTITY_INT8_BYTES,
950}];
951
952pub const IDENTITY_INT8: TosaPackedCase = TosaPackedCase {
954 name: "identity-int8",
955 dtype: PackedDType::Int8,
956 artifact: include_bytes!("data/identity-int8-v1.0.0.tosa"),
957 inputs: IDENTITY_INT8_TENSORS,
958 outputs: IDENTITY_INT8_TENSORS,
959};
960
961const MATMUL_INT8_LHS: &[u8] = &[0x80, 0xff, 0x7f, 0x05, 0xfa, 0x07];
962const MATMUL_INT8_RHS: &[u8] = &[0x08, 0xf7, 0x0a, 0x0b, 0x0c, 0xf3];
963const MATMUL_INT8_INPUTS: &[PackedTensor] = &[
964 PackedTensor {
965 shape: &[1, 2, 3],
966 bytes: MATMUL_INT8_LHS,
967 },
968 PackedTensor {
969 shape: &[1, 3, 2],
970 bytes: MATMUL_INT8_RHS,
971 },
972];
973const MATMUL_INT8_OUTPUT: &[i32] = &[538, -544, 88, -260];
974const MATMUL_INT8_OUTPUTS: &[Int32Tensor] = &[Int32Tensor {
975 shape: &[1, 2, 2],
976 values: MATMUL_INT8_OUTPUT,
977}];
978
979pub const MATMUL_INT8: TosaInt8MatmulCase = TosaInt8MatmulCase {
981 name: "matmul-int8",
982 artifact: include_bytes!("data/matmul-int8-v1.0.0.tosa"),
983 inputs: MATMUL_INT8_INPUTS,
984 zero_points: [-2, 3],
985 outputs: MATMUL_INT8_OUTPUTS,
986};
987
988const CLASSIFIER_SAMPLES_INT8: &[u8] = &[0x4E, 0x08, 0x03, 0x08, 0x4E, 0x03];
992const CLASSIFIER_WEIGHTS_INT8: &[u8] = &[0x07, 0x03, 0x03, 0x07, 0x02, 0x02];
993const CLASSIFIER_INPUTS_INT8: &[PackedTensor] = &[
994 PackedTensor {
995 shape: &[1, 2, 3],
996 bytes: CLASSIFIER_SAMPLES_INT8,
997 },
998 PackedTensor {
999 shape: &[1, 3, 2],
1000 bytes: CLASSIFIER_WEIGHTS_INT8,
1001 },
1002];
1003const CLASSIFIER_LOGITS_INT32: &[i32] = &[315, 35, 35, 315];
1004const CLASSIFIER_OUTPUTS_INT32: &[Int32Tensor] = &[Int32Tensor {
1005 shape: &[1, 2, 2],
1006 values: CLASSIFIER_LOGITS_INT32,
1007}];
1008
1009pub const QUANTIZED_CLASSIFIER_INT8: TosaInt8MatmulCase = TosaInt8MatmulCase {
1012 name: "quantized-classifier-int8",
1013 artifact: MATMUL_INT8.artifact,
1014 inputs: CLASSIFIER_INPUTS_INT8,
1015 zero_points: [-2, 3],
1016 outputs: CLASSIFIER_OUTPUTS_INT32,
1017};
1018
1019const RESCALE_INT32_INPUT: &[i32] = &[
1020 -1000, -251, -250, -249, -3, -2, -1, 0, 1, 2, 3, 249, 250, 251, 260, 1000,
1021];
1022const RESCALE_INT8_OUTPUT: &[u8] = &[
1023 0x80, 0x80, 0x80, 0x81, 0xfc, 0xfc, 0xfd, 0xfd, 0xfe, 0xfe, 0xff, 0x7a, 0x7a, 0x7b, 0x7f, 0x7f,
1024];
1025
1026pub const RESCALE_INT32_TO_INT8: TosaInt32ToInt8RescaleCase = TosaInt32ToInt8RescaleCase {
1028 name: "rescale-int32-to-int8",
1029 artifact: include_bytes!("data/rescale-int32-to-int8-v1.0.0.tosa"),
1030 input: Int32Tensor {
1031 shape: &[16],
1032 values: RESCALE_INT32_INPUT,
1033 },
1034 multiplier: 1 << 29,
1035 shift: 30,
1036 output_zero_point: -3,
1037 output: PackedTensor {
1038 shape: &[16],
1039 bytes: RESCALE_INT8_OUTPUT,
1040 },
1041};
1042
1043const IDENTITY_INT4_BYTES: &[u8] = &[0xd9, 0x0f, 0x31, 0x76];
1045const IDENTITY_INT4_TENSORS: &[PackedTensor] = &[PackedTensor {
1046 shape: &[8],
1047 bytes: IDENTITY_INT4_BYTES,
1048}];
1049
1050pub const IDENTITY_INT4: TosaPackedCase = TosaPackedCase {
1052 name: "identity-int4",
1053 dtype: PackedDType::Int4,
1054 artifact: include_bytes!("data/identity-int4-v1.0.0.tosa"),
1055 inputs: IDENTITY_INT4_TENSORS,
1056 outputs: IDENTITY_INT4_TENSORS,
1057};
1058
1059const IDENTITY_FP8E4M3_BYTES: &[u8] = &[0x00, 0x80, 0x01, 0x81, 0x38, 0xb8, 0x7e, 0x7f];
1060const IDENTITY_FP8E4M3_TENSORS: &[PackedTensor] = &[PackedTensor {
1061 shape: &[8],
1062 bytes: IDENTITY_FP8E4M3_BYTES,
1063}];
1064
1065pub const IDENTITY_FP8E4M3: TosaPackedCase = TosaPackedCase {
1067 name: "identity-fp8e4m3",
1068 dtype: PackedDType::Fp8E4M3,
1069 artifact: include_bytes!("data/identity-fp8e4m3-v1.0.0.tosa"),
1070 inputs: IDENTITY_FP8E4M3_TENSORS,
1071 outputs: IDENTITY_FP8E4M3_TENSORS,
1072};
1073
1074const IDENTITY_FP8E5M2_BYTES: &[u8] = &[0x00, 0x80, 0x01, 0x81, 0x3c, 0x7b, 0x7c, 0x7d];
1075const IDENTITY_FP8E5M2_TENSORS: &[PackedTensor] = &[PackedTensor {
1076 shape: &[8],
1077 bytes: IDENTITY_FP8E5M2_BYTES,
1078}];
1079
1080pub const IDENTITY_FP8E5M2: TosaPackedCase = TosaPackedCase {
1082 name: "identity-fp8e5m2",
1083 dtype: PackedDType::Fp8E5M2,
1084 artifact: include_bytes!("data/identity-fp8e5m2-v1.0.0.tosa"),
1085 inputs: IDENTITY_FP8E5M2_TENSORS,
1086 outputs: IDENTITY_FP8E5M2_TENSORS,
1087};
1088
1089const fn all_fp8_encodings() -> [u8; 1024] {
1090 let mut values = [0u8; 1024];
1091 let mut index = 0;
1092 while index < values.len() {
1093 values[index] = index as u8;
1094 index += 1;
1095 }
1096 values
1097}
1098
1099const fn fp8e4m3_bf16_oracle() -> [u16; 1024] {
1100 let mut values = [0u16; 1024];
1101 let mut index = 0;
1102 while index < values.len() {
1103 let bits = index as u8;
1104 let sign = ((bits & 0x80) as u16) << 8;
1105 let exponent = ((bits >> 3) & 0x0f) as u16;
1106 let fraction = (bits & 0x07) as u16;
1107 values[index] = if exponent == 0 {
1108 let subnormal = [
1109 0x0000, 0x3b00, 0x3b80, 0x3bc0, 0x3c00, 0x3c20, 0x3c40, 0x3c60,
1110 ];
1111 sign | subnormal[fraction as usize]
1112 } else if exponent == 0x0f && fraction == 0x07 {
1113 sign | 0x7fc0
1114 } else {
1115 sign | ((exponent + 120) << 7) | (fraction << 4)
1116 };
1117 index += 1;
1118 }
1119 values
1120}
1121
1122const fn fp8e5m2_bf16_oracle() -> [u16; 1024] {
1123 let mut values = [0u16; 1024];
1124 let mut index = 0;
1125 while index < values.len() {
1126 let bits = index as u8;
1127 let sign = ((bits & 0x80) as u16) << 8;
1128 let exponent = ((bits >> 2) & 0x1f) as u16;
1129 let fraction = (bits & 0x03) as u16;
1130 values[index] = if exponent == 0 {
1131 let subnormal = [0x0000, 0x3780, 0x3800, 0x3840];
1132 sign | subnormal[fraction as usize]
1133 } else if exponent == 0x1f {
1134 sign | if fraction == 0 { 0x7f80 } else { 0x7fc0 }
1135 } else {
1136 sign | ((exponent + 112) << 7) | (fraction << 5)
1137 };
1138 index += 1;
1139 }
1140 values
1141}
1142
1143const ALL_FP8_ENCODINGS: [u8; 1024] = all_fp8_encodings();
1144const CAST_FP8E4M3_OUTPUT: [u16; 1024] = fp8e4m3_bf16_oracle();
1145
1146pub const CAST_FP8E4M3_TO_BF16: TosaFp8ToBfloat16Case = TosaFp8ToBfloat16Case {
1149 name: "cast-fp8e4m3-to-bf16",
1150 input_dtype: PackedDType::Fp8E4M3,
1151 artifact: include_bytes!("data/cast-fp8e4m3-to-bf16-v1.0.0.tosa"),
1152 input: PackedTensor {
1153 shape: &[1024],
1154 bytes: &ALL_FP8_ENCODINGS,
1155 },
1156 output: Bfloat16Tensor {
1157 shape: &[1024],
1158 bits: &CAST_FP8E4M3_OUTPUT,
1159 },
1160};
1161
1162const CAST_FP8E5M2_OUTPUT: [u16; 1024] = fp8e5m2_bf16_oracle();
1163
1164pub const CAST_FP8E5M2_TO_BF16: TosaFp8ToBfloat16Case = TosaFp8ToBfloat16Case {
1167 name: "cast-fp8e5m2-to-bf16",
1168 input_dtype: PackedDType::Fp8E5M2,
1169 artifact: include_bytes!("data/cast-fp8e5m2-to-bf16-v1.0.0.tosa"),
1170 input: PackedTensor {
1171 shape: &[1024],
1172 bytes: &ALL_FP8_ENCODINGS,
1173 },
1174 output: Bfloat16Tensor {
1175 shape: &[1024],
1176 bits: &CAST_FP8E5M2_OUTPUT,
1177 },
1178};
1179
1180#[derive(Clone, Copy, Debug, PartialEq)]
1182pub enum Fp32TierTensor {
1183 Fp32(&'static [f32]),
1185 Bool(&'static [u8]),
1187 Int32(&'static [i32]),
1189}
1190
1191impl Fp32TierTensor {
1192 pub fn bytes(self) -> Vec<u8> {
1194 match self {
1195 Self::Fp32(values) => values
1196 .iter()
1197 .flat_map(|value| value.to_le_bytes())
1198 .collect(),
1199 Self::Bool(values) => values.to_vec(),
1200 Self::Int32(values) => values
1201 .iter()
1202 .flat_map(|value| value.to_le_bytes())
1203 .collect(),
1204 }
1205 }
1206
1207 pub fn byte_len(self) -> usize {
1209 match self {
1210 Self::Fp32(values) => values.len() * 4,
1211 Self::Bool(values) => values.len(),
1212 Self::Int32(values) => values.len() * 4,
1213 }
1214 }
1215}
1216
1217#[derive(Clone, Copy, Debug, PartialEq)]
1224pub struct TosaFp32OperatorCase {
1225 pub name: &'static str,
1227 pub artifact: &'static [u8],
1229 pub inputs: &'static [Fp32TierTensor],
1231 pub output: Fp32TierTensor,
1233 pub absolute_tolerance: f32,
1235 pub relative_tolerance: f32,
1237}
1238
1239impl TosaFp32OperatorCase {
1240 pub fn output_matches(self, actual: &[u8]) -> bool {
1242 if actual.len() != self.output.byte_len() {
1243 return false;
1244 }
1245 match self.output {
1246 Fp32TierTensor::Bool(expected) => actual == expected,
1247 Fp32TierTensor::Int32(expected) => actual
1248 .chunks_exact(4)
1249 .map(|bytes| i32::from_le_bytes(bytes.try_into().expect("four-byte chunk")))
1250 .eq(expected.iter().copied()),
1251 Fp32TierTensor::Fp32(expected) => {
1252 actual
1253 .chunks_exact(4)
1254 .zip(expected)
1255 .all(|(bytes, expected)| {
1256 let actual = f32::from_le_bytes(bytes.try_into().expect("four-byte chunk"));
1257 if expected.is_nan() {
1258 actual.is_nan()
1259 } else if expected.is_infinite() || *expected == 0.0 {
1260 expected.to_bits() == actual.to_bits()
1261 } else {
1262 let difference = (expected - actual).abs();
1263 difference <= self.absolute_tolerance
1264 || difference <= self.relative_tolerance * expected.abs()
1265 }
1266 })
1267 }
1268 }
1269 }
1270}
1271
1272const UNARY_INPUTS_FP32: &[Fp32TierTensor] = &[Fp32TierTensor::Fp32(&[0.5, 1.0, 2.0, 4.0])];
1273
1274const EXACT: (f32, f32) = (0.0, 0.0);
1277const TRANSCENDENTAL: (f32, f32) = (1.0e-6, 4.0e-6);
1283
1284const fn unary_fp32_case(
1285 name: &'static str,
1286 artifact: &'static [u8],
1287 output: &'static [f32],
1288 tolerance: (f32, f32),
1289) -> TosaFp32OperatorCase {
1290 TosaFp32OperatorCase {
1291 name,
1292 artifact,
1293 inputs: UNARY_INPUTS_FP32,
1294 output: Fp32TierTensor::Fp32(output),
1295 absolute_tolerance: tolerance.0,
1296 relative_tolerance: tolerance.1,
1297 }
1298}
1299
1300pub const FP32_UNARY_CASES: &[TosaFp32OperatorCase] = &[
1302 unary_fp32_case(
1303 "abs-fp32",
1304 include_bytes!("data/abs-fp32-v1.0.0.tosa"),
1305 &[0.5, 1.0, 2.0, 4.0],
1306 EXACT,
1307 ),
1308 unary_fp32_case(
1309 "ceil-fp32",
1310 include_bytes!("data/ceil-fp32-v1.0.0.tosa"),
1311 &[1.0, 1.0, 2.0, 4.0],
1312 EXACT,
1313 ),
1314 unary_fp32_case(
1315 "cos-fp32",
1316 include_bytes!("data/cos-fp32-v1.0.0.tosa"),
1317 &[0.877_582_55, 0.540_302_3, -0.416_146_84, -0.653_643_6],
1318 TRANSCENDENTAL,
1319 ),
1320 unary_fp32_case(
1321 "erf-fp32",
1322 include_bytes!("data/erf-fp32-v1.0.0.tosa"),
1323 &[0.520_499_9, 0.842_700_8, 0.995_322_3, 1.0],
1324 TRANSCENDENTAL,
1325 ),
1326 unary_fp32_case(
1327 "exp-fp32",
1328 include_bytes!("data/exp-fp32-v1.0.0.tosa"),
1329 &[1.648_721_2, 2.718_281_7, 7.389_056, 54.598_15],
1330 TRANSCENDENTAL,
1331 ),
1332 unary_fp32_case(
1333 "floor-fp32",
1334 include_bytes!("data/floor-fp32-v1.0.0.tosa"),
1335 &[0.0, 1.0, 2.0, 4.0],
1336 EXACT,
1337 ),
1338 unary_fp32_case(
1339 "log-fp32",
1340 include_bytes!("data/log-fp32-v1.0.0.tosa"),
1341 &[
1342 -core::f32::consts::LN_2,
1343 0.0,
1344 core::f32::consts::LN_2,
1345 2.0 * core::f32::consts::LN_2,
1346 ],
1347 TRANSCENDENTAL,
1348 ),
1349 unary_fp32_case(
1350 "negate-fp32",
1351 include_bytes!("data/negate-fp32-v1.0.0.tosa"),
1352 &[-0.5, -1.0, -2.0, -4.0],
1353 EXACT,
1354 ),
1355 unary_fp32_case(
1356 "reciprocal-fp32",
1357 include_bytes!("data/reciprocal-fp32-v1.0.0.tosa"),
1358 &[2.0, 1.0, 0.5, 0.25],
1359 TRANSCENDENTAL,
1360 ),
1361 unary_fp32_case(
1362 "rsqrt-fp32",
1363 include_bytes!("data/rsqrt-fp32-v1.0.0.tosa"),
1364 &[
1365 core::f32::consts::SQRT_2,
1366 1.0,
1367 core::f32::consts::FRAC_1_SQRT_2,
1368 0.5,
1369 ],
1370 TRANSCENDENTAL,
1371 ),
1372 unary_fp32_case(
1373 "sin-fp32",
1374 include_bytes!("data/sin-fp32-v1.0.0.tosa"),
1375 &[0.479_425_55, 0.841_470_96, 0.909_297_4, -0.756_802_5],
1376 TRANSCENDENTAL,
1377 ),
1378 unary_fp32_case(
1379 "sigmoid-fp32",
1380 include_bytes!("data/sigmoid-fp32-v1.0.0.tosa"),
1381 &[0.622_459_35, 0.731_058_6, 0.880_797_1, 0.982_013_76],
1382 TRANSCENDENTAL,
1383 ),
1384 unary_fp32_case(
1385 "tanh-fp32",
1386 include_bytes!("data/tanh-fp32-v1.0.0.tosa"),
1387 &[0.462_117_16, 0.761_594_2, 0.964_027_6, 0.999_329_3],
1388 TRANSCENDENTAL,
1389 ),
1390 unary_fp32_case(
1391 "clamp-fp32",
1392 include_bytes!("data/clamp-fp32-v1.0.0.tosa"),
1393 &[0.5, 1.0, 1.0, 1.0],
1394 EXACT,
1395 ),
1396];
1397
1398const BINARY_INPUTS_FP32: &[Fp32TierTensor] = &[
1399 Fp32TierTensor::Fp32(&[2.0, 4.0]),
1400 Fp32TierTensor::Fp32(&[1.0, 2.0, 3.0]),
1401];
1402
1403const fn binary_fp32_case(
1404 name: &'static str,
1405 artifact: &'static [u8],
1406 output: &'static [f32],
1407 tolerance: (f32, f32),
1408) -> TosaFp32OperatorCase {
1409 TosaFp32OperatorCase {
1410 name,
1411 artifact,
1412 inputs: BINARY_INPUTS_FP32,
1413 output: Fp32TierTensor::Fp32(output),
1414 absolute_tolerance: tolerance.0,
1415 relative_tolerance: tolerance.1,
1416 }
1417}
1418
1419pub const FP32_BINARY_CASES: &[TosaFp32OperatorCase] = &[
1421 binary_fp32_case(
1422 "add-fp32",
1423 include_bytes!("data/add-fp32-v1.0.0.tosa"),
1424 &[3.0, 4.0, 5.0, 5.0, 6.0, 7.0],
1425 EXACT,
1426 ),
1427 binary_fp32_case(
1428 "sub-fp32",
1429 include_bytes!("data/sub-fp32-v1.0.0.tosa"),
1430 &[1.0, 0.0, -1.0, 3.0, 2.0, 1.0],
1431 EXACT,
1432 ),
1433 binary_fp32_case(
1434 "mul-fp32",
1435 include_bytes!("data/mul-fp32-v1.0.0.tosa"),
1436 &[2.0, 4.0, 6.0, 4.0, 8.0, 12.0],
1437 EXACT,
1438 ),
1439 binary_fp32_case(
1440 "pow-fp32",
1441 include_bytes!("data/pow-fp32-v1.0.0.tosa"),
1442 &[2.0, 4.0, 8.0, 4.0, 16.0, 64.0],
1443 TRANSCENDENTAL,
1444 ),
1445 binary_fp32_case(
1446 "maximum-fp32",
1447 include_bytes!("data/maximum-fp32-v1.0.0.tosa"),
1448 &[2.0, 2.0, 3.0, 4.0, 4.0, 4.0],
1449 EXACT,
1450 ),
1451 binary_fp32_case(
1452 "minimum-fp32",
1453 include_bytes!("data/minimum-fp32-v1.0.0.tosa"),
1454 &[1.0, 2.0, 2.0, 1.0, 2.0, 3.0],
1455 EXACT,
1456 ),
1457];
1458
1459const COMPARISON_INPUTS_FP32: &[Fp32TierTensor] = &[
1460 Fp32TierTensor::Fp32(&[1.0, 2.0, 3.0, 4.0]),
1461 Fp32TierTensor::Fp32(&[1.0, 3.0, 2.0, 4.0]),
1462];
1463const LOGICAL_INPUTS_FP32_TIER: &[Fp32TierTensor] = &[
1464 Fp32TierTensor::Bool(&[0, 0, 1, 1]),
1465 Fp32TierTensor::Bool(&[0, 1, 0, 1]),
1466];
1467const LOGICAL_NOT_INPUT_FP32_TIER: &[Fp32TierTensor] = &[Fp32TierTensor::Bool(&[0, 0, 1, 1])];
1468const SELECT_INPUTS_FP32: &[Fp32TierTensor] = &[
1469 Fp32TierTensor::Bool(&[0, 1, 0, 1]),
1470 Fp32TierTensor::Fp32(&[1.0, 2.0, 3.0, 4.0]),
1471 Fp32TierTensor::Fp32(&[5.0, 6.0, 7.0, 8.0]),
1472];
1473
1474const fn exact_fp32_case(
1475 name: &'static str,
1476 artifact: &'static [u8],
1477 inputs: &'static [Fp32TierTensor],
1478 output: Fp32TierTensor,
1479) -> TosaFp32OperatorCase {
1480 TosaFp32OperatorCase {
1481 name,
1482 artifact,
1483 inputs,
1484 output,
1485 absolute_tolerance: 0.0,
1486 relative_tolerance: 0.0,
1487 }
1488}
1489
1490pub const FP32_LOGICAL_CASES: &[TosaFp32OperatorCase] = &[
1492 exact_fp32_case(
1493 "equal-fp32",
1494 include_bytes!("data/equal-fp32-v1.0.0.tosa"),
1495 COMPARISON_INPUTS_FP32,
1496 Fp32TierTensor::Bool(&[1, 0, 0, 1]),
1497 ),
1498 exact_fp32_case(
1499 "greater-fp32",
1500 include_bytes!("data/greater-fp32-v1.0.0.tosa"),
1501 COMPARISON_INPUTS_FP32,
1502 Fp32TierTensor::Bool(&[0, 0, 1, 0]),
1503 ),
1504 exact_fp32_case(
1505 "greater-equal-fp32",
1506 include_bytes!("data/greater-equal-fp32-v1.0.0.tosa"),
1507 COMPARISON_INPUTS_FP32,
1508 Fp32TierTensor::Bool(&[1, 0, 1, 1]),
1509 ),
1510 exact_fp32_case(
1511 "logical-and-fp32-tier",
1512 include_bytes!("data/logical-and-fp16-v1.0.0.tosa"),
1513 LOGICAL_INPUTS_FP32_TIER,
1514 Fp32TierTensor::Bool(&[0, 0, 0, 1]),
1515 ),
1516 exact_fp32_case(
1517 "logical-or-fp32-tier",
1518 include_bytes!("data/logical-or-fp16-v1.0.0.tosa"),
1519 LOGICAL_INPUTS_FP32_TIER,
1520 Fp32TierTensor::Bool(&[0, 1, 1, 1]),
1521 ),
1522 exact_fp32_case(
1523 "logical-xor-fp32-tier",
1524 include_bytes!("data/logical-xor-fp16-v1.0.0.tosa"),
1525 LOGICAL_INPUTS_FP32_TIER,
1526 Fp32TierTensor::Bool(&[0, 1, 1, 0]),
1527 ),
1528 exact_fp32_case(
1529 "logical-not-fp32-tier",
1530 include_bytes!("data/logical-not-fp16-v1.0.0.tosa"),
1531 LOGICAL_NOT_INPUT_FP32_TIER,
1532 Fp32TierTensor::Bool(&[1, 1, 0, 0]),
1533 ),
1534 exact_fp32_case(
1535 "select-fp32",
1536 include_bytes!("data/select-fp32-v1.0.0.tosa"),
1537 SELECT_INPUTS_FP32,
1538 Fp32TierTensor::Fp32(&[5.0, 2.0, 7.0, 4.0]),
1539 ),
1540];
1541
1542const REDUCTION_INPUTS_FP32: &[Fp32TierTensor] =
1543 &[Fp32TierTensor::Fp32(&[1.0, 3.0, 2.0, -1.0, 4.0, 2.0])];
1544
1545pub const FP32_REDUCTION_CASES: &[TosaFp32OperatorCase] = &[
1547 exact_fp32_case(
1548 "argmax-fp32",
1549 include_bytes!("data/argmax-fp32-v1.0.0.tosa"),
1550 REDUCTION_INPUTS_FP32,
1551 Fp32TierTensor::Int32(&[1, 1]),
1552 ),
1553 exact_fp32_case(
1554 "reduce-max-fp32",
1555 include_bytes!("data/reduce-max-fp32-v1.0.0.tosa"),
1556 REDUCTION_INPUTS_FP32,
1557 Fp32TierTensor::Fp32(&[3.0, 4.0]),
1558 ),
1559 exact_fp32_case(
1560 "reduce-min-fp32",
1561 include_bytes!("data/reduce-min-fp32-v1.0.0.tosa"),
1562 REDUCTION_INPUTS_FP32,
1563 Fp32TierTensor::Fp32(&[1.0, -1.0]),
1564 ),
1565 exact_fp32_case(
1566 "reduce-product-fp32",
1567 include_bytes!("data/reduce-product-fp32-v1.0.0.tosa"),
1568 REDUCTION_INPUTS_FP32,
1569 Fp32TierTensor::Fp32(&[6.0, -8.0]),
1570 ),
1571 exact_fp32_case(
1572 "reduce-sum-fp32",
1573 include_bytes!("data/reduce-sum-fp32-v1.0.0.tosa"),
1574 REDUCTION_INPUTS_FP32,
1575 Fp32TierTensor::Fp32(&[6.0, 5.0]),
1576 ),
1577];
1578
1579pub const FP32_MOVEMENT_CASES: &[TosaFp32OperatorCase] = &[
1581 exact_fp32_case(
1582 "const-add-fp32",
1583 include_bytes!("data/const-fp32-v1.0.0.tosa"),
1584 &[Fp32TierTensor::Fp32(&[10.0, 20.0, 30.0, 40.0])],
1585 Fp32TierTensor::Fp32(&[11.0, 22.0, 33.0, 44.0]),
1586 ),
1587 exact_fp32_case(
1588 "reshape-const-shape-fp32",
1589 include_bytes!("data/reshape-fp32-v1.0.0.tosa"),
1590 &[Fp32TierTensor::Fp32(&[1.0, 2.0, 3.0, 4.0])],
1591 Fp32TierTensor::Fp32(&[1.0, 2.0, 3.0, 4.0]),
1592 ),
1593 exact_fp32_case(
1594 "transpose-fp32",
1595 include_bytes!("data/transpose-fp32-v1.0.0.tosa"),
1596 &[Fp32TierTensor::Fp32(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0])],
1597 Fp32TierTensor::Fp32(&[1.0, 4.0, 2.0, 5.0, 3.0, 6.0]),
1598 ),
1599 exact_fp32_case(
1600 "reverse-fp32",
1601 include_bytes!("data/reverse-fp32-v1.0.0.tosa"),
1602 &[Fp32TierTensor::Fp32(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0])],
1603 Fp32TierTensor::Fp32(&[3.0, 2.0, 1.0, 6.0, 5.0, 4.0]),
1604 ),
1605 exact_fp32_case(
1606 "concat-fp32",
1607 include_bytes!("data/concat-fp32-v1.0.0.tosa"),
1608 &[
1609 Fp32TierTensor::Fp32(&[1.0, 2.0]),
1610 Fp32TierTensor::Fp32(&[3.0, 4.0]),
1611 ],
1612 Fp32TierTensor::Fp32(&[1.0, 3.0, 2.0, 4.0]),
1613 ),
1614];
1615
1616pub const LINEAR_TANH_FP32: TosaFp32OperatorCase = TosaFp32OperatorCase {
1620 name: "linear-tanh-fp32",
1621 artifact: include_bytes!("data/linear-tanh-fp32-v1.0.0.tosa"),
1622 inputs: &[
1623 Fp32TierTensor::Fp32(&[1.0, 2.0, 3.0, -1.0, 0.5, 2.0]),
1624 Fp32TierTensor::Fp32(&[1.0, 0.0, 0.0, 1.0, 1.0, -1.0]),
1625 ],
1626 output: Fp32TierTensor::Fp32(&[0.999_753_24, -0.848_283_65, 0.905_148_27, -0.941_375_55]),
1627 absolute_tolerance: TRANSCENDENTAL.0,
1628 relative_tolerance: TRANSCENDENTAL.1,
1629};
1630
1631pub const FP32_OPERATOR_CASE_GROUPS: &[&[TosaFp32OperatorCase]] = &[
1633 FP32_UNARY_CASES,
1634 FP32_BINARY_CASES,
1635 FP32_LOGICAL_CASES,
1636 FP32_REDUCTION_CASES,
1637 FP32_MOVEMENT_CASES,
1638 &[LINEAR_TANH_FP32],
1639];
1640
1641#[cfg(test)]
1642mod tests {
1643 use super::*;
1644 use virtio_accel_tosa::{
1645 DType, ExtensionSet, Level, ProfileSet, Target, Version, low_precision_storage_bytes, parse,
1646 };
1647
1648 #[test]
1649 fn fp32_operator_cases_are_valid_for_the_fp32_target_and_self_consistent() {
1650 let target = Target::new(
1651 Version::TOSA_1_0,
1652 ProfileSet::FLOATING_POINT,
1653 Level::Level8K,
1654 ExtensionSet::NONE,
1655 );
1656 let mut names = std::collections::BTreeSet::new();
1657 for case in FP32_OPERATOR_CASE_GROUPS
1658 .iter()
1659 .flat_map(|group| group.iter())
1660 {
1661 assert!(names.insert(case.name), "duplicate case {}", case.name);
1662 parse(case.artifact)
1663 .unwrap_or_else(|error| panic!("{}: parse {error:?}", case.name))
1664 .validate_for(target)
1665 .unwrap_or_else(|error| panic!("{}: semantics {error:?}", case.name));
1666 assert!(
1667 case.output_matches(&case.output.bytes()),
1668 "{}: oracle must accept itself",
1669 case.name
1670 );
1671 let mut wrong = case.output.bytes();
1672 wrong.pop();
1673 assert!(!case.output_matches(&wrong), "{}: short output", case.name);
1674 }
1675 assert_eq!(names.len(), 39);
1676 }
1677
1678 #[test]
1679 fn fp32_operator_oracle_applies_float_rules_and_exact_auxiliaries() {
1680 let float = FP32_UNARY_CASES[2];
1681 let mut close = float.output.bytes();
1682 let first = f32::from_le_bytes(close[..4].try_into().unwrap());
1684 close[..4].copy_from_slice(&f32::from_bits(first.to_bits() + 1).to_le_bytes());
1685 assert!(float.output_matches(&close));
1686 let mut far = float.output.bytes();
1687 far[..4].copy_from_slice(&(first + 1.0e-3).to_le_bytes());
1688 assert!(!float.output_matches(&far));
1689
1690 let exact = FP32_UNARY_CASES[0];
1691 let mut nudged = exact.output.bytes();
1692 let first = f32::from_le_bytes(nudged[..4].try_into().unwrap());
1693 nudged[..4].copy_from_slice(&f32::from_bits(first.to_bits() + 1).to_le_bytes());
1694 assert!(!exact.output_matches(&nudged));
1695
1696 let logical = FP32_LOGICAL_CASES[0];
1697 assert!(logical.output_matches(&[1, 0, 0, 1]));
1698 assert!(!logical.output_matches(&[1, 0, 0, 2]));
1699 let argmax = FP32_REDUCTION_CASES[0];
1700 assert!(argmax.output_matches(&Fp32TierTensor::Int32(&[1, 1]).bytes()));
1701 assert!(!argmax.output_matches(&Fp32TierTensor::Int32(&[1, 0]).bytes()));
1702 }
1703
1704 #[test]
1705 fn matmul_oracle_checks_shape_values_and_signed_zero() {
1706 assert!(MATMUL_FP32.output_matches(0, MATMUL_OUTPUT));
1707 assert!(!MATMUL_FP32.output_matches(0, &[58.0, 64.0]));
1708 assert!(!MATMUL_FP32.output_matches(1, MATMUL_OUTPUT));
1709
1710 let zero = TosaFloat32Case {
1711 outputs: &[Float32Tensor {
1712 shape: &[1],
1713 values: &[-0.0],
1714 }],
1715 ..MATMUL_FP32
1716 };
1717 assert!(zero.output_matches(0, &[-0.0]));
1718 assert!(!zero.output_matches(0, &[0.0]));
1719 }
1720
1721 #[test]
1722 fn max_pool_oracle_preserves_nhwc_order() {
1723 assert!(MAX_POOL2D_FP32.output_matches(0, MAX_POOL2D_OUTPUT));
1724 assert!(
1725 !MAX_POOL2D_FP32.output_matches(0, &[6.0, 8.0, 14.0, 16.0, 106.0, 108.0, 114.0, 116.0])
1726 );
1727 }
1728
1729 #[test]
1730 fn bf16_max_pool_oracle_is_exact_and_nhwc_ordered() {
1731 assert!(MAX_POOL2D_BF16.output_matches(0, MAX_POOL2D_OUTPUT_BF16_BITS));
1732 assert!(!MAX_POOL2D_BF16.output_matches(
1733 0,
1734 &[
1735 0x40c0, 0x4100, 0x4160, 0x4180, 0x42d4, 0x42d8, 0x42e4, 0x42e8
1736 ]
1737 ));
1738 parse(MAX_POOL2D_BF16.artifact)
1739 .unwrap()
1740 .validate_for(Target::new(
1741 Version::TOSA_1_0,
1742 ProfileSet::FLOATING_POINT,
1743 Level::Level8K,
1744 ExtensionSet::BF16,
1745 ))
1746 .unwrap();
1747 }
1748
1749 #[test]
1750 fn identity_edge_oracle_handles_nonfinite_and_signed_zero_values() {
1751 assert!(IDENTITY_EDGES_FP32.output_matches(0, IDENTITY_EDGE_VALUES));
1752 let mut wrong_zero = IDENTITY_EDGE_VALUES.to_vec();
1753 wrong_zero[2] = 0.0;
1754 assert!(!IDENTITY_EDGES_FP32.output_matches(0, &wrong_zero));
1755 }
1756
1757 #[test]
1758 fn fp16_oracle_is_exact_except_for_nan_payloads() {
1759 assert!(MATMUL_FP16.output_matches(0, MATMUL_OUTPUT_FP16_BITS));
1760 assert!(MOCK_LINEAR_CLASSIFIER_FP16.output_matches(0, MOCK_CLASSIFIER_LOGITS_FP16_BITS));
1761 assert!(!MATMUL_FP16.output_matches(0, &[0x5340, 0x5400]));
1762 assert!(!MATMUL_FP16.output_matches(1, MATMUL_OUTPUT_FP16_BITS));
1763
1764 let mut canonicalized_nan = IDENTITY_EDGE_FP16_BITS.to_vec();
1765 canonicalized_nan[0] = 0x7fff;
1766 assert!(IDENTITY_EDGES_FP16.output_matches(0, &canonicalized_nan));
1767 canonicalized_nan[2] = 0x0000;
1768 assert!(!IDENTITY_EDGES_FP16.output_matches(0, &canonicalized_nan));
1769 }
1770
1771 #[test]
1772 fn fp16_max_pool_oracle_preserves_nhwc_order() {
1773 assert!(MAX_POOL2D_FP16.output_matches(0, MAX_POOL2D_OUTPUT_FP16_BITS));
1774 assert!(!MAX_POOL2D_FP16.output_matches(
1775 0,
1776 &[
1777 0x4600, 0x4800, 0x4b00, 0x4c00, 0x56a0, 0x56c0, 0x5720, 0x5740
1778 ]
1779 ));
1780 }
1781
1782 #[test]
1783 fn packed_oracles_preserve_storage_and_int4_layout() {
1784 for case in [
1785 IDENTITY_INT4,
1786 IDENTITY_INT8,
1787 IDENTITY_FP8E4M3,
1788 IDENTITY_FP8E5M2,
1789 ] {
1790 let tensor = case.inputs[0];
1791 let elements = tensor.shape.iter().product();
1792 assert_eq!(case.dtype.storage_bytes(elements), Some(tensor.bytes.len()));
1793 assert!(case.output_matches(0, tensor.bytes));
1794 assert!(!case.output_matches(1, tensor.bytes));
1795 }
1796 assert_eq!(IDENTITY_INT4.inputs[0].bytes, &[0xd9, 0x0f, 0x31, 0x76]);
1797
1798 let mut canonicalized_e4m3_nan = IDENTITY_FP8E4M3.outputs[0].bytes.to_vec();
1799 canonicalized_e4m3_nan[7] = 0xff;
1800 assert!(IDENTITY_FP8E4M3.output_matches(0, &canonicalized_e4m3_nan));
1801 let mut canonicalized_e5m2_nan = IDENTITY_FP8E5M2.outputs[0].bytes.to_vec();
1802 canonicalized_e5m2_nan[7] = 0x7f;
1803 assert!(IDENTITY_FP8E5M2.output_matches(0, &canonicalized_e5m2_nan));
1804 }
1805
1806 #[test]
1807 fn fp8_to_bf16_cast_oracles_are_derived_from_the_shared_exact_decoders() {
1808 use virtio_accel_tosa::{
1809 fp8e4m3_to_bf16_bits, fp8e4m3_to_f32, fp8e5m2_to_bf16_bits, fp8e5m2_to_f32,
1810 };
1811
1812 for case in [CAST_FP8E4M3_TO_BF16, CAST_FP8E5M2_TO_BF16] {
1813 assert_eq!(case.input.bytes.len(), case.output.bits.len());
1814 for (input, expected) in case.input.bytes.iter().zip(case.output.bits) {
1815 let value = match case.input_dtype {
1816 PackedDType::Fp8E4M3 => fp8e4m3_to_f32(*input),
1817 PackedDType::Fp8E5M2 => fp8e5m2_to_f32(*input),
1818 _ => unreachable!("CAST case must use FP8"),
1819 };
1820 if value.is_nan() {
1821 assert!(is_bfloat16_nan(*expected));
1822 } else {
1823 assert_eq!(*expected, (value.to_bits() >> 16) as u16);
1824 }
1825 let shared_bits = match case.input_dtype {
1826 PackedDType::Fp8E4M3 => fp8e4m3_to_bf16_bits(*input),
1827 PackedDType::Fp8E5M2 => fp8e5m2_to_bf16_bits(*input),
1828 _ => unreachable!("CAST case must use FP8"),
1829 };
1830 assert_eq!(*expected, shared_bits);
1831 }
1832 assert!(case.output_matches(case.output.bits));
1833 assert!(!case.output_matches(&case.output.bits[..8]));
1834 }
1835 }
1836
1837 #[test]
1838 fn int8_matmul_oracle_is_derived_from_the_shared_exact_dot_product() {
1839 use virtio_accel_tosa::dot_i8_i32;
1840
1841 let lhs = MATMUL_INT8.inputs[0].bytes;
1842 let rhs = MATMUL_INT8.inputs[1].bytes;
1843 let mut actual = Vec::new();
1844 for row in 0..2 {
1845 for column in 0..2 {
1846 let left = &lhs[row * 3..row * 3 + 3];
1847 let right = [rhs[column], rhs[2 + column], rhs[4 + column]];
1848 actual.push(
1849 dot_i8_i32(
1850 left,
1851 &right,
1852 MATMUL_INT8.zero_points[0],
1853 MATMUL_INT8.zero_points[1],
1854 0,
1855 )
1856 .unwrap(),
1857 );
1858 }
1859 }
1860 assert!(MATMUL_INT8.output_matches(0, &actual));
1861 assert!(!MATMUL_INT8.output_matches(0, &[538, -544]));
1862 assert!(!MATMUL_INT8.output_matches(1, &actual));
1863 }
1864
1865 #[test]
1866 fn int8_classifier_oracle_is_derived_from_the_shared_exact_dot_product() {
1867 use virtio_accel_tosa::dot_i8_i32;
1868
1869 let left = QUANTIZED_CLASSIFIER_INT8.inputs[0].bytes;
1870 let right = QUANTIZED_CLASSIFIER_INT8.inputs[1].bytes;
1871 let mut logits = Vec::new();
1872 for sample in 0..2 {
1873 for class in 0..2 {
1874 let features = &left[sample * 3..sample * 3 + 3];
1875 let weights = [right[class], right[2 + class], right[4 + class]];
1876 logits.push(
1877 dot_i8_i32(
1878 features,
1879 &weights,
1880 QUANTIZED_CLASSIFIER_INT8.zero_points[0],
1881 QUANTIZED_CLASSIFIER_INT8.zero_points[1],
1882 0,
1883 )
1884 .unwrap(),
1885 );
1886 }
1887 }
1888 assert!(QUANTIZED_CLASSIFIER_INT8.output_matches(0, &logits));
1889 assert!(logits[0] > logits[1] && logits[3] > logits[2]);
1892 }
1893
1894 #[test]
1895 fn int32_to_int8_rescale_oracle_is_derived_from_the_shared_exact_helper() {
1896 use virtio_accel_tosa::rescale_i32_to_i8;
1897
1898 let actual: Vec<_> = RESCALE_INT32_TO_INT8
1899 .input
1900 .values
1901 .iter()
1902 .map(|value| {
1903 rescale_i32_to_i8(
1904 *value,
1905 RESCALE_INT32_TO_INT8.multiplier,
1906 RESCALE_INT32_TO_INT8.shift,
1907 RESCALE_INT32_TO_INT8.output_zero_point,
1908 false,
1909 )
1910 .unwrap() as u8
1911 })
1912 .collect();
1913 assert!(RESCALE_INT32_TO_INT8.output_matches(&actual));
1914 assert!(!RESCALE_INT32_TO_INT8.output_matches(&actual[..8]));
1915 }
1916
1917 #[test]
1918 fn packed_artifacts_are_valid_for_their_declared_tosa_profiles_and_extensions() {
1919 let integer = Target::new(
1920 Version::TOSA_1_0,
1921 ProfileSet::INTEGER,
1922 Level::Level8K,
1923 ExtensionSet::NONE,
1924 );
1925 let floating = |extension| {
1926 Target::new(
1927 Version::TOSA_1_0,
1928 ProfileSet::FLOATING_POINT,
1929 Level::Level8K,
1930 extension,
1931 )
1932 };
1933 for (case, target, dtype) in [
1934 (IDENTITY_INT8, integer, DType::INT8),
1935 (
1936 IDENTITY_INT4,
1937 Target::new(
1938 Version::TOSA_1_0,
1939 ProfileSet::INTEGER,
1940 Level::Level8K,
1941 ExtensionSet::INT4,
1942 ),
1943 DType::INT4,
1944 ),
1945 (
1946 IDENTITY_FP8E4M3,
1947 floating(ExtensionSet::FP8E4M3),
1948 DType::FP8E4M3,
1949 ),
1950 (
1951 IDENTITY_FP8E5M2,
1952 floating(ExtensionSet::FP8E5M2),
1953 DType::FP8E5M2,
1954 ),
1955 ] {
1956 parse(case.artifact).unwrap().validate_for(target).unwrap();
1957 let elements = case.inputs[0].shape.iter().product();
1958 assert_eq!(
1959 low_precision_storage_bytes(dtype, elements),
1960 Some(case.inputs[0].bytes.len())
1961 );
1962 }
1963 parse(MATMUL_INT8.artifact)
1964 .unwrap()
1965 .validate_for(integer)
1966 .unwrap();
1967 parse(RESCALE_INT32_TO_INT8.artifact)
1968 .unwrap()
1969 .validate_for(integer)
1970 .unwrap();
1971
1972 let fp8_storage = Target::new(
1973 Version::TOSA_1_0,
1974 ProfileSet::FLOATING_POINT,
1975 Level::Level8K,
1976 ExtensionSet::BF16
1977 .union(ExtensionSet::FP8E4M3)
1978 .union(ExtensionSet::FP8E5M2),
1979 );
1980 for case in [CAST_FP8E4M3_TO_BF16, CAST_FP8E5M2_TO_BF16] {
1981 parse(case.artifact)
1982 .unwrap()
1983 .validate_for(fp8_storage)
1984 .unwrap();
1985 }
1986 }
1987}