use std::any::Any;
use std::sync::{Arc, OnceLock};
use sim_kernel::{DefaultFactory, Error, Factory, Result, Symbol, Value};
use sim_lib_numbers_tensor::{Tensor, TensorLocation, TensorStorage};
use crate::model::{ModeledComputeFault, ModeledResidentSegment, ModeledTensorExecutor};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct ResidentHandle {
symbol: Symbol,
}
impl ResidentHandle {
pub(crate) fn new(id: usize) -> Self {
Self {
symbol: Symbol::qualified("compute.alloc", id.to_string()),
}
}
pub fn symbol(&self) -> &Symbol {
&self.symbol
}
}
pub(crate) struct ModeledResidentDescriptor {
pub(crate) site: Symbol,
pub(crate) allocation: ResidentHandle,
pub(crate) segments: Vec<ModeledResidentSegment>,
pub(crate) shape: Vec<usize>,
pub(crate) dtype: Symbol,
}
pub struct ModeledResidentStorage {
site: Symbol,
allocation: ResidentHandle,
segments: Arc<[ModeledResidentSegment]>,
shape: Arc<[usize]>,
dtype: Symbol,
cells: Arc<[Value]>,
executor: ModeledTensorExecutor,
fault: Option<ModeledComputeFault>,
materialized: OnceLock<Result<Arc<dyn TensorStorage>>>,
}
impl ModeledResidentStorage {
pub(crate) fn new(
descriptor: ModeledResidentDescriptor,
cells: Arc<[Value]>,
executor: ModeledTensorExecutor,
fault: Option<ModeledComputeFault>,
) -> Self {
Self {
site: descriptor.site,
allocation: descriptor.allocation,
segments: descriptor.segments.into(),
shape: descriptor.shape.into(),
dtype: descriptor.dtype,
cells,
executor,
fault,
materialized: OnceLock::new(),
}
}
pub fn allocation(&self) -> &ResidentHandle {
&self.allocation
}
pub fn segments(&self) -> &[ModeledResidentSegment] {
&self.segments
}
pub fn resident_tensor(&self) -> Option<Tensor> {
Tensor::from_storage(
self.shape.to_vec(),
self.dtype.clone(),
Arc::new(BoxedTensorStorageForResident::new(
self.dtype.clone(),
self.cells.clone(),
)),
)
.ok()
}
}
impl TensorStorage for ModeledResidentStorage {
fn dtype(&self) -> &Symbol {
&self.dtype
}
fn len(&self) -> usize {
self.cells.len()
}
fn location(&self) -> TensorLocation {
TensorLocation::Resident {
site: self.site.clone(),
allocation: self.allocation.symbol().clone(),
}
}
fn cell(&self, index: usize) -> Result<Value> {
let storage = self.materialize()?;
storage.cell(index)
}
fn materialize(&self) -> Result<Arc<dyn TensorStorage>> {
self.materialized
.get_or_init(|| {
self.executor.increment_readbacks();
if !self.executor.is_resident_active(&self.allocation) {
self.executor.increment_materialization_failures();
return Err(Error::Eval(
"modeled compute resident allocation was evicted".to_owned(),
));
}
if self.fault == Some(ModeledComputeFault::ReadbackFailure) {
self.executor.increment_materialization_failures();
return Err(Error::Eval("modeled compute readback failed".to_owned()));
}
Ok(Arc::new(BoxedTensorStorageForResident::new(
self.dtype.clone(),
self.cells.clone(),
)))
})
.clone()
}
fn as_any(&self) -> &dyn Any {
self
}
}
struct BoxedTensorStorageForResident {
dtype: Symbol,
cells: Arc<[Value]>,
}
impl BoxedTensorStorageForResident {
fn new(dtype: Symbol, cells: Arc<[Value]>) -> Self {
Self { dtype, cells }
}
}
impl TensorStorage for BoxedTensorStorageForResident {
fn dtype(&self) -> &Symbol {
&self.dtype
}
fn len(&self) -> usize {
self.cells.len()
}
fn location(&self) -> TensorLocation {
TensorLocation::Host
}
fn cell(&self, index: usize) -> Result<Value> {
self.cells
.get(index)
.cloned()
.ok_or_else(|| Error::Eval("tensor cell index was out of bounds".to_owned()))
}
fn materialize(&self) -> Result<Arc<dyn TensorStorage>> {
Ok(Arc::new(Self {
dtype: self.dtype.clone(),
cells: self.cells.clone(),
}))
}
fn as_any(&self) -> &dyn Any {
self
}
}
impl sim_kernel::Object for ResidentHandle {
fn display(&self, _cx: &mut sim_kernel::Cx) -> Result<String> {
Ok(format!("#<compute-resident {}>", self.symbol))
}
fn as_any(&self) -> &dyn Any {
self
}
}
impl sim_kernel::ObjectCompat for ResidentHandle {
fn class(&self, _cx: &mut sim_kernel::Cx) -> Result<sim_kernel::ClassRef> {
DefaultFactory.class_stub(
sim_kernel::CORE_FUNCTION_CLASS_ID,
Symbol::qualified("compute", "ResidentHandle"),
)
}
fn as_table(&self, cx: &mut sim_kernel::Cx) -> Result<Value> {
cx.factory().table(vec![(
Symbol::new("allocation"),
cx.factory().symbol(self.symbol.clone())?,
)])
}
}