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::{
crop_backward, crop_forward, crop_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 Cropping2D {
cropping: Border2D,
input_shape: Option<Vec<usize>>,
}
impl Cropping2D {
pub fn new(cropping: impl Into<Border2D>) -> Self {
Cropping2D {
cropping: cropping.into(),
input_shape: None,
}
}
}
impl Layer for Cropping2D {
fn forward(&mut self, input: &Tensor) -> Result<Tensor, Error> {
let output = crop_forward(input, &self.cropping.0, 4, "Cropping2D")?;
self.input_shape = Some(input.shape().to_vec());
Ok(output)
}
fn predict(&self, input: &Tensor) -> Result<Tensor, Error> {
crop_forward(input, &self.cropping.0, 4, "Cropping2D")
}
fn backward(&mut self, grad_output: &Tensor) -> Result<Tensor, Error> {
crop_backward(
grad_output,
self.input_shape.as_deref(),
&self.cropping.0,
"Cropping2D",
)
}
fn layer_type(&self) -> &str {
"Cropping2D"
}
fn output_shape(&self) -> String {
crop_summary(self.input_shape.as_deref(), &self.cropping.0)
}
no_trainable_parameters_layer_functions!();
}