2317 lines
86 KiB
Python
2317 lines
86 KiB
Python
"""
|
||
HF-Detail Sampling -- exponential integrator with spectral high-frequency
|
||
emphasis (HFE), tuned for realistic detail preservation.
|
||
|
||
Samplers
|
||
--------
|
||
2-stage (res_2s base): hfe_s1..s8, hfe_auto
|
||
3-stage (res_3s base): hfe3_s1..s8, hfe3_auto
|
||
4-stage (res_4s base): hfe4_s1..s8, hfe4_auto
|
||
5-stage (res_5s base): hfe5_s1..s8, hfe5_auto
|
||
|
||
s1 = no HF emphasis (vanilla integrator)
|
||
s8 = maximum sharpness
|
||
auto = per-step adaptive eta based on sigma envelope and content gate
|
||
|
||
Higher stage counts = more model evaluations per step = higher ODE
|
||
integration accuracy. The HFE enhancement is applied the same way
|
||
across all stage counts.
|
||
|
||
Schedulers (arctangent S-curve, bong_tangent-inspired)
|
||
------------------------------------------------------
|
||
atan_gentle -- mild mid-sigma concentration
|
||
atan_focused -- moderate detail-range concentration
|
||
atan_steep -- aggressive detail-range concentration
|
||
|
||
How the HFE sampler works
|
||
--------------------------
|
||
Base: 2-stage singlestep exponential integrator (res_2s) with phi-function
|
||
coefficients, giving an exact treatment of exponential decay and a second-
|
||
order correction from a midpoint evaluation.
|
||
|
||
Enhancement: the inter-stage correction delta (denoised_2 - denoised_1)
|
||
captures what the model reveals at lower noise -- texture, edges, micro-
|
||
structure. A 3x3 spatial high-pass (residual after box blur in latent
|
||
space) extracts the fine detail component, which is re-injected with extra
|
||
weight ``eta``. This compounds across every step.
|
||
|
||
How hfe_auto adapts
|
||
--------------------
|
||
eta_effective = eta_peak * sigma_envelope * content_gate
|
||
|
||
sigma_envelope: smoothstep from 0 (high noise, no emphasis) to 1 (detail
|
||
range, full emphasis). Prevents noise amplification at early steps.
|
||
|
||
content_gate: measures HF energy in the correction delta. When the model
|
||
correction is already rich in high-frequency content, the gate reduces
|
||
emphasis (detail is already there). When the correction is smooth, the
|
||
gate opens wider (detail needs boosting).
|
||
|
||
Cost: one 3x3 avg_pool per step for all variants (negligible vs model eval).
|
||
hfe_auto adds a few scalar ops on top.
|
||
"""
|
||
|
||
import math
|
||
import logging
|
||
from typing import Any, Dict, List, Optional
|
||
|
||
import torch
|
||
import torch.nn.functional as F
|
||
from tqdm.auto import trange
|
||
|
||
import comfy.samplers as comfy_samplers
|
||
|
||
LOGGER = logging.getLogger("HFDetailSampling")
|
||
|
||
|
||
# =====================================================================
|
||
# Phi functions (exponential integrator building blocks)
|
||
# =====================================================================
|
||
#
|
||
# phi1(-h) = (1 - e^{-h}) / h
|
||
# phi2(-h) = (e^{-h} - 1 + h) / h^2
|
||
#
|
||
# Taylor branches avoid catastrophic cancellation near h ~ 0.
|
||
|
||
def _phi1(h: torch.Tensor) -> torch.Tensor:
|
||
"""phi1(-h) for positive h. Scalar or broadcastable tensor."""
|
||
return torch.where(
|
||
h.abs() > 1e-4,
|
||
(1.0 - torch.exp(-h)) / h,
|
||
1.0 - h / 2.0 + h * h / 6.0,
|
||
)
|
||
|
||
|
||
def _phi2(h: torch.Tensor) -> torch.Tensor:
|
||
"""phi2(-h) for positive h. Scalar or broadcastable tensor."""
|
||
h2 = h * h
|
||
return torch.where(
|
||
h.abs() > 1e-4,
|
||
(torch.exp(-h) - 1.0 + h) / h2,
|
||
0.5 - h / 6.0 + h2 / 24.0,
|
||
)
|
||
|
||
|
||
def _phi3(h: torch.Tensor) -> torch.Tensor:
|
||
"""phi3(-h) for positive h. Needed by 4-stage and 5-stage integrators."""
|
||
h2 = h * h
|
||
h3 = h2 * h
|
||
return torch.where(
|
||
h.abs() > 1e-4,
|
||
(1.0 - torch.exp(-h) - h + h2 / 2.0) / h3,
|
||
1.0 / 6.0 - h / 24.0 + h2 / 120.0,
|
||
)
|
||
|
||
|
||
# =====================================================================
|
||
# Spectral detail extraction
|
||
# =====================================================================
|
||
|
||
def _clamp_boost(boost: torch.Tensor, eps: torch.Tensor,
|
||
max_ratio: float = 0.35) -> torch.Tensor:
|
||
"""Clamp HF boost so its RMS doesn't exceed max_ratio * eps RMS.
|
||
|
||
Prevents the HFE injection from overwhelming the denoising signal,
|
||
especially for higher stage counts where the inter-stage delta is
|
||
naturally larger.
|
||
"""
|
||
eps_rms = eps.square().mean().sqrt().clamp(min=1e-8)
|
||
boost_rms = boost.square().mean().sqrt()
|
||
if boost_rms > max_ratio * eps_rms:
|
||
boost = boost * (max_ratio * eps_rms / boost_rms)
|
||
return boost
|
||
|
||
|
||
def _ensure_4d(t: torch.Tensor):
|
||
"""Fold [B,C,T,H,W] -> [B*T,C,H,W] for 2D spatial ops.
|
||
|
||
Returns (folded_tensor, unfold_function).
|
||
For 4D input, returns (t, identity).
|
||
"""
|
||
if t.ndim == 5:
|
||
B, C, T, H, W = t.shape
|
||
return t.reshape(B * T, C, H, W), lambda x: x.reshape(B, C, T, x.shape[-2], x.shape[-1])
|
||
return t, lambda x: x
|
||
|
||
|
||
def _spatial_lowpass(t: torch.Tensor, kernel_size: int = 5) -> torch.Tensor:
|
||
"""Box-blur lowpass that handles 3D [B,C,N], 4D [B,C,H,W], and 5D [B,C,T,H,W]."""
|
||
pad = kernel_size // 2
|
||
if t.ndim == 5:
|
||
t_4d, unfold = _ensure_4d(t)
|
||
return unfold(_spatial_lowpass(t_4d, kernel_size))
|
||
if t.ndim == 4:
|
||
padded = F.pad(t, [pad] * 4, mode='reflect')
|
||
return F.avg_pool2d(padded, kernel_size, stride=1)
|
||
if t.ndim == 3:
|
||
padded = F.pad(t, [pad, pad], mode='reflect')
|
||
return F.avg_pool1d(padded, kernel_size, stride=1)
|
||
return t
|
||
|
||
|
||
def _extract_hf(t: torch.Tensor, kernel_size: int = 3) -> torch.Tensor:
|
||
"""
|
||
Spatial high-pass via residual after box blur.
|
||
|
||
Handles 3D [B,C,N], 4D [B,C,H,W], and 5D [B,C,T,H,W] latent tensors.
|
||
Returns zeros for other shapes.
|
||
"""
|
||
if t.ndim not in (3, 4, 5):
|
||
return torch.zeros_like(t)
|
||
return t - _spatial_lowpass(t, kernel_size)
|
||
|
||
|
||
# =====================================================================
|
||
# Console sigma plot
|
||
# =====================================================================
|
||
|
||
def _plot_sigmas(sigmas: torch.Tensor, name: str,
|
||
width: int = 120, height: int = 32) -> None:
|
||
"""Render a sigma schedule as an ASCII chart in the console."""
|
||
vals = sigmas.tolist()
|
||
if vals and vals[-1] == 0.0:
|
||
vals = vals[:-1]
|
||
n = len(vals)
|
||
if n < 2:
|
||
return
|
||
|
||
y_hi = max(vals)
|
||
y_lo = min(vals)
|
||
y_span = y_hi - y_lo
|
||
if y_span < 1e-12:
|
||
return
|
||
|
||
# Interpolate: for each column, compute the y value by linearly
|
||
# interpolating between the two nearest data points.
|
||
col_row = [0] * width
|
||
for c in range(width):
|
||
t = c * (n - 1) / (width - 1) # fractional step index
|
||
lo_i = min(int(t), n - 2)
|
||
hi_i = lo_i + 1
|
||
frac = t - lo_i
|
||
v = vals[lo_i] * (1.0 - frac) + vals[hi_i] * frac
|
||
r = int((y_hi - v) * (height - 1) / y_span + 0.5)
|
||
col_row[c] = max(0, min(height - 1, r))
|
||
|
||
# Build character canvas with connected line segments
|
||
grid = [[' '] * width for _ in range(height)]
|
||
for c in range(width):
|
||
if c == 0:
|
||
grid[col_row[c]][c] = '*'
|
||
else:
|
||
r0, r1 = col_row[c - 1], col_row[c]
|
||
lo_r, hi_r = min(r0, r1), max(r0, r1)
|
||
for r in range(lo_r, hi_r + 1):
|
||
grid[r][c] = '*'
|
||
|
||
# Y-axis label positions (5 evenly spaced)
|
||
label_rows = {0, height // 4, height // 2, 3 * height // 4, height - 1}
|
||
|
||
out = [
|
||
'',
|
||
f' {name} ({n} steps, sigma {y_hi:.2f} -> {y_lo:.4f})',
|
||
f' +{"-" * width}+',
|
||
]
|
||
for r in range(height):
|
||
y = y_hi - r * y_span / (height - 1)
|
||
lbl = f'{y:7.2f}' if r in label_rows else ' '
|
||
out.append(f' {lbl} |{"".join(grid[r])}|')
|
||
out.append(f' +{"-" * width}+')
|
||
|
||
# X-axis labels: 0, midpoint, end
|
||
mid_s = str(n // 2)
|
||
end_s = str(n)
|
||
gap1 = width // 2 - len(mid_s)
|
||
gap2 = width - width // 2 - len(end_s)
|
||
out.append(f' 0{" " * gap1}{mid_s}{" " * gap2}{end_s}')
|
||
out.append(f' {"step":^{width}}')
|
||
out.append('')
|
||
|
||
print('\n'.join(out))
|
||
|
||
|
||
# =====================================================================
|
||
# Fixed-eta sampler core
|
||
# =====================================================================
|
||
#
|
||
# Butcher tableau (res_2s exponential, 2-stage singlestep):
|
||
#
|
||
# 0 | 0 0
|
||
# c2 | c2*phi1(-h*c2) 0
|
||
# ----+-----------------------------
|
||
# | phi1(-h) - phi2(-h)/c2 phi2(-h)/c2
|
||
#
|
||
# Spectral sharpening:
|
||
# delta = e2 - e1 (correction signal)
|
||
# delta_hf = high_pass_3x3(delta) (fine spatial detail)
|
||
# e2' = e2 + eta * delta_hf (amplified texture/edges)
|
||
# x_next = x + h * (b1*e1 + b2*e2')
|
||
#
|
||
# eta = 0 recovers standard res_2s exactly.
|
||
|
||
@torch.no_grad()
|
||
def _sample_hfe(
|
||
model: Any,
|
||
x: torch.Tensor,
|
||
sigmas: torch.Tensor,
|
||
extra_args: Optional[Dict[str, Any]] = None,
|
||
callback: Optional[Any] = None,
|
||
disable: bool = False,
|
||
*,
|
||
c2: float = 0.5,
|
||
eta: float = 0.0,
|
||
) -> torch.Tensor:
|
||
"""
|
||
Exponential integrator with fixed-strength spectral detail sharpening.
|
||
|
||
c2: intermediate evaluation point in (0,1].
|
||
eta: high-frequency amplification strength. 0 = standard res_2s.
|
||
"""
|
||
if extra_args is None:
|
||
extra_args = {}
|
||
|
||
s_in = x.new_ones([x.shape[0]])
|
||
sigmas = sigmas.to(device=x.device, dtype=x.dtype)
|
||
total_steps = len(sigmas) - 1
|
||
|
||
for i in trange(total_steps, disable=disable):
|
||
sigma = sigmas[i]
|
||
sigma_next = sigmas[i + 1]
|
||
|
||
# Final step: sigma_next ~ 0, just return denoised
|
||
if sigma_next < 1e-6:
|
||
denoised = model(x, sigma * s_in, **extra_args)
|
||
x = denoised
|
||
if callback is not None:
|
||
callback({"i": i, "sigma": 0.0, "denoised": denoised, "x": x})
|
||
break
|
||
|
||
h = torch.log(sigma / sigma_next)
|
||
phi1 = _phi1(h)
|
||
phi2 = _phi2(h)
|
||
|
||
# --- Stage 1 ---
|
||
denoised_1 = model(x, sigma * s_in, **extra_args)
|
||
eps_1 = denoised_1 - x
|
||
|
||
# --- Intermediate point ---
|
||
sigma_mid = sigma * torch.exp(-c2 * h)
|
||
hc2 = h * c2
|
||
a21 = c2 * _phi1(hc2)
|
||
X_2 = x + h * a21 * eps_1
|
||
|
||
# --- Stage 2 ---
|
||
denoised_2 = model(X_2, sigma_mid * s_in, **extra_args)
|
||
eps_2 = denoised_2 - x
|
||
|
||
# --- Spectral HF sharpening ---
|
||
# Sigma warmup: suppress at high noise (first ~25% of steps) to
|
||
# prevent amplifying noise. Full strength from ~55% onward.
|
||
if eta > 0.0:
|
||
progress = i / max(total_steps - 1, 1)
|
||
sigma_gate = max(0.0, min(1.0, (progress - 0.25) / 0.30))
|
||
eta_step = eta * sigma_gate
|
||
|
||
if eta_step > 1e-3:
|
||
delta = eps_2 - eps_1
|
||
delta_hf = _extract_hf(delta)
|
||
eps_2 = eps_2 + _clamp_boost(eta_step * delta_hf, eps_2)
|
||
|
||
# --- Output weights (standard res_2s) ---
|
||
b2 = phi2 / c2
|
||
b1 = phi1 - b2
|
||
|
||
x = x + h * (b1 * eps_1 + b2 * eps_2)
|
||
|
||
# Guard against NaN/inf from numerical instability
|
||
if torch.isnan(x).any() or torch.isinf(x).any():
|
||
LOGGER.warning("HFE step %d: NaN/inf detected, falling back to "
|
||
"standard res_2s for remaining steps.", i)
|
||
eta = 0.0
|
||
x = x.nan_to_num(nan=0.0, posinf=1e4, neginf=-1e4)
|
||
|
||
if callback is not None:
|
||
callback({
|
||
"i": i,
|
||
"sigma": float(sigma_next),
|
||
"denoised": denoised_2,
|
||
"x": x,
|
||
})
|
||
|
||
return x
|
||
|
||
|
||
# =====================================================================
|
||
# Adaptive sampler (hfe_auto)
|
||
# =====================================================================
|
||
#
|
||
# Three things adapt per step:
|
||
#
|
||
# 1. c2 (Butcher tableau): ramps from c2_start (conservative, high sigma)
|
||
# to c2_end (aggressive, low sigma). This changes the actual ODE
|
||
# solver weights each step -- not just a scaling knob.
|
||
#
|
||
# 2. eta (HF emphasis): eta_peak * sigma_envelope * content_gate
|
||
# - sigma_envelope: smoothstep, 0 at high noise, 1 in detail range.
|
||
# - content_gate: 0 when correction is already HF-rich (model is
|
||
# producing detail on its own), 1 when smooth (needs sharpening).
|
||
# Full [0, 1] range -- no floor, so it can fully shut off.
|
||
#
|
||
# 3. kernel_size: 3x3 at low sigma (fine texture), 5x5 at high sigma
|
||
# (coarser structural detail).
|
||
|
||
@torch.no_grad()
|
||
def _sample_hfe_auto(
|
||
model: Any,
|
||
x: torch.Tensor,
|
||
sigmas: torch.Tensor,
|
||
extra_args: Optional[Dict[str, Any]] = None,
|
||
callback: Optional[Any] = None,
|
||
disable: bool = False,
|
||
*,
|
||
eta_peak: float = 0.55,
|
||
c2_start: float = 0.45,
|
||
c2_end: float = 0.85,
|
||
) -> torch.Tensor:
|
||
"""
|
||
Fully adaptive HFE sampler.
|
||
|
||
eta_peak: maximum HF amplification (reached at low sigma with smooth correction).
|
||
c2_start: intermediate eval point at high sigma (conservative).
|
||
c2_end: intermediate eval point at low sigma (aggressive detail capture).
|
||
"""
|
||
if extra_args is None:
|
||
extra_args = {}
|
||
|
||
s_in = x.new_ones([x.shape[0]])
|
||
sigmas = sigmas.to(device=x.device, dtype=x.dtype)
|
||
sigma_max = float(sigmas[0])
|
||
total_steps = len(sigmas) - 1
|
||
|
||
for i in trange(total_steps, disable=disable):
|
||
sigma = sigmas[i]
|
||
sigma_next = sigmas[i + 1]
|
||
|
||
if sigma_next < 1e-6:
|
||
denoised = model(x, sigma * s_in, **extra_args)
|
||
x = denoised
|
||
if callback is not None:
|
||
callback({"i": i, "sigma": 0.0, "denoised": denoised, "x": x})
|
||
break
|
||
|
||
# --- Adaptive c2: ramps from conservative to aggressive ---
|
||
progress = 1.0 - float(sigma) / sigma_max # 0 at start, 1 at end
|
||
c2 = c2_start + progress * (c2_end - c2_start)
|
||
|
||
h = torch.log(sigma / sigma_next)
|
||
phi1 = _phi1(h)
|
||
phi2 = _phi2(h)
|
||
|
||
# --- Stage 1 ---
|
||
denoised_1 = model(x, sigma * s_in, **extra_args)
|
||
eps_1 = denoised_1 - x
|
||
|
||
# --- Intermediate point (c2 varies per step) ---
|
||
sigma_mid = sigma * torch.exp(-c2 * h)
|
||
hc2 = h * c2
|
||
a21 = c2 * _phi1(hc2)
|
||
X_2 = x + h * a21 * eps_1
|
||
|
||
# --- Stage 2 ---
|
||
denoised_2 = model(X_2, sigma_mid * s_in, **extra_args)
|
||
eps_2 = denoised_2 - x
|
||
|
||
# --- Adaptive eta ---
|
||
# Adaptive kernel: 5x5 early (coarser detail), 3x3 late (fine texture)
|
||
ks = 5 if progress < 0.5 else 3
|
||
delta = eps_2 - eps_1
|
||
delta_hf = _extract_hf(delta, kernel_size=ks)
|
||
|
||
# Sigma envelope: smoothstep, suppresses at high noise
|
||
envelope = progress * progress * (3.0 - 2.0 * progress)
|
||
|
||
# Content gate: full range [0, 1]
|
||
# 0 = correction already HF-rich (model producing detail on its own)
|
||
# 1 = correction is smooth (detail needs boosting)
|
||
hf_energy = float((delta_hf ** 2).mean())
|
||
total_energy = float((delta ** 2).mean())
|
||
hf_ratio = hf_energy / (total_energy + 1e-8)
|
||
content_gate = max(0.0, 1.0 - 2.0 * hf_ratio)
|
||
|
||
eta_step = eta_peak * envelope * content_gate
|
||
|
||
if eta_step > 1e-3:
|
||
eps_2 = eps_2 + _clamp_boost(eta_step * delta_hf, eps_2)
|
||
|
||
# --- Output weights (change every step with adaptive c2) ---
|
||
b2 = phi2 / c2
|
||
b1 = phi1 - b2
|
||
|
||
x = x + h * (b1 * eps_1 + b2 * eps_2)
|
||
|
||
# Guard against NaN/inf
|
||
if torch.isnan(x).any() or torch.isinf(x).any():
|
||
LOGGER.warning("HFE auto step %d: NaN/inf detected, disabling "
|
||
"emphasis for remaining steps.", i)
|
||
eta_peak = 0.0
|
||
x = x.nan_to_num(nan=0.0, posinf=1e4, neginf=-1e4)
|
||
|
||
if callback is not None:
|
||
callback({
|
||
"i": i,
|
||
"sigma": float(sigma_next),
|
||
"denoised": denoised_2,
|
||
"x": x,
|
||
})
|
||
|
||
return x
|
||
|
||
|
||
# ---------------------------------------------------------------------
|
||
# Stage-count dispatch -- the multi-stage integrator cores are kept as
|
||
# distinct functions because each has its own Butcher tableau; this
|
||
# dispatcher routes a single user-facing sampler to the right core
|
||
# based on the requested *stages* count. Default eta values per stage
|
||
# come from the calibrated per-stage tunings (mirrors the old hfe_auto
|
||
# / hfe3_auto / hfe4_auto / hfe5_auto presets).
|
||
# ---------------------------------------------------------------------
|
||
|
||
_HFE_AUTO_DEFAULT_ETA = {2: 0.55, 3: 0.275, 4: 0.183, 5: 0.138}
|
||
_HFE_FIXED_STAGE_SCALE = {2: 1.0, 3: 1.0 / 2, 4: 1.0 / 3, 5: 1.0 / 4}
|
||
|
||
|
||
def _dispatch_hfe(model, x, sigmas, extra_args, callback, disable,
|
||
*, stages: int, eta: float, c2: float = 0.5):
|
||
"""Pick the right fixed-eta integrator core for the requested stages.
|
||
*eta* and *c2* are the user-supplied knobs; only *c2* is honored on
|
||
the 2-stage core (the 3/4/5-stage cores have fixed Butcher tableau).
|
||
"""
|
||
s = max(2, min(5, int(stages)))
|
||
if s == 2:
|
||
return _sample_hfe(model, x, sigmas, extra_args, callback, disable,
|
||
c2=c2, eta=eta)
|
||
if s == 3:
|
||
return _sample_hfe_3s(model, x, sigmas, extra_args, callback, disable,
|
||
eta=eta)
|
||
if s == 4:
|
||
return _sample_hfe_4s(model, x, sigmas, extra_args, callback, disable,
|
||
eta=eta)
|
||
return _sample_hfe_5s(model, x, sigmas, extra_args, callback, disable,
|
||
eta=eta)
|
||
|
||
|
||
def _dispatch_hfe_auto(model, x, sigmas, extra_args, callback, disable,
|
||
*, stages: int, eta: float,
|
||
c2_start: float = 0.45, c2_end: float = 0.85):
|
||
s = max(2, min(5, int(stages)))
|
||
if s == 2:
|
||
return _sample_hfe_auto(model, x, sigmas, extra_args, callback,
|
||
disable, eta_peak=eta,
|
||
c2_start=c2_start, c2_end=c2_end)
|
||
if s == 3:
|
||
return _sample_hfe_3s_auto(model, x, sigmas, extra_args, callback,
|
||
disable, eta_peak=eta)
|
||
if s == 4:
|
||
return _sample_hfe_4s_auto(model, x, sigmas, extra_args, callback,
|
||
disable, eta_peak=eta)
|
||
return _sample_hfe_5s_auto(model, x, sigmas, extra_args, callback,
|
||
disable, eta_peak=eta)
|
||
|
||
|
||
def sample_hfe_auto(model, x, sigmas, extra_args=None, callback=None, disable=False,
|
||
stages: int = 2,
|
||
eta: float = -1.0,
|
||
c2_start: float = 0.45, c2_end: float = 0.85):
|
||
"""Adaptive HFE -- variable c2, eta, and kernel per step.
|
||
|
||
*stages* selects the underlying exponential-integrator order (2..5).
|
||
*eta* maps to the per-stage peak HF amplification; pass < 0 to use
|
||
the calibrated default for the chosen stage count.
|
||
"""
|
||
if eta < 0:
|
||
eta = _HFE_AUTO_DEFAULT_ETA.get(int(stages), 0.55)
|
||
LOGGER.info(">>> hfe_auto sampler invoked (%d sigmas, stages=%d, eta=%.3f)",
|
||
len(sigmas), int(stages), eta)
|
||
return _dispatch_hfe_auto(
|
||
model, x, sigmas, extra_args, callback, disable,
|
||
stages=stages, eta=eta, c2_start=c2_start, c2_end=c2_end,
|
||
)
|
||
|
||
|
||
# =====================================================================
|
||
# 3-stage exponential integrator with HFE (hfe3_*)
|
||
# =====================================================================
|
||
#
|
||
# Butcher tableau (res_3s, c2=1/2, c3=1):
|
||
#
|
||
# 0 | 0 0 0
|
||
# 1/2 | a2_1 0 0
|
||
# 1 | a3_1 a3_2 0
|
||
# ----+--------------------------------
|
||
# | b1 b2 b3
|
||
#
|
||
# gamma = (3*c3^3 - 2*c3) / (c2*(2 - 3*c2)) = 4
|
||
# 3 model evaluations per step.
|
||
|
||
_3S_C2 = 0.5
|
||
_3S_C3 = 1.0
|
||
_3S_GAMMA = 4.0
|
||
|
||
|
||
@torch.no_grad()
|
||
def _sample_hfe_3s(
|
||
model: Any,
|
||
x: torch.Tensor,
|
||
sigmas: torch.Tensor,
|
||
extra_args: Optional[Dict[str, Any]] = None,
|
||
callback: Optional[Any] = None,
|
||
disable: bool = False,
|
||
*,
|
||
eta: float = 0.0,
|
||
) -> torch.Tensor:
|
||
"""3-stage exponential integrator with fixed-strength HFE."""
|
||
if extra_args is None:
|
||
extra_args = {}
|
||
|
||
s_in = x.new_ones([x.shape[0]])
|
||
sigmas = sigmas.to(device=x.device, dtype=x.dtype)
|
||
total_steps = len(sigmas) - 1
|
||
c2, c3, gamma = _3S_C2, _3S_C3, _3S_GAMMA
|
||
|
||
for i in trange(total_steps, disable=disable):
|
||
sigma = sigmas[i]
|
||
sigma_next = sigmas[i + 1]
|
||
|
||
if sigma_next < 1e-6:
|
||
denoised = model(x, sigma * s_in, **extra_args)
|
||
x = denoised
|
||
if callback is not None:
|
||
callback({"i": i, "sigma": 0.0, "denoised": denoised, "x": x})
|
||
break
|
||
|
||
h = torch.log(sigma / sigma_next)
|
||
phi1_h = _phi1(h)
|
||
phi2_h = _phi2(h)
|
||
|
||
# --- Stage 1 ---
|
||
denoised_1 = model(x, sigma * s_in, **extra_args)
|
||
eps_1 = denoised_1 - x
|
||
|
||
# --- Stage 2 (c2=1/2) ---
|
||
hc2 = h * c2
|
||
a2_1 = c2 * _phi1(hc2)
|
||
X_2 = x + h * a2_1 * eps_1
|
||
sigma_2 = sigma * torch.exp(-c2 * h)
|
||
denoised_2 = model(X_2, sigma_2 * s_in, **extra_args)
|
||
eps_2 = denoised_2 - x
|
||
|
||
# --- Stage 3 (c3=1) ---
|
||
hc3 = h * c3
|
||
a3_2 = gamma * c2 * _phi2(hc2) + (c3 ** 2 / c2) * _phi2(hc3)
|
||
a3_1 = c3 * _phi1(hc3) - a3_2
|
||
X_3 = x + h * (a3_1 * eps_1 + a3_2 * eps_2)
|
||
sigma_3 = sigma * torch.exp(-c3 * h)
|
||
denoised_3 = model(X_3, sigma_3 * s_in, **extra_args)
|
||
eps_3 = denoised_3 - x
|
||
|
||
# --- Spectral HF sharpening ---
|
||
if eta > 0.0:
|
||
progress = i / max(total_steps - 1, 1)
|
||
sigma_gate = max(0.0, min(1.0, (progress - 0.25) / 0.30))
|
||
eta_step = eta * sigma_gate
|
||
if eta_step > 1e-3:
|
||
delta = eps_3 - eps_1
|
||
delta_hf = _extract_hf(delta)
|
||
eps_3 = eps_3 + _clamp_boost(eta_step * delta_hf, eps_3)
|
||
|
||
# --- Output weights ---
|
||
b3 = phi2_h / (gamma * c2 + c3)
|
||
b2 = gamma * b3
|
||
b1 = phi1_h - b2 - b3
|
||
|
||
x = x + h * (b1 * eps_1 + b2 * eps_2 + b3 * eps_3)
|
||
|
||
if torch.isnan(x).any() or torch.isinf(x).any():
|
||
LOGGER.warning("HFE 3s step %d: NaN/inf detected, disabling "
|
||
"emphasis for remaining steps.", i)
|
||
eta = 0.0
|
||
x = x.nan_to_num(nan=0.0, posinf=1e4, neginf=-1e4)
|
||
|
||
if callback is not None:
|
||
callback({
|
||
"i": i,
|
||
"sigma": float(sigma_next),
|
||
"denoised": denoised_3,
|
||
"x": x,
|
||
})
|
||
|
||
return x
|
||
|
||
|
||
@torch.no_grad()
|
||
def _sample_hfe_3s_auto(
|
||
model: Any,
|
||
x: torch.Tensor,
|
||
sigmas: torch.Tensor,
|
||
extra_args: Optional[Dict[str, Any]] = None,
|
||
callback: Optional[Any] = None,
|
||
disable: bool = False,
|
||
*,
|
||
eta_peak: float = 0.55,
|
||
) -> torch.Tensor:
|
||
"""3-stage adaptive HFE -- per-step eta based on sigma and content."""
|
||
if extra_args is None:
|
||
extra_args = {}
|
||
|
||
s_in = x.new_ones([x.shape[0]])
|
||
sigmas = sigmas.to(device=x.device, dtype=x.dtype)
|
||
sigma_max = float(sigmas[0])
|
||
total_steps = len(sigmas) - 1
|
||
c2, c3, gamma = _3S_C2, _3S_C3, _3S_GAMMA
|
||
|
||
for i in trange(total_steps, disable=disable):
|
||
sigma = sigmas[i]
|
||
sigma_next = sigmas[i + 1]
|
||
|
||
if sigma_next < 1e-6:
|
||
denoised = model(x, sigma * s_in, **extra_args)
|
||
x = denoised
|
||
if callback is not None:
|
||
callback({"i": i, "sigma": 0.0, "denoised": denoised, "x": x})
|
||
break
|
||
|
||
progress = 1.0 - float(sigma) / sigma_max
|
||
h = torch.log(sigma / sigma_next)
|
||
phi1_h = _phi1(h)
|
||
phi2_h = _phi2(h)
|
||
|
||
# --- Stage 1 ---
|
||
denoised_1 = model(x, sigma * s_in, **extra_args)
|
||
eps_1 = denoised_1 - x
|
||
|
||
# --- Stage 2 ---
|
||
hc2 = h * c2
|
||
a2_1 = c2 * _phi1(hc2)
|
||
X_2 = x + h * a2_1 * eps_1
|
||
sigma_2 = sigma * torch.exp(-c2 * h)
|
||
denoised_2 = model(X_2, sigma_2 * s_in, **extra_args)
|
||
eps_2 = denoised_2 - x
|
||
|
||
# --- Stage 3 ---
|
||
hc3 = h * c3
|
||
a3_2 = gamma * c2 * _phi2(hc2) + (c3 ** 2 / c2) * _phi2(hc3)
|
||
a3_1 = c3 * _phi1(hc3) - a3_2
|
||
X_3 = x + h * (a3_1 * eps_1 + a3_2 * eps_2)
|
||
sigma_3 = sigma * torch.exp(-c3 * h)
|
||
denoised_3 = model(X_3, sigma_3 * s_in, **extra_args)
|
||
eps_3 = denoised_3 - x
|
||
|
||
# --- Adaptive eta ---
|
||
ks = 5 if progress < 0.5 else 3
|
||
delta = eps_3 - eps_1
|
||
delta_hf = _extract_hf(delta, kernel_size=ks)
|
||
envelope = progress * progress * (3.0 - 2.0 * progress)
|
||
hf_energy = float((delta_hf ** 2).mean())
|
||
total_energy = float((delta ** 2).mean())
|
||
hf_ratio = hf_energy / (total_energy + 1e-8)
|
||
content_gate = max(0.0, 1.0 - 2.0 * hf_ratio)
|
||
eta_step = eta_peak * envelope * content_gate
|
||
|
||
if eta_step > 1e-3:
|
||
eps_3 = eps_3 + _clamp_boost(eta_step * delta_hf, eps_3)
|
||
|
||
# --- Output weights ---
|
||
b3 = phi2_h / (gamma * c2 + c3)
|
||
b2 = gamma * b3
|
||
b1 = phi1_h - b2 - b3
|
||
|
||
x = x + h * (b1 * eps_1 + b2 * eps_2 + b3 * eps_3)
|
||
|
||
if torch.isnan(x).any() or torch.isinf(x).any():
|
||
LOGGER.warning("HFE 3s auto step %d: NaN/inf detected, disabling "
|
||
"emphasis for remaining steps.", i)
|
||
eta_peak = 0.0
|
||
x = x.nan_to_num(nan=0.0, posinf=1e4, neginf=-1e4)
|
||
|
||
if callback is not None:
|
||
callback({
|
||
"i": i,
|
||
"sigma": float(sigma_next),
|
||
"denoised": denoised_3,
|
||
"x": x,
|
||
})
|
||
|
||
return x
|
||
|
||
|
||
# sample_hfe3_auto / hfe4_auto / hfe5_auto have been folded into
|
||
# sample_hfe_auto(stages=N) -- pick the stage count via ManualSampler.
|
||
|
||
|
||
# =====================================================================
|
||
# 4-stage exponential integrator with HFE (hfe4_*)
|
||
# =====================================================================
|
||
#
|
||
# Butcher tableau (Strehmel-Weiner, c2=1/2, c3=1/2, c4=1):
|
||
#
|
||
# 0 | 0 0 0 0
|
||
# 1/2 | a2_1 0 0 0
|
||
# 1/2 | a3_1 a3_2 0 0
|
||
# 1 | a4_1 a4_2 a4_3 0
|
||
# ----+--------------------------------------------
|
||
# | b1 b2 b3 b4
|
||
#
|
||
# 4 model evaluations per step. Weak 4th order accuracy.
|
||
|
||
_4S_C2 = 0.5
|
||
_4S_C3 = 0.5
|
||
_4S_C4 = 1.0
|
||
|
||
|
||
@torch.no_grad()
|
||
def _sample_hfe_4s(
|
||
model: Any,
|
||
x: torch.Tensor,
|
||
sigmas: torch.Tensor,
|
||
extra_args: Optional[Dict[str, Any]] = None,
|
||
callback: Optional[Any] = None,
|
||
disable: bool = False,
|
||
*,
|
||
eta: float = 0.0,
|
||
) -> torch.Tensor:
|
||
"""4-stage exponential integrator with fixed-strength HFE."""
|
||
if extra_args is None:
|
||
extra_args = {}
|
||
|
||
s_in = x.new_ones([x.shape[0]])
|
||
sigmas = sigmas.to(device=x.device, dtype=x.dtype)
|
||
total_steps = len(sigmas) - 1
|
||
c2, c3, c4 = _4S_C2, _4S_C3, _4S_C4
|
||
|
||
for i in trange(total_steps, disable=disable):
|
||
sigma = sigmas[i]
|
||
sigma_next = sigmas[i + 1]
|
||
|
||
if sigma_next < 1e-6:
|
||
denoised = model(x, sigma * s_in, **extra_args)
|
||
x = denoised
|
||
if callback is not None:
|
||
callback({"i": i, "sigma": 0.0, "denoised": denoised, "x": x})
|
||
break
|
||
|
||
h = torch.log(sigma / sigma_next)
|
||
hc2 = h * c2
|
||
hc3 = h * c3
|
||
phi1_h = _phi1(h)
|
||
phi2_h = _phi2(h)
|
||
phi3_h = _phi3(h)
|
||
|
||
# --- Stage 1 ---
|
||
denoised_1 = model(x, sigma * s_in, **extra_args)
|
||
eps_1 = denoised_1 - x
|
||
|
||
# --- Stage 2 (c2=1/2) ---
|
||
a2_1 = c2 * _phi1(hc2)
|
||
X_2 = x + h * a2_1 * eps_1
|
||
sigma_2 = sigma * torch.exp(-c2 * h)
|
||
denoised_2 = model(X_2, sigma_2 * s_in, **extra_args)
|
||
eps_2 = denoised_2 - x
|
||
|
||
# --- Stage 3 (c3=1/2) ---
|
||
a3_2 = c3 * _phi2(hc3)
|
||
a3_1 = c3 * _phi1(hc3) - a3_2
|
||
X_3 = x + h * (a3_1 * eps_1 + a3_2 * eps_2)
|
||
sigma_3 = sigma * torch.exp(-c3 * h)
|
||
denoised_3 = model(X_3, sigma_3 * s_in, **extra_args)
|
||
eps_3 = denoised_3 - x
|
||
|
||
# --- Stage 4 (c4=1) ---
|
||
a4_2 = -2.0 * phi2_h
|
||
a4_3 = 4.0 * phi2_h
|
||
a4_1 = phi1_h - a4_2 - a4_3
|
||
X_4 = x + h * (a4_1 * eps_1 + a4_2 * eps_2 + a4_3 * eps_3)
|
||
sigma_4 = sigma * torch.exp(-c4 * h)
|
||
denoised_4 = model(X_4, sigma_4 * s_in, **extra_args)
|
||
eps_4 = denoised_4 - x
|
||
|
||
# --- Spectral HF sharpening ---
|
||
if eta > 0.0:
|
||
progress = i / max(total_steps - 1, 1)
|
||
sigma_gate = max(0.0, min(1.0, (progress - 0.25) / 0.30))
|
||
eta_step = eta * sigma_gate
|
||
if eta_step > 1e-3:
|
||
delta = eps_4 - eps_1
|
||
delta_hf = _extract_hf(delta)
|
||
eps_4 = eps_4 + _clamp_boost(eta_step * delta_hf, eps_4)
|
||
|
||
# --- Output weights (Strehmel-Weiner, b2=0) ---
|
||
b3 = 4.0 * phi2_h - 8.0 * phi3_h
|
||
b4 = -phi2_h + 4.0 * phi3_h
|
||
b1 = phi1_h - b3 - b4
|
||
|
||
x = x + h * (b1 * eps_1 + b3 * eps_3 + b4 * eps_4)
|
||
|
||
if torch.isnan(x).any() or torch.isinf(x).any():
|
||
LOGGER.warning("HFE 4s step %d: NaN/inf detected, disabling "
|
||
"emphasis for remaining steps.", i)
|
||
eta = 0.0
|
||
x = x.nan_to_num(nan=0.0, posinf=1e4, neginf=-1e4)
|
||
|
||
if callback is not None:
|
||
callback({
|
||
"i": i,
|
||
"sigma": float(sigma_next),
|
||
"denoised": denoised_4,
|
||
"x": x,
|
||
})
|
||
|
||
return x
|
||
|
||
|
||
@torch.no_grad()
|
||
def _sample_hfe_4s_auto(
|
||
model: Any,
|
||
x: torch.Tensor,
|
||
sigmas: torch.Tensor,
|
||
extra_args: Optional[Dict[str, Any]] = None,
|
||
callback: Optional[Any] = None,
|
||
disable: bool = False,
|
||
*,
|
||
eta_peak: float = 0.55,
|
||
) -> torch.Tensor:
|
||
"""4-stage adaptive HFE -- per-step eta based on sigma and content."""
|
||
if extra_args is None:
|
||
extra_args = {}
|
||
|
||
s_in = x.new_ones([x.shape[0]])
|
||
sigmas = sigmas.to(device=x.device, dtype=x.dtype)
|
||
sigma_max = float(sigmas[0])
|
||
total_steps = len(sigmas) - 1
|
||
c2, c3, c4 = _4S_C2, _4S_C3, _4S_C4
|
||
|
||
for i in trange(total_steps, disable=disable):
|
||
sigma = sigmas[i]
|
||
sigma_next = sigmas[i + 1]
|
||
|
||
if sigma_next < 1e-6:
|
||
denoised = model(x, sigma * s_in, **extra_args)
|
||
x = denoised
|
||
if callback is not None:
|
||
callback({"i": i, "sigma": 0.0, "denoised": denoised, "x": x})
|
||
break
|
||
|
||
progress = 1.0 - float(sigma) / sigma_max
|
||
h = torch.log(sigma / sigma_next)
|
||
hc2 = h * c2
|
||
hc3 = h * c3
|
||
phi1_h = _phi1(h)
|
||
phi2_h = _phi2(h)
|
||
phi3_h = _phi3(h)
|
||
|
||
# --- Stage 1 ---
|
||
denoised_1 = model(x, sigma * s_in, **extra_args)
|
||
eps_1 = denoised_1 - x
|
||
|
||
# --- Stage 2 ---
|
||
a2_1 = c2 * _phi1(hc2)
|
||
X_2 = x + h * a2_1 * eps_1
|
||
sigma_2 = sigma * torch.exp(-c2 * h)
|
||
denoised_2 = model(X_2, sigma_2 * s_in, **extra_args)
|
||
eps_2 = denoised_2 - x
|
||
|
||
# --- Stage 3 ---
|
||
a3_2 = c3 * _phi2(hc3)
|
||
a3_1 = c3 * _phi1(hc3) - a3_2
|
||
X_3 = x + h * (a3_1 * eps_1 + a3_2 * eps_2)
|
||
sigma_3 = sigma * torch.exp(-c3 * h)
|
||
denoised_3 = model(X_3, sigma_3 * s_in, **extra_args)
|
||
eps_3 = denoised_3 - x
|
||
|
||
# --- Stage 4 ---
|
||
a4_2 = -2.0 * phi2_h
|
||
a4_3 = 4.0 * phi2_h
|
||
a4_1 = phi1_h - a4_2 - a4_3
|
||
X_4 = x + h * (a4_1 * eps_1 + a4_2 * eps_2 + a4_3 * eps_3)
|
||
sigma_4 = sigma * torch.exp(-c4 * h)
|
||
denoised_4 = model(X_4, sigma_4 * s_in, **extra_args)
|
||
eps_4 = denoised_4 - x
|
||
|
||
# --- Adaptive eta ---
|
||
ks = 5 if progress < 0.5 else 3
|
||
delta = eps_4 - eps_1
|
||
delta_hf = _extract_hf(delta, kernel_size=ks)
|
||
envelope = progress * progress * (3.0 - 2.0 * progress)
|
||
hf_energy = float((delta_hf ** 2).mean())
|
||
total_energy = float((delta ** 2).mean())
|
||
hf_ratio = hf_energy / (total_energy + 1e-8)
|
||
content_gate = max(0.0, 1.0 - 2.0 * hf_ratio)
|
||
eta_step = eta_peak * envelope * content_gate
|
||
|
||
if eta_step > 1e-3:
|
||
eps_4 = eps_4 + _clamp_boost(eta_step * delta_hf, eps_4)
|
||
|
||
# --- Output weights ---
|
||
b2 = 0.0
|
||
b3 = 4.0 * phi2_h - 8.0 * phi3_h
|
||
b4 = -phi2_h + 4.0 * phi3_h
|
||
b1 = phi1_h - b3 - b4
|
||
|
||
x = x + h * (b1 * eps_1 + b3 * eps_3 + b4 * eps_4)
|
||
|
||
if torch.isnan(x).any() or torch.isinf(x).any():
|
||
LOGGER.warning("HFE 4s auto step %d: NaN/inf detected, disabling "
|
||
"emphasis for remaining steps.", i)
|
||
eta_peak = 0.0
|
||
x = x.nan_to_num(nan=0.0, posinf=1e4, neginf=-1e4)
|
||
|
||
if callback is not None:
|
||
callback({
|
||
"i": i,
|
||
"sigma": float(sigma_next),
|
||
"denoised": denoised_4,
|
||
"x": x,
|
||
})
|
||
|
||
return x
|
||
|
||
|
||
# sample_hfe4_auto folded into sample_hfe_auto(stages=4).
|
||
|
||
|
||
# =====================================================================
|
||
# 5-stage exponential integrator with HFE (hfe5_*)
|
||
# =====================================================================
|
||
#
|
||
# Butcher tableau (c2=1/2, c3=1/2, c4=1, c5=1/2):
|
||
#
|
||
# 0 | 0 0 0 0 0
|
||
# 1/2 | a2_1 0 0 0 0
|
||
# 1/2 | a3_1 a3_2 0 0 0
|
||
# 1 | a4_1 a4_2 a4_3 0 0
|
||
# 1/2 | a5_1 a5_2 a5_3 a5_4 0
|
||
# ----+------------------------------------
|
||
# | b1 b2 b3 b4 b5
|
||
#
|
||
# 5 model evaluations per step. Non-monotonic node placement (c5=1/2).
|
||
|
||
_5S_C2 = 0.5
|
||
_5S_C3 = 0.5
|
||
_5S_C4 = 1.0
|
||
_5S_C5 = 0.5
|
||
|
||
|
||
@torch.no_grad()
|
||
def _sample_hfe_5s(
|
||
model: Any,
|
||
x: torch.Tensor,
|
||
sigmas: torch.Tensor,
|
||
extra_args: Optional[Dict[str, Any]] = None,
|
||
callback: Optional[Any] = None,
|
||
disable: bool = False,
|
||
*,
|
||
eta: float = 0.0,
|
||
) -> torch.Tensor:
|
||
"""5-stage exponential integrator with fixed-strength HFE."""
|
||
LOGGER.info(">>> _sample_hfe_5s called with eta=%.4f", eta)
|
||
if extra_args is None:
|
||
extra_args = {}
|
||
|
||
s_in = x.new_ones([x.shape[0]])
|
||
sigmas = sigmas.to(device=x.device, dtype=x.dtype)
|
||
total_steps = len(sigmas) - 1
|
||
c2, c3, c4, c5 = _5S_C2, _5S_C3, _5S_C4, _5S_C5
|
||
|
||
for i in trange(total_steps, disable=disable):
|
||
sigma = sigmas[i]
|
||
sigma_next = sigmas[i + 1]
|
||
|
||
if sigma_next < 1e-6:
|
||
denoised = model(x, sigma * s_in, **extra_args)
|
||
x = denoised
|
||
if callback is not None:
|
||
callback({"i": i, "sigma": 0.0, "denoised": denoised, "x": x})
|
||
break
|
||
|
||
h = torch.log(sigma / sigma_next)
|
||
hc2 = h * c2
|
||
hc3 = h * c3
|
||
hc5 = h * c5
|
||
phi1_h = _phi1(h)
|
||
phi2_h = _phi2(h)
|
||
phi3_h = _phi3(h)
|
||
|
||
# --- Stage 1 ---
|
||
denoised_1 = model(x, sigma * s_in, **extra_args)
|
||
eps_1 = denoised_1 - x
|
||
|
||
# --- Stage 2 (c2=1/2) ---
|
||
a2_1 = c2 * _phi1(hc2)
|
||
X_2 = x + h * a2_1 * eps_1
|
||
sigma_2 = sigma * torch.exp(-c2 * h)
|
||
denoised_2 = model(X_2, sigma_2 * s_in, **extra_args)
|
||
eps_2 = denoised_2 - x
|
||
|
||
# --- Stage 3 (c3=1/2) ---
|
||
a3_2 = _phi2(hc3)
|
||
a3_1 = c3 * _phi1(hc3) - a3_2
|
||
X_3 = x + h * (a3_1 * eps_1 + a3_2 * eps_2)
|
||
sigma_3 = sigma * torch.exp(-c3 * h)
|
||
denoised_3 = model(X_3, sigma_3 * s_in, **extra_args)
|
||
eps_3 = denoised_3 - x
|
||
|
||
# --- Stage 4 (c4=1) ---
|
||
a4_2 = phi2_h
|
||
a4_3 = phi2_h
|
||
a4_1 = phi1_h - a4_2 - a4_3
|
||
X_4 = x + h * (a4_1 * eps_1 + a4_2 * eps_2 + a4_3 * eps_3)
|
||
sigma_4 = sigma * torch.exp(-c4 * h)
|
||
denoised_4 = model(X_4, sigma_4 * s_in, **extra_args)
|
||
eps_4 = denoised_4 - x
|
||
|
||
# --- Stage 5 (c5=1/2, non-monotonic) ---
|
||
phi2_hc5 = _phi2(hc5)
|
||
phi3_hc5 = _phi3(hc5)
|
||
a5_2 = 0.5 * phi2_hc5 - phi3_h + 0.25 * phi2_h - 0.5 * phi3_hc5
|
||
a5_3 = a5_2
|
||
a5_4 = 0.25 * phi2_hc5 - a5_2
|
||
a5_1 = c5 * _phi1(hc5) - a5_2 - a5_3 - a5_4
|
||
X_5 = x + h * (a5_1 * eps_1 + a5_2 * eps_2 + a5_3 * eps_3 + a5_4 * eps_4)
|
||
sigma_5 = sigma * torch.exp(-c5 * h)
|
||
denoised_5 = model(X_5, sigma_5 * s_in, **extra_args)
|
||
eps_5 = denoised_5 - x
|
||
|
||
# --- Spectral HF sharpening ---
|
||
if eta > 0.0:
|
||
progress = i / max(total_steps - 1, 1)
|
||
sigma_gate = max(0.0, min(1.0, (progress - 0.25) / 0.30))
|
||
eta_step = eta * sigma_gate
|
||
if eta_step > 1e-3:
|
||
delta = eps_5 - eps_1
|
||
delta_hf = _extract_hf(delta)
|
||
eps_5 = eps_5 + _clamp_boost(eta_step * delta_hf, eps_5)
|
||
|
||
# --- Output weights (b2=0, b3=0) ---
|
||
b4 = -phi2_h + 4.0 * phi3_h
|
||
b5 = 4.0 * phi2_h - 8.0 * phi3_h
|
||
b1 = phi1_h - b4 - b5
|
||
|
||
x = x + h * (b1 * eps_1 + b4 * eps_4 + b5 * eps_5)
|
||
|
||
if torch.isnan(x).any() or torch.isinf(x).any():
|
||
LOGGER.warning("HFE 5s step %d: NaN/inf detected, disabling "
|
||
"emphasis for remaining steps.", i)
|
||
eta = 0.0
|
||
x = x.nan_to_num(nan=0.0, posinf=1e4, neginf=-1e4)
|
||
|
||
if callback is not None:
|
||
callback({
|
||
"i": i,
|
||
"sigma": float(sigma_next),
|
||
"denoised": denoised_5,
|
||
"x": x,
|
||
})
|
||
|
||
return x
|
||
|
||
|
||
@torch.no_grad()
|
||
def _sample_hfe_5s_auto(
|
||
model: Any,
|
||
x: torch.Tensor,
|
||
sigmas: torch.Tensor,
|
||
extra_args: Optional[Dict[str, Any]] = None,
|
||
callback: Optional[Any] = None,
|
||
disable: bool = False,
|
||
*,
|
||
eta_peak: float = 0.55,
|
||
) -> torch.Tensor:
|
||
"""5-stage adaptive HFE -- per-step eta based on sigma and content."""
|
||
if extra_args is None:
|
||
extra_args = {}
|
||
|
||
s_in = x.new_ones([x.shape[0]])
|
||
sigmas = sigmas.to(device=x.device, dtype=x.dtype)
|
||
sigma_max = float(sigmas[0])
|
||
total_steps = len(sigmas) - 1
|
||
c2, c3, c4, c5 = _5S_C2, _5S_C3, _5S_C4, _5S_C5
|
||
|
||
for i in trange(total_steps, disable=disable):
|
||
sigma = sigmas[i]
|
||
sigma_next = sigmas[i + 1]
|
||
|
||
if sigma_next < 1e-6:
|
||
denoised = model(x, sigma * s_in, **extra_args)
|
||
x = denoised
|
||
if callback is not None:
|
||
callback({"i": i, "sigma": 0.0, "denoised": denoised, "x": x})
|
||
break
|
||
|
||
progress = 1.0 - float(sigma) / sigma_max
|
||
h = torch.log(sigma / sigma_next)
|
||
hc2 = h * c2
|
||
hc3 = h * c3
|
||
hc5 = h * c5
|
||
phi1_h = _phi1(h)
|
||
phi2_h = _phi2(h)
|
||
phi3_h = _phi3(h)
|
||
|
||
# --- Stage 1 ---
|
||
denoised_1 = model(x, sigma * s_in, **extra_args)
|
||
eps_1 = denoised_1 - x
|
||
|
||
# --- Stage 2 ---
|
||
a2_1 = c2 * _phi1(hc2)
|
||
X_2 = x + h * a2_1 * eps_1
|
||
sigma_2 = sigma * torch.exp(-c2 * h)
|
||
denoised_2 = model(X_2, sigma_2 * s_in, **extra_args)
|
||
eps_2 = denoised_2 - x
|
||
|
||
# --- Stage 3 ---
|
||
a3_2 = _phi2(hc3)
|
||
a3_1 = c3 * _phi1(hc3) - a3_2
|
||
X_3 = x + h * (a3_1 * eps_1 + a3_2 * eps_2)
|
||
sigma_3 = sigma * torch.exp(-c3 * h)
|
||
denoised_3 = model(X_3, sigma_3 * s_in, **extra_args)
|
||
eps_3 = denoised_3 - x
|
||
|
||
# --- Stage 4 ---
|
||
a4_2 = phi2_h
|
||
a4_3 = phi2_h
|
||
a4_1 = phi1_h - a4_2 - a4_3
|
||
X_4 = x + h * (a4_1 * eps_1 + a4_2 * eps_2 + a4_3 * eps_3)
|
||
sigma_4 = sigma * torch.exp(-c4 * h)
|
||
denoised_4 = model(X_4, sigma_4 * s_in, **extra_args)
|
||
eps_4 = denoised_4 - x
|
||
|
||
# --- Stage 5 ---
|
||
phi2_hc5 = _phi2(hc5)
|
||
phi3_hc5 = _phi3(hc5)
|
||
a5_2 = 0.5 * phi2_hc5 - phi3_h + 0.25 * phi2_h - 0.5 * phi3_hc5
|
||
a5_3 = a5_2
|
||
a5_4 = 0.25 * phi2_hc5 - a5_2
|
||
a5_1 = c5 * _phi1(hc5) - a5_2 - a5_3 - a5_4
|
||
X_5 = x + h * (a5_1 * eps_1 + a5_2 * eps_2 + a5_3 * eps_3 + a5_4 * eps_4)
|
||
sigma_5 = sigma * torch.exp(-c5 * h)
|
||
denoised_5 = model(X_5, sigma_5 * s_in, **extra_args)
|
||
eps_5 = denoised_5 - x
|
||
|
||
# --- Adaptive eta ---
|
||
ks = 5 if progress < 0.5 else 3
|
||
delta = eps_5 - eps_1
|
||
delta_hf = _extract_hf(delta, kernel_size=ks)
|
||
envelope = progress * progress * (3.0 - 2.0 * progress)
|
||
hf_energy = float((delta_hf ** 2).mean())
|
||
total_energy = float((delta ** 2).mean())
|
||
hf_ratio = hf_energy / (total_energy + 1e-8)
|
||
content_gate = max(0.0, 1.0 - 2.0 * hf_ratio)
|
||
eta_step = eta_peak * envelope * content_gate
|
||
|
||
if eta_step > 1e-3:
|
||
eps_5 = eps_5 + _clamp_boost(eta_step * delta_hf, eps_5)
|
||
|
||
# --- Output weights (b2=0, b3=0) ---
|
||
b4 = -phi2_h + 4.0 * phi3_h
|
||
b5 = 4.0 * phi2_h - 8.0 * phi3_h
|
||
b1 = phi1_h - b4 - b5
|
||
|
||
x = x + h * (b1 * eps_1 + b4 * eps_4 + b5 * eps_5)
|
||
|
||
if torch.isnan(x).any() or torch.isinf(x).any():
|
||
LOGGER.warning("HFE 5s auto step %d: NaN/inf detected, disabling "
|
||
"emphasis for remaining steps.", i)
|
||
eta_peak = 0.0
|
||
x = x.nan_to_num(nan=0.0, posinf=1e4, neginf=-1e4)
|
||
|
||
if callback is not None:
|
||
callback({
|
||
"i": i,
|
||
"sigma": float(sigma_next),
|
||
"denoised": denoised_5,
|
||
"x": x,
|
||
})
|
||
|
||
return x
|
||
|
||
|
||
# sample_hfe5_auto folded into sample_hfe_auto(stages=5).
|
||
|
||
|
||
# =====================================================================
|
||
# Graduated fixed-strength presets (hfe_s1..s8, hfe3_s1..s8, etc.)
|
||
# =====================================================================
|
||
#
|
||
# eta follows a power-1.5 curve from 0.00 (s1) to 0.48 (s8) so that
|
||
# the perceptual jump between adjacent levels feels roughly even.
|
||
#
|
||
# 2-stage presets also vary c2 from 0.45 to 0.80.
|
||
# 3/4/5-stage presets use fixed c values (from reference tableaux)
|
||
# and only vary eta.
|
||
|
||
_HFE_LEVELS = 8
|
||
_HFE_C2_MIN = 0.45
|
||
_HFE_C2_MAX = 0.80
|
||
_HFE_ETA_MAX = 0.48
|
||
|
||
# Map stage count -> (core function, name prefix)
|
||
_STAGE_CORES = {
|
||
2: (_sample_hfe, "hfe"),
|
||
3: (_sample_hfe_3s, "hfe3"),
|
||
4: (_sample_hfe_4s, "hfe4"),
|
||
5: (_sample_hfe_5s, "hfe5"),
|
||
}
|
||
|
||
|
||
def _make_hfe_preset(level: int):
|
||
"""Factory: a 2-stage strength preset (hfe_s<level>).
|
||
|
||
Each preset has a calibrated default ``eta`` and ``c2`` matching the
|
||
historical s1..s8 levels, but both knobs (plus ``stages``, 2..5) are
|
||
overridable kwargs so ManualSampler / extra_options can dial the
|
||
same shape up to a higher-order integrator at runtime.
|
||
"""
|
||
t = level / (_HFE_LEVELS - 1)
|
||
eta_default = _HFE_ETA_MAX * (t ** 1.5)
|
||
c2_default = _HFE_C2_MIN + t * (_HFE_C2_MAX - _HFE_C2_MIN)
|
||
|
||
# Negative-value sentinel keeps "user provided eta" distinguishable
|
||
# from "use the calibrated default". Without it, a Manual Sampler
|
||
# call like hfe_s4(stages=4, eta=0.5) would get its eta silently
|
||
# divided by (stages-1).
|
||
def sampler(model, x, sigmas, extra_args=None, callback=None, disable=False,
|
||
stages: int = 2,
|
||
eta: float = -1.0,
|
||
c2: float = -1.0):
|
||
using_eta_default = eta < 0
|
||
using_c2_default = c2 < 0
|
||
actual_eta = eta_default if using_eta_default else eta
|
||
actual_c2 = c2_default if using_c2_default else c2
|
||
# Auto-scale eta only when the caller is taking the preset default
|
||
# AND dialing stages above 2 -- preserves the original per-stage
|
||
# calibration. User-supplied eta is taken at face value.
|
||
if using_eta_default and int(stages) > 2:
|
||
actual_eta = actual_eta * _HFE_FIXED_STAGE_SCALE.get(
|
||
int(stages), 1.0)
|
||
LOGGER.info(">>> hfe_s%d preset invoked (stages=%d, eta=%.4f%s)",
|
||
level + 1, int(stages), actual_eta,
|
||
" [default-scaled]" if using_eta_default and int(stages) > 2
|
||
else "")
|
||
return _dispatch_hfe(model, x, sigmas, extra_args, callback, disable,
|
||
stages=stages, eta=actual_eta, c2=actual_c2)
|
||
|
||
sampler.__doc__ = (f"HFE strength {level + 1}/{_HFE_LEVELS}"
|
||
f" -- default eta={eta_default:.3f}, c2={c2_default:.3f}, "
|
||
f"stages=2..5 via ManualSampler")
|
||
name = f"hfe_s{level + 1}"
|
||
sampler.__name__ = f"sample_{name}"
|
||
sampler.__qualname__ = sampler.__name__
|
||
return name, sampler
|
||
|
||
|
||
# Generate just the 8 strength presets (2-stage by default; users dial
|
||
# stages 2..5 via ManualSampler -- see "stages" kwarg).
|
||
_HFE_PRESETS = {}
|
||
for _lvl in range(_HFE_LEVELS):
|
||
_name, _fn = _make_hfe_preset(_lvl)
|
||
_HFE_PRESETS[_name] = _fn
|
||
|
||
|
||
# =====================================================================
|
||
# Experimental HFE samplers (hfx_*)
|
||
# =====================================================================
|
||
#
|
||
# Four fundamentally different enhancement modes, all using a shared
|
||
# 2-stage exponential integrator base (c2=0.65, eta=0.25):
|
||
#
|
||
# sharp - Unsharp mask on eps_2 (inside integrator)
|
||
# boost - Uniform eps_2 scaling / lying sigma (inside integrator)
|
||
# detail - Post-step HF injection from denoised_2
|
||
# stochastic - Structure-aware noise injection (post-step, non-det)
|
||
# momentum - Cross-step temporal EMA on denoised (inside integrator)
|
||
# spectral - FFT frequency reshaping of eps_2 (inside integrator)
|
||
# orthogonal - Gram-Schmidt novel-info amplification (inside integrator)
|
||
# refine - ODE curvature-adaptive spatial emphasis (inside integrator)
|
||
# focus - Value-domain power-law contrast (inside integrator)
|
||
# coherence - Inter-stage FFT phase coherence gating (inside integrator)
|
||
|
||
_HFX_C2 = 0.65
|
||
_HFX_ETA = 0.25
|
||
_HFX_SDE_STRENGTH = 0.08
|
||
_HFX_BOOST_FACTOR = 0.5
|
||
_HFX_MOMENTUM_STRENGTH = 0.35
|
||
_HFX_SPECTRAL_ALPHA = 0.6
|
||
_HFX_FOCUS_GAMMA = 0.8
|
||
|
||
|
||
@torch.no_grad()
|
||
def _sample_hfx(
|
||
model: Any,
|
||
x: torch.Tensor,
|
||
sigmas: torch.Tensor,
|
||
extra_args: Optional[Dict[str, Any]] = None,
|
||
callback: Optional[Any] = None,
|
||
disable: bool = False,
|
||
*,
|
||
c2: float = _HFX_C2,
|
||
eta: float = _HFX_ETA,
|
||
mode: str = 'sharp',
|
||
# Per-mode overrides
|
||
boost_factor: Optional[float] = None,
|
||
sde_strength: Optional[float] = None,
|
||
momentum_strength: Optional[float] = None,
|
||
spectral_alpha: Optional[float] = None,
|
||
focus_gamma: Optional[float] = None,
|
||
) -> torch.Tensor:
|
||
"""
|
||
Experimental HFE sampler with 10 fundamentally different enhancement modes.
|
||
|
||
mode:
|
||
'sharp' -- Unsharp mask on eps_2 inside integrator: amplifies the
|
||
model's own fine-scale predictions.
|
||
'boost' -- Uniform eps_2 scaling inside integrator ("lying sigma"):
|
||
makes the model take a larger step toward its prediction.
|
||
'detail' -- Post-step HF injection from denoised_2: adds detail the
|
||
model predicted but the integrator smoothed away.
|
||
'stochastic' -- Structure-aware noise injection post-step: non-deterministic
|
||
variation weighted by local image detail.
|
||
'momentum' -- Cross-step temporal EMA: amplifies the direction the model's
|
||
prediction is moving between steps (temporal memory).
|
||
'spectral' -- FFT frequency reshaping of eps_2: power-law spectral boost
|
||
for precise frequency band control.
|
||
'orthogonal' -- Gram-Schmidt projection: amplifies the component of eps_2
|
||
orthogonal to eps_1 (novel information from stage 2).
|
||
'refine' -- ODE curvature-adaptive emphasis: amplifies eps_2 more where
|
||
|eps_2 - eps_1| is high (local truncation error is large).
|
||
'focus' -- Value-domain power-law contrast: nonlinear gain based on
|
||
element-wise magnitude (divisive normalization).
|
||
'coherence' -- Inter-stage FFT phase coherence gating: amplifies frequency
|
||
bins where eps_1 and eps_2 agree structurally.
|
||
"""
|
||
if extra_args is None:
|
||
extra_args = {}
|
||
|
||
_boost_f = boost_factor if boost_factor is not None else _HFX_BOOST_FACTOR
|
||
_sde_s = sde_strength if sde_strength is not None else _HFX_SDE_STRENGTH
|
||
_momentum_s = momentum_strength if momentum_strength is not None else _HFX_MOMENTUM_STRENGTH
|
||
_spectral_a = spectral_alpha if spectral_alpha is not None else _HFX_SPECTRAL_ALPHA
|
||
_focus_g = focus_gamma if focus_gamma is not None else _HFX_FOCUS_GAMMA
|
||
|
||
s_in = x.new_ones([x.shape[0]])
|
||
sigmas = sigmas.to(device=x.device, dtype=x.dtype)
|
||
total_steps = len(sigmas) - 1
|
||
|
||
# Schedule-based denoise gate: full for txt2img (sigmas[0] >= 0.7),
|
||
# quadratically suppressed for img2img (truncated schedule).
|
||
_max_sigma = float(sigmas[0])
|
||
_denoise_gate = min(1.0, (_max_sigma / 0.7) ** 2)
|
||
|
||
# Cross-step state for momentum mode
|
||
denoised_prev = None
|
||
|
||
for i in trange(total_steps, disable=disable):
|
||
sigma = sigmas[i]
|
||
sigma_next = sigmas[i + 1]
|
||
|
||
if sigma_next < 1e-6:
|
||
denoised = model(x, sigma * s_in, **extra_args)
|
||
x = denoised
|
||
if callback is not None:
|
||
callback({"i": i, "sigma": 0.0, "denoised": denoised, "x": x})
|
||
break
|
||
|
||
h = torch.log(sigma / sigma_next)
|
||
phi1_h = _phi1(h)
|
||
phi2_h = _phi2(h)
|
||
|
||
# --- Stage 1 ---
|
||
denoised_1 = model(x, sigma * s_in, **extra_args)
|
||
eps_1 = denoised_1 - x
|
||
|
||
# --- Stage 2 ---
|
||
hc2 = h * c2
|
||
a21 = c2 * _phi1(hc2)
|
||
X_2 = x + h * a21 * eps_1
|
||
sigma_mid = sigma * torch.exp(-c2 * h)
|
||
denoised_2 = model(X_2, sigma_mid * s_in, **extra_args)
|
||
denoised_2_orig = denoised_2 # Preserve for momentum tracking
|
||
eps_2 = denoised_2 - x
|
||
|
||
# --- Sigma warmup gate ---
|
||
progress = i / max(total_steps - 1, 1)
|
||
sigma_gate = max(0.0, min(1.0, (progress - 0.25) / 0.30))
|
||
eta_step = eta * sigma_gate
|
||
|
||
# Save eps_2 before any mode touches it (for per-step safety cap)
|
||
eps_2_pre = eps_2
|
||
|
||
# --- Mode: momentum (inside integrator -- temporal EMA on denoised_2) ---
|
||
if mode == 'momentum' and eta_step > 1e-3 and denoised_prev is not None:
|
||
temporal_diff = denoised_2 - denoised_prev
|
||
denoised_2 = denoised_2 + eta_step * _momentum_s * temporal_diff
|
||
eps_2 = denoised_2 - x # Recalculate eps_2 from modified denoised_2
|
||
|
||
# --- Mode: sharp (inside integrator -- modify eps_2) ---
|
||
if mode == 'sharp' and eta_step > 1e-3:
|
||
highpass = eps_2 - _spatial_lowpass(eps_2, kernel_size=5)
|
||
eps_2 = eps_2 + eta_step * 3.0 * highpass
|
||
|
||
# --- Mode: boost (inside integrator -- scale eps_2) ---
|
||
elif mode == 'boost' and eta_step > 1e-3:
|
||
eps_2 = eps_2 * (1.0 + eta_step * _boost_f)
|
||
|
||
# --- Mode: spectral (inside integrator -- FFT frequency reshaping) ---
|
||
elif mode == 'spectral' and eta_step > 1e-3:
|
||
if eps_2.ndim >= 4:
|
||
h_dim, w_dim = eps_2.shape[-2], eps_2.shape[-1]
|
||
eps_fft = torch.fft.rfft2(eps_2)
|
||
y_freq = torch.fft.fftfreq(h_dim, device=eps_2.device)
|
||
x_freq = torch.fft.rfftfreq(w_dim, device=eps_2.device)
|
||
freq = torch.sqrt(y_freq[:, None] ** 2 + x_freq[None, :] ** 2).clamp(min=1e-10)
|
||
boost = freq.pow(_spectral_a)
|
||
boost = boost / boost.mean()
|
||
eps_fft = eps_fft * (1.0 + eta_step * (boost - 1.0))
|
||
eps_2 = torch.fft.irfft2(eps_fft, s=(h_dim, w_dim))
|
||
else:
|
||
n_dim = eps_2.shape[-1]
|
||
eps_fft = torch.fft.rfft(eps_2)
|
||
freq = torch.fft.rfftfreq(n_dim, device=eps_2.device).clamp(min=1e-10)
|
||
boost = freq.pow(_spectral_a)
|
||
boost = boost / boost.mean()
|
||
eps_fft = eps_fft * (1.0 + eta_step * (boost - 1.0))
|
||
eps_2 = torch.fft.irfft(eps_fft, n=n_dim)
|
||
|
||
# --- Mode: orthogonal (inside integrator -- Gram-Schmidt) ---
|
||
elif mode == 'orthogonal' and eta_step > 1e-3:
|
||
e1 = eps_1.reshape(eps_1.shape[0], -1)
|
||
e2 = eps_2.reshape(eps_2.shape[0], -1)
|
||
e1_norm = e1 / e1.norm(dim=-1, keepdim=True).clamp(min=1e-8)
|
||
proj_coeff = (e2 * e1_norm).sum(dim=-1, keepdim=True)
|
||
projection = proj_coeff * e1_norm
|
||
ortho = (e2 - projection).reshape_as(eps_2)
|
||
eps_2 = eps_2 + eta_step * ortho
|
||
|
||
# --- Mode: refine (inside integrator -- ODE curvature-adaptive) ---
|
||
elif mode == 'refine' and eta_step > 1e-3:
|
||
curvature = (eps_2 - eps_1).abs()
|
||
curvature_smooth = _spatial_lowpass(curvature, kernel_size=5)
|
||
curvature_norm = curvature_smooth / curvature_smooth.mean().clamp(min=1e-8)
|
||
eps_2 = eps_2 * (1.0 + eta_step * (curvature_norm - 1.0))
|
||
|
||
# --- Mode: focus (inside integrator -- value-domain contrast) ---
|
||
elif mode == 'focus' and eta_step > 1e-3:
|
||
eps_mag = eps_2.abs().clamp(min=1e-8)
|
||
eps_ref = eps_mag.mean()
|
||
gain = (eps_mag / eps_ref).pow(_focus_g)
|
||
gain = gain / gain.mean() # energy preservation
|
||
eps_2 = eps_2 * (1.0 + eta_step * (gain - 1.0))
|
||
|
||
# --- Mode: coherence (inside integrator -- phase coherence gating) ---
|
||
elif mode == 'coherence' and eta_step > 1e-3:
|
||
if eps_2.ndim >= 4:
|
||
h_dim, w_dim = eps_2.shape[-2], eps_2.shape[-1]
|
||
e1_fft = torch.fft.rfft2(eps_1)
|
||
e2_fft = torch.fft.rfft2(eps_2)
|
||
phase_diff = torch.angle(e2_fft * e1_fft.conj())
|
||
coh = torch.cos(phase_diff)
|
||
gate = 1.0 + eta_step * coh
|
||
eps_2 = torch.fft.irfft2(e2_fft * gate, s=(h_dim, w_dim))
|
||
else:
|
||
n_dim = eps_2.shape[-1]
|
||
e1_fft = torch.fft.rfft(eps_1)
|
||
e2_fft = torch.fft.rfft(eps_2)
|
||
phase_diff = torch.angle(e2_fft * e1_fft.conj())
|
||
coh = torch.cos(phase_diff)
|
||
gate = 1.0 + eta_step * coh
|
||
eps_2 = torch.fft.irfft(e2_fft * gate, n=n_dim)
|
||
|
||
# --- Per-step safety cap on eps_2 modification ---
|
||
# Prevents compounding across steps from corrupting colors/structure.
|
||
# Cap delta to 10% of original eps_2 RMS per step.
|
||
if mode not in ('detail', 'stochastic') and eta_step > 1e-3:
|
||
_delta = eps_2 - eps_2_pre
|
||
_d_rms = _delta.square().mean().sqrt()
|
||
if _d_rms > 1e-8:
|
||
_ref_rms = eps_2_pre.square().mean().sqrt().clamp(min=1e-8)
|
||
_cap = 0.10 * _ref_rms
|
||
if _d_rms > _cap:
|
||
eps_2 = eps_2_pre + _delta * (_cap / _d_rms)
|
||
|
||
# --- Output weights (standard integrator step) ---
|
||
b2 = phi2_h / c2
|
||
b1 = phi1_h - b2
|
||
x = x + h * (b1 * eps_1 + b2 * eps_2)
|
||
|
||
# --- Mode: detail (post-step -- inject HF from denoised_2) ---
|
||
if mode == 'detail' and eta_step > 1e-3:
|
||
hf = denoised_2 - _spatial_lowpass(denoised_2, kernel_size=5)
|
||
sigma_scale = float(sigma)
|
||
corr_applied = eta_step * _denoise_gate * sigma_scale * hf
|
||
|
||
# Safety cap: limit to 3% of x RMS
|
||
_x_rms = x.square().mean().sqrt().clamp(min=1e-8)
|
||
_cap = 0.03 * _x_rms
|
||
_ca_rms = corr_applied.square().mean().sqrt()
|
||
if _ca_rms > _cap:
|
||
corr_applied = corr_applied * (_cap / _ca_rms)
|
||
|
||
# Apply with mean preservation (spatial dims only)
|
||
_spatial_dims = tuple(range(2, x.ndim))
|
||
x_mean = x.mean(dim=_spatial_dims, keepdim=True)
|
||
x = x + corr_applied
|
||
x = x - x.mean(dim=_spatial_dims, keepdim=True) + x_mean
|
||
|
||
# --- Mode: stochastic (post-step -- structure-aware noise) ---
|
||
elif mode == 'stochastic' and eta_step > 1e-3:
|
||
if float(sigma_next) > 0:
|
||
noise = torch.randn_like(x)
|
||
# Detail-weighted: more noise where image has structure
|
||
detail_energy = (x - _spatial_lowpass(x, kernel_size=5)).abs()
|
||
detail_norm = detail_energy / detail_energy.mean().clamp(min=1e-8)
|
||
x = x + (eta_step * float(sigma_next) * _sde_s
|
||
* noise * detail_norm)
|
||
|
||
# Update cross-step state for momentum mode
|
||
if mode == 'momentum':
|
||
denoised_prev = denoised_2_orig
|
||
|
||
# NaN/inf guard
|
||
if torch.isnan(x).any() or torch.isinf(x).any():
|
||
LOGGER.warning("HFX %s step %d: NaN/inf detected, disabling "
|
||
"emphasis for remaining steps.", mode, i)
|
||
eta = 0.0
|
||
x = x.nan_to_num(nan=0.0, posinf=1e4, neginf=-1e4)
|
||
|
||
if callback is not None:
|
||
callback({
|
||
"i": i,
|
||
"sigma": float(sigma_next),
|
||
"denoised": denoised_2,
|
||
"x": x,
|
||
})
|
||
|
||
return x
|
||
|
||
|
||
# --- Experimental sampler wrappers (base, no strength suffix) ---
|
||
|
||
def sample_hfx_sharp(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = _HFX_ETA):
|
||
"""Unsharp mask on eps_2 -- amplifies model's fine-scale predictions."""
|
||
LOGGER.info(">>> hfx_sharp sampler invoked (%d sigmas, eta=%.3f)",
|
||
len(sigmas), eta)
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
|
||
mode='sharp', eta=eta)
|
||
|
||
|
||
def sample_hfx_boost(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = _HFX_ETA):
|
||
"""Uniform eps_2 scaling (lying sigma) -- larger steps toward prediction."""
|
||
LOGGER.info(">>> hfx_boost sampler invoked (%d sigmas, eta=%.3f)",
|
||
len(sigmas), eta)
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
|
||
mode='boost', eta=eta)
|
||
|
||
|
||
def sample_hfx_detail(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = _HFX_ETA):
|
||
"""Post-step HF injection from denoised_2 -- recovers smoothed detail."""
|
||
LOGGER.info(">>> hfx_detail sampler invoked (%d sigmas, eta=%.3f)",
|
||
len(sigmas), eta)
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
|
||
mode='detail', eta=eta)
|
||
|
||
|
||
def sample_hfx_stochastic(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = _HFX_ETA):
|
||
"""Structure-aware noise injection -- non-deterministic variation."""
|
||
LOGGER.info(">>> hfx_stochastic sampler invoked (%d sigmas, eta=%.3f)",
|
||
len(sigmas), eta)
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
|
||
mode='stochastic', eta=eta)
|
||
|
||
|
||
def sample_hfx_momentum(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = _HFX_ETA):
|
||
"""Cross-step temporal EMA -- amplifies direction of prediction change."""
|
||
LOGGER.info(">>> hfx_momentum sampler invoked (%d sigmas, eta=%.3f)",
|
||
len(sigmas), eta)
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
|
||
mode='momentum', eta=eta)
|
||
|
||
|
||
def sample_hfx_spectral(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = _HFX_ETA):
|
||
"""FFT frequency reshaping -- power-law spectral boost on eps_2."""
|
||
LOGGER.info(">>> hfx_spectral sampler invoked (%d sigmas, eta=%.3f)",
|
||
len(sigmas), eta)
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
|
||
mode='spectral', eta=eta)
|
||
|
||
|
||
def sample_hfx_orthogonal(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = _HFX_ETA):
|
||
"""Gram-Schmidt projection -- amplifies novel info from stage 2."""
|
||
LOGGER.info(">>> hfx_orthogonal sampler invoked (%d sigmas, eta=%.3f)",
|
||
len(sigmas), eta)
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
|
||
mode='orthogonal', eta=eta)
|
||
|
||
|
||
def sample_hfx_refine(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = _HFX_ETA):
|
||
"""ODE curvature-adaptive emphasis -- amplifies where integrator is least accurate."""
|
||
LOGGER.info(">>> hfx_refine sampler invoked (%d sigmas, eta=%.3f)",
|
||
len(sigmas), eta)
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
|
||
mode='refine', eta=eta)
|
||
|
||
|
||
def sample_hfx_focus(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = _HFX_ETA):
|
||
"""Value-domain power-law contrast -- amplifies dominant corrections."""
|
||
LOGGER.info(">>> hfx_focus sampler invoked (%d sigmas, eta=%.3f)",
|
||
len(sigmas), eta)
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
|
||
mode='focus', eta=eta)
|
||
|
||
|
||
def sample_hfx_coherence(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = _HFX_ETA):
|
||
"""Inter-stage phase coherence gating -- trusts structurally confident frequencies."""
|
||
LOGGER.info(">>> hfx_coherence sampler invoked (%d sigmas, eta=%.3f)",
|
||
len(sigmas), eta)
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
|
||
mode='coherence', eta=eta)
|
||
|
||
|
||
# =====================================================================
|
||
# Graduated experimental presets (hfx_*_s1..s4)
|
||
# =====================================================================
|
||
#
|
||
# 4 strength tiers per mode, sweeping the key parameter for each mode.
|
||
# All use c2=0.65 (moderate).
|
||
|
||
_HFX_LEVELS = 4
|
||
|
||
# Per-mode sweep definitions: (mode, param_name, values_s1_to_s4)
|
||
_HFX_SWEEPS = {
|
||
'sharp': {
|
||
'param': 'eta',
|
||
'values': (0.10, 0.30, 0.70, 1.50),
|
||
},
|
||
'boost': {
|
||
'param': 'boost_factor',
|
||
'values': (0.2, 0.6, 1.2, 2.5),
|
||
},
|
||
'detail': {
|
||
'param': 'eta',
|
||
'values': (0.10, 0.30, 0.70, 1.50),
|
||
},
|
||
'stochastic': {
|
||
'param': 'sde_strength',
|
||
'values': (0.03, 0.10, 0.25, 0.50),
|
||
},
|
||
'momentum': {
|
||
'param': 'momentum_strength',
|
||
'values': (0.15, 0.50, 1.00, 2.00),
|
||
},
|
||
'spectral': {
|
||
'param': 'spectral_alpha',
|
||
'values': (0.3, 0.8, 1.5, 2.5),
|
||
},
|
||
'orthogonal': {
|
||
'param': 'eta',
|
||
'values': (0.10, 0.30, 0.70, 1.50),
|
||
},
|
||
'refine': {
|
||
'param': 'eta',
|
||
'values': (0.10, 0.30, 0.70, 1.50),
|
||
},
|
||
'focus': {
|
||
'param': 'focus_gamma',
|
||
'values': (0.4, 1.0, 1.8, 3.0),
|
||
},
|
||
'coherence': {
|
||
'param': 'eta',
|
||
'values': (0.10, 0.30, 0.70, 1.50),
|
||
},
|
||
}
|
||
|
||
|
||
def _make_hfx_preset(mode: str, level: int):
|
||
"""Factory: create a graduated experimental sampler.
|
||
|
||
For 'eta' sweeps, eta varies and mode-specific param stays default.
|
||
For mode-specific param sweeps, eta stays at _HFX_ETA and param varies.
|
||
"""
|
||
sweep = _HFX_SWEEPS[mode]
|
||
param = sweep['param']
|
||
value = sweep['values'][level]
|
||
|
||
if param == 'eta':
|
||
def sampler(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = value):
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback,
|
||
disable, mode=mode, eta=eta)
|
||
desc = f"eta={value:.2f}"
|
||
else:
|
||
# Mode-specific param sweep (e.g. boost_factor, sde_strength).
|
||
# Expose eta as the standard universal knob and the per-mode
|
||
# param under its native name so ManualSampler can introspect
|
||
# both.
|
||
param_default = value
|
||
|
||
def _make_param_sampler():
|
||
def sampler(model, x, sigmas, extra_args=None, callback=None,
|
||
disable=False, eta: float = _HFX_ETA,
|
||
**mode_kwargs):
|
||
kw = {param: mode_kwargs.get(param, param_default), "eta": eta}
|
||
return _sample_hfx(model, x, sigmas, extra_args, callback,
|
||
disable, mode=mode, **kw)
|
||
return sampler
|
||
sampler = _make_param_sampler()
|
||
desc = f"{param}={value}"
|
||
|
||
name = f"hfx_{mode}_s{level + 1}"
|
||
sampler.__name__ = f"sample_{name}"
|
||
sampler.__qualname__ = sampler.__name__
|
||
sampler.__doc__ = f"HFX {mode} strength {level + 1}/{_HFX_LEVELS} -- {desc}"
|
||
return name, sampler
|
||
|
||
|
||
_HFX_PRESETS = {}
|
||
for _mode in _HFX_SWEEPS:
|
||
for _lvl in range(_HFX_LEVELS):
|
||
_name, _fn = _make_hfx_preset(_mode, _lvl)
|
||
_HFX_PRESETS[_name] = _fn
|
||
|
||
|
||
# =====================================================================
|
||
# Two-stage S-curve schedulers (bong_tangent architecture)
|
||
# =====================================================================
|
||
#
|
||
# Two-stage design: stage 1 (σ_max → σ_mid) for structure,
|
||
# stage 2 (σ_mid → σ_min) for detail. Each stage applies a curve
|
||
# function with independent slope/pivot. Direct sigma mapping -- no
|
||
# Karras power-law.
|
||
#
|
||
# slope_adj = slope / (steps / 40) [step-count normalization]
|
||
|
||
|
||
# --- Curve functions ---------------------------------------------------
|
||
#
|
||
# Each takes (xs, pivot, slope) and returns raw monotone-decreasing values.
|
||
# Callers normalize the output to [0, 1].
|
||
|
||
def _curve_atan(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
|
||
"""Arctangent S-curve (same basis as bong_tangent)."""
|
||
return ((2.0 / math.pi) * torch.atan(-slope * (xs - pivot)) + 1.0) / 2.0
|
||
|
||
|
||
def _curve_logistic(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
|
||
"""Logistic sigmoid S-curve (exponential tails, sharper than atan)."""
|
||
return 1.0 / (1.0 + torch.exp(slope * (xs - pivot)))
|
||
|
||
|
||
def _curve_cosine(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
|
||
"""Cosine S-curve (smoothest, no sharp inflection).
|
||
|
||
*slope* controls sharpness via a power warp on the normalized position:
|
||
slope < 1 compresses the curve toward the middle (sharper transition),
|
||
slope > 1 spreads it toward the extremes (gentler).
|
||
"""
|
||
n = xs.shape[0]
|
||
if n <= 1:
|
||
return torch.ones_like(xs)
|
||
t = (xs - xs[0]) / (xs[-1] - xs[0]) # [0, 1]
|
||
# Power warp: slope acts as exponent (higher = gentler)
|
||
t_warped = t.clamp(0.0, 1.0).pow(max(slope, 0.01))
|
||
return (1.0 + torch.cos(math.pi * t_warped)) / 2.0
|
||
|
||
|
||
def _curve_kumaraswamy(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
|
||
"""Kumaraswamy power-law S-curve (inherently asymmetric).
|
||
|
||
Uses the Kumaraswamy CDF as a closed-form approximation of the
|
||
regularized incomplete beta function. *slope* controls the
|
||
concentration exponent (higher = stronger bend). The *pivot* sets
|
||
the balance between head and tail weighting: pivot < n/2 biases
|
||
toward early steps, pivot > n/2 biases toward late steps.
|
||
|
||
Note: this is NOT ComfyUI's BetaSchedulerNode (`scipy.stats.beta.ppf`
|
||
over the model's timestep table). Different math, different shape.
|
||
"""
|
||
n = xs.shape[0]
|
||
if n <= 1:
|
||
return torch.ones_like(xs)
|
||
t = (xs - xs[0]) / (xs[-1] - xs[0]) # [0, 1]
|
||
pivot_norm = max(min(pivot / max(n - 1, 1), 0.95), 0.05)
|
||
concentration = max(slope * 5.0, 0.1)
|
||
a = max(concentration * (1.0 - pivot_norm), 0.01)
|
||
b = max(concentration * pivot_norm, 0.01)
|
||
cdf = 1.0 - (1.0 - t.clamp(1e-7, 1.0 - 1e-7).pow(a)).pow(b)
|
||
return 1.0 - cdf
|
||
|
||
|
||
def _curve_laplacian(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
|
||
"""Laplacian (double-exponential) S-curve.
|
||
|
||
Sharper peak than logistic — creates very tight concentration at the
|
||
pivot with rapid exponential falloff on both sides.
|
||
"""
|
||
diff = slope * (xs - pivot)
|
||
# Laplace CDF: 0.5 * exp(x) for x<0, 1 - 0.5*exp(-x) for x>=0
|
||
# We want a decreasing function, so we negate the argument.
|
||
cdf = torch.where(
|
||
diff <= 0,
|
||
0.5 * torch.exp(diff),
|
||
1.0 - 0.5 * torch.exp(-diff),
|
||
)
|
||
return 1.0 - cdf # decreasing
|
||
|
||
|
||
def _curve_linear(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
|
||
"""Pure linear descent (no S-curve). Useful as a baseline reference."""
|
||
n = xs.shape[0]
|
||
if n <= 1:
|
||
return torch.ones_like(xs)
|
||
return torch.linspace(1.0, 0.0, n, dtype=xs.dtype, device=xs.device)
|
||
|
||
|
||
# --- Core two-stage engine -------------------------------------------
|
||
|
||
def _stage_sigmas(
|
||
curve_fn,
|
||
n: int,
|
||
slope: float,
|
||
pivot: float,
|
||
start: float,
|
||
end: float,
|
||
) -> torch.Tensor:
|
||
"""Apply *curve_fn* over *n* steps, mapping [start, end] sigma range.
|
||
|
||
Faithfully reproduces bong_tangent's get_bong_tangent_sigmas() logic
|
||
for any pluggable curve function.
|
||
"""
|
||
if n < 1:
|
||
return torch.tensor([], dtype=torch.float64)
|
||
|
||
xs = torch.arange(n, dtype=torch.float64)
|
||
raw = curve_fn(xs, pivot, slope)
|
||
|
||
r_max = raw[0].item()
|
||
r_min = raw[-1].item()
|
||
r_range = r_max - r_min
|
||
|
||
if r_range < 1e-12:
|
||
normalized = torch.linspace(1.0, 0.0, n, dtype=torch.float64)
|
||
else:
|
||
normalized = (raw - r_min) / r_range
|
||
|
||
return end + normalized * (start - end)
|
||
|
||
|
||
def _two_stage_sigmas(
|
||
curve_fn,
|
||
steps: int,
|
||
sigma_max: float,
|
||
sigma_min: float,
|
||
slope_1: float = 0.2,
|
||
slope_2: float = 0.2,
|
||
pivot_frac_1: float = 0.6,
|
||
pivot_frac_2: float = 0.6,
|
||
mid_frac: float = 0.5,
|
||
) -> torch.Tensor:
|
||
"""Two-stage S-curve schedule (bong_tangent architecture).
|
||
|
||
Parameters
|
||
----------
|
||
curve_fn : callable
|
||
One of ``_curve_atan``, ``_curve_logistic``, ``_curve_cosine``.
|
||
steps : int
|
||
Total sampling steps requested.
|
||
sigma_max, sigma_min : float
|
||
Sigma bounds from ``model_sampling``.
|
||
slope_1, slope_2 : float
|
||
Curve concentration per stage (before step-count normalization).
|
||
pivot_frac_1, pivot_frac_2 : float
|
||
Pivot position within the *total* step range (0‥1), matching
|
||
bong_tangent's convention where the pivot is relative to total
|
||
steps, not per-stage.
|
||
mid_frac : float
|
||
Where to split the sigma range: ``sigma_mid = sigma_min +
|
||
mid_frac * (sigma_max - sigma_min)``.
|
||
"""
|
||
# Match bong_tangent: pad by 2 then trim junction
|
||
n = steps + 2
|
||
sigma_mid = sigma_min + mid_frac * (sigma_max - sigma_min)
|
||
|
||
# Split steps the same way bong_tangent does
|
||
midpoint = int((n * pivot_frac_1 + n * pivot_frac_2) / 2)
|
||
stage_1_len = midpoint
|
||
stage_2_len = n - midpoint
|
||
|
||
# Absolute pivot indices (bong_tangent convention)
|
||
piv_1 = int(n * pivot_frac_1)
|
||
piv_2 = int(n * pivot_frac_2) - stage_1_len # relative to stage 2
|
||
|
||
# Step-count-normalized slopes
|
||
s1 = slope_1 / max(n / 40.0, 0.1)
|
||
s2 = slope_2 / max(n / 40.0, 0.1)
|
||
|
||
sigmas_1 = _stage_sigmas(curve_fn, stage_1_len, s1, piv_1,
|
||
sigma_max, sigma_mid)
|
||
sigmas_2 = _stage_sigmas(curve_fn, stage_2_len, s2, piv_2,
|
||
sigma_mid, sigma_min)
|
||
|
||
# Drop last of stage 1 (duplicate at junction)
|
||
if len(sigmas_1) > 0:
|
||
sigmas_1 = sigmas_1[:-1]
|
||
|
||
sigmas = torch.cat([sigmas_1, sigmas_2,
|
||
torch.zeros(1, dtype=torch.float64)])
|
||
return sigmas.float()
|
||
|
||
|
||
# --- Schedule wrapper -------------------------------------------------
|
||
|
||
def _tangent_schedule(
|
||
model_sampling: Any,
|
||
steps: int,
|
||
curve_fn,
|
||
slope_1: float,
|
||
slope_2: float,
|
||
pivot_frac_1: float = 0.6,
|
||
pivot_frac_2: float = 0.6,
|
||
mid_frac: float = 0.5,
|
||
name: str = '',
|
||
) -> torch.Tensor:
|
||
sigma_max = float(model_sampling.sigma_max)
|
||
sigma_min = float(model_sampling.sigma_min)
|
||
sigmas = _two_stage_sigmas(
|
||
curve_fn, steps, sigma_max, sigma_min,
|
||
slope_1=slope_1, slope_2=slope_2,
|
||
pivot_frac_1=pivot_frac_1, pivot_frac_2=pivot_frac_2,
|
||
mid_frac=mid_frac,
|
||
)
|
||
if name:
|
||
_plot_sigmas(sigmas, name)
|
||
return sigmas
|
||
|
||
|
||
# --- Presets -----------------------------------------------------------
|
||
|
||
def scheduler_atan_gentle(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Mild arctangent concentration (closest to bong_tangent defaults)."""
|
||
return _tangent_schedule(model_sampling, steps, _curve_atan,
|
||
slope_1=0.15, slope_2=0.15,
|
||
name='atan_gentle')
|
||
|
||
|
||
def scheduler_atan_focused(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Moderate arctangent concentration."""
|
||
return _tangent_schedule(model_sampling, steps, _curve_atan,
|
||
slope_1=0.25, slope_2=0.25,
|
||
name='atan_focused')
|
||
|
||
|
||
def scheduler_atan_steep(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Aggressive arctangent concentration."""
|
||
return _tangent_schedule(model_sampling, steps, _curve_atan,
|
||
slope_1=0.40, slope_2=0.40,
|
||
name='atan_steep')
|
||
|
||
|
||
def scheduler_logistic(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Logistic S-curve (sharper transitions than arctangent)."""
|
||
return _tangent_schedule(model_sampling, steps, _curve_logistic,
|
||
slope_1=0.20, slope_2=0.20,
|
||
name='logistic')
|
||
|
||
|
||
def scheduler_cosine(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Cosine S-curve (smoothest, most gradual transitions)."""
|
||
return _tangent_schedule(model_sampling, steps, _curve_cosine,
|
||
slope_1=1.0, slope_2=1.0,
|
||
name='cosine')
|
||
|
||
|
||
def scheduler_kumaraswamy(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Kumaraswamy power-law S-curve (inherently asymmetric concentration).
|
||
|
||
Distinct from ComfyUI's `beta` (BetaSchedulerNode) — see
|
||
`_curve_kumaraswamy` for the math.
|
||
"""
|
||
return _tangent_schedule(model_sampling, steps, _curve_kumaraswamy,
|
||
slope_1=0.20, slope_2=0.20,
|
||
name='kumaraswamy')
|
||
|
||
|
||
def scheduler_laplacian(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Laplacian S-curve (tightest pivot concentration, sharp falloff)."""
|
||
return _tangent_schedule(model_sampling, steps, _curve_laplacian,
|
||
slope_1=0.20, slope_2=0.20,
|
||
name='laplacian')
|
||
|
||
|
||
def scheduler_linear(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Linear two-stage (no S-curve, baseline reference)."""
|
||
return _tangent_schedule(model_sampling, steps, _curve_linear,
|
||
slope_1=0.20, slope_2=0.20,
|
||
name='linear')
|
||
|
||
|
||
# --- Asymmetric presets ------------------------------------------------
|
||
# Same curves but with different slope_1 vs slope_2 to bias toward
|
||
# structure (high-sigma lingering) or detail (low-sigma lingering).
|
||
|
||
def scheduler_atan_structure(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Arctangent biased toward structure: gentle stage 1, steep stage 2."""
|
||
return _tangent_schedule(model_sampling, steps, _curve_atan,
|
||
slope_1=0.10, slope_2=0.35,
|
||
name='atan_structure')
|
||
|
||
|
||
def scheduler_atan_detail(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Arctangent biased toward detail: steep stage 1, gentle stage 2."""
|
||
return _tangent_schedule(model_sampling, steps, _curve_atan,
|
||
slope_1=0.35, slope_2=0.10,
|
||
name='atan_detail')
|
||
|
||
|
||
def scheduler_logistic_structure(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Logistic biased toward structure: gentle stage 1, steep stage 2."""
|
||
return _tangent_schedule(model_sampling, steps, _curve_logistic,
|
||
slope_1=0.10, slope_2=0.30,
|
||
name='logistic_structure')
|
||
|
||
|
||
def scheduler_logistic_detail(model_sampling: Any, steps: int) -> torch.Tensor:
|
||
"""Logistic biased toward detail: steep stage 1, gentle stage 2."""
|
||
return _tangent_schedule(model_sampling, steps, _curve_logistic,
|
||
slope_1=0.30, slope_2=0.10,
|
||
name='logistic_detail')
|
||
|
||
|
||
# =====================================================================
|
||
# Registration
|
||
# =====================================================================
|
||
|
||
_SAMPLERS: Dict[str, Any] = {}
|
||
_SAMPLERS.update(_HFE_PRESETS) # hfe_s1..s8 (2-stage default; stages kwarg up to 5)
|
||
_SAMPLERS["hfe_auto"] = sample_hfe_auto # adaptive (stages 2..5 via kwarg)
|
||
# hfe3_auto / hfe4_auto / hfe5_auto and hfe3_s* / hfe4_s* / hfe5_s* have
|
||
# been removed -- pick stages 2..5 via ManualSampler instead.
|
||
# Experimental base samplers (10 fundamentally different modes)
|
||
_SAMPLERS["hfx_sharp"] = sample_hfx_sharp
|
||
_SAMPLERS["hfx_boost"] = sample_hfx_boost
|
||
_SAMPLERS["hfx_detail"] = sample_hfx_detail
|
||
_SAMPLERS["hfx_stochastic"] = sample_hfx_stochastic
|
||
_SAMPLERS["hfx_momentum"] = sample_hfx_momentum
|
||
_SAMPLERS["hfx_spectral"] = sample_hfx_spectral
|
||
_SAMPLERS["hfx_orthogonal"] = sample_hfx_orthogonal
|
||
_SAMPLERS["hfx_refine"] = sample_hfx_refine
|
||
_SAMPLERS["hfx_focus"] = sample_hfx_focus
|
||
_SAMPLERS["hfx_coherence"] = sample_hfx_coherence
|
||
# Graduated experimental presets (hfx_*_s1..s4 for each mode)
|
||
_SAMPLERS.update(_HFX_PRESETS)
|
||
|
||
_SCHEDULERS = {
|
||
# Symmetric atan (bong_tangent-derived)
|
||
"atan_gentle": scheduler_atan_gentle,
|
||
"atan_focused": scheduler_atan_focused,
|
||
"atan_steep": scheduler_atan_steep,
|
||
# Alternative curves
|
||
"logistic": scheduler_logistic,
|
||
"cosine": scheduler_cosine,
|
||
"kumaraswamy": scheduler_kumaraswamy,
|
||
"laplacian": scheduler_laplacian,
|
||
"linear": scheduler_linear,
|
||
# Asymmetric presets
|
||
"atan_structure": scheduler_atan_structure,
|
||
"atan_detail": scheduler_atan_detail,
|
||
"logistic_structure": scheduler_logistic_structure,
|
||
"logistic_detail": scheduler_logistic_detail,
|
||
}
|
||
|
||
# Old names from all previous versions
|
||
_OLD_NAMES = [
|
||
"euler_hfdetail", "hfdetail_power",
|
||
"hfdetail_soft", "hfdetail", "hfdetail_strong",
|
||
"res_2s_soft", "res_2s_sharp", "res_2s_crisp",
|
||
"tangent_soft", "tangent_sharp", "tangent_crisp",
|
||
"hfe_soft", "hfe_sharp", "hfe_crisp",
|
||
# Old experimental modes (replaced by sharp/boost/detail/stochastic)
|
||
"hfx_lap", "hfx_mom", "hfx_fft", "hfx_sde", "hfx_spatial",
|
||
"hfx_lap_mom", "hfx_lap_spatial", "hfx_fft_spatial",
|
||
"hfx_lap_fine", "hfx_lap_broad",
|
||
# Old graduated presets
|
||
"hfx_lap_s1", "hfx_lap_s2", "hfx_lap_s3", "hfx_lap_s4",
|
||
"hfx_mom_s1", "hfx_mom_s2", "hfx_mom_s3", "hfx_mom_s4",
|
||
"hfx_fft_s1", "hfx_fft_s2", "hfx_fft_s3", "hfx_fft_s4",
|
||
"hfx_sde_s1", "hfx_sde_s2", "hfx_sde_s3", "hfx_sde_s4",
|
||
"hfx_spatial_s1", "hfx_spatial_s2", "hfx_spatial_s3", "hfx_spatial_s4",
|
||
]
|
||
# 3/4/5-stage variants -- now reachable via sample_hfe_auto(stages=N) or
|
||
# any sample_hfe_s<level>(stages=N) through ManualSampler.
|
||
for _stages in (3, 4, 5):
|
||
_OLD_NAMES.append(f"hfe{_stages}_auto")
|
||
for _lvl in range(1, 9):
|
||
_OLD_NAMES.append(f"hfe{_stages}_s{_lvl}")
|
||
|
||
|
||
def _unregister_old() -> None:
|
||
"""Remove entries from previous versions."""
|
||
for attr in ("KSAMPLER_NAMES", "SAMPLER_NAMES", "SCHEDULER_NAMES"):
|
||
names = getattr(comfy_samplers, attr, None)
|
||
if isinstance(names, (list, tuple)):
|
||
names = list(names)
|
||
changed = False
|
||
for old in _OLD_NAMES:
|
||
if old in names:
|
||
names.remove(old)
|
||
changed = True
|
||
if changed:
|
||
setattr(comfy_samplers, attr, names)
|
||
|
||
KSampler = getattr(comfy_samplers, "KSampler", None)
|
||
if KSampler is not None and hasattr(KSampler, "SAMPLERS"):
|
||
samplers = list(getattr(KSampler, "SAMPLERS"))
|
||
changed = False
|
||
for old in _OLD_NAMES:
|
||
if old in samplers:
|
||
samplers.remove(old)
|
||
changed = True
|
||
if changed:
|
||
KSampler.SAMPLERS = samplers
|
||
|
||
kdiff = getattr(comfy_samplers, "k_diffusion_sampling", None)
|
||
if kdiff is not None:
|
||
for old in _OLD_NAMES:
|
||
attr = f"sample_{old}"
|
||
if hasattr(kdiff, attr):
|
||
delattr(kdiff, attr)
|
||
|
||
handlers = getattr(comfy_samplers, "SCHEDULER_HANDLERS", None)
|
||
if isinstance(handlers, dict):
|
||
for old in _OLD_NAMES:
|
||
handlers.pop(old, None)
|
||
|
||
|
||
def _register_samplers() -> None:
|
||
kdiff = getattr(comfy_samplers, "k_diffusion_sampling", None)
|
||
|
||
registered: List[str] = []
|
||
skipped: List[str] = []
|
||
for name, func in _SAMPLERS.items():
|
||
# Refuse to overwrite an existing sampler — even if the name is
|
||
# only present on kdiff (e.g. a built-in `sample_<name>`), bail
|
||
# so we don't shadow upstream behavior.
|
||
kdiff_attr = f"sample_{name}"
|
||
if kdiff is not None and hasattr(kdiff, kdiff_attr) \
|
||
and getattr(kdiff, kdiff_attr) is not func:
|
||
LOGGER.warning(
|
||
"RES4SHO: refusing to overwrite existing sampler '%s' "
|
||
"on k_diffusion_sampling.", name)
|
||
skipped.append(name)
|
||
continue
|
||
|
||
ksampler_names = getattr(comfy_samplers, "KSAMPLER_NAMES", None)
|
||
if isinstance(ksampler_names, (list, tuple)):
|
||
kl = list(ksampler_names)
|
||
if name not in kl:
|
||
kl.append(name)
|
||
comfy_samplers.KSAMPLER_NAMES = kl
|
||
|
||
sampler_names = getattr(comfy_samplers, "SAMPLER_NAMES", [])
|
||
if not isinstance(sampler_names, list):
|
||
sampler_names = list(sampler_names)
|
||
if name not in sampler_names:
|
||
sampler_names.append(name)
|
||
comfy_samplers.SAMPLER_NAMES = sampler_names
|
||
|
||
KSampler = getattr(comfy_samplers, "KSampler", None)
|
||
if KSampler is not None and hasattr(KSampler, "SAMPLERS"):
|
||
sl = list(getattr(KSampler, "SAMPLERS"))
|
||
if name not in sl:
|
||
sl.append(name)
|
||
KSampler.SAMPLERS = sl
|
||
|
||
if kdiff is not None:
|
||
setattr(kdiff, kdiff_attr, func)
|
||
registered.append(name)
|
||
|
||
LOGGER.info("HFE samplers registered: %s", registered)
|
||
if skipped:
|
||
LOGGER.warning("HFE samplers skipped (name collision): %s", skipped)
|
||
|
||
|
||
def _register_schedulers() -> None:
|
||
handlers = getattr(comfy_samplers, "SCHEDULER_HANDLERS", None)
|
||
|
||
registered: List[str] = []
|
||
skipped: List[str] = []
|
||
for name, func in _SCHEDULERS.items():
|
||
# Refuse to overwrite an existing handler — most importantly,
|
||
# ComfyUI's built-in `beta` (BetaSchedulerNode), which is a
|
||
# different math from anything we ship.
|
||
if isinstance(handlers, dict) and name in handlers:
|
||
LOGGER.warning(
|
||
"RES4SHO: refusing to overwrite existing scheduler "
|
||
"'%s' in SCHEDULER_HANDLERS.", name)
|
||
skipped.append(name)
|
||
continue
|
||
|
||
if isinstance(handlers, dict) and len(handlers) > 0:
|
||
any_handler = next(iter(handlers.values()))
|
||
HandlerType = type(any_handler)
|
||
handlers[name] = HandlerType(handler=func, use_ms=True)
|
||
|
||
names = getattr(comfy_samplers, "SCHEDULER_NAMES", [])
|
||
if not isinstance(names, list):
|
||
names = list(names)
|
||
if name not in names:
|
||
names.append(name)
|
||
comfy_samplers.SCHEDULER_NAMES = names
|
||
|
||
KSampler = getattr(comfy_samplers, "KSampler", None)
|
||
if KSampler is not None and hasattr(KSampler, "SCHEDULERS"):
|
||
sched_list = getattr(KSampler, "SCHEDULERS")
|
||
if not isinstance(sched_list, list):
|
||
sched_list = list(sched_list)
|
||
if name not in sched_list:
|
||
sched_list.append(name)
|
||
KSampler.SCHEDULERS = sched_list
|
||
registered.append(name)
|
||
|
||
LOGGER.info("HFE schedulers registered: %s", registered)
|
||
if skipped:
|
||
LOGGER.warning(
|
||
"HFE schedulers skipped (name collision): %s", skipped)
|
||
|
||
|
||
# =====================================================================
|
||
# Initialization
|
||
# =====================================================================
|
||
|
||
def initialize_hfdetail_extension() -> None:
|
||
try:
|
||
_unregister_old()
|
||
except Exception:
|
||
LOGGER.debug("Old HFDetail entries cleanup skipped.", exc_info=True)
|
||
|
||
try:
|
||
_register_samplers()
|
||
except Exception:
|
||
LOGGER.error("Failed to register HFE samplers.", exc_info=True)
|
||
|
||
try:
|
||
_register_schedulers()
|
||
except Exception:
|
||
LOGGER.error("Failed to register HFE schedulers.", exc_info=True)
|
||
|
||
|
||
initialize_hfdetail_extension()
|
||
|
||
NODE_CLASS_MAPPINGS: Dict[str, Any] = {}
|
||
NODE_DISPLAY_NAME_MAPPINGS: Dict[str, str] = {}
|