1use std::collections::HashMap;
2
3use crate::rules::GraphPrimitiveBuilder;
4use crate::{
5 ADKey, ADRuleError, ADRuleKind, ADRuleResult, Primitive, PrimitiveBuilder, PrimitiveValue,
6};
7use computegraph::graph::GraphBuilder;
8use computegraph::{LocalValueId, OperationRole, ValueKey, ValueRef};
9
10use crate::LinearizedGraph;
11
12pub fn linear_transpose<Op: Primitive>(
26 linear: &LinearizedGraph<Op>,
27 ctx: &mut Op::ADContext,
28) -> ADRuleResult<LinearizedGraph<Op>>
29where
30 Op::InputKey: ADKey,
31{
32 let mut builder = GraphBuilder::<Op>::new();
33 let mut cotangent_env: HashMap<ValueKey<Op>, LocalValueId> = HashMap::new();
34 let mut cotangent_seed_inputs = Vec::new();
35 let graph = linear.as_graph();
36
37 for (index, maybe_tangent_output) in linear.tangent_outputs().iter().enumerate() {
38 let tangent_output_id = match maybe_tangent_output {
39 Some(tangent_output_id) => tangent_output_id,
40 None => continue,
41 };
42
43 let source_key = graph.values()[*tangent_output_id].key.clone();
44 let seed_key = cotangent_seed_key(linear, index)?;
45 let seed_id = builder.add_input(seed_key.clone());
46 cotangent_env.insert(source_key, seed_id);
47 cotangent_seed_inputs.push((seed_key, seed_id));
48 }
49
50 for op_node in graph.operations().iter().rev() {
51 let cotangent_out: Vec<Option<LocalValueId>> = op_node
52 .outputs
53 .iter()
54 .map(|output_id| cotangent_env.get(&graph.values()[*output_id].key).copied())
55 .collect();
56 if cotangent_out.iter().all(Option::is_none) {
57 continue;
58 }
59
60 let rule_inputs: Vec<PrimitiveValue<Op>> = op_node
61 .inputs
62 .iter()
63 .map(|input| match input {
64 ValueRef::Local(local_id) => {
65 PrimitiveValue::External(graph.values()[*local_id].key.clone())
66 }
67 ValueRef::External(key) => PrimitiveValue::External(key.clone()),
68 })
69 .collect();
70
71 let mut primitive_builder = GraphPrimitiveBuilder::new(&mut builder);
72 let cotangent_in = op_node.operation.transpose_rule(
73 &mut primitive_builder,
74 &cotangent_out,
75 &rule_inputs,
76 &op_node.role,
77 ctx,
78 )?;
79 if cotangent_in.len() != rule_inputs.len() {
80 return Err(ADRuleError::invalid_input(
81 format!("{:?}", op_node.operation),
82 ADRuleKind::Transpose,
83 format!(
84 "rule returned {} cotangents for {} inputs",
85 cotangent_in.len(),
86 rule_inputs.len()
87 ),
88 ));
89 }
90
91 for (input, maybe_cotangent) in rule_inputs.iter().zip(cotangent_in) {
92 let cotangent_id = match maybe_cotangent {
93 Some(cotangent_id) => cotangent_id,
94 None => continue,
95 };
96 let input_key = match input {
97 PrimitiveValue::Local(_) => {
98 unreachable!("rule inputs are normalized to external refs")
99 }
100 PrimitiveValue::External(key) => key.clone(),
101 };
102
103 match cotangent_env.get(&input_key).copied() {
104 Some(existing_id) => {
105 let mut primitive_builder = GraphPrimitiveBuilder::new(&mut builder);
106 let sum = primitive_builder.add_primitive(
107 Op::add(),
108 vec![
109 PrimitiveValue::Local(existing_id),
110 PrimitiveValue::Local(cotangent_id),
111 ],
112 OperationRole::Linearized {
113 active_mask: vec![true, true],
114 },
115 );
116 cotangent_env.insert(input_key, sum[0]);
117 }
118 None => {
119 cotangent_env.insert(input_key, cotangent_id);
120 }
121 }
122 }
123 }
124
125 let tangent_outputs: Vec<Option<LocalValueId>> = linear
126 .tangent_inputs()
127 .iter()
128 .map(|(_, tangent_input_id)| {
129 let tangent_input_key = &graph.values()[*tangent_input_id].key;
130 cotangent_env.get(tangent_input_key).copied()
131 })
132 .collect();
133 let active_outputs: Vec<LocalValueId> = tangent_outputs.iter().filter_map(|id| *id).collect();
134 if !active_outputs.is_empty() {
135 builder.set_outputs(active_outputs);
136 }
137
138 Ok(LinearizedGraph::from_parts(
139 builder.build(),
140 cotangent_seed_inputs,
141 tangent_outputs,
142 ))
143}
144
145pub fn linear_transpose_with_builder<Op: Primitive>(
147 linear: &LinearizedGraph<Op>,
148 builder: &mut impl PrimitiveBuilder<Op>,
149 cotangent_seeds: &[Option<LocalValueId>],
150 ctx: &mut Op::ADContext,
151) -> ADRuleResult<Vec<Option<LocalValueId>>>
152where
153 Op::InputKey: ADKey,
154{
155 let mut cotangent_env: HashMap<ValueKey<Op>, LocalValueId> = HashMap::new();
156 let graph = linear.as_graph();
157
158 for (index, maybe_tangent_output) in linear.tangent_outputs().iter().enumerate() {
159 if let (Some(output_id), Some(Some(seed_id))) =
160 (maybe_tangent_output, cotangent_seeds.get(index))
161 {
162 let key = graph.values()[*output_id].key.clone();
163 cotangent_env.insert(key, *seed_id);
164 }
165 }
166
167 for op_node in graph.operations().iter().rev() {
168 let cotangent_out: Vec<Option<LocalValueId>> = op_node
169 .outputs
170 .iter()
171 .map(|output_id| cotangent_env.get(&graph.values()[*output_id].key).copied())
172 .collect();
173 if cotangent_out.iter().all(Option::is_none) {
174 continue;
175 }
176
177 let rule_inputs: Vec<PrimitiveValue<Op>> = op_node
178 .inputs
179 .iter()
180 .map(|input| match input {
181 ValueRef::Local(local_id) => {
182 PrimitiveValue::External(graph.values()[*local_id].key.clone())
183 }
184 ValueRef::External(key) => PrimitiveValue::External(key.clone()),
185 })
186 .collect();
187
188 let cotangent_in = op_node.operation.transpose_rule(
189 builder,
190 &cotangent_out,
191 &rule_inputs,
192 &op_node.role,
193 ctx,
194 )?;
195 if cotangent_in.len() != rule_inputs.len() {
196 return Err(ADRuleError::invalid_input(
197 format!("{:?}", op_node.operation),
198 ADRuleKind::Transpose,
199 format!(
200 "rule returned {} cotangents for {} inputs",
201 cotangent_in.len(),
202 rule_inputs.len()
203 ),
204 ));
205 }
206
207 for (input, maybe_cotangent) in rule_inputs.iter().zip(cotangent_in) {
208 let cotangent_id = match maybe_cotangent {
209 Some(cotangent_id) => cotangent_id,
210 None => continue,
211 };
212 let input_key = match input {
213 PrimitiveValue::Local(_) => {
214 unreachable!("rule inputs are normalized to external refs")
215 }
216 PrimitiveValue::External(key) => key.clone(),
217 };
218
219 match cotangent_env.get(&input_key).copied() {
220 Some(existing_id) => {
221 let sum = builder.add_primitive(
222 Op::add(),
223 vec![
224 PrimitiveValue::Local(existing_id),
225 PrimitiveValue::Local(cotangent_id),
226 ],
227 OperationRole::Linearized {
228 active_mask: vec![true, true],
229 },
230 );
231 cotangent_env.insert(input_key, sum[0]);
232 }
233 None => {
234 cotangent_env.insert(input_key, cotangent_id);
235 }
236 }
237 }
238 }
239
240 Ok(linear
241 .tangent_inputs()
242 .iter()
243 .map(|(_, tangent_input_id)| {
244 let tangent_input_key = &graph.values()[*tangent_input_id].key;
245 cotangent_env.get(tangent_input_key).copied()
246 })
247 .collect())
248}
249
250fn cotangent_seed_key<Op: Primitive>(
251 linear: &LinearizedGraph<Op>,
252 index: usize,
253) -> ADRuleResult<Op::InputKey>
254where
255 Op::InputKey: ADKey,
256{
257 if linear.tangent_inputs().is_empty() {
258 return Err(ADRuleError::invalid_input(
259 "tidu::linear_transpose",
260 ADRuleKind::Transpose,
261 "active tangent outputs require at least one tangent input to derive seed keys",
262 ));
263 }
264
265 let base_slot = index.min(linear.tangent_inputs().len() - 1);
266 let base_key = &linear.tangent_inputs()[base_slot].0;
267 Ok(base_key.tangent_of(u64::MAX - index as u64))
268}