Refactor: extract shared _hue_lerp and _wrap_hue_diff helpers
- Add _wrap_hue_diff() and _hue_lerp() to core/color_space.py for reuse across the codebase - Remove duplicate _hue_lerp from nodes/intervene.py, import from core - Update diagnostics.py to use shared _hue_lerp instead of inline logic - Move input_var computation outside strength loop in test_type_ii_uniformity - Move get_alpha_beta_t50() outside sigma loop in test_early_timestep_amplification
This commit is contained in:
@@ -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)
|
||||
|
||||
+7
-7
@@ -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
|
||||
|
||||
+1
-10
@@ -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.
|
||||
|
||||
|
||||
Reference in New Issue
Block a user