Fix sharpness calibration: use same base images for all blur levels
- Generate base images once upfront, then apply all blur levels to the SAME images (was generating different random images per blur level) - Add separable convolution for kernel_size > 15 (O(2k) vs O(k²)) - Add kernel caching to avoid recomputation - Add per-blur-level progress logging This matches SubspaceLab's methodology and ensures PCA receives properly paired stimuli (same image × multiple blur levels).
This commit is contained in:
+64
-28
@@ -2,6 +2,7 @@
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
@@ -36,18 +37,29 @@ class SharpnessData:
|
||||
)
|
||||
|
||||
|
||||
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]
|
||||
# Cache for Gaussian kernels to avoid recomputation
|
||||
_gaussian_kernel_cache: Dict[Tuple[int, float, torch.dtype, torch.device], torch.Tensor] = {}
|
||||
|
||||
|
||||
def _apply_gaussian_blur(images, blur_sigma):
|
||||
def _gaussian_kernel_1d(kernel_size: int, sigma: float, device: torch.device, dtype: torch.dtype) -> torch.Tensor:
|
||||
"""Create a 1D Gaussian kernel with caching."""
|
||||
cache_key = (kernel_size, sigma, dtype, device)
|
||||
if cache_key in _gaussian_kernel_cache:
|
||||
return _gaussian_kernel_cache[cache_key]
|
||||
|
||||
x = torch.arange(kernel_size, dtype=dtype, device=device) - kernel_size // 2
|
||||
gauss = torch.exp(-x.pow(2) / (2 * sigma ** 2))
|
||||
gauss = gauss / gauss.sum()
|
||||
|
||||
_gaussian_kernel_cache[cache_key] = gauss
|
||||
return gauss
|
||||
|
||||
|
||||
def _apply_gaussian_blur(images: torch.Tensor, blur_sigma: float) -> torch.Tensor:
|
||||
"""Apply Gaussian blur to a batch of images [B, C, H, W].
|
||||
|
||||
Uses separable convolution for kernel_size > 15 (O(2k) vs O(k²)).
|
||||
|
||||
Returns blurred images on same device/dtype as input.
|
||||
blur_sigma=0 returns input unchanged.
|
||||
"""
|
||||
@@ -55,26 +67,45 @@ def _apply_gaussian_blur(images, blur_sigma):
|
||||
return images
|
||||
|
||||
# Kernel size: 6*sigma rounded up to odd
|
||||
kernel_size = int(math.ceil(blur_sigma * 6)) | 1 # ensure odd
|
||||
kernel_size = int(math.ceil(blur_sigma * 6)) | 1
|
||||
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)
|
||||
device, dtype = images.device, images.dtype
|
||||
|
||||
if kernel_size <= 15:
|
||||
# Direct 2D convolution for small kernels
|
||||
kernel_1d = _gaussian_kernel_1d(kernel_size, blur_sigma, device, dtype)
|
||||
kernel_2d = kernel_1d[:, None] @ kernel_1d[None, :] # outer product
|
||||
kernel = kernel_2d.view(1, 1, kernel_size, kernel_size).expand(C, 1, -1, -1)
|
||||
pad = kernel_size // 2
|
||||
images_grouped = images.view(B * C, 1, H, W)
|
||||
blurred = F.conv2d(images_grouped, kernel, padding=pad, groups=1)
|
||||
return blurred.view(B, C, H, W)
|
||||
else:
|
||||
# Separable convolution: apply 1D Gaussian in X then Y direction
|
||||
# This reduces O(k²) to O(2k) operations
|
||||
gauss_1d = _gaussian_kernel_1d(kernel_size, blur_sigma, device, dtype)
|
||||
pad = kernel_size // 2
|
||||
|
||||
# Horizontal pass: convolve along W dimension
|
||||
kernel_h = gauss_1d.view(1, 1, 1, kernel_size).expand(C, 1, 1, kernel_size)
|
||||
blurred = F.conv2d(images, kernel_h, padding=(0, pad), groups=C)
|
||||
|
||||
# Vertical pass: convolve along H dimension
|
||||
kernel_v = gauss_1d.view(1, 1, kernel_size, 1).expand(C, 1, kernel_size, 1)
|
||||
blurred = F.conv2d(blurred, kernel_v, padding=(pad, 0), groups=C)
|
||||
|
||||
return blurred
|
||||
|
||||
|
||||
def calibrate_sharpness(vae, num_samples=64, image_size=512,
|
||||
blur_levels=(0, 1, 2, 4, 8, 16), batch_size=8):
|
||||
def calibrate_sharpness(vae, num_samples: int = 64, image_size: int = 512,
|
||||
blur_levels: Tuple[float, ...] = (0, 1, 2, 4, 8, 16),
|
||||
batch_size: int = 8) -> SharpnessData:
|
||||
"""Compute sharpness subspace data (PCA basis, mean, sign) from FLUX VAE.
|
||||
|
||||
1. Generate num_samples random noise images (spatial detail needed for blur)
|
||||
2. For each blur level, apply Gaussian blur per-batch
|
||||
2. For each blur level, apply Gaussian blur to the SAME images
|
||||
3. VAE encode → patchify → average patches → [64] vector
|
||||
4. PCA on all vectors → extract PC1 (+ PC2)
|
||||
5. Determine sign: positive strength = sharper
|
||||
@@ -88,22 +119,27 @@ def calibrate_sharpness(vae, num_samples=64, image_size=512,
|
||||
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)}")
|
||||
|
||||
# Deterministic seed for reproducibility across runs
|
||||
# Deterministic seed for reproducibility
|
||||
rng = torch.Generator().manual_seed(42)
|
||||
|
||||
# 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.
|
||||
# Step 1: Generate all base images upfront [num_samples, 3, H, W]
|
||||
# Use noise (not solid color) — spatial detail is needed for blur to have effect.
|
||||
# Per-batch generation to limit memory.
|
||||
print(f"[LCS Sharpness Calibration] Generating {num_samples} noise images...")
|
||||
base_images = torch.rand(num_samples, 3, image_size, image_size, generator=rng)
|
||||
|
||||
# Step 2+3: For each blur level, apply blur to ALL base images, then encode
|
||||
vectors = []
|
||||
blur_labels = [] # track blur sigma per vector for sign determination
|
||||
pbar = comfy.utils.ProgressBar(total_images)
|
||||
|
||||
for blur_sigma in blur_levels:
|
||||
for batch_start in range(0, num_samples, batch_size):
|
||||
actual_batch = min(batch_size, num_samples - batch_start)
|
||||
print(f"[LCS Sharpness Calibration] Processing blur σ={blur_sigma}...")
|
||||
|
||||
# Generate random noise images per-batch [B, 3, H, W]
|
||||
batch = torch.rand(actual_batch, 3, image_size, image_size, generator=rng)
|
||||
for batch_start in range(0, num_samples, batch_size):
|
||||
batch_end = min(batch_start + batch_size, num_samples)
|
||||
batch = base_images[batch_start:batch_end] # [B, 3, H, W]
|
||||
actual_batch = batch.shape[0]
|
||||
|
||||
# Apply blur (no-op for sigma=0)
|
||||
blurred = _apply_gaussian_blur(batch, blur_sigma)
|
||||
|
||||
Reference in New Issue
Block a user