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:
@@ -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
@@ -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
@@ -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
@@ -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.
|
||||
"""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user