use std::collections::BTreeSet;
use crate::{
DType, Map, RT, Tensor, ZyxError,
backend::BufferId,
graph::plan::drain_events_for_buf,
graph::{ClassId, Graph, GraphId},
kernel::Op,
shape::Dim,
slab::SlabId,
tensor::TensorId,
view::View,
};
#[cfg_attr(feature = "py", pyo3::pyclass)]
pub struct Tape {
graph_id: GraphId,
}
impl Tape {
pub fn new<'a>(params: impl IntoIterator<Item = &'a Tensor>) -> Result<Tape, ZyxError> {
let mut rt = RT.lock();
let graph_id = rt.graphs.push(Graph::new());
for p in params {
rt.promote_to_graph(p.id, graph_id)?;
}
Ok(Tape { graph_id })
}
pub fn empty() -> Tape {
Self::new(std::iter::empty()).unwrap()
}
pub fn add(&self, tensor: &Tensor) -> Result<(), ZyxError> {
let mut rt = RT.lock();
rt.promote_to_graph(tensor.id, self.graph_id)?;
Ok(())
}
pub fn extend<'a>(&self, params: impl IntoIterator<Item = &'a Tensor>) -> Result<(), ZyxError> {
let mut rt = RT.lock();
for p in params {
rt.promote_to_graph(p.id, self.graph_id)?;
}
Ok(())
}
}
impl Tape {
#[must_use]
pub fn gradient<'a>(&self, target: &Tensor, sources: impl IntoIterator<Item = &'a Tensor>) -> Vec<Tensor> {
let sources: Vec<TensorId> = sources.into_iter().map(Tensor::id).collect();
let mut rt = RT.lock();
let grads: Map<TensorId, TensorId> = rt.gradient(target.id(), sources.iter().copied().collect(), self.graph_id);
sources
.into_iter()
.map(|x: TensorId| {
let id = match grads.get(&x) {
Some(&id) => id,
None => {
let shape = rt.shape(x).into();
let dtype = rt.dtype(x);
rt.new_full(shape, dtype.zero_constant())
}
};
Tensor { id }
})
.collect()
}
pub fn realize<'a>(self, tensors: impl IntoIterator<Item = &'a Tensor>) -> Result<(), ZyxError> {
let mut rt = RT.lock();
let graph_id = self.graph_id;
let output_pairs: Vec<(TensorId, ClassId)> = tensors
.into_iter()
.map(|t| {
if rt.tensors[t.id].class_id.is_null() {
panic!("non-graph tensor in realize")
} else {
(t.id, rt.tensors[t.id].class_id)
}
})
.collect();
let output_tids: Vec<TensorId> = output_pairs.iter().map(|(tid, _)| *tid).collect();
let output_classes: Vec<ClassId> = output_pairs.iter().map(|(_, cid)| *cid).collect();
debug_assert!(rt.graphs.contains_key(graph_id));
rt.debug_assert_pre_realize(graph_id);
let output_set: BTreeSet<ClassId> = output_classes.iter().copied().collect();
let cache_key = rt.graphs[graph_id].cache_key(&output_set);
if let Some(plan) = rt.plan_cache.get(&cache_key) {
let mut class_buf: Map<ClassId, BufferId> = Map::default();
for &cid in &plan.leaf_classes {
let &tid = rt.graphs[graph_id].leaf_map.get(&cid).unwrap();
debug_assert!(rt.buffer_map.contains_key(&tid), "leaf class {cid:?} tid {tid:?} not in buffer_map");
class_buf.insert(cid, rt.buffer_map[&tid]);
}
rt.execute_plan(cache_key, &mut class_buf)?;
for (&tid, &cid) in output_tids.iter().zip(output_classes.iter()) {
rt.buffer_map.insert(tid, class_buf[&cid]);
rt.eagerify(tid);
}
rt.debug_assert_no_stray_buffers(graph_id, &output_tids);
return Ok(());
}
let plan = rt.compile_graph(graph_id, &output_set)?;
let mut class_buf: Map<ClassId, BufferId> = Map::default();
for &cid in &plan.leaf_classes {
let &tid = rt.graphs[graph_id].leaf_map.get(&cid).unwrap();
debug_assert!(rt.buffer_map.contains_key(&tid), "leaf class {cid:?} tid {tid:?} not in buffer_map");
class_buf.insert(cid, rt.buffer_map[&tid]);
}
rt.plan_cache.insert(cache_key, plan);
rt.execute_plan(cache_key, &mut class_buf)?;
for (&tid, &cid) in output_tids.iter().zip(output_classes.iter()) {
rt.buffer_map.insert(tid, class_buf[&cid]);
rt.eagerify(tid);
}
rt.debug_assert_no_stray_buffers(graph_id, &output_tids);
Ok(())
}
}
impl Drop for Tape {
fn drop(&mut self) {
let mut rt = RT.lock();
let graph_id = self.graph_id;
rt.graphs[graph_id].dead = true;
let leaves: Vec<TensorId> = rt.graphs[graph_id].leaf_map.values().copied().collect();
for tid in leaves {
if rt.tensors[tid].graph_id != graph_id {
continue;
}
if rt.tensors[tid].rc == 0 {
if let Some(buf_id) = rt.buffer_map.remove(&tid) {
let still_used = rt.buffer_map.values().any(|b| b.pool == buf_id.pool && b.buffer == buf_id.buffer);
if !still_used {
let wait_list = drain_events_for_buf(&mut rt.events, buf_id);
rt.pools[buf_id.pool].deallocate(buf_id.buffer, wait_list);
}
}
rt.tensors.remove(tid);
} else {
rt.tensors[tid].class_id = ClassId::NULL;
rt.tensors[tid].graph_id = GraphId::NULL;
rt.graphs[graph_id].ref_count -= 1;
}
}
if rt.graphs[graph_id].ref_count == 0 {
rt.remove_dead_graph(graph_id);
}
}
}
impl Tape {
pub fn freeze<'a>(self, outputs: impl IntoIterator<Item = &'a Tensor>) -> Result<FrozenTape, ZyxError> {
let mut rt = RT.lock();
let graph_id = self.graph_id;
let outputs: Vec<(ClassId, Vec<Dim>, DType)> = outputs
.into_iter()
.map(|t| {
if rt.tensors[t.id].class_id.is_null() {
panic!("non-graph tensor in realize")
} else {
(rt.tensors[t.id].class_id, rt.shape(t.id).into(), rt.dtype(t.id))
}
})
.collect();
debug_assert!(rt.graphs.contains_key(graph_id));
rt.debug_assert_pre_realize(graph_id);
let output_set: BTreeSet<ClassId> = outputs.iter().map(|x| x.0).collect();
let cache_key = rt.graphs[graph_id].cache_key(&output_set);
if rt.plan_cache.contains_key(&cache_key) {
return Ok(FrozenTape { cache_key, outputs });
}
let plan = rt.compile_graph(graph_id, &output_set)?;
rt.plan_cache.insert(cache_key, plan);
return Ok(FrozenTape { cache_key, outputs });
}
}
pub struct FrozenTape {
cache_key: u64,
outputs: Vec<(ClassId, Vec<Dim>, DType)>,
}
impl FrozenTape {
pub fn replay<'a>(&self, inputs: impl IntoIterator<Item = &'a Tensor>) -> Result<Vec<Tensor>, ZyxError> {
let mut rt = RT.lock();
let mut class_buf: Map<ClassId, BufferId> = Map::default();
for (tid, &cid) in inputs.into_iter().zip(rt.plan_cache[&self.cache_key].leaf_classes.iter()) {
class_buf.insert(cid, rt.buffer_map[&tid.id]);
}
rt.execute_plan(self.cache_key, &mut class_buf)?;
let mut outputs = Vec::new();
for (cid, shape, dtype) in self.outputs.iter() {
let view = View::contiguous(shape);
let tid = rt.new_eager_tensor(Op::LoadView(Box::new((*dtype, view))));
rt.buffer_map.insert(tid, class_buf[cid]);
outputs.push(Tensor::from_id(tid));
}
Ok(outputs)
}
}