Skip to main content

virtio_accel_tosa/
specialize.rs

1use alloc::vec::Vec;
2use core::fmt;
3use core::num::NonZeroUsize;
4
5use crate::ValueId;
6
7const SHAPE_TAG: u64 = 0x5348_4150_4500_0001;
8const CTC_TAG: u64 = 0x4354_4300_0000_0002;
9
10/// Failure while constructing a canonical specialization key.
11#[derive(Clone, Copy, Debug, PartialEq, Eq)]
12pub enum SpecializationError {
13    ValuesOutOfOrder,
14    TooManyWords,
15    AllocationFailed,
16}
17
18impl fmt::Display for SpecializationError {
19    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
20        write!(formatter, "{self:?}")
21    }
22}
23
24/// Exact, collision-safe key for one set of dynamic shapes and CTC bytes.
25///
26/// The fingerprint accelerates rejection; equality always compares the canonical words as well.
27#[derive(Clone, Debug, PartialEq, Eq)]
28pub struct SpecializationKey {
29    words: Vec<u64>,
30    fingerprint: [u64; 2],
31}
32
33impl SpecializationKey {
34    pub fn words(&self) -> &[u64] {
35        &self.words
36    }
37
38    pub const fn fingerprint(&self) -> [u64; 2] {
39        self.fingerprint
40    }
41}
42
43/// Bounded builder for provider cache keys.
44///
45/// Values must be appended in increasing [`ValueId`] order, making equivalent submissions produce
46/// identical keys without a sorting allocation. Byte payloads are packed little-endian.
47#[derive(Debug)]
48pub struct SpecializationKeyBuilder {
49    words: Vec<u64>,
50    max_words: usize,
51    last_value: Option<ValueId>,
52}
53
54impl SpecializationKeyBuilder {
55    pub fn new(max_words: usize) -> Self {
56        Self {
57            words: Vec::new(),
58            max_words,
59            last_value: None,
60        }
61    }
62
63    pub fn push_shape(
64        &mut self,
65        value: ValueId,
66        dimensions: &[u64],
67    ) -> Result<(), SpecializationError> {
68        self.begin_value(value, SHAPE_TAG, dimensions.len(), dimensions.len())?;
69        self.words.extend_from_slice(dimensions);
70        Ok(())
71    }
72
73    pub fn push_ctc(&mut self, value: ValueId, bytes: &[u8]) -> Result<(), SpecializationError> {
74        let payload_words = bytes.len().div_ceil(8);
75        self.begin_value(value, CTC_TAG, bytes.len(), payload_words)?;
76        for chunk in bytes.chunks(8) {
77            let mut packed = [0_u8; 8];
78            packed[..chunk.len()].copy_from_slice(chunk);
79            self.words.push(u64::from_le_bytes(packed));
80        }
81        Ok(())
82    }
83
84    pub fn finish(self) -> SpecializationKey {
85        SpecializationKey {
86            fingerprint: fingerprint(&self.words),
87            words: self.words,
88        }
89    }
90
91    fn begin_value(
92        &mut self,
93        value: ValueId,
94        tag: u64,
95        payload_len: usize,
96        payload_words: usize,
97    ) -> Result<(), SpecializationError> {
98        if self.last_value.is_some_and(|prior| prior >= value) {
99            return Err(SpecializationError::ValuesOutOfOrder);
100        }
101        let additional = 3_usize
102            .checked_add(payload_words)
103            .ok_or(SpecializationError::TooManyWords)?;
104        // Reserve the complete record before mutating either the words or ordering state. A caller
105        // can therefore recover from a limit/allocation error and retry with a smaller value.
106        self.reserve(additional)?;
107        self.words.push(tag);
108        self.words.push(u64::from(value.get()));
109        self.words
110            .push(u64::try_from(payload_len).map_err(|_| SpecializationError::TooManyWords)?);
111        self.last_value = Some(value);
112        Ok(())
113    }
114
115    fn reserve(&mut self, additional: usize) -> Result<(), SpecializationError> {
116        if self
117            .words
118            .len()
119            .checked_add(additional)
120            .is_none_or(|required| required > self.max_words)
121        {
122            return Err(SpecializationError::TooManyWords);
123        }
124        self.words
125            .try_reserve_exact(additional)
126            .map_err(|_| SpecializationError::AllocationFailed)
127    }
128}
129
130fn fingerprint(words: &[u64]) -> [u64; 2] {
131    let mut first = 0xcbf2_9ce4_8422_2325_u64;
132    let mut second = 0x9e37_79b9_7f4a_7c15_u64;
133    for &word in words {
134        first ^= word;
135        first = first.wrapping_mul(0x0000_0100_0000_01b3);
136        second ^= word.wrapping_add(first.rotate_left(17));
137        second = second.rotate_left(27).wrapping_mul(0x94d0_49bb_1331_11eb);
138    }
139    [first, second ^ words.len() as u64]
140}
141
142#[derive(Debug)]
143struct CacheEntry<V> {
144    key: SpecializationKey,
145    value: V,
146    last_used: u64,
147}
148
149/// An insertion that could not reserve its one bounded cache entry.
150#[derive(Debug)]
151pub struct CacheInsertError<V> {
152    pub key: SpecializationKey,
153    pub value: V,
154}
155
156/// Small exact-key LRU for compiled shape specializations.
157///
158/// This portable cache is deliberately synchronization-free. A concurrent provider should put a
159/// lock or single-flight compilation state around it at its existing program/queue owner boundary.
160#[derive(Debug)]
161pub struct SpecializationCache<V> {
162    capacity: NonZeroUsize,
163    clock: u64,
164    entries: Vec<CacheEntry<V>>,
165}
166
167impl<V> SpecializationCache<V> {
168    pub fn new(capacity: NonZeroUsize) -> Self {
169        Self {
170            capacity,
171            clock: 0,
172            entries: Vec::new(),
173        }
174    }
175
176    pub const fn capacity(&self) -> usize {
177        self.capacity.get()
178    }
179
180    pub fn len(&self) -> usize {
181        self.entries.len()
182    }
183
184    pub fn is_empty(&self) -> bool {
185        self.entries.is_empty()
186    }
187
188    pub fn get(&mut self, key: &SpecializationKey) -> Option<&V> {
189        let index = self.position(key)?;
190        let stamp = self.tick();
191        self.entries[index].last_used = stamp;
192        Some(&self.entries[index].value)
193    }
194
195    pub fn get_mut(&mut self, key: &SpecializationKey) -> Option<&mut V> {
196        let index = self.position(key)?;
197        let stamp = self.tick();
198        self.entries[index].last_used = stamp;
199        Some(&mut self.entries[index].value)
200    }
201
202    /// Insert or replace a specialization, returning the replaced or evicted compiled value.
203    pub fn insert(
204        &mut self,
205        key: SpecializationKey,
206        value: V,
207    ) -> Result<Option<V>, CacheInsertError<V>> {
208        let stamp = self.tick();
209        if let Some(index) = self.position(&key) {
210            self.entries[index].last_used = stamp;
211            return Ok(Some(core::mem::replace(
212                &mut self.entries[index].value,
213                value,
214            )));
215        }
216        if self.entries.len() == self.capacity.get() {
217            let index = self
218                .entries
219                .iter()
220                .enumerate()
221                .min_by_key(|(_, entry)| entry.last_used)
222                .map(|(index, _)| index)
223                .expect("nonzero full cache");
224            let evicted = self.entries.swap_remove(index).value;
225            self.entries.push(CacheEntry {
226                key,
227                value,
228                last_used: stamp,
229            });
230            return Ok(Some(evicted));
231        }
232        if self.entries.try_reserve_exact(1).is_err() {
233            return Err(CacheInsertError { key, value });
234        }
235        self.entries.push(CacheEntry {
236            key,
237            value,
238            last_used: stamp,
239        });
240        Ok(None)
241    }
242
243    fn position(&self, key: &SpecializationKey) -> Option<usize> {
244        self.entries.iter().position(|entry| {
245            entry.key.fingerprint == key.fingerprint && entry.key.words == key.words
246        })
247    }
248
249    fn tick(&mut self) -> u64 {
250        if self.clock == u64::MAX {
251            self.entries.sort_unstable_by_key(|entry| entry.last_used);
252            for (index, entry) in self.entries.iter_mut().enumerate() {
253                entry.last_used = index as u64;
254            }
255            self.clock = self.entries.len() as u64;
256        }
257        let stamp = self.clock;
258        self.clock += 1;
259        stamp
260    }
261}
262
263#[cfg(test)]
264mod tests {
265    use super::*;
266
267    fn value(raw: u32) -> ValueId {
268        ValueId::from_raw(raw)
269    }
270
271    fn key(raw: u32) -> SpecializationKey {
272        let mut builder = SpecializationKeyBuilder::new(16);
273        builder
274            .push_shape(value(raw), &[1, u64::from(raw)])
275            .unwrap();
276        builder.finish()
277    }
278
279    #[test]
280    fn keys_are_canonical_bounded_and_collision_safe() {
281        let mut first = SpecializationKeyBuilder::new(16);
282        first.push_shape(value(1), &[1, 2, 3]).unwrap();
283        first
284            .push_ctc(value(2), &[1, 2, 3, 4, 5, 6, 7, 8, 9])
285            .unwrap();
286        let first = first.finish();
287
288        let mut second = SpecializationKeyBuilder::new(16);
289        second.push_shape(value(1), &[1, 2, 3]).unwrap();
290        second
291            .push_ctc(value(2), &[1, 2, 3, 4, 5, 6, 7, 8, 9])
292            .unwrap();
293        assert_eq!(first, second.finish());
294
295        let mut invalid = SpecializationKeyBuilder::new(8);
296        invalid.push_shape(value(2), &[1]).unwrap();
297        assert_eq!(
298            invalid.push_shape(value(1), &[1]),
299            Err(SpecializationError::ValuesOutOfOrder)
300        );
301
302        let mut retry = SpecializationKeyBuilder::new(4);
303        assert_eq!(
304            retry.push_shape(value(3), &[1, 2]),
305            Err(SpecializationError::TooManyWords)
306        );
307        retry.push_shape(value(3), &[1]).unwrap();
308        assert_eq!(retry.finish().words().len(), 4);
309    }
310
311    #[test]
312    fn cache_replaces_and_evicts_the_least_recently_used_value() {
313        let mut cache = SpecializationCache::new(NonZeroUsize::new(2).unwrap());
314        let one = key(1);
315        let two = key(2);
316        let three = key(3);
317        assert_eq!(cache.insert(one.clone(), 10).unwrap(), None);
318        assert_eq!(cache.insert(two.clone(), 20).unwrap(), None);
319        assert_eq!(cache.get(&one), Some(&10));
320        assert_eq!(cache.insert(three.clone(), 30).unwrap(), Some(20));
321        assert_eq!(cache.get(&two), None);
322        assert_eq!(cache.get(&three), Some(&30));
323        assert_eq!(cache.insert(one.clone(), 11).unwrap(), Some(10));
324        assert_eq!(cache.get(&one), Some(&11));
325    }
326}