Skip to main content

ruda_runtime/runtime/memory_management/memory_pool/
handle.rs

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