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::IxDyn;
#[derive(Debug)]
pub struct Permute {
forward_axes: Vec<usize>,
backward_axes: Vec<usize>,
input_shape: Option<Vec<usize>>,
}
impl Permute {
pub fn new(dims: Vec<usize>) -> Result<Self, Error> {
if dims.is_empty() {
return Err(Error::invalid_parameter(
"dims",
"is empty, and a permutation must name at least 1 axis",
));
}
let mut seen = vec![false; dims.len()];
for &axis in &dims {
match axis.checked_sub(1).and_then(|index| seen.get_mut(index)) {
Some(slot) if !*slot => *slot = true,
Some(_) => {
return Err(Error::invalid_parameter(
"dims",
format!("names axis {axis} more than once"),
));
}
None => {
return Err(Error::invalid_parameter(
"dims",
format!(
"holds {}, and every entry must be between 1 and {}",
axis,
dims.len()
),
));
}
}
}
let forward_axes: Vec<usize> = std::iter::once(0).chain(dims).collect();
let mut backward_axes = vec![0; forward_axes.len()];
for (output_axis, &input_axis) in forward_axes.iter().enumerate() {
backward_axes[input_axis] = output_axis;
}
Ok(Permute {
forward_axes,
backward_axes,
input_shape: None,
})
}
fn permuted_shape(&self, input_shape: &[usize]) -> Vec<usize> {
self.forward_axes.iter().map(|&a| input_shape[a]).collect()
}
fn validate(&self, input: &Tensor) -> Result<(), Error> {
if input.ndim() != self.forward_axes.len() {
return Err(Error::invalid_input(format!(
"Permute layer expects a {}D input, got a {}D tensor",
self.forward_axes.len(),
input.ndim()
)));
}
if input.is_empty() {
return Err(Error::empty_input("input tensor"));
}
Ok(())
}
}
fn permute_into(input: &Tensor, axes: &[usize]) -> Tensor {
let view = input.view().permuted_axes(IxDyn(axes));
let mut output = Tensor::zeros(view.raw_dim());
output.assign(&view);
output
}
impl Layer for Permute {
fn forward(&mut self, input: &Tensor) -> Result<Tensor, Error> {
self.validate(input)?;
self.input_shape = Some(input.shape().to_vec());
Ok(permute_into(input, &self.forward_axes))
}
fn predict(&self, input: &Tensor) -> Result<Tensor, Error> {
self.validate(input)?;
Ok(permute_into(input, &self.forward_axes))
}
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("Permute"));
};
let expected = self.permuted_shape(input_shape);
if grad_output.shape() != expected.as_slice() {
return Err(Error::shape_mismatch(expected, grad_output.shape()));
}
Ok(permute_into(grad_output, &self.backward_axes))
}
fn layer_type(&self) -> &str {
"Permute"
}
fn output_shape(&self) -> String {
match &self.input_shape {
Some(shape) => {
let axes: Vec<String> = self.permuted_shape(shape)[1..]
.iter()
.map(|e| e.to_string())
.collect();
format!("(None, {})", axes.join(", "))
}
None => "Unknown".to_string(),
}
}
no_trainable_parameters_layer_functions!();
}