use super::error::Result;
use super::image_view::{ImageView, OwnedImage};
use super::scalar::Scalar;
pub trait SourcePixel: Copy {
fn to_scalar(self) -> Scalar;
}
impl SourcePixel for u8 {
#[inline]
fn to_scalar(self) -> Scalar {
self as Scalar
}
}
impl SourcePixel for u16 {
#[inline]
fn to_scalar(self) -> Scalar {
self as Scalar
}
}
impl SourcePixel for f32 {
#[inline]
fn to_scalar(self) -> Scalar {
self
}
}
pub struct GradientField {
pub(crate) gx: OwnedImage<Scalar>,
pub(crate) gy: OwnedImage<Scalar>,
}
impl GradientField {
#[inline]
pub fn width(&self) -> usize {
self.gx.width()
}
#[inline]
pub fn height(&self) -> usize {
self.gx.height()
}
#[inline]
pub fn gx(&self) -> ImageView<'_, Scalar> {
self.gx.view()
}
#[inline]
pub fn gy(&self) -> ImageView<'_, Scalar> {
self.gy.view()
}
#[inline]
pub fn get(&self, x: usize, y: usize) -> Option<(Scalar, Scalar)> {
let gx = self.gx.get(x, y)?;
let gy = self.gy.get(x, y)?;
Some((gx, gy))
}
#[inline]
pub fn magnitude(&self, x: usize, y: usize) -> Option<Scalar> {
let (gx, gy) = self.get(x, y)?;
Some((gx * gx + gy * gy).sqrt())
}
pub fn max_magnitude(&self) -> Scalar {
let gx = self.gx.data();
let gy = self.gy.data();
gx.iter()
.zip(gy.iter())
.map(|(&x, &y)| x * x + y * y)
.fold(0.0f32, Scalar::max)
.sqrt()
}
}
pub fn sobel_gradient<P: SourcePixel>(image: &ImageView<'_, P>) -> Result<GradientField> {
let w = image.width();
let h = image.height();
let stride = image.stride();
let src = image.as_slice();
let mut gx = OwnedImage::<Scalar>::zeros(w, h)?;
let mut gy = OwnedImage::<Scalar>::zeros(w, h)?;
let gx_data = gx.data_mut();
let gy_data = gy.data_mut();
for y in 1..h - 1 {
let row_prev = (y - 1) * stride;
let row_curr = y * stride;
let row_next = (y + 1) * stride;
for x in 1..w - 1 {
let p00 = src[row_prev + x - 1].to_scalar();
let p10 = src[row_prev + x].to_scalar();
let p20 = src[row_prev + x + 1].to_scalar();
let p01 = src[row_curr + x - 1].to_scalar();
let p21 = src[row_curr + x + 1].to_scalar();
let p02 = src[row_next + x - 1].to_scalar();
let p12 = src[row_next + x].to_scalar();
let p22 = src[row_next + x + 1].to_scalar();
let dx = (-p00 + p20 - 2.0 * p01 + 2.0 * p21 - p02 + p22) / 8.0;
let dy = (-p00 - 2.0 * p10 - p20 + p02 + 2.0 * p12 + p22) / 8.0;
let idx = y * w + x;
gx_data[idx] = dx;
gy_data[idx] = dy;
}
}
Ok(GradientField { gx, gy })
}
pub fn sobel_gradient_f32(image: &ImageView<'_, f32>) -> Result<GradientField> {
sobel_gradient(image)
}
pub fn gradient_magnitude(field: &GradientField) -> Result<OwnedImage<Scalar>> {
let w = field.width();
let h = field.height();
let mut mag = OwnedImage::<Scalar>::zeros(w, h)?;
let mag_data = mag.data_mut();
let gx_data = field.gx.data();
let gy_data = field.gy.data();
for i in 0..w * h {
let gx = gx_data[i];
let gy = gy_data[i];
mag_data[i] = (gx * gx + gy * gy).sqrt();
}
Ok(mag)
}
pub fn thin_gradient(field: &GradientField) -> Result<GradientField> {
let w = field.width();
let h = field.height();
let gx_in = field.gx.data();
let gy_in = field.gy.data();
let mut mag_sq = vec![0.0f32; w * h];
for i in 0..w * h {
mag_sq[i] = gx_in[i] * gx_in[i] + gy_in[i] * gy_in[i];
}
let mut out_gx = OwnedImage::<Scalar>::zeros(w, h)?;
let mut out_gy = OwnedImage::<Scalar>::zeros(w, h)?;
let out_gx_data = out_gx.data_mut();
let out_gy_data = out_gy.data_mut();
const LO: Scalar = std::f32::consts::SQRT_2 - 1.0;
const HI: Scalar = std::f32::consts::SQRT_2 + 1.0;
for y in 1..h.saturating_sub(1) {
for x in 1..w.saturating_sub(1) {
let idx = y * w + x;
let m = mag_sq[idx];
if m == 0.0 {
continue; }
let gx = gx_in[idx];
let gy = gy_in[idx];
let ax = gx.abs();
let ay = gy.abs();
let sx: isize = if gx >= 0.0 { 1 } else { -1 };
let sy: isize = if gy >= 0.0 { 1 } else { -1 };
let (fx, fy): (isize, isize) = if ay < LO * ax {
(sx, 0) } else if ay > HI * ax {
(0, sy) } else {
(sx, sy) };
let fwd = mag_sq[(y as isize + fy) as usize * w + (x as isize + fx) as usize];
let bwd = mag_sq[(y as isize - fy) as usize * w + (x as isize - fx) as usize];
if m >= fwd && m > bwd {
out_gx_data[idx] = gx;
out_gy_data[idx] = gy;
}
}
}
Ok(GradientField {
gx: out_gx,
gy: out_gy,
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[non_exhaustive]
pub enum GradientOperator {
#[default]
Sobel,
Scharr,
}
pub fn scharr_gradient<P: SourcePixel>(image: &ImageView<'_, P>) -> Result<GradientField> {
let w = image.width();
let h = image.height();
let stride = image.stride();
let src = image.as_slice();
let mut gx = OwnedImage::<Scalar>::zeros(w, h)?;
let mut gy = OwnedImage::<Scalar>::zeros(w, h)?;
let gx_data = gx.data_mut();
let gy_data = gy.data_mut();
for y in 1..h - 1 {
let row_prev = (y - 1) * stride;
let row_curr = y * stride;
let row_next = (y + 1) * stride;
for x in 1..w - 1 {
let p00 = src[row_prev + x - 1].to_scalar();
let p10 = src[row_prev + x].to_scalar();
let p20 = src[row_prev + x + 1].to_scalar();
let p01 = src[row_curr + x - 1].to_scalar();
let p21 = src[row_curr + x + 1].to_scalar();
let p02 = src[row_next + x - 1].to_scalar();
let p12 = src[row_next + x].to_scalar();
let p22 = src[row_next + x + 1].to_scalar();
let dx =
(-3.0 * p00 + 3.0 * p20 - 10.0 * p01 + 10.0 * p21 - 3.0 * p02 + 3.0 * p22) / 32.0;
let dy =
(-3.0 * p00 - 10.0 * p10 - 3.0 * p20 + 3.0 * p02 + 10.0 * p12 + 3.0 * p22) / 32.0;
let idx = y * w + x;
gx_data[idx] = dx;
gy_data[idx] = dy;
}
}
Ok(GradientField { gx, gy })
}
pub fn scharr_gradient_f32(image: &ImageView<'_, f32>) -> Result<GradientField> {
scharr_gradient(image)
}
pub fn compute_gradient<P: SourcePixel>(
image: &ImageView<'_, P>,
operator: GradientOperator,
) -> Result<GradientField> {
match operator {
GradientOperator::Sobel => sobel_gradient(image),
GradientOperator::Scharr => scharr_gradient(image),
}
}
pub fn compute_gradient_f32(
image: &ImageView<'_, f32>,
operator: GradientOperator,
) -> Result<GradientField> {
compute_gradient(image, operator)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn gradient_of_horizontal_step() {
#[rustfmt::skip]
let data: Vec<u8> = vec![
0, 0, 255, 255, 255,
0, 0, 255, 255, 255,
0, 0, 255, 255, 255,
];
let image = ImageView::from_slice(&data, 5, 3).unwrap();
let grad = sobel_gradient(&image).unwrap();
let (gx, gy) = grad.get(2, 1).unwrap();
assert!(
gx > 30.0,
"expected strong horizontal gradient, got gx={gx}"
);
assert!(
gy.abs() < 1e-6,
"expected zero vertical gradient, got gy={gy}"
);
}
#[test]
fn gradient_of_vertical_step() {
#[rustfmt::skip]
let data: Vec<u8> = vec![
0, 0, 0,
0, 0, 0,
255, 255, 255,
255, 255, 255,
255, 255, 255,
];
let image = ImageView::from_slice(&data, 3, 5).unwrap();
let grad = sobel_gradient(&image).unwrap();
let (gx, gy) = grad.get(1, 2).unwrap();
assert!(
gx.abs() < 1e-6,
"expected zero horizontal gradient, got gx={gx}"
);
assert!(gy > 30.0, "expected strong vertical gradient, got gy={gy}");
}
#[test]
fn gradient_magnitude_computation() {
let data: Vec<u8> = vec![0; 9];
let image = ImageView::from_slice(&data, 3, 3).unwrap();
let grad = sobel_gradient(&image).unwrap();
let mag = gradient_magnitude(&grad).unwrap();
assert!(mag.data().iter().all(|&v| v == 0.0));
}
#[test]
fn gradient_field_dimensions() {
let data: Vec<u8> = vec![128; 20];
let image = ImageView::from_slice(&data, 5, 4).unwrap();
let grad = sobel_gradient(&image).unwrap();
assert_eq!(grad.width(), 5);
assert_eq!(grad.height(), 4);
}
fn field_from(w: usize, h: usize, gx: Vec<Scalar>, gy: Vec<Scalar>) -> GradientField {
GradientField {
gx: OwnedImage::from_vec(gx, w, h).unwrap(),
gy: OwnedImage::from_vec(gy, w, h).unwrap(),
}
}
#[test]
fn thin_reduces_horizontal_band_to_single_column() {
let w = 5;
let h = 5;
let row = [0.0, 1.0, 2.0, 1.0, 0.0];
let mut gx = vec![0.0f32; w * h];
for y in 0..h {
for x in 0..w {
gx[y * w + x] = row[x];
}
}
let field = field_from(w, h, gx, vec![0.0; w * h]);
let thin = thin_gradient(&field).unwrap();
for y in 0..h {
for x in 0..w {
let kept = thin.magnitude(x, y).unwrap() > 0.0;
let expect = (1..=3).contains(&y) && x == 2;
assert_eq!(
kept, expect,
"pixel ({x},{y}) kept={kept}, expected {expect}"
);
}
}
}
#[test]
fn thin_reduces_vertical_band_to_single_row() {
let w = 5;
let h = 5;
let col = [0.0, 1.0, 2.0, 1.0, 0.0];
let mut gy = vec![0.0f32; w * h];
for y in 0..h {
for x in 0..w {
gy[y * w + x] = col[y];
}
}
let field = field_from(w, h, vec![0.0; w * h], gy);
let thin = thin_gradient(&field).unwrap();
for y in 0..h {
for x in 0..w {
let kept = thin.magnitude(x, y).unwrap() > 0.0;
let expect = y == 2 && (1..=3).contains(&x);
assert_eq!(
kept, expect,
"pixel ({x},{y}) kept={kept}, expected {expect}"
);
}
}
}
#[test]
fn thin_preserves_surviving_gradient_values() {
let w = 5;
let h = 3;
let row = [0.0, 1.0, 2.0, 1.0, 0.0];
let mut gx = vec![0.0f32; w * h];
for y in 0..h {
for x in 0..w {
gx[y * w + x] = row[x];
}
}
let field = field_from(w, h, gx, vec![0.0; w * h]);
let thin = thin_gradient(&field).unwrap();
let (gx_k, gy_k) = thin.get(2, 1).unwrap();
assert_eq!((gx_k, gy_k), (2.0, 0.0));
}
#[test]
fn thin_borders_stay_zero() {
let size = 32;
let mut data = vec![0u8; size * size];
for y in 0..size {
for x in 0..size {
let dx = x as f32 - 16.0;
let dy = y as f32 - 16.0;
if dx * dx + dy * dy <= 64.0 {
data[y * size + x] = 255;
}
}
}
let image = ImageView::from_slice(&data, size, size).unwrap();
let thin = thin_gradient(&sobel_gradient(&image).unwrap()).unwrap();
for x in 0..size {
assert_eq!(thin.magnitude(x, 0).unwrap(), 0.0);
assert_eq!(thin.magnitude(x, size - 1).unwrap(), 0.0);
}
for y in 0..size {
assert_eq!(thin.magnitude(0, y).unwrap(), 0.0);
assert_eq!(thin.magnitude(size - 1, y).unwrap(), 0.0);
}
}
#[test]
fn thin_is_deterministic_and_idempotent() {
let w = 16;
let h = 12;
let mut gx = vec![0.0f32; w * h];
let mut gy = vec![0.0f32; w * h];
for i in 0..w * h {
gx[i] = ((i * 7) % 11) as f32 - 5.0;
gy[i] = ((i * 13) % 9) as f32 - 4.0;
}
let field = field_from(w, h, gx, gy);
let a = thin_gradient(&field).unwrap();
let b = thin_gradient(&field).unwrap();
assert_eq!(a.gx.data(), b.gx.data(), "gx not deterministic");
assert_eq!(a.gy.data(), b.gy.data(), "gy not deterministic");
let c = thin_gradient(&a).unwrap();
assert_eq!(a.gx.data(), c.gx.data(), "not idempotent (gx)");
assert_eq!(a.gy.data(), c.gy.data(), "not idempotent (gy)");
}
#[test]
fn thin_tiny_image_no_panic() {
for (w, h) in [(2usize, 2usize), (1, 5), (5, 1), (3, 1), (1, 1)] {
let field = field_from(w, h, vec![1.0; w * h], vec![1.0; w * h]);
let thin = thin_gradient(&field).unwrap();
assert!(thin.gx.data().iter().all(|&v| v == 0.0));
assert!(thin.gy.data().iter().all(|&v| v == 0.0));
}
}
#[test]
fn scharr_gradient_of_horizontal_step() {
#[rustfmt::skip]
let data: Vec<u8> = vec![
0, 0, 255, 255, 255,
0, 0, 255, 255, 255,
0, 0, 255, 255, 255,
];
let image = ImageView::from_slice(&data, 5, 3).unwrap();
let grad = scharr_gradient(&image).unwrap();
let (gx, gy) = grad.get(2, 1).unwrap();
assert!(
gx > 30.0,
"expected strong horizontal gradient, got gx={gx}"
);
assert!(
gy.abs() < 1e-6,
"expected zero vertical gradient, got gy={gy}"
);
}
#[test]
fn scharr_gradient_zeros_on_uniform() {
let data: Vec<u8> = vec![128; 25];
let image = ImageView::from_slice(&data, 5, 5).unwrap();
let grad = scharr_gradient(&image).unwrap();
assert!(grad.gx().as_slice().iter().all(|&v| v == 0.0));
assert!(grad.gy().as_slice().iter().all(|&v| v == 0.0));
}
#[test]
fn scharr_dimensions_match() {
let data: Vec<u8> = vec![128; 20];
let image = ImageView::from_slice(&data, 5, 4).unwrap();
let grad = scharr_gradient(&image).unwrap();
assert_eq!(grad.width(), 5);
assert_eq!(grad.height(), 4);
}
#[test]
fn scharr_gradient_f32_matches_u8() {
#[rustfmt::skip]
let data_u8: Vec<u8> = vec![
0, 0, 255, 255, 255,
0, 0, 255, 255, 255,
0, 0, 255, 255, 255,
];
let data_f32: Vec<f32> = data_u8.iter().map(|&v| v as f32).collect();
let img_u8 = ImageView::from_slice(&data_u8, 5, 3).unwrap();
let img_f32 = ImageView::from_slice(&data_f32, 5, 3).unwrap();
let grad_u8 = scharr_gradient(&img_u8).unwrap();
let grad_f32 = scharr_gradient_f32(&img_f32).unwrap();
let (gx_u8, gy_u8) = grad_u8.get(2, 1).unwrap();
let (gx_f32, gy_f32) = grad_f32.get(2, 1).unwrap();
assert!(
(gx_u8 - gx_f32).abs() < 1e-4,
"gx mismatch: u8={gx_u8} f32={gx_f32}"
);
assert!(
(gy_u8 - gy_f32).abs() < 1e-4,
"gy mismatch: u8={gy_u8} f32={gy_f32}"
);
}
#[test]
fn compute_gradient_f32_dispatches() {
#[rustfmt::skip]
let data: Vec<f32> = vec![
0.0, 0.0, 255.0, 255.0, 255.0,
0.0, 0.0, 255.0, 255.0, 255.0,
0.0, 0.0, 255.0, 255.0, 255.0,
];
let image = ImageView::from_slice(&data, 5, 3).unwrap();
let sobel = compute_gradient_f32(&image, GradientOperator::Sobel).unwrap();
let scharr = compute_gradient_f32(&image, GradientOperator::Scharr).unwrap();
let (sobel_gx, _) = sobel.get(2, 1).unwrap();
let (scharr_gx, _) = scharr.get(2, 1).unwrap();
assert!(
sobel_gx.abs() > 0.1,
"Sobel f32 gx should be nonzero: {sobel_gx}"
);
assert!(
scharr_gx.abs() > 0.1,
"Scharr f32 gx should be nonzero: {scharr_gx}"
);
}
#[test]
fn compute_gradient_dispatches_correctly() {
#[rustfmt::skip]
let data: Vec<u8> = vec![
100, 0, 0,
0, 0, 0,
0, 0, 0,
];
let image = ImageView::from_slice(&data, 3, 3).unwrap();
let sobel = compute_gradient(&image, GradientOperator::Sobel).unwrap();
let scharr = compute_gradient(&image, GradientOperator::Scharr).unwrap();
let (sx, _) = sobel.get(1, 1).unwrap();
let (cx, _) = scharr.get(1, 1).unwrap();
assert!(sx.abs() > 0.1, "Sobel gx should be nonzero: {sx}");
assert!(cx.abs() > 0.1, "Scharr gx should be nonzero: {cx}");
assert!(
(sx - cx).abs() > 0.1,
"Sobel gx ({sx}) and Scharr gx ({cx}) should differ"
);
}
}