Add V3: HF attenuation + DCT composition push for composition diversity

Experimentally validated combined mechanism:
1. Butterworth LPF erases HF spatial anchoring (composition sketch)
2. Random 4x4 DCT field applied after blur redistributes latent energy

Push runs after cleanup in single post-cfg hook — ensures signal survives.
Defaults: strength=0.5, n_periods=2, noise_type=pink.
This commit is contained in:
facok
2026-04-16 17:08:02 +08:00
parent 1f69b5a4cf
commit 724f5fa153
8 changed files with 376 additions and 480 deletions
+2
View File
@@ -0,0 +1,2 @@
__pycache__/
docs/
+32 -42
View File
@@ -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
+32 -42
View File
@@ -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 等)——使用不同的钩子,互不干扰
## 已测试模型
+7 -8
View File
@@ -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",
]
+232
View File
@@ -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
+71
View File
@@ -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)
-102
View File
@@ -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)
-286
View File
@@ -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