Skip to main content

tidu/
linear_transpose.rs

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
12/// Transpose a linearized graph, reversing linear flow.
13///
14/// Fan-out accumulation is emitted explicitly with [`crate::Primitive::add`];
15/// no duplication primitive is assumed by the graph transform.
16///
17/// # Examples
18///
19/// ```ignore
20/// let mut ctx = ();
21/// let transposed = tidu::linear_transpose(&linear, &mut ctx)?;
22/// assert_eq!(transposed.tangent_outputs().len(), linear.tangent_inputs().len());
23/// # Ok::<(), tidu::ADRuleError>(())
24/// ```
25pub 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
145/// Execute the transpose of a linearized graph using a caller-provided builder.
146pub 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}