Skip to main content

virtio_accel_tosa/
capability.rs

1use crate::{DType, Op, Target};
2
3macro_rules! flags {
4    ($name:ident, $bits:ty) => {
5        #[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
6        #[repr(transparent)]
7        pub struct $name($bits);
8
9        impl $name {
10            pub const NONE: Self = Self(0);
11
12            pub const fn union(self, other: Self) -> Self {
13                Self(self.0 | other.0)
14            }
15
16            pub const fn contains(self, other: Self) -> bool {
17                self.0 & other.0 == other.0
18            }
19
20            pub const fn bits(self) -> $bits {
21                self.0
22            }
23        }
24    };
25}
26
27flags!(ValueRoles, u8);
28
29impl ValueRoles {
30    /// Graph-visible block input.
31    pub const INPUT: Self = Self(1 << 0);
32    /// Graph-visible block output.
33    pub const OUTPUT: Self = Self(1 << 1);
34    /// Serialized constant tensor.
35    pub const CONSTANT: Self = Self(1 << 2);
36    /// Non-boundary value produced and consumed inside a graph.
37    pub const INTERMEDIATE: Self = Self(1 << 3);
38    /// Every tensor role.
39    pub const ALL: Self =
40        Self(Self::INPUT.0 | Self::OUTPUT.0 | Self::CONSTANT.0 | Self::INTERMEDIATE.0);
41}
42
43flags!(DTypeConstraints, u8);
44
45impl DTypeConstraints {
46    /// This dtype is accepted only when a constant is consumed as a compile-time operator
47    /// parameter and does not become a graph-visible provider tensor.
48    pub const PARAMETER_ONLY: Self = Self(1 << 0);
49}
50
51flags!(OperatorConstraints, u16);
52
53impl OperatorConstraints {
54    /// Every NaN-mode attribute on this operator must select propagating NaNs.
55    pub const PROPAGATING_NAN: Self = Self(1 << 0);
56    /// Pool padding values must all be zero.
57    pub const ZERO_PADDING: Self = Self(1 << 1);
58    /// Shape or permutation operands must be serialized compile-time constants.
59    pub const CONSTANT_PARAMETERS: Self = Self(1 << 2);
60    /// TOSA zero-point operands must be serialized zeros.
61    pub const ZERO_ZERO_POINTS: Self = Self(1 << 3);
62    /// The TOSA `MUL` shift operand must be a serialized zero.
63    pub const ZERO_SHIFT: Self = Self(1 << 4);
64}
65
66/// One dtype admitted in explicitly listed graph roles.
67#[derive(Clone, Copy, Debug, PartialEq, Eq)]
68pub struct DTypeCapability {
69    pub dtype: DType,
70    pub roles: ValueRoles,
71    pub constraints: DTypeConstraints,
72}
73
74impl DTypeCapability {
75    pub const fn new(dtype: DType, roles: ValueRoles) -> Self {
76        Self {
77            dtype,
78            roles,
79            constraints: DTypeConstraints::NONE,
80        }
81    }
82
83    pub const fn constrained(
84        dtype: DType,
85        roles: ValueRoles,
86        constraints: DTypeConstraints,
87    ) -> Self {
88        Self {
89            dtype,
90            roles,
91            constraints,
92        }
93    }
94}
95
96/// One implemented operator and the conservative restrictions a scheduler must preserve.
97#[derive(Clone, Copy, Debug, PartialEq, Eq)]
98pub struct OperatorCapability {
99    pub op: Op,
100    pub constraints: OperatorConstraints,
101}
102
103impl OperatorCapability {
104    pub const fn new(op: Op) -> Self {
105        Self {
106            op,
107            constraints: OperatorConstraints::NONE,
108        }
109    }
110
111    pub const fn constrained(op: Op, constraints: OperatorConstraints) -> Self {
112        Self { op, constraints }
113    }
114}
115
116/// Provider treatment of semantic runtime conditions derived during TOSA analysis.
117#[derive(Clone, Copy, Debug, PartialEq, Eq)]
118#[non_exhaustive]
119pub enum RuntimeConditionSupport {
120    /// No analyzed runtime conditions are accepted at program admission.
121    None,
122    /// Advisory `REQUIRE` conditions may remain, but mandatory dynamic conditions are rejected.
123    AdvisoryOnly,
124}
125
126/// Whole-graph structural boundary that is not reducible to an operator bit.
127#[derive(Clone, Copy, Debug, PartialEq, Eq)]
128pub struct GraphCapabilities {
129    /// Maximum regions accepted in one artifact.
130    pub max_regions: usize,
131    /// Maximum basic blocks accepted across all regions.
132    pub max_blocks: usize,
133    /// Whether graph-visible tensor dimensions may remain dynamic at load time.
134    pub dynamic_shapes: bool,
135    /// Runtime conditions the provider may retain after semantic analysis.
136    pub runtime_conditions: RuntimeConditionSupport,
137}
138
139/// Conservative semantic admission descriptor for one exact TOSA target.
140///
141/// A positive query means that scheduling a `load_program` attempt is lawful. It does not promise
142/// that concrete shapes, cross-operand relationships, resource availability, native compilation,
143/// or runtime device state will succeed. Program admission remains authoritative.
144#[derive(Clone, Copy, Debug)]
145pub struct CapabilityDescriptor {
146    pub target: Target,
147    pub dtypes: &'static [DTypeCapability],
148    pub operators: &'static [OperatorCapability],
149    pub graph: GraphCapabilities,
150}
151
152impl CapabilityDescriptor {
153    pub const fn dtype(self, dtype: DType) -> Option<DTypeCapability> {
154        let mut index = 0;
155        while index < self.dtypes.len() {
156            let capability = self.dtypes[index];
157            if capability.dtype.get() == dtype.get() {
158                return Some(capability);
159            }
160            index += 1;
161        }
162        None
163    }
164
165    pub const fn supports_dtype(self, dtype: DType, role: ValueRoles) -> bool {
166        match self.dtype(dtype) {
167            Some(capability) => capability.roles.contains(role),
168            None => false,
169        }
170    }
171
172    pub const fn operator(self, op: Op) -> Option<OperatorCapability> {
173        let mut index = 0;
174        while index < self.operators.len() {
175            let capability = self.operators[index];
176            if capability.op.get() == op.get() {
177                return Some(capability);
178            }
179            index += 1;
180        }
181        None
182    }
183
184    pub const fn supports_operator(self, op: Op) -> bool {
185        self.operator(op).is_some()
186    }
187}
188
189/// Optional host-side interface implemented by concrete TOSA providers.
190///
191/// Providers return an empty slice when their native runtime/device is unavailable. Each entry is
192/// an exact target/profile tier; consumers must not combine fields from different descriptors.
193pub trait TosaCapabilityProvider {
194    fn tosa_capabilities(&self) -> &'static [CapabilityDescriptor];
195}
196
197#[cfg(test)]
198mod tests {
199    use super::*;
200    use crate::{ExtensionSet, Level, ProfileSet, Version};
201
202    const DTYPES: &[DTypeCapability] = &[
203        DTypeCapability::new(DType::FP32, ValueRoles::ALL),
204        DTypeCapability::constrained(
205            DType::INT8,
206            ValueRoles::CONSTANT,
207            DTypeConstraints::PARAMETER_ONLY,
208        ),
209    ];
210    const OPERATORS: &[OperatorCapability] = &[
211        OperatorCapability::new(Op::IDENTITY),
212        OperatorCapability::constrained(Op::MAXIMUM, OperatorConstraints::PROPAGATING_NAN),
213    ];
214    const DESCRIPTOR: CapabilityDescriptor = CapabilityDescriptor {
215        target: Target::new(
216            Version::TOSA_1_0,
217            ProfileSet::FLOATING_POINT,
218            Level::Level8K,
219            ExtensionSet::NONE,
220        ),
221        dtypes: DTYPES,
222        operators: OPERATORS,
223        graph: GraphCapabilities {
224            max_regions: 1,
225            max_blocks: 1,
226            dynamic_shapes: false,
227            runtime_conditions: RuntimeConditionSupport::None,
228        },
229    };
230
231    #[test]
232    fn role_queries_do_not_turn_parameter_constants_into_boundaries() {
233        assert!(DESCRIPTOR.supports_dtype(DType::FP32, ValueRoles::INPUT));
234        assert!(DESCRIPTOR.supports_dtype(DType::INT8, ValueRoles::CONSTANT));
235        assert!(!DESCRIPTOR.supports_dtype(DType::INT8, ValueRoles::INPUT));
236        assert!(
237            DESCRIPTOR
238                .dtype(DType::INT8)
239                .unwrap()
240                .constraints
241                .contains(DTypeConstraints::PARAMETER_ONLY)
242        );
243    }
244
245    #[test]
246    fn operator_queries_return_scheduler_visible_restrictions() {
247        assert!(DESCRIPTOR.supports_operator(Op::IDENTITY));
248        assert!(!DESCRIPTOR.supports_operator(Op::ERF));
249        assert!(
250            DESCRIPTOR
251                .operator(Op::MAXIMUM)
252                .unwrap()
253                .constraints
254                .contains(OperatorConstraints::PROPAGATING_NAN)
255        );
256    }
257}