Skip to main content

virtio_accel_conformance/
numerics.rs

1//! Device-neutral numerical acceptance cases for hardware backends.
2//!
3//! Each case couples a stable TOSA artifact with exact input shapes and a numerical oracle. Host
4//! backends consume the same bytes and values, so a provider cannot quietly substitute a
5//! backend-specific graph while claiming cross-device equivalence.
6
7/// One immutable IEEE-754 binary16 tensor in a numerical acceptance case.
8///
9/// Elements are stored as their exact bit patterns so this stable-Rust crate does not require the
10/// nightly-only primitive `f16` feature.
11#[derive(Clone, Copy, Debug, PartialEq, Eq)]
12pub struct Float16Tensor {
13    /// Static row-major shape.
14    pub shape: &'static [usize],
15    /// Row-major IEEE-754 binary16 element bits.
16    pub bits: &'static [u16],
17}
18
19/// A stable TOSA graph and binary16 numerical oracle shared by host backends.
20#[derive(Clone, Copy, Debug, PartialEq, Eq)]
21pub struct TosaFloat16Case {
22    /// Diagnostic case name.
23    pub name: &'static str,
24    /// TOSA 1.0 FlatBuffer payload.
25    pub artifact: &'static [u8],
26    /// Block inputs in declared slot order.
27    pub inputs: &'static [Float16Tensor],
28    /// Block outputs in declared slot order.
29    pub outputs: &'static [Float16Tensor],
30}
31
32impl TosaFloat16Case {
33    /// Compare one backend output with this case's selected output oracle.
34    ///
35    /// Every non-NaN value must match bit-for-bit, including signed zero, infinities, and
36    /// subnormals. NaN payloads may be canonicalized by the accelerator.
37    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/// One immutable bfloat16 tensor represented by exact storage bits.
53#[derive(Clone, Copy, Debug, PartialEq, Eq)]
54pub struct Bfloat16Tensor {
55    /// Static row-major shape.
56    pub shape: &'static [usize],
57    /// Row-major IEEE-754 bfloat16 element bits.
58    pub bits: &'static [u16],
59}
60
61/// A stable TOSA EXT-BF16 graph and bit-exact bfloat16 numerical oracle.
62#[derive(Clone, Copy, Debug, PartialEq, Eq)]
63pub struct TosaBfloat16Case {
64    /// Diagnostic case name.
65    pub name: &'static str,
66    /// TOSA 1.0 FlatBuffer payload.
67    pub artifact: &'static [u8],
68    /// Block inputs in declared slot order.
69    pub inputs: &'static [Bfloat16Tensor],
70    /// Block outputs in declared slot order.
71    pub outputs: &'static [Bfloat16Tensor],
72}
73
74impl TosaBfloat16Case {
75    /// Compare one backend output bit-for-bit, allowing only NaN payload canonicalization.
76    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/// Raw tensor storage used by mixed-type Hexagon operator-parity cases.
92#[derive(Clone, Copy, Debug, PartialEq, Eq)]
93pub enum RawTensor {
94    /// IEEE-754 binary16 elements represented by their exact bits.
95    Fp16(&'static [u16]),
96    /// TOSA BOOL storage, one zero-or-one byte per element.
97    Bool(&'static [u8]),
98    /// Signed INT32 elements.
99    Int32(&'static [i32]),
100}
101
102impl RawTensor {
103    /// Encode the tensor in the client-visible little-endian storage layout.
104    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    /// Exact client-visible storage size.
119    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/// One mixed-type TOSA operator case with a numerical output oracle.
129#[derive(Clone, Copy, Debug, PartialEq, Eq)]
130pub struct TosaRawCase {
131    /// Diagnostic case name.
132    pub name: &'static str,
133    /// TOSA 1.0 FlatBuffer payload.
134    pub artifact: &'static [u8],
135    /// Block inputs in declared slot order.
136    pub inputs: &'static [RawTensor],
137    /// Single block output.
138    pub output: RawTensor,
139    /// Maximum accepted binary16 ULP distance; ignored for exact BOOL/INT32 outputs.
140    pub fp16_max_ulps: u16,
141}
142
143impl TosaRawCase {
144    /// Compare raw client-visible bytes with the typed numerical oracle.
145    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/// Packed scalar encoding used by a low-precision TOSA acceptance case.
186#[derive(Clone, Copy, Debug, PartialEq, Eq)]
187pub enum PackedDType {
188    /// Signed two's-complement values packed low nibble first.
189    Int4,
190    /// Signed two's-complement bytes.
191    Int8,
192    /// TOSA FP8 E4M3 bytes.
193    Fp8E4M3,
194    /// TOSA FP8 E5M2 bytes.
195    Fp8E5M2,
196}
197
198impl PackedDType {
199    /// Number of bytes needed for `elements` densely packed values.
200    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/// One immutable packed low-precision tensor.
209#[derive(Clone, Copy, Debug, PartialEq, Eq)]
210pub struct PackedTensor {
211    /// Static row-major logical shape.
212    pub shape: &'static [usize],
213    /// Densely packed elements in TOSA byte order.
214    pub bytes: &'static [u8],
215}
216
217/// A stable explicit FP8 → BF16 CAST with a bit-exact output oracle.
218#[derive(Clone, Copy, Debug, PartialEq, Eq)]
219pub struct TosaFp8ToBfloat16Case {
220    /// Diagnostic case name.
221    pub name: &'static str,
222    /// Source FP8 storage encoding.
223    pub input_dtype: PackedDType,
224    /// TOSA 1.0 FlatBuffer payload.
225    pub artifact: &'static [u8],
226    /// Single graph-visible FP8 input.
227    pub input: PackedTensor,
228    /// Single graph-visible BF16 output.
229    pub output: Bfloat16Tensor,
230}
231
232impl TosaFp8ToBfloat16Case {
233    /// Compare BF16 output bits exactly, allowing only NaN payload canonicalization.
234    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/// One immutable INT32 tensor produced by an integer-profile acceptance case.
252#[derive(Clone, Copy, Debug, PartialEq, Eq)]
253pub struct Int32Tensor {
254    /// Static row-major shape.
255    pub shape: &'static [usize],
256    /// Exact row-major tensor elements.
257    pub values: &'static [i32],
258}
259
260/// A stable TOSA INT8 matrix multiplication with exact INT32 accumulation.
261#[derive(Clone, Copy, Debug, PartialEq, Eq)]
262pub struct TosaInt8MatmulCase {
263    /// Diagnostic case name.
264    pub name: &'static str,
265    /// TOSA 1.0 FlatBuffer payload.
266    pub artifact: &'static [u8],
267    /// Signed INT8 block inputs in declared slot order.
268    pub inputs: &'static [PackedTensor],
269    /// Compile-time zero points for the left and right operands.
270    pub zero_points: [i8; 2],
271    /// Exact INT32 block outputs in declared slot order.
272    pub outputs: &'static [Int32Tensor],
273}
274
275impl TosaInt8MatmulCase {
276    /// Compare one backend output with the exact INT32 oracle.
277    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/// A stable TOSA INT32-to-INT8 RESCALE with exact fixed-point rounding and saturation.
285#[derive(Clone, Copy, Debug, PartialEq, Eq)]
286pub struct TosaInt32ToInt8RescaleCase {
287    /// Diagnostic case name.
288    pub name: &'static str,
289    /// TOSA 1.0 FlatBuffer payload.
290    pub artifact: &'static [u8],
291    /// Single graph-visible INT32 input.
292    pub input: Int32Tensor,
293    /// Non-negative scale32 fixed-point multiplier.
294    pub multiplier: i32,
295    /// TOSA scale32 right shift.
296    pub shift: i8,
297    /// Signed INT8 zero point added after scaling.
298    pub output_zero_point: i8,
299    /// Exact signed INT8 output storage.
300    pub output: PackedTensor,
301}
302
303impl TosaInt32ToInt8RescaleCase {
304    /// Compare backend output bytes exactly with the signed INT8 oracle.
305    pub fn output_matches(self, actual: &[u8]) -> bool {
306        self.output.bytes == actual
307    }
308}
309
310/// A stable TOSA graph and packed low-precision oracle shared by host backends.
311#[derive(Clone, Copy, Debug, PartialEq, Eq)]
312pub struct TosaPackedCase {
313    /// Diagnostic case name.
314    pub name: &'static str,
315    /// Tensor scalar encoding.
316    pub dtype: PackedDType,
317    /// TOSA 1.0 FlatBuffer payload.
318    pub artifact: &'static [u8],
319    /// Block inputs in declared slot order.
320    pub inputs: &'static [PackedTensor],
321    /// Block outputs in declared slot order.
322    pub outputs: &'static [PackedTensor],
323}
324
325impl TosaPackedCase {
326    /// Compare one backend output with this case's selected storage oracle.
327    ///
328    /// Integer values must match bit-for-bit. FP8 values do too, except that NaN sign and payload
329    /// may be canonicalized by an accelerator.
330    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/// One immutable FP32 tensor in a numerical acceptance case.
354#[derive(Clone, Copy, Debug, PartialEq)]
355pub struct Float32Tensor {
356    /// Static row-major shape.
357    pub shape: &'static [usize],
358    /// Row-major tensor elements.
359    pub values: &'static [f32],
360}
361
362/// A stable TOSA graph and FP32 numerical oracle shared by host backends.
363#[derive(Clone, Copy, Debug, PartialEq)]
364pub struct TosaFloat32Case {
365    /// Diagnostic case name.
366    pub name: &'static str,
367    /// TOSA 1.0 FlatBuffer payload.
368    pub artifact: &'static [u8],
369    /// Block inputs in declared slot order.
370    pub inputs: &'static [Float32Tensor],
371    /// Block outputs in declared slot order.
372    pub outputs: &'static [Float32Tensor],
373    /// Maximum absolute error accepted for finite values.
374    pub absolute_tolerance: f32,
375    /// Maximum relative error accepted for finite values.
376    pub relative_tolerance: f32,
377}
378
379impl TosaFloat32Case {
380    /// Compare one backend output with this case's selected output oracle.
381    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
418/// FP32 batched matrix multiplication with non-square operands.
419pub 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
443/// FP32 two-channel NHWC max pooling with a 2x2 kernel and stride two.
444pub 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
470/// BF16 two-channel NHWC max pooling with a 2x2 kernel, stride two, zero padding, and an exact
471/// integer-valued oracle.
472pub 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
495/// FP32 identity over NaN, infinities, signed zeros, a subnormal, and ordinary finite values.
496pub 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
523/// Binary16 batched matrix multiplication with non-square operands.
524pub 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
557/// Binary16 broadcast addition over `[2, 1]` and `[1, 3]` inputs.
558pub 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
565/// Binary16 broadcast subtraction over `[2, 1]` and `[1, 3]` inputs.
566pub 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
573/// Binary16 broadcast multiplication with a compile-time zero shift.
574pub 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
581/// Binary16 broadcast power over `[2, 1]` and `[1, 3]` inputs.
582pub 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
589/// Binary16 broadcast maximum over `[2, 1]` and `[1, 3]` inputs.
590pub 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
597/// Binary16 broadcast minimum over `[2, 1]` and `[1, 3]` inputs.
598pub 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
622/// FP16 unary and activation cases supported by the QNN HTP operator package.
623pub 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
719/// Mixed BOOL/FP16 comparison, logical, and selection cases.
720pub 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
783/// FP16 reductions and INT32 argmax over a two-row input.
784pub 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
822/// Static constants and FP16 data-movement cases.
823pub 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, // [1.0, 2.0, 3.0]
870    0xbc00, 0x3800, 0x4000, // [-1.0, 0.5, 2.0]
871];
872const MOCK_CLASSIFIER_WEIGHTS_FP16_BITS: &[u16] = &[
873    0x3c00, 0x0000, // feature 0 -> [1.0, 0.0]
874    0x0000, 0x3c00, // feature 1 -> [0.0, 1.0]
875    0x3c00, 0xbc00, // feature 2 -> [1.0, -1.0]
876];
877const MOCK_CLASSIFIER_LOGITS_FP16_BITS: &[u16] = &[
878    0x4400, 0xbc00, // [4.0, -1.0]
879    0x3c00, 0xbe00, // [1.0, -1.5]
880];
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
896/// Two-sample, three-feature FP16 linear classifier with a direct-bound 3x2 weight matrix.
897pub 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
921/// Binary16 two-channel NHWC max pooling with a 2x2 kernel and stride two.
922pub 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
938/// Binary16 identity over NaN, infinities, signed zeros, a subnormal, and finite values.
939pub 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
952/// INT8 identity spanning negative, zero, and positive values.
953pub 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
979/// Non-square INT8 batched matrix multiplication with nonzero zero points and INT32 accumulation.
980pub 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
988// Two samples of three INT8 features times a 3x2 feature-classifier weight matrix. Each sample
989// is correctly classifiable by the logits alone (zero point left of the range) and has a clear
990// winning class.
991const 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
1009/// Exact INT8 quantized linear classifier: two samples, three features, two classes, INT32
1010/// logits with unambiguous per-sample argmax winners.
1011pub 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
1026/// Signed scale32 RESCALE covering ties, negative values, saturation, and a nonzero output point.
1027pub 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
1043// Logical values [-7, -3, -1, 0, 1, 3, 6, 7], packed low nibble first.
1044const IDENTITY_INT4_BYTES: &[u8] = &[0xd9, 0x0f, 0x31, 0x76];
1045const IDENTITY_INT4_TENSORS: &[PackedTensor] = &[PackedTensor {
1046    shape: &[8],
1047    bytes: IDENTITY_INT4_BYTES,
1048}];
1049
1050/// Packed INT4 identity spanning the TOSA-defined finite range.
1051pub 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
1065/// FP8 E4M3 identity over signed zeros, subnormals, ordinary values, finite maximum, and NaN.
1066pub 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
1080/// FP8 E5M2 identity over signed zeros, subnormals, one, finite maximum, infinity, and NaN.
1081pub 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
1146/// Explicit FP8 E4M3 → BF16 CAST spanning signed zeros, subnormals, ordinary values, finite
1147/// maximum, and NaN. All 256 byte encodings repeat four times across one XDNA conversion tile.
1148pub 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
1164/// Explicit FP8 E5M2 → BF16 CAST spanning signed zeros, subnormals, one, finite maximum,
1165/// infinity, and NaN. All 256 byte encodings repeat four times across one XDNA conversion tile.
1166pub 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/// Raw tensor storage used by the shared FP32-tier operator cases.
1181#[derive(Clone, Copy, Debug, PartialEq)]
1182pub enum Fp32TierTensor {
1183    /// IEEE-754 binary32 elements.
1184    Fp32(&'static [f32]),
1185    /// TOSA BOOL storage, one zero-or-one byte per element.
1186    Bool(&'static [u8]),
1187    /// Signed INT32 elements.
1188    Int32(&'static [i32]),
1189}
1190
1191impl Fp32TierTensor {
1192    /// Encode the tensor in the client-visible little-endian storage layout.
1193    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    /// Exact client-visible storage size.
1208    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/// One FP32-tier operator case shared by host backends: a stable TOSA 1.0 graph over FP32
1218/// tensors with BOOL and INT32 auxiliaries, and a numerical oracle for its single output.
1219///
1220/// Float outputs follow the [`TosaFloat32Case`] rules: NaN matches any NaN, infinities and
1221/// zeros compare bit-exactly (signed zero included), and finite values must fall within the
1222/// absolute or relative tolerance. BOOL and INT32 outputs compare exactly.
1223#[derive(Clone, Copy, Debug, PartialEq)]
1224pub struct TosaFp32OperatorCase {
1225    /// Diagnostic case name.
1226    pub name: &'static str,
1227    /// TOSA 1.0 FlatBuffer payload.
1228    pub artifact: &'static [u8],
1229    /// Block inputs in declared slot order.
1230    pub inputs: &'static [Fp32TierTensor],
1231    /// Single block output.
1232    pub output: Fp32TierTensor,
1233    /// Maximum absolute error accepted for finite float values.
1234    pub absolute_tolerance: f32,
1235    /// Maximum relative error accepted for finite float values.
1236    pub relative_tolerance: f32,
1237}
1238
1239impl TosaFp32OperatorCase {
1240    /// Compare raw client-visible output bytes with the typed oracle.
1241    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
1274/// Tolerance for operators every conformant device must compute exactly (sign, rounding, and
1275/// selection operators, plus additions and multiplications of exactly representable values).
1276const EXACT: (f32, f32) = (0.0, 0.0);
1277/// Tolerance for `GLSL.std.450`-class built-ins whose relative error Vulkan bounds in ulps
1278/// (`exp`: 3 + 2|x|, `log`: 3, `inversesqrt`: 2, division: 2.5), for the crate-authored
1279/// polynomial evaluations of `sin`, `cos`, `tanh`, and `erf`, and for `pow`
1280/// (4 + 3|y·log₂x| ulps). Roughly 30 ulps at unit magnitude, so a genuinely wrong kernel still
1281/// fails by orders of magnitude.
1282const 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
1300/// FP32 unary and activation operators over `[0.5, 1, 2, 4]`.
1301pub 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
1419/// FP32 broadcast binary operators over `[2, 1]` and `[1, 3]` inputs producing `[2, 3]`.
1420pub 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
1490/// FP32 comparisons, BOOL logic, and selection.
1491pub 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
1545/// FP32 reductions and INT32 argmax over a `[2, 3]` input along axis 1.
1546pub 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
1579/// Static FP32 constants and data-movement operators.
1580pub 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
1616/// Three-operator FP32 graph: `tanh(x · w + bias)` with constant zero points and a `[1, 1, 2]`
1617/// bias broadcast over the `[1, 2, 2]` product. Inputs are the mock classifier's features and
1618/// weights; the oracle is evaluated in binary64 and rounded.
1619pub 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
1631/// Every FP32-tier operator case group, for backends that iterate the whole tier.
1632pub 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        // Perturb the first element by one binary32 ulp: within tolerance.
1683        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        // The per-sample winner must be unambiguous: every row pairs one strong and one weak
1890        // logit in opposite order.
1891        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}