From 069ec7a64bba4582f11cd812b4c7b5aa914613e1 Mon Sep 17 00:00:00 2001 From: "kosinkadink1@gmail.com" Date: Fri, 13 Sep 2024 17:20:22 +0900 Subject: [PATCH] Added hook_patches to ModelPatcher for weights (model) --- comfy/hooks.py | 204 +++++++++++++++++++++++++++++++++++++++++ comfy/model_patcher.py | 203 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 407 insertions(+) create mode 100644 comfy/hooks.py diff --git a/comfy/hooks.py b/comfy/hooks.py new file mode 100644 index 000000000..5b07e6aec --- /dev/null +++ b/comfy/hooks.py @@ -0,0 +1,204 @@ +from typing import TYPE_CHECKING, List, Dict, Tuple +import enum + +if TYPE_CHECKING: + from comfy.model_base import BaseModel + +class EnumHookMode(enum.Enum): + MinVram = "minvram" + MaxSpeed = "maxspeed" + +class HookRef: + pass + +class HookWeight: + def __init__(self): + self.hook_ref = HookRef() + self.hook_keyframe = HookWeightKeyframeGroup() + + @property + def strength(self): + return self.hook_keyframe.strength + + def initialize_timesteps(self, model: 'BaseModel'): + self.hook_keyframe.initalize_timesteps(model) + + def reset(self): + self.hook_keyframe.reset() + + def clone(self): + c = HookWeight() + c.hook_ref = self.hook_ref + c.hook_keyframe = self.hook_keyframe + return c + + def __eq__(self, other: 'HookWeight'): + return self.__class__ == other.__class__ and self.hook_ref == other.hook_ref + + def __hash__(self): + return hash(self.hook_ref) + +class HookWeightGroup: + def __init__(self): + self.hooks: List[HookWeight] = [] + + def add(self, hook: HookWeight): + if hook not in self.hooks: + self.hooks.append(hook) + + def contains(self, hook: HookWeight): + return hook in self.hooks + + def clone(self): + c = HookWeightGroup() + # TODO: review if clone is necessary + for hook in self.hooks: + c.add(hook.clone()) + return c + + def clone_and_combine(self, other: 'HookWeightGroup'): + c = self.clone() + for hook in other.hooks: + c.add(hook.clone()) + return c + + def set_keyframes_on_hooks(self, hook_kf: 'HookWeightKeyframeGroup'): + hook_kf = hook_kf.clone() + for hook in self.hooks: + hook.hook_keyframe = hook_kf + + @staticmethod + def combine_all_hooks(hooks_list: List['HookWeightGroup'], require_count=1) -> 'HookWeightGroup': + actual: List[HookWeightGroup] = [] + for group in hooks_list: + if group is not None: + actual.append(group) + if len(actual) < require_count: + raise Exception(f"Need at least {require_count} hooks to combine, but only had {len(actual)}.") + # if only 1 hook, just reutnr itself without cloning + if len(actual) == 1: + return actual[0] + final_hook: HookWeightGroup = None + for hook in actual: + if final_hook is None: + final_hook = hook.clone() + else: + final_hook - final_hook.clone_and_combine() + return final_hook + +class HookWeightKeyframe: + def __init__(self, strength: float, start_percent=0.0, guarantee_steps=1): + self.strength = strength + # scheduling + self.start_percent = float(start_percent) + self.start_t = 999999999.9 + self.guarantee_steps = guarantee_steps + + def clone(self): + c = HookWeightKeyframe(strength=self.strength, + start_percent=self.start_percent, guarantee_steps=self.guarantee_steps) + c.start_t = self.start_t + return c + +class HookWeightKeyframeGroup: + def __init__(self): + self.keyframes: List[HookWeightKeyframe] = [] + self._current_keyframe: HookWeightKeyframe = None + self._current_used_steps = 0 + self._current_index = 0 + self._curr_t = -1. + + # properties shadow those of HookWeightsKeyframe + @property + def strength(self): + if self._current_keyframe is not None: + return self._current_keyframe.strength + return 1.0 + + def reset(self): + self._current_keyframe = None + self._current_used_steps = 0 + self._current_index = 0 + self.curr_t = -1. + + def add(self, keyframe: HookWeightKeyframe): + # add to end of list, then sort + self.keyframes.append(keyframe) + self.keyframes = get_sorted_list_via_attr(self.keyframes, "start_percent") + self._set_first_as_current() + + def _set_first_as_current(self): + if len(self.keyframes) > 0: + self._current_keyframe = self.keyframes[0] + else: + self._current_keyframe = None + + def has_index(self, index: int): + return index >= 0 and index < len(self.keyframes) + + def is_empty(self): + return len(self.keyframes) == 0 + + def clone(self): + c = HookWeightKeyframeGroup() + for keyframe in self.keyframes: + c.keyframes.append(keyframe) + c._set_first_as_current() + return c + + def initalize_timesteps(self, model: 'BaseModel'): + for keyframe in self.keyframes: + keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent) + + def prepare_current_keyframe(self, curr_t: float) -> bool: + if self.is_empty(): + return False + if curr_t == self._curr_t: + return False + prev_index = self._current_index + # if met guaranteed steps, look for next keyframe in case need to switch + if self._current_used_steps >= self._current_keyframe.guarantee_steps: + # if has next index, loop through and see if need to switch + if self.has_index(self._current_index+1): + for i in range(self._current_index+1, len(self.keyframes)): + eval_c = self.keyframes[i] + # check if start_t is greater or equal to curr_t + # NOTE: t is in terms of sigmas, not percent, so bigger number = earlier step in sampling + if eval_c.start_t >= curr_t: + self._current_index = i + self._current_keyframe = eval_c + self._current_used_steps = 0 + # if guarantee_steps greater than zero, stop searching for other keyframes + if self._current_keyframe.guarantee_steps > 0: + break + # if eval_c is outside the percent range, stop looking further + else: break + # update steps current context is used + self._current_used_steps += 1 + # update current timestep this was performed on + self._curr_t = curr_t + # return True if keyframe changed, False if no change + return prev_index != self._current_index + +def get_sorted_list_via_attr(objects: List, attr: str) -> List: + if not objects: + return objects + elif len(objects) <= 1: + return [x for x in objects] + # now that we know we have to sort, do it following these rules: + # a) if objects have same value of attribute, maintain their relative order + # b) perform sorting of the groups of objects with same attributes + unique_attrs = {} + for o in objects: + val_attr = getattr(o, attr) + attr_list: List = unique_attrs.get(val_attr, list()) + attr_list.append(o) + if val_attr not in unique_attrs: + unique_attrs[val_attr] = attr_list + # now that we have the unique attr values grouped together in relative order, sort them by key + sorted_attrs = dict(sorted(unique_attrs.items())) + # now flatten out the dict into a list to return + sorted_list = [] + for object_list in sorted_attrs.values(): + sorted_list.extend(object_list) + return sorted_list diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index a89ad4fd7..6995f406d 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -16,6 +16,7 @@ along with this program. If not, see . """ +from typing import Dict, List, Tuple, Optional import torch import copy import inspect @@ -28,6 +29,7 @@ import comfy.utils import comfy.float import comfy.model_management import comfy.lora +import comfy.hooks from comfy.comfy_types import UnetWrapperFunction def string_to_seed(data): @@ -111,6 +113,13 @@ class ModelPatcher: self.weight_inplace_update = weight_inplace_update self.patches_uuid = uuid.uuid4() + self.hook_patches: Dict[comfy.hooks.HookRef] = {} + self.hook_backup: Dict[str, Tuple[torch.Tensor, torch.device]] = {} + self.cached_hook_patches: Dict[comfy.hooks.HookWeightGroup, Dict[str, torch.Tensor]] = {} + self.current_hooks: Optional[comfy.hooks.HookWeightGroup] = None + # TODO: hook_mode should be entirely removed; behavior should be determined by remaining VRAM/memory + self.hook_mode = comfy.hooks.EnumHookMode.MaxSpeed + if not hasattr(self.model, 'model_loaded_weight_memory'): self.model.model_loaded_weight_memory = 0 @@ -143,6 +152,20 @@ class ModelPatcher: n.model_options = copy.deepcopy(self.model_options) n.backup = self.backup n.object_patches_backup = self.object_patches_backup + + # hooks + for hook_ref in self.hook_patches: + n.hook_patches[hook_ref] = {} + for k in self.hook_patches[hook_ref]: + n.hook_patches[hook_ref][k] = self.hook_patches[hook_ref][k][:] + # TODO: do we really need to clone cached_hook_patches/current_hooks? + for group in self.cached_hook_patches: + n.cached_hook_patches[group] = {} + for k in self.cached_hook_patches[group]: + n.cached_hook_patches[group][k] = self.cached_hook_patches[group][k] + n.hook_backup = self.hook_backup + n.current_hooks = self.current_hooks + n.hook_mode = self.hook_mode return n def is_clone(self, other): @@ -530,3 +553,183 @@ class ModelPatcher: def calculate_weight(self, patches, weight, key, intermediate_dtype=torch.float32): print("WARNING the ModelPatcher.calculate_weight function is deprecated, please use: comfy.lora.calculate_weight instead") return comfy.lora.calculate_weight(patches, weight, key, intermediate_dtype=intermediate_dtype) + + def set_hook_mode(self, hook_mode: comfy.hooks.EnumHookMode): + self.hook_mode = hook_mode + + def prepare_hook_patches_current_keyframe(self, t: torch.Tensor, hook_groups: List[comfy.hooks.HookWeightGroup]): + curr_t = t[0] + for hook_group in hook_groups: + for hook in hook_group.hooks: + changed = hook.hook_keyframe.prepare_current_keyframe(curr_t=curr_t) + # if keyframe changed, remove any cached LoraHookGroups that contain hook with the same hook_ref; + # this will cause the weights to be recalculated when sampling + if changed: + # reset current_lora_hooks if contains lora hook that changed + if self.current_hooks is not None: + for current_hook in self.current_hooks.hooks: + if current_hook == hook: + self.current_hooks = None + break + for cached_group in list(self.cached_hook_patches.keys()): + if cached_group.contains(hook): + self.cached_hook_patches.pop(cached_group) + + def add_hook_patches(self, hook: comfy.hooks.HookWeight, patches, strength_patch=1.0, strength_model=1.0): + # NOTE: this mirrors behavior of add_patches func + current_hook_patches: Dict[str,List] = self.hook_patches.get(hook.hook_ref, {}) + p = set() + model_sd = self.model.state_dict() + for k in patches: + offset = None + function = None + if isinstance(k, str): + key = k + else: + offset = k[1] + key = k[0] + if len(k) > 2: + function = k[2] + + if key in model_sd: + p.add(k) + current_patches: List[Tuple] = current_hook_patches.get(key, []) + current_patches.append((strength_patch, patches[k], strength_model, offset, function)) + current_hook_patches[key] = current_patches + self.hook_patches[hook.hook_ref] = current_hook_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 add_hooked_patches_as_diffs(self, hook: comfy.hooks.HookWeight, patches: Dict, strength_patch=1.0, strength_model=1.0): + # NOTE: this mirrors behavior of add_patches func + current_hooked_patches: Dict[str,List] = self.hooked_patches.get(hook.hook_ref, {}) + p = set() + model_sd = self.model.state_dict() + for k in patches: + offset = None + function = None + if isinstance(k, str): + key = k + else: + offset = k[1] + key = k[0] + if len(k) > 2: + function = k[2] + + if key in model_sd: + p.add(k) + current_patches: List[Tuple] = current_hooked_patches.get(key, []) + # take difference between desired weight and existing weight to get diff + # TODO: create fix for fp8; cast to torch32 first? + current_patches.append((strength_patch, (patches[k]-comfy.utils.get_attr(self.model, key),), strength_model, offset, function)) + current_hooked_patches[key] = current_patches + self.hook_patches[hook.hook_ref] = 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_hook_patches(self, hooks: comfy.hooks.HookWeightGroup): + # combined_patches will contain weights of all relevant hooks, per key + combined_patches = {} + if hooks is not None: + for hook in hooks.hooks: + hook_patches: Dict = self.hook_patches.get(hook.hook_ref, {}) + for key in hook_patches.keys(): + current_patches: List[Tuple] = combined_patches.get(key, []) + if math.isclose(hook.strength, 1.0): + current_patches.extend(hook_patches[key]) + else: + # patches are stored as tuples: (strength_patch, (tuple_with_weights,), strength_model) + for patch in hook_patches[key]: + new_patch = List(patch) + new_patch[0] *= hook.strength + current_patches.append(Tuple(new_patch)) + combined_patches[key] = current_patches + return combined_patches + + def apply_hooks(self, hooks: comfy.hooks.HookWeightGroup): + if self.current_hooks == hooks: + return + self.patch_hooks(hooks=hooks) + + def patch_hooks(self, hooks: comfy.hooks.HookWeightGroup): + self.unpatch_hooks() + model_sd = self.model_state_dict() + # if have cached weights for hooks, use it + cached_weights = self.cached_hook_patches.get(hooks, None) + if cached_weights is not None: + for key in cached_weights: + if key not in model_sd: + print(f"WARNING cached hook could not patch. key does not exist in model: {key}") + continue + self.patch_cached_hook_weights(cached_weights=cached_weights, key=key) + else: + relevant_patches = self.get_combined_hook_patches(hooks=hooks) + for key in relevant_patches: + if key not in model_sd: + print(f"WARNING cached hook would not patch. key does not exist in model: {key}") + continue + self.patch_hook_weight_to_device(hooks=hooks, combined_patches=relevant_patches, key=key) + self.current_hooks = hooks + + def patch_cached_hook_weights(self, cached_weights: Dict, key: str): + if key not in self.hook_backup: + weight: torch.Tensor = comfy.utils.get_attr(self.model, key) + target_device = self.offload_device + if self.hook_mode == comfy.hooks.EnumHookMode.MaxSpeed: + target_device = weight.device + self.hook_backup[key] = (weight.to(device=target_device, copy=self.weight_inplace_update), weight.device) + if self.weight_inplace_update: + comfy.utils.copy_to_param(self.model, key, cached_weights[key]) + else: + comfy.utils.set_attr_param(self.model, key, cached_weights[key]) + + def clear_cached_hook_weights(self): + self.cached_hook_patches.clear() + self.current_hooks = None + + def patch_hook_weight_to_device(self, hooks: comfy.hooks.HookWeightGroup, combined_patches: dict, key: str): + if key not in combined_patches: + return + weight: torch.Tensor = comfy.utils.get_attr(self.model, key) + if key not in self.hook_backup: + target_device = self.offload_device + if self.hook_mode == comfy.hooks.EnumHookMode.MaxSpeed: + target_device = weight.device + self.hook_backup[key] = (weight.to(device=target_device, copy=self.weight_inplace_update), weight.device) + + # TODO: properly handle lowvram situations for cached hook patches + temp_weight = comfy.model_management.cast_to_device(weight, weight.device, torch.float32, copy=True) + out_weight = comfy.lora.calculate_weight(combined_patches[key], temp_weight, key).to(weight.dtype) + if self.hook_mode == comfy.hooks.EnumHookMode.MaxSpeed: + self.cached_hook_patches.setdefault(hooks, {}) + self.cached_hook_patches[hooks][key] = out_weight + if self.weight_inplace_update: + comfy.utils.copy_to_param(self.model, key, out_weight) + else: + comfy.utils.set_attr_param(self.model, key, out_weight) + + def unpatch_hooks(self) -> None: + if len(self.hook_backup) == 0: + return + keys = list(self.hook_backup.keys()) + if self.weight_inplace_update: + for k in keys: + if self.hook_mode == comfy.hooks.EnumHookMode.MaxSpeed: # does not need to be cast; device already matches + comfy.utils.copy_to_param(self.model, k, self.hook_backup[k][0]) + else: + comfy.utils.copy_to_param(self.model, k, self.hook_backup[k][0].to(device=self.hook_backup[k][1])) + else: + for k in keys: + if self.hook_mode == comfy.hooks.EnumHookMode.MaxSpeed: + comfy.utils.copy_to_param(self.model, k, self.hook_backup[k][0]) + else: + comfy.utils.copy_to_param(self.model, k, self.hook_backup[k][0].to(device=self.hook_backup[k][1])) + + self.hook_backup.clear() + self.current_hooks = None # TODO: should this be clear_cached_hooked_weights instead? + + def clean_hooks(self): + self.unpatch_hooks() + self.clear_cached_hook_weights()