remove old code from ControlNetAdvanced (forgot to do it last commit)

This commit is contained in:
Jedrzej Kosinski
2023-09-01 05:34:26 -05:00
parent 565fdf44a9
commit dcbccd43de
-72
View File
@@ -228,78 +228,6 @@ class ControlNetAdvanced(ControlBase):
o[i] += prev_val
return out
def get_control_old(self, x_noisy, t, cond, batched_number):
control_prev = None
if self.previous_controlnet is not None:
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
if self.timestep_range is not None:
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
if control_prev is not None:
return control_prev
else:
return {}
output_dtype = x_noisy.dtype
if self.cond_hint is None or x_noisy.shape[2] * 8 != self.cond_hint.shape[2] or x_noisy.shape[3] * 8 != self.cond_hint.shape[3]:
if self.cond_hint is not None:
del self.cond_hint
self.cond_hint = None
self.cond_hint = utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(self.control_model.dtype).to(self.device)
if x_noisy.shape[0] != self.cond_hint.shape[0]:
self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number)
if self.control_model.dtype == torch.float16:
precision_scope = torch.autocast
else:
precision_scope = contextlib.nullcontext
# TODO: select based on progress in diffusion
current_timestep_keyframe = self.timestep_keyframes[0]
with precision_scope(model_management.get_autocast_device(self.device)):
context = torch.cat(cond['c_crossattn'], 1)
y = cond.get('c_adm', None)
control = self.control_model(x=x_noisy, hint=self.cond_hint, timesteps=t, context=context, y=y)
out = {'middle':[], 'output': []}
autocast_enabled = torch.is_autocast_enabled()
for i in range(len(control)):
if i == (len(control) - 1):
key = 'middle'
index = 0
else:
key = 'output'
index = i
x = control[i]
if self.global_average_pooling:
x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3])
if current_timestep_keyframe.latent_keyframes is not None:
# get batch indeces to zero out, AKA latents that should not be influenced by ControlNet
indeces_to_zero = set(range(x.size()[0]//2))
for keyframe in current_timestep_keyframe.latent_keyframes:
if keyframe.batch_index in indeces_to_zero:
indeces_to_zero.remove(keyframe.batch_index)
# zero them out by multiplying by zero
for batch_index in indeces_to_zero:
x[batch_index] *= 0.0
x[(x.size()[0]//2) + batch_index] *= 0.0
x *= self.strength * self.weights[i]
if x.dtype != output_dtype and not autocast_enabled:
x = x.to(output_dtype)
if control_prev is not None and key in control_prev:
prev = control_prev[key][index]
if prev is not None:
x += prev
out[key].append(x)
if control_prev is not None and 'input' in control_prev:
out['input'] = control_prev['input']
return out
def copy(self):
c = ControlNetAdvanced(self.control_model, self.timestep_keyframes, global_average_pooling=self.global_average_pooling)
self.copy_to(c)