mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-10-11 21:57:08 +08:00
Added hook_patches to ModelPatcher for weights (model)
This commit is contained in:
parent
7c6eb4fb29
commit
069ec7a64b
204
comfy/hooks.py
Normal file
204
comfy/hooks.py
Normal file
@ -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
|
||||||
@ -16,6 +16,7 @@
|
|||||||
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from typing import Dict, List, Tuple, Optional
|
||||||
import torch
|
import torch
|
||||||
import copy
|
import copy
|
||||||
import inspect
|
import inspect
|
||||||
@ -28,6 +29,7 @@ import comfy.utils
|
|||||||
import comfy.float
|
import comfy.float
|
||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
import comfy.lora
|
import comfy.lora
|
||||||
|
import comfy.hooks
|
||||||
from comfy.comfy_types import UnetWrapperFunction
|
from comfy.comfy_types import UnetWrapperFunction
|
||||||
|
|
||||||
def string_to_seed(data):
|
def string_to_seed(data):
|
||||||
@ -111,6 +113,13 @@ class ModelPatcher:
|
|||||||
self.weight_inplace_update = weight_inplace_update
|
self.weight_inplace_update = weight_inplace_update
|
||||||
self.patches_uuid = uuid.uuid4()
|
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'):
|
if not hasattr(self.model, 'model_loaded_weight_memory'):
|
||||||
self.model.model_loaded_weight_memory = 0
|
self.model.model_loaded_weight_memory = 0
|
||||||
|
|
||||||
@ -143,6 +152,20 @@ class ModelPatcher:
|
|||||||
n.model_options = copy.deepcopy(self.model_options)
|
n.model_options = copy.deepcopy(self.model_options)
|
||||||
n.backup = self.backup
|
n.backup = self.backup
|
||||||
n.object_patches_backup = self.object_patches_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
|
return n
|
||||||
|
|
||||||
def is_clone(self, other):
|
def is_clone(self, other):
|
||||||
@ -530,3 +553,183 @@ class ModelPatcher:
|
|||||||
def calculate_weight(self, patches, weight, key, intermediate_dtype=torch.float32):
|
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")
|
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)
|
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()
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user