use nalgebra::{DMatrix, DVector};
#[derive(Clone)]
pub struct Linear {
weights: DMatrix<f32>,
bias: DVector<f32>,
apply_relu: bool
}
impl Linear {
pub fn new(weights_vector: Vec<f32>, bias_vector: Vec<f32>, apply_relu: bool) -> Self {
let weights = DMatrix::from_vec(bias_vector.len(), weights_vector.len() / bias_vector.len(), weights_vector);
let bias = DVector::from_vec(bias_vector);
Self { weights, bias, apply_relu }
}
pub fn forward(&self, input: &DVector<f32>) -> DVector<f32> {
let mut out = (&self.weights * input) + &self.bias;
if self.apply_relu {
out = out.map(relu);
}
out
}
}
#[derive(Clone)]
pub struct EmbeddingBag {
vectors: Vec<DVector<f32>>,
bias: DVector<f32>,
apply_relu: bool,
obs_shape: Vec<usize>,
conv_dim: usize
}
impl EmbeddingBag {
pub fn new(vec_vectors: Vec<Vec<f32>>, bias_vector: Vec<f32>, apply_relu: bool, obs_shape: Vec<usize>, conv_dim: usize) -> Self {
let vectors = vec_vectors.into_iter().map(|vec| DVector::from_vec(vec)).collect();
let bias = DVector::from_vec(bias_vector);
Self { vectors, bias, apply_relu, obs_shape, conv_dim }
}
pub fn forward(&self, input: &Vec<usize>) -> DVector<f32> {
let mut out = self.bias.clone();
if self.obs_shape.len() == 1 {
for &i in input.iter() {
out += &self.vectors[i];
}
} else if self.obs_shape.len() == 2 {
let v_size = self.vectors[0].len();
for &i in input.iter() {
let mut row = i / self.obs_shape[1];
let mut col = i % self.obs_shape[1];
if self.conv_dim == 1 {(row, col) = (col, row);}
let mut out_slice = out.rows_mut(col * v_size, v_size);
out_slice += &self.vectors[row];
}
}
if self.apply_relu {
out = out.map(relu);
}
out
}
}
fn relu(x: f32) -> f32 {
if x > 0.0 { x } else { 0.0 }
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_linear_forward() {
let linear = Linear::new(vec![1.0, 2.0, 3.0, 4.0], vec![1.0, 1.0], false);
let input = DVector::from_vec(vec![1.0, 2.0]);
let out = linear.forward(&input);
assert_eq!(out, DVector::from_vec(vec![8.0, 11.0]));
}
#[test]
fn test_linear_forward_relu() {
let linear = Linear::new(vec![-1.0, -2.0, 0.0, 1.0], vec![0.0, 0.0], true);
let input = DVector::from_vec(vec![1.0, 2.0]);
let out = linear.forward(&input);
assert_eq!(out, DVector::from_vec(vec![0.0, 0.0]));
}
#[test]
fn test_embedding_bag_forward() {
let emb = EmbeddingBag::new(
vec![vec![1.0, 2.0], vec![3.0, 4.0]],
vec![0.0, 0.0],
false,
vec![2],
0,
);
let out = emb.forward(&vec![0, 1]);
assert_eq!(out, DVector::from_vec(vec![4.0, 6.0]));
}
}