Skip to main content

ruda_runtime/runtime/memory_management/memory_pool/
handle.rs

1use crate::runtime::memory_management::MemoryHandle;
2use alloc::sync::Arc;
3#[cfg(target_has_atomic = "64")]
4use core::sync::atomic::{AtomicU64, Ordering};
5#[cfg(not(target_has_atomic = "64"))]
6use spin::Mutex;
7
8/// Managed Memory handle
9#[derive(Debug)]
10pub struct ManagedMemoryHandle {
11    descriptor: Arc<ManagedMemoryDescriptor>,
12    // Holds only the reference counts of the handle.
13    handle_count: Arc<()>,
14}
15
16/// Binding of a memory handle
17#[derive(Debug)]
18pub struct ManagedMemoryBinding {
19    descriptor: Arc<ManagedMemoryDescriptor>,
20}
21
22impl Clone for ManagedMemoryHandle {
23    fn clone(&self) -> Self {
24        Self {
25            descriptor: self.descriptor.clone(),
26            handle_count: self.handle_count.clone(),
27        }
28    }
29}
30
31/// Managed memory descriptor.
32///
33/// Host-side diagnostics can read the location while the device thread updates
34/// it. Protect the whole location so all readers observe a consistent snapshot.
35/// Send/Sync are derived from the fields; no unchecked Cell sharing is needed.
36pub(crate) struct ManagedMemoryDescriptor {
37    pub(crate) id: ManagedMemoryId,
38    #[cfg(target_has_atomic = "64")]
39    location: AtomicU64,
40    #[cfg(not(target_has_atomic = "64"))]
41    location: Mutex<MemoryLocation>,
42}
43
44impl core::fmt::Debug for ManagedMemoryDescriptor {
45    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
46        f.debug_struct("ManagedMemoryDescriptor")
47            .field("id", &self.id)
48            .field("location", &self.location())
49            .finish()
50    }
51}
52
53#[derive(Debug, PartialEq, Eq, Clone, Copy, Hash)]
54/// Managed memory unique identifier.
55pub struct ManagedMemoryId {
56    pub(crate) value: usize,
57}
58
59impl PartialEq for ManagedMemoryDescriptor {
60    fn eq(&self, other: &Self) -> bool {
61        self.id == other.id
62    }
63}
64
65impl Eq for ManagedMemoryDescriptor {}
66
67#[derive(Clone, Copy, Debug)]
68/// Defines where the [`ManagedMemoryId`] is located.
69pub(crate) struct MemoryLocation {
70    /// The memory pool index in the global memory management.
71    pub pool: u8,
72    /// The memory page index in a memory pool.
73    pub page: u16,
74    /// The memory slice index in a memory page.
75    pub slice: u32,
76    /// Whether the memory location is known/initialized.
77    pub init: u8,
78}
79
80impl ManagedMemoryDescriptor {
81    /// Update the memory location for the given [`ManagedMemoryId`].
82    pub(crate) fn update_location(&self, location: MemoryLocation) {
83        #[cfg(target_has_atomic = "64")]
84        {
85            self.location.store(location.to_bits(), Ordering::Release);
86        }
87        #[cfg(not(target_has_atomic = "64"))]
88        {
89            *self.location.lock() = location;
90        }
91    }
92
93    /// Update only the slice position for the given [`ManagedMemoryId`].
94    pub(crate) fn update_slice(&self, slice: u32) {
95        self.modify(|location| MemoryLocation { slice, ..location });
96    }
97
98    /// Update only the memory page position for the given [`ManagedMemoryId`].
99    pub fn update_page(&self, page: u16) {
100        self.modify(|location| MemoryLocation { page, ..location });
101    }
102
103    /// Retrieves the current location.
104    pub(crate) fn location(&self) -> MemoryLocation {
105        #[cfg(target_has_atomic = "64")]
106        {
107            MemoryLocation::from_bits(self.location.load(Ordering::Acquire))
108        }
109        #[cfg(not(target_has_atomic = "64"))]
110        {
111            *self.location.lock()
112        }
113    }
114
115    pub(crate) fn slice(&self) -> usize {
116        self.location().slice as usize
117    }
118
119    pub(crate) fn page(&self) -> usize {
120        self.location().page as usize
121    }
122
123    fn modify(&self, update: impl Fn(MemoryLocation) -> MemoryLocation) {
124        #[cfg(target_has_atomic = "64")]
125        {
126            let _ = self.location.fetch_update(Ordering::AcqRel, Ordering::Acquire, |bits| {
127                Some(update(MemoryLocation::from_bits(bits)).to_bits())
128            });
129        }
130        #[cfg(not(target_has_atomic = "64"))]
131        {
132            let mut location = self.location.lock();
133            *location = update(*location);
134        }
135    }
136}
137
138impl MemoryLocation {
139    #[cfg(target_has_atomic = "64")]
140    fn to_bits(self) -> u64 {
141        self.pool as u64
142            | (self.page as u64) << 8
143            | (self.slice as u64) << 24
144            | (self.init as u64) << 56
145    }
146
147    #[cfg(target_has_atomic = "64")]
148    fn from_bits(bits: u64) -> Self {
149        Self {
150            pool: bits as u8,
151            page: (bits >> 8) as u16,
152            slice: (bits >> 24) as u32,
153            init: (bits >> 56) as u8,
154        }
155    }
156
157    /// Creates a new memory location.
158    pub(crate) fn new(pool: u8, page: u16, slice: u32) -> Self {
159        Self {
160            pool,
161            page,
162            slice,
163            init: 1,
164        }
165    }
166
167    /// Creates a new uninitialized memory location.
168    pub(crate) fn uninit() -> Self {
169        Self {
170            pool: 0,
171            page: 0,
172            slice: 0,
173            init: 0,
174        }
175    }
176}
177
178impl ManagedMemoryHandle {
179    /// Creates a new managed memory handle.
180    pub fn new() -> Self {
181        let value = Self::gen_id();
182
183        Self {
184            descriptor: Arc::new(ManagedMemoryDescriptor {
185                id: ManagedMemoryId { value },
186                #[cfg(target_has_atomic = "64")]
187                location: AtomicU64::new(MemoryLocation::uninit().to_bits()),
188                #[cfg(not(target_has_atomic = "64"))]
189                location: Mutex::new(MemoryLocation::uninit()),
190            }),
191            handle_count: Arc::new(()),
192        }
193    }
194
195    /// Retrieves the descriptor for the current handle.
196    pub(crate) fn descriptor(&self) -> &ManagedMemoryDescriptor {
197        &self.descriptor
198    }
199
200    /// Return whether the current handle can be modified in-place.
201    pub fn can_mut(&self) -> bool {
202        Arc::strong_count(&self.handle_count) <= 2
203    }
204
205    /// Return whether the current handle is free.
206    pub fn is_free(&self) -> bool {
207        Arc::strong_count(&self.descriptor) <= 1
208    }
209
210    /// Returns the binding for the current handle.
211    pub fn binding(self) -> ManagedMemoryBinding {
212        ManagedMemoryBinding {
213            descriptor: self.descriptor.clone(),
214        }
215    }
216
217    fn gen_id() -> usize {
218        static COUNTER: core::sync::atomic::AtomicUsize = core::sync::atomic::AtomicUsize::new(0);
219        let value = COUNTER.fetch_add(1, core::sync::atomic::Ordering::Relaxed);
220        if value == usize::MAX {
221            core::panic!("Memory ID overflowed");
222        }
223        value
224    }
225}
226
227impl ManagedMemoryBinding {
228    /// Stable allocation identity, independent of the device address or view.
229    /// This is an identity token, not a pointer or proof that a resource is ready.
230    pub fn id(&self) -> ManagedMemoryId { self.descriptor.id }
231
232    /// Retrieves the descriptor for the current binding.
233    pub(crate) fn descriptor(&self) -> &ManagedMemoryDescriptor {
234        &self.descriptor
235    }
236}
237
238impl Default for ManagedMemoryHandle {
239    fn default() -> Self {
240        Self::new()
241    }
242}
243
244impl Clone for ManagedMemoryBinding {
245    fn clone(&self) -> Self {
246        Self {
247            descriptor: self.descriptor.clone(),
248        }
249    }
250}
251
252impl MemoryHandle<ManagedMemoryBinding> for ManagedMemoryHandle {
253    fn can_mut(&self) -> bool {
254        self.can_mut()
255    }
256
257    fn binding(self) -> ManagedMemoryBinding {
258        self.binding()
259    }
260}
261
262/// Calculates a best-effort heuristic for the alignment of row-aligned tensors.
263/// Prefers contiguous alignments for unit dimensions, 16-byte minimum alignment for non-unit,
264/// scaling with input size up to `buffer_align`.
265pub fn optimal_align(shape: usize, elem_size: usize, buffer_align: usize) -> usize {
266    if shape == 1 {
267        elem_size
268    } else {
269        (shape * elem_size)
270            .next_power_of_two()
271            .clamp(16, buffer_align)
272    }
273}
274
275#[cfg(test)]
276mod tests {
277    use super::*;
278
279    #[test]
280    fn test_memory_id_mutability() {
281        let handle1 = ManagedMemoryHandle::new();
282        handle1.descriptor().update_slice(4);
283        assert_eq!(handle1.descriptor().slice(), 4);
284
285        let handle2 = ManagedMemoryHandle::new();
286        handle2
287            .clone()
288            .descriptor()
289            .update_location(handle1.descriptor().location());
290        assert_eq!(handle2.descriptor().slice(), 4);
291    }
292
293    #[test]
294    fn test_location_visible_through_shared_arc() {
295        let handle = ManagedMemoryHandle::new();
296        let handle2 = handle.clone();
297
298        let location = MemoryLocation::new(1, 2, 3);
299        handle.descriptor().update_location(location);
300
301        assert_eq!(handle2.descriptor().location().pool, 1);
302        assert_eq!(handle2.descriptor().location().page, 2);
303        assert_eq!(handle2.descriptor().location().slice, 3);
304        assert_eq!(handle2.descriptor().location().init, 1);
305
306        handle.descriptor().update_slice(42);
307        assert_eq!(handle2.descriptor().slice(), 42);
308    }
309
310    #[test]
311    #[cfg(feature = "runtime-std")]
312    fn concurrent_debug_reads_consistent_location_snapshots() {
313        let handle = ManagedMemoryHandle::new();
314        let writer = handle.clone();
315        let task = std::thread::spawn(move || {
316            for i in 1..=128_u32 {
317                writer.descriptor().update_location(MemoryLocation::new(i as u8, i as u16, i));
318                std::thread::yield_now();
319            }
320        });
321        for _ in 0..128 {
322            let _ = alloc::format!("{handle:?}");
323            let location = handle.descriptor().location();
324            assert_eq!(location.page as u32, location.slice);
325            assert_eq!(location.pool as u32, location.slice);
326            std::thread::yield_now();
327        }
328        task.join().unwrap();
329        assert_eq!(handle.descriptor().slice(), 128);
330    }
331
332    #[test]
333    #[cfg(feature = "runtime-std")]
334    fn concurrent_field_updates_do_not_overwrite_each_other() {
335        let handle = ManagedMemoryHandle::new();
336        let first = handle.clone();
337        let second = handle.clone();
338        let page_writer = std::thread::spawn(move || {
339            for page in 1..=128 { first.descriptor().update_page(page); }
340        });
341        let slice_writer = std::thread::spawn(move || {
342            for slice in 1..=128 { second.descriptor().update_slice(slice); }
343        });
344        page_writer.join().unwrap();
345        slice_writer.join().unwrap();
346        assert_eq!(handle.descriptor().page(), 128);
347        assert_eq!(handle.descriptor().slice(), 128);
348    }
349
350}