Skip to main content

virtio_accel_vulkan/
lower.rs

1//! TOSA admission for the Vulkan backend: the advertised targets, the capability descriptor, and
2//! the hardware-free lowering of an admitted graph into a [`ProgramPlan`] of kernel dispatches.
3//!
4//! This module compiles and unit-tests on every host. It owns every decision about which TOSA
5//! graphs the backend executes and every index computation the kernels will perform; `native`
6//! only turns an accepted plan into Vulkan objects. Because the kernels address storage buffers
7//! with the geometry planned here and no robust-buffer-access mode is relied upon, every shape,
8//! axis, permutation, and pooling window is re-derived and checked against the declared tensor
9//! shapes before a plan is produced: a graph whose declared shapes disagree with its operator
10//! semantics is rejected, never dispatched.
11//!
12//! The FP32 tier (ADR 0004, ADR 0007) admits static single-block graphs over the 42 operators the
13//! Core ML and OpenVINO backends share, with `BOOL` and `INT32` auxiliaries where TOSA defines
14//! them; the FP16 tier (ADR 0008) admits the same graphs over binary16 tensors on every device —
15//! its conversions are crate-owned integer and binary32 code, so it needs no device feature.
16//! `CONST` tensors and intermediates live in a per-program arena; `RESHAPE` and `IDENTITY` of
17//! arena tensors are views, not copies. The provisional integer target stays declared but admits
18//! nothing until its per-device gating is ratified.
19
20// Builds forced to the placeholder (`VIRTIO_ACCEL_VULKAN=0`, or an OS outside the loader host
21// set) still type-check and unit-test this admission path; only the native module calls it.
22#![cfg_attr(not(va_vulkan), allow(dead_code))]
23
24use std::collections::HashMap;
25use std::fmt;
26
27use virtio_accel_tosa::{
28    AnalysisError, AnalyzedValueKind, CapabilityDescriptor, DType, DTypeCapability,
29    DTypeConstraints, Error as ParseError, ExtensionSet, GraphCapabilities, Level,
30    NanPropagationMode, Op, OpAttributes, OperatorCapability, OperatorConstraints, OperatorId,
31    OptimizationHints, ProfileSet, RuntimeCondition, RuntimeConditionSupport, Target, TosaAnalysis,
32    ValueId, ValueRoles, Version, parse,
33};
34
35use crate::shader::{
36    ElementwiseOp, ElementwiseSpec, Fp8Format, MAX_RANK, MoveGeometry, NanMode, Operand,
37    PoolGeometry, ReduceOp, Storage, matmul_spec, max_pool_spec, move_spec, reduce_spec,
38};
39
40/// The FP32 base tier: TOSA 1.0, floating-point profile, level 8K, no extensions.
41pub const VULKAN_TOSA_TARGET: Target = Target::new(
42    Version::TOSA_1_0,
43    ProfileSet::FLOATING_POINT,
44    Level::Level8K,
45    ExtensionSet::NONE,
46);
47
48/// The FP8 tier's target: TOSA 1.0, floating-point profile, level 8K, both FP8 extensions.
49///
50/// Unlike the FP16 tier, this cannot share [`VULKAN_TOSA_TARGET`]'s identity. FP16 lives in the
51/// base floating-point profile, so narrowing dtypes under one target was enough (ADR 0008); FP8
52/// legality is gated on `ExtensionSet::FP8E4M3` / `FP8E5M2`, so the envelope itself differs and
53/// the tier needs a target of its own — the same split `virtio-accel-xdna` makes (ADR 0009).
54pub const VULKAN_TOSA_FP8_TARGET: Target = Target::new(
55    Version::TOSA_1_0,
56    ProfileSet::FLOATING_POINT,
57    Level::Level8K,
58    ExtensionSet::NONE
59        .union(ExtensionSet::FP8E4M3)
60        .union(ExtensionSet::FP8E5M2),
61);
62
63/// The provisional integer tier: TOSA 1.0, integer profile, level 8K, no extensions.
64///
65/// Declared (ADR 0004) but not yet advertised: no capability descriptor names it and admission
66/// rejects it until `shaderInt8` gating and the operator subset table close wayfinder ticket 5.
67pub const VULKAN_TOSA_INTEGER_TARGET: Target = Target::new(
68    Version::TOSA_1_0,
69    ProfileSet::INTEGER,
70    Level::Level8K,
71    ExtensionSet::NONE,
72);
73
74const FLOAT_DTYPES: &[DTypeCapability] = &[
75    DTypeCapability::new(DType::FP32, ValueRoles::ALL),
76    DTypeCapability::new(DType::BOOL, ValueRoles::ALL),
77    DTypeCapability::new(DType::INT32, ValueRoles::ALL),
78    // The TOSA 1.0 `MUL` shift operand is an INT8 constant consumed at admission (it must be
79    // zero); INT8 never becomes a graph-visible provider tensor.
80    DTypeCapability::constrained(
81        DType::INT8,
82        ValueRoles::CONSTANT,
83        DTypeConstraints::PARAMETER_ONLY,
84    ),
85];
86
87/// The FP16 tier's dtypes: the FP32 tier plus binary16 tensors in every role. FP16 shares the
88/// floating-point target identity with FP32 (the tier is dtype-narrowed, exactly the Hexagon
89/// pattern), and it is the descriptor every native instance advertises (ADR 0008).
90const FLOAT16_DTYPES: &[DTypeCapability] = &[
91    DTypeCapability::new(DType::FP32, ValueRoles::ALL),
92    DTypeCapability::new(DType::FP16, ValueRoles::ALL),
93    DTypeCapability::new(DType::BOOL, ValueRoles::ALL),
94    DTypeCapability::new(DType::INT32, ValueRoles::ALL),
95    DTypeCapability::constrained(
96        DType::INT8,
97        ValueRoles::CONSTANT,
98        DTypeConstraints::PARAMETER_ONLY,
99    ),
100];
101
102/// The 42 operators shared with the Core ML and OpenVINO FP32 tiers. NaN modes and pool padding
103/// are unconstrained: the kernels implement both `PROPAGATE` and `IGNORE` literally and exclude
104/// padded taps from the window.
105const FLOAT_OPERATORS: &[OperatorCapability] = &[
106    OperatorCapability::new(Op::ARGMAX),
107    OperatorCapability::constrained(Op::MATMUL, OperatorConstraints::ZERO_ZERO_POINTS),
108    OperatorCapability::new(Op::MAX_POOL2D),
109    OperatorCapability::new(Op::CLAMP),
110    OperatorCapability::new(Op::ERF),
111    OperatorCapability::new(Op::SIGMOID),
112    OperatorCapability::new(Op::TANH),
113    OperatorCapability::new(Op::ADD),
114    OperatorCapability::new(Op::LOGICAL_AND),
115    OperatorCapability::new(Op::LOGICAL_OR),
116    OperatorCapability::new(Op::LOGICAL_XOR),
117    OperatorCapability::new(Op::MAXIMUM),
118    OperatorCapability::new(Op::MINIMUM),
119    OperatorCapability::constrained(Op::MUL, OperatorConstraints::ZERO_SHIFT),
120    OperatorCapability::new(Op::POW),
121    OperatorCapability::new(Op::SUB),
122    OperatorCapability::new(Op::ABS),
123    OperatorCapability::new(Op::CEIL),
124    OperatorCapability::new(Op::COS),
125    OperatorCapability::new(Op::EXP),
126    OperatorCapability::new(Op::FLOOR),
127    OperatorCapability::new(Op::LOG),
128    OperatorCapability::new(Op::LOGICAL_NOT),
129    OperatorCapability::constrained(Op::NEGATE, OperatorConstraints::ZERO_ZERO_POINTS),
130    OperatorCapability::new(Op::RECIPROCAL),
131    OperatorCapability::new(Op::RSQRT),
132    OperatorCapability::new(Op::SIN),
133    OperatorCapability::new(Op::SELECT),
134    OperatorCapability::new(Op::EQUAL),
135    OperatorCapability::new(Op::GREATER),
136    OperatorCapability::new(Op::GREATER_EQUAL),
137    OperatorCapability::new(Op::REDUCE_MAX),
138    OperatorCapability::new(Op::REDUCE_MIN),
139    OperatorCapability::new(Op::REDUCE_PRODUCT),
140    OperatorCapability::new(Op::REDUCE_SUM),
141    OperatorCapability::new(Op::CONCAT),
142    OperatorCapability::constrained(Op::RESHAPE, OperatorConstraints::CONSTANT_PARAMETERS),
143    OperatorCapability::new(Op::REVERSE),
144    OperatorCapability::new(Op::TRANSPOSE),
145    OperatorCapability::new(Op::CONST),
146    OperatorCapability::new(Op::CONST_SHAPE),
147    OperatorCapability::new(Op::IDENTITY),
148];
149
150/// The FP32 tier's admitted boundary: exactly what the generated kernels execute.
151pub const VULKAN_TOSA_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
152    target: VULKAN_TOSA_TARGET,
153    dtypes: FLOAT_DTYPES,
154    operators: FLOAT_OPERATORS,
155    graph: GraphCapabilities {
156        max_regions: 1,
157        max_blocks: 1,
158        dynamic_shapes: false,
159        runtime_conditions: RuntimeConditionSupport::None,
160    },
161};
162
163/// The FP16 tier's admitted boundary (ADR 0008): the same 42 operators and graph envelope with
164/// binary16 tensors in every role. The tier needs no device feature — the kernels' binary16
165/// conversions are crate-owned integer and binary32 code, the implementation choice TOSA 1.0
166/// §1.10.3 names explicitly — so it is advertised on every device the backend opens, with
167/// numerics identical to the FP32 tier's everywhere.
168pub const VULKAN_TOSA_FP16_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
169    target: VULKAN_TOSA_TARGET,
170    dtypes: FLOAT16_DTYPES,
171    operators: FLOAT_OPERATORS,
172    graph: GraphCapabilities {
173        max_regions: 1,
174        max_blocks: 1,
175        dynamic_shapes: false,
176        runtime_conditions: RuntimeConditionSupport::None,
177    },
178};
179
180/// The FP8 tier's dtypes: both FP8 encodings in every role, plus the FP16 that TOSA assigns as
181/// the result of an FP8 MATMUL, and INT32 for shape operands.
182const FLOAT8_DTYPES: &[DTypeCapability] = &[
183    DTypeCapability::new(DType::FP8E4M3, ValueRoles::ALL),
184    DTypeCapability::new(DType::FP8E5M2, ValueRoles::ALL),
185    DTypeCapability::new(DType::FP16, ValueRoles::ALL),
186    DTypeCapability::new(DType::INT32, ValueRoles::ALL),
187];
188
189/// The FP8 tier's operators: what TOSA admits for FP8 *and* this crate executes. Deliberately a
190/// subset, not [`FLOAT_OPERATORS`] — TOSA admits no FP8 elementwise operator at all (no
191/// arithmetic, comparison, selection, reduction or transcendental lane takes FP8), so the tier
192/// is MATMUL plus the data movement that feeds it. A capability descriptor cannot express
193/// per-operator dtype legality, so listing the 42-operator table here would advertise FP8 lanes
194/// that admission would then reject.
195///
196/// `CAST`, `MAX_POOL2D` and `ARGMAX` are admitted by TOSA for FP8 and are deliberately absent:
197/// they need kernels this tier does not yet carry.
198const FLOAT8_OPERATORS: &[OperatorCapability] = &[
199    OperatorCapability::constrained(Op::MATMUL, OperatorConstraints::ZERO_ZERO_POINTS),
200    OperatorCapability::new(Op::CONCAT),
201    OperatorCapability::constrained(Op::RESHAPE, OperatorConstraints::CONSTANT_PARAMETERS),
202    OperatorCapability::new(Op::REVERSE),
203    OperatorCapability::new(Op::TRANSPOSE),
204    OperatorCapability::new(Op::CONST),
205    OperatorCapability::new(Op::CONST_SHAPE),
206    OperatorCapability::new(Op::IDENTITY),
207    OperatorCapability::new(Op::CAST),
208    OperatorCapability::new(Op::MAX_POOL2D),
209    OperatorCapability::new(Op::ARGMAX),
210];
211
212/// The FP8 tier's admitted boundary (ADR 0009): `(FP8, FP8) -> FP16` MATMUL and exact FP8 data
213/// movement. Like the FP16 tier it needs no device feature — the widening is crate-owned integer
214/// and binary32 code and nothing writes FP8 except a raw byte copy — so it is advertised on
215/// every device the backend opens, with numerics identical everywhere.
216pub const VULKAN_TOSA_FP8_CAPABILITY: CapabilityDescriptor = CapabilityDescriptor {
217    target: VULKAN_TOSA_FP8_TARGET,
218    dtypes: FLOAT8_DTYPES,
219    operators: FLOAT8_OPERATORS,
220    graph: GraphCapabilities {
221        max_regions: 1,
222        max_blocks: 1,
223        dynamic_shapes: false,
224        runtime_conditions: RuntimeConditionSupport::None,
225    },
226};
227
228/// Whether the FP32 tier admits `op`.
229pub const fn supports_tosa_operator(op: Op) -> bool {
230    VULKAN_TOSA_CAPABILITY.supports_operator(op)
231}
232
233/// Whether the advertised FP16 tier exposes `dtype` at a program boundary.
234pub const fn supports_tosa_dtype(dtype: DType) -> bool {
235    VULKAN_TOSA_FP16_CAPABILITY.supports_dtype(dtype, ValueRoles::INPUT)
236        || VULKAN_TOSA_FP16_CAPABILITY.supports_dtype(dtype, ValueRoles::OUTPUT)
237        || VULKAN_TOSA_FP8_CAPABILITY.supports_dtype(dtype, ValueRoles::INPUT)
238        || VULKAN_TOSA_FP8_CAPABILITY.supports_dtype(dtype, ValueRoles::OUTPUT)
239}
240
241/// Why an artifact was not admitted.
242#[derive(Clone, Copy, Debug, PartialEq, Eq)]
243pub enum LoweringError {
244    Parse(ParseError),
245    Analysis(AnalysisError),
246    /// The target is not one this backend advertises.
247    UnsupportedTarget,
248    /// The graph shape (regions, blocks, boundary, operator structure, attributes, or declared
249    /// tensor shapes) is outside the tier.
250    UnsupportedGraph,
251    UnsupportedType(DType),
252    UnsupportedOperator(Op),
253    /// A static shape does not fit the kernels' 32-bit element domain, or the program needs more
254    /// arena or slots than the plan can express.
255    ResourceLimit,
256}
257
258impl fmt::Display for LoweringError {
259    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
260        write!(formatter, "{self:?}")
261    }
262}
263
264impl std::error::Error for LoweringError {}
265
266/// Arena regions are aligned generously so every tensor start is cache-line aligned and any
267/// storage-buffer offset alignment a device reports (at most 256) is satisfied.
268pub(crate) const ARENA_ALIGNMENT: u64 = 256;
269
270/// Device-independent kernel selection; `native` maps it onto a [`crate::shader::KernelKey`]
271/// once the device's workgroup tuning and descriptor-array length are known.
272#[derive(Clone, Copy, Debug, PartialEq, Eq)]
273pub(crate) enum KernelSpec {
274    /// Native packed E2M1 weights with one E4M3 scale per sixteen weights.
275    Nvfp4Matmul {
276        /// Shared weights and enough tokens make the cooperative-matrix path eligible. Device
277        /// tuning still falls back when the exact matrix shape is unavailable.
278        cooperative: bool,
279    },
280    Elementwise {
281        op: ElementwiseOp,
282        float: Storage,
283        broadcast: bool,
284    },
285    Reduce {
286        op: ReduceOp,
287        float: Storage,
288    },
289    Matmul {
290        input: Storage,
291        output: Storage,
292    },
293    /// The streaming split-`k` MATMUL, selected when `m ≤ shader::STREAM_ROWS`; its lhs is
294    /// always binary32 words (widened by an inserted `Cast` when the tensor is narrower).
295    MatmulStream {
296        rhs: Storage,
297        output: Storage,
298    },
299    Cast {
300        input: Storage,
301        output: Storage,
302    },
303    MaxPool {
304        nan_mode: NanMode,
305        float: Storage,
306    },
307    Move {
308        storage: Storage,
309        contiguous: bool,
310    },
311}
312
313/// Dispatch geometry before device tuning.
314#[derive(Clone, Copy, Debug, PartialEq, Eq)]
315pub(crate) enum Work {
316    /// A grid-stride kernel over this many items.
317    Linear(u32),
318    /// A tiled MATMUL over `m × n` outputs per batch.
319    Matmul { m: u32, n: u32, batch: u32 },
320    /// A streaming MATMUL over `n` columns per batch.
321    MatmulStream { n: u32, batch: u32 },
322    /// One workgroup reduces each `[token, output-row]` dot product.
323    Nvfp4Matmul { m: u32, n: u32 },
324}
325
326/// One recorded `vkCmdDispatch`.
327#[derive(Clone, Debug, PartialEq, Eq)]
328pub(crate) struct DispatchPlan {
329    pub kernel: KernelSpec,
330    /// Specialization constants in the kernel's declared order.
331    pub spec: Vec<u32>,
332    pub work: Work,
333    /// A compute→compute memory barrier must precede this dispatch: it reads a tensor an
334    /// earlier dispatch wrote, or writes memory an earlier dispatch touched.
335    pub barrier_before: bool,
336}
337
338/// One binding slot of an admitted program.
339#[derive(Clone, Copy, Debug, PartialEq, Eq)]
340pub(crate) struct SlotPlan {
341    pub slot: u32,
342    pub role: SlotRole,
343    /// Exact tensor bytes: the required length of a binding over this slot.
344    pub byte_len: u64,
345    pub storage: Storage,
346}
347
348#[derive(Clone, Copy, Debug, PartialEq, Eq)]
349pub(crate) enum SlotRole {
350    Input,
351    Output,
352}
353
354/// A serialized constant to upload into the arena at `load_program`.
355#[derive(Clone, Debug, PartialEq, Eq)]
356pub(crate) struct ConstantPlan {
357    pub offset: u64,
358    pub bytes: Vec<u8>,
359}
360
361/// A hardware-free execution plan for one admitted TOSA graph.
362///
363/// Slots follow the workspace convention shared with the other TOSA backends: block inputs take
364/// slots `0..inputs`, block outputs follow in declared order. Bound slots occupy descriptor
365/// array elements `0..slots.len()`; the arena, when present, sits at index `slots.len()`.
366#[derive(Clone, Debug, PartialEq, Eq)]
367pub(crate) struct ProgramPlan {
368    pub slots: Vec<SlotPlan>,
369    /// Bytes of program-owned storage (constants and intermediates); zero when none is needed.
370    pub arena_bytes: u64,
371    pub constants: Vec<ConstantPlan>,
372    pub dispatches: Vec<DispatchPlan>,
373}
374
375impl ProgramPlan {
376    /// Descriptor-array element the arena is bound at.
377    pub fn arena_buffer_index(&self) -> u32 {
378        self.slots.len() as u32
379    }
380
381    /// Descriptor-array elements the program addresses.
382    pub fn buffer_count(&self) -> u32 {
383        self.slots.len() as u32 + u32::from(self.arena_bytes != 0)
384    }
385
386    #[cfg(test)]
387    fn slot(&self, slot: u32) -> Option<&SlotPlan> {
388        self.slots.iter().find(|plan| plan.slot == slot)
389    }
390}
391
392/// Admit `bytes` for `target` and produce its plan, or explain the rejection.
393pub(crate) fn lower_tosa(bytes: &[u8], target: Target) -> Result<ProgramPlan, LoweringError> {
394    if target != VULKAN_TOSA_TARGET && target != VULKAN_TOSA_FP8_TARGET {
395        return Err(LoweringError::UnsupportedTarget);
396    }
397    let model = parse(bytes).map_err(LoweringError::Parse)?;
398    let analysis = model.analyze_for(target).map_err(LoweringError::Analysis)?;
399    let capability = if target == VULKAN_TOSA_FP8_TARGET {
400        VULKAN_TOSA_FP8_CAPABILITY
401    } else {
402        VULKAN_TOSA_CAPABILITY
403    };
404    Lowering::new(&analysis, capability)?.run()
405}
406
407// ---------------------------------------------------------------------------------------------
408// Tensor bookkeeping
409// ---------------------------------------------------------------------------------------------
410
411/// Static shape summary of a tensor value.
412#[derive(Clone, Debug, PartialEq, Eq)]
413struct TensorShape {
414    dtype: DType,
415    dims: Vec<u32>,
416    elements: u32,
417}
418
419impl TensorShape {
420    fn storage(&self) -> Storage {
421        storage_of(self.dtype)
422    }
423
424    fn byte_len(&self) -> u64 {
425        u64::from(self.elements) * scalar_bytes(self.dtype)
426    }
427
428    fn rank(&self) -> usize {
429        self.dims.len()
430    }
431
432    /// Row-major element strides.
433    fn strides(&self) -> Vec<u32> {
434        let mut strides = vec![1_u32; self.dims.len()];
435        let mut acc = 1_u32;
436        for d in (0..self.dims.len()).rev() {
437            strides[d] = acc;
438            acc = acc.wrapping_mul(self.dims[d]);
439        }
440        strides
441    }
442
443    /// Dims padded with leading ones to `MAX_RANK`.
444    fn padded_dims(&self) -> [u32; MAX_RANK] {
445        pad_leading(&self.dims, 1)
446    }
447}
448
449fn pad_leading(values: &[u32], fill: u32) -> [u32; MAX_RANK] {
450    let mut padded = [fill; MAX_RANK];
451    let offset = MAX_RANK - values.len();
452    padded[offset..].copy_from_slice(values);
453    padded
454}
455
456fn storage_of(dtype: DType) -> Storage {
457    match dtype {
458        DType::BOOL => Storage::Byte,
459        DType::FP8E4M3 => Storage::Quarter(Fp8Format::E4M3),
460        DType::FP8E5M2 => Storage::Quarter(Fp8Format::E5M2),
461        DType::FP16 => Storage::Half,
462        _ => Storage::Word,
463    }
464}
465
466fn scalar_bytes(dtype: DType) -> u64 {
467    match dtype {
468        DType::BOOL | DType::FP8E4M3 | DType::FP8E5M2 => 1,
469        DType::FP16 => 2,
470        _ => 4,
471    }
472}
473
474/// Where a tensor's bytes live during execution.
475#[derive(Clone, Copy, Debug, PartialEq, Eq)]
476enum Location {
477    Slot(u32),
478    /// An arena region (index into `Lowering::regions`).
479    Region(usize),
480}
481
482/// Memory identity for hazard tracking. Arena tensors are keyed by the bytes they occupy, not by
483/// region index: lifetime packing hands a freed region's bytes to later tensors, possibly
484/// straddling several earlier regions, and a dispatch touching any of those bytes must be ordered
485/// after every earlier dispatch that touched them.
486#[derive(Clone, Copy, Debug, PartialEq, Eq)]
487enum MemKey {
488    Slot(u32),
489    Arena { offset: u64, end: u64 },
490}
491
492impl MemKey {
493    fn overlaps(self, other: Self) -> bool {
494        match (self, other) {
495            (Self::Slot(a), Self::Slot(b)) => a == b,
496            (
497                Self::Arena { offset, end },
498                Self::Arena {
499                    offset: other_offset,
500                    end: other_end,
501                },
502            ) => offset < other_end && other_offset < end,
503            _ => false,
504        }
505    }
506}
507
508#[derive(Clone, Copy, Debug)]
509struct Region {
510    offset: u64,
511    bytes: u64,
512    /// Last execution position at which the region is read; the region is free afterwards.
513    live_end: u32,
514}
515
516struct Lowering<'a, 'b> {
517    analysis: &'b TosaAnalysis<'a>,
518    /// The tier being lowered. Admission consults this tier's operator list, so what a target
519    /// admits is exactly what its descriptor advertises.
520    capability: CapabilityDescriptor,
521    inputs: &'b [ValueId],
522    outputs: &'b [ValueId],
523    order: &'b [OperatorId],
524    shapes: HashMap<ValueId, TensorShape>,
525    locations: HashMap<ValueId, Location>,
526    /// Last execution position at which each value is consumed by a live operator.
527    last_use: HashMap<ValueId, u32>,
528    regions: Vec<Region>,
529    arena_bytes: u64,
530    constants: Vec<ConstantPlan>,
531    dispatches: Vec<DispatchPlan>,
532    slots: Vec<SlotPlan>,
533    /// Hazard tracking since the last barrier: memory written (with the writing value) and read.
534    written: Vec<(MemKey, ValueId)>,
535    read: Vec<MemKey>,
536    position: u32,
537}
538
539impl<'a, 'b> Lowering<'a, 'b> {
540    fn new(
541        analysis: &'b TosaAnalysis<'a>,
542        capability: CapabilityDescriptor,
543    ) -> Result<Self, LoweringError> {
544        if analysis.regions().len() != 1 || analysis.blocks().len() != 1 {
545            return Err(LoweringError::UnsupportedGraph);
546        }
547        // `PowDomain` is the IEEE NaN case the POW kernel produces itself; every other runtime
548        // condition needs error detection this backend does not perform.
549        if analysis
550            .conditions()
551            .iter()
552            .any(|condition| !matches!(condition, RuntimeCondition::PowDomain { .. }))
553        {
554            return Err(LoweringError::UnsupportedGraph);
555        }
556        let block = analysis.blocks()[0].id();
557        let inputs = analysis.block_inputs(block);
558        let outputs = analysis.block_outputs(block);
559        let order = analysis.execution_order(block);
560        if inputs.is_empty()
561            || outputs.is_empty()
562            || inputs.iter().any(|input| outputs.contains(input))
563            || outputs
564                .iter()
565                .any(|output| analysis.serialized_constant(*output).is_some())
566        {
567            return Err(LoweringError::UnsupportedGraph);
568        }
569        let mut duplicates = inputs.iter().chain(outputs).collect::<Vec<_>>();
570        duplicates.sort_unstable_by_key(|value| value.get());
571        if duplicates.windows(2).any(|pair| pair[0] == pair[1]) {
572            return Err(LoweringError::UnsupportedGraph);
573        }
574        if inputs.len() + outputs.len() > u32::MAX as usize {
575            return Err(LoweringError::ResourceLimit);
576        }
577        Ok(Self {
578            analysis,
579            capability,
580            inputs,
581            outputs,
582            order,
583            shapes: HashMap::new(),
584            locations: HashMap::new(),
585            last_use: HashMap::new(),
586            regions: Vec::new(),
587            arena_bytes: 0,
588            constants: Vec::new(),
589            dispatches: Vec::new(),
590            slots: Vec::new(),
591            written: Vec::new(),
592            read: Vec::new(),
593            position: 0,
594        })
595    }
596
597    fn run(mut self) -> Result<ProgramPlan, LoweringError> {
598        // Boundary slots: inputs first, outputs after, in declared order.
599        for (index, value) in self.inputs.iter().chain(self.outputs).enumerate() {
600            let slot = index as u32;
601            let shape = self.shape(*value)?;
602            self.slots.push(SlotPlan {
603                slot,
604                role: if index < self.inputs.len() {
605                    SlotRole::Input
606                } else {
607                    SlotRole::Output
608                },
609                byte_len: shape.byte_len(),
610                storage: shape.storage(),
611            });
612            self.locations.insert(*value, Location::Slot(slot));
613        }
614
615        // Liveness: the last live consumer of every value, and how many consumers it has.
616        let mut consumers: HashMap<ValueId, u32> = HashMap::new();
617        for (position, operator_id) in self.order.iter().enumerate() {
618            let operator = self.analysis.operator(*operator_id);
619            if self.skipped(operator) {
620                continue;
621            }
622            for input in self.analysis.operator_inputs(*operator_id) {
623                self.last_use.insert(*input, position as u32);
624                *consumers.entry(*input).or_default() += 1;
625            }
626        }
627        self.alias_outputs(&consumers)?;
628
629        for (position, operator_id) in self.order.iter().enumerate() {
630            self.position = position as u32;
631            let operator = self.analysis.operator(*operator_id);
632            if self.skipped(operator) {
633                continue;
634            }
635            let op = operator.op();
636            if !self.capability.supports_operator(op) {
637                return Err(LoweringError::UnsupportedOperator(op));
638            }
639            let operator_inputs = self.analysis.operator_inputs(*operator_id);
640            let operator_outputs = self.analysis.operator_outputs(*operator_id);
641            let [output] = operator_outputs else {
642                return Err(LoweringError::UnsupportedGraph);
643            };
644            let output = *output;
645            if self.analysis.serialized_constant(output).is_some() {
646                return Err(LoweringError::UnsupportedGraph);
647            }
648            match op {
649                Op::IDENTITY => self.lower_copy(operator_inputs, output, 1)?,
650                Op::RESHAPE => self.lower_reshape(operator_inputs, output)?,
651                Op::TRANSPOSE => self.lower_transpose(operator, operator_inputs, output)?,
652                Op::REVERSE => self.lower_reverse(operator, operator_inputs, output)?,
653                Op::CONCAT => self.lower_concat(operator, operator_inputs, output)?,
654                Op::CAST => self.lower_cast(operator_inputs, output)?,
655                Op::MATMUL => self.lower_matmul(operator_inputs, output)?,
656                Op::MAX_POOL2D => self.lower_max_pool(operator, operator_inputs, output)?,
657                Op::ARGMAX
658                | Op::REDUCE_MAX
659                | Op::REDUCE_MIN
660                | Op::REDUCE_PRODUCT
661                | Op::REDUCE_SUM => self.lower_reduce(operator, operator_inputs, output)?,
662                _ => self.lower_elementwise(operator, operator_inputs, output)?,
663            }
664        }
665
666        // Every block output must have been produced by a live operator.
667        for output in self.outputs {
668            if self.analysis.value(*output).producer().is_none() {
669                return Err(LoweringError::UnsupportedGraph);
670            }
671        }
672        Ok(ProgramPlan {
673            slots: self.slots,
674            arena_bytes: self.arena_bytes,
675            constants: self.constants,
676            dispatches: self.dispatches,
677        })
678    }
679
680    /// Producers of serialized constants and operators the analysis proved dead emit nothing.
681    fn skipped(&self, operator: &virtio_accel_tosa::AnalyzedOperator<'_>) -> bool {
682        matches!(operator.op(), Op::CONST | Op::CONST_SHAPE)
683            || operator.hints().contains(OptimizationHints::DEAD)
684    }
685
686    // -- shapes -------------------------------------------------------------------------------
687
688    fn shape(&mut self, value: ValueId) -> Result<TensorShape, LoweringError> {
689        if let Some(shape) = self.shapes.get(&value) {
690            return Ok(shape.clone());
691        }
692        let AnalyzedValueKind::Tensor(tensor) = self.analysis.value(value).kind() else {
693            return Err(LoweringError::UnsupportedGraph);
694        };
695        let dtype = tensor.dtype();
696        if !matches!(
697            dtype,
698            DType::FP32
699                | DType::FP16
700                | DType::BOOL
701                | DType::INT32
702                | DType::FP8E4M3
703                | DType::FP8E5M2
704        ) {
705            return Err(LoweringError::UnsupportedType(dtype));
706        }
707        let rank = tensor.rank().ok_or(LoweringError::UnsupportedGraph)?;
708        if rank > MAX_RANK {
709            return Err(LoweringError::UnsupportedGraph);
710        }
711        let mut dims = Vec::with_capacity(rank);
712        let mut elements = 1_u64;
713        for dimension in tensor.dimensions() {
714            let dimension = u32::try_from(dimension)
715                .ok()
716                .filter(|dimension| *dimension > 0)
717                .ok_or(LoweringError::UnsupportedGraph)?;
718            dims.push(dimension);
719            elements = elements
720                .checked_mul(u64::from(dimension))
721                .ok_or(LoweringError::ResourceLimit)?;
722        }
723        let elements = u32::try_from(elements).map_err(|_| LoweringError::ResourceLimit)?;
724        let shape = TensorShape {
725            dtype,
726            dims,
727            elements,
728        };
729        self.shapes.insert(value, shape.clone());
730        Ok(shape)
731    }
732
733    /// The shape of a floating-point tensor (`FP32`, or `FP16` where the tier is advertised).
734    fn float_shape(&mut self, value: ValueId) -> Result<TensorShape, LoweringError> {
735        let shape = self.shape(value)?;
736        if !matches!(shape.dtype, DType::FP32 | DType::FP16) {
737            return Err(LoweringError::UnsupportedType(shape.dtype));
738        }
739        Ok(shape)
740    }
741
742    /// The shape of a float operand the tier can load and store: `FP32`/`FP16`, or either FP8
743    /// encoding under the FP8 target. Distinct from [`float_shape`](Self::float_shape), which is
744    /// the narrower set TOSA admits for float *arithmetic*.
745    fn convertible_float_shape(&mut self, value: ValueId) -> Result<TensorShape, LoweringError> {
746        let shape = self.shape(value)?;
747        if !matches!(
748            shape.dtype,
749            DType::FP32 | DType::FP16 | DType::FP8E4M3 | DType::FP8E5M2
750        ) {
751            return Err(LoweringError::UnsupportedType(shape.dtype));
752        }
753        Ok(shape)
754    }
755
756    fn typed_shape(&mut self, value: ValueId, dtype: DType) -> Result<TensorShape, LoweringError> {
757        let shape = self.shape(value)?;
758        if shape.dtype != dtype {
759            return Err(LoweringError::UnsupportedType(shape.dtype));
760        }
761        Ok(shape)
762    }
763
764    // -- arena --------------------------------------------------------------------------------
765
766    /// Allocate a region of `bytes` live through `live_end` at the lowest offset free of every
767    /// overlapping-lifetime region. The region is first written by the dispatch at the current
768    /// position, so regions whose lifetime ended earlier are free.
769    fn allocate_region(&mut self, bytes: u64, live_end: u32) -> Result<usize, LoweringError> {
770        self.allocate_region_written_at(bytes, live_end, self.position)
771    }
772
773    /// [`allocate_region`](Self::allocate_region) for bytes first written at `written_at`
774    /// rather than at the current position. A region whose lifetime ended before `written_at`
775    /// is free; one that ends at or after it is still overlapped, because the dispatch that
776    /// writes it runs after the new bytes exist.
777    fn allocate_region_written_at(
778        &mut self,
779        bytes: u64,
780        live_end: u32,
781        written_at: u32,
782    ) -> Result<usize, LoweringError> {
783        let bytes = bytes.max(1).div_ceil(ARENA_ALIGNMENT) * ARENA_ALIGNMENT;
784        let mut candidates = vec![0_u64];
785        let overlapping: Vec<Region> = self
786            .regions
787            .iter()
788            .filter(|region| region.live_end >= written_at)
789            .copied()
790            .collect();
791        for region in &overlapping {
792            candidates.push(region.offset + region.bytes);
793        }
794        candidates.sort_unstable();
795        let offset = candidates
796            .into_iter()
797            .find(|candidate| {
798                let end = candidate + bytes;
799                overlapping.iter().all(|region| {
800                    end <= region.offset || *candidate >= region.offset + region.bytes
801                })
802            })
803            .ok_or(LoweringError::ResourceLimit)?;
804        let end = offset
805            .checked_add(bytes)
806            .ok_or(LoweringError::ResourceLimit)?;
807        // Word offsets must fit the kernels' `u32` operand base.
808        if end / 4 > u64::from(u32::MAX) {
809            return Err(LoweringError::ResourceLimit);
810        }
811        self.arena_bytes = self.arena_bytes.max(end);
812        self.regions.push(Region {
813            offset,
814            bytes,
815            live_end,
816        });
817        Ok(self.regions.len() - 1)
818    }
819
820    /// A program output produced by `IDENTITY` or `RESHAPE` from an intermediate that nothing
821    /// else reads is the intermediate under another shape: its producer writes the output slot
822    /// directly and the copy operator becomes a no-op (ADR 0012). Inputs, outputs and constants
823    /// already have their bytes elsewhere and are left alone; a second consumer would need the
824    /// intermediate to outlive the output slot's binding contract, so it is left alone too.
825    fn alias_outputs(&mut self, consumers: &HashMap<ValueId, u32>) -> Result<(), LoweringError> {
826        for output in self.outputs {
827            let Some(producer) = self.analysis.value(*output).producer() else {
828                continue;
829            };
830            let operator = self.analysis.operator(producer);
831            if self.skipped(operator) || !matches!(operator.op(), Op::IDENTITY | Op::RESHAPE) {
832                continue;
833            }
834            let Some(source) = self.analysis.operator_inputs(producer).first().copied() else {
835                continue;
836            };
837            if self.locations.contains_key(&source)
838                || self.analysis.serialized_constant(source).is_some()
839                || consumers.get(&source) != Some(&1)
840            {
841                continue;
842            }
843            let (from, to) = (self.shape(source)?, self.shape(*output)?);
844            if from.dtype != to.dtype || from.elements != to.elements {
845                continue;
846            }
847            let slot = self.locations[output];
848            self.locations.insert(source, slot);
849        }
850        Ok(())
851    }
852
853    fn value_last_use(&self, value: ValueId) -> u32 {
854        self.last_use.get(&value).copied().unwrap_or(self.position)
855    }
856
857    /// The location of an operator input: a bound slot, an arena intermediate, or a constant
858    /// uploaded into the arena on first use.
859    fn input_location(&mut self, value: ValueId) -> Result<Location, LoweringError> {
860        if let Some(location) = self.locations.get(&value) {
861            return Ok(*location);
862        }
863        let shape = self.shape(value)?;
864        let bytes = self
865            .analysis
866            .serialized_constant(value)
867            .ok_or(LoweringError::UnsupportedGraph)?;
868        if bytes.len() as u64 != shape.byte_len() {
869            return Err(LoweringError::UnsupportedGraph);
870        }
871        // Constants are uploaded at `load_program`, before the first dispatch, and stay live
872        // for the whole program: their bytes are written at position 0, so a region an
873        // earlier intermediate has already died in is *not* free — the dispatch that wrote
874        // that intermediate runs after the upload and would overwrite the constant.
875        let region = self.allocate_region_written_at(shape.byte_len(), u32::MAX, 0)?;
876        self.constants.push(ConstantPlan {
877            offset: self.regions[region].offset,
878            bytes: bytes.to_vec(),
879        });
880        let location = Location::Region(region);
881        self.locations.insert(value, location);
882        Ok(location)
883    }
884
885    /// The location an operator output is written to: its slot, or a fresh arena region.
886    fn output_location(&mut self, value: ValueId) -> Result<Location, LoweringError> {
887        if let Some(location) = self.locations.get(&value) {
888            return Ok(*location);
889        }
890        let shape = self.shape(value)?;
891        let live_end = self.value_last_use(value);
892        let region = self.allocate_region(shape.byte_len(), live_end)?;
893        let location = Location::Region(region);
894        self.locations.insert(value, location);
895        Ok(location)
896    }
897
898    fn operand(&self, location: Location) -> Operand {
899        match location {
900            Location::Slot(slot) => Operand {
901                buffer: slot,
902                base: 0,
903            },
904            Location::Region(index) => Operand {
905                buffer: self.slots.len() as u32,
906                base: (self.regions[index].offset / 4) as u32,
907            },
908        }
909    }
910
911    // -- dispatch recording -------------------------------------------------------------------
912
913    fn mem_key(&self, location: Location) -> MemKey {
914        match location {
915            Location::Slot(slot) => MemKey::Slot(slot),
916            Location::Region(index) => {
917                let region = &self.regions[index];
918                MemKey::Arena {
919                    offset: region.offset,
920                    end: region.offset + region.bytes,
921                }
922            }
923        }
924    }
925
926    fn dispatch(
927        &mut self,
928        kernel: KernelSpec,
929        spec: Vec<u32>,
930        work: Work,
931        reads: &[Location],
932        writes: Location,
933        written_value: ValueId,
934    ) {
935        let reads: Vec<MemKey> = reads.iter().map(|read| self.mem_key(*read)).collect();
936        let writes = self.mem_key(writes);
937        let raw = reads.iter().any(|read| {
938            self.written
939                .iter()
940                .any(|(written, _)| written.overlaps(*read))
941        });
942        let war = self.read.iter().any(|read| read.overlaps(writes));
943        // Segments of one `CONCAT` write disjoint parts of one tensor and need no ordering; any
944        // other overlap with earlier written bytes does.
945        let waw = self
946            .written
947            .iter()
948            .any(|(written, value)| written.overlaps(writes) && *value != written_value);
949        let barrier_before = raw || war || waw;
950        if barrier_before {
951            self.written.clear();
952            self.read.clear();
953        }
954        self.read.extend_from_slice(&reads);
955        self.written.push((writes, written_value));
956        self.dispatches.push(DispatchPlan {
957            kernel,
958            spec,
959            work,
960            barrier_before,
961        });
962    }
963
964    // -- operators ----------------------------------------------------------------------------
965
966    /// `IDENTITY` (and any copy of `expected_inputs` operands whose first is the source): a
967    /// view whenever the result is an intermediate — over an arena region, whose lifetime the
968    /// view extends, or directly over a bound input slot — and a contiguous copy only when the
969    /// result is a program output that no producer could write directly (ADR 0012).
970    fn lower_copy(
971        &mut self,
972        inputs: &[ValueId],
973        output: ValueId,
974        expected_inputs: usize,
975    ) -> Result<(), LoweringError> {
976        if inputs.len() != expected_inputs {
977            return Err(LoweringError::UnsupportedGraph);
978        }
979        let source = self.shape(inputs[0])?;
980        let target = self.shape(output)?;
981        if source.dtype != target.dtype || source.elements != target.elements {
982            return Err(LoweringError::UnsupportedGraph);
983        }
984        let from = self.input_location(inputs[0])?;
985        if !self.locations.contains_key(&output) {
986            // An intermediate: alias the source's bytes. An arena region's lifetime extends
987            // over the view; a bound input slot is read-only for the whole submission.
988            if let Location::Region(region) = from {
989                let live_end = self.value_last_use(output);
990                self.regions[region].live_end = self.regions[region].live_end.max(live_end);
991            }
992            self.locations.insert(output, from);
993            return Ok(());
994        }
995        let to = self.output_location(output)?;
996        if to == from {
997            // The source was aliased to this output slot ahead of time (`alias_outputs`) and
998            // its producer has already written the bytes where they belong.
999            return Ok(());
1000        }
1001        let storage = source.storage();
1002        let geometry = MoveGeometry {
1003            count: source.elements,
1004            dims: source.padded_dims(),
1005            in_strides: pad_leading(&source.strides(), 0),
1006            in_offset: 0,
1007            out_strides: pad_leading(&source.strides(), 0),
1008            out_offset: 0,
1009        };
1010        let spec = move_spec(self.operand(from), self.operand(to), geometry, true);
1011        // A contiguous copy runs one invocation per destination word (`shader`), so the work
1012        // item count is words, not elements.
1013        self.dispatch(
1014            KernelSpec::Move {
1015                storage,
1016                contiguous: true,
1017            },
1018            spec,
1019            Work::Linear(source.elements.div_ceil(storage.lanes())),
1020            &[from],
1021            to,
1022            output,
1023        );
1024        Ok(())
1025    }
1026
1027    fn lower_reshape(&mut self, inputs: &[ValueId], output: ValueId) -> Result<(), LoweringError> {
1028        let [source, shape] = inputs else {
1029            return Err(LoweringError::UnsupportedGraph);
1030        };
1031        let AnalyzedValueKind::Shape(shape) = self.analysis.value(*shape).kind() else {
1032            return Err(LoweringError::UnsupportedGraph);
1033        };
1034        let values = shape.values().ok_or(LoweringError::UnsupportedGraph)?;
1035        let target = self.shape(output)?;
1036        let declared: Vec<i64> = values.collect();
1037        if declared.len() != target.rank()
1038            || declared
1039                .iter()
1040                .zip(&target.dims)
1041                .any(|(declared, dim)| *declared != i64::from(*dim))
1042        {
1043            return Err(LoweringError::UnsupportedGraph);
1044        }
1045        self.lower_copy(&[*source], output, 1)
1046    }
1047
1048    fn lower_transpose(
1049        &mut self,
1050        operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1051        inputs: &[ValueId],
1052        output: ValueId,
1053    ) -> Result<(), LoweringError> {
1054        let [input] = inputs else {
1055            return Err(LoweringError::UnsupportedGraph);
1056        };
1057        let OpAttributes::Transpose { perms } = operator.source().attributes() else {
1058            return Err(LoweringError::UnsupportedGraph);
1059        };
1060        let source = self.shape(*input)?;
1061        let target = self.shape(output)?;
1062        let perms: Vec<usize> = perms
1063            .iter()
1064            .map(|perm| usize::try_from(perm).map_err(|_| LoweringError::UnsupportedGraph))
1065            .collect::<Result<_, _>>()?;
1066        if perms.len() != source.rank()
1067            || target.rank() != source.rank()
1068            || source.dtype != target.dtype
1069        {
1070            return Err(LoweringError::UnsupportedGraph);
1071        }
1072        let mut seen = vec![false; perms.len()];
1073        for (d, perm) in perms.iter().enumerate() {
1074            if *perm >= perms.len() || seen[*perm] || target.dims[d] != source.dims[*perm] {
1075                return Err(LoweringError::UnsupportedGraph);
1076            }
1077            seen[*perm] = true;
1078        }
1079        let source_strides = source.strides();
1080        let in_strides: Vec<u32> = perms.iter().map(|perm| source_strides[*perm]).collect();
1081        let geometry = MoveGeometry {
1082            count: target.elements,
1083            dims: target.padded_dims(),
1084            in_strides: pad_leading(&in_strides, 0),
1085            in_offset: 0,
1086            out_strides: pad_leading(&target.strides(), 0),
1087            out_offset: 0,
1088        };
1089        self.lower_move(*input, output, source.storage(), geometry)
1090    }
1091
1092    fn lower_reverse(
1093        &mut self,
1094        operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1095        inputs: &[ValueId],
1096        output: ValueId,
1097    ) -> Result<(), LoweringError> {
1098        let [input] = inputs else {
1099            return Err(LoweringError::UnsupportedGraph);
1100        };
1101        let OpAttributes::Reverse { axis } = operator.source().attributes() else {
1102            return Err(LoweringError::UnsupportedGraph);
1103        };
1104        let source = self.shape(*input)?;
1105        let target = self.shape(output)?;
1106        if source != target {
1107            return Err(LoweringError::UnsupportedGraph);
1108        }
1109        let axis = usize::try_from(axis)
1110            .ok()
1111            .filter(|axis| *axis < source.rank())
1112            .ok_or(LoweringError::UnsupportedGraph)?;
1113        let strides = source.strides();
1114        let mut in_strides = strides.clone();
1115        in_strides[axis] = strides[axis].wrapping_neg();
1116        let in_offset = (source.dims[axis] - 1).wrapping_mul(strides[axis]);
1117        let geometry = MoveGeometry {
1118            count: source.elements,
1119            dims: source.padded_dims(),
1120            in_strides: pad_leading(&in_strides, 0),
1121            in_offset,
1122            out_strides: pad_leading(&strides, 0),
1123            out_offset: 0,
1124        };
1125        self.lower_move(*input, output, source.storage(), geometry)
1126    }
1127
1128    fn lower_concat(
1129        &mut self,
1130        operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1131        inputs: &[ValueId],
1132        output: ValueId,
1133    ) -> Result<(), LoweringError> {
1134        if inputs.is_empty() {
1135            return Err(LoweringError::UnsupportedGraph);
1136        }
1137        let OpAttributes::Concat { axis } = operator.source().attributes() else {
1138            return Err(LoweringError::UnsupportedGraph);
1139        };
1140        let target = self.shape(output)?;
1141        let axis = usize::try_from(axis)
1142            .ok()
1143            .filter(|axis| *axis < target.rank())
1144            .ok_or(LoweringError::UnsupportedGraph)?;
1145        let mut total = 0_u64;
1146        let mut sources = Vec::with_capacity(inputs.len());
1147        for input in inputs {
1148            let source = self.shape(*input)?;
1149            if source.dtype != target.dtype
1150                || source.rank() != target.rank()
1151                || source
1152                    .dims
1153                    .iter()
1154                    .zip(&target.dims)
1155                    .enumerate()
1156                    .any(|(d, (source, target))| d != axis && source != target)
1157            {
1158                return Err(LoweringError::UnsupportedGraph);
1159            }
1160            total += u64::from(source.dims[axis]);
1161            sources.push(source);
1162        }
1163        if total != u64::from(target.dims[axis]) {
1164            return Err(LoweringError::UnsupportedGraph);
1165        }
1166        let out_strides = target.strides();
1167        let to = self.output_location(output)?;
1168        let mut offset = 0_u32;
1169        for (input, source) in inputs.iter().zip(&sources) {
1170            let geometry = MoveGeometry {
1171                count: source.elements,
1172                dims: source.padded_dims(),
1173                in_strides: pad_leading(&source.strides(), 0),
1174                in_offset: 0,
1175                out_strides: pad_leading(&out_strides, 0),
1176                out_offset: offset.wrapping_mul(out_strides[axis]),
1177            };
1178            offset += source.dims[axis];
1179            let from = self.input_location(*input)?;
1180            let spec = move_spec(self.operand(from), self.operand(to), geometry, false);
1181            self.dispatch(
1182                KernelSpec::Move {
1183                    storage: source.storage(),
1184                    contiguous: false,
1185                },
1186                spec,
1187                Work::Linear(source.elements),
1188                &[from],
1189                to,
1190                output,
1191            );
1192        }
1193        Ok(())
1194    }
1195
1196    fn lower_move(
1197        &mut self,
1198        input: ValueId,
1199        output: ValueId,
1200        storage: Storage,
1201        geometry: MoveGeometry,
1202    ) -> Result<(), LoweringError> {
1203        let from = self.input_location(input)?;
1204        let to = self.output_location(output)?;
1205        let spec = move_spec(self.operand(from), self.operand(to), geometry, false);
1206        self.dispatch(
1207            KernelSpec::Move {
1208                storage,
1209                contiguous: false,
1210            },
1211            spec,
1212            Work::Linear(geometry.count),
1213            &[from],
1214            to,
1215            output,
1216        );
1217        Ok(())
1218    }
1219
1220    /// `CAST` between float dtypes: one conversion dispatch, elementwise and shape-preserving.
1221    /// This is the tier's only narrowing path, and the reason a chain of FP8 matmuls needs no
1222    /// host round trip between layers.
1223    fn lower_cast(&mut self, inputs: &[ValueId], output: ValueId) -> Result<(), LoweringError> {
1224        let [input] = inputs else {
1225            return Err(LoweringError::UnsupportedGraph);
1226        };
1227        let source = self.shape(*input)?;
1228        let destination = self.shape(output)?;
1229        if source.dims != destination.dims {
1230            return Err(LoweringError::UnsupportedGraph);
1231        }
1232        for dtype in [source.dtype, destination.dtype] {
1233            if !matches!(
1234                dtype,
1235                DType::FP32 | DType::FP16 | DType::FP8E4M3 | DType::FP8E5M2
1236            ) {
1237                return Err(LoweringError::UnsupportedType(dtype));
1238            }
1239        }
1240        let from = self.input_location(*input)?;
1241        let to = self.output_location(output)?;
1242        let strides = pad_leading(&source.strides(), 0);
1243        let geometry = MoveGeometry {
1244            count: source.elements,
1245            dims: source.padded_dims(),
1246            in_strides: strides,
1247            in_offset: 0,
1248            out_strides: strides,
1249            out_offset: 0,
1250        };
1251        let spec = move_spec(self.operand(from), self.operand(to), geometry, true);
1252        // One invocation per destination word (`shader::assemble_contiguous_lanes`).
1253        self.dispatch(
1254            KernelSpec::Cast {
1255                input: source.storage(),
1256                output: destination.storage(),
1257            },
1258            spec,
1259            Work::Linear(source.elements.div_ceil(destination.storage().lanes())),
1260            &[from],
1261            to,
1262            output,
1263        );
1264        Ok(())
1265    }
1266
1267    fn lower_matmul(&mut self, inputs: &[ValueId], output: ValueId) -> Result<(), LoweringError> {
1268        let [lhs, rhs, lhs_zp, rhs_zp] = inputs else {
1269            return Err(LoweringError::UnsupportedGraph);
1270        };
1271        let lhs_shape = self.convertible_float_shape(*lhs)?;
1272        let rhs_shape = self.typed_shape(*rhs, lhs_shape.dtype)?;
1273        // TOSA defines FP8 MATMUL as `(FP8, FP8) -> FP16`; every other admitted operand type
1274        // keeps its own dtype.
1275        let result_dtype = match lhs_shape.dtype {
1276            DType::FP8E4M3 | DType::FP8E5M2 => DType::FP16,
1277            dtype => dtype,
1278        };
1279        let out_shape = self.typed_shape(output, result_dtype)?;
1280        // TOSA 1.0 floating-point MATMUL admits only zero zero-points: the two trailing inputs
1281        // must be `CONST` tensors whose serialized payload is all-zero (signed zero included).
1282        for zero_point in [lhs_zp, rhs_zp] {
1283            self.require_zero_constant(*zero_point, lhs_shape.dtype)?;
1284        }
1285        let ([batch_l, m, k], [batch_r, k_r, n], [batch_o, m_o, n_o]) = (
1286            lhs_shape.dims.as_slice(),
1287            rhs_shape.dims.as_slice(),
1288            out_shape.dims.as_slice(),
1289        ) else {
1290            return Err(LoweringError::UnsupportedGraph);
1291        };
1292        if batch_l != batch_r || batch_l != batch_o || k != k_r || m != m_o || n != n_o {
1293            return Err(LoweringError::UnsupportedGraph);
1294        }
1295        let (batch, m, n, k) = (*batch_l, *m, *n, *k);
1296        let from_lhs = self.input_location(*lhs)?;
1297        let from_rhs = self.input_location(*rhs)?;
1298        let to = self.output_location(output)?;
1299        let (input, output_storage) = (lhs_shape.storage(), out_shape.storage());
1300        if m <= crate::shader::STREAM_ROWS {
1301            // A few rows against a wide matrix is bandwidth-bound and gets the streaming kernel,
1302            // which reads the lhs as binary32 words: a narrower lhs is widened first into an
1303            // arena intermediate that lives only for this operator (`batch · m · k` words, small
1304            // next to the weights it saves the kernel from widening in every lane).
1305            let lhs_words = if input == Storage::Word {
1306                from_lhs
1307            } else {
1308                let count = batch
1309                    .checked_mul(m)
1310                    .and_then(|rows| rows.checked_mul(k))
1311                    .ok_or(LoweringError::ResourceLimit)?;
1312                let region = self.allocate_region(u64::from(count) * 4, self.position)?;
1313                let widened = Location::Region(region);
1314                let geometry = MoveGeometry {
1315                    count,
1316                    dims: [1; MAX_RANK],
1317                    in_strides: [0; MAX_RANK],
1318                    in_offset: 0,
1319                    out_strides: [0; MAX_RANK],
1320                    out_offset: 0,
1321                };
1322                let spec = move_spec(
1323                    self.operand(from_lhs),
1324                    self.operand(widened),
1325                    geometry,
1326                    true,
1327                );
1328                self.dispatch(
1329                    KernelSpec::Cast {
1330                        input,
1331                        output: Storage::Word,
1332                    },
1333                    spec,
1334                    Work::Linear(count),
1335                    &[from_lhs],
1336                    widened,
1337                    output,
1338                );
1339                widened
1340            };
1341            let spec = matmul_spec(
1342                self.operand(lhs_words),
1343                self.operand(from_rhs),
1344                self.operand(to),
1345                m,
1346                n,
1347                k,
1348                batch,
1349            );
1350            self.dispatch(
1351                KernelSpec::MatmulStream {
1352                    rhs: input,
1353                    output: output_storage,
1354                },
1355                spec,
1356                Work::MatmulStream { n, batch },
1357                &[lhs_words, from_rhs],
1358                to,
1359                output,
1360            );
1361        } else {
1362            let spec = matmul_spec(
1363                self.operand(from_lhs),
1364                self.operand(from_rhs),
1365                self.operand(to),
1366                m,
1367                n,
1368                k,
1369                batch,
1370            );
1371            self.dispatch(
1372                KernelSpec::Matmul {
1373                    input,
1374                    output: output_storage,
1375                },
1376                spec,
1377                Work::Matmul { m, n, batch },
1378                &[from_lhs, from_rhs],
1379                to,
1380                output,
1381            );
1382        }
1383        Ok(())
1384    }
1385
1386    fn lower_max_pool(
1387        &mut self,
1388        operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1389        inputs: &[ValueId],
1390        output: ValueId,
1391    ) -> Result<(), LoweringError> {
1392        let [input] = inputs else {
1393            return Err(LoweringError::UnsupportedGraph);
1394        };
1395        let OpAttributes::MaxPool2d {
1396            kernel,
1397            stride,
1398            pad,
1399            nan_mode,
1400        } = operator.source().attributes()
1401        else {
1402            return Err(LoweringError::UnsupportedGraph);
1403        };
1404        let nan_mode = nan_mode_of(nan_mode)?;
1405        // TOSA admits FP8 MAX_POOL2D: pooling selects an existing encoding rather than
1406        // computing a new one, so the widen/narrow round trip through the kernel is exact.
1407        let source = self.convertible_float_shape(*input)?;
1408        let target = self.typed_shape(output, source.dtype)?;
1409        let ([batch, height, width, channels], [batch_o, out_height, out_width, channels_o]) =
1410            (source.dims.as_slice(), target.dims.as_slice())
1411        else {
1412            return Err(LoweringError::UnsupportedGraph);
1413        };
1414        let attribute = |list: virtio_accel_tosa::I32List<'_>, len: usize| {
1415            let values: Vec<u32> = list
1416                .iter()
1417                .map(|value| u32::try_from(value).map_err(|_| LoweringError::UnsupportedGraph))
1418                .collect::<Result<_, _>>()?;
1419            if values.len() != len {
1420                return Err(LoweringError::UnsupportedGraph);
1421            }
1422            Ok(values)
1423        };
1424        let kernel = attribute(kernel, 2)?;
1425        let stride = attribute(stride, 2)?;
1426        let pad = attribute(pad, 4)?;
1427        if kernel.contains(&0) || stride.contains(&0) {
1428            return Err(LoweringError::UnsupportedGraph);
1429        }
1430        // TOSA: every pad is smaller than its kernel extent, so no window is entirely padding,
1431        // and the padded extent divides exactly into output positions.
1432        let [pad_top, pad_bottom, pad_left, pad_right] = pad[..] else {
1433            return Err(LoweringError::UnsupportedGraph);
1434        };
1435        if pad_top >= kernel[0]
1436            || pad_bottom >= kernel[0]
1437            || pad_left >= kernel[1]
1438            || pad_right >= kernel[1]
1439        {
1440            return Err(LoweringError::UnsupportedGraph);
1441        }
1442        let output_extent = |input: u32, pad_a: u32, pad_b: u32, kernel: u32, stride: u32| {
1443            let padded = u64::from(input) + u64::from(pad_a) + u64::from(pad_b);
1444            let span = padded.checked_sub(u64::from(kernel))?;
1445            if span % u64::from(stride) != 0 {
1446                return None;
1447            }
1448            u32::try_from(span / u64::from(stride) + 1).ok()
1449        };
1450        let expected_height = output_extent(*height, pad_top, pad_bottom, kernel[0], stride[0])
1451            .ok_or(LoweringError::UnsupportedGraph)?;
1452        let expected_width = output_extent(*width, pad_left, pad_right, kernel[1], stride[1])
1453            .ok_or(LoweringError::UnsupportedGraph)?;
1454        if batch != batch_o
1455            || channels != channels_o
1456            || expected_height != *out_height
1457            || expected_width != *out_width
1458        {
1459            return Err(LoweringError::UnsupportedGraph);
1460        }
1461        // The kernel's window arithmetic stays inside u32.
1462        let last_row = u64::from(*out_height - 1) * u64::from(stride[0]) + u64::from(kernel[0]);
1463        let last_col = u64::from(*out_width - 1) * u64::from(stride[1]) + u64::from(kernel[1]);
1464        if last_row > u64::from(u32::MAX) || last_col > u64::from(u32::MAX) {
1465            return Err(LoweringError::ResourceLimit);
1466        }
1467        let geometry = PoolGeometry {
1468            batch: *batch,
1469            height: *height,
1470            width: *width,
1471            channels: *channels,
1472            out_height: *out_height,
1473            out_width: *out_width,
1474            kernel: [kernel[0], kernel[1]],
1475            stride: [stride[0], stride[1]],
1476            pad_top,
1477            pad_left,
1478        };
1479        let from = self.input_location(*input)?;
1480        let to = self.output_location(output)?;
1481        let spec = max_pool_spec(self.operand(from), self.operand(to), geometry);
1482        self.dispatch(
1483            KernelSpec::MaxPool {
1484                nan_mode,
1485                float: source.storage(),
1486            },
1487            spec,
1488            Work::Linear(target.elements),
1489            &[from],
1490            to,
1491            output,
1492        );
1493        Ok(())
1494    }
1495
1496    fn lower_reduce(
1497        &mut self,
1498        operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1499        inputs: &[ValueId],
1500        output: ValueId,
1501    ) -> Result<(), LoweringError> {
1502        let [input] = inputs else {
1503            return Err(LoweringError::UnsupportedGraph);
1504        };
1505        let (axis, op, argmax) = match operator.source().attributes() {
1506            OpAttributes::ArgMax { axis, nan_mode } => {
1507                (axis, ReduceOp::ArgMax(nan_mode_of(nan_mode)?), true)
1508            }
1509            OpAttributes::ReduceMax { axis, nan_mode } => {
1510                (axis, ReduceOp::Max(nan_mode_of(nan_mode)?), false)
1511            }
1512            OpAttributes::ReduceMin { axis, nan_mode } => {
1513                (axis, ReduceOp::Min(nan_mode_of(nan_mode)?), false)
1514            }
1515            OpAttributes::ReduceProduct { axis } => (axis, ReduceOp::Product, false),
1516            OpAttributes::ReduceSum { axis } => (axis, ReduceOp::Sum, false),
1517            _ => return Err(LoweringError::UnsupportedGraph),
1518        };
1519        // TOSA admits FP8 for ARGMAX (comparing widened values, emitting an INT32 index) but
1520        // for none of the REDUCE lanes, which stay on the float-arithmetic dtypes.
1521        let source = if argmax {
1522            self.convertible_float_shape(*input)?
1523        } else {
1524            self.float_shape(*input)?
1525        };
1526        let target = self.typed_shape(output, if argmax { DType::INT32 } else { source.dtype })?;
1527        let axis = usize::try_from(axis)
1528            .ok()
1529            .filter(|axis| *axis < source.rank())
1530            .ok_or(LoweringError::UnsupportedGraph)?;
1531        let mut expected = source.dims.clone();
1532        if argmax {
1533            expected.remove(axis);
1534        } else {
1535            expected[axis] = 1;
1536        }
1537        if target.dims != expected {
1538            return Err(LoweringError::UnsupportedGraph);
1539        }
1540        let outer = source.dims[..axis].iter().product::<u32>();
1541        let inner = source.dims[axis + 1..].iter().product::<u32>();
1542        let from = self.input_location(*input)?;
1543        let to = self.output_location(output)?;
1544        let spec = reduce_spec(
1545            self.operand(from),
1546            self.operand(to),
1547            outer,
1548            source.dims[axis],
1549            inner,
1550        );
1551        self.dispatch(
1552            KernelSpec::Reduce {
1553                op,
1554                float: source.storage(),
1555            },
1556            spec,
1557            Work::Linear(target.elements),
1558            &[from],
1559            to,
1560            output,
1561        );
1562        Ok(())
1563    }
1564
1565    fn lower_elementwise(
1566        &mut self,
1567        operator: &virtio_accel_tosa::AnalyzedOperator<'_>,
1568        inputs: &[ValueId],
1569        output: ValueId,
1570    ) -> Result<(), LoweringError> {
1571        let op = operator.op();
1572        let attributes = operator.source().attributes();
1573        let mut clamp = None;
1574        let (lane, tensor_inputs): (ElementwiseOp, Vec<ValueId>) = match op {
1575            Op::ABS => (ElementwiseOp::Abs, inputs.to_vec()),
1576            Op::CEIL => (ElementwiseOp::Ceil, inputs.to_vec()),
1577            Op::COS => (ElementwiseOp::Cos, inputs.to_vec()),
1578            Op::ERF => (ElementwiseOp::Erf, inputs.to_vec()),
1579            Op::EXP => (ElementwiseOp::Exp, inputs.to_vec()),
1580            Op::FLOOR => (ElementwiseOp::Floor, inputs.to_vec()),
1581            Op::LOG => (ElementwiseOp::Log, inputs.to_vec()),
1582            Op::RECIPROCAL => (ElementwiseOp::Reciprocal, inputs.to_vec()),
1583            Op::RSQRT => (ElementwiseOp::Rsqrt, inputs.to_vec()),
1584            Op::SIN => (ElementwiseOp::Sin, inputs.to_vec()),
1585            Op::SIGMOID => (ElementwiseOp::Sigmoid, inputs.to_vec()),
1586            Op::TANH => (ElementwiseOp::Tanh, inputs.to_vec()),
1587            Op::NEGATE => {
1588                let [value, input_zp, output_zp] = inputs else {
1589                    return Err(LoweringError::UnsupportedGraph);
1590                };
1591                let dtype = self.shape(*value)?.dtype;
1592                self.require_zero_constant(*input_zp, dtype)?;
1593                self.require_zero_constant(*output_zp, dtype)?;
1594                (ElementwiseOp::Negate, vec![*value])
1595            }
1596            Op::CLAMP => {
1597                let OpAttributes::Clamp {
1598                    min_val,
1599                    max_val,
1600                    nan_mode,
1601                } = attributes
1602                else {
1603                    return Err(LoweringError::UnsupportedGraph);
1604                };
1605                // Bounds serialize in the tensor dtype: four bytes for FP32, two for FP16. The
1606                // kernel applies the clamp at binary32 and narrows once (ADR 0008), so the
1607                // specialization words always carry binary32 bit patterns; a binary16 bound is
1608                // widened host-side, exactly. The ordering check runs on the same values.
1609                let bound = |bytes: &[u8]| -> Result<u32, LoweringError> {
1610                    let value = match bytes.len() {
1611                        4 => f32::from_le_bytes(
1612                            bytes
1613                                .try_into()
1614                                .map_err(|_| LoweringError::UnsupportedGraph)?,
1615                        ),
1616                        2 => crate::shader::f16_to_f32(u16::from_le_bytes(
1617                            bytes
1618                                .try_into()
1619                                .map_err(|_| LoweringError::UnsupportedGraph)?,
1620                        )),
1621                        _ => return Err(LoweringError::UnsupportedGraph),
1622                    };
1623                    if value.is_nan() {
1624                        return Err(LoweringError::UnsupportedGraph);
1625                    }
1626                    Ok(value.to_bits())
1627                };
1628                let lo = bound(min_val)?;
1629                let hi = bound(max_val)?;
1630                if f32::from_bits(hi) < f32::from_bits(lo) {
1631                    return Err(LoweringError::UnsupportedGraph);
1632                }
1633                clamp = Some([lo, hi]);
1634                (
1635                    ElementwiseOp::Clamp(nan_mode_of(nan_mode)?),
1636                    inputs.to_vec(),
1637                )
1638            }
1639            Op::ADD => (ElementwiseOp::Add, inputs.to_vec()),
1640            Op::SUB => (ElementwiseOp::Sub, inputs.to_vec()),
1641            Op::POW => (ElementwiseOp::Pow, inputs.to_vec()),
1642            Op::MUL => {
1643                let [lhs, rhs, shift] = inputs else {
1644                    return Err(LoweringError::UnsupportedGraph);
1645                };
1646                self.require_zero_constant(*shift, DType::INT8)?;
1647                (ElementwiseOp::Mul, vec![*lhs, *rhs])
1648            }
1649            Op::MAXIMUM => {
1650                let OpAttributes::Maximum { nan_mode } = attributes else {
1651                    return Err(LoweringError::UnsupportedGraph);
1652                };
1653                (
1654                    ElementwiseOp::Maximum(nan_mode_of(nan_mode)?),
1655                    inputs.to_vec(),
1656                )
1657            }
1658            Op::MINIMUM => {
1659                let OpAttributes::Minimum { nan_mode } = attributes else {
1660                    return Err(LoweringError::UnsupportedGraph);
1661                };
1662                (
1663                    ElementwiseOp::Minimum(nan_mode_of(nan_mode)?),
1664                    inputs.to_vec(),
1665                )
1666            }
1667            Op::EQUAL => (ElementwiseOp::Equal, inputs.to_vec()),
1668            Op::GREATER => (ElementwiseOp::Greater, inputs.to_vec()),
1669            Op::GREATER_EQUAL => (ElementwiseOp::GreaterEqual, inputs.to_vec()),
1670            Op::LOGICAL_AND => (ElementwiseOp::LogicalAnd, inputs.to_vec()),
1671            Op::LOGICAL_OR => (ElementwiseOp::LogicalOr, inputs.to_vec()),
1672            Op::LOGICAL_XOR => (ElementwiseOp::LogicalXor, inputs.to_vec()),
1673            Op::LOGICAL_NOT => (ElementwiseOp::LogicalNot, inputs.to_vec()),
1674            Op::SELECT => (ElementwiseOp::Select, inputs.to_vec()),
1675            other => return Err(LoweringError::UnsupportedOperator(other)),
1676        };
1677        self.emit_elementwise(lane, &tensor_inputs, output, clamp)
1678    }
1679
1680    fn emit_elementwise(
1681        &mut self,
1682        lane: ElementwiseOp,
1683        inputs: &[ValueId],
1684        output: ValueId,
1685        clamp: Option<[u32; 2]>,
1686    ) -> Result<(), LoweringError> {
1687        let lanes = lane.inputs();
1688        if inputs.len() != lanes.len() {
1689            return Err(LoweringError::UnsupportedGraph);
1690        }
1691        let target = self.shape(output)?;
1692        // Every float lane (`Word` in the lane table) must agree on one float dtype, FP32 or
1693        // FP16; pure-`BOOL` lanes have no float storage and the kernel's float width is unused.
1694        let mut float_dtype: Option<DType> = None;
1695        let mut unify = |dtype: DType| -> Result<(), LoweringError> {
1696            if !matches!(dtype, DType::FP32 | DType::FP16) {
1697                return Err(LoweringError::UnsupportedType(dtype));
1698            }
1699            match float_dtype {
1700                Some(existing) if existing != dtype => Err(LoweringError::UnsupportedGraph),
1701                _ => {
1702                    float_dtype = Some(dtype);
1703                    Ok(())
1704                }
1705            }
1706        };
1707        let mut shapes = Vec::with_capacity(inputs.len());
1708        for (input, storage) in inputs.iter().zip(lanes) {
1709            let shape = self.shape(*input)?;
1710            match storage {
1711                Storage::Word => unify(shape.dtype)?,
1712                Storage::Byte if shape.dtype != DType::BOOL => {
1713                    return Err(LoweringError::UnsupportedType(shape.dtype));
1714                }
1715                _ => {}
1716            }
1717            if shape.rank() != target.rank()
1718                || shape
1719                    .dims
1720                    .iter()
1721                    .zip(&target.dims)
1722                    .any(|(dim, out)| *dim != *out && *dim != 1)
1723            {
1724                return Err(LoweringError::UnsupportedGraph);
1725            }
1726            shapes.push(shape);
1727        }
1728        match lane.output() {
1729            Storage::Word => unify(target.dtype)?,
1730            Storage::Byte if target.dtype != DType::BOOL => {
1731                return Err(LoweringError::UnsupportedType(target.dtype));
1732            }
1733            _ => {}
1734        }
1735        let float = storage_of(float_dtype.unwrap_or(DType::FP32));
1736        let broadcast = shapes.iter().any(|shape| shape.dims != target.dims);
1737        let mut strides = Vec::with_capacity(inputs.len());
1738        for shape in &shapes {
1739            let own = shape.strides();
1740            let mut broadcast_strides = vec![0_u32; shape.rank()];
1741            for d in 0..shape.rank() {
1742                broadcast_strides[d] = if shape.dims[d] == 1 && target.dims[d] != 1 {
1743                    0
1744                } else {
1745                    own[d]
1746                };
1747            }
1748            strides.push(pad_leading(&broadcast_strides, 0));
1749        }
1750        let mut reads = Vec::with_capacity(inputs.len());
1751        let mut operands = Vec::with_capacity(inputs.len());
1752        for input in inputs {
1753            let location = self.input_location(*input)?;
1754            reads.push(location);
1755            operands.push(self.operand(location));
1756        }
1757        let to = self.output_location(output)?;
1758        let spec = ElementwiseSpec {
1759            count: target.elements,
1760            inputs: &operands,
1761            output: self.operand(to),
1762            dims: target.padded_dims(),
1763            strides: &strides,
1764            clamp,
1765        }
1766        .words(broadcast);
1767        self.dispatch(
1768            KernelSpec::Elementwise {
1769                op: lane,
1770                float,
1771                broadcast,
1772            },
1773            spec,
1774            Work::Linear(target.elements),
1775            &reads,
1776            to,
1777            output,
1778        );
1779        Ok(())
1780    }
1781
1782    /// A serialized constant of `dtype` whose every scalar is zero (signed zero included).
1783    fn require_zero_constant(&mut self, value: ValueId, dtype: DType) -> Result<(), LoweringError> {
1784        let AnalyzedValueKind::Tensor(tensor) = self.analysis.value(value).kind() else {
1785            return Err(LoweringError::UnsupportedGraph);
1786        };
1787        if tensor.dtype() != dtype {
1788            return Err(LoweringError::UnsupportedGraph);
1789        }
1790        let bytes = self
1791            .analysis
1792            .serialized_constant(value)
1793            .ok_or(LoweringError::UnsupportedGraph)?;
1794        let zero = match dtype {
1795            DType::FP32 => {
1796                bytes.len() % 4 == 0
1797                    && bytes.chunks_exact(4).all(|chunk| {
1798                        u32::from_le_bytes(chunk.try_into().expect("four bytes")) & 0x7fff_ffff == 0
1799                    })
1800            }
1801            DType::FP16 => {
1802                bytes.len() % 2 == 0
1803                    && bytes.chunks_exact(2).all(|chunk| {
1804                        u16::from_le_bytes(chunk.try_into().expect("two bytes")) & 0x7fff == 0
1805                    })
1806            }
1807            // Both FP8 encodings put the sign in bit 7, so this admits signed zero and nothing
1808            // else, exactly as the FP32 and FP16 arms do.
1809            DType::FP8E4M3 | DType::FP8E5M2 => bytes.iter().all(|byte| byte & 0x7f == 0),
1810            _ => bytes.iter().all(|byte| *byte == 0),
1811        };
1812        if bytes.is_empty() || !zero {
1813            return Err(LoweringError::UnsupportedGraph);
1814        }
1815        Ok(())
1816    }
1817}
1818
1819fn nan_mode_of(mode: NanPropagationMode) -> Result<NanMode, LoweringError> {
1820    if mode == NanPropagationMode::PROPAGATE {
1821        Ok(NanMode::Propagate)
1822    } else if mode == NanPropagationMode::IGNORE {
1823        Ok(NanMode::Ignore)
1824    } else {
1825        Err(LoweringError::UnsupportedGraph)
1826    }
1827}
1828
1829#[cfg(test)]
1830mod tests {
1831    use super::*;
1832    use virtio_accel_conformance::numerics::{
1833        HEXAGON_UNARY_FP16_CASES, IDENTITY_EDGES_FP16, IDENTITY_EDGES_FP32, IDENTITY_INT8,
1834        MATMUL_FP16, MATMUL_FP32, MAX_POOL2D_FP16, MAX_POOL2D_FP32,
1835    };
1836
1837    const IDENTITY_FP32_LOCAL: &[u8] = include_bytes!("../tests/data/identity-fp32-v1.0.0.tosa");
1838
1839    #[test]
1840    fn targets_validate_and_round_trip() {
1841        for target in [VULKAN_TOSA_TARGET, VULKAN_TOSA_INTEGER_TARGET] {
1842            assert_eq!(target.validate(), Ok(target));
1843            assert_eq!(Target::from_identity(target.to_identity()), Ok(target));
1844        }
1845        assert_ne!(VULKAN_TOSA_TARGET, VULKAN_TOSA_INTEGER_TARGET);
1846    }
1847
1848    #[test]
1849    fn capability_names_the_shared_fp32_operator_set() {
1850        assert_eq!(FLOAT_OPERATORS.len(), 42);
1851        for op in [
1852            Op::IDENTITY,
1853            Op::MATMUL,
1854            Op::MAX_POOL2D,
1855            Op::ARGMAX,
1856            Op::ERF,
1857            Op::CONCAT,
1858            Op::TRANSPOSE,
1859        ] {
1860            assert!(supports_tosa_operator(op), "{op:?}");
1861        }
1862        for op in [Op::CONV2D, Op::AVG_POOL2D, Op::CAST, Op::RESCALE, Op::PAD] {
1863            assert!(!supports_tosa_operator(op), "{op:?}");
1864        }
1865        assert!(supports_tosa_dtype(DType::FP32));
1866        assert!(supports_tosa_dtype(DType::FP16));
1867        assert!(supports_tosa_dtype(DType::BOOL));
1868        assert!(supports_tosa_dtype(DType::INT32));
1869        // INT8 is a compile-time parameter only (the `MUL` shift), never a boundary dtype.
1870        assert!(!supports_tosa_dtype(DType::INT8));
1871        assert!(VULKAN_TOSA_CAPABILITY.supports_dtype(DType::INT8, ValueRoles::CONSTANT));
1872        assert!(!VULKAN_TOSA_CAPABILITY.supports_dtype(DType::INT8, ValueRoles::INTERMEDIATE));
1873        assert_eq!(VULKAN_TOSA_CAPABILITY.target, VULKAN_TOSA_TARGET);
1874    }
1875
1876    #[test]
1877    fn lowers_the_local_fp32_identity_artifact() {
1878        let plan = lower_tosa(IDENTITY_FP32_LOCAL, VULKAN_TOSA_TARGET).unwrap();
1879        assert_eq!(plan.slots.len(), 2);
1880        assert_eq!(plan.slot(0).unwrap().role, SlotRole::Input);
1881        assert_eq!(plan.slot(1).unwrap().role, SlotRole::Output);
1882        assert_eq!(plan.slot(0).unwrap().byte_len, 4);
1883        assert_eq!(plan.slot(1).unwrap().storage, Storage::Word);
1884        assert!(plan.slot(2).is_none());
1885        assert_eq!(plan.arena_bytes, 0);
1886        assert!(plan.constants.is_empty());
1887        assert_eq!(plan.dispatches.len(), 1);
1888        let dispatch = &plan.dispatches[0];
1889        assert_eq!(
1890            dispatch.kernel,
1891            KernelSpec::Move {
1892                storage: Storage::Word,
1893                contiguous: true
1894            }
1895        );
1896        assert_eq!(dispatch.work, Work::Linear(1));
1897        assert!(!dispatch.barrier_before);
1898        // input (buffer 0, base 0), output (buffer 1, base 0), count 1
1899        assert_eq!(dispatch.spec, vec![0, 0, 1, 0, 1]);
1900    }
1901
1902    #[test]
1903    fn lowers_the_shared_fp32_edge_identity_artifact() {
1904        let plan = lower_tosa(IDENTITY_EDGES_FP32.artifact, VULKAN_TOSA_TARGET).unwrap();
1905        let expected = IDENTITY_EDGES_FP32.inputs[0].values.len();
1906        assert_eq!(plan.dispatches[0].work, Work::Linear(expected as u32));
1907        assert_eq!(plan.slot(1).unwrap().byte_len as usize, expected * 4);
1908    }
1909
1910    #[test]
1911    fn rejects_other_targets_before_parsing() {
1912        assert_eq!(
1913            lower_tosa(IDENTITY_FP32_LOCAL, VULKAN_TOSA_INTEGER_TARGET),
1914            Err(LoweringError::UnsupportedTarget)
1915        );
1916        assert_eq!(
1917            lower_tosa(IDENTITY_INT8.artifact, VULKAN_TOSA_INTEGER_TARGET),
1918            Err(LoweringError::UnsupportedTarget)
1919        );
1920    }
1921
1922    #[test]
1923    fn rejects_mistyped_identity_graphs_loudly() {
1924        // INT8 identity under the floating-point target: never relabeled.
1925        assert!(matches!(
1926            lower_tosa(IDENTITY_INT8.artifact, VULKAN_TOSA_TARGET),
1927            Err(LoweringError::UnsupportedType(DType::INT8) | LoweringError::Analysis(_))
1928        ));
1929    }
1930
1931    #[test]
1932    fn fp16_capability_extends_the_fp32_boundary() {
1933        for dtype in [DType::FP32, DType::FP16, DType::BOOL, DType::INT32] {
1934            assert!(
1935                VULKAN_TOSA_FP16_CAPABILITY.supports_dtype(dtype, ValueRoles::INPUT),
1936                "{dtype:?}"
1937            );
1938            assert!(
1939                VULKAN_TOSA_FP16_CAPABILITY.supports_dtype(dtype, ValueRoles::OUTPUT),
1940                "{dtype:?}"
1941            );
1942        }
1943        assert_eq!(VULKAN_TOSA_FP16_CAPABILITY.target, VULKAN_TOSA_TARGET);
1944        assert_eq!(VULKAN_TOSA_FP16_CAPABILITY.operators, FLOAT_OPERATORS);
1945        assert!(!VULKAN_TOSA_CAPABILITY.supports_dtype(DType::FP16, ValueRoles::INPUT));
1946    }
1947
1948    #[test]
1949    fn lowers_the_shared_fp16_artifacts() {
1950        let plan = lower_tosa(IDENTITY_EDGES_FP16.artifact, VULKAN_TOSA_TARGET).unwrap();
1951        let expected = IDENTITY_EDGES_FP16.inputs[0].bits.len() as u32;
1952        assert_eq!(plan.slot(0).unwrap().byte_len, u64::from(expected) * 2);
1953        assert_eq!(plan.slot(0).unwrap().storage, Storage::Half);
1954        assert_eq!(plan.dispatches.len(), 1);
1955        assert_eq!(
1956            plan.dispatches[0].kernel,
1957            KernelSpec::Move {
1958                storage: Storage::Half,
1959                contiguous: true
1960            }
1961        );
1962        // Two binary16 lanes per destination word, one invocation per word.
1963        assert_eq!(plan.dispatches[0].work, Work::Linear(expected.div_ceil(2)));
1964
1965        // Two rows: the streaming kernel, its FP16 lhs widened first into the arena.
1966        let plan = lower_tosa(MATMUL_FP16.artifact, VULKAN_TOSA_TARGET).unwrap();
1967        assert_eq!(plan.dispatches.len(), 2);
1968        assert_eq!(
1969            plan.dispatches[0].kernel,
1970            KernelSpec::Cast {
1971                input: Storage::Half,
1972                output: Storage::Word
1973            }
1974        );
1975        assert_eq!(plan.dispatches[0].work, Work::Linear(6));
1976        assert_eq!(
1977            plan.dispatches[1].kernel,
1978            KernelSpec::MatmulStream {
1979                rhs: Storage::Half,
1980                output: Storage::Half
1981            }
1982        );
1983        assert!(plan.dispatches[1].barrier_before, "reads the widened lhs");
1984        assert_eq!(plan.arena_bytes, ARENA_ALIGNMENT);
1985        assert_eq!(plan.slot(0).unwrap().byte_len, 6 * 2);
1986        assert_eq!(plan.slot(2).unwrap().byte_len, 4 * 2);
1987
1988        let plan = lower_tosa(MAX_POOL2D_FP16.artifact, VULKAN_TOSA_TARGET).unwrap();
1989        assert_eq!(
1990            plan.dispatches[0].kernel,
1991            KernelSpec::MaxPool {
1992                nan_mode: NanMode::Propagate,
1993                float: Storage::Half
1994            }
1995        );
1996    }
1997
1998    #[test]
1999    fn fp16_clamp_bounds_arrive_widened_to_binary32() {
2000        // The shared clamp-fp16 fixture clamps [0.5, 1.0, 2.0, 4.0] to [-1.0, 1.0]; the kernel
2001        // applies the clamp at binary32 and narrows once (ADR 0008), so the specialization
2002        // words carry the widened binary32 patterns.
2003        let clamp = HEXAGON_UNARY_FP16_CASES
2004            .iter()
2005            .find(|case| case.name == "clamp-fp16")
2006            .expect("the shared clamp-fp16 case");
2007        let plan = lower_tosa(clamp.artifact, VULKAN_TOSA_TARGET).unwrap();
2008        let dispatch = &plan.dispatches[0];
2009        assert_eq!(
2010            dispatch.kernel,
2011            KernelSpec::Elementwise {
2012                op: ElementwiseOp::Clamp(NanMode::Propagate),
2013                float: Storage::Half,
2014                broadcast: false
2015            }
2016        );
2017        assert_eq!(dispatch.spec[dispatch.spec.len() - 2], (-1.0_f32).to_bits());
2018        assert_eq!(dispatch.spec[dispatch.spec.len() - 1], 1.0_f32.to_bits());
2019    }
2020
2021    #[test]
2022    fn fp16_negate_zero_points_are_consumed_at_admission() {
2023        let negate = HEXAGON_UNARY_FP16_CASES
2024            .iter()
2025            .find(|case| case.name == "negate-fp16")
2026            .expect("the shared negate-fp16 case");
2027        let plan = lower_tosa(negate.artifact, VULKAN_TOSA_TARGET).unwrap();
2028        assert_eq!(
2029            plan.dispatches[0].kernel,
2030            KernelSpec::Elementwise {
2031                op: ElementwiseOp::Negate,
2032                float: Storage::Half,
2033                broadcast: false
2034            }
2035        );
2036        assert_eq!(plan.arena_bytes, 0);
2037    }
2038
2039    #[test]
2040    fn lowers_the_shared_fp32_matmul_artifact() {
2041        let plan = lower_tosa(MATMUL_FP32.artifact, VULKAN_TOSA_TARGET).unwrap();
2042        assert_eq!(plan.slots.len(), 3);
2043        assert_eq!(plan.slot(0).unwrap().byte_len, 6 * 4);
2044        assert_eq!(plan.slot(1).unwrap().byte_len, 6 * 4);
2045        assert_eq!(plan.slot(2).unwrap().role, SlotRole::Output);
2046        assert_eq!(plan.slot(2).unwrap().byte_len, 4 * 4);
2047        assert_eq!(plan.dispatches.len(), 1);
2048        let dispatch = &plan.dispatches[0];
2049        // Two rows: the streaming kernel, reading the FP32 lhs directly.
2050        assert_eq!(
2051            dispatch.kernel,
2052            KernelSpec::MatmulStream {
2053                rhs: Storage::Word,
2054                output: Storage::Word
2055            }
2056        );
2057        assert_eq!(dispatch.work, Work::MatmulStream { n: 2, batch: 1 });
2058        // lhs, rhs, out operands then m, n, k, batch.
2059        assert_eq!(dispatch.spec, vec![0, 0, 1, 0, 2, 0, 2, 2, 3, 1]);
2060        // The zero-point constants are consumed at admission, never uploaded.
2061        assert_eq!(plan.arena_bytes, 0);
2062    }
2063
2064    #[test]
2065    fn lowers_the_shared_fp32_max_pool_artifact() {
2066        let plan = lower_tosa(MAX_POOL2D_FP32.artifact, VULKAN_TOSA_TARGET).unwrap();
2067        let dispatch = &plan.dispatches[0];
2068        assert_eq!(
2069            dispatch.kernel,
2070            KernelSpec::MaxPool {
2071                nan_mode: NanMode::Propagate,
2072                float: Storage::Word
2073            }
2074        );
2075        assert_eq!(dispatch.work, Work::Linear(8));
2076        // in, out, N, H, W, C, OH, OW, KH, KW, SH, SW, PT, PL
2077        assert_eq!(
2078            dispatch.spec,
2079            vec![0, 0, 1, 0, 1, 4, 4, 2, 2, 2, 2, 2, 2, 2, 0, 0]
2080        );
2081    }
2082
2083    #[test]
2084    fn rejects_garbage_as_a_parse_error() {
2085        assert!(matches!(
2086            lower_tosa(b"not a flatbuffer", VULKAN_TOSA_TARGET),
2087            Err(LoweringError::Parse(_))
2088        ));
2089    }
2090
2091    #[test]
2092    fn arena_regions_pack_by_lifetime() {
2093        // Build a lowering with no graph behind it purely to exercise the allocator.
2094        let bytes = IDENTITY_FP32_LOCAL;
2095        let model = parse(bytes).unwrap();
2096        let analysis = model.analyze_for(VULKAN_TOSA_TARGET).unwrap();
2097        let mut lowering = Lowering::new(&analysis, VULKAN_TOSA_CAPABILITY).unwrap();
2098        lowering.position = 0;
2099        let a = lowering.allocate_region(100, 1).unwrap();
2100        let b = lowering.allocate_region(100, 5).unwrap();
2101        assert_eq!(lowering.regions[a].offset, 0);
2102        assert_eq!(lowering.regions[b].offset, ARENA_ALIGNMENT);
2103        // Position 2: `a` is dead, its space is reused before the arena grows.
2104        lowering.position = 2;
2105        let c = lowering.allocate_region(ARENA_ALIGNMENT, 9).unwrap();
2106        assert_eq!(lowering.regions[c].offset, 0);
2107        let d = lowering.allocate_region(1, 9).unwrap();
2108        assert_eq!(lowering.regions[d].offset, 2 * ARENA_ALIGNMENT);
2109        assert_eq!(lowering.arena_bytes, 3 * ARENA_ALIGNMENT);
2110    }
2111
2112    /// A constant first used at position 2 is uploaded before position 0's dispatch runs, so the
2113    /// bytes of `a` (dead since position 1) are not free for it: that dispatch would overwrite
2114    /// the constant. An intermediate first written at position 2 may still take them.
2115    #[test]
2116    fn constants_never_share_bytes_with_earlier_intermediates() {
2117        let model = parse(IDENTITY_FP32_LOCAL).unwrap();
2118        let analysis = model.analyze_for(VULKAN_TOSA_TARGET).unwrap();
2119        let mut lowering = Lowering::new(&analysis, VULKAN_TOSA_CAPABILITY).unwrap();
2120        lowering.position = 0;
2121        let a = lowering.allocate_region(64, 1).unwrap();
2122        assert_eq!(lowering.regions[a].offset, 0);
2123        lowering.position = 2;
2124        let constant = lowering
2125            .allocate_region_written_at(64, u32::MAX, 0)
2126            .unwrap();
2127        assert_eq!(lowering.regions[constant].offset, ARENA_ALIGNMENT);
2128        let intermediate = lowering.allocate_region(64, 3).unwrap();
2129        assert_eq!(lowering.regions[intermediate].offset, 0);
2130    }
2131
2132    /// Region `a` is read at position 1 and dies; at position 2 the packer hands its bytes to a
2133    /// new tensor `b`. The dispatch writing `b` reads nothing that was written since the last
2134    /// barrier, so only byte-range hazard tracking can order it after `a`'s reader.
2135    #[test]
2136    fn reused_arena_bytes_force_a_barrier_before_the_new_writer() {
2137        let model = parse(IDENTITY_FP32_LOCAL).unwrap();
2138        let analysis = model.analyze_for(VULKAN_TOSA_TARGET).unwrap();
2139        let mut lowering = Lowering::new(&analysis, VULKAN_TOSA_CAPABILITY).unwrap();
2140        let values: Vec<ValueId> = analysis.values().iter().map(|value| value.id()).collect();
2141        let (a_value, b_value) = (values[0], values[1]);
2142        let kernel = KernelSpec::Move {
2143            storage: Storage::Word,
2144            contiguous: true,
2145        };
2146
2147        // Position 0: a producer writes `a`.
2148        lowering.position = 0;
2149        let a = lowering.allocate_region(64, 1).unwrap();
2150        lowering.dispatch(
2151            kernel,
2152            Vec::new(),
2153            Work::Linear(16),
2154            &[Location::Slot(0)],
2155            Location::Region(a),
2156            a_value,
2157        );
2158        // Position 1: the last reader of `a` writes a bound slot.
2159        lowering.position = 1;
2160        lowering.dispatch(
2161            kernel,
2162            Vec::new(),
2163            Work::Linear(16),
2164            &[Location::Region(a)],
2165            Location::Slot(1),
2166            b_value,
2167        );
2168        // Position 2: `a` is dead, so `b` is packed into its bytes and written from a slot.
2169        lowering.position = 2;
2170        let b = lowering.allocate_region(64, 3).unwrap();
2171        assert_eq!(lowering.regions[b].offset, lowering.regions[a].offset);
2172        assert_ne!(a, b, "a fresh region index over the same bytes");
2173        lowering.dispatch(
2174            kernel,
2175            Vec::new(),
2176            Work::Linear(16),
2177            &[Location::Slot(0)],
2178            Location::Region(b),
2179            b_value,
2180        );
2181
2182        let barriers: Vec<bool> = lowering
2183            .dispatches
2184            .iter()
2185            .map(|dispatch| dispatch.barrier_before)
2186            .collect();
2187        assert_eq!(
2188            barriers,
2189            [false, true, true],
2190            "the reader depends on the producer (RAW); the new writer on the reader (WAR)"
2191        );
2192
2193        // Disjoint arena bytes carry no hazard: a writer into a fresh region needs no barrier.
2194        lowering.position = 3;
2195        let c = lowering.allocate_region(64, 4).unwrap();
2196        assert_ne!(lowering.regions[c].offset, lowering.regions[b].offset);
2197        lowering.dispatch(
2198            kernel,
2199            Vec::new(),
2200            Work::Linear(16),
2201            &[Location::Slot(0)],
2202            Location::Region(c),
2203            a_value,
2204        );
2205        assert!(!lowering.dispatches[3].barrier_before);
2206        assert!(
2207            MemKey::Arena {
2208                offset: 0,
2209                end: 256
2210            }
2211            .overlaps(MemKey::Arena {
2212                offset: 255,
2213                end: 512
2214            })
2215        );
2216        assert!(
2217            !MemKey::Arena {
2218                offset: 0,
2219                end: 256
2220            }
2221            .overlaps(MemKey::Arena {
2222                offset: 256,
2223                end: 512
2224            })
2225        );
2226        assert!(!MemKey::Slot(0).overlaps(MemKey::Arena {
2227            offset: 0,
2228            end: 256
2229        }));
2230    }
2231
2232    #[test]
2233    fn padding_and_strides_follow_row_major_order() {
2234        let shape = TensorShape {
2235            dtype: DType::FP32,
2236            dims: vec![2, 3, 4],
2237            elements: 24,
2238        };
2239        assert_eq!(shape.strides(), vec![12, 4, 1]);
2240        assert_eq!(shape.padded_dims(), [1, 1, 1, 2, 3, 4]);
2241        assert_eq!(pad_leading(&shape.strides(), 0), [0, 0, 0, 12, 4, 1]);
2242    }
2243}