Add auto mode with drift-driven intensity to LCS Color Anchor

Default mode is now "auto", which infers the concrete mode from
connected inputs (reference+vae → reference, mask → smooth,
nothing → self_anchor) and derives intensity from runtime drift
signals instead of requiring manual tuning.
This commit is contained in:
facok
2026-03-22 03:29:27 +08:00
parent ef74e1d5c9
commit bab475edc0
2 changed files with 448 additions and 0 deletions
+109
View File
@@ -0,0 +1,109 @@
"""Schedule-aware adaptive logic for LCS color anchoring.
Derives intervention windows, strength envelopes, and phase assignments
from the sigma schedule's amplification factor (beta_50 / beta_t), replacing
all manually-tuned step/strength parameters with data-driven decisions.
"""
import math
import torch
from .defaults import get_beta_table
def compute_amplification(sigma_val, device=None):
"""Compute amplification factor A = max_k(beta_50[k] / beta_t(sigma)[k]).
The amplification factor measures how much the normalization step inflates
noise relative to signal. High A means corrections are dangerous (amplified
noise dominates), low A means corrections are safe.
sigma_val: float in [0, 1] (FLUX sigma, 1=noise, 0=clean)
Returns: float amplification factor
"""
beta_table = get_beta_table() # [51, 3]
beta_50 = beta_table[50] # [3]
# Convert sigma to paper timestep
t = 50.0 * (1.0 - max(0.0, min(1.0, sigma_val)))
t = max(0.0, min(50.0, t))
t_low = int(t)
t_high = min(t_low + 1, 50)
frac = t - t_low
beta_t = (1.0 - frac) * beta_table[t_low] + frac * beta_table[t_high]
# Per-component ratio, take max
beta_t_safe = beta_t.clamp(min=1e-8)
ratios = beta_50 / beta_t_safe # [3]
return ratios.max().item()
def compute_step_phases(sigmas, mode):
"""Assign a phase to each sampling step based on amplification factor.
Physics-derived constants (not empirical):
A_MAX = 10.0 — above: normalization amplifies noise >10x → skip
A_WARMUP = 5.0 — self_anchor only: observe phase for EMA buildup
SIGMA_MIN = 0.15 — below: final detail refinement → skip
sigmas: 1D tensor of sigma values for each step (length N+1, last is 0)
mode: "smooth", "reference", or "self_anchor"
Returns: list of N strings, each "skip" / "observe" / "correct"
"""
A_MAX = 10.0
A_WARMUP = 5.0
SIGMA_MIN = 0.15
n_steps = len(sigmas) - 1 # last sigma is terminal (0)
phases = []
for i in range(n_steps):
sigma_val = float(sigmas[i])
# Final refinement — skip
if sigma_val < SIGMA_MIN:
phases.append("skip")
continue
amp = compute_amplification(sigma_val)
# Too noisy — skip
if amp > A_MAX:
phases.append("skip")
continue
# Self-anchor warmup zone
if mode == "self_anchor" and amp > A_WARMUP:
phases.append("observe")
continue
phases.append("correct")
return phases
def estimate_intensity(drift_signal):
"""Map drift magnitude to intensity in [0.15, 0.6]."""
DRIFT_SCALE = 0.2
INTENSITY_MIN = 0.15
INTENSITY_MAX = 0.6
return max(INTENSITY_MIN, min(INTENSITY_MAX, drift_signal / DRIFT_SCALE))
def compute_strength_envelope(n_correction_steps):
"""Sinusoidal bell envelope over correction steps.
sin(pi * i / (n-1)) for i in 0..n-1
Prevents abrupt on/off at phase boundaries.
Single step returns [1.0].
Returns: 1D tensor of length n_correction_steps
"""
if n_correction_steps <= 0:
return torch.zeros(0)
if n_correction_steps == 1:
return torch.ones(1)
n = n_correction_steps
indices = torch.arange(n, dtype=torch.float32)
return torch.sin(math.pi * indices / (n - 1))
+339
View File
@@ -0,0 +1,339 @@
"""Color anchor node: correct color drift during sampling.
Adaptive version — all scheduling and filtering parameters are derived
from runtime signals (sigma schedule, local color statistics, robust
outlier detection). User controls: mode + intensity.
"""
import torch
import torch.nn.functional as F
from comfy_api.latest import io
from ..core.adaptive import compute_step_phases, compute_strength_envelope, estimate_intensity
from ..core.bilateral import bilateral_filter_lcs, estimate_bilateral_params
from ..core.relationships import (
compute_local_relationships,
detect_anomalies_adaptive,
infer_color_from_neighbors,
)
from ..core.patchify import patchify, unpatchify
from ..core.sampling import (
find_step_index,
denoised_to_raw,
raw_to_denoised,
unpack_video_if_needed,
repack_video_if_needed,
downsample_mask,
)
from ..core.timestep import get_alpha_beta, get_alpha_beta_t50, normalize_to_t50, denormalize_from_t50
LCS_DATA = io.Custom("LCS_DATA")
def _encode_reference_to_lcs(reference_image, vae, lcs_data):
"""VAE-encode reference image to LCS coordinates.
reference_image: [B, H, W, 3] (BHWC ComfyUI format)
Returns (c_ref [1, L, 3], h_len, w_len) in t=50 space.
"""
latent = vae.encode(reference_image[:1, :, :, :3])
patches, h_len, w_len, _ = patchify(latent)
if patches is None:
return None, 0, 0
device = patches.device
dtype = patches.dtype
ld = lcs_data.to(device, dtype)
c_ref = (patches - ld.mean) @ ld.basis # [1, L, 3]
return c_ref, h_len, w_len
def _resize_color_field(c, src_h, src_w, dst_h, dst_w):
"""Bilinear resize of [B, L, 3] color field for resolution mismatch."""
if src_h == dst_h and src_w == dst_w:
return c
B = c.shape[0]
grid = c.reshape(B, src_h, src_w, 3).permute(0, 3, 1, 2)
resized = F.interpolate(grid, size=(dst_h, dst_w), mode="bilinear", align_corners=False)
return resized.permute(0, 2, 3, 1).reshape(B, -1, 3)
def _build_adaptive_anchor_fn(lcs_data, mode, intensity, mask,
c_ref=None, ref_h=0, ref_w=0, r_ref=None,
auto_intensity=False):
"""Build unified post_cfg_function for all anchor modes.
Phase assignment and strength scheduling are derived from the sigma
schedule on the first hook call. All filter/threshold parameters are
estimated from the data at each step.
Closure state auto-resets per graph execution (new closure = new dict).
"""
state = {
"phases": None,
"envelope": None,
"correction_index": 0,
"r_ema": None,
"c_ema": None,
"prev_c_mean": None,
"drift_samples": [],
"auto_intensity_val": None,
}
def post_cfg_fn(args):
denoised = args["denoised"]
sigma = args["sigma"]
model = args["model"]
# --- Lazy init: compute phases and envelope from sigma schedule ---
if state["phases"] is None:
sigmas = args["model_options"]["transformer_options"]["sample_sigmas"]
state["phases"] = compute_step_phases(sigmas, mode)
n_correct = sum(1 for p in state["phases"] if p == "correct")
state["envelope"] = compute_strength_envelope(n_correct)
state["correction_index"] = 0
# Find current step index
sigmas = args["model_options"]["transformer_options"]["sample_sigmas"]
step_index = find_step_index(sigma, sigmas)
# Look up phase (guard against out-of-range)
if step_index >= len(state["phases"]):
return denoised
phase = state["phases"][step_index]
# Skip phase — return unchanged
if phase == "skip":
return denoised
# --- Common pipeline: unpack → raw → patchify → project → normalize ---
working, pack_info = unpack_video_if_needed(denoised, args)
sigma_val = float(sigma.flatten()[0])
device = working.device
dtype = working.dtype
ld = lcs_data.to(device, dtype)
B_mat = ld.basis
mu = ld.mean
raw = denoised_to_raw(working, model)
patches, h_len, w_len, extra_shape = patchify(raw)
if patches is None:
return denoised
projection = (patches - mu) @ B_mat # [B, L, 3]
reconstruction = projection @ B_mat.T + mu
residual = patches - reconstruction
alpha_t, beta_t = get_alpha_beta(sigma_val, device=device)
alpha_t, beta_t = alpha_t.to(dtype), beta_t.to(dtype)
alpha_50, beta_50 = get_alpha_beta_t50(device=device)
alpha_50, beta_50 = alpha_50.to(dtype), beta_50.to(dtype)
c_norm = normalize_to_t50(projection, alpha_t, beta_t, alpha_50, beta_50)
# --- Observe phase (self_anchor warmup): update EMA, return unchanged ---
if phase == "observe":
r_current = compute_local_relationships(c_norm, h_len, w_len)
decay = 0.8
if state["r_ema"] is None:
state["r_ema"] = r_current.detach().clone()
state["c_ema"] = c_norm.detach().clone()
else:
state["r_ema"] = decay * state["r_ema"] + (1 - decay) * r_current.detach()
state["c_ema"] = decay * state["c_ema"] + (1 - decay) * c_norm.detach()
# Collect step-to-step drift for auto_intensity (self_anchor)
c_mean_now = c_norm.detach().mean(dim=1, keepdim=True)
if auto_intensity and state["prev_c_mean"] is not None:
drift = (c_mean_now - state["prev_c_mean"]).abs().mean().item()
state["drift_samples"].append(drift)
state["prev_c_mean"] = c_mean_now
return denoised
# --- Correct phase ---
# Auto-intensity: compute on first correction step, cache for rest
effective_intensity = intensity
if auto_intensity:
if state["auto_intensity_val"] is None:
if mode == "self_anchor" and state["drift_samples"]:
drift_signal = sum(state["drift_samples"]) / len(state["drift_samples"])
elif mode == "reference":
c_ref_dev = c_ref.to(device=device, dtype=dtype)
if ref_h != h_len or ref_w != w_len:
c_ref_meas = _resize_color_field(c_ref_dev, ref_h, ref_w, h_len, w_len)
else:
c_ref_meas = c_ref_dev
drift_signal = (c_norm - c_ref_meas).abs().mean().item()
elif mode == "smooth":
sigma_s, sigma_c = estimate_bilateral_params(c_norm, h_len, w_len)
c_filt = bilateral_filter_lcs(c_norm, h_len, w_len, sigma_s, sigma_c)
drift_signal = (c_filt - c_norm).abs().mean().item()
else:
drift_signal = 0.2 # fallback
state["auto_intensity_val"] = estimate_intensity(drift_signal)
effective_intensity = state["auto_intensity_val"]
# Compute step strength from envelope
ci = state["correction_index"]
envelope = state["envelope"]
if ci < len(envelope):
step_strength = effective_intensity * float(envelope[ci])
else:
step_strength = effective_intensity
state["correction_index"] = ci + 1
# Self-anchor convergence damping
if mode == "self_anchor" and state["prev_c_mean"] is not None:
c_mean_now = c_norm.detach().mean(dim=1, keepdim=True)
delta = (c_mean_now - state["prev_c_mean"]).abs().mean().item()
step_strength *= min(delta / 0.1, 1.0)
# Mode-specific correction
if mode == "smooth":
sigma_s, sigma_c = estimate_bilateral_params(c_norm, h_len, w_len)
c_filtered = bilateral_filter_lcs(c_norm, h_len, w_len, sigma_s, sigma_c)
new_c_norm = c_norm + step_strength * (c_filtered - c_norm)
elif mode == "reference":
c_ref_dev = c_ref.to(device=device, dtype=dtype)
r_ref_dev = r_ref.to(device=device, dtype=dtype)
if ref_h != h_len or ref_w != w_len:
c_ref_resized = _resize_color_field(c_ref_dev, ref_h, ref_w, h_len, w_len)
r_ref_resized = compute_local_relationships(c_ref_resized, h_len, w_len)
else:
c_ref_resized = c_ref_dev
r_ref_resized = r_ref_dev
B_size = c_norm.shape[0]
c_ref_exp = c_ref_resized.expand(B_size, -1, -1)
r_ref_exp = r_ref_resized.expand(B_size, -1, -1)
r_current = compute_local_relationships(c_norm, h_len, w_len)
anomaly_mag = detect_anomalies_adaptive(r_current, r_ref_exp)
correction = c_ref_exp - c_norm
new_c_norm = c_norm + step_strength * anomaly_mag * correction
else: # self_anchor
r_current = compute_local_relationships(c_norm, h_len, w_len)
if state["r_ema"] is None:
# No warmup data yet — seed EMA and skip
state["r_ema"] = r_current.detach().clone()
state["c_ema"] = c_norm.detach().clone()
state["prev_c_mean"] = c_norm.detach().mean(dim=1, keepdim=True)
return denoised
anomaly_mag = detect_anomalies_adaptive(r_current, state["r_ema"])
c_corrected = infer_color_from_neighbors(
c_norm, state["r_ema"], anomaly_mag, h_len, w_len
)
new_c_norm = c_norm + step_strength * (c_corrected - c_norm)
# Update EMA (slow decay during correction)
decay = 0.95
state["r_ema"] = decay * state["r_ema"] + (1 - decay) * r_current.detach()
state["c_ema"] = decay * state["c_ema"] + (1 - decay) * c_norm.detach()
state["prev_c_mean"] = c_norm.detach().mean(dim=1, keepdim=True)
# --- Apply mask ---
if mask is not None:
mask_flat = downsample_mask(mask, h_len, w_len, device, dtype)
if mask_flat.shape[1] != new_c_norm.shape[1]:
mask_flat = mask_flat[:, :new_c_norm.shape[1], :]
new_c_norm = c_norm + mask_flat * (new_c_norm - c_norm)
# --- Denormalize → reconstruct → unpatchify → repack ---
new_projection = denormalize_from_t50(new_c_norm, alpha_t, beta_t, alpha_50, beta_50)
patches_new = new_projection @ B_mat.T + mu + residual
raw_new = unpatchify(patches_new, h_len, w_len, extra_shape)
modified = raw_to_denoised(raw_new, model).to(dtype)
return repack_video_if_needed(modified, pack_info)
return post_cfg_fn
class LCSColorAnchor(io.ComfyNode):
"""Correct color drift during sampling by anchoring local color relationships.
Four modes:
- auto: Infer mode from connected inputs and intensity from drift signals
- smooth: Bilateral filter smooths color discontinuities (inpainting boundaries)
- reference: Anchor to a reference image's color relationships
- self_anchor: Build internal color model during warmup, then correct drift
All scheduling and filter parameters are derived adaptively from the sigma
schedule and image content. In auto mode, intensity is also derived automatically.
"""
@classmethod
def define_schema(cls) -> io.Schema:
return io.Schema(
node_id="LCSColorAnchor",
display_name="LCS Color Anchor",
category="LCS/intervention",
description="Correct color drift during sampling by anchoring local color relationships",
inputs=[
io.Model.Input("model"),
LCS_DATA.Input("lcs_data", tooltip="Calibration data from LCSLoadData"),
io.Combo.Input("mode", options=["auto", "smooth", "reference", "self_anchor"],
default="auto",
tooltip="auto: infer mode and intensity from connected inputs; smooth: bilateral filter; reference: anchor to image; self_anchor: auto-detect drift"),
io.Float.Input("intensity", default=0.5, min=0.0, max=1.0, step=0.05,
tooltip="Correction intensity (0 = none, 1 = full)"),
io.Vae.Input("vae", optional=True,
tooltip="Required for reference mode (VAE-encodes reference image)"),
io.Image.Input("reference_image", optional=True,
tooltip="Reference image for reference mode"),
io.Mask.Input("mask", optional=True,
tooltip="Optional mask for localized correction"),
],
outputs=[
io.Model.Output(display_name="model"),
],
)
@classmethod
def execute(cls, model, lcs_data, mode, intensity,
vae=None, reference_image=None, mask=None) -> io.NodeOutput:
"""Clone model, attach adaptive color anchor hook."""
m = model.clone()
# Resolve auto mode based on connected inputs
auto_intensity = False
if mode == "auto":
auto_intensity = True
if reference_image is not None and vae is not None:
mode = "reference"
elif mask is not None:
mode = "smooth"
else:
mode = "self_anchor"
if not auto_intensity and intensity < 1e-6:
return io.NodeOutput(m)
c_ref = None
ref_h = 0
ref_w = 0
r_ref = None
if mode == "reference":
if vae is None or reference_image is None:
print("[LCS Color Anchor] Reference mode requires vae and reference_image — skipping.")
return io.NodeOutput(m)
c_ref, ref_h, ref_w = _encode_reference_to_lcs(reference_image, vae, lcs_data)
if c_ref is None:
print("[LCS Color Anchor] Failed to encode reference image — skipping.")
return io.NodeOutput(m)
r_ref = compute_local_relationships(c_ref, ref_h, ref_w)
hook = _build_adaptive_anchor_fn(
lcs_data, mode, intensity, mask,
c_ref=c_ref, ref_h=ref_h, ref_w=ref_w, r_ref=r_ref,
auto_intensity=auto_intensity,
)
m.set_model_sampler_post_cfg_function(hook)
return io.NodeOutput(m)