Skip to main content

ruda_runtime/runtime/memory_management/memory_pool/
handle.rs

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