Skip to main content

broadcast_shape

Function broadcast_shape 

pub fn broadcast_shape<B>(
    grad: <B as BackendTypes>::FloatTensorPrimitive,
    shape: &Shape,
) -> <B as BackendTypes>::FloatTensorPrimitive
where B: Backend,
Expand description

Make sure the grad tensor has the given shape.

If broadcasting happened during the forward pass, the gradients will be sum along the broadcasted dimension.