Made CLIP LoRA hooks work as intended, proper way to force ModelPatcherAndInjector to cause model reload when hooks present

This commit is contained in:
Jedrzej Kosinski
2024-03-29 06:22:30 -05:00
parent d45bc705a4
commit 1e8fba9384
4 changed files with 213 additions and 43 deletions
+209 -39
View File
@@ -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 = {}
+2 -2
View File
@@ -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] 🎭🅐🅓",
-2
View File
@@ -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
+2
View File
@@ -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: