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
12pub type ForwardTangentInput<Op> = (
14 <Op as GraphOperation>::InputKey,
15 Option<Arc<<Op as GraphOperation>::Operand>>,
16);
17
18pub trait ForwardExecutor<Op: Primitive>
20where
21 Op::InputKey: ADKey,
22{
23 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 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 fn add_operands(&mut self, a: &Arc<Op::Operand>, b: &Arc<Op::Operand>) -> Arc<Op::Operand>;
48}
49
50pub 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}