Skip to main content

tidu/eager/
forward.rs

1use std::collections::{HashMap, HashSet};
2use std::sync::Arc;
3
4use crate::{ADKey, ADRuleResult, Primitive};
5use computegraph::{GraphOperation, ValueKey};
6
7use crate::LinearizedGraph;
8
9use super::record::RecordedGraph;
10use super::trace::{Trace, TraceNode};
11
12/// Tangent seed aligned with one linearized graph input.
13pub type ForwardTangentInput<Op> = (
14    <Op as GraphOperation>::InputKey,
15    Option<Arc<<Op as GraphOperation>::Operand>>,
16);
17
18/// Downstream execution hooks for eager forward-mode AD.
19pub trait ForwardExecutor<Op: Primitive>
20where
21    Op::InputKey: ADKey,
22{
23    /// Linearize one recorded eager graph node for forward-mode propagation.
24    ///
25    /// The default implementation calls [`RecordedGraph::linearize`]. Runtime
26    /// owners may override this to memoize per-node transform results without
27    /// taking trace traversal ownership away from tidu.
28    fn linearize_recorded_graph(
29        &mut self,
30        graph: &RecordedGraph<Op>,
31        output_slots: &[usize],
32        ctx: &mut Op::ADContext,
33    ) -> ADRuleResult<Arc<LinearizedGraph<Op>>> {
34        graph.linearize(output_slots, ctx).map(Arc::new)
35    }
36
37    /// Run a linearized graph with concrete tangent seeds.
38    fn run_linearized_forward(
39        &mut self,
40        linear: &LinearizedGraph<Op>,
41        tangent_in: &[ForwardTangentInput<Op>],
42        external_data: &HashMap<ValueKey<Op>, Arc<Op::Operand>>,
43        ctx: &mut Op::ADContext,
44    ) -> ADRuleResult<Vec<Option<Arc<Op::Operand>>>>;
45
46    /// Add two concrete operands for tangent accumulation.
47    fn add_operands(&mut self, a: &Arc<Op::Operand>, b: &Arc<Op::Operand>) -> Arc<Op::Operand>;
48}
49
50/// Execute forward-mode AD over an eager trace for one output value.
51pub fn try_forward<Op: Primitive>(
52    output_key: &ValueKey<Op>,
53    output_trace: Option<&Trace<Op>>,
54    tangent_seeds: &HashMap<ValueKey<Op>, Arc<Op::Operand>>,
55    executor: &mut impl ForwardExecutor<Op>,
56    ctx: &mut Op::ADContext,
57) -> ADRuleResult<Option<Arc<Op::Operand>>>
58where
59    Op::InputKey: ADKey,
60{
61    let sorted_nodes = topo_sort_trace(output_trace);
62    let mut tangents = tangent_seeds.clone();
63
64    for node in &sorted_nodes {
65        let mut graph_input_tangents: HashMap<Op::InputKey, Option<Arc<Op::Operand>>> =
66            HashMap::new();
67        let mut has_active_input = false;
68        for (input_key, edge) in node
69            .computation()
70            .input_keys()
71            .iter()
72            .zip(node.input_edges().iter())
73        {
74            let tangent = tangents.get(&edge.key).cloned();
75            if tangent.is_some() {
76                has_active_input = true;
77            }
78            graph_input_tangents.insert(input_key.clone(), tangent);
79        }
80        if !has_active_input {
81            continue;
82        }
83
84        let output_slots: Vec<usize> = (0..node.primal_out_keys().len()).collect();
85        let linear = executor.linearize_recorded_graph(node.computation(), &output_slots, ctx)?;
86        let tangent_in: Vec<_> = linear
87            .tangent_inputs()
88            .iter()
89            .map(|(key, _)| {
90                (
91                    key.clone(),
92                    graph_input_tangents.get(key).cloned().unwrap_or(None),
93                )
94            })
95            .collect();
96        let tangent_out = executor.run_linearized_forward(
97            linear.as_ref(),
98            &tangent_in,
99            node.saved_data(),
100            ctx,
101        )?;
102        assert_eq!(
103            tangent_out.len(),
104            output_slots.len(),
105            "eager forward executor returned {} tangents for {} output slots",
106            tangent_out.len(),
107            output_slots.len()
108        );
109
110        for (slot, maybe_tangent) in output_slots.into_iter().zip(tangent_out) {
111            let tangent = match maybe_tangent {
112                Some(tangent) => tangent,
113                None => continue,
114            };
115            let key = node.primal_out_keys()[slot].clone();
116            let accumulated = match tangents.remove(&key) {
117                Some(existing) => executor.add_operands(&existing, &tangent),
118                None => tangent,
119            };
120            tangents.insert(key, accumulated);
121        }
122    }
123
124    Ok(tangents.get(output_key).cloned())
125}
126
127fn topo_sort_trace<Op: GraphOperation>(
128    output_trace: Option<&Trace<Op>>,
129) -> Vec<Arc<TraceNode<Op>>> {
130    fn visit<Op: GraphOperation>(
131        node: &Arc<TraceNode<Op>>,
132        visited: &mut HashSet<*const TraceNode<Op>>,
133        order: &mut Vec<Arc<TraceNode<Op>>>,
134    ) {
135        let ptr = Arc::as_ptr(node);
136        if !visited.insert(ptr) {
137            return;
138        }
139
140        for edge in node.input_edges() {
141            if let Some(parent) = &edge.node {
142                visit(parent, visited, order);
143            }
144        }
145
146        order.push(node.clone());
147    }
148
149    let mut visited = HashSet::new();
150    let mut order = Vec::new();
151    if let Some(trace) = output_trace {
152        visit(trace.node(), &mut visited, &mut order);
153    }
154    order
155}