ruda_tensor/ops/modules/
linear.rs1use alloc::vec;
2
3use crate::tensor::FloatTensor;
4use crate::{Backend, TensorMetadata};
5use ruda_core::tensor::Shape;
6
7pub(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 let weight = unsqueeze_leading::<B>(weight, x_ndims);
22 let output = B::float_matmul(x, weight);
23
24 match bias {
25 Some(bias) => {
26 let bias = unsqueeze_leading::<B>(bias, x_ndims);
28 B::float_add(output, bias)
29 }
30 None => output,
31 }
32}
33
34fn 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
48pub(crate) fn linear_x_backward<B: Backend>(
52 weight: FloatTensor<B>,
53 output_grad: FloatTensor<B>,
54) -> FloatTensor<B> {
55 let weight = B::float_swap_dims(weight, 0, 1);
57 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
63pub(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 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
90pub(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 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}