fix: separate alpha channel when color processing

This commit is contained in:
Sonny Box
2026-04-05 20:48:19 -07:00
parent 46e8e6a1eb
commit 3d04ffd83b
2 changed files with 32 additions and 6 deletions
+31 -5
View File
@@ -28,6 +28,28 @@ def _apply_3x3(img: torch.Tensor, mat: torch.Tensor) -> torch.Tensor:
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.
@@ -180,11 +202,13 @@ def _apply_white_balance_cat(
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(img_srgb))
lin = _srgb_to_linear(_clamp01(rgb_srgb))
# 2) linear RGB -> XYZ
rgb2xyz = _RGB_TO_XYZ.to(device=device, dtype=dtype)
@@ -231,8 +255,8 @@ def _apply_white_balance_cat(
lin_out = _apply_3x3(xyz_adapted, xyz2rgb)
# 6) linear -> sRGB
out = _linear_to_srgb(lin_out)
return _clamp01(out)
out_rgb = _clamp01(_linear_to_srgb(lin_out))
return _recombine_rgb_and_extra_channels(out_rgb, extras)
def _apply_brightness_contrast_gamma(
@@ -344,7 +368,9 @@ def _apply_saturation_hue(
saturation: float = 1.0,
hue_degrees: float = 0.0,
) -> torch.Tensor:
hsv = _rgb_to_hsv(img)
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)
return _hsv_to_rgb(hsv)
out_rgb = _hsv_to_rgb(hsv)
return _recombine_rgb_and_extra_channels(out_rgb, extras)
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-supernodes"
description = "A miscellaneous node pack of custom implementations needed for workflows released by my SuperComfy YouTube channel."
version = "0.2.8"
version = "0.2.9"
license = { file = "LICENSE" }
dependencies = ["matplotlib"]