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 pub const INPUT: Self = Self(1 << 0);
32 pub const OUTPUT: Self = Self(1 << 1);
34 pub const CONSTANT: Self = Self(1 << 2);
36 pub const INTERMEDIATE: Self = Self(1 << 3);
38 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 pub const PARAMETER_ONLY: Self = Self(1 << 0);
49}
50
51flags!(OperatorConstraints, u16);
52
53impl OperatorConstraints {
54 pub const PROPAGATING_NAN: Self = Self(1 << 0);
56 pub const ZERO_PADDING: Self = Self(1 << 1);
58 pub const CONSTANT_PARAMETERS: Self = Self(1 << 2);
60 pub const ZERO_ZERO_POINTS: Self = Self(1 << 3);
62 pub const ZERO_SHIFT: Self = Self(1 << 4);
64}
65
66#[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#[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#[derive(Clone, Copy, Debug, PartialEq, Eq)]
118#[non_exhaustive]
119pub enum RuntimeConditionSupport {
120 None,
122 AdvisoryOnly,
124}
125
126#[derive(Clone, Copy, Debug, PartialEq, Eq)]
128pub struct GraphCapabilities {
129 pub max_regions: usize,
131 pub max_blocks: usize,
133 pub dynamic_shapes: bool,
135 pub runtime_conditions: RuntimeConditionSupport,
137}
138
139#[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
189pub 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}