diff --git a/fabric/fabric.py b/fabric/fabric.py index edd73e0..f3ff59d 100644 --- a/fabric/fabric.py +++ b/fabric/fabric.py @@ -2,7 +2,7 @@ import warnings import torch import comfy from nodes import KSamplerAdvanced, CLIPTextEncode -from .unet import q_sample +from .unet import q_sample, get_timesteps def ksampler_fabric(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, clip, pos_weight, neg_weight, feedback_percent, pos_latents=None, neg_latents=None): @@ -68,6 +68,13 @@ def fabric_sample(model, add_noise, noise_seed, steps, cfg, sampler_name, schedu all_latents = all_latents.to(device) print(f"[FABRIC] {num_pos} positive latents, {num_neg} negative latents") + # + # Map steps to timesteps + # + timesteps = get_timesteps(model_patched, steps, sampler_name, scheduler, denoise, device) + feedback_start_ts = timesteps[feedback_start] + feedback_end_ts = timesteps[min(feedback_end, len(timesteps) - 1)] + # # Precompute hidden states # @@ -134,11 +141,16 @@ def fabric_sample(model, add_noise, noise_seed, steps, cfg, sampler_name, schedu nonlocal cond_or_uncond nonlocal num_pos, num_neg nonlocal model_patched + nonlocal feedback_start_ts, feedback_end_ts input = params['input'] ts = params['timestep'] c = params['c'] + # Normal pass if not in feedback range + if not (feedback_end_ts.item() <= ts[0].item() <= feedback_start_ts.item()): + return model_func(input, ts, **c) + # Save cond_or_uncond index for attention patch cond_or_uncond = params['cond_or_uncond'] diff --git a/fabric/unet.py b/fabric/unet.py index 8d57db4..0a90170 100644 --- a/fabric/unet.py +++ b/fabric/unet.py @@ -16,10 +16,6 @@ def q_sample(model, x_start, t): extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise) -# -# UNUSED -# - def get_timesteps(model, steps, sampler, scheduler, denoise, device): real_model = model.model sampler = comfy.samplers.KSampler( @@ -28,6 +24,10 @@ def get_timesteps(model, steps, sampler, scheduler, denoise, device): ) return sampler.model_wrap.sigma_to_discrete_timestep(sampler.sigmas) +# +# UNUSED +# + def forward(model, steps, sampler, scheduler, denoise, device, zs, ts, pos, neg, seed): real_model = model.model diff --git a/nodes.py b/nodes.py index 644192a..bc2c25f 100644 --- a/nodes.py +++ b/nodes.py @@ -12,8 +12,8 @@ class KSamplerFABRIC: "null_neg": ("CONDITIONING",), "pos_weight": ("FLOAT", {"default": 1., "min": 0., "max": 1., "step": 0.01}), "neg_weight": ("FLOAT", {"default": 1., "min": 0., "max": 1., "step": 0.01}), - "feedback_start": ("INT", {"default": 1, "min": 1, "max": 10000, "step": 1}), - "feedback_end": ("INT", {"default": 20, "min": 1, "max": 10000, "step": 1}), + "feedback_start": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1}), + "feedback_end": ("INT", {"default": 20, "min": 0, "max": 10000, "step": 1}), }, "optional": { "pos_latents": ("LATENT",), @@ -44,8 +44,8 @@ class KSamplerAdvFABRIC: "null_neg": ("CONDITIONING",), "pos_weight": ("FLOAT", {"default": 1., "min": 0., "max": 1., "step": 0.01}), "neg_weight": ("FLOAT", {"default": 1., "min": 0., "max": 1., "step": 0.01}), - "feedback_start": ("INT", {"default": 1, "min": 1, "max": 10000, "step": 1}), - "feedback_end": ("INT", {"default": 20, "min": 1, "max": 10000, "step": 1}), + "feedback_start": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1}), + "feedback_end": ("INT", {"default": 10000, "min": 0, "max": 10000, "step": 1}), }, "optional": { "pos_latents": ("LATENT",),