diff --git a/control.py b/control.py index e052a11..ff725dd 100644 --- a/control.py +++ b/control.py @@ -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)