Skip to main content

virtio_accel_tosa/
artifact.rs

1use core::fmt;
2
3use crate::Version;
4use virtio_accel_core::{ArtifactFormat, TargetIdentity};
5
6/// Raw TOSA FlatBuffer payload (`"TOSA"` as a big-endian four-character code).
7pub 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/// Set of TOSA base profiles implemented by a target.
16#[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/// TOSA implementation level.
51#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
52#[repr(u32)]
53pub enum Level {
54    /// No finite level is claimed; provider-specific limits still apply.
55    Unbounded = 0,
56    /// TOSA Level 8K.
57    Level8K = 1,
58}
59
60/// Argument ceilings assigned by a TOSA 1.0 implementation level.
61#[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/// TOSA 1.0 profile-extension bits.
110#[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/// Device-neutral target requirements carried in `virtio-accel`'s opaque target words.
155#[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    /// Check that every extension is paired with one of its permitted base profiles.
197    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}