use smallvec::{SmallVec, smallvec};
use crate::op::{Map, Op};
use crate::{Element, MapOperation, Tensor};
use super::candidates::Candidate;
use super::pattern::Pattern;
use super::view::View;
#[derive(Debug, Clone)]
pub(crate) struct BatchNormalization {
pub(crate) input: usize,
pub(crate) scale: usize,
pub(crate) shift: usize,
pub(crate) epsilon: usize,
pub(crate) mean: usize,
pub(crate) variance: usize,
}
impl BatchNormalization {
pub(crate) fn reads(&self) -> [usize; 4] {
[self.input, self.scale, self.shift, self.epsilon]
}
pub(crate) fn apply<E: Element>(
&self,
values: &[Tensor<E>],
) -> (Tensor<E>, Tensor<E>, Tensor<E>) {
values[self.input].batch_normalized(
&values[self.scale],
&values[self.shift],
&values[self.epsilon],
)
}
}
struct Tail {
group: BatchNormalization,
interiors: SmallVec<[usize; 8]>,
centered: usize,
}
fn match_tail<E: Element>(index: usize, view: &View<Tensor<E>>) -> Option<Tail> {
let Some(Op::Add(_)) = view.op(index) else {
return None;
};
if view.shape(index).rank() != 2 {
return None;
}
let scaled = view.operand(index, 0);
let shift_bcast = view.operand(index, 1);
let Some(Op::BroadcastAlong(shift_along)) = view.op(shift_bcast) else {
return None;
};
let Some(Op::Mul(_)) = view.op(scaled) else {
return None;
};
let shift = view.sole_operand(shift_bcast);
let normalized = view.operand(scaled, 0);
let scale_bcast = view.operand(scaled, 1);
let Some(Op::BroadcastAlong(scale_along)) = view.op(scale_bcast) else {
return None;
};
if shift_along.axis != 0 || scale_along.axis != 0 {
return None;
}
let scale = view.sole_operand(scale_bcast);
let Some(Op::Div(_)) = view.op(normalized) else {
return None;
};
let centered = view.operand(normalized, 0);
let dev_bcast = view.operand(normalized, 1);
let Some(Op::BroadcastAlong(dev_along)) = view.op(dev_bcast) else {
return None;
};
if dev_along.axis != 0 {
return None;
}
let deviation = view.sole_operand(dev_bcast);
let Some(Op::Map(Map {
op: MapOperation::Sqrt,
})) = view.op(deviation)
else {
return None;
};
let var_plus = view.sole_operand(deviation);
let Some(Op::Add(_)) = view.op(var_plus) else {
return None;
};
let variance = view.operand(var_plus, 0);
let eps_bcast = view.operand(var_plus, 1);
let Some(Op::Broadcast(_)) = view.op(eps_bcast) else {
return None;
};
let epsilon = view.sole_operand(eps_bcast);
let Some(Op::Leaf(_)) = view.op(epsilon) else {
return None;
};
if view.shape(epsilon).volume() != 1 {
return None;
}
let Some(Op::Sub(_)) = view.op(centered) else {
return None;
};
let input = view.operand(centered, 0);
let mean_bcast = view.operand(centered, 1);
let Some(Op::BroadcastAlong(mean_along)) = view.op(mean_bcast) else {
return None;
};
if mean_along.axis != 0 {
return None;
}
let mean = view.sole_operand(mean_bcast);
Some(Tail {
group: BatchNormalization {
input,
scale,
shift,
epsilon,
mean,
variance,
},
interiors: smallvec![
scaled,
shift_bcast,
scale_bcast,
normalized,
dev_bcast,
deviation,
var_plus,
eps_bcast,
centered,
mean_bcast,
],
centered,
})
}
fn mean_along_of<E: Element>(node: usize, view: &View<Tensor<E>>) -> Option<(usize, usize, usize)> {
let Some(Op::Div(_)) = view.op(node) else {
return None;
};
let sum = view.operand(node, 0);
let count = view.operand(node, 1);
let Some(Op::SumAlong(along)) = view.op(sum) else {
return None;
};
if along.axis != 0 {
return None;
}
let Some(Op::Leaf(leaf)) = view.op(count) else {
return None;
};
let source = view.sole_operand(sum);
let batch = view.shape(source).axes()[0];
if !leaf.0.is_counted(view.shape(node), batch) {
return None;
}
Some((source, sum, count))
}
pub(crate) fn match_training<E: Element>(
index: usize,
view: &View<Tensor<E>>,
) -> Option<Candidate> {
let mut tail = match_tail(index, view)?;
let (mean_source, mean_sum, mean_count) = mean_along_of(tail.group.mean, view)?;
if mean_source != tail.group.input {
return None;
}
let (squared, var_sum, var_count) = mean_along_of(tail.group.variance, view)?;
let Some(Op::Mul(_)) = view.op(squared) else {
return None;
};
if view.operand(squared, 0) != tail.centered || view.operand(squared, 1) != tail.centered {
return None;
}
tail.interiors
.extend_from_slice(&[mean_sum, mean_count, squared, var_sum, var_count]);
let pattern = Pattern::BatchNormTraining(tail.group);
let named = pattern.named();
Some(Candidate {
pattern,
root: index,
interiors: tail.interiors,
named,
})
}
pub(crate) fn match_inference<E: Element>(
index: usize,
view: &View<Tensor<E>>,
) -> Option<Candidate> {
let tail = match_tail(index, view)?;
Some(Candidate {
pattern: Pattern::BatchNormInference(tail.group),
root: index,
interiors: tail.interiors,
named: SmallVec::new(),
})
}
#[cfg(test)]
#[path = "tests/batch_norm_tests.rs"]
mod tests;