use crate::Sigma;
pub trait SeparableWeights: Sized {
fn h_weights(&self) -> &[f32];
fn h_anchor(&self) -> usize;
fn v_weights(&self) -> &[f32];
fn v_anchor(&self) -> usize;
#[must_use]
fn flipped(&self) -> Self;
}
#[derive(Clone, Debug)]
pub struct SeparableKernel<const HK: usize, const VK: usize> {
h_weights: [f32; HK],
h_anchor: usize,
v_weights: [f32; VK],
v_anchor: usize,
}
impl<const HK: usize, const VK: usize> SeparableKernel<HK, VK> {
const _ASSERT_NONZERO: () = {
assert!(
HK > 0,
"SeparableKernel: horizontal kernel length HK must be > 0"
);
assert!(
VK > 0,
"SeparableKernel: vertical kernel length VK must be > 0"
);
};
pub fn new(h_weights: [f32; HK], v_weights: [f32; VK]) -> Self {
let () = Self::_ASSERT_NONZERO;
Self {
h_weights,
h_anchor: HK / 2,
v_weights,
v_anchor: VK / 2,
}
}
pub fn with_anchors(
h_weights: [f32; HK],
h_anchor: usize,
v_weights: [f32; VK],
v_anchor: usize,
) -> Self {
let () = Self::_ASSERT_NONZERO;
assert!(
h_anchor < HK,
"h_anchor ({h_anchor}) out of bounds for horizontal kernel of size {HK}"
);
assert!(
v_anchor < VK,
"v_anchor ({v_anchor}) out of bounds for vertical kernel of size {VK}"
);
Self {
h_weights,
h_anchor,
v_weights,
v_anchor,
}
}
pub fn h_weights(&self) -> &[f32; HK] {
&self.h_weights
}
pub fn v_weights(&self) -> &[f32; VK] {
&self.v_weights
}
pub fn h_anchor(&self) -> usize {
self.h_anchor
}
pub fn v_anchor(&self) -> usize {
self.v_anchor
}
pub fn flipped(&self) -> Self {
let mut h = self.h_weights;
h.reverse();
let mut v = self.v_weights;
v.reverse();
Self {
h_weights: h,
h_anchor: HK - 1 - self.h_anchor,
v_weights: v,
v_anchor: VK - 1 - self.v_anchor,
}
}
}
impl<const HK: usize, const VK: usize> SeparableWeights for SeparableKernel<HK, VK> {
#[inline]
fn h_weights(&self) -> &[f32] {
&self.h_weights
}
#[inline]
fn h_anchor(&self) -> usize {
self.h_anchor
}
#[inline]
fn v_weights(&self) -> &[f32] {
&self.v_weights
}
#[inline]
fn v_anchor(&self) -> usize {
self.v_anchor
}
#[inline]
fn flipped(&self) -> Self {
SeparableKernel::flipped(self)
}
}
impl<const K: usize> SeparableKernel<K, K> {
pub fn symmetric(weights: [f32; K]) -> Self {
let () = Self::_ASSERT_NONZERO;
Self {
h_weights: weights,
h_anchor: K / 2,
v_weights: weights,
v_anchor: K / 2,
}
}
}
impl SeparableKernel<3, 3> {
pub fn gaussian_3() -> Self {
Self::symmetric([0.25, 0.5, 0.25])
}
pub fn box_blur_3() -> Self {
let w = 1.0 / 3.0;
Self::symmetric([w, w, w])
}
}
impl SeparableKernel<5, 5> {
pub fn gaussian_5() -> Self {
Self::symmetric([0.0625, 0.25, 0.375, 0.25, 0.0625])
}
pub fn box_blur_5() -> Self {
let w = 1.0 / 5.0;
Self::symmetric([w, w, w, w, w])
}
}
pub const MAX_RADIUS: usize = 64;
const MAX_GAUSSIAN_TAPS: usize = 2 * MAX_RADIUS + 1;
#[derive(Clone)]
pub struct GaussianKernel1D {
weights: [f32; MAX_GAUSSIAN_TAPS],
len: usize,
anchor: usize,
}
impl GaussianKernel1D {
pub fn weights(&self) -> &[f32] {
&self.weights[..self.len]
}
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn anchor(&self) -> usize {
self.anchor
}
pub fn radius(&self) -> usize {
self.anchor
}
}
impl SeparableWeights for GaussianKernel1D {
#[inline]
fn h_weights(&self) -> &[f32] {
self.weights()
}
#[inline]
fn h_anchor(&self) -> usize {
self.anchor
}
#[inline]
fn v_weights(&self) -> &[f32] {
self.weights()
}
#[inline]
fn v_anchor(&self) -> usize {
self.anchor
}
#[inline]
fn flipped(&self) -> Self {
self.clone()
}
}
impl core::fmt::Debug for GaussianKernel1D {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("GaussianKernel1D")
.field("len", &self.len)
.field("anchor", &self.anchor)
.field("weights", &self.weights())
.finish()
}
}
fn gaussian_radius(sigma: Sigma, truncate: f32) -> usize {
assert!(
truncate > 0.0,
"gaussian kernel: truncate must be > 0.0 (got {truncate})"
);
(truncate * sigma.get() + 0.5).floor() as usize
}
#[must_use]
pub fn gaussian_kernel_size(sigma: Sigma, truncate: f32) -> usize {
2 * gaussian_radius(sigma, truncate) + 1
}
#[must_use]
pub fn gaussian_kernel_1d(sigma: Sigma, truncate: f32) -> GaussianKernel1D {
let radius = gaussian_radius(sigma, truncate);
let sigma = sigma.get();
assert!(
radius <= MAX_RADIUS,
"gaussian kernel: sigma {sigma} (truncate {truncate}) needs radius {radius}, \
which exceeds MAX_RADIUS ({MAX_RADIUS}); the largest supported sigma is {}",
MAX_RADIUS as f32 / truncate
);
let n = 2 * radius + 1;
let mut weights = [0.0f32; MAX_GAUSSIAN_TAPS];
let inv_two_sigma_sq = 1.0 / (2.0 * sigma * sigma);
let mut sum = 0.0f32;
for (i, w) in weights[..n].iter_mut().enumerate() {
let d = i as f32 - radius as f32;
let value = (-d * d * inv_two_sigma_sq).exp();
*w = value;
sum += value;
}
let inv_sum = 1.0 / sum;
for w in weights[..n].iter_mut() {
*w *= inv_sum;
}
GaussianKernel1D {
weights,
len: n,
anchor: radius,
}
}
impl<const HK: usize, const VK: usize> PartialEq for SeparableKernel<HK, VK> {
fn eq(&self, other: &Self) -> bool {
self.h_anchor == other.h_anchor
&& self.v_anchor == other.v_anchor
&& self.h_weights == other.h_weights
&& self.v_weights == other.v_weights
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::image::ImageView;
use crate::sigma;
#[test]
fn new_centered_anchors() {
let k = SeparableKernel::new([1.0, 2.0, 1.0], [1.0, 4.0, 6.0, 4.0, 1.0]);
assert_eq!(k.h_anchor(), 1); assert_eq!(k.v_anchor(), 2); assert_eq!(k.h_weights(), &[1.0, 2.0, 1.0]);
assert_eq!(k.v_weights(), &[1.0, 4.0, 6.0, 4.0, 1.0]);
}
#[test]
fn new_even_sizes_center_left() {
let k = SeparableKernel::new([1.0; 4], [1.0; 2]);
assert_eq!(k.h_anchor(), 2); assert_eq!(k.v_anchor(), 1); }
#[test]
fn with_anchors_explicit() {
let k = SeparableKernel::with_anchors([1.0, 2.0, 3.0], 0, [4.0, 5.0], 1);
assert_eq!(k.h_anchor(), 0);
assert_eq!(k.v_anchor(), 1);
assert_eq!(k.h_weights(), &[1.0, 2.0, 3.0]);
assert_eq!(k.v_weights(), &[4.0, 5.0]);
}
#[test]
#[should_panic(expected = "h_anchor")]
fn with_anchors_h_out_of_bounds() {
SeparableKernel::with_anchors([1.0, 2.0, 3.0], 3, [1.0], 0);
}
#[test]
#[should_panic(expected = "v_anchor")]
fn with_anchors_v_out_of_bounds() {
SeparableKernel::with_anchors([1.0], 0, [1.0, 2.0], 2);
}
#[test]
fn symmetric_constructor() {
let k = SeparableKernel::symmetric([1.0, 2.0, 1.0]);
assert_eq!(k.h_weights(), k.v_weights());
assert_eq!(k.h_anchor(), k.v_anchor());
assert_eq!(k.h_anchor(), 1);
}
#[test]
fn flipped_reverses_weights() {
let k = SeparableKernel::new([1.0, 2.0, 3.0], [4.0, 5.0]);
let f = k.flipped();
assert_eq!(f.h_weights(), &[3.0, 2.0, 1.0]);
assert_eq!(f.v_weights(), &[5.0, 4.0]);
}
#[test]
fn flipped_mirrors_anchors() {
let k = SeparableKernel::with_anchors([1.0, 2.0, 3.0], 0, [4.0, 5.0, 6.0], 2);
let f = k.flipped();
assert_eq!(f.h_anchor(), 2); assert_eq!(f.v_anchor(), 0); }
#[test]
fn flipped_centered_anchor_stays_centered() {
let k = SeparableKernel::gaussian_3();
let f = k.flipped();
assert_eq!(f.h_anchor(), 1);
assert_eq!(f.v_anchor(), 1);
}
#[test]
fn flipped_involution() {
let k = SeparableKernel::with_anchors([1.0, 2.0, 3.0], 0, [4.0, 5.0], 1);
let ff = k.flipped().flipped();
assert_eq!(k, ff);
}
#[test]
fn flipped_symmetric_kernel_unchanged() {
let k = SeparableKernel::box_blur_3();
let f = k.flipped();
assert_eq!(k, f);
}
#[test]
fn flipped_gaussian_5_symmetric() {
let k = SeparableKernel::gaussian_5();
let f = k.flipped();
assert_eq!(k, f);
}
#[test]
fn gaussian_3_weights() {
let k = SeparableKernel::gaussian_3();
assert_eq!(k.h_weights(), &[0.25, 0.5, 0.25]);
assert_eq!(k.v_weights(), &[0.25, 0.5, 0.25]);
assert_eq!(k.h_anchor(), 1);
assert_eq!(k.v_anchor(), 1);
assert!((k.h_weights().iter().sum::<f32>() - 1.0).abs() < 1e-7);
}
#[test]
fn gaussian_5_weights() {
let k = SeparableKernel::gaussian_5();
assert_eq!(k.h_weights(), &[0.0625, 0.25, 0.375, 0.25, 0.0625]);
assert_eq!(k.v_weights(), &[0.0625, 0.25, 0.375, 0.25, 0.0625]);
assert_eq!(k.h_anchor(), 2);
assert_eq!(k.v_anchor(), 2);
assert!((k.h_weights().iter().sum::<f32>() - 1.0).abs() < 1e-7);
}
#[test]
fn box_blur_3_weights() {
let k = SeparableKernel::box_blur_3();
let third = 1.0f32 / 3.0;
for &w in k.h_weights() {
assert!((w - third).abs() < 1e-7);
}
for &w in k.v_weights() {
assert!((w - third).abs() < 1e-7);
}
}
#[test]
fn box_blur_5_weights() {
let k = SeparableKernel::box_blur_5();
let fifth = 1.0f32 / 5.0;
for &w in k.h_weights() {
assert!((w - fifth).abs() < 1e-7);
}
for &w in k.v_weights() {
assert!((w - fifth).abs() < 1e-7);
}
}
#[test]
fn gaussian_3_outer_product_matches_neighborhood() {
let sep = SeparableKernel::gaussian_3();
let full = crate::image::Neighborhood::<f32, 3, 3>::gaussian_3x3();
for y in 0..3 {
for x in 0..3 {
let outer = sep.h_weights()[x] * sep.v_weights()[y];
let expected = full.weights().pixel_at(x, y) / 16.0;
assert!(
(outer - expected).abs() < 1e-6,
"mismatch at ({x}, {y}): outer={outer}, expected={expected}"
);
}
}
}
#[test]
fn gaussian_5_outer_product_matches_neighborhood() {
let sep = SeparableKernel::gaussian_5();
let full = crate::image::Neighborhood::<f32, 5, 5>::gaussian_5x5();
for y in 0..5 {
for x in 0..5 {
let outer = sep.h_weights()[x] * sep.v_weights()[y];
let expected = full.weights().pixel_at(x, y) / 256.0;
assert!(
(outer - expected).abs() < 1e-6,
"mismatch at ({x}, {y}): outer={outer}, expected={expected}"
);
}
}
}
#[test]
fn box_blur_3_outer_product_matches_neighborhood() {
let sep = SeparableKernel::box_blur_3();
let full = crate::image::Neighborhood::<f32, 3, 3>::box_blur_3x3();
for y in 0..3 {
for x in 0..3 {
let outer = sep.h_weights()[x] * sep.v_weights()[y];
let expected = full.weights().pixel_at(x, y);
assert!(
(outer - expected).abs() < 1e-6,
"mismatch at ({x}, {y}): outer={outer}, expected={expected}"
);
}
}
}
#[test]
fn box_blur_5_outer_product_matches_neighborhood() {
let sep = SeparableKernel::box_blur_5();
let full = crate::image::Neighborhood::<f32, 5, 5>::box_blur_5x5();
for y in 0..5 {
for x in 0..5 {
let outer = sep.h_weights()[x] * sep.v_weights()[y];
let expected = full.weights().pixel_at(x, y);
assert!(
(outer - expected).abs() < 1e-6,
"mismatch at ({x}, {y}): outer={outer}, expected={expected}"
);
}
}
}
#[test]
fn clone_produces_equal_kernel() {
let k = SeparableKernel::with_anchors([1.0, 2.0, 3.0], 0, [4.0, 5.0], 1);
let c = k.clone();
assert_eq!(k, c);
}
#[test]
fn debug_format_contains_weights() {
let k = SeparableKernel::new([1.0, 2.0], [3.0]);
let dbg = format!("{k:?}");
assert!(dbg.contains("SeparableKernel"));
assert!(dbg.contains("h_weights"));
assert!(dbg.contains("v_weights"));
}
#[test]
fn partial_eq_different_weights() {
let a = SeparableKernel::new([1.0, 2.0, 3.0], [1.0]);
let b = SeparableKernel::new([3.0, 2.0, 1.0], [1.0]);
assert_ne!(a, b);
}
#[test]
fn partial_eq_different_anchors() {
let a = SeparableKernel::with_anchors([1.0, 2.0, 3.0], 0, [1.0], 0);
let b = SeparableKernel::with_anchors([1.0, 2.0, 3.0], 2, [1.0], 0);
assert_ne!(a, b);
}
#[test]
fn asymmetric_3x5() {
let k = SeparableKernel::new([1.0, 2.0, 1.0], [1.0, 4.0, 6.0, 4.0, 1.0]);
assert_eq!(k.h_anchor(), 1);
assert_eq!(k.v_anchor(), 2);
let f = k.flipped();
assert_eq!(f.h_weights(), &[1.0, 2.0, 1.0]);
assert_eq!(f.v_weights(), &[1.0, 4.0, 6.0, 4.0, 1.0]);
}
#[test]
fn asymmetric_weights_flip() {
let k = SeparableKernel::new([1.0, 0.0, 0.0], [0.0, 1.0]);
let f = k.flipped();
assert_eq!(f.h_weights(), &[0.0, 0.0, 1.0]);
assert_eq!(f.v_weights(), &[1.0, 0.0]);
}
#[test]
fn identity_1x1() {
let k = SeparableKernel::new([1.0], [1.0]);
assert_eq!(k.h_anchor(), 0);
assert_eq!(k.v_anchor(), 0);
let f = k.flipped();
assert_eq!(f, k);
}
#[test]
fn gaussian_kernel_weights_sum_to_one() {
for &sigma in &[0.5f32, 0.8, 1.0, 1.7, 3.0, 8.0] {
let k = gaussian_kernel_1d(Sigma::new(sigma).unwrap(), 4.0);
let sum: f32 = k.weights().iter().sum();
assert!(
(sum - 1.0).abs() < 1e-6,
"sigma {sigma}: weights sum to {sum}, expected 1.0"
);
}
}
#[test]
fn gaussian_kernel_weights_symmetric() {
let k = gaussian_kernel_1d(sigma!(1.5), 4.0);
let w = k.weights();
let n = w.len();
for i in 0..n {
assert!(
(w[i] - w[n - 1 - i]).abs() < 1e-7,
"asymmetry at tap {i}: {} vs {}",
w[i],
w[n - 1 - i]
);
}
}
#[test]
fn gaussian_kernel_matches_reference_formula() {
let sigma = 1.3f32;
let truncate = 4.0f32;
let k = gaussian_kernel_1d(Sigma::new(sigma).unwrap(), truncate);
let radius = k.radius();
let n = k.len();
let mut reference = vec![0.0f32; n];
let mut sum = 0.0f32;
for (i, r) in reference.iter_mut().enumerate() {
let d = i as f32 - radius as f32;
*r = (-(d * d) / (2.0 * sigma * sigma)).exp();
sum += *r;
}
for r in &mut reference {
*r /= sum;
}
for (i, (&got, &want)) in k.weights().iter().zip(&reference).enumerate() {
assert!((got - want).abs() < 1e-6, "tap {i}: got {got}, want {want}");
}
}
#[test]
fn gaussian_kernel_size_follows_truncate() {
assert_eq!(gaussian_kernel_size(sigma!(1.0), 4.0), 9); assert_eq!(gaussian_kernel_size(sigma!(2.0), 3.0), 13); assert_eq!(gaussian_kernel_size(sigma!(1.0), 3.0), 7); assert_eq!(
gaussian_kernel_1d(sigma!(1.0), 4.0).len(),
gaussian_kernel_size(sigma!(1.0), 4.0)
);
}
#[test]
fn gaussian_kernel_tiny_sigma_is_identity() {
let k = gaussian_kernel_1d(sigma!(0.05), 4.0);
assert_eq!(k.len(), 1);
assert_eq!(k.radius(), 0);
assert_eq!(k.anchor(), 0);
assert!((k.weights()[0] - 1.0).abs() < 1e-7);
}
#[test]
#[should_panic(expected = "truncate must be > 0.0")]
fn gaussian_kernel_zero_truncate_panics() {
let _ = gaussian_kernel_1d(sigma!(1.0), 0.0);
}
#[test]
#[should_panic(expected = "exceeds MAX_RADIUS")]
fn gaussian_kernel_over_radius_panics() {
let _ = gaussian_kernel_1d(sigma!(20.0), 4.0);
}
#[test]
fn gaussian_kernel_size_reports_over_max_radius_without_panicking() {
assert_eq!(gaussian_kernel_size(sigma!(20.0), 4.0), 161);
assert!(gaussian_kernel_size(sigma!(20.0), 4.0) > 2 * MAX_RADIUS + 1);
assert_eq!(gaussian_kernel_size(sigma!(16.0), 4.0), 2 * MAX_RADIUS + 1);
}
#[test]
fn gaussian_kernel_at_max_radius_is_ok() {
let k = gaussian_kernel_1d(sigma!(16.0), 4.0);
assert_eq!(k.radius(), MAX_RADIUS);
assert_eq!(k.len(), 2 * MAX_RADIUS + 1);
let sum: f32 = k.weights().iter().sum();
assert!((sum - 1.0).abs() < 1e-5);
}
#[test]
fn gaussian_kernel_is_its_own_flip() {
for sigma in [0.05f32, 0.8, 1.5, 4.0] {
for truncate in [3.0f32, 4.0] {
let k = gaussian_kernel_1d(Sigma::new(sigma).unwrap(), truncate);
let w = k.weights();
let n = w.len();
for i in 0..n {
assert!(
(w[i] - w[n - 1 - i]).abs() < f32::EPSILON,
"σ={sigma} t={truncate}: tap {i} != tap {}",
n - 1 - i,
);
}
assert_eq!(
k.anchor(),
n - 1 - k.anchor(),
"σ={sigma}: anchor off-centre"
);
let f = SeparableWeights::flipped(&k);
assert_eq!(f.h_weights(), k.h_weights());
assert_eq!(f.v_weights(), k.v_weights());
assert_eq!(f.h_anchor(), k.h_anchor());
assert_eq!(f.v_anchor(), k.v_anchor());
}
}
}
#[test]
fn gaussian_kernel_reports_the_same_taps_on_both_axes() {
let k = gaussian_kernel_1d(sigma!(1.5), 4.0);
assert_eq!(k.h_weights(), k.v_weights());
assert_eq!(k.h_anchor(), k.v_anchor());
assert_eq!(k.h_weights(), k.weights());
assert_eq!(k.h_anchor(), k.anchor());
}
#[test]
fn separable_kernel_trait_form_matches_its_inherent_methods() {
let k = SeparableKernel::with_anchors([1.0, 2.0, 3.0], 0, [4.0, 5.0], 1);
assert_eq!(SeparableWeights::h_weights(&k), &k.h_weights()[..]);
assert_eq!(SeparableWeights::v_weights(&k), &k.v_weights()[..]);
assert_eq!(SeparableWeights::h_anchor(&k), k.h_anchor());
assert_eq!(SeparableWeights::v_anchor(&k), k.v_anchor());
let inherent = k.flipped();
let via_trait = SeparableWeights::flipped(&k);
assert_eq!(
SeparableWeights::h_weights(&via_trait),
&inherent.h_weights()[..]
);
assert_eq!(
SeparableWeights::v_weights(&via_trait),
&inherent.v_weights()[..]
);
assert_eq!(via_trait.h_anchor(), inherent.h_anchor());
assert_eq!(via_trait.v_anchor(), inherent.v_anchor());
}
#[test]
fn both_kernel_types_satisfy_the_trait_contract() {
fn axes<K: SeparableWeights>(k: &K) -> (usize, usize, usize, usize) {
assert!(!k.h_weights().is_empty(), "h axis must be non-empty");
assert!(!k.v_weights().is_empty(), "v axis must be non-empty");
assert!(k.h_anchor() < k.h_weights().len(), "h anchor out of bounds");
assert!(k.v_anchor() < k.v_weights().len(), "v anchor out of bounds");
(
k.h_weights().len(),
k.h_anchor(),
k.v_weights().len(),
k.v_anchor(),
)
}
assert_eq!(axes(&SeparableKernel::gaussian_5()), (5, 2, 5, 2));
assert_eq!(axes(&SeparableKernel::box_blur_3()), (3, 1, 3, 1));
assert_eq!(axes(&gaussian_kernel_1d(sigma!(1.0), 4.0)), (9, 4, 9, 4));
assert_eq!(axes(&gaussian_kernel_1d(sigma!(0.05), 4.0)), (1, 0, 1, 0));
}
}