From 9b2a19eef98f9716aa770cb99250437681b5a211 Mon Sep 17 00:00:00 2001 From: blepping Date: Sun, 24 Mar 2024 04:01:18 -0600 Subject: [PATCH 1/2] Fall back to using c_crossattn if y key does not exist --- __init__.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/__init__.py b/__init__.py index 048765c..9092103 100644 --- a/__init__.py +++ b/__init__.py @@ -94,11 +94,12 @@ class CADS: c = args["c"] if noise_scale > 0.0: + apply_to = c.get(key, c["c_crossattn"]) gamma = cads_gamma(timestep) - for i in range(c[key].size(dim=0)): + for i in range(apply_to.size(dim=0)): if cond_or_uncond[i % len(cond_or_uncond)] == skip: continue - c[key][i] = cads_noise(gamma, c[key][i]) + apply_to[i] = cads_noise(gamma, apply_to[i]) if previous_wrapper: return previous_wrapper(apply_model, args) From b77cee21a6220f14d0e4172e868fd0a009cc52f2 Mon Sep 17 00:00:00 2001 From: blepping Date: Sun, 24 Mar 2024 04:06:24 -0600 Subject: [PATCH 2/2] Change variable name to avoid confusion with apply_to --- __init__.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/__init__.py b/__init__.py index 9092103..f7f19a5 100644 --- a/__init__.py +++ b/__init__.py @@ -94,12 +94,12 @@ class CADS: c = args["c"] if noise_scale > 0.0: - apply_to = c.get(key, c["c_crossattn"]) + noise_target = c.get(key, c["c_crossattn"]) gamma = cads_gamma(timestep) - for i in range(apply_to.size(dim=0)): + for i in range(noise_target.size(dim=0)): if cond_or_uncond[i % len(cond_or_uncond)] == skip: continue - apply_to[i] = cads_noise(gamma, apply_to[i]) + noise_target[i] = cads_noise(gamma, noise_target[i]) if previous_wrapper: return previous_wrapper(apply_model, args)