diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..083fe14 --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ +__pycache__/ +docs/ diff --git a/README.md b/README.md index 3b4f809..4d4012b 100644 --- a/README.md +++ b/README.md @@ -1,6 +1,6 @@ # ComfyUI-DiversityBoost -Restore composition diversity for distilled diffusion models. Training-free, single-step frequency-domain phase injection. +Restore composition diversity for distilled diffusion models. Training-free, single-step, zero model modification. > [中文版 README](README_zh.md) @@ -12,20 +12,22 @@ Root cause: distillation freezes the spatial distribution of token norms across ## The Fix -DiversityBoost injects seed-dependent low-frequency phase from the initial noise into the model's denoised prediction at **step 0 only**, via a post-CFG hook. +DiversityBoost applies two mechanisms at **step 0 only**, via a single post-CFG hook: -In the frequency domain, **amplitude** encodes energy distribution (naturalness), while **phase** encodes spatial arrangement (composition). By rotating only low-frequency phase — while attenuating high-frequency amplitude to prevent deformities — different seeds produce genuinely different compositions again. +1. **HF attenuation** (Butterworth LPF) — erases the high-frequency spatial anchoring from the model's fully-committed prediction, producing a blurry "composition sketch" +2. **DCT composition push** — applies a random low-frequency spatial field that redistributes energy across the latent, nudging the model toward different compositions + +The model then freely reconstructs coherent details at subsequent steps, with per-seed noise driving different reconstruction paths. Zero model modification. Zero training. One node. ## How It Works -1. **FFT** the model's step-0 prediction and the initial noise -2. Compute a **shared rotation field** (channel-mean phasors) — preserves inter-channel phase exactly (no color fringing) -3. Apply a **6th-order Butterworth low-pass filter** — only composition-scale frequencies are touched; object-scale details are protected -4. **Tanh soft-cap** per-bin rotation — decouples diversity strength from worst-case rotation risk -5. **Attenuate high-frequency amplitude** — prevents "frequency shearing" (body moves but fingers stay pinned) -6. **IFFT** back to spatial domain +1. Convert the model's step-0 prediction to raw latent space +2. **Butterworth LPF** in frequency domain — attenuate high-frequency amplitude (6th-order, elliptical, resolution-independent) +3. **DCT spatial field** — synthesize a random 4×4 low-frequency field (zero DC, pink noise weighted), normalize to unit std, scale by strength +4. **Multiplicative push** — `blurred × (1 + field)`, clamped to prevent dead zones +5. Convert back Step 0 only. The model's own attractor handles the rest. @@ -54,51 +56,39 @@ Default settings work well. No tuning needed for most use cases. | Parameter | Default | Range | Description | |-----------|---------|-------|-------------| -| **strength** | 1.0 | 0.0 – 1.0 | Phase rotation strength. 1.0 is safe with the default max_rotation cap. | -| **n_periods** | 2 | 1 – 10 | Max spatial periods to affect. Resolution-independent. | -| **max_rotation** | 1.5708 (π/2) | 0.0 – 3.14 | Per-bin rotation budget (radians). Caps worst-case rotation via tanh. | +| **strength** | 0.5 | 0.0 – 2.0 | Composition push amplitude. 0 = cleanup only, 0.5 = moderate, 1.0 = strong. | +| **clamp** | 1.0 | 0.1 – 3.0 | Upper bound for the multiplicative scale factor. Scale is clamped to [0.1, 1+clamp]. Higher = allow stronger push. | +| **noise_type** | pink | pink / white / blue | Frequency spectrum of random DCT coefficients. Pink boosts low-freq composition modes (recommended). | +| **n_periods** | 2 | 1 – 10 | Butterworth cutoff — spatial periods to preserve. Lower = more HF erased = more diversity. | +| **dc_preserve** | 0.0 | 0.0 – 1.0 | DC amplitude preservation. 0 = max diversity (tone varies per seed). 1 = preserve original brightness. | | **energy_compensate** | False | — | Rescale output RMS to match original. Off by default. | ### n_periods Guide -| Value | Effect | -|-------|--------| -| 1 | Ultra-conservative (global balance only) | -| 2 | Conservative (default, recommended) | -| 3 | Balanced (moderate diversity) | -| 4 | Aggressive (more diversity, higher risk) | - -The cutoff adapts to image resolution automatically: `freq_cutoff = n_periods / max(H, W)`. This means `n_periods=2` always affects the same composition-scale frequencies regardless of whether you're generating 512x512 or 2048x2048. - -### max_rotation Presets +n_periods sets the Butterworth -3dB point at FFT bin N. Frequencies above this are strongly attenuated. | Value | Effect | |-------|--------| -| 1.00 | Conservative — less diversity, very safe | -| 1.57 (π/2) | Balanced (default, recommended) | -| 2.00 | Aggressive — more diversity, some risk | -| 0.00 | Disabled — original uncapped behavior | +| 1 | Most aggressive — erases nearly all spatial structure including most DCT composition signal | +| 2 | Recommended — clean frequency gap between DCT composition modes (bin ~1.5) and object-scale details (bin ~5+) | +| 3 | Moderate — preserves more mid-frequency detail | +| 4+ | Mild — less HF erased, less diversity | -## DiversityBoost vs Dummy Token +### strength Guide -Both aim to restore diversity in distilled models, but they work at fundamentally different levels: - -| | DiversityBoost | Dummy Token | -|---|---|---| -| **Mechanism** | Rotate low-frequency phase in frequency domain | Add/modify padding tokens to shift attention context | -| **Target** | Spatial arrangement (composition skeleton) | Global attention bias (indirect) | -| **Composition change** | Direct and controllable (precise rotation angles) | Indirect, random (butterfly effect) | -| **Diversity source** | Each seed's unique noise phase | Random padding token content | -| **Safety** | Amplitude exactly preserved; Butterworth + tanh double-bounded | No mathematical guarantees | -| **Prompt adherence** | Unaffected (operates on spatial structure, not semantics) | May degrade (alters attention distribution) | - -In short: DiversityBoost operates precisely on composition-scale phase in the frequency domain, with mathematically bounded safety guarantees. Dummy token injects perturbation at the token level — the effect is indirect and unpredictable. +| Value | Effect | +|-------|--------| +| 0.0 | HF cleanup only (no composition push) | +| 0.3 | Subtle composition variation | +| 0.5 | Moderate (default, recommended) | +| 1.0 | Strong composition changes | ## Tips -- **Start with defaults** — strength=1.0, n_periods=2, max_rotation=π/2 is a safe baseline -- **Want more diversity?** Raise `n_periods` first (more frequency bins affected), then `max_rotation` (higher per-bin ceiling) -- **Compatible** with other model patches (ComfyUI-LCS color control, ControlNet, etc.) — operates on a different hook +- **Start with defaults** — strength=0.5, n_periods=2, noise_type=pink is a safe baseline +- **Want more diversity?** Raise `strength`. Keep `n_periods=2` — lowering to 1 kills most DCT composition signal +- **Want cleanup only?** Set `strength=0` — pure HF attenuation, no composition push +- **Compatible** with other model patches (ControlNet, etc.) — operates on a different hook ## Tested Models diff --git a/README_zh.md b/README_zh.md index 60f1318..ea4eaf0 100644 --- a/README_zh.md +++ b/README_zh.md @@ -1,6 +1,6 @@ # ComfyUI-DiversityBoost -恢复蒸馏扩散模型的构图多样性。免训练、单步频域相位注入。 +恢复蒸馏扩散模型的构图多样性。免训练、单步执行、零模型修改。 > [English README](README.md) @@ -12,20 +12,22 @@ ## 解决方案 -DiversityBoost 在**第 0 步**通过 post-CFG 钩子,将初始噪声中 seed 特有的低频相位注入模型的去噪预测。 +DiversityBoost 在**第 0 步**通过单个 post-CFG 钩子应用两个机制: -在频域中,**振幅**编码能量分布(自然性),**相位**编码空间排列(构图)。只旋转低频相位——同时衰减高频振幅以防止畸变——不同 seed 就能重新产生不同构图。 +1. **高频衰减**(Butterworth 低通滤波)——擦除模型预测中的高频空间锚定,产生模糊的"构图草稿" +2. **DCT 构图推动**——施加一个随机低频空间场,重新分配潜空间的能量分布,引导模型朝不同构图方向重建 + +后续步骤中,模型在每个 seed 的噪声驱动下自由重建连贯的细节,产生不同的构图。 零模型修改。零训练。一个节点。 ## 工作原理 -1. 对模型第 0 步预测和初始噪声做 **FFT** -2. 计算**共享旋转场**(通道均值相量)——精确保持通道间相位关系(无色差) -3. 应用 **6 阶 Butterworth 低通滤波器**——只影响构图尺度频率,物体尺度细节受保护 -4. **tanh 软限幅**逐 bin 旋转——将多样性强度与最坏情况旋转风险解耦 -5. **衰减高频振幅**——防止"频率剪切"(身体移动但手指钉在原位) -6. **IFFT** 回到空间域 +1. 将模型第 0 步预测转换到原始潜空间 +2. **Butterworth 低通滤波**——在频域中衰减高频振幅(6 阶、椭圆归一化、分辨率无关) +3. **DCT 空间场**——合成随机 4×4 低频场(零 DC、pink 噪声加权),归一化到单位标准差,按 strength 缩放 +4. **乘法推动**——`模糊结果 × (1 + field)`,钳位防止死区 +5. 转换回原空间 只在第 0 步执行。后续步骤由模型自身的吸引子接管。 @@ -54,51 +56,39 @@ MODEL → [Diversity Boost] → MODEL → KSampler | 参数 | 默认值 | 范围 | 说明 | |------|--------|------|------| -| **strength** | 1.0 | 0.0 – 1.0 | 相位旋转强度。在默认 max_rotation 限幅下,1.0 是安全的。 | -| **n_periods** | 2 | 1 – 10 | 影响的最大空间周期数。自动适配分辨率。 | -| **max_rotation** | 1.5708 (π/2) | 0.0 – 3.14 | 逐 bin 旋转预算(弧度)。通过 tanh 限制最坏情况旋转。 | +| **strength** | 0.5 | 0.0 – 2.0 | 构图推动幅度。0 = 仅清理,0.5 = 适中,1.0 = 强烈。 | +| **clamp** | 1.0 | 0.1 – 3.0 | 乘法缩放因子的上限。scale 被钳位到 [0.1, 1+clamp]。越高 = 允许更强的推动。 | +| **noise_type** | pink | pink / white / blue | 随机 DCT 系数的频谱类型。pink 增强低频构图模式(推荐)。 | +| **n_periods** | 2 | 1 – 10 | Butterworth 截止频率——保留的空间周期数。越低 = 擦除越多高频 = 更多多样性。 | +| **dc_preserve** | 0.0 | 0.0 – 1.0 | DC 振幅保留。0 = 最大多样性(色调随 seed 变化),1 = 保留原始亮度。 | | **energy_compensate** | False | — | 将输出 RMS 缩放至与原始预测一致。默认关闭。 | ### n_periods 参考 -| 值 | 效果 | -|----|------| -| 1 | 最保守(仅全局平衡) | -| 2 | 保守(默认,推荐) | -| 3 | 均衡(适度多样性) | -| 4 | 激进(更多多样性,有风险) | - -截止频率自动适配图像分辨率:`freq_cutoff = n_periods / max(H, W)`。无论生成 512x512 还是 2048x2048,`n_periods=2` 始终影响相同的构图尺度频率。 - -### max_rotation 预设 +n_periods 设定 Butterworth 滤波器的 -3dB 点在 FFT 第 N 个 bin。高于此频率的成分被强烈衰减。 | 值 | 效果 | |----|------| -| 1.00 | 保守——多样性较低,非常安全 | -| 1.57 (π/2) | 均衡(默认,推荐) | -| 2.00 | 激进——更多多样性,有一定风险 | -| 0.00 | 禁用——无限幅(原始行为) | +| 1 | 最激进——擦除几乎所有空间结构,包括大部分 DCT 构图信号 | +| 2 | 推荐——DCT 构图模式(bin ~1.5)与物体尺度细节(bin ~5+)之间有干净的频率间隙 | +| 3 | 适度——保留更多中频细节 | +| 4+ | 温和——擦除较少高频,多样性较低 | -## DiversityBoost vs Dummy Token +### strength 参考 -两者都旨在恢复蒸馏模型的多样性,但作用层面完全不同: - -| | DiversityBoost | Dummy Token | -|---|---|---| -| **机制** | 在频域旋转低频相位 | 添加/修改 padding token 偏移 attention 上下文 | -| **作用目标** | 空间排列(构图骨架) | 全局 attention 偏置(间接) | -| **构图改变** | 直接、可控(精确旋转角度) | 间接、随机(蝴蝶效应) | -| **多样性来源** | 每个 seed 独有的噪声相位 | 随机 padding token 内容 | -| **安全性** | 振幅精确保持;Butterworth + tanh 双重限幅 | 无数学保证 | -| **prompt 遵循度** | 不受影响(操作空间结构,不涉及语义) | 可能下降(改变 attention 分布) | - -简单说:DiversityBoost 在频域精确操作构图尺度的相位,有数学边界保证安全性。Dummy token 在 token 层面注入扰动,效果间接且不可预测。 +| 值 | 效果 | +|----|------| +| 0.0 | 仅高频清理(无构图推动) | +| 0.3 | 微妙的构图变化 | +| 0.5 | 适中(默认,推荐) | +| 1.0 | 强烈的构图变化 | ## 使用建议 -- **从默认值开始**——strength=1.0、n_periods=2、max_rotation=π/2 是安全的基线 -- **想要更多多样性?** 先提高 `n_periods`(影响更多频率 bin),再提高 `max_rotation`(提高逐 bin 上限) -- **兼容**其他模型补丁(ComfyUI-LCS 颜色控制、ControlNet 等)——使用不同的钩子,互不干扰 +- **从默认值开始**——strength=0.5、n_periods=2、noise_type=pink 是安全的基线 +- **想要更多多样性?** 提高 `strength`。保持 `n_periods=2`——降到 1 会杀掉大部分 DCT 构图信号 +- **只需清理?** 设置 `strength=0`——纯高频衰减,无构图推动 +- **兼容**其他模型补丁(ControlNet 等)——使用不同的钩子,互不干扰 ## 已测试模型 diff --git a/__init__.py b/__init__.py index 68107dc..ae04173 100644 --- a/__init__.py +++ b/__init__.py @@ -1,20 +1,18 @@ """ComfyUI-DiversityBoost: Restore composition diversity for distilled diffusion models. -Injects seed-dependent low-frequency phase from initial noise into the model's -denoised prediction, making different seeds produce different compositions instead -of identical layouts. Training-free, single-step (step 0 only), zero model modification. +HF attenuation + DCT composition push at step 0. Training-free, single-step, +zero model modification. """ -# V3 ComfyExtension entry point from comfy_api.latest import ComfyExtension, io -from .node import DiversityBoost +from .core_node import DiversityBoostCore class DiversityBoostExtension(ComfyExtension): """V3 ComfyExtension providing the DiversityBoost node.""" async def get_node_list(self) -> list[type[io.ComfyNode]]: - return [DiversityBoost] + return [DiversityBoostCore] async def comfy_entrypoint() -> DiversityBoostExtension: @@ -24,11 +22,11 @@ async def comfy_entrypoint() -> DiversityBoostExtension: # V2 backward compatibility NODE_CLASS_MAPPINGS = { - "DiversityBoost": DiversityBoost, + "DiversityBoostCore": DiversityBoostCore, } NODE_DISPLAY_NAME_MAPPINGS = { - "DiversityBoost": "Diversity Boost", + "DiversityBoostCore": "Diversity Boost", } __all__ = [ @@ -36,4 +34,5 @@ __all__ = [ "NODE_DISPLAY_NAME_MAPPINGS", "DiversityBoostExtension", "comfy_entrypoint", + "DiversityBoostCore", ] diff --git a/core.py b/core.py new file mode 100644 index 0000000..3ce2bc3 --- /dev/null +++ b/core.py @@ -0,0 +1,232 @@ +"""DiversityBoost core — HF attenuation + DCT composition push. + +Restores composition diversity lost during distillation via two mechanisms +applied in a single post-cfg hook at step 0: +1. HF attenuation (Butterworth LPF) — blurry "composition sketch" +2. DCT composition push — multiplicative low-freq spatial field +""" + +import logging +import math +from functools import lru_cache + +import torch + +from .sampling import ( + denoised_to_raw, + raw_to_denoised, + find_step_index, + unpack_video_if_needed, + repack_video_if_needed, +) + +log = logging.getLogger("ComfyUI-DiversityBoost") + + +# --------------------------------------------------------------------------- +# 2D DCT basis (orthonormal, cached) +# --------------------------------------------------------------------------- + +@lru_cache(maxsize=4) +def _build_dct_basis_2d(H, W, n_modes_h=4, n_modes_w=4): + """Build orthonormal 2D DCT-II basis matrix [H*W, n_modes_h*n_modes_w]. + + Same matrix for analysis (projection) and synthesis (reconstruction): + synthesis: field = basis @ coeffs + """ + def _dct1d(N, n_modes): + n = torch.arange(N, dtype=torch.float64) + k = torch.arange(n_modes, dtype=torch.float64) + phi = torch.cos(math.pi * k[None, :] * (n[:, None] + 0.5) / N) + norm = torch.full((n_modes,), math.sqrt(2.0 / N), dtype=torch.float64) + norm[0] = 1.0 / math.sqrt(N) + return phi * norm[None, :] + + phi_h = _dct1d(H, n_modes_h) + phi_w = _dct1d(W, n_modes_w) + basis_2d = torch.einsum('hu,wv->hwuv', phi_h, phi_w) + basis_2d = basis_2d.reshape(H * W, n_modes_h * n_modes_w) + return basis_2d.float() + + +# --------------------------------------------------------------------------- +# Noise frequency weights +# --------------------------------------------------------------------------- + +def _build_pink_weights(n_h, n_w): + """1/f amplitude weights: lower frequencies dominate.""" + weights = [] + for u in range(n_h): + for v in range(n_w): + freq_sq = u * u + v * v + weights.append(0.0 if freq_sq == 0 else 1.0 / (freq_sq ** 0.25)) + return torch.tensor(weights, dtype=torch.float32) + + +def _build_blue_weights(n_h, n_w): + """f-proportional weights: higher frequencies dominate.""" + weights = [] + for u in range(n_h): + for v in range(n_w): + freq_sq = u * u + v * v + weights.append(0.0 if freq_sq == 0 else (freq_sq ** 0.25)) + return torch.tensor(weights, dtype=torch.float32) + + +def _build_noise_weights(noise_type, n_h, n_w): + """Build frequency weights for given noise type, or None for white.""" + if noise_type == "pink": + return _build_pink_weights(n_h, n_w) + elif noise_type == "blue": + return _build_blue_weights(n_h, n_w) + return None + + +# --------------------------------------------------------------------------- +# Butterworth LPF +# --------------------------------------------------------------------------- + +def _build_freq_mask(H, W, n_periods, device): + """Elliptical Butterworth LPF mask for rfft2 output [1, 1, H, W//2+1].""" + freq_y = torch.fft.fftfreq(H, device=device).unsqueeze(1) + freq_x = torch.fft.rfftfreq(W, device=device).unsqueeze(0) + r_norm = torch.sqrt((freq_y * H / n_periods) ** 2 + + (freq_x * W / n_periods) ** 2) + mask = 1.0 / torch.sqrt(1.0 + r_norm.pow(12)) + mask[0, 0] = 0.0 + return mask.unsqueeze(0).unsqueeze(0) + + +# --------------------------------------------------------------------------- +# Combined hook: HF attenuation → DCT composition push +# --------------------------------------------------------------------------- + +def build_diversity_fn(strength=0.5, clamp_val=1.0, noise_type="pink", + n_periods=2, dc_preserve=0.0, + energy_compensate=False): + """Build a post_cfg_function that attenuates HF then applies DCT push. + + Execution order within step 0: + 1. Convert to raw latent space + 2. HF attenuation (Butterworth LPF) + 3. DCT composition push (4×4 random spatial field, multiplicative) + 4. Convert back + + The push operates on the blurred x0_hat in latent pixel space [B,C,H,W], + not token space. This ensures nothing downstream can erase the signal. + + Parameters: + strength: push amplitude (0-2). 0 = no push (cleanup only). + clamp_val: safety clamp for field values. + noise_type: "pink", "white", or "blue" frequency weighting. + n_periods: Butterworth cutoff (spatial periods to preserve). + dc_preserve: DC amplitude preservation [0, 1]. + energy_compensate: rescale output RMS to match original prediction. + """ + n_modes_h, n_modes_w = 4, 4 + n_modes = n_modes_h * n_modes_w + freq_weights = _build_noise_weights(noise_type, n_modes_h, n_modes_w) + + state = { + "amp_scale": None, + "basis_2d": None, + "freq_weights": None, + "cached_hw": None, + } + + def diversity_hook(args): + denoised = args["denoised"] + sigma = args["sigma"] + model = args["model"] + model_options = args["model_options"] + + # --- Step 0 only --- + sample_sigmas = model_options.get("transformer_options", {}).get("sample_sigmas") + if sample_sigmas is None: + return denoised + step_index = find_step_index(sigma, sample_sigmas) + if step_index != 0: + return denoised + + # --- Unpack video if needed --- + working, pack_info = unpack_video_if_needed(denoised, args) + + # --- Convert to raw space --- + raw_pred = denoised_to_raw(working, model) + B, C, H, W = raw_pred.shape + device = raw_pred.device + orig_dtype = raw_pred.dtype + + # --- Build or reuse cached tensors --- + if state["cached_hw"] != (H, W): + freq_mask = _build_freq_mask(H, W, n_periods, device) + amp = freq_mask.clone() + amp[:, :, 0, 0] = dc_preserve + state["amp_scale"] = amp + if strength > 1e-6: + basis_cpu = _build_dct_basis_2d(H, W, n_modes_h, n_modes_w) + state["basis_2d"] = basis_cpu.to(device=device) + if freq_weights is not None: + state["freq_weights"] = freq_weights.to(device=device) + state["cached_hw"] = (H, W) + amp_scale = state["amp_scale"].to(device=device) + + # --- Step 1: HF attenuation --- + F_pred = torch.fft.rfft2(raw_pred.float()) + F_blurred = F_pred * amp_scale + raw_blurred = torch.fft.irfft2(F_blurred, s=(H, W)) + + # --- Step 2: DCT composition push on blurred result --- + if strength > 1e-6: + coeffs = torch.randn(B, n_modes, device=device, dtype=torch.float32) + coeffs[:, 0] = 0.0 + + if state["freq_weights"] is not None: + coeffs = coeffs * state["freq_weights"] + + field = torch.einsum('nk,bk->bn', state["basis_2d"], coeffs) + field = field.reshape(B, H, W) + + field_std = field.reshape(B, -1).std(dim=1).clamp(min=1e-8) + field = field / field_std[:, None, None] + field = field * strength + + scale = (1.0 + field).clamp(min=0.10, max=1.0 + clamp_val).unsqueeze(1) + raw_new = raw_blurred * scale + else: + raw_new = raw_blurred + scale = None + + # --- Energy compensation --- + if energy_compensate: + pred_rms = raw_pred.float().pow(2).mean(dim=(-2, -1), keepdim=True).sqrt().clamp(min=1e-8) + new_rms = raw_new.pow(2).mean(dim=(-2, -1), keepdim=True).sqrt().clamp(min=1e-8) + raw_new = raw_new * (pred_rms / new_rms) + + # --- Logging --- + if log.isEnabledFor(logging.INFO): + with torch.no_grad(): + delta = (raw_new - raw_pred.float()) + delta_rms = delta.pow(2).mean().sqrt().item() + pred_rms_val = raw_pred.float().pow(2).mean().sqrt().item() + push_info = "" + if scale is not None: + s_flat = scale.squeeze(1) + push_info = ( + f" push=[{s_flat.min().item():.4f}, {s_flat.max().item():.4f}]" + f" strength={strength:.3f} clamp={clamp_val:.2f} noise={noise_type}" + ) + log.info( + "[DiversityBoost] step=0 n_periods=%d dc=%.2f" + "%s shape=%s delta_rms=%.4f pred_rms=%.4f ratio=%.4f", + n_periods, dc_preserve, + push_info, list(raw_pred.shape), + delta_rms, pred_rms_val, + delta_rms / max(pred_rms_val, 1e-8), + ) + + # --- Convert back --- + modified = raw_to_denoised(raw_new, model).to(dtype=orig_dtype) + return repack_video_if_needed(modified, pack_info) + + return diversity_hook diff --git a/core_node.py b/core_node.py new file mode 100644 index 0000000..9232af3 --- /dev/null +++ b/core_node.py @@ -0,0 +1,71 @@ +"""DiversityBoost node — HF attenuation + DCT composition push.""" + +import time + +from comfy_api.latest import io + +from .core import build_diversity_fn + + +class DiversityBoostCore(io.ComfyNode): + """Restore composition diversity for distilled diffusion models. + + Single post-cfg hook at step 0: first attenuates HF amplitude + (Butterworth LPF), then applies a random low-frequency DCT spatial + field to the blurred result. Push runs AFTER cleanup so its signal + cannot be erased by downstream processing. + """ + + @classmethod + def define_schema(cls) -> io.Schema: + return io.Schema( + node_id="DiversityBoostCore", + display_name="Diversity Boost", + category="sampling", + description="Restore composition diversity for distilled models. " + "HF attenuation + DCT composition push at step 0.", + inputs=[ + io.Model.Input("model"), + io.Float.Input("strength", default=0.50, min=0.0, max=2.0, step=0.05, + tooltip="Composition push amplitude. " + "0 = cleanup only. 0.5 = moderate. 1.0 = strong."), + io.Float.Input("clamp", default=1.0, min=0.1, max=3.0, step=0.1, + tooltip="Safety clamp for DCT field values."), + io.Combo.Input("noise_type", + options=["pink", "white", "blue"], + default="pink", + tooltip="Frequency spectrum of random DCT coefficients. " + "pink = stronger composition push (recommended)."), + io.Int.Input("n_periods", default=2, min=1, max=10, step=1, + tooltip="Butterworth cutoff. 2 = preserves DCT signal."), + io.Float.Input("dc_preserve", default=0.0, min=0.0, max=1.0, step=0.1, + tooltip="DC amplitude preservation (0 = max diversity)."), + io.Boolean.Input("energy_compensate", default=False, + tooltip="Rescale output energy to match original."), + ], + outputs=[ + io.Model.Output(display_name="model"), + ], + ) + + @classmethod + def fingerprint_inputs(cls, **kwargs): + return time.time() + + @classmethod + def execute(cls, model, strength, clamp, noise_type, + n_periods, dc_preserve, energy_compensate) -> io.NodeOutput: + m = model.clone() + + m.set_model_sampler_post_cfg_function( + build_diversity_fn( + strength=strength, + clamp_val=clamp, + noise_type=noise_type, + n_periods=n_periods, + dc_preserve=dc_preserve, + energy_compensate=energy_compensate, + ), + ) + + return io.NodeOutput(m) diff --git a/node.py b/node.py deleted file mode 100644 index 9f8f30a..0000000 --- a/node.py +++ /dev/null @@ -1,102 +0,0 @@ -"""DiversityBoost node — ComfyUI node for frequency-domain composition diversity.""" - -import time - -from comfy_api.latest import io - -from .phase_inject import build_phase_injection_fn - - -class DiversityBoost(io.ComfyNode): - """Restore composition diversity lost during distillation. - - Injects seed-dependent low-frequency phase from initial noise into model - prediction, while attenuating high-frequency amplitude to prevent - 'frequency shearing' (limb deformity from spatially incoherent LF/HF). - Operates at step 0 only via post-CFG hook. - """ - - @classmethod - def define_schema(cls) -> io.Schema: - return io.Schema( - node_id="DiversityBoost", - display_name="Diversity Boost", - category="sampling", - description="Restore composition diversity for distilled models. " - "Injects seed-dependent low-frequency phase from initial " - "noise into model prediction. Different seeds produce " - "different compositions instead of identical layouts. " - "Step 0 only, zero risk to model internals.", - inputs=[ - io.Model.Input("model"), - io.Float.Input("strength", default=1.00, min=0.0, max=1.0, step=0.05, - tooltip="Phase rotation strength (gamma). " - "0 = no effect. 1 = full rotation toward noise phase. " - "With max_rotation cap, 1.0 is safe — " - "the cap prevents any bin from over-rotating. " - "To increase diversity further, raise freq_cutoff " - "(more bins) or max_rotation (higher ceiling)."), - io.Int.Input("n_periods", default=2, min=1, max=10, step=1, - tooltip="Max spatial periods to affect (Butterworth LPF). " - "Frequencies with ≤ this many full cycles across " - "the frame are modified; higher frequencies are " - "protected. Resolution-independent: the cutoff " - "adapts automatically to any image size. " - "1 = ultra-conservative (global balance only), " - "2 = conservative (default, recommended), " - "3 = balanced (moderate diversity), " - "4 = aggressive (more diversity, higher risk)."), - io.Float.Input("max_rotation", default=1.5708, min=0.0, max=3.14, step=0.05, - tooltip="Per-bin rotation budget in radians (0 = no cap). " - "Caps the maximum phase rotation at any single " - "frequency bin via tanh soft saturation. " - "Decouples diversity strength from tail risk. " - "Default π/2 ≈ 1.57 is the natural scale " - "(mean |θ| of the noise prior). At this value: " - "composition bins (e.g. bin (1,0)) max 87°, " - "transition bins (e.g. bin (0,3)) max 68°, " - "object bins (e.g. bin (8,0)) max 2°. " - "1.00 = conservative (less diversity, very safe), " - "1.57 = balanced (π/2, recommended), " - "2.00 = aggressive (more diversity, some risk), " - "0.00 = disabled (original uncapped behavior)."), - io.Boolean.Input("energy_compensate", default=False, - tooltip="Rescale output RMS to match original prediction. " - "When hf_preserve < 1, high-freq amplitude is " - "attenuated, reducing total energy. Enable this " - "to compensate by scaling the result back to the " - "original energy level. Off by default."), - io.Float.Input("dc_preserve", default=0.0, min=0.0, max=1.0, step=0.1, - tooltip="DC amplitude preservation (0-1). " - "DC is the image mean (overall brightness/color). " - "0 = let DC be attenuated with HF (model rebuilds " - "brightness freely, often better quality). " - "1 = fully preserve DC (original brightness). " - "Default 0."), - ], - outputs=[ - io.Model.Output(display_name="model"), - ], - ) - - @classmethod - def fingerprint_inputs(cls, **kwargs): - return time.time() - - @classmethod - def execute(cls, model, strength, n_periods, max_rotation, energy_compensate, dc_preserve) -> io.NodeOutput: - m = model.clone() - - if strength > 1e-6: - m.set_model_sampler_post_cfg_function( - build_phase_injection_fn( - strength=strength, - n_periods=n_periods, - max_rotation=max_rotation, - hf_preserve=0.0, - energy_compensate=energy_compensate, - dc_preserve=dc_preserve, - ), - ) - - return io.NodeOutput(m) diff --git a/phase_inject.py b/phase_inject.py deleted file mode 100644 index 89ad2ac..0000000 --- a/phase_inject.py +++ /dev/null @@ -1,286 +0,0 @@ -"""DiversityBoost — Frequency-domain phase injection for composition diversity. - -Restores composition diversity lost during distillation by injecting the initial -noise's low-frequency phase into the model's denoised prediction. Operates at -step 0 only via a post-CFG hook — completely external, zero risk to model internals. - -Core idea: in frequency domain, amplitude encodes energy distribution ("naturalness") -while phase encodes spatial arrangement ("composition"). By rotating the prediction's -low-frequency phase toward the noise's phase — while attenuating high-frequency -amplitude to prevent "frequency shearing" — we restore per-seed composition diversity -that distillation collapsed. - -Key design decisions: - - Slerp (spherical interpolation): rotate pred's phase toward noise's phase by - γ·θ on the unit circle. Unlike lerp+renorm, Δφ is strictly linear in strength - — no topological singularity when opposing phasors cancel at γ≈0.5. - - Shared rotation field: compute a single θ per spatial-freq bin from channel-mean - phasors of pred and noise, then broadcast to all channels. Inter-channel - relative phase is exactly preserved → zero chromatic fringing. - - Per-bin rotation budget (φ_max): θ ~ Uniform(-π,π) has CV=1/√3 ≈ 58%, so - the tail (max rotation) is always 2× the mean. Without a cap, strength - simultaneously controls diversity and tail risk at a fixed ratio — making - the usable range extremely narrow ("weak → sudden collapse"). The tanh - soft-cap φ_eff = φ_max·tanh(φ_raw/φ_max) decouples mean diversity from - worst-case rotation, widening the usable strength range to [0, 1]. - - 6th-order Butterworth LPF (not Gaussian): steep transition band cleanly - separates composition-scale frequencies (≤4 spatial periods across the frame) - from object-scale frequencies (≥5 periods). Uses elliptical normalization - so n_periods applies independently per axis — resolution and aspect-ratio - independent. A Gaussian with σ=0.05 leaks 0.73 at r=0.04, causing - probabilistic distortion of object structure. - - High-frequency attenuation (hf_preserve): distilled models produce a - fully-formed x0_hat at step 0 where HF details are spatially bound to LF - structure. Rotating only LF phase while preserving HF amplitude causes - "frequency shearing" — the body moves but fingers/edges stay pinned at - the original position, forcing the model to generate deformities. - Attenuating HF amplitude makes x0_hat resemble a teacher's blurry step-0 - prediction, letting subsequent steps freely reconstruct coherent details. -""" - -import logging -import math - -import torch - -from .sampling import ( - denoised_to_raw, - raw_to_denoised, - find_step_index, - unpack_video_if_needed, - repack_video_if_needed, -) - -log = logging.getLogger("ComfyUI-DiversityBoost") - - -def _build_freq_mask(H, W, n_periods, device): - """Build an elliptical Butterworth low-pass frequency mask for rfft2 output. - - Returns a real-valued mask [1, 1, H, W//2+1] with: - - 6th-order Butterworth rolloff per-axis normalized - - DC component zeroed (preserves image mean / energy conservation) - - Uses elliptical normalization: each axis is normalized by its own cutoff - (n_periods/H for vertical, n_periods/W for horizontal). This ensures - "n_periods=2" means ≤2 full cycles in EACH direction, regardless of - aspect ratio. A circular mask with cutoff = n/max(H,W) would under-count - bins along the shorter axis. - """ - freq_y = torch.fft.fftfreq(H, device=device).unsqueeze(1) # [H, 1] - freq_x = torch.fft.rfftfreq(W, device=device).unsqueeze(0) # [1, W//2+1] - - # Elliptical normalized radius: each axis scaled by its own dimension. - # r_norm = sqrt((fy * H / n)^2 + (fx * W / n)^2) - # r_norm = 1.0 at the cutoff boundary in each axis direction. - r_norm = torch.sqrt((freq_y * H / n_periods) ** 2 + - (freq_x * W / n_periods) ** 2) - - # 6th-order Butterworth: 1 / sqrt(1 + r_norm^12). - # r_norm = 1 → mask = 0.707 (cutoff), r_norm = 2 → mask ≈ 0.016. - mask = 1.0 / torch.sqrt(1.0 + r_norm.pow(12)) - - # Zero DC: preserve mean intensity - mask[0, 0] = 0.0 - - return mask.unsqueeze(0).unsqueeze(0) # [1, 1, H, W//2+1] - - -def _phase_inject(raw_pred, raw_noise, strength, freq_mask, max_rotation, - hf_preserve, energy_compensate, dc_preserve): - """Low-frequency phase injection with optional high-frequency attenuation. - - Given: - raw_pred — model's denoised prediction x0_hat [B, C, H, W] - raw_noise — initial noise z_T (≈ args["input"] at step 0) [B, C, H, W] - strength — blend factor gamma in [0, 1] - freq_mask — Butterworth LPF mask from _build_freq_mask [1, 1, H, W//2+1] - max_rotation — per-bin rotation budget φ_max in radians (>0), or 0 to disable - hf_preserve — high-frequency amplitude preservation factor D in [0, 1]. - 1.0 = preserve all HF amplitude (original behavior). - 0.0 = attenuate HF amplitude to zero (maximum blur). - energy_compensate — if True, rescale output RMS to match original prediction. - dc_preserve — DC amplitude preservation factor in [0, 1]. - 0.0 = DC attenuated like HF (model rebuilds brightness). - 1.0 = DC fully preserved (original brightness). - - Returns: - raw_new — modified prediction with injected phase [B, C, H, W] - """ - H, W = raw_pred.shape[-2:] - - # Step 1: Forward FFT - F_pred = torch.fft.rfft2(raw_pred.float()) # [B, C, H, W//2+1] complex - F_noise = torch.fft.rfft2(raw_noise.float()) # [B, C, H, W//2+1] complex - - # Step 2: Shared rotation field — compute a single θ per spatial-freq bin, - # then broadcast to all channels. This guarantees inter-channel relative - # phase is EXACTLY preserved (identical rotation for all C channels at each bin). - unit_pred_per_ch = F_pred / F_pred.abs().clamp(min=1e-8) # [B, C, H, W//2+1] - unit_pred_shared = unit_pred_per_ch.mean(dim=1, keepdim=True) # [B, 1, H, W//2+1] - unit_pred_shared = unit_pred_shared / unit_pred_shared.abs().clamp(min=1e-8) - - unit_noise_per_ch = F_noise / F_noise.abs().clamp(min=1e-8) - unit_noise_shared = unit_noise_per_ch.mean(dim=1, keepdim=True) # [B, 1, H, W//2+1] - unit_noise_shared = unit_noise_shared / unit_noise_shared.abs().clamp(min=1e-8) - - # θ_shared: shortest arc from shared-pred to shared-noise phase. - theta = (unit_pred_shared.conj() * unit_noise_shared).angle() # [B, 1, H, W//2+1] - - # Step 3: Per-bin rotation with tanh soft-cap (rotation budget). - gamma_mask = strength * freq_mask # [1, 1, H, W//2+1] - phi_raw = gamma_mask * theta # [B, 1, H, W//2+1] - - if max_rotation > 0: - phi_eff = max_rotation * torch.tanh(phi_raw / max_rotation) # [B, 1, H, W//2+1] - else: - phi_eff = phi_raw # no cap — original behavior - - rotation = torch.exp(1j * phi_eff) # [B, 1, H, W//2+1] - - # Step 4: Apply rotation + optional high-frequency amplitude attenuation. - # S(r) = M(r) + D·(1 - M(r)): low-freq bins (M≈1) get S≈1 (preserved), - # high-freq bins (M≈0) get S=D (attenuated). - if hf_preserve < 1.0 - 1e-6: - amp_scale = freq_mask + hf_preserve * (1.0 - freq_mask) # [1, 1, H, W//2+1] - # DC amplitude: dc_preserve controls how much DC is kept. - # freq_mask[0,0]=0 → amp_scale[0,0]=hf_preserve without this line. - # dc_preserve=1 → fully preserve, dc_preserve=0 → same as hf_preserve. - amp_scale[:, :, 0, 0] = dc_preserve - F_new = F_pred * (rotation * amp_scale) - else: - F_new = F_pred * rotation # amplitude exactly preserved - - # Step 5: Inverse FFT back to spatial domain - raw_new = torch.fft.irfft2(F_new, s=(H, W)) - - # Optional energy compensation: when hf_preserve < 1, amp_scale reduces - # total energy. This rescales per-sample RMS back to the original level. - if energy_compensate and hf_preserve < 1.0 - 1e-6: - pred_rms = raw_pred.float().pow(2).mean(dim=(-2, -1), keepdim=True).sqrt().clamp(min=1e-8) - new_rms = raw_new.pow(2).mean(dim=(-2, -1), keepdim=True).sqrt().clamp(min=1e-8) - raw_new = raw_new * (pred_rms / new_rms) - - return raw_new.to(dtype=raw_pred.dtype) - - -def build_phase_injection_fn(strength=1.0, n_periods=2, max_rotation=1.5708, - hf_preserve=0.0, energy_compensate=False, - dc_preserve=0.0): - """Build a post_cfg_function that injects noise phase into the denoised prediction. - - Parameters: - strength: blend factor gamma (0 = no effect, 1 = full injection). - n_periods: max spatial periods to affect (integer). - The Butterworth cutoff uses elliptical normalization: - r_norm = sqrt((ky/n)^2 + (kx/n)^2), where ky,kx are - integer cycle counts per axis. Resolution and aspect-ratio - independent — n_periods=2 always covers the same bins. - - 1: ultra-conservative (global balance only) - - 2: conservative (default, recommended) - - 3: balanced (moderate diversity) - - 4: aggressive (more diversity, higher risk) - max_rotation: per-bin rotation budget φ_max in radians (0 = no cap). - Default π/2 ≈ 1.5708. - hf_preserve: high-frequency amplitude preservation factor D in [0, 1]. - D=0 produces a blurry "composition sketch", D=1 preserves all HF. - energy_compensate: if True, rescale output RMS to match original prediction. - dc_preserve: DC amplitude preservation factor in [0, 1]. - 0 = DC attenuated (model rebuilds brightness freely). - 1 = DC fully preserved (original brightness kept). - - Returns: - A closure suitable for model.set_model_sampler_post_cfg_function(). - """ - - state = { - "freq_mask": None, - "cached_hw": None, - } - - def phase_inject_hook(args): - denoised = args["denoised"] - sigma = args["sigma"] - model = args["model"] - model_options = args["model_options"] - x_t = args["input"] - - # --- Step 0 only --- - sample_sigmas = model_options.get("transformer_options", {}).get("sample_sigmas") - if sample_sigmas is None: - return denoised - - step_index = find_step_index(sigma, sample_sigmas) - if step_index != 0: - return denoised - - # --- Handle LTXAV packed format --- - working, pack_info = unpack_video_if_needed(denoised, args) - x_t_working, _ = unpack_video_if_needed(x_t, args) - - # --- Convert to raw space (undo process_in) --- - raw_pred = denoised_to_raw(working, model) - raw_noise = denoised_to_raw(x_t_working, model) - - # --- Build or reuse frequency mask --- - H, W = raw_pred.shape[-2:] - if state["cached_hw"] != (H, W): - state["freq_mask"] = _build_freq_mask(H, W, n_periods, raw_pred.device) - state["cached_hw"] = (H, W) - freq_mask = state["freq_mask"].to(device=raw_pred.device) - - # --- Core: phase injection --- - raw_new = _phase_inject(raw_pred, raw_noise, strength, freq_mask, - max_rotation, hf_preserve, energy_compensate, - dc_preserve) - - # --- Logging --- - if log.isEnabledFor(logging.INFO): - with torch.no_grad(): - delta = (raw_new - raw_pred).float() - B = raw_pred.shape[0] - - # Per-batch-item diagnostics - per_item = [] - for b in range(B): - d_rms = delta[b].pow(2).mean().sqrt().item() - p_rms = raw_pred[b].float().pow(2).mean().sqrt().item() - per_item.append((d_rms, p_rms)) - - delta_rms = delta.pow(2).mean().sqrt().item() - pred_rms = raw_pred.float().pow(2).mean().sqrt().item() - - if B > 1: - per_item_str = " ".join( - f"b{b}={d:.4f}/{p:.4f}" for b, (d, p) in enumerate(per_item) - ) - log.info( - "[DiversityBoost] step=0 strength=%.3f n_periods=%d " - "max_rotation=%.3f hf_preserve=%.3f " - "shape=%s delta_rms=%.4f pred_rms=%.4f ratio=%.4f " - "per_batch(delta/pred): %s", - strength, n_periods, - max_rotation, hf_preserve, - list(raw_pred.shape), - delta_rms, pred_rms, - delta_rms / max(pred_rms, 1e-8), - per_item_str, - ) - else: - log.info( - "[DiversityBoost] step=0 strength=%.3f n_periods=%d " - "max_rotation=%.3f hf_preserve=%.3f " - "shape=%s delta_rms=%.4f pred_rms=%.4f ratio=%.4f", - strength, n_periods, - max_rotation, hf_preserve, - list(raw_pred.shape), - delta_rms, pred_rms, - delta_rms / max(pred_rms, 1e-8), - ) - - # --- Convert back to process_in space --- - modified = raw_to_denoised(raw_new, model).to(dtype=denoised.dtype) - - return repack_video_if_needed(modified, pack_info) - - return phase_inject_hook