use crate::error::Error;
use crate::neural_network::Tensor;
use crate::neural_network::layers::TrainingParameters;
use crate::neural_network::layers::layer_weight::LayerWeight;
use crate::neural_network::layers::no_trainable_parameters_layer_functions;
use crate::neural_network::traits::Layer;
use ndarray::{Axis, IxDyn};
#[derive(Debug)]
pub struct RepeatVector {
n: usize,
input_shape: Option<Vec<usize>>,
}
impl RepeatVector {
pub fn new(n: usize) -> Result<Self, Error> {
if n == 0 {
return Err(Error::invalid_parameter(
"n",
"is 0, and the output must hold at least 1 step",
));
}
Ok(RepeatVector {
n,
input_shape: None,
})
}
fn validate(&self, input: &Tensor) -> Result<(), Error> {
if input.ndim() != 2 {
return Err(Error::invalid_input(format!(
"RepeatVector layer expects a 2D input [batch_size, features], got a {}D tensor",
input.ndim()
)));
}
if input.is_empty() {
return Err(Error::empty_input("input tensor"));
}
Ok(())
}
fn repeat(&self, input: &Tensor) -> Tensor {
let shape = input.shape();
let mut output = Tensor::zeros(IxDyn(&[shape[0], self.n, shape[1]]));
output.assign(&input.view().insert_axis(Axis(1)));
output
}
}
impl Layer for RepeatVector {
fn forward(&mut self, input: &Tensor) -> Result<Tensor, Error> {
self.validate(input)?;
self.input_shape = Some(input.shape().to_vec());
Ok(self.repeat(input))
}
fn predict(&self, input: &Tensor) -> Result<Tensor, Error> {
self.validate(input)?;
Ok(self.repeat(input))
}
fn backward(&mut self, grad_output: &Tensor) -> Result<Tensor, Error> {
let Some(input_shape) = &self.input_shape else {
return Err(Error::forward_pass_not_run("RepeatVector"));
};
let expected = [input_shape[0], self.n, input_shape[1]];
if grad_output.shape() != expected {
return Err(Error::shape_mismatch(expected, grad_output.shape()));
}
Ok(grad_output.sum_axis(Axis(1)))
}
fn layer_type(&self) -> &str {
"RepeatVector"
}
fn output_shape(&self) -> String {
match &self.input_shape {
Some(shape) => format!("(None, {}, {})", self.n, shape[1]),
None => "Unknown".to_string(),
}
}
no_trainable_parameters_layer_functions!();
}