use std::sync::OnceLock;
use image::{DynamicImage, Rgb, RgbImage};
static REFERENCE_TILE_BYTES: &[u8] = include_bytes!("../assets/reference_tile.jpg");
const STD_EPSILON: f64 = 1e-6;
const XN: f64 = 0.95047;
const YN: f64 = 1.0;
const ZN: f64 = 1.08883;
fn srgb_to_linear(c: f64) -> f64 {
if c <= 0.04045 {
c / 12.92
} else {
((c + 0.055) / 1.055).powf(2.4)
}
}
fn linear_to_srgb(c: f64) -> f64 {
if c <= 0.0031308 {
c * 12.92
} else {
1.055 * c.powf(1.0 / 2.4) - 0.055
}
}
fn xyz_forward(t: f64) -> f64 {
const DELTA: f64 = 6.0 / 29.0;
if t > DELTA * DELTA * DELTA {
t.cbrt()
} else {
t / (3.0 * DELTA * DELTA) + 4.0 / 29.0
}
}
fn xyz_inverse(t: f64) -> f64 {
const DELTA: f64 = 6.0 / 29.0;
if t > DELTA {
t * t * t
} else {
3.0 * DELTA * DELTA * (t - 4.0 / 29.0)
}
}
fn rgb_to_lab(pixel: Rgb<u8>) -> [f64; 3] {
let [r, g, b] = pixel.0;
let r = srgb_to_linear(r as f64 / 255.0);
let g = srgb_to_linear(g as f64 / 255.0);
let b = srgb_to_linear(b as f64 / 255.0);
let x = 0.4124564 * r + 0.3575761 * g + 0.1804375 * b;
let y = 0.2126729 * r + 0.7151522 * g + 0.0721750 * b;
let z = 0.0193339 * r + 0.1191920 * g + 0.9503041 * b;
let fx = xyz_forward(x / XN);
let fy = xyz_forward(y / YN);
let fz = xyz_forward(z / ZN);
let l = 116.0 * fy - 16.0;
let a = 500.0 * (fx - fy);
let b = 200.0 * (fy - fz);
[l, a, b]
}
fn lab_to_rgb(lab: [f64; 3]) -> Rgb<u8> {
let [l, a, b] = lab;
let fy = (l + 16.0) / 116.0;
let fx = fy + a / 500.0;
let fz = fy - b / 200.0;
let x = XN * xyz_inverse(fx);
let y = YN * xyz_inverse(fy);
let z = ZN * xyz_inverse(fz);
let r = 3.2404542 * x - 1.5371385 * y - 0.4985314 * z;
let g = -0.9692660 * x + 1.8760108 * y + 0.0415560 * z;
let b = 0.0556434 * x - 0.2040259 * y + 1.0572252 * z;
let to_u8 = |c: f64| (linear_to_srgb(c).clamp(0.0, 1.0) * 255.0).round() as u8;
Rgb([to_u8(r), to_u8(g), to_u8(b)])
}
fn lab_mean_std(lab_pixels: &[[f64; 3]]) -> ([f64; 3], [f64; 3]) {
let count = lab_pixels.len() as f64;
let mut mean = [0.0; 3];
for pixel in lab_pixels {
for c in 0..3 {
mean[c] += pixel[c];
}
}
for m in &mut mean {
*m /= count;
}
let mut variance = [0.0; 3];
for pixel in lab_pixels {
for c in 0..3 {
let diff = pixel[c] - mean[c];
variance[c] += diff * diff;
}
}
let mut std = [0.0; 3];
for c in 0..3 {
std[c] = (variance[c] / count).sqrt();
}
(mean, std)
}
fn target_stats() -> &'static ([f64; 3], [f64; 3]) {
static TARGET_STATS: OnceLock<([f64; 3], [f64; 3])> = OnceLock::new();
TARGET_STATS.get_or_init(|| {
let reference = image::load_from_memory(REFERENCE_TILE_BYTES)
.expect("bundled reference tile is a valid image");
let lab_pixels: Vec<[f64; 3]> = reference
.to_rgb8()
.pixels()
.map(|&p| rgb_to_lab(p))
.collect();
lab_mean_std(&lab_pixels)
})
}
pub fn normalize_reinhard(image: &DynamicImage) -> DynamicImage {
let (target_mean, target_std) = target_stats();
let rgb_image = image.to_rgb8();
let lab_pixels: Vec<[f64; 3]> = rgb_image.pixels().map(|&p| rgb_to_lab(p)).collect();
let (mean, std) = lab_mean_std(&lab_pixels);
let (width, height) = rgb_image.dimensions();
let mut output = RgbImage::new(width, height);
for (pixel, lab_pixel) in output.pixels_mut().zip(lab_pixels.iter()) {
let mut normalized = [0.0; 3];
for c in 0..3 {
let scale = if std[c] > STD_EPSILON {
target_std[c] / std[c]
} else {
1.0
};
normalized[c] = (lab_pixel[c] - mean[c]) * scale + target_mean[c];
}
*pixel = lab_to_rgb(normalized);
}
DynamicImage::ImageRgb8(output)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn srgb_linear_round_trips_at_the_extremes() {
assert!((srgb_to_linear(0.0) - 0.0).abs() < 1e-9);
assert!((srgb_to_linear(1.0) - 1.0).abs() < 1e-9);
assert!((linear_to_srgb(0.0) - 0.0).abs() < 1e-9);
assert!((linear_to_srgb(1.0) - 1.0).abs() < 1e-9);
for c in [0.02, 0.2, 0.5, 0.9] {
let round_tripped = linear_to_srgb(srgb_to_linear(c));
assert!((round_tripped - c).abs() < 1e-6, "c={c}");
}
}
#[test]
fn rgb_to_lab_round_trips_within_rounding_tolerance() {
for pixel in [
Rgb([0, 0, 0]),
Rgb([255, 255, 255]),
Rgb([200, 50, 50]),
Rgb([30, 180, 90]),
] {
let lab = rgb_to_lab(pixel);
let round_tripped = lab_to_rgb(lab);
for c in 0..3 {
let diff = (round_tripped.0[c] as i16 - pixel.0[c] as i16).abs();
assert!(diff <= 1, "pixel={pixel:?} round_tripped={round_tripped:?}");
}
}
}
#[test]
fn lab_mean_std_matches_hand_computed_values() {
let pixels = [[0.0, 0.0, 0.0], [2.0, 2.0, 2.0]];
let (mean, std) = lab_mean_std(&pixels);
assert_eq!(mean, [1.0, 1.0, 1.0]);
assert_eq!(std, [1.0, 1.0, 1.0]);
}
#[test]
fn normalize_reinhard_preserves_image_dimensions() {
let image = DynamicImage::ImageRgb8(RgbImage::from_pixel(4, 3, Rgb([120, 80, 60])));
let normalized = normalize_reinhard(&image);
assert_eq!(normalized.width(), 4);
assert_eq!(normalized.height(), 3);
}
}