virtio_accel_xdna/experiments/
bfp_experiment.rs1use virtio_accel_core::{ArtifactFormat, BackendError, TargetIdentity};
18
19pub const XDNA_BFP_EXPERIMENT_FORMAT: ArtifactFormat = match ArtifactFormat::new(0x5842_4650) {
22 Some(format) => format,
23 None => unreachable!(),
24};
25
26pub const XDNA_BFP_EXPERIMENT_TARGET_IDENTITY: TargetIdentity = TargetIdentity([
29 u32::from_le_bytes(*b"XBFP"),
30 1,
31 0,
32 0,
33 0,
34 0,
35 0,
36 0,
37 0,
38 0,
39 0,
40 0,
41]);
42
43const MAGIC: [u8; 4] = *b"XBFP";
44const VERSION: u32 = 1;
45const FLAVOR_MXINT8_MATMUL: u32 = 1;
46const HEADER_LEN: usize = 4 + 4 + 4 + 4 + 4 + 4 + 8 + 8;
47
48pub const UNIT_BYTES: u64 = 72;
50
51#[derive(Debug)]
53pub struct BfpExperimentArtifact<'a> {
54 pub m: u32,
55 pub k: u32,
56 pub n: u32,
57 pub xclbin: &'a [u8],
58 pub insts: &'a [u8],
59}
60
61impl<'a> BfpExperimentArtifact<'a> {
62 pub fn parse(bytes: &'a [u8]) -> Result<Self, BackendError> {
66 if bytes.len() < HEADER_LEN || bytes[0..4] != MAGIC {
67 return Err(BackendError::InvalidArgument);
68 }
69 let word =
70 |at: usize| u32::from_le_bytes(bytes[at..at + 4].try_into().expect("header word"));
71 if word(4) != VERSION {
72 return Err(BackendError::Incompatible);
73 }
74 if word(8) != FLAVOR_MXINT8_MATMUL {
75 return Err(BackendError::Unsupported);
76 }
77 let (m, k, n) = (word(12), word(16), word(20));
78 let xclbin_len = u64::from_le_bytes(bytes[24..32].try_into().expect("header word"));
79 let insts_len = u64::from_le_bytes(bytes[32..40].try_into().expect("header word"));
80
81 let xclbin_end = (HEADER_LEN as u64)
82 .checked_add(xclbin_len)
83 .ok_or(BackendError::InvalidArgument)?;
84 let total = xclbin_end
85 .checked_add(insts_len)
86 .ok_or(BackendError::InvalidArgument)?;
87 if total != bytes.len() as u64 || insts_len % 4 != 0 || xclbin_len == 0 || insts_len == 0 {
88 return Err(BackendError::InvalidArgument);
89 }
90
91 if m != 8 || n != 8 || !(32..=512).contains(&k) || k % 32 != 0 {
94 return Err(BackendError::Unsupported);
95 }
96
97 let xclbin_end = usize::try_from(xclbin_end).map_err(|_| BackendError::InvalidArgument)?;
98 Ok(Self {
99 m,
100 k,
101 n,
102 xclbin: &bytes[HEADER_LEN..xclbin_end],
103 insts: &bytes[xclbin_end..],
104 })
105 }
106
107 pub fn slot_bytes(&self) -> ([u64; 2], [u64; 1]) {
110 let operand = u64::from(self.k) / 8 * UNIT_BYTES;
111 let output = u64::from(self.m) * u64::from(self.n) * 4;
112 ([operand, operand], [output])
113 }
114
115 pub fn to_precompiled_container(&self) -> Vec<u8> {
118 let (inputs, outputs) = self.slot_bytes();
119 crate::artifact::encode("MLIR_AIE", &inputs, &outputs, self.xclbin, self.insts)
120 }
121}
122
123pub fn encode(m: u32, k: u32, n: u32, xclbin: &[u8], insts: &[u8]) -> Vec<u8> {
126 let mut out = Vec::with_capacity(HEADER_LEN + xclbin.len() + insts.len());
127 out.extend_from_slice(&MAGIC);
128 out.extend_from_slice(&VERSION.to_le_bytes());
129 out.extend_from_slice(&FLAVOR_MXINT8_MATMUL.to_le_bytes());
130 out.extend_from_slice(&m.to_le_bytes());
131 out.extend_from_slice(&k.to_le_bytes());
132 out.extend_from_slice(&n.to_le_bytes());
133 out.extend_from_slice(&(xclbin.len() as u64).to_le_bytes());
134 out.extend_from_slice(&(insts.len() as u64).to_le_bytes());
135 out.extend_from_slice(xclbin);
136 out.extend_from_slice(insts);
137 out
138}
139
140#[cfg(test)]
141mod tests {
142 use super::*;
143
144 fn sample(k: u32) -> Vec<u8> {
145 encode(8, k, 8, &[0xAA; 16], &[0xBB; 8])
146 }
147
148 #[test]
149 fn round_trips_and_derives_the_slot_plan() {
150 let bytes = sample(512);
151 let parsed = BfpExperimentArtifact::parse(&bytes).expect("valid container");
152 assert_eq!((parsed.m, parsed.k, parsed.n), (8, 512, 8));
153 assert_eq!(parsed.xclbin, &[0xAA; 16]);
154 assert_eq!(parsed.insts, &[0xBB; 8]);
155 let (inputs, outputs) = parsed.slot_bytes();
156 assert_eq!(inputs, [4608, 4608]);
157 assert_eq!(outputs, [256]);
158
159 let container = parsed.to_precompiled_container();
160 let inner =
161 crate::artifact::PrecompiledArtifact::parse(&container).expect("valid translation");
162 assert_eq!(inner.entry, "MLIR_AIE");
163 assert_eq!(inner.slot_bytes, [4608, 4608, 256]);
164 assert_eq!((inner.inputs, inner.outputs), (2, 1));
165 }
166
167 #[test]
168 fn rejects_bad_magic_version_flavor_and_framing() {
169 let mut bytes = sample(64);
170 bytes[0] = b'Y';
171 assert_eq!(
172 BfpExperimentArtifact::parse(&bytes).unwrap_err(),
173 BackendError::InvalidArgument
174 );
175 let mut bytes = sample(64);
176 bytes[4] = 9;
177 assert_eq!(
178 BfpExperimentArtifact::parse(&bytes).unwrap_err(),
179 BackendError::Incompatible
180 );
181 let mut bytes = sample(64);
182 bytes[8] = 2;
183 assert_eq!(
184 BfpExperimentArtifact::parse(&bytes).unwrap_err(),
185 BackendError::Unsupported
186 );
187 let bytes = sample(64);
188 assert_eq!(
189 BfpExperimentArtifact::parse(&bytes[..bytes.len() - 1]).unwrap_err(),
190 BackendError::InvalidArgument
191 );
192 assert_eq!(
193 BfpExperimentArtifact::parse(&bytes[..HEADER_LEN - 1]).unwrap_err(),
194 BackendError::InvalidArgument
195 );
196 }
197
198 #[test]
199 fn rejects_shapes_outside_the_proven_envelope() {
200 for (m, k, n) in [
201 (8, 24, 8),
202 (8, 544, 8),
203 (8, 33, 8),
204 (16, 64, 8),
205 (8, 64, 16),
206 ] {
207 let bytes = encode(m, k, n, &[1; 4], &[2; 4]);
208 assert_eq!(
209 BfpExperimentArtifact::parse(&bytes).unwrap_err(),
210 BackendError::Unsupported,
211 "m={m} k={k} n={n}"
212 );
213 }
214 }
215
216 #[test]
217 fn format_and_identity_collide_with_nothing_released() {
218 assert_ne!(
219 XDNA_BFP_EXPERIMENT_FORMAT,
220 crate::artifact::XDNA_PRECOMPILED_FORMAT
221 );
222 assert_ne!(
223 XDNA_BFP_EXPERIMENT_FORMAT.get(),
224 virtio_accel_tosa::ARTIFACT_FORMAT.get()
225 );
226 for target in [
227 crate::XDNA_TOSA_TARGET,
228 crate::XDNA_TOSA_INTEGER_TARGET,
229 crate::XDNA_TOSA_FP8_TARGET,
230 ] {
231 assert_ne!(target.to_identity(), XDNA_BFP_EXPERIMENT_TARGET_IDENTITY);
232 }
233 }
234}