From 2f855bbeeadecaad84aeb6d70d6d561e7baa2b06 Mon Sep 17 00:00:00 2001 From: Clybius Date: Sun, 19 Jul 2026 09:04:19 -0500 Subject: [PATCH] Feat/Fix: clyb_geomextrap sampler (3 NFE) & Apple MPS dtype fix --- __init__.py | 2 + clyb_Guidance.py | 15 ++- clyb_Samplers.py | 310 ++++++++++++++++++++++++++++++++++++++++++++++- 3 files changed, 319 insertions(+), 8 deletions(-) diff --git a/__init__.py b/__init__.py index 99476a8..70c5884 100644 --- a/__init__.py +++ b/__init__.py @@ -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", diff --git a/clyb_Guidance.py b/clyb_Guidance.py index dd11cc9..394db35 100644 --- a/clyb_Guidance.py +++ b/clyb_Guidance.py @@ -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 diff --git a/clyb_Samplers.py b/clyb_Samplers.py index 5c39ff4..7074009 100644 --- a/clyb_Samplers.py +++ b/clyb_Samplers.py @@ -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", },