Skip to main content

tidu/
linearize.rs

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
12/// Linearize a resolved computation graph, producing a linear graph.
13///
14/// The transform walks the reachable DAG from `outputs` in dependency-first
15/// order and delegates primitive-specific JVP generation to
16/// [`crate::Primitive::jvp_rule`].
17///
18/// # Examples
19///
20/// ```ignore
21/// use computegraph::resolve::resolve;
22/// use tidu::linearize;
23///
24/// let view = resolve(vec![primal_graph]);
25/// let mut ctx = ();
26/// let aliases = std::collections::HashMap::new();
27/// let linear = linearize(&view, &[output_key], &[input_key], 1, &mut ctx, &aliases)?;
28/// assert_eq!(linear.tangent_outputs().len(), 1);
29/// # Ok::<(), tidu::ADRuleError>(())
30/// ```
31pub 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}