use std::{collections::BTreeMap, sync::Arc};
use cudarc::driver::{CudaContext, CudaSlice, CudaStream, DriverError};
fn next_rng(state: &mut u64) -> u64 {
*state = state.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
*state >> 33
}
fn main() {
let context = CudaContext::new(0).unwrap();
let stream = context.default_stream();
let mut seg = Segment::new(1 << 20, &stream).unwrap(); let mut ledger: Vec<(usize, usize)> = Vec::new(); let mut state = 12345u64;
for _ in 0..5000 {
if ledger.is_empty() || next_rng(&mut state) % 4 != 0 {
let want = 256 * (1 + next_rng(&mut state) as usize % 512); if let Some(off) = seg.alloc_raw(want) {
ledger.push((off, want));
}
} else {
let idx = (next_rng(&mut state) as usize) % ledger.len();
let (off, len) = ledger.remove(idx);
seg.free_raw(off, len);
}
seg.assert_consistency();
}
for (off, len) in ledger.drain(..) {
seg.free_raw(off, len);
}
seg.assert_consistency();
assert_eq!(seg.live(), 0);
assert_eq!(seg.free_map().len(), 1, "all frees must coalesce into one run");
assert_eq!(seg.free_map().get(&0), Some(&seg.total()), "segment must be whole again");
}
pub struct Segment {
storage: Arc<CudaSlice<u8>>,
total: usize,
live: usize,
free: BTreeMap<usize, usize>,
}
impl Segment {
pub fn new(total: usize, stream: &Arc<CudaStream>) -> Result<Self, DriverError> {
let total = align_up(total.max(1), 256);
let storage = Arc::new(unsafe { stream.alloc::<u8>(total)? });
let mut free = BTreeMap::new();
free.insert(0, total);
Ok(Segment { storage, total, live: 0, free })
}
pub fn total(&self) -> usize {
self.total
}
pub fn live(&self) -> usize {
self.live
}
pub fn storage(&self) -> &CudaSlice<u8> {
&self.storage
}
pub fn free_map(&self) -> &BTreeMap<usize, usize> {
&self.free
}
pub fn alloc_raw(&mut self, bytes: usize) -> Option<usize> {
let bytes = align_up(bytes, 256);
if bytes == 0 || bytes > self.total {
return None;
}
let mut best: Option<(usize, usize)> = None;
for (&offset, &size) in self.free.iter() {
if size >= bytes {
match best {
None => best = Some((offset, size)),
Some((_, bs)) if size < bs => best = Some((offset, size)),
_ => {}
}
}
}
let (offset, origin) = best?;
self.free.remove(&offset);
let leftover = origin - bytes;
if leftover > 0 {
self.free.insert(offset + bytes, leftover);
}
self.live += bytes;
Some(offset)
}
pub fn free_raw(&mut self, offset: usize, bytes: usize) {
assert!(offset % 256 == 0, "offset must be 256-aligned");
let bytes = align_up(bytes, 256);
assert!(bytes > 0, "cannot free 0 bytes");
assert!(offset + bytes <= self.total, "free range out of segment");
assert!(self.live >= bytes, "live underflow / double free");
let mut start = offset;
if let Some((&k, &v)) = self.free.range(..offset).next_back() {
assert!(k + v <= offset, "double free: overlaps left free run");
if k + v == offset {
start = k; }
}
let mut end = offset + bytes;
if let Some((&k, &v)) = self.free.range(offset..).next() {
assert!(k >= offset + bytes, "double free: overlaps right free run");
if k == offset + bytes {
end = k + v; }
}
self.free.retain(|&k, _| k < start || k >= end);
self.free.insert(start, end - start);
self.live -= bytes;
}
pub fn assert_consistency(&self) {
let mut sum_free = 0usize;
let mut prev_key: Option<usize> = None;
for (&k, &v) in self.free.iter() {
assert!(v > 0, "zero-length free run at {k}");
assert!(k + v <= self.total, "free run [{k}, {}) exceeds total {}", k + v, self.total);
if let Some(pk) = prev_key {
assert!(
pk + self.free[&pk] < k,
"free runs at {pk} and {k} are touching (coalesce missed)"
);
}
sum_free += v;
prev_key = Some(k);
}
assert_eq!(self.live + sum_free, self.total, "conservation broken: live+free != total");
}
}
fn align_up(v: usize, a: usize) -> usize {
(v + a - 1) & !(a - 1)
}