diff --git a/core/color_space.py b/core/color_space.py index 5ec536e..909eda2 100644 --- a/core/color_space.py +++ b/core/color_space.py @@ -30,6 +30,16 @@ def _bicone_factor(l, clamp_min=None): return factor +def _wrap_hue_diff(diff): + """Wrap hue differences to the shortest path on the unit circle [-0.5, 0.5].""" + return diff - (diff > 0.5).float() + (diff < -0.5).float() + + +def _hue_lerp(h1, h2, t): + """Lerp hues on the circle [0,1], taking the shortest path.""" + return (h1 + t * _wrap_hue_diff(h2 - h1)) % 1.0 + + def _chromatic_plane_basis(a): """Build orthonormal basis (a_unit, e1, e2) for the chromatic plane perpendicular to a.""" a_unit = a / (a.norm() + 1e-10) diff --git a/core/diagnostics.py b/core/diagnostics.py index e38ee93..6e7f756 100644 --- a/core/diagnostics.py +++ b/core/diagnostics.py @@ -6,7 +6,7 @@ cause image blurriness or quality degradation during LCS intervention. import torch import math -from .color_space import decode_lcs_to_hsl, encode_hsl_to_lcs +from .color_space import decode_lcs_to_hsl, encode_hsl_to_lcs, _hue_lerp from .timestep import get_alpha_beta, get_alpha_beta_t50, normalize_to_t50, denormalize_from_t50 # Test constants @@ -111,12 +111,13 @@ def test_type_ii_uniformity(anchor_lcs, anchor_angles): s_new = torch.full_like(s_cur, t_s) l_new = torch.full_like(l_cur, t_l) + # Compute input variance once (patches never changes) + input_var = patches.var(dim=0).mean().item() + # Test different strengths for strength in _TEST_STRENGTHS: - # Hue lerp (wrap to [-0.5, 0.5]) - diff = t_h - h_cur - diff = diff - (diff > 0.5).float() + (diff < -0.5).float() - h_interp = (h_cur + strength * diff) % 1.0 + # Hue lerp using shared helper + h_interp = _hue_lerp(h_cur, h_new, strength) s_interp = (s_cur + strength * (s_new - s_cur)).clamp(0, 1) l_interp = (l_cur + strength * (l_new - l_cur)).clamp(0, 1) @@ -124,7 +125,6 @@ def test_type_ii_uniformity(anchor_lcs, anchor_angles): new_patches = encode_hsl_to_lcs(h_interp, s_interp, l_interp, anchor_lcs, anchor_angles) # Measure variance loss - input_var = patches.var(dim=0).mean().item() output_var = new_patches.var(dim=0).mean().item() var_ratio = output_var / (input_var + 1e-10) @@ -145,10 +145,10 @@ def test_early_timestep_amplification(): """ # Typical LCS coordinate magnitude at t=50 c_ref = torch.tensor(_T50_REFERENCE_COORD, dtype=torch.float32) + alpha_50, beta_50 = get_alpha_beta_t50() # Constant across all sigmas for sigma in [1.0, 0.99, 0.95, 0.90, 0.85, 0.80, 0.50, 0.0]: alpha_t, beta_t = get_alpha_beta(sigma) - alpha_50, beta_50 = get_alpha_beta_t50() # Simulate a noisy observation at timestep t # In diffusion, the observation is alpha_t * clean + beta_t * noise diff --git a/nodes/intervene.py b/nodes/intervene.py index 87d6664..d63df9f 100644 --- a/nodes/intervene.py +++ b/nodes/intervene.py @@ -7,7 +7,7 @@ from comfy_api.latest import io from ..core.lcs_data import LCSData from ..core.patchify import patchify, unpatchify from ..core.timestep import get_alpha_beta, get_alpha_beta_t50, normalize_to_t50, denormalize_from_t50 -from ..core.color_space import hex_to_hsl, encode_hsl_to_lcs, decode_lcs_to_hsl +from ..core.color_space import hex_to_hsl, encode_hsl_to_lcs, decode_lcs_to_hsl, _hue_lerp LCS_DATA = io.Custom("LCS_DATA") @@ -179,15 +179,6 @@ def _build_post_cfg_fn(lcs_data, target_colors_hsl, strength, mode, start_step, return post_cfg_fn -def _hue_lerp(h1, h2, t): - """Lerp hues on the circle [0,1], taking the shortest path.""" - diff = h2 - h1 - # Wrap to [-0.5, 0.5] - diff = diff - (diff > 0.5).float() + (diff < -0.5).float() - result = h1 + t * diff - return result % 1.0 - - class LCSColorIntervene(io.ComfyNode): """Steer colors during FLUX generation via the Latent Color Subspace.