Fix sharpness calibration and simplify intervention hook

- Use random noise images instead of solid-color (blur had no effect
  on spatially uniform images, making calibration a no-op)
- Generate images per-batch to avoid 200 MB upfront allocation
- Simplify hook algebraically: patches + delta * pc1_dir eliminates
  projection/reconstruction/residual intermediates (~24 MB per step)
- Move SCALE_FACTOR, SHIFT_FACTOR, find_step_index to core/sampling.py
- Rename sd → shd to avoid confusion with state_dict convention
- Remove redundant [:,:,:,:3] slice and unused imports
This commit is contained in:
facok
2026-03-20 13:12:46 +08:00
parent cf9edf0239
commit 03fb376b8c
4 changed files with 79 additions and 84 deletions
+21
View File
@@ -0,0 +1,21 @@
"""Shared sampling utilities for LCS intervention hooks."""
import torch
# FLUX VAE process_in ↔ raw space conversion constants
SCALE_FACTOR = 0.3611
SHIFT_FACTOR = 0.1159
def find_step_index(sigma, sigmas):
"""Find the step index for a given sigma value in the sigma schedule.
Uses torch.isclose for robust matching across dtype differences (e.g. bfloat16
sigma vs float32 sample_sigmas), with argmin fallback for edge cases.
"""
sigma_val = sigma.flatten()[0].float()
sigmas_f = sigmas.float()
matched = torch.isclose(sigmas_f, sigma_val, rtol=1e-3, atol=1e-5).nonzero()
if len(matched) > 0:
return matched[0].item()
return (sigmas_f - sigma_val).abs().argmin().item()
+17 -22
View File
@@ -5,7 +5,6 @@ from dataclasses import dataclass
import torch
import torch.nn.functional as F
import comfy.model_management
import comfy.utils
from .patchify import patchify
@@ -74,8 +73,8 @@ def calibrate_sharpness(vae, num_samples=64, image_size=512,
blur_levels=(0, 1, 2, 4, 8, 16), batch_size=8):
"""Compute sharpness subspace data (PCA basis, mean, sign) from FLUX VAE.
1. Generate num_samples random grayscale solid-color images
2. For each blur level, apply Gaussian blur
1. Generate num_samples random noise images (spatial detail needed for blur)
2. For each blur level, apply Gaussian blur per-batch
3. VAE encode → patchify → average patches → [64] vector
4. PCA on all vectors → extract PC1 (+ PC2)
5. Determine sign: positive strength = sharper
@@ -83,48 +82,44 @@ def calibrate_sharpness(vae, num_samples=64, image_size=512,
Returns: SharpnessData
"""
device = comfy.model_management.intermediate_device()
n_levels = len(blur_levels)
total_images = num_samples * n_levels
print(f"\n[LCS Sharpness Calibration] Starting: {num_samples} images × {n_levels} blur levels = {total_images} samples")
print(f"[LCS Sharpness Calibration] Blur sigmas: {list(blur_levels)}")
# Step 1: Generate random grayscale base images [num_samples, 3, H, W]
# Use deterministic seed for reproducibility across runs
# Deterministic seed for reproducibility across runs
rng = torch.Generator().manual_seed(42)
gray_values = torch.rand(num_samples, 1, 1, 1, generator=rng)
# [num_samples, 3, H, W] in BCHW for blur, will convert to BHWC for VAE
base_images = gray_values.expand(num_samples, 3, image_size, image_size).contiguous()
# Step 2+3: For each blur level, blur all images, VAE encode, collect patch vectors
# Generate and process per-batch to avoid large upfront allocation.
# Random noise (not solid color) — spatial detail is needed for blur to
# have a measurable effect on VAE encoding.
vectors = []
blur_labels = [] # track blur sigma per vector for sign determination
pbar = comfy.utils.ProgressBar(total_images)
for blur_sigma in blur_levels:
# Apply blur
blurred = _apply_gaussian_blur(base_images, blur_sigma)
# Encode in batches
for batch_start in range(0, num_samples, batch_size):
batch_end = min(batch_start + batch_size, num_samples)
batch = blurred[batch_start:batch_end] # [B, 3, H, W] BCHW
actual_batch = batch.shape[0]
actual_batch = min(batch_size, num_samples - batch_start)
# Generate random noise images per-batch [B, 3, H, W]
batch = torch.rand(actual_batch, 3, image_size, image_size, generator=rng)
# Apply blur (no-op for sigma=0)
blurred = _apply_gaussian_blur(batch, blur_sigma)
# Convert BCHW → BHWC for ComfyUI VAE
imgs_bhwc = batch.permute(0, 2, 3, 1).contiguous().cpu()
imgs_bhwc = blurred.permute(0, 2, 3, 1).contiguous().cpu()
# VAE encode → [B, 16, H/8, W/8]
latent = vae.encode(imgs_bhwc[:, :, :, :3])
latent = vae.encode(imgs_bhwc)
# Patchify → [B, L, 64], average across patches → [B, 64]
patches, _, _ = patchify(latent)
avg = patches.mean(dim=1).cpu()
for j in range(actual_batch):
vectors.append(avg[j])
blur_labels.append(blur_sigma)
vectors.extend(avg.unbind(0))
blur_labels.extend([blur_sigma] * actual_batch)
pbar.update(actual_batch)
+3 -17
View File
@@ -6,28 +6,14 @@ from comfy_api.latest import io
from ..core.lcs_data import LCSData
from ..core.patchify import patchify, unpatchify
from ..core.sampling import SCALE_FACTOR, SHIFT_FACTOR, find_step_index
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, _hue_lerp
LCS_DATA = io.Custom("LCS_DATA")
# FLUX VAE constants
SCALE_FACTOR = 0.3611
SHIFT_FACTOR = 0.1159
def _find_step_index(sigma, sigmas):
"""Find the step index for a given sigma value in the sigma schedule.
Uses torch.isclose for robust matching across dtype differences (e.g. bfloat16
sigma vs float32 sample_sigmas), with argmin fallback for edge cases.
"""
sigma_val = sigma.flatten()[0].float()
sigmas_f = sigmas.float()
matched = torch.isclose(sigmas_f, sigma_val, rtol=1e-3, atol=1e-5).nonzero()
if len(matched) > 0:
return matched[0].item()
return (sigmas_f - sigma_val).abs().argmin().item()
# Backward-compat alias for any external code that imported _find_step_index
_find_step_index = find_step_index
def _build_post_cfg_fn(lcs_data, target_colors_hsl, strength, mode, start_step, end_step, mask):
+38 -45
View File
@@ -10,7 +10,7 @@ from safetensors.torch import save_file, load_file
from ..core.sharpness import SharpnessData, calibrate_sharpness
from ..core.calibration import vae_fingerprint
from ..core.patchify import patchify, unpatchify
from .intervene import _find_step_index, SCALE_FACTOR, SHIFT_FACTOR
from ..core.sampling import SCALE_FACTOR, SHIFT_FACTOR, find_step_index
SHARPNESS_DATA = io.Custom("SHARPNESS_DATA")
DATA_DIR = os.path.join(os.path.dirname(os.path.dirname(__file__)), "data")
@@ -75,11 +75,29 @@ class LCSSharpnessCalibrate(io.ComfyNode):
return io.NodeOutput(data)
def _downsample_mask(mask, h_len, w_len, device, dtype):
"""Downsample a mask to patch grid and flatten to [1, L, 1]."""
mask_dev = mask.to(device=device, dtype=dtype)
if mask_dev.ndim == 3:
mask_dev = mask_dev[:1]
if mask_dev.ndim == 2:
mask_4d = mask_dev.unsqueeze(0).unsqueeze(0) # [1, 1, H, W]
elif mask_dev.ndim == 3:
mask_4d = mask_dev.unsqueeze(1) # [B, 1, H, W]
else:
mask_4d = mask_dev
mask_resized = F.interpolate(
mask_4d, size=(h_len, w_len), mode="bilinear", align_corners=False
)
return mask_resized.reshape(1, -1, 1) # [1, L, 1]
def _build_sharpness_fn(sharpness_data, strength, start_step, end_step, mask):
"""Build the post_cfg_function closure for sharpness intervention.
Edits PC1 of the sharpness subspace. No timestep normalization —
the edit is applied directly in raw patch space.
Algebraically simplified: adding delta along PC1 direction preserves all
other dimensions by construction, so no explicit projection/residual needed.
patches_new = patches + delta * pc1_direction
"""
def post_cfg_fn(args):
denoised = args["denoised"] # [B, 16, H, W] in process_in space
@@ -87,7 +105,7 @@ def _build_sharpness_fn(sharpness_data, strength, start_step, end_step, mask):
# Step gating
sigmas = args["model_options"]["transformer_options"]["sample_sigmas"]
step_index = _find_step_index(sigma, sigmas)
step_index = find_step_index(sigma, sigmas)
if step_index < start_step or step_index > end_step:
return denoised
@@ -95,9 +113,7 @@ def _build_sharpness_fn(sharpness_data, strength, start_step, end_step, mask):
device = denoised.device
dtype = denoised.dtype
sd = sharpness_data.to(device, dtype)
B_mat = sd.basis # [64, K]
mu = sd.mean # [64]
shd = sharpness_data.to(device, dtype)
# Convert from process_in to raw VAE space
raw = denoised / SCALE_FACTOR + SHIFT_FACTOR # [B, 16, H, W]
@@ -105,48 +121,26 @@ def _build_sharpness_fn(sharpness_data, strength, start_step, end_step, mask):
# Patchify
patches, h_len, w_len = patchify(raw) # [B, L, 64]
# Project to sharpness subspace
projection = (patches - mu) @ B_mat # [B, L, K]
# Sharpness edit: add delta along PC1 direction.
# Since we only shift along one basis vector, the residual (all other
# dimensions) is preserved automatically — no need to project, compute
# residual, and reconstruct.
delta = strength * shd.sign * shd.pc1_std
pc1_dir = shd.basis[:, 0] # [64]
# Compute residual (preserve non-sharpness dimensions)
reconstruction = projection @ B_mat.T + mu # [B, L, 64]
residual = patches - reconstruction # [B, L, 64]
# Edit PC1: shift by strength * sign * pc1_std
delta = strength * sd.sign * sd.pc1_std
new_projection = projection.clone()
new_projection[..., 0] = new_projection[..., 0] + delta
# Apply mask if provided
if mask is not None:
mask_dev = mask.to(device=device, dtype=dtype)
if mask_dev.ndim == 3:
mask_dev = mask_dev[:1]
if mask_dev.ndim == 2:
mask_4d = mask_dev.unsqueeze(0).unsqueeze(0) # [1, 1, H, W]
elif mask_dev.ndim == 3:
mask_4d = mask_dev.unsqueeze(1) # [B, 1, H, W]
else:
mask_4d = mask_dev
mask_resized = F.interpolate(
mask_4d, size=(h_len, w_len), mode="bilinear", align_corners=False
)
mask_flat = mask_resized.reshape(1, -1, 1) # [1, L, 1]
if mask_flat.shape[1] != new_projection.shape[1]:
mask_flat = mask_flat[:, :new_projection.shape[1], :]
# Blend: masked areas get intervention, unmasked keep original
new_projection = projection + mask_flat * (new_projection - projection)
# Reconstruct patches with residual preservation
patches_new = new_projection @ B_mat.T + mu + residual # [B, L, 64]
mask_flat = _downsample_mask(mask, h_len, w_len, device, dtype)
if mask_flat.shape[1] != patches.shape[1]:
mask_flat = mask_flat[:, :patches.shape[1], :]
patches_new = patches + (mask_flat * delta) * pc1_dir
else:
patches_new = patches + delta * pc1_dir
# Unpatchify
raw_new = unpatchify(patches_new, h_len, w_len) # [B, 16, H, W]
# Convert back to process_in space
modified = (raw_new - SHIFT_FACTOR) * SCALE_FACTOR
return modified.to(dtype)
return ((raw_new - SHIFT_FACTOR) * SCALE_FACTOR).to(dtype)
return post_cfg_fn
@@ -154,9 +148,8 @@ def _build_sharpness_fn(sharpness_data, strength, start_step, end_step, mask):
class LCSSharpnessIntervene(io.ComfyNode):
"""Control sharpness during FLUX generation via the sharpness subspace.
Installs a post-CFG hook that projects the denoised prediction into the
sharpness subspace (PC1), shifts it by the requested strength,
preserves residual dimensions, and writes back.
Installs a post-CFG hook that adds a scaled shift along the sharpness
PC1 direction, preserving all other latent structure by construction.
Positive strength = sharper, negative = blurrier.
"""