use crate::nn::module::Module;
use crate::tensor::Tensor;
use crate::error::Result;
use ndarray::IxDyn;
use ndarray_rand::rand_distr::Uniform;
use ndarray_rand::RandomExt;
pub struct Embedding {
pub weights: Tensor,
pub vocab_size: usize,
pub embedding_dim: usize,
}
impl Embedding {
pub fn new(vocab_size: usize, embedding_dim: usize) -> Self {
let weights_data = ndarray::ArrayD::random(
IxDyn(&[vocab_size, embedding_dim]),
Uniform::new(-1.0, 1.0),
);
let weights = Tensor::new(weights_data, true);
Self {
weights,
vocab_size,
embedding_dim,
}
}
}
impl Module for Embedding {
fn forward(&self, inputs: &Tensor) -> Result<Tensor> {
inputs.embedding(&self.weights)
}
fn parameters(&self) -> Vec<Tensor> {
vec![self.weights.clone()]
}
}