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
4 changed files with 133 additions and 86 deletions
-21
View File
@@ -1,21 +0,0 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- master
paths:
- "pyproject.toml"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
-8
View File
@@ -11,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.
+133 -43
View File
@@ -6,56 +6,43 @@ cos = torch.nn.CosineSimilarity(dim=1)
class AdaptiveGuider(comfy.samplers.CFGGuider):
threshold_timestep = 0
uz_scale = 0.0
def set_cfg(self, cfg):
self.cfg = cfg
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_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 self.threshold_timestep > ts:
if self.uz_scale > 0.0:
model_options = model_options.copy()
model_options["sampler_cfg_function"] = self.zero_cond
return comfy.samplers.sampling_function(
self.inner_model, x, timestep, uncond, cond, 1.0, model_options=model_options, seed=seed
)
self.threshold_timestep = 0
uncond_pred, cond_pred = comfy.samplers.calc_cond_batch(
self.inner_model, [uncond, cond], x, timestep, model_options
)
if not self.threshold >= 1.0:
# Is this reshape correct? It at least gives a scalar value...
sim = cos(cond_pred.reshape(1, -1), uncond_pred.reshape(1, -1)).item()
if sim >= self.threshold:
print("AdaptiveGuidance: Cosine similarity", sim, "exceeds threshold, setting CFG to 1.0")
self.threshold_timestep = ts
return comfy.samplers.cfg_function(
self.inner_model,
cond_pred,
uncond_pred,
self.cfg,
x,
timestep,
model_options=model_options,
cond=cond,
uncond=uncond,
)
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 AdaptiveGuidance:
@@ -68,8 +55,7 @@ class AdaptiveGuidance:
"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}),
},
"optional": {"uncond_zero_scale": ("FLOAT", {"default": 0.0, "max": 2.0, "step": 0.01})},
}
}
RETURN_TYPES = ("GUIDER",)
@@ -77,14 +63,118 @@ class AdaptiveGuidance:
CATEGORY = "sampling/custom_sampling/guiders"
def patch(self, model, positive, negative, threshold, cfg, uncond_zero_scale=0.0):
def patch(self, model, positive, negative, threshold, cfg):
g = AdaptiveGuider(model)
g.set_conds(positive, negative)
g.set_threshold(threshold)
g.set_uncond_zero_scale(uncond_zero_scale)
g.set_cfg(cfg)
g.set_threshold(threshold)
return (g,)
NODE_CLASS_MAPPINGS = {"AdaptiveGuidance": AdaptiveGuidance}
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")
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
)
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,
uncond_pred,
cond_pred,
self.cfg,
x,
timestep,
model_options=model_options,
cond=cond,
uncond=uncond,
)
NODE_CLASS_MAPPINGS = {
"AdaptiveGuidance": AdaptiveGuidance,
"LinearAdaptiveGuidance": LinearAdaptiveGuidance,
}
-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.1.0"
license = "GPL-3.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 = ""