From 7bc45daaf2667b655b7b4e58cbf321a5c9809006 Mon Sep 17 00:00:00 2001
From: kijai <40791699+kijai@users.noreply.github.com>
Date: Tue, 2 Dec 2025 14:50:26 +0200
Subject: [PATCH] Cleanup whitespaces
---
nodes_model_loading.py | 158 ++++++++++++++++++++---------------------
1 file changed, 79 insertions(+), 79 deletions(-)
diff --git a/nodes_model_loading.py b/nodes_model_loading.py
index 77d7056..250177c 100644
--- a/nodes_model_loading.py
+++ b/nodes_model_loading.py
@@ -93,7 +93,7 @@ def filter_state_dict_by_blocks(state_dict, blocks_mapping, layer_filter=[]):
for key in state_dict:
if not any(filter_str in key for filter_str in layer_filters):
if 'blocks.' in key:
-
+
block_pattern = key.split('diffusion_model.')[1].split('.', 2)[0:2]
block_key = f'{block_pattern[0]}.{block_pattern[1]}.'
@@ -101,7 +101,7 @@ def filter_state_dict_by_blocks(state_dict, blocks_mapping, layer_filter=[]):
filtered_dict[key] = state_dict[key]
else:
filtered_dict[key] = state_dict[key]
-
+
for key in filtered_dict:
print(key)
@@ -139,19 +139,19 @@ def standardize_lora_key_format(lora_sd):
parts = k.split('.')
main_part = parts[0] # e.g. lora_unet__blocks_0_cross_attn_k
weight_type = '.'.join(parts[1:]) if len(parts) > 1 else None # e.g. lora_down.weight
-
+
# Process the main part - convert from underscore to dot format
if 'blocks_' in main_part:
# Extract components
components = main_part[len('lora_unet__'):].split('_')
-
+
# Start with diffusion_model
new_key = "diffusion_model"
-
+
# Add blocks.N
if components[0] == 'blocks':
new_key += f".blocks.{components[1]}"
-
+
# Handle different module types
idx = 2
if idx < len(components):
@@ -164,19 +164,19 @@ def standardize_lora_key_format(lora_sd):
elif components[idx] == 'ffn':
new_key += ".ffn"
idx += 1
-
+
# Add the component (k, q, v, o) and handle img suffix
if idx < len(components):
component = components[idx]
idx += 1
-
+
# Check for img suffix
if idx < len(components) and components[idx] == 'img':
component += '_img'
idx += 1
-
+
new_key += f".{component}"
-
+
# Handle weight type - this is the critical fix
if weight_type:
if weight_type == 'alpha':
@@ -191,12 +191,12 @@ def standardize_lora_key_format(lora_sd):
# Add .weight suffix if missing
if not new_key.endswith('.weight'):
new_key += '.weight'
-
+
k = new_key
else:
# For other lora_unet__ formats (head, embeddings, etc.)
new_key = main_part.replace('lora_unet__', 'diffusion_model.')
-
+
# Fix specific component naming patterns
new_key = new_key.replace('_self_attn', '.self_attn')
new_key = new_key.replace('_cross_attn', '.cross_attn')
@@ -207,7 +207,7 @@ def standardize_lora_key_format(lora_sd):
new_key = new_key.replace('text_embedding', 'text.embedding')
new_key = new_key.replace('time_embedding', 'time.embedding')
new_key = new_key.replace('time_projection', 'time.projection')
-
+
# Replace remaining underscores with dots, carefully
parts = new_key.split('.')
final_parts = []
@@ -217,7 +217,7 @@ def standardize_lora_key_format(lora_sd):
else:
final_parts.append(part.replace('_', '.'))
new_key = '.'.join(final_parts)
-
+
# Handle weight type
if weight_type:
if weight_type == 'alpha':
@@ -230,9 +230,9 @@ def standardize_lora_key_format(lora_sd):
new_key += f'.{weight_type}'
if not new_key.endswith('.weight'):
new_key += '.weight'
-
+
k = new_key
-
+
# Handle special embedded components
special_components = {
'time.projection': 'time_projection',
@@ -247,7 +247,7 @@ def standardize_lora_key_format(lora_sd):
# Fix diffusion.model -> diffusion_model
if k.startswith('diffusion.model.'):
k = k.replace('diffusion.model.', 'diffusion_model.')
-
+
# Finetrainer format
if '.attn1.' in k:
k = k.replace('.attn1.', '.cross_attn.')
@@ -261,7 +261,7 @@ def standardize_lora_key_format(lora_sd):
k = k.replace('.to_q.', '.q.')
k = k.replace('.to_v.', '.v.')
k = k.replace('.to_out.0.', '.o.')
-
+
if "img_attn.proj" in k:
k = k.replace("img_attn.proj", "img_attn_proj")
if "img_attn.qkv" in k:
@@ -429,7 +429,7 @@ class WanVideoLoraSelect:
f"
| Metadata |
"
f"{metadata_rows}"
f""
- f"",
+ f"",
unique_id
)
except Exception as e:
@@ -450,7 +450,7 @@ class WanVideoLoraSelect:
loras_list.append(lora)
return (loras_list,)
-
+
class WanVideoLoraSelectByName(WanVideoLoraSelect):
@classmethod
def INPUT_TYPES(s):
@@ -469,7 +469,7 @@ class WanVideoLoraSelectByName(WanVideoLoraSelect):
"unique_id": "UNIQUE_ID",
},
}
-
+
def getlorapath(self, lora_name, strength, unique_id, blocks={}, prev_lora=None, low_mem_load=False, merge_loras=True):
lora_list = folder_paths.get_filename_list("loras")
lora_path = "none"
@@ -480,7 +480,7 @@ class WanVideoLoraSelectByName(WanVideoLoraSelect):
return super().getlorapath(
lora_path, strength, unique_id, blocks=blocks, prev_lora=prev_lora, low_mem_load=low_mem_load, merge_loras=merge_loras
)
-
+
class WanVideoLoraSelectMulti:
@classmethod
def INPUT_TYPES(s):
@@ -514,8 +514,8 @@ class WanVideoLoraSelectMulti:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Select a LoRA model from ComfyUI/models/loras"
- def getlorapath(self, lora_0, strength_0, lora_1, strength_1, lora_2, strength_2,
- lora_3, strength_3, lora_4, strength_4, blocks={}, prev_lora=None,
+ def getlorapath(self, lora_0, strength_0, lora_1, strength_1, lora_2, strength_2,
+ lora_3, strength_3, lora_4, strength_4, blocks={}, prev_lora=None,
low_mem_load=False, merge_loras=True):
if not merge_loras:
low_mem_load = False # Unmerged LoRAs don't need low_mem_load
@@ -543,7 +543,7 @@ class WanVideoLoraSelectMulti:
if len(loras_list) == 0:
return None,
return (loras_list,)
-
+
class WanVideoVACEModelSelect:
@classmethod
def INPUT_TYPES(s):
@@ -562,7 +562,7 @@ class WanVideoVACEModelSelect:
def getvacepath(self, vace_model):
vace_model = [{"path": folder_paths.get_full_path_or_raise("diffusion_models", vace_model)}]
return (vace_model,)
-
+
class WanVideoExtraModelSelect:
@classmethod
def INPUT_TYPES(s):
@@ -696,9 +696,9 @@ def load_lora_for_models_mod(model, lora, strength_model):
key_map = {}
if model is not None:
key_map = model_lora_keys_unet(model.model, key_map)
-
+
loaded = comfy.lora.load_lora(lora, key_map)
-
+
new_modelpatcher = model.clone()
k = add_patches(new_modelpatcher, loaded, strength_model)
k = set(k)
@@ -712,7 +712,7 @@ class WanVideoSetLoRAs:
@classmethod
def INPUT_TYPES(s):
return {
- "required":
+ "required":
{
"model": ("WANVIDEOMODEL", ),
},
@@ -731,15 +731,15 @@ class WanVideoSetLoRAs:
def setlora(self, model, lora=None):
if lora is None:
return (model,)
-
+
patcher = model.clone()
-
+
merge_loras = False
for l in lora:
merge_loras = l.get("merge_loras", True)
if merge_loras is True:
raise ValueError("Set LoRA node does not use low_mem_load and can't merge LoRAs, disable 'merge_loras' in the LoRA select node.")
-
+
patcher.model_options['transformer_options']["lora_scheduling_enabled"] = False
for l in lora:
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
@@ -763,12 +763,12 @@ class WanVideoSetLoRAs:
# Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys
if not any('img' in k for k in model.model.diffusion_model.state_dict().keys()):
lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k}
-
+
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
raise NotImplementedError("Control LoRA patching is not implemented in this node.")
patcher = load_lora_for_models_mod(patcher, lora_sd, lora_strength)
-
+
del lora_sd
return (patcher,)
@@ -784,9 +784,9 @@ def rename_fuser_block(name):
new_name = name.replace(f"face_adapter.fuser_blocks.{fuser_block_num}.", f"blocks.{main_block_num}.fuser_block.")
return new_name
-def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
+def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None, compile_args=None):
- params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding",
+ params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding",
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob", "face_encoder", "fuser_block"}
param_count = sum(1 for _ in transformer.named_parameters())
pbar = ProgressBar(param_count)
@@ -831,7 +831,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
elif vace_block_idx is not None:
if vace_block_idx >= len(transformer.vace_blocks) - block_swap_args.get("vace_blocks_to_swap", 0):
load_device = offload_device
-
+
is_gguf_quant = tensor.tensor_type not in [GGMLQuantizationType.F32, GGMLQuantizationType.F16]
weights = torch.from_numpy(tensor.data.copy()).to(load_device)
sd[name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights
@@ -869,7 +869,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
# GGUF: skip GGUFParameter params
if gguf and isinstance(param, GGUFParameter):
continue
-
+
key = name.replace("_orig_mod.", "")
value=sd[key]
keep_fp32 = ["patch_embedding", "motion_encoder", "condition_embedding"]
@@ -914,7 +914,7 @@ def patch_control_lora(transformer, device):
in_cls = transformer.patch_embedding.__class__ # nn.Conv3d
old_in_dim = transformer.in_dim # 16
new_in_dim = 32
-
+
new_in = in_cls(
new_in_dim,
transformer.patch_embedding.out_channels,
@@ -922,18 +922,18 @@ def patch_control_lora(transformer, device):
transformer.patch_embedding.stride,
transformer.patch_embedding.padding,
).to(device=device, dtype=torch.float32)
-
+
new_in.weight.zero_()
new_in.bias.zero_()
-
+
new_in.weight[:, :old_in_dim].copy_(transformer.patch_embedding.weight)
new_in.bias.copy_(transformer.patch_embedding.bias)
-
+
transformer.patch_embedding = new_in
transformer.expanded_patch_embedding = new_in
def patch_stand_in_lora(transformer, lora_sd, transformer_load_device, base_dtype, lora_strength):
- if "diffusion_model.blocks.0.self_attn.q_loras.down.weight" in lora_sd:
+ if "diffusion_model.blocks.0.self_attn.q_loras.down.weight" in lora_sd:
log.info("Stand-In LoRA detected")
for block in transformer.blocks:
block.self_attn.q_loras = LoRALinearLayer(transformer.dim, transformer.dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength)
@@ -975,7 +975,7 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
# Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys
#if not any('img' in k for k in sd.keys()):
# lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k}
-
+
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
control_lora = True
#stand-in LoRA patch
@@ -984,7 +984,7 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False):
# normal LoRA patch
else:
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
-
+
del lora_sd
return patcher, control_lora, unianimate_sd
@@ -997,7 +997,7 @@ class WanVideoModelLoader:
"model": (folder_paths.get_filename_list("unet_gguf") + folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}),
"base_precision": (["fp32", "bf16", "fp16", "fp16_fast"], {"default": "bf16"}),
- "quantization": (["disabled", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e4m3fn_scaled", "fp8_e4m3fn_scaled_fast", "fp8_e5m2", "fp8_e5m2_fast", "fp8_e5m2_scaled", "fp8_e5m2_scaled_fast"], {"default": "disabled",
+ "quantization": (["disabled", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e4m3fn_scaled", "fp8_e4m3fn_scaled_fast", "fp8_e5m2", "fp8_e5m2_fast", "fp8_e5m2_scaled", "fp8_e5m2_scaled_fast"], {"default": "disabled",
"tooltip": "Optional quantization method, 'disabled' acts as autoselect based by weights. Scaled modes only work with matching weights, _fast modes (fp8 matmul) require CUDA compute capability >= 8.9 (NVIDIA 4000 series and up), e4m3fn generally can not be torch.compiled on compute capability < 8.9 (3000 series and under)"}),
"load_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
},
@@ -1064,7 +1064,7 @@ class WanVideoModelLoader:
transformer_load_device = offload_device
base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision]
-
+
if base_precision == "fp16_fast":
if hasattr(torch.backends.cuda.matmul, "allow_fp16_accumulation"):
torch.backends.cuda.matmul.allow_fp16_accumulation = True
@@ -1076,7 +1076,7 @@ class WanVideoModelLoader:
torch.backends.cuda.matmul.allow_fp16_accumulation = False
except:
pass
-
+
model_path = folder_paths.get_full_path_or_raise("diffusion_models", model)
@@ -1095,7 +1095,7 @@ class WanVideoModelLoader:
sd = {key.replace("video_model.", "", 1).replace("modulation.modulation", "modulation"): value for key, value in sd.items()}
if any(key.startswith("audio_model.") for key in sd.keys()) and any(key.startswith("blocks.") for key in sd.keys()):
extra_audio_model = True
-
+
is_wananimate = "pose_patch_embedding.weight" in sd
# rename WanAnimate face fuser block keys to insert into main blocks instead
@@ -1305,20 +1305,20 @@ class WanVideoModelLoader:
model_variant = "i2v_720"
elif model_type == "t2v":
model_variant = "14B"
-
+
if dim == 1536:
model_variant = "1_3B"
if dim == 3072:
log.info(f"5B model detected, no Teacache or MagCache coefficients available, consider using EasyCache for this model")
-
+
if "high" in model.lower() or "low" in model.lower():
if "i2v" in model.lower():
model_variant = "i2v_14B_2.2"
else:
model_variant = "14B_2.2"
-
+
log.info(f"Model variant detected: {model_variant}")
-
+
TRANSFORMER_CONFIG= {
"dim": dim,
"in_features": in_features,
@@ -1367,7 +1367,7 @@ class WanVideoModelLoader:
log.info("Ovi extra audio model detected, initializing...")
TRANSFORMER_CONFIG.update({
"patch_size": [1],
- "in_dim": 20,
+ "in_dim": 20,
"out_dim": 20,
})
@@ -1437,7 +1437,7 @@ class WanVideoModelLoader:
# init audio module
from .multitalk.multitalk import SingleStreamMultiAttention
from .wanvideo.modules.model import WanLayerNorm
-
+
for block in transformer.blocks:
with init_empty_weights():
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
@@ -1461,10 +1461,10 @@ class WanVideoModelLoader:
del extra_reader
else:
extra_sd_temp = load_torch_file(extra_model_path, device=transformer_load_device, safe_load=True)
-
+
for k, v in extra_sd_temp.items():
extra_sd[k.replace("audio_proj.", "multitalk_audio_proj.")] = v
-
+
sd.update(extra_sd)
del extra_sd
@@ -1486,7 +1486,7 @@ class WanVideoModelLoader:
transformer.add_conv_in = torch.nn.Conv3d(add_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
transformer.add_proj = zero_module(torch.nn.Linear(inner_dim, inner_dim))
transformer.attn_conv_in = torch.nn.Conv3d(attn_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
-
+
# Bindweave text_projection
if "text_projection.0.weight" in sd:
log.info("Bindweave model detected, adding text_projection to the model")
@@ -1519,7 +1519,7 @@ class WanVideoModelLoader:
comfy_model.load_device = transformer_load_device
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
patcher.model.is_patched = False
-
+
scale_weights = {}
if "fp8" in quantization:
for k, v in sd.items():
@@ -1532,7 +1532,7 @@ class WanVideoModelLoader:
weight_dtype = torch.float8_e5m2
else:
weight_dtype = base_dtype
-
+
params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
control_lora = False
@@ -1546,19 +1546,19 @@ class WanVideoModelLoader:
log.info("Merging UniAnimate weights to the model...")
sd.update(unianimate_sd)
del unianimate_sd
-
+
if not gguf:
if lora is not None and merge_loras:
if not lora_low_mem_load:
load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device)
-
+
if control_lora:
patch_control_lora(patcher.model.diffusion_model, device)
- patcher.model.is_patched = True
-
+ patcher.model.is_patched = True
+
log.info("Merging LoRA to the model...")
patcher = apply_lora(
- patcher, device, transformer_load_device, params_to_keep=params_to_keep, dtype=weight_dtype, base_dtype=base_dtype, state_dict=sd,
+ patcher, device, transformer_load_device, params_to_keep=params_to_keep, dtype=weight_dtype, base_dtype=base_dtype, state_dict=sd,
low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights)
if not control_lora:
scale_weights.clear()
@@ -1588,7 +1588,7 @@ class WanVideoModelLoader:
offload_params = int(total_params_in_model * offload_percent)
params_to_keep = total_params_in_model - offload_params
log.info(f"Selected params to offload: {offload_params}")
-
+
enable_vram_management(
patcher.model.diffusion_model,
module_map = {
@@ -1647,7 +1647,7 @@ class WanVideoModelLoader:
if model._model() == patcher:
mm.current_loaded_models.remove(model)
return (patcher,)
-
+
# class WanVideoSaveModel:
# @classmethod
# def INPUT_TYPES(s):
@@ -1679,7 +1679,7 @@ class WanVideoModelLoader:
# os.makedirs(os.path.dirname(output_path), exist_ok=True)
# save_file(model_sd, output_path)
# return ()
-
+
#region load VAE
class WanVideoVAELoader:
@@ -1742,7 +1742,7 @@ class WanVideoTinyVAELoader:
"model_name": (folder_paths.get_filename_list("vae_approx"), {"tooltip": "These models are loaded from 'ComfyUI/models/vae_approx'"}),
},
"optional": {
- "precision": (["fp16", "fp32", "bf16"], {"default": "fp16"}),
+ "precision": (["fp16", "fp32", "bf16"], {"default": "fp16"}),
"parallel": ("BOOLEAN", {"default": False, "tooltip": "uses more memory but is faster"}),
}
}
@@ -1759,7 +1759,7 @@ class WanVideoTinyVAELoader:
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_path = folder_paths.get_full_path_or_raise("vae_approx", model_name)
vae_sd = load_torch_file(model_path, safe_load=True)
-
+
vae = TAEHV(vae_sd, parallel=parallel, dtype=dtype)
vae.to(device=offload_device, dtype=dtype)
@@ -1804,7 +1804,7 @@ class LoadWanVideoT5TextEncoder:
if v.dtype == torch.float8_e4m3fn:
quantization = "fp8_e4m3fn"
break
-
+
if "token_embedding.weight" not in sd and "shared.weight" not in sd:
raise ValueError("Invalid T5 text encoder model, this node expects the 'umt5-xxl' model")
if "scaled_fp8" in sd:
@@ -1814,13 +1814,13 @@ class LoadWanVideoT5TextEncoder:
if "shared.weight" in sd:
log.info("Converting T5 text encoder model to the expected format...")
converted_sd = {}
-
+
for key, value in sd.items():
# Handle encoder block patterns
if key.startswith('encoder.block.'):
parts = key.split('.')
block_num = parts[2]
-
+
# Self-attention components
if 'layer.0.SelfAttention' in key:
if key.endswith('.k.weight'):
@@ -1835,13 +1835,13 @@ class LoadWanVideoT5TextEncoder:
new_key = f"blocks.{block_num}.pos_embedding.embedding.weight"
else:
new_key = key
-
+
# Layer norms
elif 'layer.0.layer_norm' in key:
new_key = f"blocks.{block_num}.norm1.weight"
elif 'layer.1.layer_norm' in key:
new_key = f"blocks.{block_num}.norm2.weight"
-
+
# Feed-forward components
elif 'layer.1.DenseReluDense' in key:
if 'wi_0' in key:
@@ -1876,9 +1876,9 @@ class LoadWanVideoT5TextEncoder:
"dtype": dtype,
"name": model_name,
}
-
+
return (text_encoder,)
-
+
class LoadWanVideoClipTextEncoder:
@classmethod
def INPUT_TYPES(s):
@@ -1894,7 +1894,7 @@ class LoadWanVideoClipTextEncoder:
}
}
- RETURN_TYPES = ("CLIP_VISION",)
+ RETURN_TYPES = ("CLIP_VISION",)
RETURN_NAMES = ("wan_clip_vision", )
FUNCTION = "loadmodel"
CATEGORY = "WanVideoWrapper"
@@ -1916,7 +1916,7 @@ class LoadWanVideoClipTextEncoder:
clip_model = CLIPModel(dtype=dtype, device=device, state_dict=sd)
clip_model.model.to(text_encoder_load_device)
del sd
-
+
return (clip_model,)
NODE_CLASS_MAPPINGS = {