1use crate::error::{ForgeError, Result};
11use crate::ops::{self, MatmulSpec};
12use crate::tensor::Tensor;
13
14#[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 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 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 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 ops::matmul(
92 dy,
93 &bt,
94 None,
95 MatmulSpec {
96 trans_b: true,
97 alpha,
98 ..Default::default()
99 },
100 )?
101 } else {
102 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 ops::matmul(
116 &at,
117 dy,
118 None,
119 MatmulSpec {
120 trans_a: true,
121 alpha,
122 ..Default::default()
123 },
124 )?
125 } else {
126 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, >, dy, eps)?;
165 Ok(vec![(ix, dx), (ig, dgamma), (ib, dbeta)])
166 }),
167 ))
168 }
169
170 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 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 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 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 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; };
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}