diff --git a/README.md b/README.md index ce2e0b0..44505e8 100644 --- a/README.md +++ b/README.md @@ -16,14 +16,11 @@ The node sets a unet wrapper function, but attempts to preserve any existing wra The `rescale` parameter applies optional normalization to the noised conditioning. It's disabled at 0. -`apply_to` allows you to apply the noise selectively. +`apply_to` allows you to apply the noise selectively. `key` selects where to add the noise. # Bugs +Noise was previously applied to cross attention. It's now applied by default to the regular conditioning `y`, which seems to make more sense. Use the `key` parameter to restore the old behaviour. The implementation might not be correct at all; I'm not 100% clear on the math as to where the noise is actually supposed to be added. and I couldn't make it produce quite the same results as the A1111 node. The algorithm still seems to help with variety though. - -Not tested with SDXL. Might do weird things. - -I'm not sure if the rescale parameter does anything useful, but feel free to experiment. diff --git a/__init__.py b/__init__.py index 1a27281..32d41c8 100644 --- a/__init__.py +++ b/__init__.py @@ -31,6 +31,7 @@ class CADS: "start_step": ("INT", {"min": -1, "max": 10000, "default": -1}), "total_steps": ("INT", {"min": -1, "max": 10000, "default": -1}), "apply_to": (["both", "cond", "uncond"],), + "key": (["y", "c_crossattn"],), }, } @@ -39,7 +40,7 @@ class CADS: CATEGORY = "utils" - def do(self, model, noise_scale, t1, t2, rescale=0.0, start_step=-1, total_steps=-1, apply_to="both"): + def do(self, model, noise_scale, t1, t2, rescale=0.0, start_step=-1, total_steps=-1, apply_to="both", key="y"): previous_wrapper = model.model_options.get("model_function_wrapper") im = model.model.model_sampling @@ -94,10 +95,10 @@ class CADS: if noise_scale > 0.0: gamma = cads_gamma(timestep) - for i in range(c["c_crossattn"].size(dim=0)): + for i in range(c[key].size(dim=0)): if cond_or_uncond[i % len(cond_or_uncond)] == skip: continue - c["c_crossattn"][i] = cads_noise(gamma, c["c_crossattn"][i]) + c[key][i] = cads_noise(gamma, c[key][i]) if previous_wrapper: return previous_wrapper(apply_model, args) @@ -120,8 +121,6 @@ class CADS: m = model.clone() m.set_model_unet_function_wrapper(apply_cads) - # Alternative implementation. Doesn't seem to do the right thing - # m.set_model_sampler_cfg_function(apply_cads_cfg) return (m,)