377 lines
9.6 KiB
Python
377 lines
9.6 KiB
Python
import torch
|
|
|
|
|
|
def _clamp01(x: torch.Tensor) -> torch.Tensor:
|
|
return x.clamp(0.0, 1.0)
|
|
|
|
|
|
def _srgb_to_linear(x: torch.Tensor) -> torch.Tensor:
|
|
# x in [0..1]
|
|
return torch.where(x <= 0.04045, x / 12.92, ((x + 0.055) / 1.055) ** 2.4)
|
|
|
|
|
|
def _linear_to_srgb(x: torch.Tensor) -> torch.Tensor:
|
|
# x in [0..1] ideally, but allow out-of-range before clamp
|
|
return torch.where(
|
|
x <= 0.0031308,
|
|
x * 12.92,
|
|
1.055 * torch.pow(torch.clamp_min(x, 0.0), 1.0 / 2.4) - 0.055,
|
|
)
|
|
|
|
|
|
def _apply_3x3(img: torch.Tensor, mat: torch.Tensor) -> torch.Tensor:
|
|
"""
|
|
img: [...,3]
|
|
mat: [3,3]
|
|
returns [...,3]
|
|
"""
|
|
return torch.einsum("...c,dc->...d", img, mat)
|
|
|
|
|
|
def _split_rgb_and_extra_channels(
|
|
img: torch.Tensor,
|
|
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
|
"""Split image into RGB channels and optional extra channels (e.g. alpha)."""
|
|
channels = img.shape[-1]
|
|
if channels < 3:
|
|
raise ValueError(
|
|
f"Expected at least 3 channels in the last dimension, got {channels}."
|
|
)
|
|
rgb = img[..., :3]
|
|
extras = img[..., 3:] if channels > 3 else None
|
|
return rgb, extras
|
|
|
|
|
|
def _recombine_rgb_and_extra_channels(
|
|
rgb: torch.Tensor, extras: torch.Tensor | None
|
|
) -> torch.Tensor:
|
|
if extras is None:
|
|
return rgb
|
|
return torch.cat((rgb, extras), dim=-1)
|
|
|
|
|
|
def _kelvin_to_xy_approx(k: float) -> tuple[float, float]:
|
|
"""
|
|
Practical approximation for CCT (Kelvin) -> CIE xy chromaticity.
|
|
Good enough for a WB slider; not "scientific-grade", but stable and common in tooling.
|
|
|
|
Valid-ish range: 1650K..25000K (we clamp to that).
|
|
"""
|
|
k = float(max(1650.0, min(25000.0, k)))
|
|
t = k
|
|
|
|
# x approximation
|
|
if t <= 4000.0:
|
|
x = (
|
|
(-0.2661239e9 / (t**3))
|
|
- (0.2343580e6 / (t**2))
|
|
+ (0.8776956e3 / t)
|
|
+ 0.179910
|
|
)
|
|
else:
|
|
x = (
|
|
(-3.0258469e9 / (t**3))
|
|
+ (2.1070379e6 / (t**2))
|
|
+ (0.2226347e3 / t)
|
|
+ 0.240390
|
|
)
|
|
|
|
# y approximation (piecewise in x)
|
|
if t <= 2222.0:
|
|
y = (
|
|
-1.1063814 * (x**3)
|
|
- 1.34811020 * (x**2)
|
|
+ 2.18555832 * x
|
|
- 0.20219683
|
|
)
|
|
elif t <= 4000.0:
|
|
y = (
|
|
-0.9549476 * (x**3)
|
|
- 1.37418593 * (x**2)
|
|
+ 2.09137015 * x
|
|
- 0.16748867
|
|
)
|
|
else:
|
|
y = (
|
|
3.0817580 * (x**3)
|
|
- 5.87338670 * (x**2)
|
|
+ 3.75112997 * x
|
|
- 0.37001483
|
|
)
|
|
|
|
# clamp to sane range
|
|
x = float(max(1e-6, min(0.999999, x)))
|
|
y = float(max(1e-6, min(0.999999, y)))
|
|
return x, y
|
|
|
|
|
|
def _xy_to_uv(x: float, y: float) -> tuple[float, float]:
|
|
# CIE 1960 UCS u,v (not u',v')
|
|
denom = (-2.0 * x) + (12.0 * y) + 3.0
|
|
if abs(denom) < 1e-8:
|
|
return 0.0, 0.0
|
|
u = (4.0 * x) / denom
|
|
v = (6.0 * y) / denom
|
|
return u, v
|
|
|
|
|
|
def _uv_to_xy(u: float, v: float) -> tuple[float, float]:
|
|
# inverse of CIE 1960 u,v
|
|
denom = (2.0 * u) - (8.0 * v) + 4.0
|
|
if abs(denom) < 1e-8:
|
|
return 0.3127, 0.3290 # fallback D65-ish
|
|
x = (3.0 * u) / denom
|
|
y = (2.0 * v) / denom
|
|
x = float(max(1e-6, min(0.999999, x)))
|
|
y = float(max(1e-6, min(0.999999, y)))
|
|
return x, y
|
|
|
|
|
|
def _xy_to_XYZ_white(x: float, y: float) -> tuple[float, float, float]:
|
|
X = x / y
|
|
Y = 1.0
|
|
Z = (1.0 - x - y) / y
|
|
return X, Y, Z
|
|
|
|
|
|
def _bradford_adaptation_matrix(
|
|
src_white_XYZ: torch.Tensor, # [3]
|
|
dst_white_XYZ: torch.Tensor, # [3]
|
|
device: torch.device,
|
|
dtype: torch.dtype,
|
|
) -> torch.Tensor:
|
|
"""
|
|
Returns A: [3,3] Bradford chromatic adaptation matrix for XYZ values.
|
|
"""
|
|
M = torch.tensor(
|
|
[
|
|
[0.8951, 0.2664, -0.1614],
|
|
[-0.7502, 1.7135, 0.0367],
|
|
[0.0389, -0.0685, 1.0296],
|
|
],
|
|
device=device,
|
|
dtype=dtype,
|
|
)
|
|
Minv = torch.tensor(
|
|
[
|
|
[0.9869929, -0.1470543, 0.1599627],
|
|
[0.4323053, 0.5183603, 0.0492912],
|
|
[-0.0085287, 0.0400428, 0.9684867],
|
|
],
|
|
device=device,
|
|
dtype=dtype,
|
|
)
|
|
|
|
src_LMS = M @ src_white_XYZ
|
|
dst_LMS = M @ dst_white_XYZ
|
|
|
|
# Avoid divide-by-zero
|
|
scale = dst_LMS / torch.clamp_min(src_LMS, 1e-8)
|
|
D = torch.diag(scale)
|
|
|
|
A = Minv @ D @ M
|
|
return A
|
|
|
|
|
|
_RGB_TO_XYZ = torch.tensor(
|
|
[
|
|
[0.4124564, 0.3575761, 0.1804375],
|
|
[0.2126729, 0.7151522, 0.0721750],
|
|
[0.0193339, 0.1191920, 0.9503041],
|
|
],
|
|
dtype=torch.float32,
|
|
)
|
|
|
|
_XYZ_TO_RGB = torch.tensor(
|
|
[
|
|
[3.2404542, -1.5371385, -0.4985314],
|
|
[-0.9692660, 1.8760108, 0.0415560],
|
|
[0.0556434, -0.2040259, 1.0572252],
|
|
],
|
|
dtype=torch.float32,
|
|
)
|
|
|
|
|
|
def _apply_white_balance_cat(
|
|
img_srgb: torch.Tensor,
|
|
temperature_k: float,
|
|
tint: float,
|
|
) -> torch.Tensor:
|
|
"""
|
|
img_srgb: [B,H,W,3] float in [0..1], assumed sRGB-ish
|
|
temperature_k: 1650..25000 typical slider
|
|
tint: -1..1 (green..magenta). Implemented as a small shift in CIE 1960 v.
|
|
"""
|
|
rgb_srgb, extras = _split_rgb_and_extra_channels(img_srgb)
|
|
|
|
device = img_srgb.device
|
|
dtype = img_srgb.dtype
|
|
|
|
# 1) sRGB -> linear
|
|
lin = _srgb_to_linear(_clamp01(rgb_srgb))
|
|
|
|
# 2) linear RGB -> XYZ
|
|
rgb2xyz = _RGB_TO_XYZ.to(device=device, dtype=dtype)
|
|
xyz = _apply_3x3(lin, rgb2xyz)
|
|
|
|
# 3) build src/dst white points in XYZ
|
|
# Source assumed D65 (sRGB reference white)
|
|
# D65 xy:
|
|
src_x, src_y = 0.3127, 0.3290
|
|
src_X, src_Y, src_Z = _xy_to_XYZ_white(src_x, src_y)
|
|
|
|
# Destination from Kelvin + tint offset
|
|
dst_x, dst_y = _kelvin_to_xy_approx(float(temperature_k))
|
|
|
|
# Calculate baseline D65 offset so 6500K is perfectly neutral (D65)
|
|
base_x, base_y = _kelvin_to_xy_approx(6500.0)
|
|
base_u, base_v = _xy_to_uv(base_x, base_y)
|
|
d65_u, d65_v = _xy_to_uv(src_x, src_y)
|
|
u_offset = d65_u - base_u
|
|
v_offset = d65_v - base_v
|
|
|
|
# Tint: shift in UCS v direction (green<->magenta feel)
|
|
# Scale chosen to be "good enough" and not insane.
|
|
# If you want stronger/weaker, tweak 0.05.
|
|
u, v = _xy_to_uv(dst_x, dst_y)
|
|
u = u + u_offset
|
|
v = v + v_offset + float(tint) * 0.05
|
|
# Clamp to sane bounds
|
|
v = float(max(1e-6, min(0.999999, v)))
|
|
dst_x, dst_y = _uv_to_xy(u, v)
|
|
dst_X, dst_Y, dst_Z = _xy_to_XYZ_white(dst_x, dst_y)
|
|
|
|
src_white = torch.tensor([src_X, src_Y, src_Z], device=device, dtype=dtype)
|
|
dst_white = torch.tensor([dst_X, dst_Y, dst_Z], device=device, dtype=dtype)
|
|
|
|
# 4) Bradford adaptation in XYZ
|
|
A = _bradford_adaptation_matrix(
|
|
src_white, dst_white, device=device, dtype=dtype
|
|
)
|
|
xyz_adapted = _apply_3x3(xyz, A)
|
|
|
|
# 5) XYZ -> linear RGB
|
|
xyz2rgb = _XYZ_TO_RGB.to(device=device, dtype=dtype)
|
|
lin_out = _apply_3x3(xyz_adapted, xyz2rgb)
|
|
|
|
# 6) linear -> sRGB
|
|
out_rgb = _clamp01(_linear_to_srgb(lin_out))
|
|
return _recombine_rgb_and_extra_channels(out_rgb, extras)
|
|
|
|
|
|
def _apply_brightness_contrast_gamma(
|
|
img: torch.Tensor,
|
|
brightness: float = 1.0,
|
|
contrast: float = 1.0,
|
|
gamma: float = 1.0,
|
|
) -> torch.Tensor:
|
|
out = img * float(brightness)
|
|
out = (out - 0.5) * float(contrast) + 0.5
|
|
out = _clamp01(out)
|
|
|
|
g = float(gamma)
|
|
if g != 1.0:
|
|
out = torch.pow(out.clamp_min(1e-8), 1.0 / g)
|
|
|
|
return _clamp01(out)
|
|
|
|
|
|
def _rgb_to_hsv(rgb: torch.Tensor) -> torch.Tensor:
|
|
r, g, b = rgb.unbind(dim=-1)
|
|
maxc = torch.max(rgb, dim=-1).values
|
|
minc = torch.min(rgb, dim=-1).values
|
|
v = maxc
|
|
delt = maxc - minc
|
|
|
|
s = torch.where(maxc > 0.0, delt / (maxc + 1e-8), torch.zeros_like(maxc))
|
|
h = torch.zeros_like(maxc)
|
|
|
|
mask = delt > 1e-8
|
|
delt_safe = delt + 1e-8
|
|
|
|
rc = (maxc - r) / delt_safe
|
|
gc = (maxc - g) / delt_safe
|
|
bc = (maxc - b) / delt_safe
|
|
|
|
h_r = (bc - gc) % 6.0
|
|
h_g = rc - bc + 2.0
|
|
h_b = gc - rc + 4.0
|
|
|
|
is_r = (maxc == r) & mask
|
|
is_g = (maxc == g) & mask
|
|
is_b = (maxc == b) & mask
|
|
|
|
h = torch.where(is_r, h_r, h)
|
|
h = torch.where(is_g, h_g, h)
|
|
h = torch.where(is_b, h_b, h)
|
|
|
|
h = (h / 6.0) % 1.0
|
|
return torch.stack((h, s, v), dim=-1)
|
|
|
|
|
|
def _hsv_to_rgb(hsv: torch.Tensor) -> torch.Tensor:
|
|
h, s, v = hsv.unbind(dim=-1)
|
|
h6 = (h % 1.0) * 6.0
|
|
i = torch.floor(h6).to(torch.int64)
|
|
f = h6 - torch.floor(h6)
|
|
|
|
p = v * (1.0 - s)
|
|
q = v * (1.0 - s * f)
|
|
t = v * (1.0 - s * (1.0 - f))
|
|
|
|
i_mod = i % 6
|
|
r = torch.where(
|
|
i_mod == 0,
|
|
v,
|
|
torch.where(
|
|
i_mod == 1,
|
|
q,
|
|
torch.where(
|
|
i_mod == 2,
|
|
p,
|
|
torch.where(i_mod == 3, p, torch.where(i_mod == 4, t, v)),
|
|
),
|
|
),
|
|
)
|
|
g = torch.where(
|
|
i_mod == 0,
|
|
t,
|
|
torch.where(
|
|
i_mod == 1,
|
|
v,
|
|
torch.where(
|
|
i_mod == 2,
|
|
v,
|
|
torch.where(i_mod == 3, q, torch.where(i_mod == 4, p, p)),
|
|
),
|
|
),
|
|
)
|
|
b = torch.where(
|
|
i_mod == 0,
|
|
p,
|
|
torch.where(
|
|
i_mod == 1,
|
|
p,
|
|
torch.where(
|
|
i_mod == 2,
|
|
t,
|
|
torch.where(i_mod == 3, v, torch.where(i_mod == 4, v, q)),
|
|
),
|
|
),
|
|
)
|
|
|
|
return _clamp01(torch.stack((r, g, b), dim=-1))
|
|
|
|
|
|
def _apply_saturation_hue(
|
|
img: torch.Tensor,
|
|
saturation: float = 1.0,
|
|
hue_degrees: float = 0.0,
|
|
) -> torch.Tensor:
|
|
rgb, extras = _split_rgb_and_extra_channels(img)
|
|
hsv = _rgb_to_hsv(rgb)
|
|
hsv[..., 0] = (hsv[..., 0] + (float(hue_degrees) / 360.0)) % 1.0
|
|
hsv[..., 1] = (hsv[..., 1] * float(saturation)).clamp(0.0, 1.0)
|
|
out_rgb = _hsv_to_rgb(hsv)
|
|
return _recombine_rgb_and_extra_channels(out_rgb, extras)
|