use crate::traits::{Tensor, TensorOps};
use crate::SmeltError;
#[derive(Clone)]
pub struct Linear<T: Tensor> {
weight: T,
bias: T,
}
impl<T: Tensor + TensorOps<T>> Linear<T> {
pub fn new(weight: T, bias: T) -> Self {
Self { weight, bias }
}
pub fn forward(&self, tensor: &T, out: &mut T) -> Result<(), SmeltError> {
T::matmul_t(tensor, &self.weight, out)?;
T::add(&self.bias, out)?;
Ok(())
}
pub fn weight(&self) -> &T {
&self.weight
}
pub fn bias(&self) -> &T {
&self.bias
}
}
#[derive(Clone)]
pub struct LinearT<T: Tensor> {
weight: T,
bias: T,
}
impl<T: Tensor + TensorOps<T>> LinearT<T> {
pub fn new(weight: T, bias: T) -> Self {
Self { weight, bias }
}
pub fn forward(&self, tensor: &T, out: &mut T) -> Result<(), SmeltError> {
T::matmul_t(tensor, &self.weight, out)?;
T::add(&self.bias, out)?;
Ok(())
}
}
#[derive(Clone)]
pub struct UnbiasedLinear<T: Tensor> {
weight: T,
}
impl<T: Tensor + TensorOps<T>> UnbiasedLinear<T> {
pub fn new(weight: T) -> Self {
Self { weight }
}
pub fn forward(&self, tensor: &T, out: &mut T) -> Result<(), SmeltError> {
T::matmul_t(tensor, &self.weight, out)?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::cpu::f32::Tensor;
#[test]
fn test_linear() {
let zeros = Tensor::zeros(vec![2, 2]);
let weights = Tensor::zeros(vec![3, 2]);
let bias = Tensor::zeros(vec![3]);
let mut out = Tensor::zeros(vec![2, 3]);
let linear = Linear::new(weights, bias);
linear.forward(&zeros, &mut out).unwrap();
}
}