Expand scheduler and adjust curves with a bias on modern models.

This commit is contained in:
Jordan Thompson
2026-04-29 13:52:56 -07:00
parent 628ae6aad5
commit f2bf267955
2 changed files with 269 additions and 145 deletions
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "RES4SHO"
version = "1.1.1"
version = "1.1.4"
description = "High-Frequency Detail Sampling based on Res Sampling for ComfyUI"
readme = "README.md"
requires-python = ">=3.10"
+268 -144
View File
@@ -166,7 +166,7 @@ def _extract_hf(t: torch.Tensor, kernel_size: int = 3) -> torch.Tensor:
# =====================================================================
def _plot_sigmas(sigmas: torch.Tensor, name: str,
width: int = 64, height: int = 16) -> None:
width: int = 120, height: int = 32) -> None:
"""Render a sigma schedule as an ASCII chart in the console."""
vals = sigmas.tolist()
if vals and vals[-1] == 0.0:
@@ -1705,31 +1705,117 @@ for _mode in _HFX_SWEEPS:
# =====================================================================
# Tangent S-curve scheduler (bong_tangent-inspired)
# Two-stage S-curve schedulers (bong_tangent architecture)
# =====================================================================
#
# sigma(i) from an arctangent S-curve that concentrates steps around a
# pivot point. Higher slope = sharper bend = more step density at pivot.
# Two-stage design: stage 1 (σ_max → σ_mid) for structure,
# stage 2 (σ_mid → σ_min) for detail. Each stage applies a curve
# function with independent slope/pivot. Direct sigma mapping -- no
# Karras power-law.
#
# slope_adj = slope / (steps / 40) [normalization for step count]
# slope_adj = slope / (steps / 40) [step-count normalization]
def _tangent_sigmas(
steps: int,
sigma_max: float,
sigma_min: float,
# --- Curve functions ---------------------------------------------------
#
# Each takes (xs, pivot, slope) and returns raw monotone-decreasing values.
# Callers normalize the output to [0, 1].
def _curve_atan(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
"""Arctangent S-curve (same basis as bong_tangent)."""
return ((2.0 / math.pi) * torch.atan(-slope * (xs - pivot)) + 1.0) / 2.0
def _curve_logistic(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
"""Logistic sigmoid S-curve (exponential tails, sharper than atan)."""
return 1.0 / (1.0 + torch.exp(slope * (xs - pivot)))
def _curve_cosine(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
"""Cosine S-curve (smoothest, no sharp inflection).
*slope* controls sharpness via a power warp on the normalized position:
slope < 1 compresses the curve toward the middle (sharper transition),
slope > 1 spreads it toward the extremes (gentler).
"""
n = xs.shape[0]
if n <= 1:
return torch.ones_like(xs)
t = (xs - xs[0]) / (xs[-1] - xs[0]) # [0, 1]
# Power warp: slope acts as exponent (higher = gentler)
t_warped = t.clamp(0.0, 1.0).pow(max(slope, 0.01))
return (1.0 + torch.cos(math.pi * t_warped)) / 2.0
def _curve_beta(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
"""Beta / power-law S-curve (inherently asymmetric).
Uses the regularized incomplete beta function CDF. *slope* controls
the concentration exponent (higher = stronger bend). The *pivot*
sets the balance between head and tail weighting: pivot < n/2 biases
toward early steps, pivot > n/2 biases toward late steps.
"""
n = xs.shape[0]
if n <= 1:
return torch.ones_like(xs)
t = (xs - xs[0]) / (xs[-1] - xs[0]) # [0, 1]
# Derive alpha/beta from slope and pivot position
pivot_norm = max(min(pivot / max(n - 1, 1), 0.95), 0.05)
concentration = max(slope * 5.0, 0.1) # scale slope to useful range
alpha = concentration * (1.0 - pivot_norm)
beta_p = concentration * pivot_norm
# Beta CDF via element-wise power approximation (Kumaraswamy)
a = max(alpha, 0.01)
b = max(beta_p, 0.01)
cdf = 1.0 - (1.0 - t.clamp(1e-7, 1.0 - 1e-7).pow(a)).pow(b)
return 1.0 - cdf # decreasing
def _curve_laplacian(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
"""Laplacian (double-exponential) S-curve.
Sharper peak than logistic — creates very tight concentration at the
pivot with rapid exponential falloff on both sides.
"""
diff = slope * (xs - pivot)
# Laplace CDF: 0.5 * exp(x) for x<0, 1 - 0.5*exp(-x) for x>=0
# We want a decreasing function, so we negate the argument.
cdf = torch.where(
diff <= 0,
0.5 * torch.exp(diff),
1.0 - 0.5 * torch.exp(-diff),
)
return 1.0 - cdf # decreasing
def _curve_linear(xs: torch.Tensor, pivot: float, slope: float) -> torch.Tensor:
"""Pure linear descent (no S-curve). Useful as a baseline reference."""
n = xs.shape[0]
if n <= 1:
return torch.ones_like(xs)
return torch.linspace(1.0, 0.0, n, dtype=xs.dtype, device=xs.device)
# --- Core two-stage engine -------------------------------------------
def _stage_sigmas(
curve_fn,
n: int,
slope: float,
pivot_frac: float,
rho: float = 7.0,
pivot: float,
start: float,
end: float,
) -> torch.Tensor:
n = steps
if n < 1:
return torch.zeros(1, dtype=torch.float32)
"""Apply *curve_fn* over *n* steps, mapping [start, end] sigma range.
pivot = pivot_frac * (n - 1)
slope_adj = slope / max(n / 40.0, 0.1)
Faithfully reproduces bong_tangent's get_bong_tangent_sigmas() logic
for any pluggable curve function.
"""
if n < 1:
return torch.tensor([], dtype=torch.float64)
xs = torch.arange(n, dtype=torch.float64)
raw = ((2.0 / math.pi) * torch.atan(-slope_adj * (xs - pivot)) + 1.0) / 2.0
raw = curve_fn(xs, pivot, slope)
r_max = raw[0].item()
r_min = raw[-1].item()
@@ -1740,157 +1826,185 @@ def _tangent_sigmas(
else:
normalized = (raw - r_min) / r_range
# Karras power-law spacing: interpolate in sigma^(1/rho) space
inv_rho = 1.0 / rho
lo = sigma_min ** inv_rho
hi = sigma_max ** inv_rho
sigmas = (lo + normalized * (hi - lo)) ** rho
return end + normalized * (start - end)
sigmas = torch.cat([sigmas, torch.zeros(1, dtype=torch.float64)])
def _two_stage_sigmas(
curve_fn,
steps: int,
sigma_max: float,
sigma_min: float,
slope_1: float = 0.2,
slope_2: float = 0.2,
pivot_frac_1: float = 0.6,
pivot_frac_2: float = 0.6,
mid_frac: float = 0.5,
) -> torch.Tensor:
"""Two-stage S-curve schedule (bong_tangent architecture).
Parameters
----------
curve_fn : callable
One of ``_curve_atan``, ``_curve_logistic``, ``_curve_cosine``.
steps : int
Total sampling steps requested.
sigma_max, sigma_min : float
Sigma bounds from ``model_sampling``.
slope_1, slope_2 : float
Curve concentration per stage (before step-count normalization).
pivot_frac_1, pivot_frac_2 : float
Pivot position within the *total* step range (0‥1), matching
bong_tangent's convention where the pivot is relative to total
steps, not per-stage.
mid_frac : float
Where to split the sigma range: ``sigma_mid = sigma_min +
mid_frac * (sigma_max - sigma_min)``.
"""
# Match bong_tangent: pad by 2 then trim junction
n = steps + 2
sigma_mid = sigma_min + mid_frac * (sigma_max - sigma_min)
# Split steps the same way bong_tangent does
midpoint = int((n * pivot_frac_1 + n * pivot_frac_2) / 2)
stage_1_len = midpoint
stage_2_len = n - midpoint
# Absolute pivot indices (bong_tangent convention)
piv_1 = int(n * pivot_frac_1)
piv_2 = int(n * pivot_frac_2) - stage_1_len # relative to stage 2
# Step-count-normalized slopes
s1 = slope_1 / max(n / 40.0, 0.1)
s2 = slope_2 / max(n / 40.0, 0.1)
sigmas_1 = _stage_sigmas(curve_fn, stage_1_len, s1, piv_1,
sigma_max, sigma_mid)
sigmas_2 = _stage_sigmas(curve_fn, stage_2_len, s2, piv_2,
sigma_mid, sigma_min)
# Drop last of stage 1 (duplicate at junction)
if len(sigmas_1) > 0:
sigmas_1 = sigmas_1[:-1]
sigmas = torch.cat([sigmas_1, sigmas_2,
torch.zeros(1, dtype=torch.float64)])
return sigmas.float()
# --- Schedule wrapper -------------------------------------------------
def _tangent_schedule(
model_sampling: Any,
steps: int,
slope: float,
pivot_frac: float,
curve_fn,
slope_1: float,
slope_2: float,
pivot_frac_1: float = 0.6,
pivot_frac_2: float = 0.6,
mid_frac: float = 0.5,
name: str = '',
) -> torch.Tensor:
sigma_max = float(model_sampling.sigma_max)
sigma_min = float(model_sampling.sigma_min)
sigmas = _tangent_sigmas(steps, sigma_max, sigma_min, slope, pivot_frac)
sigmas = _two_stage_sigmas(
curve_fn, steps, sigma_max, sigma_min,
slope_1=slope_1, slope_2=slope_2,
pivot_frac_1=pivot_frac_1, pivot_frac_2=pivot_frac_2,
mid_frac=mid_frac,
)
if name:
_plot_sigmas(sigmas, name)
return sigmas
# --- Presets -----------------------------------------------------------
def scheduler_atan_gentle(model_sampling: Any, steps: int) -> torch.Tensor:
return _tangent_schedule(model_sampling, steps, slope=0.7, pivot_frac=0.35,
"""Mild arctangent concentration (closest to bong_tangent defaults)."""
return _tangent_schedule(model_sampling, steps, _curve_atan,
slope_1=0.15, slope_2=0.15,
name='atan_gentle')
def scheduler_atan_focused(model_sampling: Any, steps: int) -> torch.Tensor:
return _tangent_schedule(model_sampling, steps, slope=1.1, pivot_frac=0.40,
"""Moderate arctangent concentration."""
return _tangent_schedule(model_sampling, steps, _curve_atan,
slope_1=0.25, slope_2=0.25,
name='atan_focused')
def scheduler_atan_steep(model_sampling: Any, steps: int) -> torch.Tensor:
return _tangent_schedule(model_sampling, steps, slope=1.6, pivot_frac=0.45,
"""Aggressive arctangent concentration."""
return _tangent_schedule(model_sampling, steps, _curve_atan,
slope_1=0.40, slope_2=0.40,
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,
rho: float = 7.0,
) -> 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
# Karras power-law spacing: interpolate in sigma^(1/rho) space
inv_rho = 1.0 / rho
lo = sigma_min ** inv_rho
hi = sigma_max ** inv_rho
sigmas = (lo + normalized * (hi - lo)) ** rho
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
"""Logistic S-curve (sharper transitions than arctangent)."""
return _tangent_schedule(model_sampling, steps, _curve_logistic,
slope_1=0.20, slope_2=0.20,
name='logistic')
def scheduler_cosine(model_sampling: Any, steps: int) -> torch.Tensor:
"""Cosine S-curve (smoothest, most gradual transitions)."""
return _tangent_schedule(model_sampling, steps, _curve_cosine,
slope_1=1.0, slope_2=1.0,
name='cosine')
def scheduler_beta(model_sampling: Any, steps: int) -> torch.Tensor:
"""Beta power-law S-curve (inherently asymmetric concentration)."""
return _tangent_schedule(model_sampling, steps, _curve_beta,
slope_1=0.20, slope_2=0.20,
name='beta')
def scheduler_laplacian(model_sampling: Any, steps: int) -> torch.Tensor:
"""Laplacian S-curve (tightest pivot concentration, sharp falloff)."""
return _tangent_schedule(model_sampling, steps, _curve_laplacian,
slope_1=0.20, slope_2=0.20,
name='laplacian')
def scheduler_linear(model_sampling: Any, steps: int) -> torch.Tensor:
"""Linear two-stage (no S-curve, baseline reference)."""
return _tangent_schedule(model_sampling, steps, _curve_linear,
slope_1=0.20, slope_2=0.20,
name='linear')
# --- Asymmetric presets ------------------------------------------------
# Same curves but with different slope_1 vs slope_2 to bias toward
# structure (high-sigma lingering) or detail (low-sigma lingering).
def scheduler_atan_structure(model_sampling: Any, steps: int) -> torch.Tensor:
"""Arctangent biased toward structure: gentle stage 1, steep stage 2."""
return _tangent_schedule(model_sampling, steps, _curve_atan,
slope_1=0.10, slope_2=0.35,
name='atan_structure')
def scheduler_atan_detail(model_sampling: Any, steps: int) -> torch.Tensor:
"""Arctangent biased toward detail: steep stage 1, gentle stage 2."""
return _tangent_schedule(model_sampling, steps, _curve_atan,
slope_1=0.35, slope_2=0.10,
name='atan_detail')
def scheduler_logistic_structure(model_sampling: Any, steps: int) -> torch.Tensor:
"""Logistic biased toward structure: gentle stage 1, steep stage 2."""
return _tangent_schedule(model_sampling, steps, _curve_logistic,
slope_1=0.10, slope_2=0.30,
name='logistic_structure')
def scheduler_logistic_detail(model_sampling: Any, steps: int) -> torch.Tensor:
"""Logistic biased toward detail: steep stage 1, gentle stage 2."""
return _tangent_schedule(model_sampling, steps, _curve_logistic,
slope_1=0.30, slope_2=0.10,
name='logistic_detail')
# =====================================================================
@@ -1918,11 +2032,21 @@ _SAMPLERS["hfx_coherence"] = sample_hfx_coherence
_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,
# Symmetric atan (bong_tangent-derived)
"atan_gentle": scheduler_atan_gentle,
"atan_focused": scheduler_atan_focused,
"atan_steep": scheduler_atan_steep,
# Alternative curves
"logistic": scheduler_logistic,
"cosine": scheduler_cosine,
"beta": scheduler_beta,
"laplacian": scheduler_laplacian,
"linear": scheduler_linear,
# Asymmetric presets
"atan_structure": scheduler_atan_structure,
"atan_detail": scheduler_atan_detail,
"logistic_structure": scheduler_logistic_structure,
"logistic_detail": scheduler_logistic_detail,
}
# Old names from all previous versions