diff --git a/core/sampling.py b/core/sampling.py new file mode 100644 index 0000000..8e8ccfc --- /dev/null +++ b/core/sampling.py @@ -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() diff --git a/core/sharpness.py b/core/sharpness.py index 1f54b7f..34a2018 100644 --- a/core/sharpness.py +++ b/core/sharpness.py @@ -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) diff --git a/nodes/intervene.py b/nodes/intervene.py index d63df9f..d9f6f18 100644 --- a/nodes/intervene.py +++ b/nodes/intervene.py @@ -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): diff --git a/nodes/sharpen.py b/nodes/sharpen.py index 2c0353f..874f90c 100644 --- a/nodes/sharpen.py +++ b/nodes/sharpen.py @@ -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. """