Files
WASasquatch-RES4SHO/sampling.py
T
2026-04-27 08:25:43 -07:00

1928 lines
66 KiB
Python

# -*- coding: utf-8 -*-
"""
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, 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 _extract_hf(t: torch.Tensor, kernel_size: int = 3) -> torch.Tensor:
"""
Spatial high-pass via residual after box blur.
For a 4D [B,C,H,W] latent tensor this isolates edges, texture, and
micro-structure. Returns zeros for non-4D inputs.
"""
if t.ndim != 4:
return torch.zeros_like(t)
pad = kernel_size // 2
padded = F.pad(t, [pad, pad, pad, pad], mode='reflect')
low = F.avg_pool2d(padded, kernel_size, stride=1)
return t - low
def _extract_hf_pyramid(t: torch.Tensor, levels: int = 3) -> list:
"""
Laplacian pyramid: decompose into multiple frequency bands.
Returns a list of [fine, medium, coarse] band tensors.
Each band captures progressively lower spatial frequencies via
increasing kernel sizes (3, 5, 7, ...).
"""
if t.ndim != 4:
return [torch.zeros_like(t)] * levels
bands = []
current = t
for lvl in range(levels):
ks = 3 + 2 * lvl # 3, 5, 7
pad = ks // 2
padded = F.pad(current, [pad, pad, pad, pad], mode='reflect')
blurred = F.avg_pool2d(padded, ks, stride=1)
bands.append(current - blurred)
current = blurred
return bands
def _extract_hf_fft(t: torch.Tensor, cutoff: float = 0.3) -> torch.Tensor:
"""
FFT high-pass filter with smooth ramp.
Extracts frequencies above ``cutoff`` (fraction of Nyquist).
Ramps linearly from 0 at cutoff to 1 at 0.5 (Nyquist).
"""
if t.ndim != 4:
return torch.zeros_like(t)
H, W = t.shape[2], t.shape[3]
freq = torch.fft.rfft2(t)
fy = torch.fft.fftfreq(H, device=t.device).unsqueeze(1)
fx = torch.fft.rfftfreq(W, device=t.device).unsqueeze(0)
freq_mag = torch.sqrt(fy ** 2 + fx ** 2)
ramp_width = max(0.5 - cutoff, 1e-6)
mask = torch.clamp((freq_mag - cutoff) / ramp_width, 0.0, 1.0)
return torch.fft.irfft2(freq * mask, s=(H, W))
def _spatial_gate(delta: torch.Tensor, window: int = 7) -> torch.Tensor:
"""
Per-pixel gate based on local energy of the correction delta.
Returns a [0, 1] spatial map: 1 in high-variance regions (faces,
text, fine objects) where emphasis helps; 0 in smooth areas (sky,
gradients) where emphasis would add noise.
"""
if delta.ndim != 4:
return torch.ones_like(delta)
energy = delta ** 2
pad = window // 2
padded = F.pad(energy, [pad, pad, pad, pad], mode='reflect')
local_energy = F.avg_pool2d(padded, window, stride=1)
e_max = local_energy.amax(dim=(-1, -2), keepdim=True).clamp(min=1e-8)
return local_energy / e_max
# =====================================================================
# Console sigma plot
# =====================================================================
def _plot_sigmas(sigmas: torch.Tensor, name: str,
width: int = 64, height: int = 16) -> 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
# Build character canvas
grid = [[' '] * width for _ in range(height)]
for i, v in enumerate(vals):
c = int(i * (width - 1) / (n - 1) + 0.5)
r = int((y_hi - v) * (height - 1) / y_span + 0.5)
grid[max(0, min(height - 1, r))][max(0, min(width - 1, 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 + eta_step * delta_hf
# --- 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 + eta_step * delta_hf
# --- 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
def sample_hfe_auto(model, x, sigmas, extra_args=None, callback=None, disable=False):
"""Adaptive HFE -- variable c2, eta, and kernel per step."""
LOGGER.info(">>> hfe_auto sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfe_auto(
model, x, sigmas, extra_args, callback, disable,
eta_peak=0.55, c2_start=0.45, c2_end=0.85,
)
# =====================================================================
# 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 + eta_step * delta_hf
# --- 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 + eta_step * delta_hf
# --- 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
def sample_hfe3_auto(model, x, sigmas, extra_args=None, callback=None, disable=False):
"""3-stage adaptive HFE."""
LOGGER.info(">>> hfe3_auto sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfe_3s_auto(model, x, sigmas, extra_args, callback, disable,
eta_peak=0.55)
# =====================================================================
# 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 + eta_step * delta_hf
# --- 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 + eta_step * delta_hf
# --- 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
def sample_hfe4_auto(model, x, sigmas, extra_args=None, callback=None, disable=False):
"""4-stage adaptive HFE."""
LOGGER.info(">>> hfe4_auto sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfe_4s_auto(model, x, sigmas, extra_args, callback, disable,
eta_peak=0.55)
# =====================================================================
# 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 + eta_step * delta_hf
# --- 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 + eta_step * delta_hf
# --- 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
def sample_hfe5_auto(model, x, sigmas, extra_args=None, callback=None, disable=False):
"""5-stage adaptive HFE."""
LOGGER.info(">>> hfe5_auto sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfe_5s_auto(model, x, sigmas, extra_args, callback, disable,
eta_peak=0.55)
# =====================================================================
# 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, stages: int = 2):
"""Factory: create a fixed-strength HFE sampler for the given level and stage count."""
core_fn, prefix = _STAGE_CORES[stages]
t = level / (_HFE_LEVELS - 1)
eta = _HFE_ETA_MAX * (t ** 1.5)
if stages == 2:
# 2-stage: also vary c2
c2 = _HFE_C2_MIN + t * (_HFE_C2_MAX - _HFE_C2_MIN)
def sampler(model, x, sigmas, extra_args=None, callback=None, disable=False,
_c2=c2, _eta=eta):
return core_fn(model, x, sigmas, extra_args, callback, disable,
c2=_c2, eta=_eta)
sampler.__doc__ = (f"HFE {stages}s strength {level + 1}/{_HFE_LEVELS}"
f" -- c2={c2:.3f}, eta={eta:.3f}")
else:
# 3/4/5-stage: fixed c values, only vary eta
def sampler(model, x, sigmas, extra_args=None, callback=None, disable=False,
_eta=eta, _prefix=prefix, _level=level):
LOGGER.info(">>> %s_s%d preset invoked, passing eta=%.4f",
_prefix, _level + 1, _eta)
return core_fn(model, x, sigmas, extra_args, callback, disable,
eta=_eta)
sampler.__doc__ = (f"HFE {stages}s strength {level + 1}/{_HFE_LEVELS}"
f" -- eta={eta:.3f}")
name = f"{prefix}_s{level + 1}"
sampler.__name__ = f"sample_{name}"
sampler.__qualname__ = sampler.__name__
return name, sampler
# Generate all presets: hfe_s1..s8, hfe3_s1..s8, hfe4_s1..s8, hfe5_s1..s8
_HFE_PRESETS = {}
for _stages in (2, 3, 4, 5):
for _lvl in range(_HFE_LEVELS):
_name, _fn = _make_hfe_preset(_lvl, _stages)
_HFE_PRESETS[_name] = _fn
# =====================================================================
# Experimental HFE samplers (hfx_*)
# =====================================================================
#
# Each variant modifies HOW high-frequency detail is extracted and/or
# applied, using a shared 2-stage exponential integrator base.
#
# All use fixed moderate strength (c2=0.65, eta=0.25) for direct
# comparison against hfe_s5.
_HFX_C2 = 0.65
_HFX_ETA = 0.25
_HFX_SDE_STRENGTH = 0.08
_HFX_MOM_BETA = 0.7
_HFX_FFT_CUTOFF = 0.3
_HFX_LAP_WEIGHTS = (1.5, 1.0, 0.5)
@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 = 'lap',
# Per-mode overrides (use module defaults when None)
lap_weights: Optional[tuple] = None,
mom_beta: Optional[float] = None,
fft_cutoff: Optional[float] = None,
sde_strength: Optional[float] = None,
spatial_window: Optional[int] = None,
) -> torch.Tensor:
"""
Generic experimental HFE sampler.
mode:
'lap' -- Laplacian pyramid multi-scale (3 bands, weighted)
'mom' -- correction momentum (EMA across steps)
'fft' -- FFT spectral high-pass with smooth cutoff
'sde' -- stochastic HF noise injection after update
'spatial' -- spatially-adaptive per-pixel gating
Hybrid modes (combine two techniques):
'lap_mom' -- Laplacian pyramid + momentum accumulation
'lap_spatial' -- Laplacian pyramid + spatial gating
'fft_spatial' -- FFT spectral + spatial gating
"""
if extra_args is None:
extra_args = {}
# Resolve per-mode defaults
_lap_w = lap_weights or _HFX_LAP_WEIGHTS
_mom_b = mom_beta if mom_beta is not None else _HFX_MOM_BETA
_fft_c = fft_cutoff if fft_cutoff is not None else _HFX_FFT_CUTOFF
_sde_s = sde_strength if sde_strength is not None else _HFX_SDE_STRENGTH
_sp_win = spatial_window if spatial_window is not None else 7
s_in = x.new_ones([x.shape[0]])
sigmas = sigmas.to(device=x.device, dtype=x.dtype)
total_steps = len(sigmas) - 1
# Per-mode state
momentum_buf = 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)
eps_2 = denoised_2 - x
# --- Experimental HFE ---
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 and mode != 'sde':
delta = eps_2 - eps_1
if mode == 'lap':
bands = _extract_hf_pyramid(delta, levels=3)
correction = sum(w * b for w, b in zip(_lap_w, bands))
eps_2 = eps_2 + eta_step * correction
elif mode == 'mom':
delta_hf = _extract_hf(delta)
if momentum_buf is None:
momentum_buf = delta_hf.clone()
else:
momentum_buf = (_mom_b * momentum_buf
+ (1.0 - _mom_b) * delta_hf)
eps_2 = eps_2 + eta_step * momentum_buf
elif mode == 'fft':
delta_hf = _extract_hf_fft(delta, cutoff=_fft_c)
eps_2 = eps_2 + eta_step * delta_hf
elif mode == 'spatial':
delta_hf = _extract_hf(delta)
gate = _spatial_gate(delta, window=_sp_win)
eps_2 = eps_2 + eta_step * gate * delta_hf
elif mode == 'lap_mom':
bands = _extract_hf_pyramid(delta, levels=3)
correction = sum(w * b for w, b in zip(_lap_w, bands))
if momentum_buf is None:
momentum_buf = correction.clone()
else:
momentum_buf = (_mom_b * momentum_buf
+ (1.0 - _mom_b) * correction)
eps_2 = eps_2 + eta_step * momentum_buf
elif mode == 'lap_spatial':
bands = _extract_hf_pyramid(delta, levels=3)
correction = sum(w * b for w, b in zip(_lap_w, bands))
gate = _spatial_gate(delta, window=_sp_win)
eps_2 = eps_2 + eta_step * gate * correction
elif mode == 'fft_spatial':
delta_hf = _extract_hf_fft(delta, cutoff=_fft_c)
gate = _spatial_gate(delta, window=_sp_win)
eps_2 = eps_2 + eta_step * gate * delta_hf
# --- Output weights ---
b2 = phi2_h / c2
b1 = phi1_h - b2
x = x + h * (b1 * eps_1 + b2 * eps_2)
# --- SDE: post-update HF noise injection ---
if mode == 'sde' and eta_step > 1e-3:
noise = torch.randn_like(x)
noise_hf = _extract_hf(noise)
x = x + (_sde_s * float(sigma_next) * sigma_gate * noise_hf)
# 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_lap(model, x, sigmas, extra_args=None, callback=None,
disable=False):
"""Laplacian pyramid multi-scale HFE (experimental)."""
LOGGER.info(">>> hfx_lap sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
mode='lap')
def sample_hfx_mom(model, x, sigmas, extra_args=None, callback=None,
disable=False):
"""Correction momentum HFE (experimental)."""
LOGGER.info(">>> hfx_mom sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
mode='mom')
def sample_hfx_fft(model, x, sigmas, extra_args=None, callback=None,
disable=False):
"""FFT spectral shaping HFE (experimental)."""
LOGGER.info(">>> hfx_fft sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
mode='fft')
def sample_hfx_sde(model, x, sigmas, extra_args=None, callback=None,
disable=False):
"""Stochastic HF injection HFE (experimental)."""
LOGGER.info(">>> hfx_sde sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
mode='sde')
def sample_hfx_spatial(model, x, sigmas, extra_args=None, callback=None,
disable=False):
"""Spatially-adaptive gating HFE (experimental)."""
LOGGER.info(">>> hfx_spatial sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
mode='spatial')
# --- Hybrid sampler wrappers ---
def sample_hfx_lap_mom(model, x, sigmas, extra_args=None, callback=None,
disable=False):
"""Laplacian pyramid + momentum HFE (experimental hybrid)."""
LOGGER.info(">>> hfx_lap_mom sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
mode='lap_mom')
def sample_hfx_lap_spatial(model, x, sigmas, extra_args=None, callback=None,
disable=False):
"""Laplacian pyramid + spatial gating HFE (experimental hybrid)."""
LOGGER.info(">>> hfx_lap_spatial sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
mode='lap_spatial')
def sample_hfx_fft_spatial(model, x, sigmas, extra_args=None, callback=None,
disable=False):
"""FFT spectral + spatial gating HFE (experimental hybrid)."""
LOGGER.info(">>> hfx_fft_spatial sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
mode='fft_spatial')
# --- Band profile variants for hfx_lap ---
def sample_hfx_lap_fine(model, x, sigmas, extra_args=None, callback=None,
disable=False):
"""Laplacian pyramid fine-detail emphasis (experimental)."""
LOGGER.info(">>> hfx_lap_fine sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
mode='lap', lap_weights=(2.5, 0.8, 0.2))
def sample_hfx_lap_broad(model, x, sigmas, extra_args=None, callback=None,
disable=False):
"""Laplacian pyramid broad/even emphasis (experimental)."""
LOGGER.info(">>> hfx_lap_broad sampler invoked (%d sigmas)", len(sigmas))
return _sample_hfx(model, x, sigmas, extra_args, callback, disable,
mode='lap', lap_weights=(1.0, 1.2, 1.0))
# =====================================================================
# 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 = {
'lap': {
'param': 'eta',
'values': (0.10, 0.20, 0.35, 0.50),
},
'mom': {
'param': 'mom_beta',
'values': (0.40, 0.55, 0.70, 0.85),
},
'fft': {
'param': 'fft_cutoff',
'values': (0.15, 0.25, 0.35, 0.45),
},
'sde': {
'param': 'sde_strength',
'values': (0.03, 0.06, 0.10, 0.15),
},
'spatial': {
'param': 'eta',
'values': (0.10, 0.20, 0.35, 0.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_fn, x, sigmas, extra_args=None, callback=None,
disable=False, _eta=value, _mode=mode):
return _sample_hfx(model_fn, x, sigmas, extra_args, callback,
disable, mode=_mode, eta=_eta)
desc = f"eta={value:.2f}"
else:
kwarg = {param: value}
def sampler(model_fn, x, sigmas, extra_args=None, callback=None,
disable=False, _mode=mode, _kw=kwarg):
return _sample_hfx(model_fn, x, sigmas, extra_args, callback,
disable, mode=_mode, **_kw)
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
# =====================================================================
# Tangent S-curve scheduler (bong_tangent-inspired)
# =====================================================================
#
# sigma(i) from an arctangent S-curve that concentrates steps around a
# pivot point. Higher slope = sharper bend = more step density at pivot.
#
# slope_adj = slope / (steps / 40) [normalization for step count]
def _tangent_sigmas(
steps: int,
sigma_max: float,
sigma_min: float,
slope: float,
pivot_frac: float,
) -> torch.Tensor:
n = steps
if n < 1:
return torch.zeros(1, dtype=torch.float32)
pivot = pivot_frac * (n - 1)
slope_adj = slope / max(n / 40.0, 0.1)
xs = torch.arange(n, dtype=torch.float64)
raw = ((2.0 / math.pi) * torch.atan(-slope_adj * (xs - pivot)) + 1.0) / 2.0
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
sigmas = normalized * (sigma_max - sigma_min) + sigma_min
sigmas = torch.cat([sigmas, torch.zeros(1, dtype=torch.float64)])
return sigmas.float()
def _tangent_schedule(
model_sampling: Any,
steps: int,
slope: float,
pivot_frac: float,
name: str = '',
) -> torch.Tensor:
sigma_max = float(model_sampling.sigma_max)
sigma_min = float(model_sampling.sigma_min)
sigmas = _tangent_sigmas(steps, sigma_max, sigma_min, slope, pivot_frac)
if name:
_plot_sigmas(sigmas, name)
return sigmas
def scheduler_atan_gentle(model_sampling: Any, steps: int) -> torch.Tensor:
return _tangent_schedule(model_sampling, steps, slope=0.7, pivot_frac=0.35,
name='atan_gentle')
def scheduler_atan_focused(model_sampling: Any, steps: int) -> torch.Tensor:
return _tangent_schedule(model_sampling, steps, slope=1.1, pivot_frac=0.40,
name='atan_focused')
def scheduler_atan_steep(model_sampling: Any, steps: int) -> torch.Tensor:
return _tangent_schedule(model_sampling, steps, slope=1.6, pivot_frac=0.45,
name='atan_steep')
# =====================================================================
# Experimental schedulers
# =====================================================================
def _karras_tangent_sigmas(
steps: int,
sigma_max: float,
sigma_min: float,
rho: float = 7.0,
bend: float = 0.35,
pivot_frac: float = 0.40,
) -> torch.Tensor:
"""
Karras-Tangent hybrid schedule.
Base: Karras optimal spacing (rho=7).
Enhancement: warp the time ramp with an arctangent bend to concentrate
more steps in the detail-forming sigma range. bend=0 gives pure Karras,
bend=1 gives pure tangent warp.
"""
n = steps
if n < 1:
return torch.zeros(1, dtype=torch.float32)
t_lin = torch.linspace(0.0, 1.0, n, dtype=torch.float64)
# Tangent warp of the time ramp
pivot = pivot_frac
slope = 1.2 / max(n / 40.0, 0.1)
raw = ((2.0 / math.pi)
* torch.atan(-slope * (t_lin * (n - 1) - pivot * (n - 1)))
+ 1.0) / 2.0
r_max = raw[0].item()
r_min = raw[-1].item()
r_range = r_max - r_min
if r_range < 1e-12:
t_tan = 1.0 - t_lin
else:
t_tan = (raw - r_min) / r_range # [1, 0] normalized
# Blend linear descent [1->0] with tangent warp
t_blend = (1.0 - bend) * (1.0 - t_lin) + bend * t_tan
# Karras formula: sigma = (sig_min^(1/rho) + t*(sig_max^(1/rho)-sig_min^(1/rho)))^rho
inv_rho = 1.0 / rho
lo = sigma_min ** inv_rho
hi = sigma_max ** inv_rho
sigmas = (lo + t_blend * (hi - lo)) ** rho
sigmas = torch.cat([sigmas, torch.zeros(1, dtype=torch.float64)])
return sigmas.float()
def _logistic_sigmas(
steps: int,
sigma_max: float,
sigma_min: float,
steepness: float = 8.0,
midpoint: float = 0.4,
) -> torch.Tensor:
"""
Logistic (sigmoid) S-curve schedule.
Exponential tails (vs algebraic for atan) give a sharper transition
through the detail range with flatter extremes.
"""
n = steps
if n < 1:
return torch.zeros(1, dtype=torch.float32)
t = torch.linspace(0.0, 1.0, n, dtype=torch.float64)
# Sigmoid: 1 / (1 + exp(k*(t - m)))
raw = 1.0 / (1.0 + torch.exp(steepness * (t - midpoint)))
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
sigmas = normalized * (sigma_max - sigma_min) + sigma_min
sigmas = torch.cat([sigmas, torch.zeros(1, dtype=torch.float64)])
return sigmas.float()
def scheduler_karras_tan(model_sampling: Any, steps: int) -> torch.Tensor:
"""Karras-Tangent hybrid schedule (experimental)."""
sigma_max = float(model_sampling.sigma_max)
sigma_min = float(model_sampling.sigma_min)
sigmas = _karras_tangent_sigmas(steps, sigma_max, sigma_min)
_plot_sigmas(sigmas, 'karras_tan')
return sigmas
def scheduler_logistic(model_sampling: Any, steps: int) -> torch.Tensor:
"""Logistic sigmoid S-curve schedule (experimental)."""
sigma_max = float(model_sampling.sigma_max)
sigma_min = float(model_sampling.sigma_min)
sigmas = _logistic_sigmas(steps, sigma_max, sigma_min)
_plot_sigmas(sigmas, 'logistic')
return sigmas
# =====================================================================
# Registration
# =====================================================================
_SAMPLERS: Dict[str, Any] = {}
_SAMPLERS.update(_HFE_PRESETS) # hfe_s1..s8, hfe3_s1..s8, hfe4_s1..s8, hfe5_s1..s8
_SAMPLERS["hfe_auto"] = sample_hfe_auto # 2-stage adaptive
_SAMPLERS["hfe3_auto"] = sample_hfe3_auto # 3-stage adaptive
_SAMPLERS["hfe4_auto"] = sample_hfe4_auto # 4-stage adaptive
_SAMPLERS["hfe5_auto"] = sample_hfe5_auto # 5-stage adaptive
# Experimental base samplers
_SAMPLERS["hfx_lap"] = sample_hfx_lap
_SAMPLERS["hfx_mom"] = sample_hfx_mom
_SAMPLERS["hfx_fft"] = sample_hfx_fft
_SAMPLERS["hfx_sde"] = sample_hfx_sde
_SAMPLERS["hfx_spatial"] = sample_hfx_spatial
# Hybrid combinators
_SAMPLERS["hfx_lap_mom"] = sample_hfx_lap_mom
_SAMPLERS["hfx_lap_spatial"] = sample_hfx_lap_spatial
_SAMPLERS["hfx_fft_spatial"] = sample_hfx_fft_spatial
# Band profile variants
_SAMPLERS["hfx_lap_fine"] = sample_hfx_lap_fine
_SAMPLERS["hfx_lap_broad"] = sample_hfx_lap_broad
# Graduated experimental presets (hfx_*_s1..s4 for each mode)
_SAMPLERS.update(_HFX_PRESETS)
_SCHEDULERS = {
"atan_gentle": scheduler_atan_gentle,
"atan_focused": scheduler_atan_focused,
"atan_steep": scheduler_atan_steep,
"karras_tan": scheduler_karras_tan,
"logistic": scheduler_logistic,
}
# 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",
]
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)
for name, func in _SAMPLERS.items():
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:
attr = f"sample_{name}"
setattr(kdiff, attr, func)
LOGGER.info("HFE samplers registered: %s", list(_SAMPLERS.keys()))
def _register_schedulers() -> None:
handlers = getattr(comfy_samplers, "SCHEDULER_HANDLERS", None)
for name, func in _SCHEDULERS.items():
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
LOGGER.info("HFE schedulers registered: %s", list(_SCHEDULERS.keys()))
# =====================================================================
# 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] = {}