Skip to main content

forge/autograd/
mod.rs

1//! Tape-based reverse-mode autograd (roadmap v4, Stage 8).
2//!
3//! Define-by-run: a training forward records one `Node` per op on the
4//! [`Tape`]; [`Tape::backward`] sweeps the nodes in reverse, composing each
5//! backward pass out of the already-verified forward ops (matmul with
6//! trans_a/trans_b, add, sum_rows) plus the dedicated backward kernels
7//! (gelu_bwd, layernorm_bwd, softmax_bwd, scatter_add). One implementation
8//! serves both backends because every rule goes through `ops`.
9
10use crate::error::{ForgeError, Result};
11use crate::ops::{self, MatmulSpec};
12use crate::tensor::Tensor;
13
14/// A tensor tracked on the tape.
15#[derive(Clone)]
16pub struct TVar {
17    pub id: usize,
18    pub t: Tensor,
19}
20
21type BackwardFn = Box<dyn Fn(&Tensor) -> Result<Vec<(usize, Tensor)>>>;
22
23struct Node {
24    /// None for leaves (parameters, inputs).
25    backward: Option<BackwardFn>,
26}
27
28#[derive(Default)]
29pub struct Tape {
30    nodes: Vec<Node>,
31}
32
33impl Tape {
34    pub fn new() -> Tape {
35        Tape::default()
36    }
37
38    /// Register a leaf (parameter or input); its gradient is what
39    /// `backward` reports.
40    pub fn leaf(&mut self, t: Tensor) -> TVar {
41        let id = self.nodes.len();
42        self.nodes.push(Node { backward: None });
43        TVar { id, t }
44    }
45
46    fn record(&mut self, t: Tensor, f: BackwardFn) -> TVar {
47        let id = self.nodes.len();
48        self.nodes.push(Node { backward: Some(f) });
49        TVar { id, t }
50    }
51
52    pub fn add(&mut self, a: &TVar, b: &TVar) -> Result<TVar> {
53        let out = ops::add(&a.t, &b.t)?;
54        let (ia, ib) = (a.id, b.id);
55        Ok(self.record(
56            out,
57            Box::new(move |dy| Ok(vec![(ia, dy.clone()), (ib, dy.clone())])),
58        ))
59    }
60
61    /// Matmul with the specs the training forward uses: `trans_a == false`,
62    /// `b_rows == None`, operands of equal rank (2x2 or batched 3x3).
63    pub fn matmul(
64        &mut self,
65        a: &TVar,
66        b: &TVar,
67        bias: Option<&TVar>,
68        spec: MatmulSpec,
69    ) -> Result<TVar> {
70        if spec.trans_a || spec.b_rows.is_some() {
71            return Err(ForgeError::Shape(
72                "tape matmul: trans_a/b_rows are backward-only specs".into(),
73            ));
74        }
75        if a.t.shape().rank() != b.t.shape().rank() {
76            return Err(ForgeError::Shape(
77                "tape matmul requires equal-rank operands (no broadcast grads)".into(),
78            ));
79        }
80        let out = ops::matmul(&a.t, &b.t, bias.map(|b| &b.t), spec)?;
81        let (ia, ib) = (a.id, b.id);
82        let ibias = bias.map(|b| b.id);
83        let (at, bt) = (a.t.clone(), b.t.clone());
84        let alpha = spec.alpha;
85        let trans_b = spec.trans_b;
86        Ok(self.record(
87            out,
88            Box::new(move |dy| {
89                let da = if !trans_b {
90                    // Y = aAB: dA = a dY Bt
91                    ops::matmul(
92                        dy,
93                        &bt,
94                        None,
95                        MatmulSpec {
96                            trans_b: true,
97                            alpha,
98                            ..Default::default()
99                        },
100                    )?
101                } else {
102                    // Y = aABt: dA = a dY B
103                    ops::matmul(
104                        dy,
105                        &bt,
106                        None,
107                        MatmulSpec {
108                            alpha,
109                            ..Default::default()
110                        },
111                    )?
112                };
113                let db = if !trans_b {
114                    // dB = a At dY
115                    ops::matmul(
116                        &at,
117                        dy,
118                        None,
119                        MatmulSpec {
120                            trans_a: true,
121                            alpha,
122                            ..Default::default()
123                        },
124                    )?
125                } else {
126                    // dB = a dYt A
127                    ops::matmul(
128                        dy,
129                        &at,
130                        None,
131                        MatmulSpec {
132                            trans_a: true,
133                            alpha,
134                            ..Default::default()
135                        },
136                    )?
137                };
138                let mut grads = vec![(ia, da), (ib, db)];
139                if let Some(ibias) = ibias {
140                    grads.push((ibias, ops::sum_rows(dy)?));
141                }
142                Ok(grads)
143            }),
144        ))
145    }
146
147    pub fn gelu(&mut self, x: &TVar) -> Result<TVar> {
148        let out = ops::gelu(&x.t)?;
149        let ix = x.id;
150        let xt = x.t.clone();
151        Ok(self.record(
152            out,
153            Box::new(move |dy| Ok(vec![(ix, ops::gelu_bwd(&xt, dy)?)])),
154        ))
155    }
156
157    pub fn layernorm(&mut self, x: &TVar, gamma: &TVar, beta: &TVar, eps: f32) -> Result<TVar> {
158        let out = ops::layernorm(&x.t, &gamma.t, &beta.t, eps)?;
159        let (ix, ig, ib) = (x.id, gamma.id, beta.id);
160        let (xt, gt) = (x.t.clone(), gamma.t.clone());
161        Ok(self.record(
162            out,
163            Box::new(move |dy| {
164                let (dx, dgamma, dbeta) = ops::layernorm_bwd(&xt, &gt, dy, eps)?;
165                Ok(vec![(ix, dx), (ig, dgamma), (ib, dbeta)])
166            }),
167        ))
168    }
169
170    /// Softmax over the last dim (optionally causal). Saves the output for
171    /// backward; the causal mask needs no bookkeeping there because masked
172    /// outputs are exact zeros.
173    pub fn softmax(&mut self, x: &TVar, causal: bool, off: usize) -> Result<TVar> {
174        let out = ops::softmax(&x.t, causal, off)?;
175        let ix = x.id;
176        let yt = out.clone();
177        Ok(self.record(
178            out,
179            Box::new(move |dy| Ok(vec![(ix, ops::softmax_bwd(&yt, dy)?)])),
180        ))
181    }
182
183    pub fn split_heads(&mut self, qkv: &TVar, n_head: usize) -> Result<(TVar, TVar, TVar)> {
184        let (q, k, v) = ops::split_heads(&qkv.t, n_head)?;
185        let iq = qkv.id;
186        let mk = move |which: usize| -> BackwardFn {
187            Box::new(move |dy| Ok(vec![(iq, ops::unsplit_head(dy, which)?)]))
188        };
189        Ok((
190            self.record(q, mk(0)),
191            self.record(k, mk(1)),
192            self.record(v, mk(2)),
193        ))
194    }
195
196    pub fn merge_heads(&mut self, x: &TVar) -> Result<TVar> {
197        let h = x.t.shape().dim(0);
198        let out = ops::merge_heads(&x.t)?;
199        let ix = x.id;
200        Ok(self.record(
201            out,
202            Box::new(move |dy| Ok(vec![(ix, ops::unmerge_heads(dy, h)?)])),
203        ))
204    }
205
206    /// Fused token + positional embedding over a single-chunk wte.
207    /// dwte comes from scatter-add over the token ids; dwpe is the upstream
208    /// gradient placed at rows `pos..pos+t`.
209    pub fn embedding(&mut self, ids: &Tensor, wte: &TVar, wpe: &TVar, pos: usize) -> Result<TVar> {
210        let out = ops::embedding(ids, &wte.t, Some(&wpe.t), pos)?;
211        let (iw, ip) = (wte.id, wpe.id);
212        let ids = ids.clone();
213        let wte_shape = wte.t.shape().clone();
214        let (n_ctx, c) = (wpe.t.shape().dim(0), wpe.t.shape().dim(1));
215        Ok(self.record(
216            out,
217            Box::new(move |dy| {
218                let device = dy.device();
219                let mut dwte = Tensor::zeros(wte_shape.clone(), &device)?;
220                ops::scatter_add_rows(&mut dwte, &ids, dy)?;
221                // dwpe: place dy rows at pos.. via kv_append on a rank-3
222                // view created *before* the write (CPU copy-on-write).
223                let t = ids.shape().numel();
224                let mut dwpe3 = Tensor::zeros([1, n_ctx, c], &device)?;
225                ops::kv_append(&mut dwpe3, &dy.reshape([1, t, c])?, pos)?;
226                Ok(vec![(iw, dwte), (ip, dwpe3.reshape([n_ctx, c])?)])
227            }),
228        ))
229    }
230
231    /// Inverted dropout; the deterministic counter RNG regenerates the same
232    /// mask in backward, so nothing is saved.
233    pub fn dropout(&mut self, x: &TVar, p: f32, seed: u32) -> Result<TVar> {
234        if p == 0.0 {
235            return Ok(x.clone());
236        }
237        let out = ops::dropout(&x.t, p, seed)?;
238        let ix = x.id;
239        Ok(self.record(
240            out,
241            Box::new(move |dy| Ok(vec![(ix, ops::dropout(dy, p, seed)?)])),
242        ))
243    }
244
245    /// Reverse sweep from `root` seeded with `seed_grad` (dL/droot).
246    /// Returns per-node gradients; leaves keep theirs, interior gradients
247    /// are freed as soon as they have been consumed.
248    pub fn backward(&mut self, root: &TVar, seed_grad: Tensor) -> Result<Vec<Option<Tensor>>> {
249        let mut grads: Vec<Option<Tensor>> = (0..self.nodes.len()).map(|_| None).collect();
250        grads[root.id] = Some(seed_grad);
251        for id in (0..=root.id).rev() {
252            if grads[id].is_none() {
253                continue;
254            }
255            let Some(f) = self.nodes[id].backward.take() else {
256                continue; // leaf: keep the gradient
257            };
258            let dy = grads[id].take().expect("checked above");
259            for (pid, g) in f(&dy)? {
260                debug_assert!(pid < id, "backward edge must point to an earlier node");
261                grads[pid] = Some(match grads[pid].take() {
262                    Some(acc) => ops::add(&acc, &g)?,
263                    None => g,
264                });
265            }
266        }
267        Ok(grads)
268    }
269}