Feat/Fix: clyb_geomextrap sampler (3 NFE) & Apple MPS dtype fix
This commit is contained in:
@@ -12,6 +12,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
# Samplers
|
||||
"SamplerClyb_BDF": clyb_Samplers.SamplerClyb_BDF,
|
||||
"SamplerTaylorFlow": clyb_Samplers.SamplerTaylorFlow,
|
||||
"SamplerClyb_GeomExtrap": clyb_Samplers.SamplerClyb_GeomExtrap,
|
||||
"SamplerWrapperCFGPP": clyb_Samplers.SamplerWrapperCFGPP,
|
||||
# Schedulers
|
||||
"InverseSquaredScheduler": clyb_Schedulers.InverseSquaredScheduler,
|
||||
@@ -27,6 +28,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
# Samplers
|
||||
"SamplerClyb_BDF": "SamplerClyb_BDF",
|
||||
"SamplerTaylorFlow": "SamplerTaylorFlow",
|
||||
"SamplerClyb_GeomExtrap": "SamplerClyb_GeomExtrap",
|
||||
"SamplerWrapperCFGPP": "SamplerWrapperCFGPP",
|
||||
# Schedulers
|
||||
"InverseSquaredScheduler": "InverseSquaredScheduler",
|
||||
|
||||
+10
-5
@@ -203,9 +203,14 @@ class ClybGuidance:
|
||||
|
||||
# 1. Move to Frequency domain using 2D Fast Fourier Transform
|
||||
# We use norm='ortho' to ensure the transform is unitary and preserves energy.
|
||||
fft_cond = torch.fft.fftshift(torch.fft.fftn(cond.to(torch.float64), norm='ortho'))
|
||||
fft_uncond = torch.fft.fftshift(torch.fft.fftn(uncond.to(torch.float64), norm='ortho'))
|
||||
fft_out = torch.fft.fftshift(torch.fft.fftn(out.to(torch.float64), norm='ortho'))
|
||||
device = cond.device
|
||||
is_mps = device.type == "mps" if isinstance(device, torch.device) else "mps" in str(device)
|
||||
precision_dtype = torch.float32 if is_mps else torch.float64
|
||||
complex_dtype = torch.cfloat if is_mps else torch.cdouble
|
||||
|
||||
fft_cond = torch.fft.fftshift(torch.fft.fftn(cond.to(precision_dtype), norm='ortho'))
|
||||
fft_uncond = torch.fft.fftshift(torch.fft.fftn(uncond.to(precision_dtype), norm='ortho'))
|
||||
fft_out = torch.fft.fftshift(torch.fft.fftn(out.to(precision_dtype), norm='ortho'))
|
||||
|
||||
# 1. Create the 2D Hann window kernel for convolution
|
||||
hann_1d = torch.signal.windows.hann(5, device=cond.device)
|
||||
@@ -226,7 +231,7 @@ class ClybGuidance:
|
||||
fft_cond_real = fft_cond_flat.real
|
||||
fft_uncond_real = fft_uncond_flat.real
|
||||
#guidance_direction = (cond - uncond)
|
||||
local_avg_magnitude = F.conv1d(fft_cond_real, kernel.to(torch.float64), padding='same')
|
||||
local_avg_magnitude = F.conv1d(fft_cond_real, kernel.to(precision_dtype), padding='same')
|
||||
|
||||
# 3. Normalize the magnitude map for each image in the batch to the [0, 1] range
|
||||
# This makes the `strength` parameter behave consistently across different images.
|
||||
@@ -250,7 +255,7 @@ class ClybGuidance:
|
||||
|
||||
print(local_scale)
|
||||
|
||||
guided_tensor = fft_cond + (local_scale.to(torch.cdouble) * (fft_cond_flat - fft_uncond_flat)).view(fft_cond.shape)
|
||||
guided_tensor = fft_cond + (local_scale.to(complex_dtype) * (fft_cond_flat - fft_uncond_flat)).view(fft_cond.shape)
|
||||
|
||||
guided_tensor = torch.fft.ifftshift(guided_tensor)
|
||||
guided_tensor = torch.fft.ifftn(guided_tensor, norm='ortho').real
|
||||
|
||||
+307
-3
@@ -333,8 +333,9 @@ def sampler_taylor_flow(
|
||||
denoised_cur = model(x, sigma_cur * s_in, **extra_args)
|
||||
|
||||
# 2. Build Vandermonde matrix from historical timesteps
|
||||
precision_dtype = torch.float32 if (device.type == "mps" if isinstance(device, torch.device) else "mps" in str(device)) else torch.float64
|
||||
R_p = _construct_vandermonde_flow(
|
||||
history, sigma_cur, order, device, torch.float64
|
||||
history, sigma_cur, order, device, precision_dtype
|
||||
)
|
||||
|
||||
# 3. Solve for B coefficients
|
||||
@@ -491,6 +492,308 @@ def sample_cfgpp(model, x, sigmas, extra_args=None, callback=None, disable=None,
|
||||
)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# GEOM_EXTRAP SAMPLER - 3-NFE per step sampler that uses two geometric
|
||||
# midpoints between sigma_cur and sigma_down to approximate a higher-order
|
||||
# denoised prediction, then integrates that prediction into the noisy x
|
||||
# latent. Uses the same selectable sigma_calc branches as SamplerTaylorFlow.
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def _compute_ancestral_sigmas(sigma_cur, sigma_next, eta, sigma_calc, history, order):
|
||||
"""
|
||||
Compute (sigma_down, sigma_up) for the requested sigma_calc method.
|
||||
Replicates the four branches from sampler_taylor_flow (clyb_Samplers.py:268-314)
|
||||
verbatim so behaviour matches that sampler.
|
||||
"""
|
||||
h_n = sigma_next - sigma_cur
|
||||
if sigma_calc == "clyb":
|
||||
sigma_down = (
|
||||
sigma_next ** 2 / (1.0 + math.log(1.0 + abs(sigma_next - sigma_cur)) * eta)
|
||||
) ** 0.5
|
||||
sigma_up = (sigma_next ** 2 - sigma_down ** 2) ** 0.5
|
||||
elif sigma_calc == "taylor-expansion":
|
||||
step_ratio = abs(h_n) / max(sigma_next, 1e-8)
|
||||
taylor_factor = math.exp(-eta * step_ratio)
|
||||
quadratic_correction = 1.0 - eta * 0.5 * step_ratio ** 2
|
||||
sigma_down = sigma_next * taylor_factor * quadratic_correction
|
||||
sigma_up = sigma_next * max(0.0, 1.0 - taylor_factor ** 2) ** 0.5
|
||||
elif sigma_calc == "ancestral":
|
||||
sigma_down, sigma_up = get_ancestral_step(sigma_cur, sigma_next, eta)
|
||||
elif sigma_calc == "adaptive":
|
||||
window_size = min(order, len(history))
|
||||
if window_size >= 2:
|
||||
history_list = list(history)
|
||||
recent = history_list[-window_size:]
|
||||
denoised_list = [d.float() for _, d in recent]
|
||||
stacked = torch.stack(denoised_list)
|
||||
mean_d = stacked.mean(dim=0)
|
||||
var_val = ((stacked - mean_d) ** 2).mean().item()
|
||||
norm_val = mean_d.pow(2).mean().item()
|
||||
eps = 1e-8
|
||||
if math.isfinite(var_val) and math.isfinite(norm_val):
|
||||
normalized_metric = var_val / (var_val + abs(norm_val) + eps)
|
||||
normalized_metric = min(1.0, max(0.0, normalized_metric))
|
||||
else:
|
||||
normalized_metric = 0.0
|
||||
sigma_down = sigma_next * (1.0 - eta * normalized_metric)
|
||||
sigma_down = max(0.0, min(sigma_next, sigma_down))
|
||||
sigma_up = math.sqrt(max(0.0, sigma_next ** 2 - sigma_down ** 2))
|
||||
else:
|
||||
sigma_down = (
|
||||
sigma_next ** 2
|
||||
/ (1.0 + math.log(1.0 + abs(sigma_next - sigma_cur)) * eta)
|
||||
) ** 0.5
|
||||
sigma_up = (sigma_next ** 2 - sigma_down ** 2) ** 0.5
|
||||
else:
|
||||
sigma_down = (
|
||||
sigma_next ** 2 / (1.0 + math.log(1.0 + abs(sigma_next - sigma_cur)) * eta)
|
||||
) ** 0.5
|
||||
sigma_up = (sigma_next ** 2 - sigma_down ** 2) ** 0.5
|
||||
return sigma_down, sigma_up
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sampler_geom_extrap(
|
||||
model,
|
||||
x,
|
||||
sigmas,
|
||||
extra_args=None,
|
||||
callback=None,
|
||||
disable=None,
|
||||
eta=1.0,
|
||||
s_noise=1.0,
|
||||
noise_sampler=None,
|
||||
flow=False,
|
||||
sigma_calc="clyb",
|
||||
):
|
||||
"""
|
||||
Geometric-Midpoint Extrapolation sampler.
|
||||
|
||||
Per-step procedure (3 NFEs per step):
|
||||
1. Sample at sigma_cur.
|
||||
2. Integrate into x_mid1 at sigma_gm1 = sqrt(sigma_cur * sigma_down).
|
||||
3. Sample at sigma_gm1.
|
||||
4. Linearly extrapolate through (denoised1, denoised2) to predict at sigma_down.
|
||||
5. Integrate denoised_pred into x_mid2 at sigma_gm2 = sqrt(sigma_gm1 * sigma_down).
|
||||
6. Sample at sigma_gm2.
|
||||
7. If not the final step, do a quadratic (3-point divided-difference) extrapolation
|
||||
through (denoised1, denoised2, denoised3) to predict at sigma_down and integrate
|
||||
that into x. If the final step, integrate denoised3 directly into x.
|
||||
|
||||
All intermediate sigmas (sigma_gm1, sigma_gm2) and lerp weight denominators are
|
||||
clamped to >= 1e-4 for numerical stability near sigma = 0.
|
||||
"""
|
||||
extra_args = {} if extra_args is None else extra_args
|
||||
seed = extra_args.get("seed", None)
|
||||
noise_sampler = (
|
||||
default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler
|
||||
)
|
||||
s_in = x.new_ones([x.shape[0]])
|
||||
|
||||
if len(sigmas) <= 1:
|
||||
return x
|
||||
|
||||
# History buffer is needed for the "adaptive" sigma_calc branch.
|
||||
history = collections.deque(maxlen=16)
|
||||
# The last iteration index is len(sigmas) - 2 (since the loop goes 0..len(sigmas)-2,
|
||||
# and sigmas[-1] is 0). On that step sigma_next == 0 and sigma_down == 0, so the
|
||||
# 3-point extrapolation degenerates (denominators collapse). We skip extrapolation
|
||||
# there and use denoised3 directly.
|
||||
is_final_step = len(sigmas) - 2 if len(sigmas) >= 2 else 0
|
||||
|
||||
for i in trange(len(sigmas) - 1, disable=disable):
|
||||
sigma_cur, sigma_next = sigmas[i], sigmas[i + 1]
|
||||
h_n = sigma_next - sigma_cur
|
||||
|
||||
# ---- Ancestral sigma_down / sigma_up via the same branches as taylor_flow ----
|
||||
sigma_down, sigma_up = _compute_ancestral_sigmas(
|
||||
sigma_cur, sigma_next, eta, sigma_calc, history, order=16
|
||||
)
|
||||
|
||||
# ---- Flow model coefficients (identical to sampler_taylor_flow lines 316-327) ----
|
||||
alpha_ip1 = None
|
||||
alpha_down = None
|
||||
renoise_coeff = None
|
||||
alpha_ratio = 1.0
|
||||
if flow:
|
||||
alpha_ip1 = 1.0 - sigma_next
|
||||
alpha_down = 1.0 - sigma_down
|
||||
renoise_coeff = (
|
||||
sigma_next ** 2 - sigma_down ** 2 * alpha_ip1 ** 2 / alpha_down ** 2
|
||||
) ** 0.5
|
||||
alpha_ratio = alpha_ip1 / alpha_down if alpha_down != 0 else 1.0
|
||||
|
||||
# ---- Geometric midpoints, clamped for numerical stability ----
|
||||
sigma_gm1 = (sigma_cur * sigma_down).clamp_min(1e-4).sqrt()
|
||||
sigma_gm2 = (sigma_gm1 * sigma_down).clamp_min(1e-4).sqrt()
|
||||
|
||||
# ---- NFE 1: sample at sigma_cur ----
|
||||
denoised1 = model(x, sigma_cur * s_in, **extra_args)
|
||||
|
||||
# ---- Integrate denoised1 into a noisy latent at sigma_gm1 ----
|
||||
w_gm1 = (sigma_gm1 / sigma_cur).clamp_min(1e-4)
|
||||
x_mid1 = denoised1.lerp(x, weight=w_gm1)
|
||||
|
||||
# ---- NFE 2: sample at sigma_gm1 ----
|
||||
denoised2 = model(x_mid1, sigma_gm1 * s_in, **extra_args)
|
||||
|
||||
# ---- Linear extrapolation through (denoised1, denoised2) to predict at sigma_down ----
|
||||
slope_12 = (denoised2 - denoised1) / (sigma_gm1 - sigma_cur)
|
||||
denoised_pred = denoised2 + slope_12 * (sigma_down - sigma_gm1)
|
||||
|
||||
# ---- Integrate denoised_pred into a noisy latent at sigma_gm2 (from original x) ----
|
||||
w_gm2 = (sigma_gm2 / sigma_gm1).clamp_min(1e-4)
|
||||
x_mid2 = denoised_pred.lerp(x_mid1, weight=w_gm2)
|
||||
|
||||
# ---- NFE 3: sample at sigma_gm2 ----
|
||||
denoised3 = model(x_mid2, sigma_gm2 * s_in, **extra_args)
|
||||
|
||||
# ---- Final vs non-final step ----
|
||||
if i < is_final_step:
|
||||
# Quadratic (3-point) extrapolation via Newton divided differences
|
||||
d1 = (denoised2 - denoised1) / (sigma_gm1 - sigma_cur)
|
||||
d2 = (denoised3 - denoised2) / (sigma_gm2 - sigma_gm1)
|
||||
d2_div = (d2 - d1) / (sigma_gm2 - sigma_cur)
|
||||
denoised_final = (
|
||||
denoised3
|
||||
+ d2 * (sigma_down - sigma_gm2)
|
||||
+ d2_div * (sigma_down - sigma_gm2) * (sigma_down - sigma_gm1)
|
||||
)
|
||||
else:
|
||||
# Final step: sigma_down = 0, just use denoised3 directly.
|
||||
denoised_final = denoised3
|
||||
|
||||
# ---- Integrate the final denoised prediction into x at sigma_down ----
|
||||
w_down = (sigma_down / sigma_gm2).clamp_min(1e-4)
|
||||
x = denoised_final.lerp(x_mid2, weight=w_down)
|
||||
|
||||
# ---- Update history for "adaptive" sigma_calc branch on subsequent steps ----
|
||||
history.append((sigma_cur, denoised_final))
|
||||
|
||||
# ---- Ancestral noise injection (identical to sampler_taylor_flow lines 369-374) ----
|
||||
if sigma_next > 0 and eta > 0:
|
||||
noise = noise_sampler(sigma_cur, sigma_next) * s_noise
|
||||
if flow:
|
||||
x = alpha_ratio * x + noise * renoise_coeff
|
||||
else:
|
||||
x = x + noise * sigma_up
|
||||
|
||||
# ---- Callback ----
|
||||
if callback is not None:
|
||||
callback(
|
||||
{
|
||||
"x": x,
|
||||
"i": i,
|
||||
"sigma": sigma_cur,
|
||||
"sigma_hat": sigma_cur,
|
||||
"denoised": denoised_final,
|
||||
}
|
||||
)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def sample_geom_extrap(
|
||||
model,
|
||||
x,
|
||||
sigmas,
|
||||
extra_args=None,
|
||||
callback=None,
|
||||
disable=None,
|
||||
eta=1.0,
|
||||
s_noise=1.0,
|
||||
noise_sampler=None,
|
||||
sigma_calc="clyb",
|
||||
):
|
||||
"""Wrapper that detects flow vs non-flow then calls sampler_geom_extrap."""
|
||||
flow = False
|
||||
if isinstance(
|
||||
model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST
|
||||
):
|
||||
flow = True
|
||||
return sampler_geom_extrap(
|
||||
model,
|
||||
x,
|
||||
sigmas,
|
||||
extra_args=extra_args,
|
||||
callback=callback,
|
||||
disable=disable,
|
||||
eta=eta,
|
||||
s_noise=s_noise,
|
||||
noise_sampler=noise_sampler,
|
||||
flow=flow,
|
||||
sigma_calc=sigma_calc,
|
||||
)
|
||||
|
||||
|
||||
class SamplerClyb_GeomExtrap:
|
||||
"""
|
||||
Geometric-Midpoint Extrapolation sampler.
|
||||
|
||||
3-NFE per step sampler that uses two geometric midpoints between sigma_cur and
|
||||
sigma_down to build a 3-point quadratic extrapolation of the denoised prediction
|
||||
at sigma_down, then integrates that prediction into the noisy x latent.
|
||||
|
||||
Compatible with both flow-matching (FLUX, SD3, Chroma) and non-flow models.
|
||||
|
||||
Parameters:
|
||||
- eta: Ancestral sampling stochasticity (0=deterministic, 1=full stochastic)
|
||||
- s_noise: Noise scaling factor
|
||||
- sigma_calc: Ancestral sigma calculation method (clyb, taylor-expansion, ancestral, adaptive)
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"eta": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 100.0,
|
||||
"step": 0.01,
|
||||
"tooltip": "Ancestral sampling stochasticity (0=deterministic, 1=full stochastic)",
|
||||
},
|
||||
),
|
||||
"s_noise": (
|
||||
"FLOAT",
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 100.0,
|
||||
"step": 0.01,
|
||||
"tooltip": "Noise scaling factor",
|
||||
},
|
||||
),
|
||||
"sigma_calc": (
|
||||
["clyb", "taylor-expansion", "ancestral", "adaptive"],
|
||||
{
|
||||
"default": "clyb",
|
||||
"tooltip": "Ancestral sigma calculation method: clyb (original log-based), taylor-expansion (exponential+quadratic), ancestral (standard k-diffusion), adaptive (history-based convergence-aware)",
|
||||
},
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SAMPLER",)
|
||||
CATEGORY = "sampling/custom_sampling/samplers"
|
||||
FUNCTION = "get_sampler"
|
||||
|
||||
def get_sampler(self, eta, s_noise, sigma_calc):
|
||||
sampler = comfy.samplers.ksampler(
|
||||
"geom_extrap",
|
||||
{
|
||||
"eta": eta,
|
||||
"s_noise": s_noise,
|
||||
"sigma_calc": sigma_calc,
|
||||
},
|
||||
)
|
||||
return (sampler,)
|
||||
|
||||
|
||||
# The following function adds the samplers during initialization, in __init__.py
|
||||
def add_samplers():
|
||||
from comfy.samplers import KSampler, k_diffusion_sampling
|
||||
@@ -527,6 +830,7 @@ def add_samplers():
|
||||
extra_samplers = {
|
||||
"clyb_bdf": sample_clyb_bdf,
|
||||
"taylor_flow": sample_taylor_flow,
|
||||
"geom_extrap": sample_geom_extrap,
|
||||
}
|
||||
|
||||
# Wrappers are NOT in the standard sampler dropdown. They are only reachable
|
||||
@@ -604,7 +908,7 @@ class SamplerTaylorFlow:
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 1.0,
|
||||
"max": 100.0,
|
||||
"step": 0.01,
|
||||
"tooltip": "Ancestral sampling stochasticity (0=deterministic, 1=full stochastic)",
|
||||
},
|
||||
@@ -614,7 +918,7 @@ class SamplerTaylorFlow:
|
||||
{
|
||||
"default": 1.0,
|
||||
"min": 0.0,
|
||||
"max": 2.0,
|
||||
"max": 100.0,
|
||||
"step": 0.01,
|
||||
"tooltip": "Noise scaling factor",
|
||||
},
|
||||
|
||||
Reference in New Issue
Block a user