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
+118 -9
View File
@@ -7,10 +7,17 @@ cos = torch.nn.CosineSimilarity(dim=1)
class AdaptiveGuider(comfy.samplers.CFGGuider): class AdaptiveGuider(comfy.samplers.CFGGuider):
threshold_timestep = 0 threshold_timestep = 0
def set_cfg(self, cfg, threshold): def set_threshold(self, threshold):
self.cfg = cfg
self.threshold = threshold self.threshold = threshold
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): def predict_noise(self, x, timestep, model_options={}, seed=None):
cond = self.conds.get("positive") cond = self.conds.get("positive")
uncond = self.conds.get("negative") uncond = self.conds.get("negative")
@@ -24,11 +31,7 @@ class AdaptiveGuider(comfy.samplers.CFGGuider):
uncond_pred, cond_pred = comfy.samplers.calc_cond_batch( uncond_pred, cond_pred = comfy.samplers.calc_cond_batch(
self.inner_model, [uncond, cond], x, timestep, model_options self.inner_model, [uncond, cond], x, timestep, model_options
) )
# Is this reshape correct? It at least gives a scalar value... self.check_cos_sim()
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( return comfy.samplers.cfg_function(
self.inner_model, self.inner_model,
cond_pred, cond_pred,
@@ -63,9 +66,115 @@ class AdaptiveGuidance:
def patch(self, model, positive, negative, threshold, cfg): def patch(self, model, positive, negative, threshold, cfg):
g = AdaptiveGuider(model) g = AdaptiveGuider(model)
g.set_conds(positive, negative) g.set_conds(positive, negative)
g.set_cfg(cfg, threshold) g.set_cfg(cfg)
g.set_threshold(threshold)
return (g,) 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,
}