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.
This commit is contained in:
facok
2026-03-20 13:01:23 +08:00
parent 28fa9450e0
commit cf9edf0239
4 changed files with 372 additions and 0 deletions
+3
View File
@@ -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,
]
+169
View File
@@ -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,
)
+5
View File
@@ -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",
}
+195
View File
@@ -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)