1use std::collections::{HashMap, HashSet};
2use std::sync::Arc;
3
4use crate::rules::GraphPrimitiveBuilder;
5use crate::{ADKey, ADRuleError, ADRuleKind, ADRuleResult, DiffPassId, Primitive};
6use computegraph::graph::GraphBuilder;
7use computegraph::resolve::{ResolvedView, ValueDef};
8use computegraph::{GraphOperation, LocalValueId, OperationKey, ValueKey};
9
10use crate::LinearizedGraph;
11
12pub fn linearize<Op: Primitive>(
32 view: &ResolvedView<Op>,
33 outputs: &[ValueKey<Op>],
34 wrt: &[Op::InputKey],
35 pass: DiffPassId,
36 ctx: &mut Op::ADContext,
37 aliases: &HashMap<Op::InputKey, ValueKey<Op>>,
38) -> ADRuleResult<LinearizedGraph<Op>>
39where
40 Op::InputKey: ADKey,
41{
42 let mut builder = GraphBuilder::<Op>::new();
43 let topo_keys = topological_order(view, outputs, aliases);
44 let mut tangent_env: HashMap<ValueKey<Op>, Option<LocalValueId>> = HashMap::new();
45 let mut processed_ops = HashSet::new();
46
47 let mut tangent_inputs = Vec::with_capacity(wrt.len());
48 for wrt_key in wrt {
49 let tangent_key = wrt_key.tangent_of(pass);
50 let tangent_id = builder.add_input(tangent_key);
51 tangent_env.insert(ValueKey::Input(wrt_key.clone()), Some(tangent_id));
52 tangent_inputs.push((wrt_key.clone(), tangent_id));
53 }
54
55 for key in topo_keys {
56 if tangent_env.contains_key(&key) {
57 continue;
58 }
59
60 let val_def = match view.resolve_value(&key) {
61 Some(val_def) => val_def,
62 None => continue,
63 };
64
65 match val_def {
66 ValueDef::Input { key: ref input_key } => {
67 if let Some(aliased_key) = aliases.get(input_key) {
68 let aliased_tangent = tangent_env.get(aliased_key).copied().flatten();
69 tangent_env.insert(key, aliased_tangent);
70 } else {
71 tangent_env.insert(key, None);
72 }
73 }
74 ValueDef::Produced {
75 operation,
76 input_keys,
77 role,
78 ..
79 } => {
80 let global_op_key =
81 OperationKey::new(operation.clone(), input_keys.clone(), role.clone());
82 if !processed_ops.insert(global_op_key.clone()) {
83 continue;
84 }
85
86 let tangent_in: Vec<Option<LocalValueId>> = input_keys
87 .iter()
88 .map(|input_key| tangent_env.get(input_key).copied().flatten())
89 .collect();
90 let output_keys = output_keys(&global_op_key, operation.output_count());
91
92 if tangent_in.iter().all(Option::is_none) {
93 for output_key in output_keys {
94 tangent_env.insert(output_key, None);
95 }
96 continue;
97 }
98
99 let mut primitive_builder = GraphPrimitiveBuilder::new(&mut builder);
100 let tangent_out = operation.jvp_rule(
101 &mut primitive_builder,
102 &input_keys,
103 &output_keys,
104 &tangent_in,
105 ctx,
106 )?;
107 if tangent_out.len() != output_keys.len() {
108 return Err(ADRuleError::invalid_input(
109 format!("{:?}", operation),
110 ADRuleKind::Jvp,
111 format!(
112 "rule returned {} tangents for {} outputs",
113 tangent_out.len(),
114 output_keys.len()
115 ),
116 ));
117 }
118
119 for (output_key, tangent_output) in output_keys.into_iter().zip(tangent_out) {
120 tangent_env.insert(output_key, tangent_output);
121 }
122 }
123 }
124 }
125
126 let tangent_outputs: Vec<Option<LocalValueId>> = outputs
127 .iter()
128 .map(|key| tangent_env.get(key).copied().flatten())
129 .collect();
130 let active_outputs: Vec<LocalValueId> = tangent_outputs.iter().filter_map(|id| *id).collect();
131 if !active_outputs.is_empty() {
132 builder.set_outputs(active_outputs);
133 }
134
135 Ok(LinearizedGraph::from_parts(
136 builder.build(),
137 tangent_inputs,
138 tangent_outputs,
139 ))
140}
141
142fn output_keys<Op: GraphOperation>(
143 op_key: &OperationKey<Op>,
144 output_count: usize,
145) -> Vec<ValueKey<Op>> {
146 let op_key = Arc::new(op_key.clone());
147 (0..output_count)
148 .map(|output_slot| ValueKey::Derived {
149 operation: Arc::clone(&op_key),
150 output_slot: output_slot as u8,
151 })
152 .collect()
153}
154
155fn topological_order<Op: GraphOperation>(
156 view: &ResolvedView<Op>,
157 outputs: &[ValueKey<Op>],
158 aliases: &HashMap<Op::InputKey, ValueKey<Op>>,
159) -> Vec<ValueKey<Op>> {
160 fn visit<Op: GraphOperation>(
161 key: &ValueKey<Op>,
162 view: &ResolvedView<Op>,
163 aliases: &HashMap<Op::InputKey, ValueKey<Op>>,
164 visited: &mut HashSet<ValueKey<Op>>,
165 order: &mut Vec<ValueKey<Op>>,
166 ) {
167 if !visited.insert(key.clone()) {
168 return;
169 }
170
171 match view.resolve_value(key) {
172 Some(ValueDef::Produced { input_keys, .. }) => {
173 for input_key in input_keys {
174 visit(&input_key, view, aliases, visited, order);
175 }
176 }
177 Some(ValueDef::Input { key: input_key }) => {
178 if let Some(aliased_key) = aliases.get(&input_key) {
179 visit(aliased_key, view, aliases, visited, order);
180 }
181 }
182 None => {}
183 }
184
185 order.push(key.clone());
186 }
187
188 let mut visited = HashSet::new();
189 let mut order = Vec::new();
190 for output_key in outputs {
191 visit(output_key, view, aliases, &mut visited, &mut order);
192 }
193 order
194}