Skip to main content

terra_texture_rs/
stretch.rs

1//! Standard-deviation contrast stretch, as a fused single-pass kernel.
2//!
3//! Mirrors `TerraTexture.stretch`'s `stretch_std()`: values are mapped
4//! linearly so that `mean - n_std·std` becomes 0 and `mean + n_std·std`
5//! becomes 1, then clipped to \[0, 1\]. NaN in means NaN out, and NaNs
6//! are excluded from the mean and standard deviation.
7//!
8//! # Why this is faster than numpy
9//!
10//! The numpy version calls `np.nanmean()` and then `np.nanstd()`, which
11//! recomputes the mean internally, so the array is reduced several times
12//! before the elementwise subtract/divide/clip. Here one reduction pass
13//! collects the sum, sum of squares and count together. (Deduplicating
14//! only the mean call in numpy gives about 1.1×; numpy's C reductions are
15//! already efficient, so the win is fewer passes, not language speed.)
16//!
17//! # Precision
18//!
19//! Statistics are accumulated in `f64` for robustness on large arrays.
20//! This deliberately does **not** bit-match numpy's `float32` pairwise
21//! summation; it is verified against it within a statistical tolerance
22//! in `tests/test_stretch_rust.py`, which is the right bar for a
23//! mean/std computation.
24
25use ndarray::{Array2, ArrayView2, Zip};
26use rayon::prelude::*;
27
28use crate::common::PARALLEL_THRESHOLD;
29
30/// Stretch one value: `v` (`f32`, may be NaN), the lower bound `lo` and
31/// the range `denom = hi - lo + 1e-12`. Returns `(v - lo) / denom`
32/// clamped to \[0, 1\], or NaN if `v` is NaN (matching `np.clip`, which
33/// propagates NaN).
34#[inline]
35fn stretch_pixel(v: f32, lo: f32, denom: f32) -> f32 {
36    if v.is_nan() {
37        f32::NAN // matches np.clip((arr - lo) / denom, 0, 1): NaN propagates, never clamped away
38    } else {
39        ((v - lo) / denom).clamp(0.0, 1.0)
40    }
41}
42
43/// Single-threaded NaN-skipping reduction.
44///
45/// Returns `(sum, sum_of_squares, count)` as `(f64, f64, u64)`, over the
46/// non-NaN elements of `arr` only.
47fn sum_stats_serial(arr: ArrayView2<f32>) -> (f64, f64, u64) {
48    let mut sum = 0.0f64;
49    let mut sumsq = 0.0f64;
50    let mut count = 0u64;
51    for &v in arr.iter() {
52        if !v.is_nan() {
53            let vd = v as f64;
54            sum += vd;
55            sumsq += vd * vd;
56            count += 1;
57        }
58    }
59    (sum, sumsq, count)
60}
61
62/// Parallel version of [`sum_stats_serial`], same return type.
63///
64/// Uses rayon's fold + reduce over the array's flat slice, which needs a
65/// standard-layout (C-contiguous) array. If `arr` isn't one, falls back
66/// to the serial path rather than panicking. In practice it always is,
67/// because `stretch.py`'s dispatch passes an `ascontiguousarray()`'d
68/// array.
69fn sum_stats_parallel(arr: ArrayView2<f32>) -> (f64, f64, u64) {
70    // rayon's fold+reduce needs a flat parallel iterator; arrays here are
71    // standard/C-contiguous (stretch.py's dispatch passes a
72    // `.ascontiguousarray()`'d array), so `.as_slice()` is reliably
73    // `Some`. Falls back to the serial path in the (should-be-unreachable
74    // in practice) case it isn't, rather than panicking.
75    match arr.as_slice() {
76        Some(slice) => slice
77            .par_iter()
78            .fold(
79                || (0.0f64, 0.0f64, 0u64),
80                |(s, sq, c), &v| {
81                    if v.is_nan() {
82                        (s, sq, c)
83                    } else {
84                        let vd = v as f64;
85                        (s + vd, sq + vd * vd, c + 1)
86                    }
87                },
88            )
89            .reduce(
90                || (0.0f64, 0.0f64, 0u64),
91                |(s1, sq1, c1), (s2, sq2, c2)| (s1 + s2, sq1 + sq2, c1 + c2),
92            ),
93        None => sum_stats_serial(arr),
94    }
95}
96
97/// Contrast-stretch an array to \[0, 1\] using mean ± `n_std` standard
98/// deviations, ignoring NaNs.
99///
100/// # Arguments
101///
102/// * `arr` - `ArrayView2<f32>`, shape (H, W): input values, any range.
103///   NaN is allowed and marks missing data. Any memory layout, but the
104///   parallel reduction is only used on C-contiguous input.
105/// * `n_std` - `f32`: half-width of the stretch window in standard
106///   deviations (e.g. `2.0` maps mean − 2σ → 0 and mean + 2σ → 1).
107/// * `out` - `&mut Array2<f32>`, shape (H, W): overwritten with the
108///   stretched values in \[0, 1\], and NaN wherever `arr` is NaN.
109///
110/// Uses the population standard deviation (numpy's default `ddof=0`).
111/// If every element is NaN, or the array is empty, mean and std are
112/// taken as 0.
113///
114/// A constant array has std 0, so the window collapses to a point and
115/// only the `1e-12` guard keeps the division finite. Every value then
116/// equals the mean, so every output is 0, as in the numpy version.
117///
118/// Runs serially below [`PARALLEL_THRESHOLD`]
119/// elements and in parallel at or above it.
120///
121/// # Panics
122///
123/// If `out` does not have the same shape as `arr`.
124///
125/// # Example
126///
127/// ```
128/// use ndarray::{array, Array2};
129/// use terra_texture_rs::stretch_std_core;
130///
131/// let arr = array![[1.0_f32, 2.0], [3.0, f32::NAN]];
132/// let mut out = Array2::<f32>::zeros(arr.raw_dim());
133/// stretch_std_core(arr.view(), 1.0, &mut out);
134///
135/// assert!((out[[0, 1]] - 0.5).abs() < 1e-6); // the mean maps to 0.5
136/// assert!(out[[1, 1]].is_nan());             // NaN passes through
137/// ```
138pub fn stretch_std_core(arr: ArrayView2<f32>, n_std: f32, out: &mut Array2<f32>) {
139    let n = arr.len();
140    let (sum, sumsq, count) = if n >= PARALLEL_THRESHOLD {
141        sum_stats_parallel(arr)
142    } else {
143        sum_stats_serial(arr)
144    };
145
146    let mean_f64 = if count > 0 { sum / count as f64 } else { 0.0 };
147    let variance_f64 = if count > 0 {
148        (sumsq / count as f64 - mean_f64 * mean_f64).max(0.0) // guard tiny negative from float error
149    } else {
150        0.0
151    };
152    let mean = mean_f64 as f32;
153    let std = variance_f64.sqrt() as f32;
154    let lo = mean - n_std * std;
155    let hi = mean + n_std * std;
156    let denom = hi - lo + 1e-12;
157
158    let combine = |o: &mut f32, &v: &f32| *o = stretch_pixel(v, lo, denom);
159    let z = Zip::from(out).and(&arr);
160    if n >= PARALLEL_THRESHOLD {
161        z.par_for_each(combine);
162    } else {
163        z.for_each(combine);
164    }
165}