use std::collections::HashMap;
use std::sync::Arc;
use computegraph::{GraphOperation, ValueKey};
use super::record::{EagerRecordError, EagerRecordResult, RecordedGraph};
pub struct Trace<Op: GraphOperation> {
node: Arc<TraceNode<Op>>,
}
impl<Op: GraphOperation> Clone for Trace<Op> {
fn clone(&self) -> Self {
Self {
node: self.node.clone(),
}
}
}
impl<Op: GraphOperation> Trace<Op> {
pub(crate) fn new(node: Arc<TraceNode<Op>>) -> Self {
Self { node }
}
pub(crate) fn node(&self) -> &Arc<TraceNode<Op>> {
&self.node
}
pub fn saved_values(&self) -> &HashMap<ValueKey<Op>, Arc<Op::Operand>> {
self.node.saved_data()
}
}
pub(crate) struct TraceNode<Op: GraphOperation> {
computation: RecordedGraph<Op>,
primal_out_keys: Vec<ValueKey<Op>>,
saved_data: HashMap<ValueKey<Op>, Arc<Op::Operand>>,
input_edges: Vec<TraceEdge<Op>>,
}
impl<Op: GraphOperation> TraceNode<Op> {
pub(crate) fn new(
computation: RecordedGraph<Op>,
primal_out_keys: Vec<ValueKey<Op>>,
saved_data: HashMap<ValueKey<Op>, Arc<Op::Operand>>,
input_edges: Vec<TraceEdge<Op>>,
) -> EagerRecordResult<Self> {
if primal_out_keys.len() != computation.output_keys().len() {
return Err(EagerRecordError::count_mismatch(
"TraceNode primal output keys",
computation.output_keys().len(),
primal_out_keys.len(),
));
}
if input_edges.len() != computation.input_keys().len() {
return Err(EagerRecordError::count_mismatch(
"TraceNode input edges",
computation.input_keys().len(),
input_edges.len(),
));
}
Ok(Self {
computation,
primal_out_keys,
saved_data,
input_edges,
})
}
pub(crate) fn computation(&self) -> &RecordedGraph<Op> {
&self.computation
}
pub(crate) fn primal_out_keys(&self) -> &[ValueKey<Op>] {
&self.primal_out_keys
}
pub(crate) fn saved_data(&self) -> &HashMap<ValueKey<Op>, Arc<Op::Operand>> {
&self.saved_data
}
pub(crate) fn input_edges(&self) -> &[TraceEdge<Op>] {
&self.input_edges
}
}
pub(crate) struct TraceEdge<Op: GraphOperation> {
pub(crate) node: Option<Arc<TraceNode<Op>>>,
pub(crate) key: ValueKey<Op>,
pub(crate) requires_grad: bool,
}
impl<Op: GraphOperation> TraceEdge<Op> {
pub(crate) fn new(
node: Option<Arc<TraceNode<Op>>>,
key: ValueKey<Op>,
requires_grad: bool,
) -> Self {
Self {
node,
key,
requires_grad,
}
}
}