use std::collections::HashMap;
use std::fmt::Debug;
use std::sync::Arc;
use crate::tensor::backend::Backend;
use crate::tensor::graph::{
NodeKind, TensorGraphBaked, TensorGraphCacheNode, TensorGraphEdge, TensorGraphNode,
};
use crate::tensor::planner::alias::{self, AliasKind, AliasMap};
use crate::tensor::planner::runtime::{ExecKind, Slot};
use crate::tensor::planner::sort::topological_sort;
use crate::tensor::planner::{get_id, runtime};
#[derive(Debug, Clone)]
pub(crate) enum OutputKind {
Buffer(usize),
InPlaceIdx(usize),
Reference(usize),
Allocate(usize),
}
pub(crate) enum ComputeKind<'a, T, B: Backend> {
Leaf {
edge: &'a Arc<TensorGraphEdge<T, B>>,
},
Op {
node: &'a TensorGraphNode<T, B>,
output: OutputKind,
resolved_inputs: Vec<usize>,
dealloc_after: Vec<usize>,
},
CachedOp {
cache: &'a Arc<TensorGraphCacheNode<T, B>>,
output: OutputKind,
resolved_inputs: Vec<usize>,
dealloc_after: Vec<usize>,
},
Baked {
baked: &'a Arc<TensorGraphBaked<T, B>>,
resolved_inputs: Vec<usize>,
dealloc_after: Vec<usize>,
},
}
#[derive(Debug)]
pub(crate) struct OpPlan<'a, T, B: Backend> {
pub(crate) node: &'a NodeKind<T, B>,
pub(crate) resolved_inputs: Vec<&'a NodeKind<T, B>>,
pub(crate) end: Option<usize>,
}
#[inline]
fn extend_slot_life(slot_end1: Option<usize>, slot_end2: Option<usize>) -> Option<usize> {
slot_end1.and_then(|e1| slot_end2.map(|e2| e1.max(e2)))
}
#[inline]
fn build_resolved_inputs<T, B: Backend>(resolved_inputs: &[&NodeKind<T, B>]) -> Vec<usize> {
resolved_inputs.iter().map(|inp| get_id(*inp)).collect()
}
#[inline]
fn resolve_inputs<'a, T, B: Backend>(
inputs: &'a [NodeKind<T, B>],
alias_map: &AliasMap<'a, T, B>,
) -> Vec<&'a NodeKind<T, B>> {
inputs.iter().map(|i| alias_map.resolve(i)).collect()
}
#[inline]
fn track_lifetimes<T, B: Backend>(
resolved: &[&NodeKind<T, B>],
pos: usize,
id_op: &HashMap<usize, usize>,
ops: &mut [OpPlan<'_, T, B>],
) {
for inp in resolved {
if let Some(&op_idx) = id_op.get(&get_id(inp)) {
ops[op_idx].end = Some(pos);
}
}
}
struct PlanState<'a, T, B: Backend> {
plan: Vec<ComputeKind<'a, T, B>>,
slots: Vec<Slot>,
id_slot_map: HashMap<usize, usize>,
ref_deallocs: Vec<(usize, Option<usize>)>,
}
impl<'a, T, B: Backend> PlanState<'a, T, B> {
fn new() -> Self {
Self {
plan: Vec::with_capacity(32),
slots: Vec::with_capacity(32),
id_slot_map: HashMap::with_capacity(32),
ref_deallocs: Vec::with_capacity(8),
}
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(
level = "trace",
skip(self, node, resolved_inputs),
fields(
node_id = node.id,
output_len = node.layout.len(),
slots_available = self.slots.len()
)
)
)]
#[inline]
fn plan_node(
&mut self,
op_start: usize,
op_end: Option<usize>,
node: &'a TensorGraphNode<T, B>,
resolved_inputs: &[&NodeKind<T, B>],
) {
match runtime::classify(
&node.op,
resolved_inputs,
&node.layout,
op_start,
&self.slots,
&self.id_slot_map,
) {
ExecKind::Allocate => {
self.id_slot_map.insert(node.id, self.slots.len());
self.slots.push(Slot {
id: node.id,
len: node.layout.len(),
end: op_end,
});
let resolved_inputs = build_resolved_inputs(resolved_inputs);
self.plan.push(ComputeKind::Op {
node,
output: OutputKind::Allocate(node.layout.len()),
resolved_inputs,
dealloc_after: Vec::new(),
});
}
ExecKind::UseSlot { slot_idx } => {
self.id_slot_map.insert(node.id, slot_idx);
self.slots[slot_idx].end = op_end;
let resolved_inputs = build_resolved_inputs(resolved_inputs);
self.plan.push(ComputeKind::Op {
node,
output: OutputKind::Buffer(self.slots[slot_idx].id),
resolved_inputs,
dealloc_after: Vec::new(),
});
self.slots[slot_idx].id = node.id;
}
ExecKind::InPlace {
slot_idx,
input_idx,
} => {
self.id_slot_map.insert(node.id, slot_idx);
self.slots[slot_idx].end = op_end;
let resolved_inputs = build_resolved_inputs(resolved_inputs);
self.plan.push(ComputeKind::Op {
node,
output: OutputKind::InPlaceIdx(input_idx),
resolved_inputs,
dealloc_after: Vec::new(),
});
self.slots[slot_idx].id = node.id;
}
ExecKind::ReferenceEternal { input_idx } => {
let resolved_inputs = build_resolved_inputs(resolved_inputs);
self.plan.push(ComputeKind::Op {
node,
output: OutputKind::Reference(input_idx),
resolved_inputs,
dealloc_after: Vec::new(),
});
}
ExecKind::ReferenceSlot {
slot_idx,
input_idx,
} => {
let extended_end = extend_slot_life(self.slots[slot_idx].end, op_end);
self.slots[slot_idx].end = extended_end;
self.id_slot_map.insert(node.id, slot_idx);
self.ref_deallocs.push((node.id, extended_end));
let resolved_inputs = build_resolved_inputs(resolved_inputs);
self.plan.push(ComputeKind::Op {
node,
output: OutputKind::Reference(input_idx),
resolved_inputs,
dealloc_after: Vec::new(),
});
}
}
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(
level = "trace",
skip(self, cache, resolved_inputs),
fields(
node_id = cache.get_node().id,
output_len = cache.get_node().layout.len(),
cache_filled = cache.is_cache_filled(),
slots_available = self.slots.len()
)
)
)]
fn plan_cache_node(
&mut self,
op_start: usize,
cache: &'a Arc<TensorGraphCacheNode<T, B>>,
resolved_inputs: &[&NodeKind<T, B>],
) {
let node = cache.get_node();
if cache.is_cache_filled() {
let resolved_inputs = build_resolved_inputs(resolved_inputs);
self.plan.push(ComputeKind::CachedOp {
cache,
output: OutputKind::Allocate(0),
resolved_inputs,
dealloc_after: Vec::new(),
});
return;
}
match runtime::classify(
&node.op,
resolved_inputs,
&node.layout,
op_start,
&self.slots,
&self.id_slot_map,
) {
ExecKind::Allocate => {
let resolved_inputs = build_resolved_inputs(resolved_inputs);
self.plan.push(ComputeKind::CachedOp {
cache,
output: OutputKind::Allocate(node.layout.len()),
resolved_inputs,
dealloc_after: Vec::new(),
});
}
ExecKind::UseSlot { slot_idx } => {
self.slots[slot_idx].end = None;
let resolved_inputs = build_resolved_inputs(resolved_inputs);
self.plan.push(ComputeKind::CachedOp {
cache,
output: OutputKind::Buffer(self.slots[slot_idx].id),
resolved_inputs,
dealloc_after: Vec::new(),
});
}
ExecKind::InPlace {
slot_idx,
input_idx,
} => {
self.slots[slot_idx].end = None;
let resolved_inputs = build_resolved_inputs(resolved_inputs);
self.plan.push(ComputeKind::CachedOp {
cache,
output: OutputKind::InPlaceIdx(input_idx),
resolved_inputs,
dealloc_after: Vec::new(),
});
}
ExecKind::ReferenceEternal { input_idx } => {
let resolved_inputs = build_resolved_inputs(resolved_inputs);
self.plan.push(ComputeKind::CachedOp {
cache,
output: OutputKind::Reference(input_idx),
resolved_inputs,
dealloc_after: Vec::new(),
});
}
ExecKind::ReferenceSlot {
slot_idx,
input_idx,
} => {
self.slots[slot_idx].end = None;
self.ref_deallocs.push((node.id, None));
let resolved_inputs = build_resolved_inputs(resolved_inputs);
self.plan.push(ComputeKind::CachedOp {
cache,
output: OutputKind::Reference(input_idx),
resolved_inputs,
dealloc_after: Vec::new(),
});
}
}
}
}
struct RootNode<'a, T, B: Backend> {
id: usize,
resolved_inputs: Vec<&'a NodeKind<T, B>>,
}
struct PrePlan<'a, T, B: Backend> {
pre_plan: Vec<OpPlan<'a, T, B>>,
root: RootNode<'a, T, B>,
external_inputs: Vec<usize>,
}
fn pre_plan<'a, T: PartialEq + Clone, B: Backend>(
base_node: &'a TensorGraphNode<T, B>,
) -> PrePlan<'a, T, B> {
let dag_iter = topological_sort(base_node);
let mut id_op: HashMap<usize, usize> = HashMap::with_capacity(32);
let mut ops: Vec<OpPlan<'_, T, B>> = Vec::with_capacity(32);
let mut alias_map: AliasMap<'_, T, B> = AliasMap::new();
let mut external_inputs: Vec<usize> = Vec::with_capacity(8);
for node in dag_iter {
match node {
NodeKind::Edge(e) => {
id_op.insert(e.id, ops.len());
ops.push(OpPlan {
node,
resolved_inputs: Vec::new(),
end: None,
});
}
NodeKind::Slot(s) => {
external_inputs.push(s.id);
}
NodeKind::Node(n) => match alias::classify(&n.op, &n.inputs, &alias_map) {
AliasKind::NoAlias => {
let resolved_inputs = resolve_inputs(&n.inputs, &alias_map);
let pos = ops.len();
id_op.insert(n.id, pos);
track_lifetimes(&resolved_inputs, pos, &id_op, &mut ops);
ops.push(OpPlan {
node,
resolved_inputs,
end: None,
});
}
AliasKind::Takeover(parent, tag) => {
let resolved_inputs = resolve_inputs(&n.inputs, &alias_map);
let pos = ops.len();
id_op.insert(n.id, ops.len());
track_lifetimes(&resolved_inputs, pos, &id_op, &mut ops);
ops.push(OpPlan {
node,
resolved_inputs,
end: None,
});
alias_map.takeover(parent, node, tag);
}
AliasKind::Alias(target, tag) => {
alias_map.insert(n.id, target, tag);
}
},
NodeKind::Cache(cache) => {
let n = cache.get_node();
match alias::classify_cache(&n.inputs, &alias_map) {
AliasKind::Alias(target, tag) => {
alias_map.insert(n.id, target, tag);
}
AliasKind::Takeover(old_owner, tag) => {
let resolved_inputs = resolve_inputs(&n.inputs, &alias_map);
let pos = ops.len();
id_op.insert(n.id, ops.len());
track_lifetimes(&resolved_inputs, pos, &id_op, &mut ops);
ops.push(OpPlan {
node,
resolved_inputs,
end: None,
});
alias_map.takeover(old_owner, node, tag);
}
_ => unreachable!("classify_cache always aliases or takes over"),
}
}
NodeKind::Baked(baked) => {
let resolved_inputs = resolve_inputs(&baked.inputs, &alias_map);
let pos = ops.len();
id_op.insert(baked.id, pos);
track_lifetimes(&resolved_inputs, pos, &id_op, &mut ops);
ops.push(OpPlan {
node,
resolved_inputs,
end: None,
});
}
}
}
let root_resolved = resolve_inputs(&base_node.inputs, &alias_map);
let root_id = match alias::classify(&base_node.op, &base_node.inputs, &alias_map) {
AliasKind::Alias(target, _) => {
let id = get_id(target);
if let Some(&op_idx) = id_op.get(&id) {
ops[op_idx].end = None;
}
id
}
_ => {
let root_pos = ops.len();
track_lifetimes(&root_resolved, root_pos, &id_op, &mut ops);
base_node.id
}
};
PrePlan {
pre_plan: ops,
root: RootNode {
id: root_id,
resolved_inputs: root_resolved,
},
external_inputs,
}
}
pub(crate) struct Plan<'a, T, B: Backend> {
pub(crate) plan: Vec<ComputeKind<'a, T, B>>,
pub(crate) root_id: usize,
pub(crate) external_inputs: Vec<usize>,
}
pub(crate) struct CorePlan<'a, T, B: Backend> {
pub(crate) plan: Vec<ComputeKind<'a, T, B>>,
pub(crate) root_id: usize,
pub(crate) external_inputs: Vec<usize>,
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(
level = "debug",
skip(base_node),
fields(
node_id = base_node.id,
ops_count = tracing::field::Empty,
slots_count = tracing::field::Empty,
dealloc_edges = tracing::field::Empty,
ref_deallocs_count = tracing::field::Empty
)
)
)]
#[inline]
pub(crate) fn core_plan_computation<T: PartialEq + Clone, B: Backend>(
base_node: &TensorGraphNode<T, B>,
) -> CorePlan<'_, T, B> {
let mut state: PlanState<'_, T, B> = PlanState::new();
let PrePlan {
pre_plan,
root,
external_inputs,
} = pre_plan(base_node);
let ops_len = pre_plan.len();
for (i, op) in pre_plan.into_iter().enumerate() {
match op.node {
NodeKind::Edge(e) => {
state.plan.push(ComputeKind::Leaf { edge: e });
}
NodeKind::Node(node) => {
state.plan_node(i, op.end, node, &op.resolved_inputs);
}
NodeKind::Cache(cache) => {
state.plan_cache_node(i, cache, &op.resolved_inputs);
}
NodeKind::Baked(baked) => {
state.plan.push(ComputeKind::Baked {
baked,
resolved_inputs: op.resolved_inputs.iter().map(|n| get_id(*n)).collect(),
dealloc_after: Vec::new(),
});
}
NodeKind::Slot(_) => unreachable!("slots are pre-plan only nodes"),
}
}
if root.id == base_node.id {
state.plan_node(ops_len, None, base_node, &root.resolved_inputs);
}
let PlanState {
mut plan,
slots,
ref_deallocs,
..
} = state;
for (node_id, dealloc_at) in &ref_deallocs {
let Some(end) = dealloc_at else { continue };
match &mut plan[*end] {
ComputeKind::Op { dealloc_after, .. }
| ComputeKind::CachedOp { dealloc_after, .. }
| ComputeKind::Baked { dealloc_after, .. } => dealloc_after.push(*node_id),
ComputeKind::Leaf { .. } => unreachable!(),
}
}
for slot in slots.into_iter() {
let Some(end) = slot.end else { continue };
match &mut plan[end] {
ComputeKind::Op { dealloc_after, .. }
| ComputeKind::CachedOp { dealloc_after, .. }
| ComputeKind::Baked { dealloc_after, .. } => dealloc_after.push(slot.id),
ComputeKind::Leaf { .. } => unreachable!(),
}
}
CorePlan {
plan,
root_id: root.id,
external_inputs,
}
}
pub(crate) fn plan_computation<T: PartialEq + Clone, B: Backend>(
base_node: &TensorGraphNode<T, B>,
) -> Plan<'_, T, B> {
let CorePlan {
plan,
root_id,
external_inputs,
..
} = core_plan_computation(base_node);
Plan {
plan,
root_id,
external_inputs,
}
}