Author SHA1 Message Date
asagi4 e22008619b Linear adaptive guidance experiment
I think the implementation mostly matches the paper, but without proper betas
for the linear estimator, this doesn't actually do anything useful, and my
math-fu isn't good enough to calculate them.
2024-04-17 20:34:12 +03:00
5 changed files with 144 additions and 1441 deletions
-25
View File
@@ -1,25 +0,0 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- master
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'asagi4' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
-10
View File
@@ -4,8 +4,6 @@ An implementation of adaptive guidance for ComfyUI
See https://bcv-uniandes.github.io/adaptiveguidance-wp/
Import [this workflow](example_workflows/AGExample.json?raw=1) into ComfyUI to compare Adaptive Guidance vs. normal CFG.
## What
There's an `AdaptiveGuidance` node (under `sampling/custom_sampling/guiders`) that can be used with `SamplerCustomAdvanced`. Normally, you should keep the threshold quite high, between `0.99` and `1.0`
@@ -13,11 +11,3 @@ There's an `AdaptiveGuidance` node (under `sampling/custom_sampling/guiders`) th
The node calculates the cosine similarity between the u-net's conditional and unconditional output ("positive" and "negative" prompts) and once the similarity crosses the specified threshold, it sets CFG to 1.0, effectively skipping negative prompt calculations and speeding up inference.
I'm not sure if the cosine similarity calculation matches the original paper since I had to translate from maths to Python, but it appears to work.
### Uncond zero
Set uncond_zero_scale to > 0 to enable "uncond zero" CFG *after* the normal CFG gets disabled. Stolen from https://github.com/Extraltodeus/Uncond-Zero-for-ComfyUI
It seems to work slightly better than just running without CFG, but YMMV
Note: this functionality is unstable and will probably change, so using it means your workflows likely won't be perfectly reproducible.
+144 -263
View File
@@ -1,69 +1,170 @@
import comfy.samplers
import comfy_extras.nodes_perpneg
import torch
cos = torch.nn.CosineSimilarity(dim=1)
# shared structure for adaptive guiders
class AdaptiveGuider(object):
cfg_start_timestep = 1000.0
class AdaptiveGuider(comfy.samplers.CFGGuider):
threshold_timestep = 0
uz_scale = 0.0
def set_threshold(self, threshold, start_at):
self.cfg_start_timestep = start_at
def set_threshold(self, threshold):
self.threshold = threshold
def set_uncond_zero_scale(self, scale):
self.uz_scale = scale
def zero_cond(self, args):
cond = args["cond_denoised"]
x = args["input"]
x -= x.mean()
cond -= cond.mean()
return x - (cond / cond.std() ** 0.5) * self.uz_scale
def check_similarity(self, ts, cond_pred, uncond_pred):
if not self.threshold >= 1.0:
sim = cos(cond_pred.reshape(1, -1), uncond_pred.reshape(1, -1)).item()
if sim >= self.threshold:
print(f"AdaptiveGuider: Cosine similarity {sim:.4f} exceeds threshold, setting CFG to 1.0")
self.threshold_timestep = ts
def check_cos_sim(self, ts, cond_pred, uncond_pred):
# Is this reshape correct? It at least gives a scalar value...
sim = cos(cond_pred.reshape(1, -1), uncond_pred.reshape(1, -1)).item()
sim = round(sim, 4)
if sim > self.threshold:
print("AdaptiveGuidance: Cosine similarity", sim, "exceeds threshold, setting CFG to 1.0")
self.threshold_timestep = ts
def predict_noise(self, x, timestep, model_options={}, seed=None):
cond = self.conds.get("positive")
uncond = self.conds.get("negative")
ts = timestep[0].item()
if ts > self.cfg_start_timestep or self.threshold_timestep > ts or self.cfg == 1.0:
if self.uz_scale > 0.0:
model_options = model_options.copy()
model_options["sampler_cfg_function"] = self.zero_cond
cond = self.conds.get("positive")
uncond = self.conds.get("negative")
if self.threshold_timestep > ts:
return comfy.samplers.sampling_function(
self.inner_model, x, timestep, uncond, cond, 1.0, model_options=model_options, seed=seed
)
self.threshold_timestep = 0
conds = self.calc_conds(x, timestep, model_options)
self.check_similarity(ts, conds[0], conds[1])
return self.calc_cfg(conds, x, timestep, model_options)
else:
self.threshold_timestep = 0
uncond_pred, cond_pred = comfy.samplers.calc_cond_batch(
self.inner_model, [uncond, cond], x, timestep, model_options
)
self.check_cos_sim()
return comfy.samplers.cfg_function(
self.inner_model,
cond_pred,
uncond_pred,
self.cfg,
x,
timestep,
model_options=model_options,
cond=cond,
uncond=uncond,
)
class Guider_AdaptiveGuidance(AdaptiveGuider, comfy.samplers.CFGGuider):
def calc_conds(self, x, timestep, model_options):
class AdaptiveGuidance:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING",),
"threshold": ("FLOAT", {"default": 0.990, "min": 0.90, "max": 1.0, "step": 0.001, "round": 0.001}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step": 0.1, "round": 0.01}),
}
}
RETURN_TYPES = ("GUIDER",)
FUNCTION = "patch"
CATEGORY = "sampling/custom_sampling/guiders"
def patch(self, model, positive, negative, threshold, cfg):
g = AdaptiveGuider(model)
g.set_conds(positive, negative)
g.set_cfg(cfg)
g.set_threshold(threshold)
return (g,)
class LinearAdaptiveGuidance:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING",),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step": 0.1, "round": 0.01}),
"threshold": ("FLOAT", {"default": 0.990, "min": 0.90, "max": 1.0, "step": 0.001, "round": 0.001}),
"betas_cond": ("STRING", {"default": "0.4,0.2,0.05"}),
"betas_uncond": ("STRING", {"default": "0.4,0.2,0.05"}),
}
}
RETURN_TYPES = ("GUIDER",)
FUNCTION = "patch"
CATEGORY = "sampling/custom_sampling/guiders"
def patch(self, model, positive, negative, cfg, threshold, betas_cond, betas_uncond):
g = LinearAdaptiveGuider(model)
g.set_conds(positive, negative)
g.set_cfg(cfg)
g.set_threshold(threshold)
def split_floats(string):
return [float(x.strip()) for x in string.split(",")]
g.set_betas(split_floats(betas_cond), split_floats(betas_uncond))
return (g,)
class LinearAdaptiveGuider(AdaptiveGuider):
last_seen_sigma = 0
def set_betas(self, beta_cond, beta_uncond):
self.beta_cond = beta_cond
self.beta_uncond = beta_uncond
def get_beta(self, beta_list):
idx = min(self.counter - 1, len(beta_list) - 1)
return beta_list[idx]
def initialize(self):
self.cond_results = []
self.uncond_results = []
self.counter = 0
def predict_linear(self):
return torch.stack(self.cond_results, dim=0).sum(dim=0) + torch.stack(self.uncond_results, dim=0).sum(dim=0)
def predict_noise(self, x, timestep, model_options={}, seed=None):
cond = self.conds.get("positive")
uncond = self.conds.get("negative")
return comfy.samplers.calc_cond_batch(self.inner_model, [cond, uncond], x, timestep, model_options)
ts = timestep[0].item()
# Not exactly correct, but will work
if self.last_seen_sigma < ts:
self.initialize()
self.last_seen_sigma = ts
self.counter += 1
if ts < self.threshold_timestep:
return comfy.samplers.sampling_function(
self.inner_model, x, timestep, uncond, cond, 1.0, model_options=model_options, seed=seed
)
def calc_cfg(self, conds, x, timestep, model_options):
cond = self.conds.get("positive")
uncond = self.conds.get("negative")
cond_pred, uncond_pred = conds
else:
self.threshold_timestep = 0
bc = self.get_beta(self.beta_cond)
buc = self.get_beta(self.beta_uncond)
print(f"LinearAdaptive: {bc=} {buc=}")
if self.counter % 2 != 0:
# cfg step
print("LinearAdaptive: Full CFG step")
uncond_pred, cond_pred = comfy.samplers.calc_cond_batch(
self.inner_model, [uncond, cond], x, timestep, model_options
)
self.cond_results.append(cond_pred * bc)
self.uncond_results.append(uncond_pred * buc)
else:
# non-cfg step
print("LinearAdaptive: Estimated CFG step")
cond_pred = comfy.samplers.calc_cond_batch(self.inner_model, [cond], x, timestep, model_options)[0]
self.cond_results.append(cond_pred * bc)
uncond_pred = self.predict_linear()
self.uncond_results.append(uncond_pred * buc)
self.check_cos_sim(ts, cond_pred, uncond_pred)
return comfy.samplers.cfg_function(
self.inner_model,
cond_pred,
uncond_pred,
cond_pred,
self.cfg,
x,
timestep,
@@ -73,227 +174,7 @@ class Guider_AdaptiveGuidance(AdaptiveGuider, comfy.samplers.CFGGuider):
)
class Guider_PerpNegAG(AdaptiveGuider, comfy_extras.nodes_perpneg.Guider_PerpNeg):
def calc_conds(self, x, timestep, model_options):
cond = self.conds.get("positive")
uncond = self.conds.get("negative")
empty_cond = self.conds.get("empty_negative_prompt")
return comfy.samplers.calc_cond_batch(self.inner_model, [cond, uncond, empty_cond], x, timestep, model_options)
def calc_cfg(self, conds, x, timestep, model_options):
cond = self.conds.get("positive")
uncond = self.conds.get("negative")
empty_cond = self.conds.get("empty_negative_prompt")
cond_pred, uncond_pred, empty_cond_pred = conds
cfg_result = comfy_extras.nodes_perpneg.perp_neg(
x, cond_pred, uncond_pred, empty_cond_pred, self.neg_scale, self.cfg
)
for fn in model_options.get("sampler_post_cfg_function", []):
args = {
"denoised": cfg_result,
"cond": cond,
"uncond": uncond,
"model": self.inner_model,
"uncond_denoised": uncond_pred,
"cond_denoised": cond_pred,
"sigma": timestep,
"model_options": model_options,
"input": x,
# not in the original call in samplers.py:cfg_function, but made available for future hooks
"empty_cond": empty_cond,
"empty_cond_denoised": empty_cond_pred,
}
cfg_result = fn(args)
return cfg_result
class AdaptiveGuidanceGuider:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING",),
"threshold": ("FLOAT", {"default": 0.990, "min": 0.90, "max": 1.0, "step": 0.0001, "round": 0.0001}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step": 0.1, "round": 0.01}),
},
"optional": {
"uncond_zero_scale": ("FLOAT", {"default": 0.0, "max": 2.0, "step": 0.01}),
"cfg_start_pct": ("FLOAT", {"default": 0.0, "max": 1.0, "step": 0.01}),
},
}
RETURN_TYPES = ("GUIDER",)
FUNCTION = "get_guider"
CATEGORY = "sampling/custom_sampling/guiders"
def get_guider(self, model, positive, negative, threshold, cfg, uncond_zero_scale=0.0, cfg_start_pct=0.0):
cfg_start_timestep = model.get_model_object("model_sampling").percent_to_sigma(cfg_start_pct)
g = Guider_AdaptiveGuidance(model)
g.set_conds(positive, negative)
g.set_threshold(threshold, cfg_start_timestep)
g.set_uncond_zero_scale(uncond_zero_scale)
g.set_cfg(cfg)
return (g,)
class PerpNegAGGuider:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING",),
"empty_conditioning": ("CONDITIONING",),
"threshold": ("FLOAT", {"default": 0.990, "min": 0.90, "max": 1.0, "step": 0.0001, "round": 0.0001}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0, "step": 0.1, "round": 0.01}),
"neg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01}),
},
"optional": {
"uncond_zero_scale": ("FLOAT", {"default": 0.0, "max": 2.0, "step": 0.01}),
"cfg_start_pct": ("FLOAT", {"default": 0.0, "max": 1.0, "step": 0.01}),
},
}
RETURN_TYPES = ("GUIDER",)
FUNCTION = "get_guider"
CATEGORY = "sampling/custom_sampling/guiders"
def get_guider(
self,
model,
positive,
negative,
empty_conditioning,
threshold,
cfg,
neg_scale,
uncond_zero_scale=0.0,
cfg_start_pct=0.0,
):
cfg_start_timestep = model.get_model_object("model_sampling").percent_to_sigma(cfg_start_pct)
g = Guider_PerpNegAG(model)
g.set_conds(positive, negative, empty_conditioning)
g.set_threshold(threshold, cfg_start_timestep)
g.set_uncond_zero_scale(uncond_zero_scale)
g.set_cfg(cfg, neg_scale)
return (g,)
def project(a, b):
dtype = a.dtype
a, b = a.double(), b.double()
b = torch.nn.functional.normalize(b, dim=[-1, -2, -3])
a_par = (a * b).sum(dim=[-1, -2, -3], keepdim=True) * b
a_orth = a - a_par
return a_par.to(dtype), a_orth.to(dtype)
class AdaptiveProjectedGuidanceFunction:
def __init__(self, momentum, eta, norm_threshold, adaptive_momentum=0, mode="normal"):
self.eta = eta
self.norm_threshold = norm_threshold
self.current_step = 999.0
self.init_momentum = momentum
self.momentum = momentum
self.running_average = 0.0
self.mode = mode
self.adaptive_momentum = adaptive_momentum
def __call__(self, args):
if "denoised" == self.mode:
cond = args["cond_denoised"]
uncond = args["uncond_denoised"]
else:
cond = args["cond"]
uncond = args["uncond"]
cfg_scale = args["cond_scale"]
sigma = args["sigma"][0].item()
step = args["model"].model_sampling.timestep(args["sigma"])[0].item()
x_orig = args["input"]
if self.mode == "vpred":
sigma = step
x = x_orig / (sigma * sigma + 1.0)
cond = ((x - (x_orig - cond)) * (sigma**2 + 1.0) ** 0.5) / (sigma)
uncond = ((x - (x_orig - uncond)) * (sigma**2 + 1.0) ** 0.5) / (sigma)
if self.current_step < step:
self.current_step = 999.0
self.running_average = 0.0
self.momentum = self.init_momentum
else:
scale = self.init_momentum
if self.adaptive_momentum > 0:
scale -= scale * (self.adaptive_momentum**4) * (1000 - step)
if self.init_momentum < 0 and scale > 0:
scale = 0
elif self.init_momentum > 0 and scale < 0:
scale = 0
self.momentum = scale
self.current_step = step
diff = cond - uncond
new_average = self.momentum * self.running_average
self.running_average = diff + new_average
diff = self.running_average
if self.norm_threshold > 0.0:
diff_norm = diff.norm(p=2, dim=[-1, -2, -3], keepdim=True)
scale_factor = torch.minimum(torch.ones_like(diff), self.norm_threshold / diff_norm)
diff = diff * scale_factor
diff_parallel, diff_orthogonal = project(diff, cond)
pred = cond + (cfg_scale - 1) * (diff_orthogonal + self.eta * diff_parallel)
if "denoised" == self.mode:
pred = x_orig - pred
elif "vpred" == self.mode:
pred = x_orig - (x - pred * sigma / (sigma * sigma + 1.0) ** 0.5)
return pred
class AdaptiveProjectedGuidance:
@classmethod
def INPUT_TYPES(s):
return {
"required": {"model": ("MODEL",)},
"optional": {
"momentum": ("FLOAT", {"default": 0.5, "min": -1.0, "max": 1.0, "step": 0.01}),
"eta": ("FLOAT", {"default": 1, "min": 0.0, "max": 1.0, "step": 0.01}),
"norm_threshold": ("FLOAT", {"default": 15.0, "min": 0.0, "max": 50.0, "step": 0.1}),
"mode": (["normal", "denoised", "vpred"],),
"adaptive_momentum": ("FLOAT", {"default": 0.18, "min": 0, "max": 1.0, "step": 0.01}),
},
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "apply"
CATEGORY = "_for_testing"
def apply(self, model, momentum=0.5, eta=1.0, norm_threshold=15.0, mode="normal", adaptive_momentum=0.18):
fn = AdaptiveProjectedGuidanceFunction(momentum, eta, norm_threshold, adaptive_momentum, mode)
m = model.clone()
m.set_model_sampler_cfg_function(fn)
return (m,)
NODE_CLASS_MAPPINGS = {
"AdaptiveGuidance": AdaptiveGuidanceGuider,
"PerpNegAdaptiveGuidanceGuider": PerpNegAGGuider,
"AdaptiveProjectedGuidance": AdaptiveProjectedGuidance,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"AdaptiveGuidance": "AdaptiveGuider",
"PerpNegAdaptiveGuidanceGuider": "PerpNegAdaptiveGuider",
"AdaptiveGuidance": AdaptiveGuidance,
"LinearAdaptiveGuidance": LinearAdaptiveGuidance,
}
File diff suppressed because it is too large Load Diff
-14
View File
@@ -1,14 +0,0 @@
[project]
name = "comfyui-adaptive-guidance"
description = "An implementation of adaptive guidance for ComfyUI\nSee https://bcv-uniandes.github.io/adaptiveguidance-wp/"
version = "0.4.0"
license = { text = "GNU General Public License v3.0" }
[project.urls]
Repository = "https://github.com/asagi4/ComfyUI-Adaptive-Guidance"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "asagi4"
DisplayName = "ComfyUI Adaptive Guidance"
Icon = ""