1use core::fmt;
2
3use crate::Version;
4use virtio_accel_core::{ArtifactFormat, TargetIdentity};
5
6pub const ARTIFACT_FORMAT: ArtifactFormat = match ArtifactFormat::new(0x544f_5341) {
8 Some(format) => format,
9 None => panic!("TOSA artifact format must be nonzero"),
10};
11
12const TARGET_MAGIC: u32 = 0x544f_5341;
13const TARGET_ABI: u32 = 1;
14
15#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
17#[repr(transparent)]
18pub struct ProfileSet(u32);
19
20impl ProfileSet {
21 pub const INTEGER: Self = Self(1 << 0);
22 pub const FLOATING_POINT: Self = Self(1 << 1);
23 pub const ALL: Self = Self(Self::INTEGER.0 | Self::FLOATING_POINT.0);
24
25 pub const fn from_bits(bits: u32) -> Option<Self> {
26 if bits != 0 && bits & !Self::ALL.0 == 0 {
27 Some(Self(bits))
28 } else {
29 None
30 }
31 }
32
33 pub const fn bits(self) -> u32 {
34 self.0
35 }
36
37 pub const fn union(self, other: Self) -> Self {
38 Self(self.0 | other.0)
39 }
40
41 pub const fn contains(self, other: Self) -> bool {
42 self.0 & other.0 == other.0
43 }
44
45 pub const fn intersects(self, other: Self) -> bool {
46 self.0 & other.0 != 0
47 }
48}
49
50#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
52#[repr(u32)]
53pub enum Level {
54 Unbounded = 0,
56 Level8K = 1,
58}
59
60#[derive(Clone, Copy, Debug, PartialEq, Eq)]
62pub struct LevelLimits {
63 pub max_rank: usize,
64 pub max_kernel: i32,
65 pub max_stride: i32,
66 pub max_scale: i32,
67 pub max_log2_size: u32,
68 pub max_nesting: usize,
69 pub max_tensor_list_size: usize,
70}
71
72impl Level {
73 pub const fn limits(self) -> LevelLimits {
74 match self {
75 Self::Unbounded => LevelLimits {
76 max_rank: 32,
77 max_kernel: i32::MAX,
78 max_stride: i32::MAX,
79 max_scale: 2_048,
80 max_log2_size: 63,
81 max_nesting: 256,
82 max_tensor_list_size: 256,
83 },
84 Self::Level8K => LevelLimits {
85 max_rank: 6,
86 max_kernel: 8_192,
87 max_stride: 8_192,
88 max_scale: 256,
89 max_log2_size: 31,
90 max_nesting: 6,
91 max_tensor_list_size: 64,
92 },
93 }
94 }
95}
96
97impl TryFrom<u32> for Level {
98 type Error = TargetError;
99
100 fn try_from(value: u32) -> Result<Self, Self::Error> {
101 match value {
102 0 => Ok(Self::Unbounded),
103 1 => Ok(Self::Level8K),
104 _ => Err(TargetError::UnknownLevel(value)),
105 }
106 }
107}
108
109#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
111#[repr(transparent)]
112pub struct ExtensionSet(u64);
113
114impl ExtensionSet {
115 pub const NONE: Self = Self(0);
116 pub const INT16: Self = Self(1 << 0);
117 pub const INT4: Self = Self(1 << 1);
118 pub const BF16: Self = Self(1 << 2);
119 pub const FP8E4M3: Self = Self(1 << 3);
120 pub const FP8E5M2: Self = Self(1 << 4);
121 pub const FFT: Self = Self(1 << 5);
122 pub const VARIABLE: Self = Self(1 << 6);
123 pub const CONTROL_FLOW: Self = Self(1 << 7);
124 pub const DYNAMIC: Self = Self(1 << 8);
125 pub const DOUBLE_ROUND: Self = Self(1 << 9);
126 pub const INEXACT_ROUND: Self = Self(1 << 10);
127 pub const ALL: Self = Self((1 << 11) - 1);
128
129 pub const fn from_bits(bits: u64) -> Option<Self> {
130 if bits & !Self::ALL.0 == 0 {
131 Some(Self(bits))
132 } else {
133 None
134 }
135 }
136
137 pub const fn bits(self) -> u64 {
138 self.0
139 }
140
141 pub const fn union(self, other: Self) -> Self {
142 Self(self.0 | other.0)
143 }
144
145 pub const fn contains(self, other: Self) -> bool {
146 self.0 & other.0 == other.0
147 }
148
149 pub const fn intersects(self, other: Self) -> bool {
150 self.0 & other.0 != 0
151 }
152}
153
154#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
156pub struct Target {
157 pub version: Version,
158 pub profiles: ProfileSet,
159 pub level: Level,
160 pub extensions: ExtensionSet,
161}
162
163impl Target {
164 pub const fn new(
165 version: Version,
166 profiles: ProfileSet,
167 level: Level,
168 extensions: ExtensionSet,
169 ) -> Self {
170 Self {
171 version,
172 profiles,
173 level,
174 extensions,
175 }
176 }
177
178 pub const fn to_identity(self) -> TargetIdentity {
179 let extensions = self.extensions.bits();
180 TargetIdentity([
181 TARGET_MAGIC,
182 TARGET_ABI,
183 self.version.major as u32,
184 self.version.minor as u32,
185 self.version.patch as u32,
186 self.profiles.bits(),
187 self.level as u32,
188 extensions as u32,
189 (extensions >> 32) as u32,
190 0,
191 0,
192 0,
193 ])
194 }
195
196 pub fn validate(self) -> Result<Self, TargetError> {
198 let integer_only = ExtensionSet::INT16
199 .union(ExtensionSet::INT4)
200 .union(ExtensionSet::DOUBLE_ROUND)
201 .union(ExtensionSet::INEXACT_ROUND);
202 let floating_only = ExtensionSet::BF16
203 .union(ExtensionSet::FP8E4M3)
204 .union(ExtensionSet::FP8E5M2)
205 .union(ExtensionSet::FFT);
206 if self.extensions.intersects(integer_only)
207 && !self.profiles.intersects(ProfileSet::INTEGER)
208 {
209 return Err(TargetError::ExtensionProfileMismatch {
210 extensions: ExtensionSet(self.extensions.bits() & integer_only.bits()),
211 required_profiles: ProfileSet::INTEGER,
212 });
213 }
214 if self.extensions.intersects(floating_only)
215 && !self.profiles.intersects(ProfileSet::FLOATING_POINT)
216 {
217 return Err(TargetError::ExtensionProfileMismatch {
218 extensions: ExtensionSet(self.extensions.bits() & floating_only.bits()),
219 required_profiles: ProfileSet::FLOATING_POINT,
220 });
221 }
222 Ok(self)
223 }
224
225 pub fn from_identity(identity: TargetIdentity) -> Result<Self, TargetError> {
226 let words = identity.0;
227 if words[0] != TARGET_MAGIC {
228 return Err(TargetError::WrongMagic(words[0]));
229 }
230 if words[1] != TARGET_ABI {
231 return Err(TargetError::UnknownAbi(words[1]));
232 }
233 if words[2] > u16::MAX as u32 || words[3] > u16::MAX as u32 || words[4] > u16::MAX as u32 {
234 return Err(TargetError::VersionOutOfRange);
235 }
236 let profiles =
237 ProfileSet::from_bits(words[5]).ok_or(TargetError::InvalidProfiles(words[5]))?;
238 let level = Level::try_from(words[6])?;
239 let extension_bits = u64::from(words[7]) | (u64::from(words[8]) << 32);
240 let extensions = ExtensionSet::from_bits(extension_bits)
241 .ok_or(TargetError::UnknownExtensions(extension_bits))?;
242 if words[9..].iter().any(|word| *word != 0) {
243 return Err(TargetError::ReservedWords);
244 }
245 Self {
246 version: Version::new(words[2] as u16, words[3] as u16, words[4] as u16),
247 profiles,
248 level,
249 extensions,
250 }
251 .validate()
252 }
253}
254
255#[derive(Clone, Copy, Debug, PartialEq, Eq)]
256pub enum TargetError {
257 WrongMagic(u32),
258 UnknownAbi(u32),
259 VersionOutOfRange,
260 VersionMismatch {
261 target: Version,
262 model: Version,
263 },
264 InvalidProfiles(u32),
265 UnknownLevel(u32),
266 UnknownExtensions(u64),
267 ExtensionProfileMismatch {
268 extensions: ExtensionSet,
269 required_profiles: ProfileSet,
270 },
271 ReservedWords,
272}
273
274impl fmt::Display for TargetError {
275 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
276 write!(formatter, "{self:?}")
277 }
278}
279
280#[derive(Clone, Copy, Debug, PartialEq, Eq)]
281pub enum ArtifactError {
282 VersionMismatch { model: Version, target: Version },
283}
284
285impl fmt::Display for ArtifactError {
286 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
287 write!(formatter, "{self:?}")
288 }
289}