From 1e8fba93844a7426cb28ff281ddffdb498785650 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 29 Mar 2024 06:22:30 -0500 Subject: [PATCH] Made CLIP LoRA hooks work as intended, proper way to force ModelPatcherAndInjector to cause model reload when hooks present --- animatediff/model_injection.py | 248 +++++++++++++++++++++++++++------ animatediff/nodes.py | 4 +- animatediff/sampling.py | 2 - animatediff/utils_motion.py | 2 + 4 files changed, 213 insertions(+), 43 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 86b13fe..affb13c 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -93,6 +93,23 @@ class ModelPatcherAndInjector(ModelPatcher): else: return ModelPatcherAndInjector(model) + def clone_has_same_weights(self, clone: 'ModelPatcherCLIPHooks'): + returned = super().clone_has_same_weights(clone) + if not returned: + return returned + # currently, hook patches require that model gets loaded when sampled, so always say is not a clone if hooks present + if len(self.hooked_patches) > 0 or len(self.hooked_replace_patches) > 0: + return False + if type(self) != type(clone): + return False + if self.current_lora_hooks != clone.current_lora_hooks: + return False + if self.hooked_patches.keys() != clone.hooked_patches.keys(): + return False + if self.hooked_replace_patches.keys() != clone.hooked_replace_patches.keys(): + return False + return returned + def set_lora_hook_mode(self, lora_hook_mode: str): self.lora_hook_mode = lora_hook_mode @@ -208,7 +225,6 @@ class ModelPatcherAndInjector(ModelPatcher): # unpatch hooks, if needed self.unpatch_hooked(device_to=device_to) # finally, patch hooks - # TODO: handle lowvram self.patch_hooked(lora_hooks=lora_hooks, device_to=device_to) def patch_hooked(self, lora_hooks: LoraHookGroup, device_to=None, patch_weights=True) -> None: @@ -235,6 +251,7 @@ class ModelPatcherAndInjector(ModelPatcher): logger.warning(f"Cached LoraHook hook could not patch. key doesn't exist in model: {key}") self.patch_cached_hooked_weight(cached_weights=cached_weights, key=key) else: + # TODO: handle lowvram # get combined patches of relevant lora_hooks relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks) replace_patches = self.get_hooked_replace_patches(lora_hooks=lora_hooks) @@ -255,6 +272,7 @@ class ModelPatcherAndInjector(ModelPatcher): pass def patch_cached_hooked_weight(self, cached_weights: dict, key: str): + # TODO: handle lowvram inplace_update = self.weight_inplace_update target_device = self.offload_device if self.lora_hook_mode == LoraHookMode.MAX_SPEED: @@ -359,14 +377,13 @@ class ModelPatcherAndInjector(ModelPatcher): class CLIPWithHooks(CLIP): def __init__(self, clip: Union[CLIP, 'CLIPWithHooks']): super().__init__(no_init=True) - self.patcher = ModelPatcherAndInjector.create_from(clip.patcher) - self.patcher.set_lora_hook_mode(lora_hook_mode=LoraHookMode.MIN_VRAM) + self.patcher = ModelPatcherCLIPHooks.create_from(clip.patcher) self.cond_stage_model = clip.cond_stage_model self.tokenizer = clip.tokenizer self.layer_idx = clip.layer_idx self.desired_hooks: LoraHookGroup = None if hasattr(clip, "desired_hooks"): - self.desired_hooks = clip.desired_hooks + self.set_desired_hooks(clip.desired_hooks) def clone(self): cloned = CLIPWithHooks(clip=self) @@ -374,6 +391,7 @@ class CLIPWithHooks(CLIP): def set_desired_hooks(self, lora_hooks: LoraHookGroup): self.desired_hooks = lora_hooks + self.patcher.set_desired_hooks(lora_hooks=lora_hooks) def add_hooked_patches(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0): return self.patcher.add_hooked_patches(lora_hook=lora_hook, patches=patches, strength_patch=strength_patch, strength_model=strength_model) @@ -381,43 +399,195 @@ class CLIPWithHooks(CLIP): def add_hooked_replace_patches(self, lora_hook: LoraHook, patches): return self.patcher.add_hooked_replace_patches(lora_hook=lora_hook, patches=patches) - # def load_model(self, *args, **kwargs): - # #comfy.model_management.cleanup_models() - # #return super().load_model(*args, **kwargs) - # self.patcher.unpatch_hooked() - # returned = super().load_model(*args, **kwargs) - # # apply desired hooks - # self.patcher.patch_hooked(lora_hooks=self.desired_hooks) - # return returned + # def encode_from_tokens(self, tokens, return_pooled=False): + # comfy.model_management.cleanup_models() + # return super().encode_from_tokens(tokens, return_pooled) - def encode_from_tokens(self, *args, **kwargs): - # to work properly, need to hack (and then unhack) cond_stage_model's encode_token_weights function - def encode_token_weights_factory(orig_encode_token_weights: Callable, clip: CLIPWithHooks): - def encode_token_weights_wrapper_hooked(*args, **kwargs): - try: - # yeah, not sure why, but first patching is screwed up if things change... BUT only the first time. - # so we apply it twice to avoid it... this will be future me's problem. - # the first result will still end up slightly different than intended, but it's very close. - if clip.desired_hooks is not None: - temp_desired_hooks = clip.desired_hooks.clone() if clip.desired_hooks is not None else None - clip.patcher.clear_cached_hooked_weights() - clip.patcher.unpatch_hooked() - clip.patcher.apply_lora_hooks(temp_desired_hooks) - orig_encode_token_weights(*args, **kwargs) - clip.patcher.clear_cached_hooked_weights() - clip.patcher.apply_lora_hooks(clip.desired_hooks) - return orig_encode_token_weights(*args, **kwargs) - finally: - clip.patcher.unpatch_hooked() - clip.patcher.clear_cached_hooked_weights() - return encode_token_weights_wrapper_hooked - try: - orig_encode_token_weights = self.cond_stage_model.encode_token_weights - self.cond_stage_model.encode_token_weights = encode_token_weights_factory(orig_encode_token_weights, self) - return super().encode_from_tokens(*args, **kwargs) - finally: - self.cond_stage_model.encode_token_weights = orig_encode_token_weights + +class ModelPatcherCLIPHooks(ModelPatcher): + def __init__(self, m: ModelPatcher): + # replicate ModelPatcher.clone() to initialize + super().__init__(m.model, m.load_device, m.offload_device, m.size, m.current_device, weight_inplace_update=m.weight_inplace_update) + self.patches = {} + for k in m.patches: + self.patches[k] = m.patches[k][:] + if hasattr(m, "patches_uuid"): + self.patches_uuid = m.patches_uuid + + self.object_patches = m.object_patches.copy() + self.model_options = copy.deepcopy(m.model_options) + self.model_keys = m.model_keys + if hasattr(m, "backup"): + self.backup = m.backup + if hasattr(m, "object_patches_backup"): + self.object_patches_backup = m.object_patches_backup + # lora hook stuff + self.hooked_patches = {} # binds LoraHook to specific keys + self.patches_backup = {} + self.hooked_backup = {} + + self.current_lora_hooks = None + self.desired_lora_hooks = None + self.lora_hook_mode = LoraHookMode.MAX_SPEED + # replace hook stuff + self.hooked_replace_patches = {} # binds LoraHook to specific replacement keys + + def clone(self): + cloned = ModelPatcherCLIPHooks(self) + # copy lora hooks + for hook in self.hooked_patches: + cloned.hooked_patches[hook] = {} + for k in self.hooked_patches[hook]: + cloned.hooked_patches[hook][k] = self.hooked_patches[hook][k][:] + # copy replace lora hooks + for hook in self.hooked_replace_patches: + cloned.hooked_replace_patches[hook] = {} + for k in self.hooked_replace_patches[hook]: + cloned.hooked_replace_patches[hook][k] = self.hooked_replace_patches[hook][k] + cloned.patches_backup = self.patches_backup + cloned.hooked_backup = self.hooked_backup + cloned.current_lora_hooks = self.current_lora_hooks + cloned.desired_lora_hooks = self.desired_lora_hooks + cloned.lora_hook_mode = self.lora_hook_mode + return cloned + + @classmethod + def create_from(cls, model: Union[ModelPatcher, 'ModelPatcherCLIPHooks']): + if isinstance(model, ModelPatcherCLIPHooks): + return model.clone() + return ModelPatcherCLIPHooks(model) + + def clone_has_same_weights(self, clone: 'ModelPatcherCLIPHooks'): + returned = super().clone_has_same_weights(clone) + if not returned: + return returned + if type(self) != type(clone): + return False + if self.desired_lora_hooks != clone.desired_lora_hooks: + return False + if self.current_lora_hooks != clone.current_lora_hooks: + return False + if self.hooked_patches.keys() != clone.hooked_patches.keys(): + return False + if self.hooked_replace_patches.keys() != clone.hooked_replace_patches.keys(): + return False + return returned + + def set_desired_hooks(self, lora_hooks: LoraHookGroup): + self.desired_lora_hooks = lora_hooks + + def add_hooked_patches(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0): + ''' + Based on add_patches, but for hooked weights. + ''' + # TODO: make this work with timestep scheduling + current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook, {}) + p = set() + for key in patches: + if key in self.model_keys: + p.add(key) + current_patches: list[tuple] = current_hooked_patches.get(key, []) + current_patches.append((strength_patch, patches[key], strength_model)) + current_hooked_patches[key] = current_patches + self.hooked_patches[lora_hook] = current_hooked_patches + # since should care about these patches too to determine if same model, reroll patches_uuid + self.patches_uuid = uuid.uuid4() + return list(p) + + def get_combined_hooked_patches(self, lora_hooks: LoraHookGroup): + ''' + Returns patches for selected lora_hooks. + ''' + # combined_patches will contain weights of all relevant lora_hooks, per key + combined_patches = {} + if lora_hooks is not None: + for hook in lora_hooks.hooks: + hook_patches: dict = self.hooked_patches.get(hook, {}) + for key in hook_patches.keys(): + current_patches: list[tuple] = combined_patches.get(key, []) + current_patches.extend(hook_patches[key]) + combined_patches[key] = current_patches + return combined_patches + + def add_hooked_replace_patches(self, lora_hook: LoraHook, patches: dict): + self.hooked_replace_patches.setdefault(lora_hook, {}) + p = set() + for key in patches: + if key in self.model_keys: + p.add(key) + self.hooked_replace_patches[lora_hook][key] = patches[key] + # since should care about these patches too to determine if same model, reroll patches_uuid + self.patches_uuid = uuid.uuid4() + return list(p) + + def get_hooked_replace_patches(self, lora_hooks: LoraHookGroup): + # return first hook found in hooked_replace_patches + patches = {} + if lora_hooks is not None: + for hook in lora_hooks.hooks: + if hook in self.hooked_replace_patches: + patches = self.hooked_replace_patches[hook] + break + return patches + + def patch_hooked_replace_weight_to_device(self, model_sd: dict, replace_patches: dict): + # first handle replace_patches + for key in replace_patches: + if key not in model_sd: + logger.warning(f"CLIP LoraHook could not replace patch. key doesn't exist in model: {key}") + continue + weight: Tensor = comfy.utils.get_attr(self.model, key) + inplace_update = self.weight_inplace_update + target_device = self.current_device + if key not in self.hooked_backup: + self.hooked_backup[key] = weight.to(device=target_device, copy=inplace_update) + out_weight = replace_patches[key].to(target_device) + if inplace_update: + comfy.utils.copy_to_param(self.model, key, out_weight) + else: + comfy.utils.set_attr_param(self.model, key, out_weight) + + def patch_model(self, device_to=None, patch_weights=True, *args, **kwargs): + if self.desired_lora_hooks is not None: + self.patches_backup = self.patches.copy() + # first, handle replace patches # TODO: make work properly for CLIP + replace_patches = self.get_hooked_replace_patches(lora_hooks=self.desired_lora_hooks) + if len(replace_patches) > 0: + model_sd = self.model_state_dict() + self.patch_hooked_replace_weight_to_device(model_sd=model_sd, replace_patches=replace_patches) + # then, handle usual patches + relevant_patches = self.get_combined_hooked_patches(lora_hooks=self.desired_lora_hooks) + for key in relevant_patches: + self.patches.setdefault(key, []) + self.patches[key].extend(relevant_patches[key]) + self.current_lora_hooks = self.desired_lora_hooks + return super().patch_model(device_to, patch_weights, *args, **kwargs) + + def unpatch_model(self, device_to=None, unpatch_weights=True, *args, **kwargs): + try: + return super().unpatch_model(device_to, unpatch_weights, *args, **kwargs) + finally: + self.patches = self.patches_backup.copy() + self.patches_backup.clear() + # handle replace patches + keys = list(self.hooked_backup.keys()) + if self.weight_inplace_update: + for k in keys: + if device_to is None: + comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k].to(device=self.current_device)) + else: + comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k]) + else: + for k in keys: + if device_to is None: + comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k].to(device=self.current_device)) + else: + comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k]) + # clear hooked_backup + self.hooked_backup.clear() + self.current_lora_hooks = None + def load_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcherAndInjector], clip: CLIP, lora: dict[str, Tensor], lora_hook: LoraHook, strength_model: float, strength_clip: float): key_map = {} diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 6fd87f8..4ef4bc2 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -135,13 +135,13 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Conditioning "ADE_RegisterLoraHook": "Register LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_RegisterLoraHookModelOnly": "Register LoRA Hook (Model Only) πŸŽ­πŸ…πŸ…“", - #"ADE_RegisterModelAsLoraHook": "πŸ”¬Register Model as LoRA Hook πŸŽ­πŸ…πŸ…“", # CLIP does not work properly on first run + #"ADE_RegisterModelAsLoraHook": "Register Model as LoRA HookπŸ”¬ πŸŽ­πŸ…πŸ…“", # CLIP does not work properly on first run "ADE_RegisterModelAsLoraHookModelOnly": "Register Model as LoRA Hook (MO) πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooks": "Combine LoRA Hooks [2] πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooksFour": "Combine LoRA Hooks [4] πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooksEight": "Combine LoRA Hooks [8] πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToConditioning": "Set Model LoRA Hook πŸŽ­πŸ…πŸ…“", - "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA HookπŸ”¬ πŸŽ­πŸ…πŸ…“", + "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA Hook πŸŽ­πŸ…πŸ…“", # Noise Layer Nodes "ADE_NoiseLayerAdd": "Noise Layer [Add] πŸŽ­πŸ…πŸ…“", "ADE_NoiseLayerAddWeighted": "Noise Layer [Add Weighted] πŸŽ­πŸ…πŸ…“", diff --git a/animatediff/sampling.py b/animatediff/sampling.py index e202179..f483ed1 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -277,8 +277,6 @@ def motion_sample_factory(orig_comfy_sample: Callable, is_custom: bool=False) -> cached_noise = None function_injections = FunctionInjectionHolder() try: - if len(model.hooked_patches) > 0: - model_management.cleanup_models() if model.sample_settings.custom_cfg is not None: model = model.sample_settings.custom_cfg.patch_model(model) # clone params from model diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index 1ed0cd6..5b8f1f6 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -253,7 +253,9 @@ class ADKeyframeGroup: class LoraHookMode: MIN_VRAM = "min_vram" + MIN_VRAM_LOWVRAM = "min_vram_lowvram" MAX_SPEED = "max_speed" + MAX_SPEED_LOWVRAM = "max_speed_lowvram" class LoraHook: