diff --git a/adv_control/control.py b/adv_control/control.py index ee9e3ba..d3ebc3b 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -64,7 +64,7 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase): # prepare mask_cond_hint self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, dtype=dtype) - context = cond['c_crossattn'] + context = cond.get('crossattn_controlnet', cond['c_crossattn']) # uses 'y' in new ComfyUI update y = cond.get('y', None) if y is None: # TODO: remove this in the future since no longer used by newest ComfyUI @@ -288,6 +288,60 @@ class SparseCtrlAdvanced(ControlNetAdvanced): return c +class ReferenceAdvanced(ControlBase, AdvancedControlBase): + def __init__(self, timestep_keyframes: TimestepKeyframeGroup, device=None): + super().__init__(device) + AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite(), require_model=True) + # TODO: save attn patches here + + def patch_model(self, model: ModelPatcher): + # TODO: do model patching here + pass + + def pre_run_advanced(self, *args, **kwargs): + AdvancedControlBase.pre_run_advanced(self, *args, **kwargs) + # TODO: set control on patches + + def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int): + # normal ControlNet stuff + 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]: + return control_prev + + dtype = x_noisy.dtype + # prepare cond_hint + if self.sub_idxs is not None or 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 + # if self.cond_hint_original length greater or equal to real latent count, subdivide it before scaling + if self.sub_idxs is not None and self.cond_hint_original.size(0) >= self.full_latent_length: + self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.device) + else: + self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(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) + # prepare mask + self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number) + # done preparing; model patches will take care of everything now. + # return normal controlnet stuff + return control_prev + + def cleanup_advanced(self): + super().cleanup_advanced() + # TODO: cleanup patches here + + def copy(self): + c = ReferenceAdvanced(self.timestep_keyframes) + self.copy_to(c) + self.copy_to_advanced(c) + return c + + class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): # This ControlNet is more of an attention patch than a traditional controlnet def __init__(self, patch_attn1: LLLitePatch, patch_attn2: LLLitePatch, timestep_keyframes: TimestepKeyframeGroup, device=None): @@ -297,7 +351,6 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): self.patch_attn2 = patch_attn2.set_control(self) self.latent_dims_div2 = None self.latent_dims_div4 = None - def patch_model(self, model: ModelPatcher): model.set_model_attn1_patch(self.patch_attn1) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py new file mode 100644 index 0000000..e69de29