Made CLIP LoRA hooks work as intended, proper way to force ModelPatcherAndInjector to cause model reload when hooks present
This commit is contained in:
+209
-39
@@ -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 = {}
|
||||
|
||||
@@ -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] 🎭🅐🅓",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user