Skip to main content

ruda_tensor/ops/modules/
linear.rs

1use alloc::vec;
2
3use crate::tensor::FloatTensor;
4use crate::{Backend, TensorMetadata};
5use ruda_core::tensor::Shape;
6
7/// Default [linear](crate::ops::ModuleOps::linear) forward implementation.
8///
9/// Computes `y = x @ weight [+ bias]`.
10///
11/// Weight `[d_input, d_output]` and bias `[d_output]` are broadcast to match x's rank
12/// before the matmul.
13pub(crate) fn linear<B: Backend>(
14    x: FloatTensor<B>,
15    weight: FloatTensor<B>,
16    bias: Option<FloatTensor<B>>,
17) -> FloatTensor<B> {
18    let x_ndims = x.shape().num_dims();
19
20    // Reshape weight [d_input, d_output] -> [1, ..., 1, d_input, d_output] for batch matmul.
21    let weight = unsqueeze_leading::<B>(weight, x_ndims);
22    let output = B::float_matmul(x, weight);
23
24    match bias {
25        Some(bias) => {
26            // Reshape bias [d_output] -> [1, ..., 1, d_output] to match output rank.
27            let bias = unsqueeze_leading::<B>(bias, x_ndims);
28            B::float_add(output, bias)
29        }
30        None => output,
31    }
32}
33
34/// Reshape a tensor by prepending size-1 dimensions until it has `target_ndims` dimensions.
35fn unsqueeze_leading<B: Backend>(tensor: FloatTensor<B>, target_ndims: usize) -> FloatTensor<B> {
36    let shape = tensor.shape();
37    let ndims = shape.num_dims();
38    if ndims >= target_ndims {
39        return tensor;
40    }
41    let mut new_dims = vec![1usize; target_ndims - ndims];
42    for i in 0..ndims {
43        new_dims.push(shape[i]);
44    }
45    B::float_reshape(tensor, Shape::from(new_dims))
46}
47
48/// Default [linear_x_backward](crate::ops::ModuleOps::linear_x_backward) implementation.
49///
50/// Computes `dx = output_grad @ weight^T`.
51pub(crate) fn linear_x_backward<B: Backend>(
52    weight: FloatTensor<B>,
53    output_grad: FloatTensor<B>,
54) -> FloatTensor<B> {
55    // weight is [d_input, d_output], transpose to [d_output, d_input]
56    let weight = B::float_swap_dims(weight, 0, 1);
57    // Unsqueeze to match output_grad rank for batch matmul
58    let grad_ndims = output_grad.shape().num_dims();
59    let weight = unsqueeze_leading::<B>(weight, grad_ndims);
60    B::float_matmul(output_grad, weight)
61}
62
63/// Default [linear_weight_backward](crate::ops::ModuleOps::linear_weight_backward) implementation.
64///
65/// Computes `dW = x^T @ output_grad`, summed over batch dimensions.
66pub(crate) fn linear_weight_backward<B: Backend>(
67    x: FloatTensor<B>,
68    output_grad: FloatTensor<B>,
69) -> FloatTensor<B> {
70    let ndims = x.shape().num_dims();
71    let x = B::float_swap_dims(x, ndims - 2, ndims - 1);
72    let mut grad = B::float_matmul(x, output_grad);
73
74    // Sum over all batch dimensions (all dims except the last two).
75    // float_sum_dim preserves rank (keepdim), so sum each batch dim at its index.
76    let ndims = grad.shape().num_dims();
77    if ndims > 2 {
78        for dim in 0..ndims - 2 {
79            grad = B::float_sum_dim(grad, dim);
80        }
81        let shape = grad.shape();
82        let d0 = shape[ndims - 2];
83        let d1 = shape[ndims - 1];
84        B::float_reshape(grad, Shape::new([d0, d1]))
85    } else {
86        grad
87    }
88}
89
90/// Default [linear_bias_backward](crate::ops::ModuleOps::linear_bias_backward) implementation.
91///
92/// Computes `db = sum(output_grad)` over all dimensions except the last.
93pub(crate) fn linear_bias_backward<B: Backend>(output_grad: FloatTensor<B>) -> FloatTensor<B> {
94    let ndims = output_grad.shape().num_dims();
95    let mut grad = output_grad;
96
97    // Sum over all dims except the last (the output feature dim).
98    // float_sum_dim preserves rank (keepdim), so sum each dim at its index.
99    for dim in 0..ndims - 1 {
100        grad = B::float_sum_dim(grad, dim);
101    }
102
103    let shape = grad.shape();
104    let d_output = shape[ndims - 1];
105    B::float_reshape(grad, Shape::new([d_output]))
106}