use crate::error::Error;
use crate::neural_network::Tensor;
use crate::neural_network::layers::TrainingParameters;
use crate::neural_network::layers::border::Border2D;
use crate::neural_network::layers::border::pad_crop_engine::{
pad_backward, pad_forward, pad_summary,
};
use crate::neural_network::layers::layer_weight::LayerWeight;
use crate::neural_network::layers::no_trainable_parameters_layer_functions;
use crate::neural_network::traits::Layer;
#[derive(Debug)]
pub struct ZeroPadding2D {
padding: Border2D,
input_shape: Option<Vec<usize>>,
}
impl ZeroPadding2D {
pub fn new(padding: impl Into<Border2D>) -> Self {
ZeroPadding2D {
padding: padding.into(),
input_shape: None,
}
}
}
impl Layer for ZeroPadding2D {
fn forward(&mut self, input: &Tensor) -> Result<Tensor, Error> {
let output = pad_forward(input, &self.padding.0, 4, "ZeroPadding2D")?;
self.input_shape = Some(input.shape().to_vec());
Ok(output)
}
fn predict(&self, input: &Tensor) -> Result<Tensor, Error> {
pad_forward(input, &self.padding.0, 4, "ZeroPadding2D")
}
fn backward(&mut self, grad_output: &Tensor) -> Result<Tensor, Error> {
pad_backward(
grad_output,
self.input_shape.as_deref(),
&self.padding.0,
"ZeroPadding2D",
)
}
fn layer_type(&self) -> &str {
"ZeroPadding2D"
}
fn output_shape(&self) -> String {
pad_summary(self.input_shape.as_deref(), &self.padding.0)
}
no_trainable_parameters_layer_functions!();
}