""" 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 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 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] = {}