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