use crate::error::Error;
use crate::neural_network::Tensor;
use crate::neural_network::layers::TrainingParameters;
use crate::neural_network::layers::activation::format_shape;
use crate::neural_network::layers::layer_weight::{LayerWeight, PReLULayerWeight};
use crate::neural_network::layers::validation::validate_weight_shape;
use crate::neural_network::traits::{Layer, ParamGrad};
use crate::parallel_gates::cheap_map_parallel_threshold;
use ndarray::{ArrayD, Axis, Zip};
use std::borrow::Cow;
#[derive(Debug)]
pub struct PReLU {
input_shape: Vec<usize>,
shared_axes: Vec<usize>,
alpha_init: f32,
alpha: ArrayD<f32>,
input_cache: Option<Tensor>,
grad_alpha: Option<ArrayD<f32>>,
}
impl PReLU {
pub fn new(input_shape: Vec<usize>, alpha: f32) -> Result<Self, Error> {
if input_shape.len() < 2 {
return Err(Error::invalid_input(format!(
"PReLU layer expects an input_shape of rank 2 or more, with the batch axis \
first, got rank {}",
input_shape.len()
)));
}
if let Some(axis) = input_shape.iter().position(|&extent| extent == 0) {
return Err(Error::invalid_input(format!(
"PReLU layer expects every dimension of input_shape to be 1 or more: axis \
{axis} has extent 0"
)));
}
if !alpha.is_finite() {
return Err(Error::invalid_parameter(
"alpha",
"must be finite, because a non-finite slope makes every negative element \
non-finite",
));
}
let alpha_array = ArrayD::from_elem(input_shape[1..].to_vec(), alpha);
Ok(Self {
input_shape,
shared_axes: Vec::new(),
alpha_init: alpha,
alpha: alpha_array,
input_cache: None,
grad_alpha: None,
})
}
pub fn with_shared_axes(mut self, shared_axes: Vec<usize>) -> Result<Self, Error> {
let rank = self.input_shape.len();
let mut sorted = shared_axes;
sorted.sort_unstable();
for (position, &axis) in sorted.iter().enumerate() {
if axis == 0 {
return Err(Error::invalid_parameter(
"shared_axes",
"holds axis 0, which is the batch axis. Every slope is already shared over \
the batch",
));
}
if axis >= rank {
return Err(Error::invalid_parameter(
"shared_axes",
format!("holds axis {axis}, and the input has rank {rank}"),
));
}
if position > 0 && sorted[position - 1] == axis {
return Err(Error::invalid_parameter(
"shared_axes",
format!("holds axis {axis} more than 1 time"),
));
}
}
let mut param_shape = self.input_shape[1..].to_vec();
for &axis in &sorted {
param_shape[axis - 1] = 1;
}
self.alpha = ArrayD::from_elem(param_shape, self.alpha_init);
self.grad_alpha = None;
self.shared_axes = sorted;
Ok(self)
}
pub fn set_weights(&mut self, alpha: ArrayD<f32>) -> Result<(), Error> {
validate_weight_shape("alpha", self.alpha.shape(), alpha.shape())?;
self.alpha = alpha.as_standard_layout().into_owned();
Ok(())
}
fn validate_input(&self, input: &Tensor) -> Result<(), Error> {
if input.is_empty() {
return Err(Error::empty_input("input tensor"));
}
if input.ndim() != self.input_shape.len() {
return Err(Error::invalid_input(format!(
"PReLU layer expects an input of rank {}, got rank {}",
self.input_shape.len(),
input.ndim()
)));
}
for axis in 1..self.input_shape.len() {
if self.shared_axes.contains(&axis) {
continue;
}
if input.shape()[axis] != self.input_shape[axis] {
return Err(Error::invalid_input(format!(
"PReLU layer holds 1 slope per position of axis {axis}, so that axis must \
have extent {}, got {}. Add the axis to shared_axes to accept any extent",
self.input_shape[axis],
input.shape()[axis]
)));
}
}
Ok(())
}
fn activate(&self, input: &Tensor) -> Result<Tensor, Error> {
self.validate_input(input)?;
let slopes = self
.alpha
.broadcast(input.raw_dim())
.expect("validate_input accepts only a shape the slope array covers");
let mut output = Tensor::zeros(input.raw_dim());
let p_relu = |out: &mut f32, &x: &f32, &a: &f32| {
*out = if x >= 0.0 { x } else { a * x };
};
if input.len() >= cheap_map_parallel_threshold() {
Zip::from(&mut output)
.and(input)
.and(&slopes)
.par_for_each(p_relu);
} else {
Zip::from(&mut output)
.and(input)
.and(&slopes)
.for_each(p_relu);
}
Ok(output)
}
}
impl Layer for PReLU {
fn forward(&mut self, input: &Tensor) -> Result<Tensor, Error> {
let output = self.activate(input)?;
self.input_cache = Some(input.clone());
Ok(output)
}
fn predict(&self, input: &Tensor) -> Result<Tensor, Error> {
self.activate(input)
}
fn backward(&mut self, grad_output: &Tensor) -> Result<Tensor, Error> {
let Self {
shared_axes,
alpha,
input_cache,
grad_alpha,
..
} = self;
let Some(input) = input_cache.as_ref() else {
return Err(Error::forward_pass_not_run("PReLU"));
};
if grad_output.shape() != input.shape() {
return Err(Error::shape_mismatch(input.shape(), grad_output.shape()));
}
let slopes = alpha
.broadcast(input.raw_dim())
.expect("the cached input has a shape the slope array covers");
let mut grad_input = Tensor::zeros(input.raw_dim());
let mut contribution = Tensor::zeros(input.raw_dim());
let split = |dx: &mut f32, share: &mut f32, &x: &f32, &g: &f32, &a: &f32| {
if x > 0.0 {
*dx = g;
} else if x < 0.0 {
*dx = g * a;
*share = g * x;
}
};
if input.len() >= cheap_map_parallel_threshold() {
Zip::from(&mut grad_input)
.and(&mut contribution)
.and(input)
.and(grad_output)
.and(&slopes)
.par_for_each(split);
} else {
Zip::from(&mut grad_input)
.and(&mut contribution)
.and(input)
.and(grad_output)
.and(&slopes)
.for_each(split);
}
let mut reduced = contribution.sum_axis(Axis(0));
for &axis in shared_axes.iter() {
let target = Axis(axis - 1);
reduced = reduced.sum_axis(target).insert_axis(target);
}
let grad = grad_alpha.get_or_insert_with(|| ArrayD::zeros(alpha.raw_dim()));
grad.assign(&reduced);
Ok(grad_input)
}
fn layer_type(&self) -> &str {
"PReLU"
}
fn output_shape(&self) -> String {
match &self.input_cache {
Some(input) => format_shape(input.shape()),
None => format_shape(&self.input_shape),
}
}
fn param_count(&self) -> TrainingParameters {
TrainingParameters::Trainable(self.alpha.len())
}
fn parameters(&mut self) -> Vec<ParamGrad<'_>> {
let Self {
alpha, grad_alpha, ..
} = self;
let mut params = Vec::new();
if let Some(grad) = grad_alpha.as_ref() {
params.push(ParamGrad::no_decay(
alpha
.as_slice_mut()
.expect("the slopes are kept in C order"),
grad.as_slice()
.expect("the gradient buffer is kept in C order"),
));
}
params
}
fn get_weights(&self) -> LayerWeight<'_> {
LayerWeight::PReLU(PReLULayerWeight {
alpha: Cow::Borrowed(&self.alpha),
})
}
}