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:
@@ -0,0 +1,2 @@
|
||||
__pycache__/
|
||||
docs/
|
||||
@@ -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
@@ -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
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
|
||||
Reference in New Issue
Block a user