From cf9edf02395cb7637112edda99f6ea13da0614ad Mon Sep 17 00:00:00 2001 From: facok <128763816+facok@users.noreply.github.com> Date: Fri, 20 Mar 2026 13:01:23 +0800 Subject: [PATCH] Add sharpness intervention via PCA subspace in FLUX VAE patch space New calibration (core/sharpness.py) generates blur stimuli, VAE-encodes, and extracts PC1 as the sharpness direction. New nodes (nodes/sharpen.py) provide LCSSharpnessCalibrate (auto-cached per-VAE) and LCSSharpnessIntervene (post-CFG hook with strength, step window, mask). Positive strength = sharper, negative = blurrier. --- __init__.py | 3 + core/sharpness.py | 169 ++++++++++++++++++++++++++++++++++++++++ nodes/__init__.py | 5 ++ nodes/sharpen.py | 195 ++++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 372 insertions(+) create mode 100644 core/sharpness.py create mode 100644 nodes/sharpen.py diff --git a/__init__.py b/__init__.py index acf57c5..72c12c6 100644 --- a/__init__.py +++ b/__init__.py @@ -8,6 +8,7 @@ from comfy_api.latest import ComfyExtension, io from .nodes.calibrate import LCSLoadData from .nodes.intervene import LCSColorIntervene, LCSColorBatch, LCSToneAdjust from .nodes.observe import LCSPreviewColors, LCSStepObserver +from .nodes.sharpen import LCSSharpnessCalibrate, LCSSharpnessIntervene class LCSExtension(ComfyExtension): @@ -22,6 +23,8 @@ class LCSExtension(ComfyExtension): LCSToneAdjust, LCSPreviewColors, LCSStepObserver, + LCSSharpnessCalibrate, + LCSSharpnessIntervene, ] diff --git a/core/sharpness.py b/core/sharpness.py new file mode 100644 index 0000000..1f54b7f --- /dev/null +++ b/core/sharpness.py @@ -0,0 +1,169 @@ +"""Sharpness subspace calibration: PCA on blur stimuli in FLUX VAE patch space.""" + +import math +from dataclasses import dataclass + +import torch +import torch.nn.functional as F +import comfy.model_management +import comfy.utils + +from .patchify import patchify + + +@dataclass +class SharpnessData: + """Calibration data for the sharpness subspace. + + Produced by PCA on FLUX VAE-encoded images at varying blur levels. + PC1 captures ~93% of sharpness/blur variance. + """ + + basis: torch.Tensor # [64, K] PCA basis (columns), K typically 1-2 + mean: torch.Tensor # [64] PCA mean + pc1_std: float # Standard deviation along PC1 (for normalizing strength) + sign: float # +1 or -1: ensures positive strength = sharper + + def to(self, device, dtype=None): + """Move all tensors to device/dtype.""" + kw = {"device": device} + if dtype is not None: + kw["dtype"] = dtype + return SharpnessData( + basis=self.basis.to(**kw), + mean=self.mean.to(**kw), + pc1_std=self.pc1_std, + sign=self.sign, + ) + + +def _gaussian_kernel_2d(kernel_size, sigma): + """Create a 2D Gaussian kernel [1, 1, K, K] for F.conv2d.""" + ax = torch.arange(kernel_size, dtype=torch.float32) - (kernel_size - 1) / 2.0 + gauss = torch.exp(-0.5 * (ax / sigma) ** 2) + kernel_1d = gauss / gauss.sum() + kernel_2d = kernel_1d.unsqueeze(1) @ kernel_1d.unsqueeze(0) # outer product + return kernel_2d.unsqueeze(0).unsqueeze(0) # [1, 1, K, K] + + +def _apply_gaussian_blur(images, blur_sigma): + """Apply Gaussian blur to a batch of images [B, C, H, W]. + + Returns blurred images on same device/dtype as input. + blur_sigma=0 returns input unchanged. + """ + if blur_sigma < 1e-6: + return images + + # Kernel size: 6*sigma rounded up to odd + kernel_size = int(math.ceil(blur_sigma * 6)) | 1 # ensure odd + kernel_size = max(kernel_size, 3) + + kernel = _gaussian_kernel_2d(kernel_size, blur_sigma) + kernel = kernel.to(device=images.device, dtype=images.dtype) + + pad = kernel_size // 2 + B, C, H, W = images.shape + # Apply per-channel via groups + images_grouped = images.reshape(B * C, 1, H, W) + blurred = F.conv2d(images_grouped, kernel, padding=pad) + return blurred.reshape(B, C, H, W) + + +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 + 3. VAE encode → patchify → average patches → [64] vector + 4. PCA on all vectors → extract PC1 (+ PC2) + 5. Determine sign: positive strength = sharper + 6. Compute pc1_std from spread of PC1 scores + + 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 + 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 + 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] + + # Convert BCHW → BHWC for ComfyUI VAE + imgs_bhwc = batch.permute(0, 2, 3, 1).contiguous().cpu() + + # VAE encode → [B, 16, H/8, W/8] + latent = vae.encode(imgs_bhwc[:, :, :, :3]) + + # 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) + + pbar.update(actual_batch) + + # Stack all vectors: [N, 64] + X = torch.stack(vectors, dim=0).float() + blur_labels_t = torch.tensor(blur_labels, dtype=torch.float32) + print(f"[LCS Sharpness Calibration] Collected {X.shape[0]} vectors of dimension {X.shape[1]}") + + # Step 4: PCA + print("[LCS Sharpness Calibration] Computing PCA...") + mean = X.mean(dim=0) # [64] + X_centered = X - mean + U, S, Vh = torch.linalg.svd(X_centered, full_matrices=False) + # Top 2 components + basis = Vh[:2].T # [64, 2] + + # Variance explained + total_var = (S ** 2).sum() + explained = (S[:2] ** 2) / total_var + print(f"[LCS Sharpness Calibration] PC1: {explained[0]:.1%}, PC2: {explained[1]:.1%} ({(explained[0]+explained[1]):.1%} total)") + + # Step 5: Determine sign convention + # Project all vectors onto PC1 + pc1_scores = X_centered @ basis[:, 0] # [N] + + # Correlate PC1 score with blur sigma + # If positive correlation (more blur = higher score), flip sign + correlation = torch.corrcoef(torch.stack([pc1_scores, blur_labels_t]))[0, 1] + sign = -1.0 if correlation > 0 else 1.0 + print(f"[LCS Sharpness Calibration] PC1-blur correlation: {correlation:.3f} → sign = {sign:+.0f}") + + # Step 6: Compute pc1_std + pc1_std = float(pc1_scores.std()) + print(f"[LCS Sharpness Calibration] PC1 std: {pc1_std:.4f}") + print(f"[LCS Sharpness Calibration] Complete! Basis shape: {basis.shape}") + + return SharpnessData( + basis=basis, + mean=mean, + pc1_std=pc1_std, + sign=sign, + ) diff --git a/nodes/__init__.py b/nodes/__init__.py index b3be44b..eaa0ffd 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -3,6 +3,7 @@ from .calibrate import LCSLoadData from .intervene import LCSColorIntervene, LCSColorBatch, LCSToneAdjust from .observe import LCSPreviewColors, LCSStepObserver +from .sharpen import LCSSharpnessCalibrate, LCSSharpnessIntervene NODE_CLASS_MAPPINGS = { "LCSLoadData": LCSLoadData, @@ -11,6 +12,8 @@ NODE_CLASS_MAPPINGS = { "LCSToneAdjust": LCSToneAdjust, "LCSPreviewColors": LCSPreviewColors, "LCSStepObserver": LCSStepObserver, + "LCSSharpnessCalibrate": LCSSharpnessCalibrate, + "LCSSharpnessIntervene": LCSSharpnessIntervene, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -20,4 +23,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "LCSToneAdjust": "LCS Tone Adjust", "LCSPreviewColors": "LCS Preview Colors", "LCSStepObserver": "LCS Step Observer", + "LCSSharpnessCalibrate": "LCS Sharpness Calibrate", + "LCSSharpnessIntervene": "LCS Sharpness Intervene", } diff --git a/nodes/sharpen.py b/nodes/sharpen.py new file mode 100644 index 0000000..2c0353f --- /dev/null +++ b/nodes/sharpen.py @@ -0,0 +1,195 @@ +"""Sharpness nodes: LCSSharpnessCalibrate and LCSSharpnessIntervene.""" + +import os + +import torch +import torch.nn.functional as F +from comfy_api.latest import io +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 + +SHARPNESS_DATA = io.Custom("SHARPNESS_DATA") +DATA_DIR = os.path.join(os.path.dirname(os.path.dirname(__file__)), "data") + + +def _save_sharpness(data: SharpnessData, path: str): + """Save SharpnessData to safetensors file.""" + os.makedirs(os.path.dirname(path), exist_ok=True) + save_file({ + "basis": data.basis.contiguous(), + "mean": data.mean.contiguous(), + "pc1_std": torch.tensor([data.pc1_std]), + "sign": torch.tensor([data.sign]), + }, path) + + +def _load_sharpness(path: str) -> SharpnessData: + """Load SharpnessData from safetensors file.""" + d = load_file(path) + return SharpnessData( + basis=d["basis"], + mean=d["mean"], + pc1_std=float(d["pc1_std"].item()), + sign=float(d["sign"].item()), + ) + + +class LCSSharpnessCalibrate(io.ComfyNode): + """Calibrate the sharpness subspace for a VAE. + + Generates blur stimuli at varying sigma levels, VAE-encodes them, + and runs PCA to find the sharpness direction in 64D patch space. + Result is cached per-VAE fingerprint. + """ + + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="LCSSharpnessCalibrate", + display_name="LCS Sharpness Calibrate", + category="LCS/calibration", + description="Auto-calibrate and cache sharpness subspace data per-VAE", + inputs=[ + io.Vae.Input("vae", tooltip="VAE model (calibration is cached per-VAE)"), + ], + outputs=[ + SHARPNESS_DATA.Output(display_name="sharpness_data"), + ], + ) + + @classmethod + def execute(cls, vae) -> io.NodeOutput: + fp = vae_fingerprint(vae) + cache_path = os.path.join(DATA_DIR, f"sharpness_{fp}.safetensors") + + if os.path.exists(cache_path): + data = _load_sharpness(cache_path) + else: + data = calibrate_sharpness(vae) + _save_sharpness(data, cache_path) + + return io.NodeOutput(data) + + +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. + """ + def post_cfg_fn(args): + denoised = args["denoised"] # [B, 16, H, W] in process_in space + sigma = args["sigma"] + + # Step gating + sigmas = args["model_options"]["transformer_options"]["sample_sigmas"] + step_index = _find_step_index(sigma, sigmas) + + if step_index < start_step or step_index > end_step: + return denoised + + device = denoised.device + dtype = denoised.dtype + + sd = sharpness_data.to(device, dtype) + B_mat = sd.basis # [64, K] + mu = sd.mean # [64] + + # Convert from process_in to raw VAE space + raw = denoised / SCALE_FACTOR + SHIFT_FACTOR # [B, 16, H, W] + + # Patchify + patches, h_len, w_len = patchify(raw) # [B, L, 64] + + # Project to sharpness subspace + projection = (patches - mu) @ B_mat # [B, L, K] + + # 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] + + # 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 post_cfg_fn + + +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. + Positive strength = sharper, negative = blurrier. + """ + + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="LCSSharpnessIntervene", + display_name="LCS Sharpness Intervene", + category="LCS/intervention", + description="Control sharpness during FLUX generation (positive = sharper, negative = blurrier)", + inputs=[ + io.Model.Input("model"), + SHARPNESS_DATA.Input("sharpness_data", tooltip="Calibration data from LCSSharpnessCalibrate"), + io.Float.Input("strength", default=0.0, min=-2.0, max=2.0, step=0.05, + tooltip="Sharpness strength (>0 = sharper, <0 = blurrier, 0 = no change)"), + io.Int.Input("start_step", default=5, min=0, max=50, + tooltip="First step to apply sharpness intervention"), + io.Int.Input("end_step", default=15, min=0, max=50, + tooltip="Last step to apply sharpness intervention"), + io.Mask.Input("mask", optional=True, + tooltip="Optional mask for localized sharpness control"), + ], + outputs=[ + io.Model.Output(display_name="model"), + ], + ) + + @classmethod + def execute(cls, model, sharpness_data, strength, start_step, end_step, + mask=None) -> io.NodeOutput: + m = model.clone() + # Skip hook when strength is zero (true no-op) + if abs(strength) > 1e-6: + hook = _build_sharpness_fn(sharpness_data, strength, start_step, end_step, mask) + m.set_model_sampler_post_cfg_function(hook) + return io.NodeOutput(m)