support LoRAs

This commit is contained in:
kijai
2025-02-26 08:43:28 +02:00
parent f5add63007
commit fd730326c9
2 changed files with 93 additions and 4 deletions
+91 -2
View File
@@ -123,6 +123,73 @@ class WanVideoModelConfig:
# denoiser is handled by extension
self.unet_config["disable_unet_model_creation"] = True
def filter_state_dict_by_blocks(state_dict, blocks_mapping):
filtered_dict = {}
for key in state_dict:
if 'double_blocks.' in key or 'single_blocks.' in key:
block_pattern = key.split('diffusion_model.')[1].split('.', 2)[0:2]
block_key = f'{block_pattern[0]}.{block_pattern[1]}.'
if block_key in blocks_mapping:
filtered_dict[key] = state_dict[key]
return filtered_dict
def standardize_lora_key_format(lora_sd):
new_sd = {}
for k, v in lora_sd.items():
# Diffusers format
if k.startswith('transformer.'):
k = k.replace('transformer.', 'diffusion_model.')
if "img_attn.proj" in k:
k = k.replace("img_attn.proj", "img_attn_proj")
if "img_attn.qkv" in k:
k = k.replace("img_attn.qkv", "img_attn_qkv")
if "txt_attn.proj" in k:
k = k.replace("txt_attn.proj ", "txt_attn_proj")
if "txt_attn.qkv" in k:
k = k.replace("txt_attn.qkv", "txt_attn_qkv")
new_sd[k] = v
return new_sd
class WanVideoLoraSelect:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"lora": (folder_paths.get_filename_list("loras"),
{"tooltip": "LORA models are expected to be in ComfyUI/models/loras with .safetensors extension"}),
"strength": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.0001, "tooltip": "LORA strength, set to 0.0 to unmerge the LORA"}),
},
"optional": {
"prev_lora":("WANVIDLORA", {"default": None, "tooltip": "For loading multiple LoRAs"}),
"blocks":("SELECTEDBLOCKS", ),
}
}
RETURN_TYPES = ("WANVIDLORA",)
RETURN_NAMES = ("lora", )
FUNCTION = "getlorapath"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Select a LoRA model from ComfyUI/models/loras"
def getlorapath(self, lora, strength, blocks=None, prev_lora=None, fuse_lora=False):
loras_list = []
lora = {
"path": folder_paths.get_full_path("loras", lora),
"strength": strength,
"name": lora.split(".")[0],
"blocks": blocks
}
if prev_lora is not None:
loras_list.extend(prev_lora)
loras_list.append(lora)
return (loras_list,)
#region Model loading
class WanVideoModelLoader:
@@ -145,6 +212,8 @@ class WanVideoModelLoader:
], {"default": "sdpa"}),
"compile_args": ("WANCOMPILEARGS", ),
"block_swap_args": ("BLOCKSWAPARGS", ),
"lora": ("WANVIDLORA", {"default": None}),
}
}
@@ -154,7 +223,7 @@ class WanVideoModelLoader:
CATEGORY = "WanVideoWrapper"
def loadmodel(self, model, base_precision, load_device, quantization,
compile_args=None, attention_mode="sdpa", block_swap_args=None):
compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None):
transformer = None
mm.unload_all_models()
mm.soft_empty_cache()
@@ -230,6 +299,24 @@ class WanVideoModelLoader:
comfy_model.load_device = transformer_load_device
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
if lora is not None:
from comfy.sd import load_lora_for_models
for l in lora:
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
lora_path = l["path"]
lora_strength = l["strength"]
lora_sd = load_torch_file(lora_path, safe_load=True)
lora_sd = standardize_lora_key_format(lora_sd)
if l["blocks"]:
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"])
#for k in lora_sd.keys():
# print(k)
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
comfy.model_management.load_models_gpu([patcher])
del sd
gc.collect()
mm.soft_empty_cache()
@@ -484,8 +571,8 @@ class LoadWanVideoClipTextEncoder:
model_path = folder_paths.get_full_path("text_encoders", model_name)
sd = load_torch_file(model_path, safe_load=True)
clip_model = CLIPModel(dtype=dtype, device=text_encoder_load_device, state_dict=sd, tokenizer_path=tokenizer_path)
del sd
return (clip_model,)
@@ -1077,6 +1164,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoTorchCompileSettings": WanVideoTorchCompileSettings,
"WanVideoLatentPreview": WanVideoLatentPreview,
"WanVideoEmptyEmbeds": WanVideoEmptyEmbeds,
"WanVideoLoraSelect": WanVideoLoraSelect,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoSampler": "WanVideo Sampler",
@@ -1093,4 +1181,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoTorchCompileSettings": "WanVideo Torch Compile Settings",
"WanVideoLatentPreview": "WanVideo Latent Preview",
"WanVideoEmptyEmbeds": "WanVideo Empty Embeds",
"WanVideoLoraSelect": "WanVideo Lora Select",
}
+2 -2
View File
@@ -497,10 +497,10 @@ class WanModel(ModelMixin, ConfigMixin):
for b, block in enumerate(self.blocks):
if b > self.blocks_to_swap:
print(f"Moving double_block {b} to main device")
print(f"Moving transformer block {b} to main device")
block.to(self.main_device)
else:
print(f"Moving double_block {b} to offload_device")
print(f"Moving transformer block {b} to offload_device")
block.to(self.offload_device)
def forward(