Skip to main content

virtio_accel_xdna/
lower.rs

1//! Portable TOSA admission for the XDNA backend.
2//!
3//! This module compiles on every host (no HRX, no `unsafe`). It declares the backend's advertised
4//! `Target` constants (issue #82) and [`admit`]s a TOSA artifact into a [`CompilerSpec`] — the
5//! validated, integers-and-enums-only description the compiler helper (issue #84) turns into an
6//! amdxdna artifact. Anything outside the advertised subset is rejected here, before any subprocess
7//! runs. Graph lowering for compute tiers grows on top of this; the compilable subsets today are
8//! the BF16 IDENTITY (a DMA copy), BF16 → FP32 MATMUL (issue #90), BF16 NHWC MAX_POOL2D
9//! (issue #91), explicit FP8 → BF16 CAST storage conversion (issue #109), and exact INT8
10//! IDENTITY, zero-point-aware INT8 → INT32 MATMUL (issue #144), and exact static
11//! INT32 → INT8 RESCALE (issue #147).
12
13use virtio_accel_tosa::{
14    AnalyzedValueKind, CapabilityDescriptor, DType, DTypeCapability, ExtensionSet,
15    GraphCapabilities, Level, NanPropagationMode, Op, OpAttributes, OperatorCapability,
16    OperatorConstraints, ProfileSet, RoundingMode, RuntimeConditionSupport, Target, ValueRoles,
17    Version, parse,
18};
19
20/// The BF16 floating-point tier: TOSA 1.0, floating-point profile, level 8K, BF16 extension.
21///
22/// XDNA2 executes BF16 with FP32 accumulation natively; FP32/FP16 have no compute path and are
23/// rejected at admission rather than silently run as BF16 (issue #82).
24pub const XDNA_TOSA_TARGET: Target = Target::new(
25    Version::TOSA_1_0,
26    ProfileSet::FLOATING_POINT,
27    Level::Level8K,
28    ExtensionSet::BF16,
29);
30
31/// The FP8 storage tier: graph-visible E4M3/E5M2 inputs explicitly cast to BF16 on the NPU.
32///
33/// This is separate from [`XDNA_TOSA_TARGET`] so adding the storage tier does not change the
34/// identity of the existing BF16 target. FP8 is not advertised as a compute dtype: every admitted
35/// graph has an FP8 block input, one explicit CAST, and a BF16 block output.
36pub const XDNA_TOSA_FP8_TARGET: Target = Target::new(
37    Version::TOSA_1_0,
38    ProfileSet::FLOATING_POINT,
39    Level::Level8K,
40    ExtensionSet::BF16
41        .union(ExtensionSet::FP8E4M3)
42        .union(ExtensionSet::FP8E5M2),
43);
44
45/// The integer tier: TOSA 1.0, integer profile, level 8K, no extensions.
46///
47/// INT8 identity and zero-point-aware MATMUL with exact INT32 results, kept on a separate target from the
48/// floating-point tier exactly as the OpenVINO backend separates its FP and INTEGER targets.
49pub const XDNA_TOSA_INTEGER_TARGET: Target = Target::new(
50    Version::TOSA_1_0,
51    ProfileSet::INTEGER,
52    Level::Level8K,
53    ExtensionSet::NONE,
54);
55
56const BF16_DTYPES: &[DTypeCapability] = &[
57    DTypeCapability::new(DType::BF16, ValueRoles::ALL),
58    DTypeCapability::new(DType::FP32, ValueRoles::OUTPUT),
59];
60
61const BF16_OPERATORS: &[OperatorCapability] = &[
62    OperatorCapability::new(Op::CONST),
63    OperatorCapability::new(Op::IDENTITY),
64    OperatorCapability::constrained(Op::MATMUL, OperatorConstraints::ZERO_ZERO_POINTS),
65    OperatorCapability::constrained(
66        Op::MAX_POOL2D,
67        OperatorConstraints::PROPAGATING_NAN.union(OperatorConstraints::ZERO_PADDING),
68    ),
69];
70
71/// Conservative capability boundary for the implemented XDNA BF16 execution tier.
72pub const XDNA_TOSA_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
73    target: XDNA_TOSA_TARGET,
74    dtypes: BF16_DTYPES,
75    operators: BF16_OPERATORS,
76    graph: GraphCapabilities {
77        max_regions: 1,
78        max_blocks: 1,
79        dynamic_shapes: false,
80        runtime_conditions: RuntimeConditionSupport::None,
81    },
82};
83
84// The roles must cover every admitted tier's boundary *and* interior. The standalone CAST tier
85// ends at a BF16 block output; the fused MATMUL tier keeps the same explicit BF16 promotion as a
86// graph-interior value and ends at FP32, so BF16 carries `INTERMEDIATE` and FP32 carries `OUTPUT`.
87const FP8_STORAGE_DTYPES: &[DTypeCapability] = &[
88    DTypeCapability::new(DType::FP8E4M3, ValueRoles::INPUT),
89    DTypeCapability::new(DType::FP8E5M2, ValueRoles::INPUT),
90    DTypeCapability::new(
91        DType::BF16,
92        ValueRoles::OUTPUT.union(ValueRoles::INTERMEDIATE),
93    ),
94    DTypeCapability::new(DType::FP32, ValueRoles::OUTPUT),
95];
96
97const FP8_STORAGE_OPERATORS: &[OperatorCapability] = &[
98    OperatorCapability::new(Op::CAST),
99    OperatorCapability::new(Op::CONST),
100    OperatorCapability::constrained(Op::MATMUL, OperatorConstraints::ZERO_ZERO_POINTS),
101];
102
103/// Conservative capability boundary for explicit FP8 storage conversion.
104pub const XDNA_TOSA_FP8_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
105    target: XDNA_TOSA_FP8_TARGET,
106    dtypes: FP8_STORAGE_DTYPES,
107    operators: FP8_STORAGE_OPERATORS,
108    graph: GraphCapabilities {
109        max_regions: 1,
110        max_blocks: 1,
111        dynamic_shapes: false,
112        runtime_conditions: RuntimeConditionSupport::None,
113    },
114};
115
116// This is OpenVINO's semantic surface for its integer target, plus the one role its integer tier
117// does not need: OpenVINO only ever *produces* INT32, while this backend also *consumes* it as the
118// RESCALE input below. A dtype omitted from `INPUT` is unroutable through the standard
119// `INPUT || OUTPUT` capability filter, so the advertised roles have to cover every admitted tier.
120// XDNA's admission functions below impose an additional, hardware-specific static-memory envelope.
121const INTEGER_DTYPES: &[DTypeCapability] = &[
122    DTypeCapability::new(DType::INT8, ValueRoles::ALL),
123    DTypeCapability::new(
124        DType::INT32,
125        ValueRoles::INPUT
126            .union(ValueRoles::OUTPUT)
127            .union(ValueRoles::CONSTANT)
128            .union(ValueRoles::INTERMEDIATE),
129    ),
130];
131
132const INTEGER_OPERATORS: &[OperatorCapability] = &[
133    OperatorCapability::new(Op::CONST),
134    OperatorCapability::new(Op::IDENTITY),
135    OperatorCapability::new(Op::MATMUL),
136    OperatorCapability::new(Op::RESCALE),
137];
138
139/// Conservative exact integer-profile capability boundary for XDNA lowering.
140///
141/// This mirrors OpenVINO's dtype and operator declaration. XDNA additionally admits only bounded,
142/// static specializations that fit the proven one-core implementation; this is a hardware
143/// envelope, not a semantic fallback.
144pub const XDNA_TOSA_INTEGER_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
145    target: XDNA_TOSA_INTEGER_TARGET,
146    dtypes: INTEGER_DTYPES,
147    operators: INTEGER_OPERATORS,
148    graph: GraphCapabilities {
149        max_regions: 1,
150        max_blocks: 1,
151        dynamic_shapes: false,
152        runtime_conditions: RuntimeConditionSupport::None,
153    },
154};
155
156/// The DMA line size the IDENTITY template transfers in; an admitted element count must be a
157/// positive multiple of it. The Rust compiler driver carries this value into the helper spec.
158pub(crate) const IDENTITY_LINE_SIZE: usize = 1024;
159
160/// The fixed element tile converted by the FP8 → BF16 kernel.
161pub(crate) const FP8_CAST_LINE_SIZE: usize = 1024;
162
163/// Largest direct-DMA line used by the INT8 IDENTITY template.
164///
165/// Small corpus tensors use one shorter line. Every line is four-byte aligned because the AIE DMA
166/// rejects shorter granularity; larger tensors must divide into complete 1,024-byte lines so the
167/// runtime sequence never needs a hidden tail copy.
168pub(crate) const INT8_IDENTITY_MAX_LINE_SIZE: usize = 1024;
169
170/// The one tested MATMUL compute tile (`m`, `k`, `n`), proven on npu2 (AIE2P).
171///
172/// The bf16→fp32 kernel's micro-tile is (4, 8, 8); this L1-fitting macro-tile is a multiple of it,
173/// and every admitted `(M, K, N)` is a positive multiple of this tile — the single tiling the
174/// helper compiles and the hardware tests exercise. Untested shapes are rejected (issue #90).
175/// The FP32 output is 4 B/element, so this tile is smaller than a same-shape bf16 tile would be, to
176/// keep the double-buffered C tile plus the A/B tiles inside the compute core's ~64 KiB L1. The
177/// Rust compiler driver carries this tested tile into the helper spec.
178pub(crate) const MATMUL_TILE_M: usize = 32;
179pub(crate) const MATMUL_TILE_K: usize = 64;
180pub(crate) const MATMUL_TILE_N: usize = 32;
181
182/// Largest admitted MATMUL dimension. The tested tiling generalizes across multiples of the tile,
183/// but only within this envelope; larger shapes are a later generalization and are rejected now.
184pub(crate) const MATMUL_MAX_DIM: usize = 512;
185
186/// Largest combined input/output footprint for the exact INT8 MATMUL specialization.
187///
188/// The initial kernel keeps complete A, B, and C tensors in one AIE2P core. Limiting their
189/// undoubled footprint to 16 KiB leaves room for depth-two FIFOs, code, and bookkeeping in its
190/// roughly 64 KiB local memory. This is stricter than OpenVINO because it is an XDNA memory limit.
191pub(crate) const INT8_MATMUL_MAX_TOTAL_BYTES: usize = 16 * 1024;
192
193/// Largest combined direct-bound INT32 input and padded INT8 output for one RESCALE worker.
194///
195/// Like the first INT8 MATMUL path, the worker keeps complete depth-two objects in one AIE2P
196/// core. This bound leaves room for code and bookkeeping and is rejected before compilation.
197pub(crate) const INT8_RESCALE_MAX_TOTAL_BYTES: usize = 16 * 1024;
198
199/// Maximum admitted pooling kernel and stride dimensions. Larger windows are unproven on the
200/// scalar AIE2P kernel and are rejected before the compiler subprocess runs.
201pub(crate) const MAX_POOL_MAX_KERNEL: usize = 8;
202pub(crate) const MAX_POOL_MAX_STRIDE: usize = 8;
203
204/// Maximum combined input/output BF16 elements for one pooling specialization.
205///
206/// The worker keeps depth-two input and output objects in a compute core's roughly 64 KiB local
207/// memory. Capping the undoubled footprint at 16 KiB leaves half of L1 for code and bookkeeping.
208pub(crate) const MAX_POOL_MAX_TOTAL_ELEMENTS: usize = 8 * 1024;
209
210/// A validated operator specialization ready for the compiler helper. Each variant names its input
211/// and output dtypes; the closed shape is integers only, so no guest bytes cross the boundary.
212#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
213pub enum CompilerSpec {
214    /// BF16 → BF16 elementwise copy of `elements` values (a positive multiple of
215    /// 1,024 values).
216    Identity { elements: usize },
217    /// Exact INT8 → INT8 elementwise copy. `line_size` is the complete DMA line selected by
218    /// admission; no tail staging or dtype conversion is permitted.
219    Int8Identity { elements: usize, line_size: usize },
220    /// Explicit FP8 storage conversion to BF16. Every finite source value is exactly
221    /// representable, so the output is bit-exact except for permitted NaN canonicalization.
222    Fp8ToBf16 { format: Fp8Format, elements: usize },
223    /// BF16 × BF16 → FP32 matrix multiply `C[M, N] = A[M, K] · B[K, N]` (batch 1). Each of `m`,
224    /// `k`, `n` is a positive multiple of the corresponding MATMUL tile dimension and at most
225    /// 512. The FP32 output is the TOSA-mandated accumulator (issue #82).
226    Matmul { m: usize, k: usize, n: usize },
227    /// Fused FP8 × FP8 → FP32 matrix multiply (batch 1): the graph's explicit BF16 promotion is
228    /// performed on the compute core, per L1 tile, instead of through DDR.
229    ///
230    /// Numerically identical to [`CompilerSpec::Fp8ToBf16`] followed by [`CompilerSpec::Matmul`] —
231    /// FP8 → BF16 is exact for every encoding and the multiply is the same BF16 → FP32 kernel — so
232    /// fusing is a placement choice, not a change of numerical contract. It removes the
233    /// caller-visible BF16 tensors and their DDR round trip.
234    Fp8Matmul {
235        format: Fp8Format,
236        m: usize,
237        k: usize,
238        n: usize,
239    },
240    /// Exact zero-point-aware INT8 × INT8 → INT32 matrix multiply (batch 1).
241    ///
242    /// The serialized TOSA zero points are part of the specialization and therefore also part of
243    /// the compiler cache key. Arithmetic is `(a - a_zp) * (b - b_zp)` accumulated in INT32.
244    Int8Matmul {
245        m: usize,
246        k: usize,
247        n: usize,
248        left_zero_point: i8,
249        right_zero_point: i8,
250    },
251    /// Exact signed INT32 → INT8 scale32 RESCALE with one shared multiplier and shift.
252    Int32ToInt8Rescale {
253        elements: usize,
254        multiplier: i32,
255        shift: i8,
256        output_zero_point: i8,
257    },
258    /// Batch-1 BF16 NHWC MAX_POOL2D with zero padding and propagating NaNs. The complete static
259    /// specialization is carried to the helper so no TOSA bytes cross the subprocess boundary.
260    MaxPool2d {
261        input_h: usize,
262        input_w: usize,
263        channels: usize,
264        output_h: usize,
265        output_w: usize,
266        kernel_h: usize,
267        kernel_w: usize,
268        stride_h: usize,
269        stride_w: usize,
270    },
271}
272
273/// The two TOSA/OCP FP8 storage encodings accepted by the conversion template.
274#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
275pub enum Fp8Format {
276    E4M3,
277    E5M2,
278}
279
280/// Why a TOSA artifact is not admissible to the compilable subset.
281#[derive(Clone, Copy, Debug, PartialEq, Eq)]
282pub enum AdmitError {
283    /// The bytes are not a valid TOSA artifact.
284    Parse,
285    /// The graph is valid TOSA but not for the requested target.
286    Analysis,
287    /// The graph is admissible TOSA but outside the compilable subset.
288    Unsupported,
289}
290
291impl core::fmt::Display for AdmitError {
292    fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
293        write!(formatter, "{self:?}")
294    }
295}
296
297impl std::error::Error for AdmitError {}
298
299/// The one authoritative mapping to wire-level error codes, shared by `load_program` and
300/// `compile_artifact` so the two paths can never drift. The codes match the OpenVINO backend's
301/// classification of the same failure classes (its `lowering_error`): a malformed or semantically
302/// invalid artifact is the guest's mistake (`InvalidArgument`); a valid graph this backend cannot
303/// execute is `Unsupported`.
304impl From<AdmitError> for virtio_accel_core::BackendError {
305    fn from(error: AdmitError) -> Self {
306        match error {
307            AdmitError::Parse | AdmitError::Analysis => Self::InvalidArgument,
308            AdmitError::Unsupported => Self::Unsupported,
309        }
310    }
311}
312
313/// Admit a TOSA artifact for `target`, returning the specialization the helper compiles.
314///
315/// The BF16 target admits IDENTITY (a DMA copy), BF16 → FP32 MATMUL, and BF16 NHWC MAX_POOL2D.
316/// The separate FP8 storage target admits one explicit FP8 → BF16 CAST. The integer target admits
317/// exact INT8 IDENTITY, zero-point-aware INT8 → INT32 MATMUL, and static signed INT32 → INT8
318/// RESCALE. Everything else is rejected without running the compiler. Each template admits only
319/// graphs whose **dataflow** matches what
320/// the compiled kernel executes — the IDENTITY
321/// template requires every operator to be IDENTITY (no constants: with a single block input, every
322/// value then provably carries that input's bytes), and the MATMUL template requires the operator's
323/// operands to be exactly the block inputs (constants may exist only as the two zero-points).
324/// Without these checks a semantically different graph (say, a constant-output IDENTITY or a
325/// constant-weights MATMUL) would compile to a kernel that reads runtime buffers the graph never
326/// asked for, returning well-formed but wrong data. Semantic and target validity — including that
327/// BF16 MATMUL zero-points are constant zero — is enforced by
328/// [`analyze_for`](virtio_accel_tosa::Model::analyze_for) before these structural checks.
329pub fn admit(bytes: &[u8], target: Target) -> Result<CompilerSpec, AdmitError> {
330    if target != XDNA_TOSA_TARGET
331        && target != XDNA_TOSA_FP8_TARGET
332        && target != XDNA_TOSA_INTEGER_TARGET
333    {
334        return Err(AdmitError::Unsupported);
335    }
336    let model = parse(bytes).map_err(|_| AdmitError::Parse)?;
337    let analysis = model
338        .analyze_for(target)
339        .map_err(|_| AdmitError::Analysis)?;
340
341    if analysis.regions().len() != 1
342        || analysis.blocks().len() != 1
343        || !analysis.conditions().is_empty()
344    {
345        return Err(AdmitError::Unsupported);
346    }
347    let block = analysis.blocks()[0].id();
348
349    // Classify in one pass. IDENTITY and CAST tolerate no other operator kind (not even CONST);
350    // MATMUL tolerates exactly one MATMUL plus CONST operators, which `admit_matmul` then pins down
351    // to the two zero-points.
352    let mut matmul = None;
353    let mut max_pool = None;
354    let mut casts: [Option<virtio_accel_tosa::OperatorId>; 2] = [None, None];
355    let mut cast_count = 0usize;
356    let mut rescale = None;
357    let mut identities = 0usize;
358    let mut constants = 0usize;
359    for operator in analysis.execution_order(block) {
360        match analysis.operator(*operator).op() {
361            Op::IDENTITY => identities += 1,
362            Op::CONST => constants += 1,
363            Op::MATMUL if matmul.is_none() => matmul = Some(*operator),
364            Op::MAX_POOL2D if max_pool.is_none() => max_pool = Some(*operator),
365            Op::CAST if cast_count < casts.len() => {
366                casts[cast_count] = Some(*operator);
367                cast_count += 1;
368            }
369            Op::RESCALE if rescale.is_none() => rescale = Some(*operator),
370            _ => return Err(AdmitError::Unsupported),
371        }
372    }
373    match (
374        target, matmul, max_pool, cast_count, rescale, identities, constants,
375    ) {
376        // All-IDENTITY (zero operators included: the block output then *is* the block input, and a
377        // DMA copy is exact for it).
378        (XDNA_TOSA_TARGET, None, None, 0, None, _, 0) => admit_identity(&analysis, block),
379        (XDNA_TOSA_TARGET, Some(matmul), None, 0, None, 0, _) => {
380            admit_matmul(&analysis, block, matmul)
381        }
382        (XDNA_TOSA_TARGET, None, Some(max_pool), 0, None, 0, 0) => {
383            admit_max_pool2d(&analysis, block, max_pool)
384        }
385        (XDNA_TOSA_FP8_TARGET, None, None, 1, None, 0, 0) => {
386            admit_fp8_to_bf16(&analysis, block, casts[0].ok_or(AdmitError::Unsupported)?)
387        }
388        // Fused: both MATMUL operands are promoted from FP8 by their own explicit CAST.
389        (XDNA_TOSA_FP8_TARGET, Some(matmul), None, 2, None, 0, _) => admit_fp8_matmul(
390            &analysis,
391            block,
392            matmul,
393            [
394                casts[0].ok_or(AdmitError::Unsupported)?,
395                casts[1].ok_or(AdmitError::Unsupported)?,
396            ],
397        ),
398        (XDNA_TOSA_INTEGER_TARGET, None, None, 0, None, _, 0) => {
399            admit_int8_identity(&analysis, block)
400        }
401        (XDNA_TOSA_INTEGER_TARGET, Some(matmul), None, 0, None, 0, _) => {
402            admit_int8_matmul(&analysis, block, matmul)
403        }
404        (XDNA_TOSA_INTEGER_TARGET, None, None, 0, Some(rescale), 0, 4) => {
405            admit_int32_to_int8_rescale(&analysis, block, rescale)
406        }
407        _ => Err(AdmitError::Unsupported),
408    }
409}
410
411/// Admit exact INT8 IDENTITY, including the shared eight-byte corpus fixture.
412///
413/// Unlike BF16's established 1,024-element template, a small INT8 tensor is one complete DMA line.
414/// Larger tensors must be an exact multiple of the maximum line so no tail is staged on the host.
415fn admit_int8_identity(
416    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
417    block: virtio_accel_tosa::BlockId,
418) -> Result<CompilerSpec, AdmitError> {
419    let inputs = analysis.block_inputs(block);
420    let outputs = analysis.block_outputs(block);
421    if inputs.len() != 1 || outputs.len() != 1 {
422        return Err(AdmitError::Unsupported);
423    }
424    for value in analysis.values() {
425        if let AnalyzedValueKind::Tensor(tensor) = value.kind() {
426            if tensor.dtype() != DType::INT8 {
427                return Err(AdmitError::Unsupported);
428            }
429        }
430    }
431
432    let elements = tensor_elements(analysis, outputs[0])?;
433    let line_size = elements.min(INT8_IDENTITY_MAX_LINE_SIZE);
434    if line_size % 4 != 0 || (elements > INT8_IDENTITY_MAX_LINE_SIZE && elements % line_size != 0) {
435        return Err(AdmitError::Unsupported);
436    }
437    Ok(CompilerSpec::Int8Identity {
438        elements,
439        line_size,
440    })
441}
442
443/// Admit exact zero-point-aware INT8 MATMUL with the same graph semantics as OpenVINO.
444///
445/// XDNA specializes the two serialized rank-1 INT8 zero points into the device kernel because the
446/// AIE kernel API has no provider graph in which to insert OpenVINO's widen-and-subtract nodes.
447/// Both paths compute the same TOSA expression; XDNA never adjusts values on the host.
448fn admit_int8_matmul(
449    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
450    block: virtio_accel_tosa::BlockId,
451    matmul: virtio_accel_tosa::OperatorId,
452) -> Result<CompilerSpec, AdmitError> {
453    let inputs = analysis.operator_inputs(matmul);
454    let outputs = analysis.operator_outputs(matmul);
455    if inputs.len() != 4
456        || outputs.len() != 1
457        // The compiled kernel always declares two independent input slots, so one value feeding
458        // both operands would let a caller bind different buffers and compute `A * B` for a graph
459        // that says `X * X`.
460        || inputs[0] == inputs[1]
461        || analysis.block_inputs(block) != [inputs[0], inputs[1]]
462        || analysis.block_outputs(block) != [outputs[0]]
463    {
464        return Err(AdmitError::Unsupported);
465    }
466    for operator in analysis.execution_order(block) {
467        if analysis.operator(*operator).op() != Op::CONST {
468            continue;
469        }
470        for produced in analysis.operator_outputs(*operator) {
471            if *produced != inputs[2] && *produced != inputs[3] {
472                return Err(AdmitError::Unsupported);
473            }
474        }
475    }
476
477    let lhs = matmul_dims(analysis, inputs[0], DType::INT8)?;
478    let rhs = matmul_dims(analysis, inputs[1], DType::INT8)?;
479    let out = matmul_dims(analysis, outputs[0], DType::INT32)?;
480    let ([1, m, k], [1, k2, n], [1, m2, n2]) = (lhs, rhs, out) else {
481        return Err(AdmitError::Unsupported);
482    };
483    if k != k2 || m != m2 || n != n2 || [m, k, n].iter().any(|dim| *dim > MATMUL_MAX_DIM) {
484        return Err(AdmitError::Unsupported);
485    }
486
487    // AIE DMA descriptors transfer whole 32-bit words. The direct-binding ABI therefore rounds
488    // each INT8 input slot up to four bytes; the kernel ignores those explicit padding bytes.
489    let lhs_bytes = align_to_four(m.checked_mul(k).ok_or(AdmitError::Unsupported)?)?;
490    let rhs_bytes = align_to_four(k.checked_mul(n).ok_or(AdmitError::Unsupported)?)?;
491    let output_bytes = m
492        .checked_mul(n)
493        .and_then(|elements| elements.checked_mul(4))
494        .ok_or(AdmitError::Unsupported)?;
495    if lhs_bytes
496        .checked_add(rhs_bytes)
497        .and_then(|bytes| bytes.checked_add(output_bytes))
498        .is_none_or(|bytes| bytes > INT8_MATMUL_MAX_TOTAL_BYTES)
499    {
500        return Err(AdmitError::Unsupported);
501    }
502
503    Ok(CompilerSpec::Int8Matmul {
504        m,
505        k,
506        n,
507        left_zero_point: int8_zero_point(analysis, inputs[2])?,
508        right_zero_point: int8_zero_point(analysis, inputs[3])?,
509    })
510}
511
512/// Admit the released signed scale32 RESCALE row without OpenVINO-style implicit conversion.
513///
514/// OpenVINO does not advertise RESCALE in its conservative integer tier. XDNA deliberately extends
515/// the operator surface here, while retaining the same separate target, static lowering boundary,
516/// direct binding, and reject-don't-fallback behavior. Per-channel, unsigned, double-round, and
517/// inexact-round forms remain unsupported until they have separate kernels and corpus proof.
518fn admit_int32_to_int8_rescale(
519    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
520    block: virtio_accel_tosa::BlockId,
521    rescale: virtio_accel_tosa::OperatorId,
522) -> Result<CompilerSpec, AdmitError> {
523    let inputs = analysis.operator_inputs(rescale);
524    let outputs = analysis.operator_outputs(rescale);
525    if inputs.len() != 5
526        || outputs.len() != 1
527        || analysis.block_inputs(block) != [inputs[0]]
528        || analysis.block_outputs(block) != [outputs[0]]
529    {
530        return Err(AdmitError::Unsupported);
531    }
532    for operator in analysis.execution_order(block) {
533        if analysis.operator(*operator).op() != Op::CONST {
534            continue;
535        }
536        for produced in analysis.operator_outputs(*operator) {
537            if !inputs[1..].contains(produced) {
538                return Err(AdmitError::Unsupported);
539            }
540        }
541    }
542
543    let AnalyzedValueKind::Tensor(input) = analysis.value(inputs[0]).kind() else {
544        return Err(AdmitError::Unsupported);
545    };
546    let AnalyzedValueKind::Tensor(output) = analysis.value(outputs[0]).kind() else {
547        return Err(AdmitError::Unsupported);
548    };
549    if input.dtype() != DType::INT32
550        || output.dtype() != DType::INT8
551        || !input.dimensions().eq(output.dimensions())
552    {
553        return Err(AdmitError::Unsupported);
554    }
555    let OpAttributes::Rescale {
556        scale32,
557        rounding_mode,
558        per_channel,
559        input_unsigned,
560        output_unsigned,
561    } = analysis.operator(rescale).source().attributes()
562    else {
563        return Err(AdmitError::Unsupported);
564    };
565    if !scale32
566        || rounding_mode != RoundingMode::SINGLE_ROUND
567        || per_channel
568        || input_unsigned
569        || output_unsigned
570    {
571        return Err(AdmitError::Unsupported);
572    }
573
574    let multiplier = int32_constant(analysis, inputs[1])?;
575    let shift = int8_zero_point(analysis, inputs[2])?;
576    let input_zero_point = int32_constant(analysis, inputs[3])?;
577    let output_zero_point = int8_zero_point(analysis, inputs[4])?;
578    if multiplier < 0 || !(2..=62).contains(&shift) || input_zero_point != 0 {
579        return Err(AdmitError::Unsupported);
580    }
581
582    let elements = tensor_elements(analysis, inputs[0])?;
583    let input_bytes = elements.checked_mul(4).ok_or(AdmitError::Unsupported)?;
584    let output_bytes = align_to_four(elements)?;
585    if input_bytes
586        .checked_add(output_bytes)
587        .is_none_or(|bytes| bytes > INT8_RESCALE_MAX_TOTAL_BYTES)
588    {
589        return Err(AdmitError::Unsupported);
590    }
591    Ok(CompilerSpec::Int32ToInt8Rescale {
592        elements,
593        multiplier,
594        shift,
595        output_zero_point,
596    })
597}
598
599fn align_to_four(bytes: usize) -> Result<usize, AdmitError> {
600    bytes
601        .checked_add(3)
602        .map(|bytes| bytes & !3)
603        .ok_or(AdmitError::Unsupported)
604}
605
606fn int32_constant(
607    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
608    value: virtio_accel_tosa::ValueId,
609) -> Result<i32, AdmitError> {
610    let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
611        return Err(AdmitError::Unsupported);
612    };
613    if tensor.dtype() != DType::INT32 || tensor.dimensions().ne([1]) {
614        return Err(AdmitError::Unsupported);
615    }
616    let bytes: [u8; 4] = analysis
617        .serialized_constant(value)
618        .ok_or(AdmitError::Unsupported)?
619        .try_into()
620        .map_err(|_| AdmitError::Unsupported)?;
621    Ok(i32::from_le_bytes(bytes))
622}
623
624fn int8_zero_point(
625    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
626    value: virtio_accel_tosa::ValueId,
627) -> Result<i8, AdmitError> {
628    let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
629        return Err(AdmitError::Unsupported);
630    };
631    if tensor.dtype() != DType::INT8 || tensor.dimensions().ne([1]) {
632        return Err(AdmitError::Unsupported);
633    }
634    let [byte] = analysis
635        .serialized_constant(value)
636        .ok_or(AdmitError::Unsupported)?
637    else {
638        return Err(AdmitError::Unsupported);
639    };
640    Ok(*byte as i8)
641}
642
643fn tensor_elements(
644    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
645    value: virtio_accel_tosa::ValueId,
646) -> Result<usize, AdmitError> {
647    let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
648        return Err(AdmitError::Unsupported);
649    };
650    let mut elements = 1usize;
651    for dimension in tensor.dimensions() {
652        let dimension = usize::try_from(dimension).map_err(|_| AdmitError::Unsupported)?;
653        if dimension == 0 {
654            return Err(AdmitError::Unsupported);
655        }
656        elements = elements
657            .checked_mul(dimension)
658            .ok_or(AdmitError::Unsupported)?;
659    }
660    Ok(elements)
661}
662
663/// Admit one explicit FP8 → BF16 CAST directly connecting the block input and output.
664///
665/// FP8 is storage only: no FP8 arithmetic is inferred or hidden. The guest-visible CAST is the
666/// exact point where the NPU expands each element, after which existing BF16 compute kernels can
667/// consume the result in a subsequent program.
668fn admit_fp8_to_bf16(
669    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
670    block: virtio_accel_tosa::BlockId,
671    cast: virtio_accel_tosa::OperatorId,
672) -> Result<CompilerSpec, AdmitError> {
673    let inputs = analysis.operator_inputs(cast);
674    let outputs = analysis.operator_outputs(cast);
675    if inputs.len() != 1
676        || outputs.len() != 1
677        || analysis.block_inputs(block) != [inputs[0]]
678        || analysis.block_outputs(block) != [outputs[0]]
679    {
680        return Err(AdmitError::Unsupported);
681    }
682
683    let AnalyzedValueKind::Tensor(input) = analysis.value(inputs[0]).kind() else {
684        return Err(AdmitError::Unsupported);
685    };
686    let AnalyzedValueKind::Tensor(output) = analysis.value(outputs[0]).kind() else {
687        return Err(AdmitError::Unsupported);
688    };
689    let format = match input.dtype() {
690        DType::FP8E4M3 => Fp8Format::E4M3,
691        DType::FP8E5M2 => Fp8Format::E5M2,
692        _ => return Err(AdmitError::Unsupported),
693    };
694    if output.dtype() != DType::BF16 || input.dimensions().ne(output.dimensions()) {
695        return Err(AdmitError::Unsupported);
696    }
697
698    let mut elements = 1usize;
699    for dimension in output.dimensions() {
700        let dimension = usize::try_from(dimension).map_err(|_| AdmitError::Unsupported)?;
701        if dimension == 0 {
702            return Err(AdmitError::Unsupported);
703        }
704        elements = elements
705            .checked_mul(dimension)
706            .ok_or(AdmitError::Unsupported)?;
707    }
708    if elements % FP8_CAST_LINE_SIZE != 0 {
709        return Err(AdmitError::Unsupported);
710    }
711
712    Ok(CompilerSpec::Fp8ToBf16 { format, elements })
713}
714
715/// Admit the BF16 IDENTITY subset: one BF16 input and output, IDENTITY operators only (already
716/// established by the caller — with no constants and one block input, every value in the block
717/// carries the input's bytes, so the output equals the input under any operator arrangement),
718/// every tensor BF16, and an element count that is a positive multiple of [`IDENTITY_LINE_SIZE`].
719fn admit_identity(
720    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
721    block: virtio_accel_tosa::BlockId,
722) -> Result<CompilerSpec, AdmitError> {
723    let inputs = analysis.block_inputs(block);
724    let outputs = analysis.block_outputs(block);
725    if inputs.len() != 1 || outputs.len() != 1 {
726        return Err(AdmitError::Unsupported);
727    }
728    for value in analysis.values() {
729        if let AnalyzedValueKind::Tensor(tensor) = value.kind() {
730            if tensor.dtype() != DType::BF16 {
731                return Err(AdmitError::Unsupported);
732            }
733        }
734    }
735
736    let AnalyzedValueKind::Tensor(output) = analysis.value(outputs[0]).kind() else {
737        return Err(AdmitError::Unsupported);
738    };
739    output.rank().ok_or(AdmitError::Unsupported)?;
740    let mut elements: usize = 1;
741    for dimension in output.dimensions() {
742        // Negative (dynamic) dimensions must not sign-extend into a huge count.
743        let dimension = usize::try_from(dimension).map_err(|_| AdmitError::Unsupported)?;
744        elements = elements
745            .checked_mul(dimension)
746            .ok_or(AdmitError::Unsupported)?;
747    }
748    if elements == 0 || elements % IDENTITY_LINE_SIZE != 0 {
749        return Err(AdmitError::Unsupported);
750    }
751
752    Ok(CompilerSpec::Identity { elements })
753}
754
755/// Admit the BF16 → FP32 MATMUL subset: `lhs`/`rhs` BF16 rank-3 `[1, M, K]`/`[1, K, N]`, output
756/// FP32 rank-3 `[1, M, N]`, and each of `M`, `K`, `N` a positive multiple of the tested tile within
757/// [`MATMUL_MAX_DIM`]. TOSA's MATMUL carries four inputs — `lhs`, `rhs`, and the two zero-points —
758/// and one output; the zero-points' constant-zero requirement was already enforced by analysis.
759///
760/// The dataflow must be exactly the compiled kernel's: `lhs`/`rhs` are the block inputs (in slot
761/// order — the runtime binds A to slot 0 and B to slot 1), the MATMUL output is the block output,
762/// and every CONST in the graph produces only the two zero-points. A CONST-produced operand (baked
763/// weights) or a CONST-produced block output is a semantically different program and is rejected.
764fn admit_matmul(
765    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
766    block: virtio_accel_tosa::BlockId,
767    matmul: virtio_accel_tosa::OperatorId,
768) -> Result<CompilerSpec, AdmitError> {
769    let inputs = analysis.operator_inputs(matmul);
770    let outputs = analysis.operator_outputs(matmul);
771    if inputs.len() != 4 || outputs.len() != 1 {
772        return Err(AdmitError::Unsupported);
773    }
774
775    // The operator's dataflow must be the block's: lhs/rhs are the block inputs in binding order,
776    // and the MATMUL result is the block output. One value feeding both operands is rejected: the
777    // compiled kernel declares two independent slots, so a caller could bind different buffers and
778    // compute `A * B` for a graph that says `X * X`.
779    if inputs[0] == inputs[1]
780        || analysis.block_inputs(block) != [inputs[0], inputs[1]]
781        || analysis.block_outputs(block) != [outputs[0]]
782    {
783        return Err(AdmitError::Unsupported);
784    }
785    // Every CONST feeds only the zero-points (operands 2 and 3); any other constant value would be
786    // graph state the compiled kernel cannot reproduce.
787    for operator in analysis.execution_order(block) {
788        if analysis.operator(*operator).op() != Op::CONST {
789            continue;
790        }
791        for produced in analysis.operator_outputs(*operator) {
792            if *produced != inputs[2] && *produced != inputs[3] {
793                return Err(AdmitError::Unsupported);
794            }
795        }
796    }
797
798    let lhs = matmul_dims(analysis, inputs[0], DType::BF16)?;
799    let rhs = matmul_dims(analysis, inputs[1], DType::BF16)?;
800    let out = matmul_dims(analysis, outputs[0], DType::FP32)?;
801
802    // Batch 1 only (the tested tiling), and the shared dimensions must agree: A[1,M,K], B[1,K,N],
803    // C[1,M,N].
804    let ([1, m, k], [1, k2, n], [1, m2, n2]) = (lhs, rhs, out) else {
805        return Err(AdmitError::Unsupported);
806    };
807    if k != k2 || m != m2 || n != n2 {
808        return Err(AdmitError::Unsupported);
809    }
810    if !tile_admissible(m, MATMUL_TILE_M)
811        || !tile_admissible(k, MATMUL_TILE_K)
812        || !tile_admissible(n, MATMUL_TILE_N)
813    {
814        return Err(AdmitError::Unsupported);
815    }
816
817    Ok(CompilerSpec::Matmul { m, k, n })
818}
819
820/// Admit a fused FP8 MATMUL: `CAST(A_fp8) . CAST(B_fp8) -> FP32`, batch 1.
821///
822/// This is the only admitted tier with graph-interior values, so the dataflow is pinned exactly:
823/// each MATMUL operand must be produced by its own CAST, each CAST must consume a distinct block
824/// input, and the promoted BF16 values must not escape as block outputs. Anything looser would let
825/// the compiled kernel — which binds two FP8 inputs and one FP32 output and promotes internally —
826/// stand in for a graph it does not implement.
827///
828/// The promotion stays explicit in the graph, exactly as the standalone CAST tier requires; only
829/// its *placement* changes, from a DDR round trip to core-local scratch. The arithmetic is
830/// unchanged, so results are bit-identical to running CAST and MATMUL as separate programs.
831fn admit_fp8_matmul(
832    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
833    block: virtio_accel_tosa::BlockId,
834    matmul: virtio_accel_tosa::OperatorId,
835    casts: [virtio_accel_tosa::OperatorId; 2],
836) -> Result<CompilerSpec, AdmitError> {
837    let inputs = analysis.operator_inputs(matmul);
838    let outputs = analysis.operator_outputs(matmul);
839    if inputs.len() != 4 || outputs.len() != 1 {
840        return Err(AdmitError::Unsupported);
841    }
842    if analysis.block_outputs(block) != [outputs[0]] {
843        return Err(AdmitError::Unsupported);
844    }
845
846    // Pair each CAST with the MATMUL operand it produces; the two must cover lhs and rhs exactly
847    // once each, in binding order.
848    let mut promoted: [Option<virtio_accel_tosa::ValueId>; 2] = [None, None];
849    for cast in casts {
850        let cast_inputs = analysis.operator_inputs(cast);
851        let cast_outputs = analysis.operator_outputs(cast);
852        if cast_inputs.len() != 1 || cast_outputs.len() != 1 {
853            return Err(AdmitError::Unsupported);
854        }
855        let operand = if cast_outputs[0] == inputs[0] {
856            0
857        } else if cast_outputs[0] == inputs[1] {
858            1
859        } else {
860            // A CAST feeding anything but a MATMUL operand is graph the kernel cannot reproduce.
861            return Err(AdmitError::Unsupported);
862        };
863        if promoted[operand].is_some() {
864            return Err(AdmitError::Unsupported);
865        }
866        // The promoted value is interior: it feeds the multiply and must not also be a block
867        // output, which would require materializing the BF16 tensor this tier exists to avoid.
868        if analysis.block_outputs(block).contains(&cast_outputs[0]) {
869            return Err(AdmitError::Unsupported);
870        }
871        promoted[operand] = Some(cast_inputs[0]);
872    }
873    let ([Some(lhs_storage), Some(rhs_storage)], _) = (promoted, ()) else {
874        return Err(AdmitError::Unsupported);
875    };
876
877    // The block's dataflow: the two FP8 storage tensors are the block inputs in binding order. One
878    // value feeding both operands is rejected for the same reason as the BF16 tier.
879    if lhs_storage == rhs_storage || analysis.block_inputs(block) != [lhs_storage, rhs_storage] {
880        return Err(AdmitError::Unsupported);
881    }
882    // Every CONST feeds only the zero points (operands 2 and 3).
883    for operator in analysis.execution_order(block) {
884        if analysis.operator(*operator).op() != Op::CONST {
885            continue;
886        }
887        for produced in analysis.operator_outputs(*operator) {
888            if *produced != inputs[2] && *produced != inputs[3] {
889                return Err(AdmitError::Unsupported);
890            }
891        }
892    }
893
894    // Both operands must carry the same FP8 encoding: the compiled kernel instantiates one decoder.
895    let lhs_format = fp8_storage_format(analysis, lhs_storage)?;
896    let rhs_format = fp8_storage_format(analysis, rhs_storage)?;
897    if lhs_format != rhs_format {
898        return Err(AdmitError::Unsupported);
899    }
900
901    let lhs = matmul_dims(analysis, lhs_storage, fp8_dtype(lhs_format))?;
902    let rhs = matmul_dims(analysis, rhs_storage, fp8_dtype(rhs_format))?;
903    let lhs_bf16 = matmul_dims(analysis, inputs[0], DType::BF16)?;
904    let rhs_bf16 = matmul_dims(analysis, inputs[1], DType::BF16)?;
905    let out = matmul_dims(analysis, outputs[0], DType::FP32)?;
906
907    // Promotion is elementwise, so each CAST must preserve its operand's shape exactly.
908    if lhs != lhs_bf16 || rhs != rhs_bf16 {
909        return Err(AdmitError::Unsupported);
910    }
911    let ([1, m, k], [1, k2, n], [1, m2, n2]) = (lhs, rhs, out) else {
912        return Err(AdmitError::Unsupported);
913    };
914    if k != k2 || m != m2 || n != n2 {
915        return Err(AdmitError::Unsupported);
916    }
917    if !tile_admissible(m, MATMUL_TILE_M)
918        || !tile_admissible(k, MATMUL_TILE_K)
919        || !tile_admissible(n, MATMUL_TILE_N)
920    {
921        return Err(AdmitError::Unsupported);
922    }
923
924    Ok(CompilerSpec::Fp8Matmul {
925        format: lhs_format,
926        m,
927        k,
928        n,
929    })
930}
931
932/// The FP8 storage encoding of `value`, or `Unsupported` for any other dtype.
933fn fp8_storage_format(
934    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
935    value: virtio_accel_tosa::ValueId,
936) -> Result<Fp8Format, AdmitError> {
937    let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
938        return Err(AdmitError::Unsupported);
939    };
940    match tensor.dtype() {
941        DType::FP8E4M3 => Ok(Fp8Format::E4M3),
942        DType::FP8E5M2 => Ok(Fp8Format::E5M2),
943        _ => Err(AdmitError::Unsupported),
944    }
945}
946
947/// The TOSA dtype of an FP8 storage encoding.
948fn fp8_dtype(format: Fp8Format) -> DType {
949    match format {
950        Fp8Format::E4M3 => DType::FP8E4M3,
951        Fp8Format::E5M2 => DType::FP8E5M2,
952    }
953}
954
955/// The rank-3 dimensions of `value`, requiring the given dtype and every dimension statically
956/// positive. Dynamic (non-positive) dimensions and non-tensor or wrong-dtype values are rejected.
957fn matmul_dims(
958    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
959    value: virtio_accel_tosa::ValueId,
960    dtype: DType,
961) -> Result<[usize; 3], AdmitError> {
962    let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
963        return Err(AdmitError::Unsupported);
964    };
965    if tensor.dtype() != dtype || tensor.rank() != Some(3) {
966        return Err(AdmitError::Unsupported);
967    }
968    let mut dims = [0usize; 3];
969    for (slot, dimension) in dims.iter_mut().zip(tensor.dimensions()) {
970        *slot = usize::try_from(dimension).map_err(|_| AdmitError::Unsupported)?;
971        if *slot == 0 {
972            return Err(AdmitError::Unsupported);
973        }
974    }
975    Ok(dims)
976}
977
978/// A dimension is admissible when it is a positive multiple of the tested tile within the envelope.
979fn tile_admissible(dim: usize, tile: usize) -> bool {
980    dim > 0 && dim % tile == 0 && dim <= MATMUL_MAX_DIM
981}
982
983/// Admit one batch-1 BF16 NHWC MAX_POOL2D directly connecting the block input and output.
984///
985/// OpenVINO accepts the same propagating-NaN and zero-padding semantic envelope. XDNA narrows it
986/// further to bounded static tensors and small positive kernels/strides because each complete
987/// tensor is double-buffered in one AIE2P compute core's local memory.
988fn admit_max_pool2d(
989    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
990    block: virtio_accel_tosa::BlockId,
991    max_pool: virtio_accel_tosa::OperatorId,
992) -> Result<CompilerSpec, AdmitError> {
993    let inputs = analysis.operator_inputs(max_pool);
994    let outputs = analysis.operator_outputs(max_pool);
995    if inputs.len() != 1
996        || outputs.len() != 1
997        || analysis.block_inputs(block) != [inputs[0]]
998        || analysis.block_outputs(block) != [outputs[0]]
999    {
1000        return Err(AdmitError::Unsupported);
1001    }
1002
1003    let OpAttributes::MaxPool2d {
1004        kernel,
1005        stride,
1006        pad,
1007        nan_mode,
1008    } = analysis.operator(max_pool).source().attributes()
1009    else {
1010        return Err(AdmitError::Unsupported);
1011    };
1012    if nan_mode != NanPropagationMode::PROPAGATE {
1013        return Err(AdmitError::Unsupported);
1014    }
1015    let kernel = exact_positive_pair(kernel.iter(), MAX_POOL_MAX_KERNEL)?;
1016    let stride = exact_positive_pair(stride.iter(), MAX_POOL_MAX_STRIDE)?;
1017    let pad: Vec<_> = pad.iter().collect();
1018    if pad != [0, 0, 0, 0] {
1019        return Err(AdmitError::Unsupported);
1020    }
1021
1022    let [batch, input_h, input_w, channels] = pool_dims(analysis, inputs[0])?;
1023    let [output_batch, output_h, output_w, output_channels] = pool_dims(analysis, outputs[0])?;
1024    if batch != 1 || output_batch != 1 || channels != output_channels {
1025        return Err(AdmitError::Unsupported);
1026    }
1027    let input_elements = input_h
1028        .checked_mul(input_w)
1029        .and_then(|elements| elements.checked_mul(channels))
1030        .ok_or(AdmitError::Unsupported)?;
1031    let output_elements = output_h
1032        .checked_mul(output_w)
1033        .and_then(|elements| elements.checked_mul(channels))
1034        .ok_or(AdmitError::Unsupported)?;
1035    if input_elements
1036        .checked_add(output_elements)
1037        .is_none_or(|total| total > MAX_POOL_MAX_TOTAL_ELEMENTS)
1038    {
1039        return Err(AdmitError::Unsupported);
1040    }
1041
1042    Ok(CompilerSpec::MaxPool2d {
1043        input_h,
1044        input_w,
1045        channels,
1046        output_h,
1047        output_w,
1048        kernel_h: kernel[0],
1049        kernel_w: kernel[1],
1050        stride_h: stride[0],
1051        stride_w: stride[1],
1052    })
1053}
1054
1055fn exact_positive_pair(
1056    values: impl Iterator<Item = i32>,
1057    maximum: usize,
1058) -> Result<[usize; 2], AdmitError> {
1059    let values: Vec<_> = values.collect();
1060    let [first, second] = values.as_slice() else {
1061        return Err(AdmitError::Unsupported);
1062    };
1063    let first = usize::try_from(*first).map_err(|_| AdmitError::Unsupported)?;
1064    let second = usize::try_from(*second).map_err(|_| AdmitError::Unsupported)?;
1065    if first == 0 || second == 0 || first > maximum || second > maximum {
1066        return Err(AdmitError::Unsupported);
1067    }
1068    Ok([first, second])
1069}
1070
1071fn pool_dims(
1072    analysis: &virtio_accel_tosa::TosaAnalysis<'_>,
1073    value: virtio_accel_tosa::ValueId,
1074) -> Result<[usize; 4], AdmitError> {
1075    let AnalyzedValueKind::Tensor(tensor) = analysis.value(value).kind() else {
1076        return Err(AdmitError::Unsupported);
1077    };
1078    if tensor.dtype() != DType::BF16 || tensor.rank() != Some(4) {
1079        return Err(AdmitError::Unsupported);
1080    }
1081    let mut dims = [0usize; 4];
1082    for (slot, dimension) in dims.iter_mut().zip(tensor.dimensions()) {
1083        *slot = usize::try_from(dimension).map_err(|_| AdmitError::Unsupported)?;
1084        if *slot == 0 {
1085            return Err(AdmitError::Unsupported);
1086        }
1087    }
1088    Ok(dims)
1089}
1090
1091#[cfg(test)]
1092mod tests {
1093    use super::*;
1094    use virtio_accel_conformance::numerics::RESCALE_INT32_TO_INT8;
1095    use virtio_accel_tosa::Target;
1096    use virtio_accel_tosa_build::{OperatorKind, OwnedGraph, OwnedOperator, OwnedTensor};
1097
1098    fn identity_graph(dtype: DType, shape: Vec<i32>) -> OwnedGraph<'static> {
1099        let mut graph = OwnedGraph::new("main");
1100        graph.push_tensor(OwnedTensor::new("x", shape.clone(), dtype));
1101        graph.push_tensor(OwnedTensor::new("y", shape, dtype));
1102        graph.push_operator(OwnedOperator::new(
1103            OperatorKind::Identity,
1104            vec!["x".into()],
1105            vec!["y".into()],
1106        ));
1107        graph.push_input("x");
1108        graph.push_output("y");
1109        graph
1110    }
1111
1112    fn fp8_cast_graph(input: DType, output: DType, shape: Vec<i32>) -> OwnedGraph<'static> {
1113        let mut graph = OwnedGraph::new("main");
1114        graph
1115            .push_tensor(OwnedTensor::new("x", shape.clone(), input))
1116            .push_tensor(OwnedTensor::new("y", shape, output))
1117            .push_operator(OwnedOperator::new(
1118                OperatorKind::Cast,
1119                vec!["x".into()],
1120                vec!["y".into()],
1121            ))
1122            .push_input("x")
1123            .push_output("y");
1124        graph
1125    }
1126
1127    /// A batch-1 MATMUL `C[1,M,N] = A[1,M,K] · B[1,K,N]` with the two constant-zero zero-points.
1128    fn matmul_graph(
1129        m: i32,
1130        k: i32,
1131        n: i32,
1132        in_dtype: DType,
1133        out_dtype: DType,
1134    ) -> OwnedGraph<'static> {
1135        matmul_graph_with_zero_points(m, k, n, in_dtype, out_dtype, 0, 0)
1136    }
1137
1138    fn matmul_graph_with_zero_points(
1139        m: i32,
1140        k: i32,
1141        n: i32,
1142        in_dtype: DType,
1143        out_dtype: DType,
1144        left_zero_point: i8,
1145        right_zero_point: i8,
1146    ) -> OwnedGraph<'static> {
1147        let zero_point = |dtype: DType, value: i8| match dtype {
1148            DType::INT8 => vec![value as u8],
1149            DType::BF16 => vec![0u8; 2],
1150            DType::FP32 => vec![0u8; 4],
1151            _ => Vec::new(),
1152        };
1153        let mut graph = OwnedGraph::new("main");
1154        graph
1155            .push_tensor(OwnedTensor::new("lhs", vec![1, m, k], in_dtype))
1156            .push_tensor(OwnedTensor::new("rhs", vec![1, k, n], in_dtype))
1157            .push_tensor(OwnedTensor::constant(
1158                "lhs_zp",
1159                vec![1],
1160                in_dtype,
1161                zero_point(in_dtype, left_zero_point),
1162            ))
1163            .push_tensor(OwnedTensor::constant(
1164                "rhs_zp",
1165                vec![1],
1166                in_dtype,
1167                zero_point(in_dtype, right_zero_point),
1168            ))
1169            .push_tensor(OwnedTensor::new("output", vec![1, m, n], out_dtype))
1170            .push_operator(OwnedOperator::new(
1171                OperatorKind::Const,
1172                vec![],
1173                vec!["lhs_zp".into()],
1174            ))
1175            .push_operator(OwnedOperator::new(
1176                OperatorKind::Const,
1177                vec![],
1178                vec!["rhs_zp".into()],
1179            ))
1180            .push_operator(OwnedOperator::new(
1181                OperatorKind::MatMul,
1182                vec!["lhs".into(), "rhs".into(), "lhs_zp".into(), "rhs_zp".into()],
1183                vec!["output".into()],
1184            ))
1185            .push_input("lhs")
1186            .push_input("rhs")
1187            .push_output("output");
1188        graph
1189    }
1190
1191    fn rescale_graph(
1192        elements: i32,
1193        multiplier: i32,
1194        shift: i8,
1195        per_channel: bool,
1196        rounding_mode: RoundingMode,
1197    ) -> OwnedGraph<'static> {
1198        let mut graph = OwnedGraph::new("main");
1199        graph
1200            .push_tensor(OwnedTensor::new("input", vec![elements], DType::INT32))
1201            .push_tensor(OwnedTensor::constant(
1202                "multiplier",
1203                vec![1],
1204                DType::INT32,
1205                multiplier.to_le_bytes().to_vec(),
1206            ))
1207            .push_tensor(OwnedTensor::constant(
1208                "shift",
1209                vec![1],
1210                DType::INT8,
1211                vec![shift as u8],
1212            ))
1213            .push_tensor(OwnedTensor::constant(
1214                "input_zp",
1215                vec![1],
1216                DType::INT32,
1217                0_i32.to_le_bytes().to_vec(),
1218            ))
1219            .push_tensor(OwnedTensor::constant(
1220                "output_zp",
1221                vec![1],
1222                DType::INT8,
1223                vec![(-3_i8) as u8],
1224            ))
1225            .push_tensor(OwnedTensor::new("output", vec![elements], DType::INT8));
1226        for parameter in ["multiplier", "shift", "input_zp", "output_zp"] {
1227            graph.push_operator(OwnedOperator::new(
1228                OperatorKind::Const,
1229                vec![],
1230                vec![parameter.into()],
1231            ));
1232        }
1233        graph
1234            .push_operator(OwnedOperator::new(
1235                OperatorKind::Rescale {
1236                    scale32: true,
1237                    rounding_mode,
1238                    per_channel,
1239                    input_unsigned: false,
1240                    output_unsigned: false,
1241                },
1242                vec![
1243                    "input".into(),
1244                    "multiplier".into(),
1245                    "shift".into(),
1246                    "input_zp".into(),
1247                    "output_zp".into(),
1248                ],
1249                vec!["output".into()],
1250            ))
1251            .push_input("input")
1252            .push_output("output");
1253        graph
1254    }
1255
1256    struct MaxPoolCase {
1257        input: [i32; 3],
1258        kernel: [i32; 2],
1259        stride: [i32; 2],
1260        pad: [i32; 4],
1261        dtype: DType,
1262        nan_mode: NanPropagationMode,
1263    }
1264
1265    fn max_pool_graph(case: MaxPoolCase) -> OwnedGraph<'static> {
1266        let [input_h, input_w, channels] = case.input;
1267        let output_h = (input_h + case.pad[0] + case.pad[1] - case.kernel[0]) / case.stride[0] + 1;
1268        let output_w = (input_w + case.pad[2] + case.pad[3] - case.kernel[1]) / case.stride[1] + 1;
1269        let mut graph = OwnedGraph::new("main");
1270        graph
1271            .push_tensor(OwnedTensor::new(
1272                "input",
1273                vec![1, input_h, input_w, channels],
1274                case.dtype,
1275            ))
1276            .push_tensor(OwnedTensor::new(
1277                "output",
1278                vec![1, output_h, output_w, channels],
1279                case.dtype,
1280            ))
1281            .push_operator(OwnedOperator::new(
1282                OperatorKind::MaxPool2d {
1283                    kernel: case.kernel,
1284                    stride: case.stride,
1285                    pad: case.pad,
1286                    nan_mode: case.nan_mode,
1287                },
1288                vec!["input".into()],
1289                vec!["output".into()],
1290            ))
1291            .push_input("input")
1292            .push_output("output");
1293        graph
1294    }
1295
1296    #[test]
1297    fn both_targets_are_coherent_and_distinct() {
1298        assert_eq!(XDNA_TOSA_TARGET.validate(), Ok(XDNA_TOSA_TARGET));
1299        assert_eq!(
1300            XDNA_TOSA_INTEGER_TARGET.validate(),
1301            Ok(XDNA_TOSA_INTEGER_TARGET)
1302        );
1303        assert_ne!(XDNA_TOSA_TARGET, XDNA_TOSA_INTEGER_TARGET);
1304        for target in [
1305            XDNA_TOSA_TARGET,
1306            XDNA_TOSA_FP8_TARGET,
1307            XDNA_TOSA_INTEGER_TARGET,
1308        ] {
1309            assert_eq!(Target::from_identity(target.to_identity()), Ok(target));
1310        }
1311    }
1312
1313    #[test]
1314    fn admits_bf16_identity() {
1315        let bytes = identity_graph(DType::BF16, vec![1, 4, 1024])
1316            .build(XDNA_TOSA_TARGET)
1317            .expect("build bf16 identity");
1318        let spec = admit(&bytes, XDNA_TOSA_TARGET).expect("admit");
1319        assert_eq!(spec, CompilerSpec::Identity { elements: 4 * 1024 });
1320    }
1321
1322    #[test]
1323    fn integer_capability_preserves_the_openvino_base_and_adds_rescale() {
1324        assert_eq!(
1325            XDNA_TOSA_INTEGER_CAPABILITY.target,
1326            XDNA_TOSA_INTEGER_TARGET
1327        );
1328        assert_eq!(XDNA_TOSA_INTEGER_CAPABILITY.dtypes, INTEGER_DTYPES);
1329        assert_eq!(XDNA_TOSA_INTEGER_CAPABILITY.operators, INTEGER_OPERATORS);
1330        for op in [Op::CONST, Op::IDENTITY, Op::MATMUL, Op::RESCALE] {
1331            assert!(XDNA_TOSA_INTEGER_CAPABILITY.supports_operator(op));
1332        }
1333        assert_eq!(XDNA_TOSA_INTEGER_CAPABILITY.graph.max_regions, 1);
1334        assert_eq!(XDNA_TOSA_INTEGER_CAPABILITY.graph.max_blocks, 1);
1335        assert_eq!(
1336            XDNA_TOSA_INTEGER_CAPABILITY.graph.runtime_conditions,
1337            RuntimeConditionSupport::None
1338        );
1339    }
1340
1341    #[test]
1342    fn admits_shared_int8_identity_shape() {
1343        let bytes = identity_graph(DType::INT8, vec![8])
1344            .build(XDNA_TOSA_INTEGER_TARGET)
1345            .expect("build int8 identity");
1346        assert_eq!(
1347            admit(&bytes, XDNA_TOSA_INTEGER_TARGET),
1348            Ok(CompilerSpec::Int8Identity {
1349                elements: 8,
1350                line_size: 8,
1351            })
1352        );
1353    }
1354
1355    #[test]
1356    fn admits_zero_point_aware_int8_matmul() {
1357        let bytes = matmul_graph_with_zero_points(2, 3, 2, DType::INT8, DType::INT32, -2, 3)
1358            .build(XDNA_TOSA_INTEGER_TARGET)
1359            .expect("build int8 matmul");
1360        assert_eq!(
1361            admit(&bytes, XDNA_TOSA_INTEGER_TARGET),
1362            Ok(CompilerSpec::Int8Matmul {
1363                m: 2,
1364                k: 3,
1365                n: 2,
1366                left_zero_point: -2,
1367                right_zero_point: 3,
1368            })
1369        );
1370    }
1371
1372    #[test]
1373    fn admits_shared_exact_int32_to_int8_rescale() {
1374        assert_eq!(
1375            admit(RESCALE_INT32_TO_INT8.artifact, XDNA_TOSA_INTEGER_TARGET),
1376            Ok(CompilerSpec::Int32ToInt8Rescale {
1377                elements: 16,
1378                multiplier: 1 << 29,
1379                shift: 30,
1380                output_zero_point: -3,
1381            })
1382        );
1383    }
1384
1385    #[test]
1386    fn rescale_rejects_unimplemented_modes_and_invalid_parameters() {
1387        let per_channel = rescale_graph(1, 1 << 29, 30, true, RoundingMode::SINGLE_ROUND)
1388            .build(XDNA_TOSA_INTEGER_TARGET)
1389            .expect("per-channel one-element RESCALE is valid TOSA");
1390        assert_eq!(
1391            admit(&per_channel, XDNA_TOSA_INTEGER_TARGET),
1392            Err(AdmitError::Unsupported)
1393        );
1394
1395        assert!(
1396            rescale_graph(16, 1 << 29, 1, false, RoundingMode::SINGLE_ROUND)
1397                .build(XDNA_TOSA_INTEGER_TARGET)
1398                .is_err(),
1399            "shift 1 violates the released RESCALE range"
1400        );
1401
1402        assert!(
1403            rescale_graph(16, 1 << 29, 30, false, RoundingMode::DOUBLE_ROUND)
1404                .build(XDNA_TOSA_INTEGER_TARGET)
1405                .is_err(),
1406            "DOUBLE_ROUND requires an extension absent from the integer target"
1407        );
1408    }
1409
1410    #[test]
1411    fn rejects_int8_shapes_outside_the_one_core_envelope() {
1412        let non_word_identity = identity_graph(DType::INT8, vec![6])
1413            .build(XDNA_TOSA_INTEGER_TARGET)
1414            .expect("build int8 identity");
1415        assert_eq!(
1416            admit(&non_word_identity, XDNA_TOSA_INTEGER_TARGET),
1417            Err(AdmitError::Unsupported)
1418        );
1419
1420        let non_divisible_identity = identity_graph(DType::INT8, vec![1025])
1421            .build(XDNA_TOSA_INTEGER_TARGET)
1422            .expect("build int8 identity");
1423        assert_eq!(
1424            admit(&non_divisible_identity, XDNA_TOSA_INTEGER_TARGET),
1425            Err(AdmitError::Unsupported)
1426        );
1427
1428        let oversized_matmul =
1429            matmul_graph_with_zero_points(64, 64, 64, DType::INT8, DType::INT32, -2, 3)
1430                .build(XDNA_TOSA_INTEGER_TARGET)
1431                .expect("build int8 matmul");
1432        assert_eq!(
1433            admit(&oversized_matmul, XDNA_TOSA_INTEGER_TARGET),
1434            Err(AdmitError::Unsupported)
1435        );
1436    }
1437
1438    #[test]
1439    fn admits_both_explicit_fp8_to_bf16_casts() {
1440        for (dtype, format) in [
1441            (DType::FP8E4M3, Fp8Format::E4M3),
1442            (DType::FP8E5M2, Fp8Format::E5M2),
1443        ] {
1444            let bytes = fp8_cast_graph(dtype, DType::BF16, vec![1, 1, 4096])
1445                .build(XDNA_TOSA_FP8_TARGET)
1446                .expect("build fp8 cast");
1447            assert_eq!(
1448                admit(&bytes, XDNA_TOSA_FP8_TARGET),
1449                Ok(CompilerSpec::Fp8ToBf16 {
1450                    format,
1451                    elements: 4096,
1452                })
1453            );
1454        }
1455    }
1456
1457    #[test]
1458    fn fp8_storage_tier_rejects_hidden_or_unsupported_conversion() {
1459        let valid = fp8_cast_graph(DType::FP8E4M3, DType::BF16, vec![1024])
1460            .build(XDNA_TOSA_FP8_TARGET)
1461            .expect("build fp8 cast");
1462        assert_eq!(admit(&valid, XDNA_TOSA_TARGET), Err(AdmitError::Analysis));
1463
1464        for graph in [
1465            fp8_cast_graph(DType::FP8E4M3, DType::FP32, vec![1024]),
1466            fp8_cast_graph(DType::FP8E5M2, DType::BF16, vec![8]),
1467        ] {
1468            let bytes = graph
1469                .build(XDNA_TOSA_FP8_TARGET)
1470                .expect("build semantically valid cast");
1471            assert_eq!(
1472                admit(&bytes, XDNA_TOSA_FP8_TARGET),
1473                Err(AdmitError::Unsupported)
1474            );
1475        }
1476    }
1477
1478    #[test]
1479    fn admits_bf16_matmul_at_tile_multiples() {
1480        // Single tile, and a larger non-square multiple of the tested tile.
1481        for (m, k, n) in [(32, 64, 32), (64, 128, 96)] {
1482            let bytes = matmul_graph(m, k, n, DType::BF16, DType::FP32)
1483                .build(XDNA_TOSA_TARGET)
1484                .expect("build bf16 matmul");
1485            let spec = admit(&bytes, XDNA_TOSA_TARGET).expect("admit");
1486            assert_eq!(
1487                spec,
1488                CompilerSpec::Matmul {
1489                    m: m as usize,
1490                    k: k as usize,
1491                    n: n as usize,
1492                }
1493            );
1494        }
1495    }
1496
1497    #[test]
1498    fn admits_bf16_nhwc_max_pool2d_corpus_shape() {
1499        let bytes = max_pool_graph(MaxPoolCase {
1500            input: [4, 4, 2],
1501            kernel: [2, 2],
1502            stride: [2, 2],
1503            pad: [0; 4],
1504            dtype: DType::BF16,
1505            nan_mode: NanPropagationMode::PROPAGATE,
1506        })
1507        .build(XDNA_TOSA_TARGET)
1508        .expect("build bf16 max pool2d");
1509        assert_eq!(
1510            admit(&bytes, XDNA_TOSA_TARGET),
1511            Ok(CompilerSpec::MaxPool2d {
1512                input_h: 4,
1513                input_w: 4,
1514                channels: 2,
1515                output_h: 2,
1516                output_w: 2,
1517                kernel_h: 2,
1518                kernel_w: 2,
1519                stride_h: 2,
1520                stride_w: 2,
1521            })
1522        );
1523    }
1524
1525    #[test]
1526    fn rejects_max_pool2d_outside_the_proven_envelope() {
1527        let cases = [
1528            max_pool_graph(MaxPoolCase {
1529                input: [4, 4, 2],
1530                kernel: [2, 2],
1531                stride: [2, 2],
1532                pad: [0; 4],
1533                dtype: DType::FP32,
1534                nan_mode: NanPropagationMode::PROPAGATE,
1535            }),
1536            max_pool_graph(MaxPoolCase {
1537                input: [4, 4, 2],
1538                kernel: [2, 2],
1539                stride: [2, 2],
1540                pad: [0; 4],
1541                dtype: DType::BF16,
1542                nan_mode: NanPropagationMode::IGNORE,
1543            }),
1544            max_pool_graph(MaxPoolCase {
1545                input: [4, 4, 2],
1546                kernel: [2, 2],
1547                stride: [2, 2],
1548                pad: [1; 4],
1549                dtype: DType::BF16,
1550                nan_mode: NanPropagationMode::PROPAGATE,
1551            }),
1552            max_pool_graph(MaxPoolCase {
1553                input: [16, 16, 2],
1554                kernel: [9, 2],
1555                stride: [1, 1],
1556                pad: [0; 4],
1557                dtype: DType::BF16,
1558                nan_mode: NanPropagationMode::PROPAGATE,
1559            }),
1560            max_pool_graph(MaxPoolCase {
1561                input: [64, 64, 2],
1562                kernel: [2, 2],
1563                stride: [2, 2],
1564                pad: [0; 4],
1565                dtype: DType::BF16,
1566                nan_mode: NanPropagationMode::PROPAGATE,
1567            }),
1568        ];
1569        for graph in cases {
1570            let bytes = graph
1571                .build(XDNA_TOSA_TARGET)
1572                .expect("build semantically valid max pool2d");
1573            assert_eq!(
1574                admit(&bytes, XDNA_TOSA_TARGET),
1575                Err(AdmitError::Unsupported)
1576            );
1577        }
1578    }
1579
1580    #[test]
1581    fn admits_zero_operator_passthrough() {
1582        // The block output *is* the block input; a DMA copy is exact for it.
1583        let mut graph = OwnedGraph::new("main");
1584        graph.push_tensor(OwnedTensor::new("x", vec![1, 4, 1024], DType::BF16));
1585        graph.push_input("x");
1586        graph.push_output("x");
1587        let bytes = graph.build(XDNA_TOSA_TARGET).expect("build passthrough");
1588        assert_eq!(
1589            admit(&bytes, XDNA_TOSA_TARGET),
1590            Ok(CompilerSpec::Identity { elements: 4 * 1024 })
1591        );
1592    }
1593
1594    #[test]
1595    fn rejects_constant_output_identity() {
1596        // TOSA semantics: the output equals the constant. The compiled IDENTITY kernel would copy
1597        // the runtime input instead — silently wrong results — so the graph must not admit.
1598        let shape = vec![1i32, 4, 1024];
1599        let mut graph = OwnedGraph::new("main");
1600        graph
1601            .push_tensor(OwnedTensor::new("x", shape.clone(), DType::BF16))
1602            .push_tensor(OwnedTensor::constant(
1603                "c",
1604                shape.clone(),
1605                DType::BF16,
1606                vec![0u8; 4 * 1024 * 2],
1607            ))
1608            .push_tensor(OwnedTensor::new("y", shape.clone(), DType::BF16))
1609            .push_tensor(OwnedTensor::new("dead", shape, DType::BF16))
1610            .push_operator(OwnedOperator::new(
1611                OperatorKind::Const,
1612                vec![],
1613                vec!["c".into()],
1614            ))
1615            .push_operator(OwnedOperator::new(
1616                OperatorKind::Identity,
1617                vec!["c".into()],
1618                vec!["y".into()],
1619            ))
1620            .push_operator(OwnedOperator::new(
1621                OperatorKind::Identity,
1622                vec!["x".into()],
1623                vec!["dead".into()],
1624            ))
1625            .push_input("x")
1626            .push_output("y");
1627        let bytes = graph
1628            .build(XDNA_TOSA_TARGET)
1629            .expect("build constant identity");
1630        assert_eq!(
1631            admit(&bytes, XDNA_TOSA_TARGET),
1632            Err(AdmitError::Unsupported)
1633        );
1634    }
1635
1636    #[test]
1637    fn rejects_constant_weights_matmul() {
1638        // A CONST-produced lhs (baked weights) is a different program from the two-runtime-input
1639        // kernel the helper compiles; admitting it would matmul against whatever lands in slot 0.
1640        let (m, k, n) = (32i32, 64i32, 32i32);
1641        let mut graph = OwnedGraph::new("main");
1642        graph
1643            .push_tensor(OwnedTensor::constant(
1644                "lhs",
1645                vec![1, m, k],
1646                DType::BF16,
1647                vec![0u8; (m * k * 2) as usize],
1648            ))
1649            .push_tensor(OwnedTensor::new("rhs", vec![1, k, n], DType::BF16))
1650            .push_tensor(OwnedTensor::constant(
1651                "lhs_zp",
1652                vec![1],
1653                DType::BF16,
1654                vec![0u8; 2],
1655            ))
1656            .push_tensor(OwnedTensor::constant(
1657                "rhs_zp",
1658                vec![1],
1659                DType::BF16,
1660                vec![0u8; 2],
1661            ))
1662            .push_tensor(OwnedTensor::new("output", vec![1, m, n], DType::FP32))
1663            .push_operator(OwnedOperator::new(
1664                OperatorKind::Const,
1665                vec![],
1666                vec!["lhs".into()],
1667            ))
1668            .push_operator(OwnedOperator::new(
1669                OperatorKind::Const,
1670                vec![],
1671                vec!["lhs_zp".into()],
1672            ))
1673            .push_operator(OwnedOperator::new(
1674                OperatorKind::Const,
1675                vec![],
1676                vec!["rhs_zp".into()],
1677            ))
1678            .push_operator(OwnedOperator::new(
1679                OperatorKind::MatMul,
1680                vec!["lhs".into(), "rhs".into(), "lhs_zp".into(), "rhs_zp".into()],
1681                vec!["output".into()],
1682            ))
1683            .push_input("rhs")
1684            .push_output("output");
1685        let bytes = graph
1686            .build(XDNA_TOSA_TARGET)
1687            .expect("build constant-weights matmul");
1688        assert_eq!(
1689            admit(&bytes, XDNA_TOSA_TARGET),
1690            Err(AdmitError::Unsupported)
1691        );
1692    }
1693
1694    #[test]
1695    fn admit_error_maps_to_the_reference_backend_error_codes() {
1696        use virtio_accel_core::BackendError;
1697        assert_eq!(
1698            BackendError::from(AdmitError::Parse),
1699            BackendError::InvalidArgument
1700        );
1701        assert_eq!(
1702            BackendError::from(AdmitError::Analysis),
1703            BackendError::InvalidArgument
1704        );
1705        assert_eq!(
1706            BackendError::from(AdmitError::Unsupported),
1707            BackendError::Unsupported
1708        );
1709    }
1710
1711    /// One value feeding both MATMUL operands is rejected. The compiled design always declares two
1712    /// independent input slots, so admitting `X * X` would let a caller bind different buffers to
1713    /// slots 0 and 1 and compute `A * B` — a result no reading of the graph produces.
1714    #[test]
1715    fn rejects_matmul_with_one_value_feeding_both_operands() {
1716        fn aliased_matmul(dim: i32, in_dtype: DType, out_dtype: DType) -> Vec<u8> {
1717            let zp = match in_dtype {
1718                DType::INT8 => vec![0u8],
1719                _ => vec![0u8; 2],
1720            };
1721            let mut graph = OwnedGraph::new("main");
1722            graph
1723                .push_tensor(OwnedTensor::new("x", vec![1, dim, dim], in_dtype))
1724                .push_tensor(OwnedTensor::constant(
1725                    "lhs_zp",
1726                    vec![1],
1727                    in_dtype,
1728                    zp.clone(),
1729                ))
1730                .push_tensor(OwnedTensor::constant("rhs_zp", vec![1], in_dtype, zp))
1731                .push_tensor(OwnedTensor::new("output", vec![1, dim, dim], out_dtype))
1732                .push_operator(OwnedOperator::new(
1733                    OperatorKind::Const,
1734                    vec![],
1735                    vec!["lhs_zp".into()],
1736                ))
1737                .push_operator(OwnedOperator::new(
1738                    OperatorKind::Const,
1739                    vec![],
1740                    vec!["rhs_zp".into()],
1741                ))
1742                .push_operator(OwnedOperator::new(
1743                    OperatorKind::MatMul,
1744                    vec!["x".into(), "x".into(), "lhs_zp".into(), "rhs_zp".into()],
1745                    vec!["output".into()],
1746                ))
1747                .push_input("x")
1748                .push_input("x")
1749                .push_output("output");
1750            let target = if in_dtype == DType::INT8 {
1751                XDNA_TOSA_INTEGER_TARGET
1752            } else {
1753                XDNA_TOSA_TARGET
1754            };
1755            graph.build(target).expect("build aliased matmul")
1756        }
1757
1758        let bf16 = aliased_matmul(64, DType::BF16, DType::FP32);
1759        assert_eq!(admit(&bf16, XDNA_TOSA_TARGET), Err(AdmitError::Unsupported));
1760        let int8 = aliased_matmul(32, DType::INT8, DType::INT32);
1761        assert_eq!(
1762            admit(&int8, XDNA_TOSA_INTEGER_TARGET),
1763            Err(AdmitError::Unsupported)
1764        );
1765    }
1766
1767    /// Build a fused FP8 MATMUL graph: two FP8 block inputs, each promoted by its own explicit
1768    /// CAST, multiplied to FP32. `escape` additionally exposes the promoted lhs as a block output.
1769    fn fp8_matmul_graph(
1770        m: i32,
1771        k: i32,
1772        n: i32,
1773        lhs_dtype: DType,
1774        rhs_dtype: DType,
1775        alias_operands: bool,
1776        escape: bool,
1777    ) -> OwnedGraph<'static> {
1778        let mut graph = OwnedGraph::new("main");
1779        let rhs_name = if alias_operands { "lhs_fp8" } else { "rhs_fp8" };
1780        graph
1781            .push_tensor(OwnedTensor::new("lhs_fp8", vec![1, m, k], lhs_dtype))
1782            .push_tensor(OwnedTensor::new("lhs_bf16", vec![1, m, k], DType::BF16))
1783            .push_tensor(OwnedTensor::new("rhs_bf16", vec![1, k, n], DType::BF16))
1784            .push_tensor(OwnedTensor::constant(
1785                "lhs_zp",
1786                vec![1],
1787                DType::BF16,
1788                vec![0u8; 2],
1789            ))
1790            .push_tensor(OwnedTensor::constant(
1791                "rhs_zp",
1792                vec![1],
1793                DType::BF16,
1794                vec![0u8; 2],
1795            ))
1796            .push_tensor(OwnedTensor::new("output", vec![1, m, n], DType::FP32));
1797        if !alias_operands {
1798            graph.push_tensor(OwnedTensor::new("rhs_fp8", vec![1, k, n], rhs_dtype));
1799        }
1800        graph
1801            .push_operator(OwnedOperator::new(
1802                OperatorKind::Cast,
1803                vec!["lhs_fp8".into()],
1804                vec!["lhs_bf16".into()],
1805            ))
1806            .push_operator(OwnedOperator::new(
1807                OperatorKind::Cast,
1808                vec![rhs_name.into()],
1809                vec!["rhs_bf16".into()],
1810            ))
1811            .push_operator(OwnedOperator::new(
1812                OperatorKind::Const,
1813                vec![],
1814                vec!["lhs_zp".into()],
1815            ))
1816            .push_operator(OwnedOperator::new(
1817                OperatorKind::Const,
1818                vec![],
1819                vec!["rhs_zp".into()],
1820            ))
1821            .push_operator(OwnedOperator::new(
1822                OperatorKind::MatMul,
1823                vec![
1824                    "lhs_bf16".into(),
1825                    "rhs_bf16".into(),
1826                    "lhs_zp".into(),
1827                    "rhs_zp".into(),
1828                ],
1829                vec!["output".into()],
1830            ))
1831            .push_input("lhs_fp8");
1832        if !alias_operands {
1833            graph.push_input("rhs_fp8");
1834        } else {
1835            graph.push_input("lhs_fp8");
1836        }
1837        graph.push_output("output");
1838        if escape {
1839            graph.push_output("lhs_bf16");
1840        }
1841        graph
1842    }
1843
1844    #[test]
1845    fn admits_fused_fp8_matmul_for_both_encodings() {
1846        for (dtype, format) in [
1847            (DType::FP8E4M3, Fp8Format::E4M3),
1848            (DType::FP8E5M2, Fp8Format::E5M2),
1849        ] {
1850            let bytes = fp8_matmul_graph(32, 64, 32, dtype, dtype, false, false)
1851                .build(XDNA_TOSA_FP8_TARGET)
1852                .expect("build fused fp8 matmul");
1853            assert_eq!(
1854                admit(&bytes, XDNA_TOSA_FP8_TARGET),
1855                Ok(CompilerSpec::Fp8Matmul {
1856                    format,
1857                    m: 32,
1858                    k: 64,
1859                    n: 32
1860                })
1861            );
1862        }
1863    }
1864
1865    /// The compiled kernel instantiates exactly one decoder, so mixing encodings across the two
1866    /// operands must not be admitted under either encoding's label.
1867    #[test]
1868    fn rejects_fused_fp8_matmul_with_mixed_encodings() {
1869        let bytes = fp8_matmul_graph(32, 64, 32, DType::FP8E4M3, DType::FP8E5M2, false, false)
1870            .build(XDNA_TOSA_FP8_TARGET)
1871            .expect("build mixed-encoding fused matmul");
1872        assert_eq!(
1873            admit(&bytes, XDNA_TOSA_FP8_TARGET),
1874            Err(AdmitError::Unsupported)
1875        );
1876    }
1877
1878    /// The promoted BF16 value is graph-interior. If the graph also demands it as a block output,
1879    /// the fused kernel cannot serve it — it never writes BF16 to DDR.
1880    #[test]
1881    fn rejects_fused_fp8_matmul_whose_promoted_operand_escapes() {
1882        let bytes = fp8_matmul_graph(32, 64, 32, DType::FP8E4M3, DType::FP8E4M3, false, true)
1883            .build(XDNA_TOSA_FP8_TARGET)
1884            .expect("build escaping fused matmul");
1885        assert_eq!(
1886            admit(&bytes, XDNA_TOSA_FP8_TARGET),
1887            Err(AdmitError::Unsupported)
1888        );
1889    }
1890
1891    /// Square shape so that one FP8 tensor can legally feed both CASTs; the rejection must then
1892    /// come from admission, not from TOSA shape validation.
1893    #[test]
1894    fn rejects_fused_fp8_matmul_with_one_value_feeding_both_operands() {
1895        let bytes = fp8_matmul_graph(64, 64, 64, DType::FP8E4M3, DType::FP8E4M3, true, false)
1896            .build(XDNA_TOSA_FP8_TARGET)
1897            .expect("build aliased fused matmul");
1898        assert_eq!(
1899            admit(&bytes, XDNA_TOSA_FP8_TARGET),
1900            Err(AdmitError::Unsupported)
1901        );
1902        // Positive control: the same shape with two distinct operands is admitted, so the
1903        // rejection above is specifically the aliasing and not the shape.
1904        let distinct = fp8_matmul_graph(64, 64, 64, DType::FP8E4M3, DType::FP8E4M3, false, false)
1905            .build(XDNA_TOSA_FP8_TARGET)
1906            .expect("build distinct fused matmul");
1907        assert!(admit(&distinct, XDNA_TOSA_FP8_TARGET).is_ok());
1908    }
1909
1910    #[test]
1911    fn rejects_fused_fp8_matmul_off_the_tested_tiling() {
1912        for (m, k, n) in [(48, 64, 32), (32, 64, MATMUL_MAX_DIM as i32 + 32)] {
1913            let bytes = fp8_matmul_graph(m, k, n, DType::FP8E4M3, DType::FP8E4M3, false, false)
1914                .build(XDNA_TOSA_FP8_TARGET)
1915                .expect("build fused matmul");
1916            assert_eq!(
1917                admit(&bytes, XDNA_TOSA_FP8_TARGET),
1918                Err(AdmitError::Unsupported)
1919            );
1920        }
1921    }
1922
1923    /// The fused tier lives on the FP8 target; the BF16 target must not admit an FP8 graph.
1924    #[test]
1925    fn fused_fp8_matmul_is_not_admitted_on_the_bf16_target() {
1926        let bytes = fp8_matmul_graph(32, 64, 32, DType::FP8E4M3, DType::FP8E4M3, false, false)
1927            .build(XDNA_TOSA_FP8_TARGET)
1928            .expect("build fused matmul");
1929        assert_eq!(admit(&bytes, XDNA_TOSA_TARGET), Err(AdmitError::Analysis));
1930    }
1931
1932    #[test]
1933    fn rejects_fp32_matmul_inputs() {
1934        // FP32-input MATMUL is admissible TOSA but has no compute path on this hardware.
1935        let bytes = matmul_graph(32, 64, 32, DType::FP32, DType::FP32)
1936            .build(XDNA_TOSA_TARGET)
1937            .expect("build fp32 matmul");
1938        assert_eq!(
1939            admit(&bytes, XDNA_TOSA_TARGET),
1940            Err(AdmitError::Unsupported)
1941        );
1942    }
1943
1944    #[test]
1945    fn rejects_matmul_shape_off_the_tested_tiling() {
1946        // M not a multiple of the tile, and a dimension past the tested envelope.
1947        for (m, k, n) in [(48, 64, 32), (32, 64, MATMUL_MAX_DIM as i32 + 32)] {
1948            let bytes = matmul_graph(m, k, n, DType::BF16, DType::FP32)
1949                .build(XDNA_TOSA_TARGET)
1950                .expect("build matmul");
1951            assert_eq!(
1952                admit(&bytes, XDNA_TOSA_TARGET),
1953                Err(AdmitError::Unsupported)
1954            );
1955        }
1956    }
1957
1958    #[test]
1959    fn rejects_fp32_identity() {
1960        // FP32 is admissible TOSA under the floating-point profile but outside the BF16 tier.
1961        let bytes = identity_graph(DType::FP32, vec![1, 4, 1024])
1962            .build(XDNA_TOSA_TARGET)
1963            .expect("build fp32 identity");
1964        assert_eq!(
1965            admit(&bytes, XDNA_TOSA_TARGET),
1966            Err(AdmitError::Unsupported)
1967        );
1968    }
1969
1970    #[test]
1971    fn rejects_non_multiple_of_line_size() {
1972        let bytes = identity_graph(DType::BF16, vec![1, 1, 100])
1973            .build(XDNA_TOSA_TARGET)
1974            .expect("build small identity");
1975        assert_eq!(
1976            admit(&bytes, XDNA_TOSA_TARGET),
1977            Err(AdmitError::Unsupported)
1978        );
1979    }
1980
1981    #[test]
1982    fn rejects_bf16_artifact_under_integer_target() {
1983        let bytes = identity_graph(DType::BF16, vec![1, 4, 1024])
1984            .build(XDNA_TOSA_TARGET)
1985            .expect("build");
1986        assert_eq!(
1987            admit(&bytes, XDNA_TOSA_INTEGER_TARGET),
1988            Err(AdmitError::Analysis)
1989        );
1990    }
1991}