remove old code from ControlNetAdvanced (forgot to do it last commit)
This commit is contained in:
-72
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user