use crate::device::metal_device;
use metal::Buffer;
use rlx_ir::{DType, Graph, NodeId};
use rlx_opt::memory::MemoryPlan;
use std::collections::HashMap;
pub struct Arena {
pub buffer: Buffer,
pub size_bytes: usize,
pub offsets: HashMap<NodeId, usize>, pub element_counts: HashMap<NodeId, usize>, pub dtypes: HashMap<NodeId, DType>, }
impl Arena {
pub fn from_plan(plan: MemoryPlan) -> Self {
Self::from_plan_with_graph(plan, None)
}
pub fn from_plan_with_graph(plan: MemoryPlan, graph: Option<&Graph>) -> Self {
let dev = metal_device().expect("Metal device required for rlx-metal arena");
let buffer = dev.alloc_shared(plan.arena_size.max(64));
let mut offsets = HashMap::with_capacity(plan.assignments.len());
let mut element_counts = HashMap::with_capacity(plan.assignments.len());
let mut dtypes = HashMap::with_capacity(plan.assignments.len());
for (node_id, slot) in &plan.assignments {
offsets.insert(*node_id, slot.offset);
let dt = graph
.map(|g| g.node(*node_id).shape.dtype())
.unwrap_or(DType::F32);
let elem_size = dt.size_bytes();
element_counts.insert(*node_id, slot.size / elem_size.max(1));
dtypes.insert(*node_id, dt);
}
Self {
buffer,
size_bytes: plan.arena_size,
offsets,
element_counts,
dtypes,
}
}
pub fn has_buffer(&self, id: NodeId) -> bool {
self.offsets.contains_key(&id)
}
pub fn byte_offset(&self, id: NodeId) -> usize {
*self.offsets.get(&id).expect("node not in arena")
}
pub fn dtype(&self, id: NodeId) -> DType {
self.dtypes.get(&id).copied().unwrap_or(DType::F32)
}
pub fn slice_mut(&mut self, id: NodeId) -> &mut [f32] {
debug_assert_eq!(self.dtype(id), DType::F32);
let off = self.byte_offset(id);
let len = *self.element_counts.get(&id).unwrap_or(&0);
unsafe {
let ptr = self.buffer.contents() as *mut u8;
std::slice::from_raw_parts_mut(ptr.add(off) as *mut f32, len)
}
}
pub fn slice(&self, id: NodeId) -> &[f32] {
debug_assert_eq!(self.dtype(id), DType::F32);
let off = self.byte_offset(id);
let len = *self.element_counts.get(&id).unwrap_or(&0);
unsafe {
let ptr = self.buffer.contents() as *const u8;
std::slice::from_raw_parts(ptr.add(off) as *const f32, len)
}
}
pub fn read_as_f32(&self, id: NodeId) -> Vec<f32> {
let dt = self.dtype(id);
let off = self.byte_offset(id);
let len = *self.element_counts.get(&id).unwrap_or(&0);
unsafe {
let base = (self.buffer.contents() as *const u8).add(off);
match dt {
DType::F32 => std::slice::from_raw_parts(base as *const f32, len).to_vec(),
DType::F16 => {
let src = std::slice::from_raw_parts(base as *const half::f16, len);
src.iter().map(|h| h.to_f32()).collect()
}
_ => std::slice::from_raw_parts(base as *const f32, len).to_vec(),
}
}
}
pub fn write_from_f32(&mut self, id: NodeId, data: &[f32]) {
let dt = self.dtype(id);
let off = self.byte_offset(id);
let cap = *self.element_counts.get(&id).unwrap_or(&0);
let len = data.len().min(cap);
unsafe {
let base = (self.buffer.contents() as *mut u8).add(off);
match dt {
DType::F32 => {
std::ptr::copy_nonoverlapping(data.as_ptr(), base as *mut f32, len);
}
DType::F16 => {
let dst = std::slice::from_raw_parts_mut(base as *mut half::f16, len);
for (i, &v) in data.iter().take(len).enumerate() {
dst[i] = half::f16::from_f32(v);
}
}
DType::BF16 => {
let dst = std::slice::from_raw_parts_mut(base as *mut half::bf16, len);
for (i, &v) in data.iter().take(len).enumerate() {
dst[i] = half::bf16::from_f32(v);
}
}
DType::I32 => {
let dst = std::slice::from_raw_parts_mut(base as *mut i32, len);
for (i, &v) in data.iter().take(len).enumerate() {
dst[i] = v as i32;
}
}
DType::I64 => {
let dst = std::slice::from_raw_parts_mut(base as *mut i64, len);
for (i, &v) in data.iter().take(len).enumerate() {
dst[i] = v as i64;
}
}
DType::U32 => {
let dst = std::slice::from_raw_parts_mut(base as *mut u32, len);
for (i, &v) in data.iter().take(len).enumerate() {
dst[i] = v as u32;
}
}
DType::I16 => {
let dst = std::slice::from_raw_parts_mut(base as *mut i16, len);
for (i, &v) in data.iter().take(len).enumerate() {
dst[i] = v as i16;
}
}
DType::I8 => {
let dst = std::slice::from_raw_parts_mut(base as *mut i8, len);
for (i, &v) in data.iter().take(len).enumerate() {
dst[i] = v as i8;
}
}
DType::U8 => {
let dst = std::slice::from_raw_parts_mut(base, len);
for (i, &v) in data.iter().take(len).enumerate() {
dst[i] = v as u8;
}
}
DType::Bool => {
let dst = std::slice::from_raw_parts_mut(base, len);
for (i, &v) in data.iter().take(len).enumerate() {
dst[i] = if v != 0.0 { 1 } else { 0 };
}
}
_ => {
std::ptr::copy_nonoverlapping(data.as_ptr(), base as *mut f32, len);
}
}
}
}
pub fn copy_node_f32(&self, dst: NodeId, src: NodeId) {
let dst_len = *self.element_counts.get(&dst).unwrap_or(&0);
let src_len = *self.element_counts.get(&src).unwrap_or(&0);
self.copy_node_f32_prefix(dst, src, dst_len.min(src_len));
}
pub fn copy_node_f32_range(
&self,
dst: NodeId,
dst_elem: usize,
src: NodeId,
src_elem: usize,
n: usize,
) {
let dst_cap = *self.element_counts.get(&dst).unwrap_or(&0);
let src_cap = *self.element_counts.get(&src).unwrap_or(&0);
if n == 0 || dst_elem + n > dst_cap || src_elem + n > src_cap {
return;
}
let dst_off = self.byte_offset(dst);
let src_off = self.byte_offset(src);
unsafe {
let base = self.buffer.contents() as *mut u8;
let src_p = (base.add(src_off) as *const f32).add(src_elem);
let dst_p = (base.add(dst_off) as *mut f32).add(dst_elem);
if src_p as *const () != dst_p as *const () {
std::ptr::copy(src_p, dst_p, n);
}
}
}
pub fn copy_node_f32_prefix(&self, dst: NodeId, src: NodeId, elems: usize) {
if elems == 0 {
return;
}
let dst_off = self.byte_offset(dst);
let src_off = self.byte_offset(src);
let dst_cap = *self.element_counts.get(&dst).unwrap_or(&0);
let src_cap = *self.element_counts.get(&src).unwrap_or(&0);
let len = elems.min(dst_cap).min(src_cap);
if len == 0 {
return;
}
unsafe {
let base = self.buffer.contents() as *mut u8;
std::ptr::copy(
base.add(src_off) as *const f32,
base.add(dst_off) as *mut f32,
len,
);
}
}
pub fn write_bytes(&mut self, id: NodeId, data: &[u8]) {
let off = self.byte_offset(id);
let cap = *self.element_counts.get(&id).unwrap_or(&0);
let len = data.len().min(cap);
unsafe {
let base = (self.buffer.contents() as *mut u8).add(off);
std::ptr::copy_nonoverlapping(data.as_ptr(), base, len);
}
}
}