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:
@@ -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,
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user