use crate::error::Error;
use crate::neural_network::Tensor;
use ndarray::{IxDyn, Slice};
fn padded_shape(input_shape: &[usize], borders: &[(usize, usize)]) -> Vec<usize> {
let mut shape = input_shape.to_vec();
for (spatial, &(before, after)) in borders.iter().enumerate() {
shape[spatial + 1] += before + after;
}
shape
}
fn cropped_shape(input_shape: &[usize], borders: &[(usize, usize)]) -> Vec<usize> {
let mut shape = input_shape.to_vec();
for (spatial, &(before, after)) in borders.iter().enumerate() {
shape[spatial + 1] -= before + after;
}
shape
}
fn pad_into(input: &Tensor, borders: &[(usize, usize)]) -> Tensor {
let input_shape = input.shape();
let mut output = Tensor::zeros(IxDyn(&padded_shape(input_shape, borders)));
output
.slice_each_axis_mut(|ax| {
let axis = ax.axis.index();
match axis.checked_sub(1).and_then(|spatial| borders.get(spatial)) {
Some(&(before, _)) => Slice::from(before..before + input_shape[axis]),
None => Slice::from(..),
}
})
.assign(input);
output
}
fn crop_out(input: &Tensor, borders: &[(usize, usize)]) -> Tensor {
let input_shape = input.shape().to_vec();
let interior = input.slice_each_axis(|ax| {
let axis = ax.axis.index();
match axis.checked_sub(1).and_then(|spatial| borders.get(spatial)) {
Some(&(before, after)) => Slice::from(before..input_shape[axis] - after),
None => Slice::from(..),
}
});
let mut output = Tensor::zeros(interior.raw_dim());
output.assign(&interior);
output
}
fn validate_input(input: &Tensor, rank: usize, layer: &'static str) -> Result<(), Error> {
if input.ndim() != rank {
return Err(Error::invalid_input(format!(
"{} layer expects a {}D input, got a {}D tensor",
layer,
rank,
input.ndim()
)));
}
if input.is_empty() {
return Err(Error::empty_input("input tensor"));
}
Ok(())
}
fn validate_crop_fits(
input_shape: &[usize],
borders: &[(usize, usize)],
layer: &'static str,
) -> Result<(), Error> {
for (spatial, &(before, after)) in borders.iter().enumerate() {
let axis = spatial + 1;
let extent = input_shape[axis];
if before + after >= extent {
return Err(Error::invalid_input(format!(
"{} layer removes {} of the {} positions on axis {}, and at least 1 must remain",
layer,
before + after,
extent,
axis
)));
}
}
Ok(())
}
fn format_shape(shape: Option<Vec<usize>>) -> String {
match shape {
Some(shape) => {
let axes: Vec<String> = shape[1..].iter().map(|e| e.to_string()).collect();
format!("(None, {})", axes.join(", "))
}
None => "Unknown".to_string(),
}
}
pub(super) fn pad_forward(
input: &Tensor,
borders: &[(usize, usize)],
rank: usize,
layer: &'static str,
) -> Result<Tensor, Error> {
validate_input(input, rank, layer)?;
Ok(pad_into(input, borders))
}
pub(super) fn pad_backward(
grad_output: &Tensor,
input_shape: Option<&[usize]>,
borders: &[(usize, usize)],
layer: &'static str,
) -> Result<Tensor, Error> {
let Some(input_shape) = input_shape else {
return Err(Error::forward_pass_not_run(layer));
};
let expected = padded_shape(input_shape, borders);
if grad_output.shape() != expected.as_slice() {
return Err(Error::shape_mismatch(expected, grad_output.shape()));
}
Ok(crop_out(grad_output, borders))
}
pub(super) fn pad_summary(input_shape: Option<&[usize]>, borders: &[(usize, usize)]) -> String {
format_shape(input_shape.map(|shape| padded_shape(shape, borders)))
}
pub(super) fn crop_forward(
input: &Tensor,
borders: &[(usize, usize)],
rank: usize,
layer: &'static str,
) -> Result<Tensor, Error> {
validate_input(input, rank, layer)?;
validate_crop_fits(input.shape(), borders, layer)?;
Ok(crop_out(input, borders))
}
pub(super) fn crop_backward(
grad_output: &Tensor,
input_shape: Option<&[usize]>,
borders: &[(usize, usize)],
layer: &'static str,
) -> Result<Tensor, Error> {
let Some(input_shape) = input_shape else {
return Err(Error::forward_pass_not_run(layer));
};
let expected = cropped_shape(input_shape, borders);
if grad_output.shape() != expected.as_slice() {
return Err(Error::shape_mismatch(expected, grad_output.shape()));
}
Ok(pad_into(grad_output, borders))
}
pub(super) fn crop_summary(input_shape: Option<&[usize]>, borders: &[(usize, usize)]) -> String {
format_shape(input_shape.map(|shape| cropped_shape(shape, borders)))
}