1use crate::{BinaryOp, Device, Float, FloatUnaryOp, GradStore, Op, ReduceOp, Shape, Tensor, TensorId, UnaryOp, no_grad};
6use std::collections::HashMap;
7
8impl<D: Device> Tensor<D, Float> {
9 pub fn backward_into(&self, grads: &mut GradStore<D>) -> crate::Result<()> {
17 no_grad!();
18 let sorted = self.sorted_nodes();
19 grads.insert(self, self.ones_like()?);
20
21 for node in sorted.iter() {
22 let op = match node.op() {
23 None => {
24 debug_assert!(node.is_leaf());
25 continue;
26 }
27 Some(op) => op,
28 };
29 let grad = grads.remove(node).expect("grad not populated");
30 backward_op(node, op, &grad, grads)?;
31 }
32
33 Ok(())
34 }
35
36 pub fn backward(&self) -> crate::Result<GradStore<D>> {
38 let mut grads = GradStore::new();
39 self.backward_into(&mut grads)?;
40 Ok(grads)
41 }
42
43 pub fn sorted_nodes(&self) -> Vec<&Tensor<D, Float>> {
45 fn walk<'a, D: Device>(
46 node: &'a Tensor<D, Float>,
47 mut nodes: Vec<&'a Tensor<D, Float>>,
48 seen: &mut HashMap<TensorId, bool>,
49 ) -> (bool, Vec<&'a Tensor<D, Float>>) {
50 if let Some(&tg) = seen.get(&node.id()) {
51 return (tg, nodes);
52 }
53 let mut track = false;
54 nodes = if node.is_leaf() {
55 track = true;
56 nodes
57 } else if let Some(op) = node.op() {
58 for input in op_inputs(op) {
59 let (tg, n) = walk(input, nodes, seen);
60 track |= tg;
61 nodes = n;
62 }
63 nodes
64 } else {
65 nodes
66 };
67 seen.insert(node.id(), track);
68 if track {
69 nodes.push(node);
70 }
71 (track, nodes)
72 }
73 let (_tg, mut nodes) = walk(self, vec![], &mut HashMap::new());
74 nodes.reverse();
75 nodes
76 }
77}
78
79fn op_inputs<D: Device>(op: &Op<D>) -> Vec<&Tensor<D, Float>> {
81 match op {
82 Op::Binary(a, b, _) | Op::Matmul(a, b) => vec![a, b],
83 Op::BinaryScalarRhs(a, _, _) | Op::BinaryScalarLhs(_, a, _) => vec![a],
84 Op::FloatUnary(a, _)
85 | Op::Unary(a, _)
86 | Op::Reduce(a, _, _)
87 | Op::Broadcast(a)
88 | Op::Narrow(a, _, _, _)
89 | Op::Slice(a, _, _, _, _)
90 | Op::Reshape(a)
91 | Op::Transpose(a, _, _)
92 | Op::Permute(a, _)
93 | Op::Copy(a)
94 | Op::Cast(a)
95 | Op::IndexSelect(a, _, _)
96 | Op::Gather(a, _, _)
97 | Op::Softmax(a, _) => vec![a],
98 Op::IndexAdd(a, _, b, _) | Op::ScatterAdd(a, _, b, _) | Op::RmsNorm(a, b, _) => vec![a, b],
99 Op::Cat(args, _) => args.iter().collect(),
100 Op::Pick(_, tv, fv) => tv.iter().chain(fv.iter()).collect(),
101 }
102}
103
104fn keepdim_shape(arg_dims: &[usize], reduced: &[usize]) -> Shape {
106 let mut dims = arg_dims.to_vec();
107 for &d in reduced {
108 dims[d] = 1;
109 }
110 Shape::from(dims)
111}
112
113fn backward_op<D: Device>(node: &Tensor<D, Float>, op: &Op<D>, grad: &Tensor<D, Float>, grads: &mut GradStore<D>) -> crate::Result<()> {
114 match op {
115 Op::Binary(lhs, rhs, BinaryOp::Add) => {
117 grads.or_insert(lhs)?.impl_add_(grad)?;
118 grads.or_insert(rhs)?.impl_add_(grad)?;
119 }
120 Op::Binary(lhs, rhs, BinaryOp::Sub) => {
121 grads.or_insert(lhs)?.impl_add_(grad)?;
122 grads.or_insert(rhs)?.impl_sub_(grad)?;
123 }
124 Op::Binary(lhs, rhs, BinaryOp::Mul) => {
125 let lg = grad.mul(rhs)?;
126 grads.or_insert(lhs)?.impl_add_(&lg)?;
127 let rg = grad.mul(lhs)?;
128 grads.or_insert(rhs)?.impl_add_(&rg)?;
129 }
130 Op::Binary(lhs, rhs, BinaryOp::Div) => {
131 let lg = grad.div(rhs)?;
132 grads.or_insert(lhs)?.impl_add_(&lg)?;
133 let rg = grad.mul(lhs)?.div(&rhs.sqr()?)?;
135 grads.or_insert(rhs)?.impl_sub_(&rg)?;
136 }
137 Op::Binary(lhs, rhs, BinaryOp::Maximum) | Op::Binary(lhs, rhs, BinaryOp::Minimum) => {
138 let mask_lhs = node.eq(lhs)?.cast_float(node.dtype())?;
139 let mask_rhs = node.eq(rhs)?.cast_float(node.dtype())?;
140 let lg = mask_lhs.mul(grad)?.div(&mask_rhs.add_scalar(1.0)?)?;
142 grads.or_insert(lhs)?.impl_add_(&lg)?;
143 let rg = mask_rhs.mul(grad)?.div(&mask_lhs.add_scalar(1.0)?)?;
144 grads.or_insert(rhs)?.impl_add_(&rg)?;
145 }
146
147 Op::BinaryScalarRhs(lhs, _, BinaryOp::Add) | Op::BinaryScalarRhs(lhs, _, BinaryOp::Sub) => {
149 grads.or_insert(lhs)?.impl_add_(grad)?;
150 }
151 Op::BinaryScalarRhs(lhs, c, BinaryOp::Mul) => {
152 let lg = grad.mul_scalar(*c)?;
153 grads.or_insert(lhs)?.impl_add_(&lg)?;
154 }
155 Op::BinaryScalarRhs(lhs, c, BinaryOp::Div) => {
156 let lg = grad.div_scalar(*c)?;
157 grads.or_insert(lhs)?.impl_add_(&lg)?;
158 }
159 Op::BinaryScalarRhs(lhs, _, BinaryOp::Maximum) | Op::BinaryScalarRhs(lhs, _, BinaryOp::Minimum) => {
160 let mask = node.eq(lhs)?.cast_float(node.dtype())?;
161 let lg = mask.mul(grad)?;
162 grads.or_insert(lhs)?.impl_add_(&lg)?;
163 }
164
165 Op::BinaryScalarLhs(_, rhs, BinaryOp::Add) => {
167 grads.or_insert(rhs)?.impl_add_(grad)?;
168 }
169 Op::BinaryScalarLhs(_, rhs, BinaryOp::Sub) => {
170 grads.or_insert(rhs)?.impl_sub_(grad)?;
171 }
172 Op::BinaryScalarLhs(c, rhs, BinaryOp::Mul) => {
173 let rg = grad.mul_scalar(*c)?;
174 grads.or_insert(rhs)?.impl_add_(&rg)?;
175 }
176 Op::BinaryScalarLhs(c, rhs, BinaryOp::Div) => {
177 let rg = grad.mul_scalar(-*c)?.div(&rhs.sqr()?)?;
179 grads.or_insert(rhs)?.impl_add_(&rg)?;
180 }
181 Op::BinaryScalarLhs(_, rhs, BinaryOp::Maximum) | Op::BinaryScalarLhs(_, rhs, BinaryOp::Minimum) => {
182 let mask = node.eq(rhs)?.cast_float(node.dtype())?;
183 let rg = mask.mul(grad)?;
184 grads.or_insert(rhs)?.impl_add_(&rg)?;
185 }
186
187 Op::FloatUnary(arg, uop) => backward_float_unary(node, arg, *uop, grad, grads)?,
188
189 Op::Unary(arg, uop) => backward_unary(node, arg, *uop, grad, grads)?,
190
191 Op::Matmul(lhs, rhs) => {
193 grads.or_insert(lhs)?.add_matmul_(grad, &rhs.transpose_last()?)?;
194 grads.or_insert(rhs)?.add_matmul_(&lhs.transpose_last()?, grad)?;
195 }
196
197 Op::Reduce(arg, rop, reduced) => backward_reduce(node, arg, *rop, reduced, grad, grads)?,
199
200 Op::Broadcast(arg) => {
202 let arg_dims = arg.dims();
203 let node_dims = node.dims();
204 let left = node_dims.len() - arg_dims.len();
205 let mut sum_dims: Vec<usize> = (0..left).collect();
206 for (d, (nd, ad)) in node_dims[left..].iter().zip(arg_dims.iter()).enumerate() {
207 if nd != ad {
208 sum_dims.push(d + left);
209 }
210 }
211 let mut arg_grad = grad.clone();
212 for &d in sum_dims.iter() {
213 arg_grad = arg_grad.sum_keepdim(d)?;
214 }
215 for _ in 0..left {
216 arg_grad = arg_grad.squeeze(0)?;
217 }
218 let g = arg_grad.broadcast_as(arg.shape().clone())?;
219 grads.or_insert(arg)?.impl_add_(&g)?;
220 }
221
222 Op::Reshape(arg) => {
224 let g = grad.reshape(arg.shape().clone())?;
225 grads.or_insert(arg)?.impl_add_(&g)?;
226 }
227 Op::Transpose(arg, d1, d2) => {
228 let g = grad.transpose(*d1, *d2)?;
229 grads.or_insert(arg)?.impl_add_(&g)?;
230 }
231 Op::Permute(arg, dims) => {
232 let mut inv = vec![0; dims.len()];
233 for (i, &d) in dims.iter().enumerate() {
234 inv[d] = i;
235 }
236 let g = grad.permute(inv)?;
237 grads.or_insert(arg)?.impl_add_(&g)?;
238 }
239 Op::Narrow(arg, dim, start, len) => {
240 let g = pad_grad_along(arg, grad, *dim, *start, *len)?;
241 grads.or_insert(arg)?.impl_add_(&g)?;
242 }
243 Op::Cat(args, dim) => {
244 let mut start = 0;
245 for arg in args {
246 let len = arg.dims()[*dim];
247 let g = grad.narrow(*dim, start, len)?;
248 grads.or_insert(arg)?.impl_add_(&g)?;
249 start += len;
250 }
251 }
252 Op::Copy(arg) => {
253 grads.or_insert(arg)?.impl_add_(grad)?;
254 }
255
256 Op::Cast(arg) => {
258 let g = grad.cast(arg.dtype())?;
259 grads.or_insert(arg)?.impl_add_(&g)?;
260 }
261
262 Op::IndexSelect(arg, indices, dim) => {
264 let acc = grads.or_insert(arg)?;
266 let updated = acc.index_add(indices, grad, *dim)?;
267 *grads.or_insert(arg)? = updated;
268 }
269 Op::Gather(arg, indices, dim) => {
270 let acc = grads.or_insert(arg)?;
271 let updated = acc.scatter_add(indices, grad, *dim)?;
272 *grads.or_insert(arg)? = updated;
273 }
274 Op::IndexAdd(init, indices, src, dim) => {
275 grads.or_insert(init)?.impl_add_(grad)?;
276 let src_grad = grad.index_select(indices, *dim)?;
277 grads.or_insert(src)?.impl_add_(&src_grad)?;
278 }
279 Op::ScatterAdd(init, indices, src, dim) => {
280 grads.or_insert(init)?.impl_add_(grad)?;
281 let src_grad = grad.gather(indices, *dim)?;
282 grads.or_insert(src)?.impl_add_(&src_grad)?;
283 }
284
285 Op::Pick(mask, tv, fv) => {
287 if let Some(tv) = tv {
288 let g = mask.pick_false(grad, 0.0)?;
289 grads.or_insert(tv)?.impl_add_(&g)?;
290 }
291 if let Some(fv) = fv {
292 let g = mask.pick_true(0.0, grad)?;
293 grads.or_insert(fv)?.impl_add_(&g)?;
294 }
295 }
296
297 Op::Softmax(input, dim) => {
299 let y = node; let gy = grad.mul(y)?;
301 let s = gy.sum_keepdim(*dim)?;
302 let g = y.mul(&grad.sub(&s.broadcast_as(grad.shape().clone())?)?)?;
303 grads.or_insert(input)?.impl_add_(&g)?;
304 }
305
306 Op::Slice(arg, dim, start, _end, step) => {
308 let arg_dtype = arg.dtype();
309
310 let body_grad: Tensor<D, Float>;
311 let body_grad_ref = if *step == 1 {
312 grad
313 } else {
314 let grad_len = grad.dims()[*dim];
315 let span_len = if grad_len > 0 { (grad_len - 1) * step + 1 } else { 0 };
316
317 let mut unsqueezed_shape = grad.dims().to_vec();
320 unsqueezed_shape.insert(*dim + 1, 1);
321 let grad_unsqueezed = grad.reshape(Shape::from(unsqueezed_shape))?;
322
323 let mut zeros_shape = grad_unsqueezed.dims().to_vec();
325 zeros_shape[*dim + 1] = step - 1;
326 let zeros_gap = Tensor::<D, Float>::zeros(Shape::from(zeros_shape), (arg.device(), arg_dtype))?;
327
328 let dilated = Tensor::cat(&[&grad_unsqueezed, &zeros_gap], *dim + 1)?;
330 let mut flattened_shape = grad.dims().to_vec();
331 flattened_shape[*dim] = grad_len * step;
332 let flattened = dilated.reshape(Shape::from(flattened_shape))?;
333
334 body_grad = flattened.narrow(*dim, 0, span_len)?;
335 &body_grad
336 };
337
338 let body_len = body_grad_ref.dims()[*dim];
339 let arg_grad = pad_grad_along(arg, body_grad_ref, *dim, *start, body_len)?;
340 grads.or_insert(arg)?.impl_add_(&arg_grad)?;
341 }
342
343 Op::RmsNorm(..) => return Err(crate::Error::BackwardNotSupported("rms_norm")),
345 }
346 Ok(())
347}
348
349fn pad_grad_along<D: Device>(
351 arg: &Tensor<D, Float>,
352 grad: &Tensor<D, Float>,
353 dim: usize,
354 start: usize,
355 len: usize,
356) -> crate::Result<Tensor<D, Float>> {
357 let arg_dims = arg.dims();
358 let make_pad = |size: usize| {
359 let mut dims = arg_dims.to_vec();
360 dims[dim] = size;
361 Tensor::<D, Float>::zeros(Shape::from(dims), (arg.device(), arg.dtype()))
362 };
363 let right = arg_dims[dim] - start - len;
364 let left_pad = if start != 0 { Some(make_pad(start)?) } else { None };
365 let right_pad = if right != 0 { Some(make_pad(right)?) } else { None };
366 match (left_pad, right_pad) {
367 (None, None) => Ok(grad.clone()),
368 (Some(l), None) => Tensor::cat(&[&l, grad], dim),
369 (None, Some(r)) => Tensor::cat(&[grad, &r], dim),
370 (Some(l), Some(r)) => Tensor::cat(&[&l, grad, &r], dim),
371 }
372}
373
374fn backward_float_unary<D: Device>(
377 node: &Tensor<D, Float>,
378 arg: &Tensor<D, Float>,
379 op: FloatUnaryOp,
380 grad: &Tensor<D, Float>,
381 grads: &mut GradStore<D>,
382) -> crate::Result<()> {
383 let contrib = match op {
385 FloatUnaryOp::Exp => grad.mul(node)?, FloatUnaryOp::Ln => grad.div(arg)?, FloatUnaryOp::Sin => grad.mul(&arg.cos()?)?, FloatUnaryOp::Cos => grad.mul(&arg.sin()?)?.neg()?, FloatUnaryOp::Tanh => {
390 let factor = node.sqr()?.neg()?.add_scalar(1.0)?;
392 grad.mul(&factor)?
393 }
394 FloatUnaryOp::Sqr => grad.mul(arg)?.mul_scalar(2.0)?, FloatUnaryOp::Sqrt => {
396 grad.mul_scalar(0.5)?.div(node)?
398 }
399 FloatUnaryOp::Recip => {
400 grad.mul(&node.sqr()?)?.neg()?
402 }
403 FloatUnaryOp::Relu => {
404 let mask = arg.gt(&arg.zeros_like()?)?.cast_float(arg.dtype())?;
406 grad.mul(&mask)?
407 }
408 FloatUnaryOp::LeakyRelu(slope) => {
409 let pos = arg.gt(&arg.zeros_like()?)?.cast_float(arg.dtype())?;
410 let neg = pos.neg()?.add_scalar(1.0)?.mul_scalar(slope)?;
411 grad.mul(&pos.add(&neg)?)?
412 }
413 FloatUnaryOp::Sigmoid => {
414 let factor = node.neg()?.add_scalar(1.0)?.mul(node)?;
416 grad.mul(&factor)?
417 }
418 FloatUnaryOp::Erf => {
419 let scale = 2.0 / std::f64::consts::PI.sqrt();
421 grad.mul(&arg.sqr()?.neg()?.exp()?)?.mul_scalar(scale)?
422 }
423 FloatUnaryOp::Silu => {
424 let sig = arg.sigmoid()?;
426 let silu = arg.mul(&sig)?;
427 let factor = sig.mul(&silu.neg()?.add_scalar(1.0)?)?.add(&silu)?;
428 grad.mul(&factor)?
429 }
430 FloatUnaryOp::Gelu => {
431 let c1 = 0.0356774;
433 let c2 = 0.797885;
434 let c3 = 0.0535161;
435 let c4 = 0.398942;
436 let x3 = arg.mul(&arg.sqr()?)?;
437 let inner = x3.mul_scalar(c1)?.add(&arg.mul_scalar(c2)?)?;
438 let tanh = inner.tanh()?;
439 let dt = x3.mul_scalar(c3)?.add(&arg.mul_scalar(c4)?)?;
440 let factor = tanh.mul_scalar(0.5)?.add(&dt.mul(&tanh.sqr()?.neg()?.add_scalar(1.0)?)?)?.add_scalar(0.5)?;
442 grad.mul(&factor)?
443 }
444 FloatUnaryOp::GeluErf => {
445 let c1 = 0.398942;
447 let sqrt2 = std::f64::consts::SQRT_2;
448 let neg_half_sq = arg.sqr()?.neg()?.div_scalar(2.0)?;
449 let scaled_exp = neg_half_sq.exp()?.mul(arg)?.mul_scalar(c1)?;
450 let erf_term = arg.div_scalar(sqrt2)?.erf()?.div_scalar(2.0)?;
451 let factor = scaled_exp.add(&erf_term)?.add_scalar(0.5)?;
452 grad.mul(&factor)?
453 }
454 FloatUnaryOp::Floor | FloatUnaryOp::Ceil | FloatUnaryOp::Round => {
455 return Err(crate::Error::BackwardNotSupported("floor/ceil/round"));
456 }
457 };
458 grads.or_insert(arg)?.impl_add_(&contrib)?;
459 Ok(())
460}
461
462fn backward_unary<D: Device>(
463 _node: &Tensor<D, Float>,
464 arg: &Tensor<D, Float>,
465 op: UnaryOp<f64>,
466 grad: &Tensor<D, Float>,
467 grads: &mut GradStore<D>,
468) -> crate::Result<()> {
469 match op {
471 UnaryOp::Neg => {
472 grads.or_insert(arg)?.impl_add_(&grad.neg()?)?;
473 }
474 UnaryOp::Abs => {
475 let g = grad.mul(&arg.sign()?)?;
476 grads.or_insert(arg)?.impl_add_(&g)?;
477 }
478 UnaryOp::Sign => {} UnaryOp::Pow(e) => {
480 let g = grad.mul(&arg.pow(e - 1.0)?)?.mul_scalar(e)?;
481 grads.or_insert(arg)?.impl_add_(&g)?;
482 }
483 UnaryOp::Affine(mul, _add) => {
484 let g = grad.mul_scalar(mul)?;
485 grads.or_insert(arg)?.impl_add_(&g)?;
486 }
487 UnaryOp::Clamp(min, max) => {
488 let dtype = arg.dtype();
489 let mut mask = arg.ones_like()?;
490 if let Some(lo) = min {
491 let t_lo = arg.zeros_like()?.add_scalar(lo)?;
492 mask = mask.mul(&arg.gt(&t_lo)?.cast_float(dtype)?)?;
493 }
494 if let Some(hi) = max {
495 let t_hi = arg.zeros_like()?.add_scalar(hi)?;
496 mask = mask.mul(&arg.lt(&t_hi)?.cast_float(dtype)?)?;
497 }
498 let g = grad.mul(&mask)?;
499 grads.or_insert(arg)?.impl_add_(&g)?;
500 }
501 };
502 Ok(())
503}
504
505fn backward_reduce<D: Device>(
507 node: &Tensor<D, Float>,
508 arg: &Tensor<D, Float>,
509 op: ReduceOp,
510 reduced: &[usize],
511 grad: &Tensor<D, Float>,
512 grads: &mut GradStore<D>,
513) -> crate::Result<()> {
514 let keep = keepdim_shape(arg.dims(), reduced);
515 match op {
516 ReduceOp::Sum => {
517 let g = grad.reshape(keep)?.broadcast_as(arg.shape().clone())?;
518 grads.or_insert(arg)?.impl_add_(&g)?;
519 }
520 ReduceOp::Mean => {
521 let n = arg.element_count() / node.element_count().max(1);
522 let g = grad.reshape(keep)?.broadcast_as(arg.shape().clone())?.div_scalar(n as f64)?;
523 grads.or_insert(arg)?.impl_add_(&g)?;
524 }
525 ReduceOp::Prod => {
526 let n_broadcast = node.reshape(keep.clone())?.broadcast_as(arg.shape().clone())?;
528 let g = grad.reshape(keep)?.broadcast_as(arg.shape().clone())?.mul(&n_broadcast)?.div(arg)?;
529 grads.or_insert(arg)?.impl_add_(&g)?;
530 }
531 ReduceOp::Max | ReduceOp::Min => {
532 let node_b = node.reshape(keepdim_shape(arg.dims(), reduced))?.broadcast_as(arg.shape().clone())?;
534 let mask = node_b.eq(arg)?.cast_float(arg.dtype())?;
535 let g = grad.reshape(keepdim_shape(arg.dims(), reduced))?.broadcast_as(arg.shape().clone())?.mul(&mask)?;
536 grads.or_insert(arg)?.impl_add_(&g)?;
537 }
538 }
539 Ok(())
540}