Skip to main content

luma_tensor/grad/
backprop.rs

1//! Reverse-mode autograd. `backward()` seeds the output with ones, walks the
2//! graph in reverse-topological order, and accumulates gradients into a
3//! [`GradStore`]. Only `Float` tensors participate.
4
5use 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    /// Accumulate gradients of `self` w.r.t. all leaf tensors that require grad
10    /// into an existing [`GradStore`].
11    ///
12    /// Unlike [`Tensor::backward`], this does not create a fresh store; leaf
13    /// gradients are added onto whatever is already present. Calling it multiple
14    /// times with different loss tensors (e.g. one per micro-batch) and then
15    /// stepping the optimizer once implements gradient accumulation.
16    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    /// Compute gradients of `self` w.r.t. all leaf tensors that require grad.
37    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    /// Reverse-topological order of nodes reachable from grad-requiring leaves.
44    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
79/// Float inputs of an op that gradient can flow through.
80fn 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
104/// Build the keepdim shape: `arg` dims with each reduced axis set to 1.
105fn 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        // ---- Binary tensor-tensor ----
116        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            // d/drhs (lhs/rhs) = -lhs/rhs^2
134            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            // split the gradient where both equal the output (scale by 1/(mask+1)).
141            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        // ---- Binary scalar-rhs ----
148        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        // ---- Binary scalar-lhs ----
166        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            // y = c / x => dy/dx = -c / x^2
178            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        // ---- Matmul ----
192        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        // ---- Reduce ----
198        Op::Reduce(arg, rop, reduced) => backward_reduce(node, arg, *rop, reduced, grad, grads)?,
199
200        // ---- Broadcast: sum grad over the broadcasted dims ----
201        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        // ---- shape movements ----
223        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        // ---- Cast: cast the gradient back to the input precision ----
257        Op::Cast(arg) => {
258            let g = grad.cast(arg.dtype())?;
259            grads.or_insert(arg)?.impl_add_(&g)?;
260        }
261
262        // ---- indexing ----
263        Op::IndexSelect(arg, indices, dim) => {
264            // scatter grad back to the selected positions.
265            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        // ---- pick: route grad through the mask to each branch ----
286        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        // ---- softmax: g_in = y * (g - sum(g*y, dim, keepdim)) ----
298        Op::Softmax(input, dim) => {
299            let y = node; // softmax output
300            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        // ---- slice: dilate (step>1) + pad back to arg shape ----
307        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                // Insert a dim of size 1 after `dim`, so that each grad element
318                // can be interleaved with (step-1) zeros.
319                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                // Build the gap zeros alongside the unsqueezed dim.
324                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                // Interleave: cat along the new dim, then flatten back.
329                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        // ---- not yet wired ----
344        Op::RmsNorm(..) => return Err(crate::Error::BackwardNotSupported("rms_norm")),
345    }
346    Ok(())
347}
348
349/// Pad `grad` with zeros along `dim` back to `arg`'s size (narrow backward).
350fn 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
374/// Gradient of unary ops. Fused kernels in luma-core are expanded here into
375/// composed tensor ops (correctness first; can be fused later).
376fn 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    // local gradient factor `local` such that d(arg) += grad * local
384    let contrib = match op {
385        FloatUnaryOp::Exp => grad.mul(node)?,               // d/dx e^x = e^x = node
386        FloatUnaryOp::Ln => grad.div(arg)?,                 // 1/x
387        FloatUnaryOp::Sin => grad.mul(&arg.cos()?)?,        // cos x
388        FloatUnaryOp::Cos => grad.mul(&arg.sin()?)?.neg()?, // -sin x
389        FloatUnaryOp::Tanh => {
390            // 1 - tanh^2 = 1 - node^2
391            let factor = node.sqr()?.neg()?.add_scalar(1.0)?;
392            grad.mul(&factor)?
393        }
394        FloatUnaryOp::Sqr => grad.mul(arg)?.mul_scalar(2.0)?, // 2x
395        FloatUnaryOp::Sqrt => {
396            // 1/(2 sqrt x) = 0.5 / node
397            grad.mul_scalar(0.5)?.div(node)?
398        }
399        FloatUnaryOp::Recip => {
400            // -1/x^2 = -node^2
401            grad.mul(&node.sqr()?)?.neg()?
402        }
403        FloatUnaryOp::Relu => {
404            // mask = arg > 0
405            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            // node * (1 - node)
415            let factor = node.neg()?.add_scalar(1.0)?.mul(node)?;
416            grad.mul(&factor)?
417        }
418        FloatUnaryOp::Erf => {
419            // 2/sqrt(pi) * exp(-x^2)
420            let scale = 2.0 / std::f64::consts::PI.sqrt();
421            grad.mul(&arg.sqr()?.neg()?.exp()?)?.mul_scalar(scale)?
422        }
423        FloatUnaryOp::Silu => {
424            // sig = sigmoid(x); silu = x*sig; d = sig*(1 - silu) + silu
425            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            // matches luma-core's tanh-approx gelu grad
432            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            // 0.5*tanh + dt*(1 - tanh^2) + 0.5
441            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            // c1 * exp(-x^2/2) * x + erf(x/sqrt2)/2 + 0.5
446            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    // ---- elementwise ops ----
470    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 => {} // gradient is zero everywhere
479        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
505/// Gradient of reductions. Reshapes grad to keepdim form then broadcasts back.
506fn 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            // d(prod)/dx_i = prod / x_i (if x_i != 0)
527            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            // route grad to the arg elements equal to the reduced value.
533            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}