use std::collections::HashMap;
use crate::topology::bounds::PayloadLike;
use crate::topology::point::PointId;
use crate::topology::sieve::FrozenSieveCsr;
use super::{AcceleratorBackend, AcceleratorError, DeviceValue};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct PlanEpochs {
pub topology: u64,
pub atlas: u64,
pub geometry: u64,
}
impl PlanEpochs {
pub fn validate(self, current: Self) -> Result<(), AcceleratorError> {
if self.topology != current.topology {
return Err(AcceleratorError::StaleTopologyPlan {
expected: self.topology,
found: current.topology,
});
}
if self.atlas != current.atlas {
return Err(AcceleratorError::StaleAtlasPlan {
expected: self.atlas,
found: current.atlas,
});
}
if self.geometry != current.geometry {
return Err(AcceleratorError::StaleGeometryPlan {
expected: self.geometry,
found: current.geometry,
});
}
Ok(())
}
}
pub struct DeviceTopology<B: AcceleratorBackend> {
pub topology_version: u64,
pub point_ids: B::Buffer<u64>,
pub cone_offsets: B::Buffer<u32>,
pub cone_points: B::Buffer<u32>,
pub support_offsets: B::Buffer<u32>,
pub support_points: B::Buffer<u32>,
pub point_count: usize,
pub incidence_count: usize,
}
pub struct DeviceMeshPlan<B: AcceleratorBackend> {
pub topology: DeviceTopology<B>,
pub index_of: HashMap<PointId, u32>,
}
impl<B: AcceleratorBackend> DeviceMeshPlan<B> {
pub fn compile<T: PayloadLike>(
backend: &B,
frozen: &FrozenSieveCsr<PointId, T>,
topology_version: u64,
) -> Result<Self, AcceleratorError> {
checked_u32(frozen.point_of.len(), "point count")?;
checked_u32(frozen.out_dsts.len(), "cone incidence count")?;
checked_u32(frozen.in_srcs.len(), "support incidence count")?;
let point_ids: Vec<u64> = frozen.point_of.iter().map(PointId::get).collect();
let upload = |result: Result<B::Buffer<u32>, B::Error>| {
result.map_err(|e| AcceleratorError::DeviceTransferFailed(e.to_string()))
};
let point_ids = backend
.upload(&point_ids)
.map_err(|e| AcceleratorError::DeviceTransferFailed(e.to_string()))?;
let cone_offsets = upload(backend.upload(&frozen.out_offsets))?;
let cone_points = upload(backend.upload(&frozen.out_dsts))?;
let support_offsets = upload(backend.upload(&frozen.in_offsets))?;
let support_points = upload(backend.upload(&frozen.in_srcs))?;
Ok(Self {
topology: DeviceTopology {
topology_version,
point_ids,
cone_offsets,
cone_points,
support_offsets,
support_points,
point_count: frozen.point_of.len(),
incidence_count: frozen.out_dsts.len(),
},
index_of: frozen.index_of.clone(),
})
}
pub fn validate_topology(&self, current: u64) -> Result<(), AcceleratorError> {
if self.topology.topology_version == current {
Ok(())
} else {
Err(AcceleratorError::StaleTopologyPlan {
expected: self.topology.topology_version,
found: current,
})
}
}
}
pub(crate) fn checked_u32(value: usize, what: &'static str) -> Result<u32, AcceleratorError> {
u32::try_from(value).map_err(|_| AcceleratorError::IndexOverflow { what, value })
}
pub(crate) fn upload<T: DeviceValue, B: AcceleratorBackend>(
backend: &B,
values: &[T],
) -> Result<B::Buffer<T>, AcceleratorError> {
backend
.upload(values)
.map_err(|e| AcceleratorError::DeviceTransferFailed(e.to_string()))
}