use std::{
cell::Cell,
ffi::{CStr, c_char, c_int, c_uint, c_void},
marker::PhantomData,
process::abort,
ptr::{NonNull, null_mut},
slice::{from_raw_parts, from_raw_parts_mut},
};
unsafe extern "C" {
fn hipMallocManaged(dev_ptr: *mut *mut c_void, size: usize, flags: c_uint) -> c_int;
fn hipFree(ptr: *mut c_void) -> c_int;
fn hipGetErrorString(hipError: c_int) -> *const c_char;
fn hipStreamSynchronize(stream: *mut c_void) -> c_int;
}
pub(crate) fn hip_check(err: c_int, file: &str, line: u32) {
if err != 0 {
let err_ptr = unsafe { hipGetErrorString(err) };
let err_str = if err_ptr.is_null() {
String::from("[lufloat error] unknown.")
} else {
unsafe { CStr::from_ptr(err_ptr) }
.to_string_lossy()
.into_owned()
};
eprintln!(
"[lufloat error] {} (Code: {}) at {}:{}.",
err_str, err, file, line
);
abort();
}
}
fn hip_malloc(size: usize) -> *mut c_void {
let mut ptr = null_mut();
let err = unsafe { hipMallocManaged(&mut ptr, size, 1) };
hip_check(err, file!(), line!());
ptr
}
fn hip_free(ptr: *mut c_void) {
let err = unsafe { hipFree(ptr) };
hip_check(err, file!(), line!());
}
pub struct Arena {
base_ptr: NonNull<u8>,
capacity: usize,
offset: Cell<usize>,
}
impl Arena {
pub fn new(len: usize) -> Self {
assert_ne!(len, 0);
assert_eq!(len % 2048, 0);
let capacity = len << 1;
Self {
base_ptr: NonNull::new(hip_malloc(capacity) as *mut u8).unwrap(),
capacity,
offset: Cell::new(0),
}
}
fn alloc(&self, size: usize) -> NonNull<u8> {
let current = self.offset.get();
let end = current + size;
assert!(end <= self.capacity);
let ptr = unsafe { self.base_ptr.as_ptr().add(current) };
self.offset.set(end);
NonNull::new(ptr).unwrap()
}
pub fn reset(&mut self) {
let err = unsafe { hipStreamSynchronize(null_mut()) };
hip_check(err, file!(), line!());
self.offset.set(0);
}
}
impl Drop for Arena {
fn drop(&mut self) {
let err = unsafe { hipStreamSynchronize(null_mut()) };
hip_check(err, file!(), line!());
hip_free(self.base_ptr.as_ptr() as *mut c_void);
}
}
pub struct UnifiedBuffer<'a> {
pub(crate) ptr: *mut u16,
pub(crate) len: usize,
_marker: PhantomData<&'a Arena>,
}
impl<'a> UnifiedBuffer<'a> {
pub fn new(arena: &'a Arena, len: usize) -> Self {
assert_ne!(len, 0);
assert_eq!(len % 2048, 0);
UnifiedBuffer {
ptr: arena.alloc(len << 1).as_ptr() as *mut u16,
len,
_marker: PhantomData,
}
}
pub fn slice(&self) -> &[u16] {
let err = unsafe { hipStreamSynchronize(null_mut()) };
hip_check(err, file!(), line!());
unsafe { from_raw_parts(self.ptr, self.len) }
}
pub fn slice_mut(&mut self) -> &mut [u16] {
let err = unsafe { hipStreamSynchronize(null_mut()) };
hip_check(err, file!(), line!());
unsafe { from_raw_parts_mut(self.ptr, self.len) }
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[should_panic]
fn arena_zero() {
let _ = Arena::new(0);
}
#[test]
#[should_panic]
fn arena_unaligned() {
let _ = Arena::new(2047);
}
#[test]
#[should_panic]
fn buffer_zero() {
let arena = Arena::new(2048);
let _ = UnifiedBuffer::new(&arena, 0);
}
#[test]
#[should_panic]
fn buffer_unaligned() {
let arena = Arena::new(2048);
let _ = UnifiedBuffer::new(&arena, 2047);
}
#[test]
#[should_panic]
fn buffer_overalloc() {
let arena = Arena::new(2048);
let _ = UnifiedBuffer::new(&arena, 4096);
}
#[test]
fn arena_reset() {
let mut arena = Arena::new(4096);
let _ = UnifiedBuffer::new(&arena, 4096);
arena.reset();
let _ = UnifiedBuffer::new(&arena, 4096);
}
}