From 7fc5ec305db730f3f0f2abe45501a0963cd7cf59 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 16 Aug 2025 16:03:15 +0300 Subject: [PATCH] Refactor model loading --- custom_linear.py | 2 +- gguf/gguf.py | 16 +- nodes.py | 85 +++++--- nodes_model_loading.py | 406 ++++++++++++++++++++++---------------- utils.py | 4 +- wanvideo/modules/model.py | 55 ++++-- wanvideo/wan_video_vae.py | 2 +- 7 files changed, 343 insertions(+), 227 deletions(-) diff --git a/custom_linear.py b/custom_linear.py index ae5e81d..391502b 100644 --- a/custom_linear.py +++ b/custom_linear.py @@ -12,7 +12,7 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s module_prefix = prefix + name + "." _replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights) - if isinstance(module, nn.Linear): + if isinstance(module, nn.Linear) and "loras" not in module_prefix: in_features = state_dict[module_prefix + "weight"].shape[1] out_features = state_dict[module_prefix + "weight"].shape[0] if scale_weights is not None: diff --git a/gguf/gguf.py b/gguf/gguf.py index 44948ad..d21b817 100644 --- a/gguf/gguf.py +++ b/gguf/gguf.py @@ -1,13 +1,25 @@ import torch import torch.nn as nn +import numpy as np from diffusers.quantizers.gguf.utils import GGUFParameter, dequantize_gguf_tensor +import gguf from diffusers.utils import is_accelerate_available from contextlib import nullcontext - +from ..utils import log if is_accelerate_available(): - import accelerate from accelerate import init_empty_weights +def load_gguf(model_path): + from gguf import GGUFReader + reader = GGUFReader(model_path) + parsed_parameters = {} + for tensor in reader.tensors: + # if the tensor is a torch supported dtype do not use GGUFParameter + is_gguf_quant = tensor.tensor_type not in [gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16] + meta_tensor = torch.empty(tensor.data.shape, dtype=torch.from_numpy(np.empty(0, dtype=tensor.data.dtype)).dtype, device='meta') + parsed_parameters[tensor.name] = GGUFParameter(meta_tensor, quant_type=tensor.tensor_type) if is_gguf_quant else meta_tensor + return parsed_parameters, reader + #based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None): def _should_convert_to_gguf(state_dict, prefix): diff --git a/nodes.py b/nodes.py index 9d2e709..b2a63bb 100644 --- a/nodes.py +++ b/nodes.py @@ -42,6 +42,17 @@ def offload_transformer(transformer): transformer.magcache_state.clear_all() transformer.easycache_state.clear_all() transformer.to(offload_device) + # for name, param in transformer.named_parameters(): + # module = transformer + # subnames = name.split('.') + # for subname in subnames[:-1]: + # module = getattr(module, subname) + # attr_name = subnames[-1] + # if param.data.is_floating_point(): + # meta_param = torch.nn.Parameter(torch.empty_like(param.data, device='meta'), requires_grad=False) + # setattr(module, attr_name, meta_param) + # else: + # pass mm.soft_empty_cache() gc.collect() @@ -348,8 +359,11 @@ class WanVideoTextEncode: raise ValueError("No cached text embeds found for prompts, please provide a T5 encoder.") if model_to_offload is not None and device == "gpu": - log.info(f"Moving video model to {offload_device}") - model_to_offload.model.to(offload_device) + try: + log.info(f"Moving video model to {offload_device}") + model_to_offload.model.to(offload_device) + except: + pass encoder = t5["model"] dtype = t5["dtype"] @@ -1080,7 +1094,7 @@ class WanVideoPhantomEmbeds: log.info(f"Phantom latents shape: {samples.shape}") - target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1 + T, + target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1, H * 8 // VAE_STRIDE[1], W * 8 // VAE_STRIDE[2]) @@ -1583,20 +1597,37 @@ class WanVideoSampler: model = model.model transformer = model.diffusion_model - dtype = model["dtype"] + dtype = model["base_dtype"] + weight_dtype = model["weight_dtype"] fp8_matmul = model["fp8_matmul"] - gguf = model["gguf"] + gguf_reader = model["gguf_reader"] control_lora = model["control_lora"] transformer_options = patcher.model_options.get("transformer_options", None) merge_loras = transformer_options["merge_loras"] + block_swap_args = transformer_options.get("block_swap_args", None) + if block_swap_args is not None: + transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False) + transformer.blocks_to_swap = block_swap_args.get("blocks_to_swap", 0) + transformer.vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", 0) + transformer.prefetch_blocks = block_swap_args.get("prefetch_blocks", 0) + transformer.block_swap_debug = block_swap_args.get("block_swap_debug", False) + transformer.offload_img_emb = block_swap_args.get("offload_img_emb", False) + transformer.offload_txt_emb = block_swap_args.get("offload_txt_emb", False) + is_5b = transformer.out_dim == 48 vae_upscale_factor = 16 if is_5b else 8 patch_linear = transformer_options.get("patch_linear", False) + from .nodes_model_loading import load_weights, load_weights_gguf + weights_assigned = False + if not merge_loras and gguf_reader is None: + load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, dtype, device, block_swap_args=block_swap_args) + weights_assigned = True - if gguf: + if gguf_reader is not None: + load_weights_gguf(transformer, gguf_reader, patcher.model["sd"], dtype, device, patcher) set_lora_params_gguf(transformer, patcher.patches) elif len(patcher.patches) != 0 and patch_linear: log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model") @@ -1867,8 +1898,6 @@ class WanVideoSampler: phantom_cfg_scale = [phantom_cfg_scale] * (steps +1) phantom_start_percent = image_embeds.get("phantom_start_percent", 0.0) phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0) - if phantom_latents is not None: - phantom_latents = phantom_latents.to(device) latent_video_length = noise.shape[1] @@ -2153,11 +2182,8 @@ class WanVideoSampler: mm.unload_all_models() mm.soft_empty_cache() gc.collect() - - if transformer_options is not None: - block_swap_args = transformer_options.get("block_swap_args", None) - if block_swap_args is not None: + if block_swap_args is not None and not weights_assigned: transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False) for name, param in transformer.named_parameters(): if "block" not in name: @@ -2187,8 +2213,8 @@ class WanVideoSampler: block.modulation = torch.nn.Parameter(block.modulation.to(device)) transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device)) - elif model["manual_offloading"]: - transformer.to(device) + #elif model["manual_offloading"]: + # transformer.to(device) # Initialize Cache if enabled transformer.enable_teacache = transformer.enable_magcache = transformer.enable_easycache = False @@ -2338,18 +2364,14 @@ class WanVideoSampler: z = torch.cat([z, recam_latents.to(z)], dim=1) use_phantom = False + phantom_ref = None if phantom_latents is not None: if (phantom_start_percent <= current_step_percentage <= phantom_end_percent) or \ (phantom_end_percent > 0 and idx == 0 and current_step_percentage >= phantom_start_percent): - - z_pos = torch.cat([z[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1) - z_phantom_img = torch.cat([z[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1) - z_neg = torch.cat([z[:,:-phantom_latents.shape[1]], torch.zeros_like(phantom_latents).to(z)], dim=1) + phantom_ref = phantom_latents.to(z) use_phantom = True if cache_state is not None and len(cache_state) != 3: cache_state.append(None) - if not use_phantom: - z_pos = z_neg = z if controlnet_latents is not None: if (controlnet_start <= current_step_percentage < controlnet_end): @@ -2375,9 +2397,9 @@ class WanVideoSampler: if minimax_latents is not None: if context_window is not None: - z_pos = z_neg = torch.cat([z, minimax_latents[:, context_window], minimax_mask_latents[:, context_window]], dim=0) + z = torch.cat([z, minimax_latents[:, context_window], minimax_mask_latents[:, context_window]], dim=0) else: - z_pos = z_neg = torch.cat([z, minimax_latents, minimax_mask_latents], dim=0) + z = torch.cat([z, minimax_latents, minimax_mask_latents], dim=0) if not multitalk_sampling and multitalk_audio_embedding is not None: audio_embedding = multitalk_audio_embedding @@ -2442,6 +2464,7 @@ class WanVideoSampler: "inner_t": [shot_len] if shot_len else None, "standin_input": standin_input, "fantasy_portrait_input": fantasy_portrait_input, + "phantom_ref": phantom_ref } batch_size = 1 @@ -2456,7 +2479,7 @@ class WanVideoSampler: if not batched_cfg: #cond noise_pred_cond, cache_state_cond = transformer( - [z_pos], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None, + [z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None, clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage, pred_id=cache_state[0] if cache_state else None, vace_data=vace_data, attn_cond=attn_cond, @@ -2477,7 +2500,7 @@ class WanVideoSampler: if not math.isclose(audio_cfg_scale[idx], 1.0): base_params['audio_proj'] = None noise_pred_uncond, cache_state_uncond = transformer( - [z_neg], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, + [z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, y=[image_cond_input] if image_cond_input is not None else None, is_uncond=True, current_step_percentage=current_step_percentage, pred_id=cache_state[1] if cache_state else None, @@ -2488,7 +2511,7 @@ class WanVideoSampler: #phantom if use_phantom and not math.isclose(phantom_cfg_scale[idx], 1.0): noise_pred_phantom, cache_state_phantom = transformer( - [z_phantom_img], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, + [z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, y=[image_cond_input] if image_cond_input is not None else None, is_uncond=True, current_step_percentage=current_step_percentage, pred_id=cache_state[2] if cache_state else None, @@ -2506,7 +2529,7 @@ class WanVideoSampler: cache_state.append(None) base_params['audio_proj'] = None noise_pred_no_audio, cache_state_audio = transformer( - [z_pos], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None, + [z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None, clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage, pred_id=cache_state[2] if cache_state else None, vace_data=vace_data, @@ -2525,7 +2548,7 @@ class WanVideoSampler: cache_state.append(None) base_params['multitalk_audio'] = torch.zeros_like(multitalk_audio_input)[-1:] noise_pred_no_audio, cache_state_audio = transformer( - [z_pos], context=negative_embeds, y=[image_cond_input] if image_cond_input is not None else None, + [z], context=negative_embeds, y=[image_cond_input] if image_cond_input is not None else None, clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage, pred_id=cache_state[2] if cache_state else None, vace_data=vace_data, @@ -3283,8 +3306,8 @@ class WanVideoSampler: if callback is not None: if recammaster is not None: callback_latent = (latent_model_input[:, :orig_noise_len].to(device) - noise_pred[:, :orig_noise_len].to(device) * t.to(device) / 1000).detach() - elif phantom_latents is not None: - callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach() + #elif phantom_latents is not None: + # callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach() else: callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach() callback(idx, callback_latent.permute(1,0,2,3), None, len(timesteps)) @@ -3303,8 +3326,8 @@ class WanVideoSampler: offload_transformer(transformer) raise e - if phantom_latents is not None: - latent = latent[:,:-phantom_latents.shape[1]] + #if phantom_latents is not None: + # latent = latent[:,:-phantom_latents.shape[1]] if cache_args is not None: cache_report(transformer, cache_args) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 9cc68f3..a0c5b59 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -708,6 +708,178 @@ class WanVideoSetLoRAs: return (patcher,) +def load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_device, block_swap_args=None): + print("block_swap_args", block_swap_args) + params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv"} + + log.info("Using accelerate to load and assign model weights to device...") + param_count = sum(1 for _ in transformer.named_parameters()) + pbar = ProgressBar(param_count) + cnt = 0 + for name, param in tqdm(transformer.named_parameters(), + desc=f"Loading transformer parameters to {transformer_load_device}", + total=param_count, + leave=True): + block_idx = vace_block_idx = None + if "vace_blocks." in name: + try: + vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0]) + except Exception: + vace_block_idx = None + elif "blocks." in name: + try: + block_idx = int(name.split("blocks.")[1].split(".")[0]) + except Exception: + block_idx = None + print("block_idx:", block_idx) + print("vace_block_idx:", vace_block_idx) + if "loras" in name: + continue + dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else weight_dtype + dtype_to_use = weight_dtype if sd[name].dtype == weight_dtype else dtype_to_use + if "modulation" in name or "norm" in name or "bias" in name: + dtype_to_use = base_dtype + if "patch_embedding" in name: + dtype_to_use = torch.float32 + + load_device = device + if block_swap_args is not None: + if block_idx is not None: + if block_idx >= len(transformer.blocks) - block_swap_args["blocks_to_swap"]: + load_device = offload_device + elif vace_block_idx is not None: + if vace_block_idx >= len(transformer.vace_blocks) - block_swap_args["vace_blocks_to_swap"]: + load_device = offload_device + + print("load_device:", load_device) + + set_module_tensor_to_device(transformer, name, device=load_device, dtype=dtype_to_use, value=sd[name]) + cnt += 1 + if cnt % 100 == 0: + pbar.update(100) + for name, param in transformer.named_parameters(): + print(name, param.dtype, param.device, param.shape) + pbar.update_absolute(param_count) + pbar.update_absolute(0) + +def load_weights_gguf(transformer, reader, sd, base_dtype, transformer_load_device, patcher): + from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter + import gguf + log.info("Using GGUF to load and assign model weights to device...") + param_count = sum(1 for _ in transformer.named_parameters()) + + sd = {} + for tensor in reader.tensors: + # if the tensor is a torch supported dtype do not use GGUFParameter + is_gguf_quant = tensor.tensor_type not in [gguf.GGMLQuantizationType.F32, gguf.GGMLQuantizationType.F16] + weights = torch.from_numpy(tensor.data.copy()).to(transformer_load_device) + sd[tensor.name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights + + patcher.model.diffusion_model = _replace_with_gguf_linear(patcher.model.diffusion_model, base_dtype, sd, patches=patcher.patches) + pbar = ProgressBar(param_count) + cnt = 0 + for name, param in tqdm(patcher.model.diffusion_model.named_parameters(), + desc=f"Loading transformer parameters to {transformer_load_device}", + total=param_count, + leave=True): + if "loras" in name: + continue + #print(name, param.dtype, param.device, param.shape) + if isinstance(param, GGUFParameter): + dtype_to_use = torch.uint8 + elif "patch_embedding" in name: + dtype_to_use = torch.float32 + else: + dtype_to_use = base_dtype + set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) + cnt += 1 + if cnt % 100 == 0: + pbar.update(100) + + #for name, param in transformer.named_parameters(): + # print(name, param.dtype, param.device, param.shape) + #patcher.load(device, full_load=True) + pbar.update_absolute(param_count) + pbar.update_absolute(0) + +def patch_control_lora(transformer, device): + log.info("Control-LoRA detected, patching model...") + + 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, + transformer.patch_embedding.kernel_size, + 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: + 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) + block.self_attn.k_loras = LoRALinearLayer(transformer.dim, transformer.dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength) + block.self_attn.v_loras = LoRALinearLayer(transformer.dim, transformer.dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength) + for lora in [block.self_attn.q_loras, block.self_attn.k_loras, block.self_attn.v_loras]: + for param in lora.parameters(): + param.requires_grad = False + for name, param in transformer.named_parameters(): + if "lora" in name: + param.data.copy_(lora_sd["diffusion_model." + name].to(param.device, dtype=param.dtype)) + +def add_lora_weights(patcher, lora, base_dtype, merge_loras=False): + #spacepxl's control LoRA patch + for l in lora: + log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}") + lora_path = l["path"] + lora_strength = l["strength"] + if isinstance(lora_strength, list): + if merge_loras: + raise ValueError("LoRA strength should be a single value when merge_loras=True") + transformer.lora_scheduling_enabled = True + if lora_strength == 0: + log.warning(f"LoRA {lora_path} has strength 0, skipping...") + continue + lora_sd = load_torch_file(lora_path, safe_load=True) + if "dwpose_embedding.0.weight" in lora_sd: #unianimate + from .unianimate.nodes import update_transformer + log.info("Unianimate LoRA detected, patching model...") + transformer = update_transformer(transformer, lora_sd) + + lora_sd = standardize_lora_key_format(lora_sd) + + if l["blocks"]: + lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"], l.get("layer_filter", [])) + + # 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} + control_lora=False + if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd: + control_lora = True + #stand-in LoRA patch + if "diffusion_model.blocks.0.self_attn.q_loras.down.weight" in lora_sd: + patch_stand_in_lora(patcher.model.diffusion_model, lora_sd, device, base_dtype, lora_strength) + # normal LoRA patch + else: + patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0) + + del lora_sd + return patcher, control_lora + #region Model loading class WanVideoModelLoader: @classmethod @@ -797,13 +969,16 @@ class WanVideoModelLoader: model_path = folder_paths.get_full_path_or_raise("diffusion_models", model) - + + gguf_reader = None if not gguf: sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True) else: - from diffusers.models.model_loading_utils import load_gguf_checkpoint - sd = load_gguf_checkpoint(model_path) - + #from diffusers.models.model_loading_utils import load_gguf_checkpoint + #sd = load_gguf_checkpoint(model_path) + from .gguf.gguf import load_gguf + sd, gguf_reader = load_gguf(model_path) + if quantization == "disabled": for k, v in sd.items(): if isinstance(v, torch.Tensor): @@ -1055,178 +1230,56 @@ class WanVideoModelLoader: model_type=comfy.model_base.ModelType.FLOW, device=device, ) - scale_weights = {} - if not gguf: - if "fp8" in quantization: - for k, v in sd.items(): - if k.endswith(".scale_weight"): - scale_weights[k] = v - if not merge_loras: - from .custom_linear import _replace_linear - transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights) - - if "fp8_e4m3fn" in quantization: - dtype = torch.float8_e4m3fn - elif "fp8_e5m2" in quantization: - dtype = torch.float8_e5m2 - else: - dtype = base_dtype - params_to_keep = {"norm", "bias", "time_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add", "ref_conv"} - if not lora_low_mem_load: - log.info("Using accelerate to load and assign model weights to device...") - param_count = sum(1 for _ in transformer.named_parameters()) - pbar = ProgressBar(param_count) - cnt = 0 - for name, param in tqdm(transformer.named_parameters(), - desc=f"Loading transformer parameters to {transformer_load_device}", - total=param_count, - leave=True): - dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype - dtype_to_use = dtype if sd[name].dtype == dtype else dtype_to_use - if "modulation" in name or "norm" in name or "bias" in name: - dtype_to_use = base_dtype - if "patch_embedding" in name: - dtype_to_use = torch.float32 - set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) - cnt += 1 - if cnt % 100 == 0: - pbar.update(100) - - #for name, param in transformer.named_parameters(): - # print(name, param.dtype, param.device, param.shape) - pbar.update_absolute(param_count) comfy_model.diffusion_model = transformer comfy_model.load_device = transformer_load_device - patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) patcher.model.is_patched = False - - control_lora = False - if lora is not None: - for l in lora: - log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}") - lora_path = l["path"] - lora_strength = l["strength"] - if isinstance(lora_strength, list): - if merge_loras: - raise ValueError("LoRA strength should be a single value when merge_loras=True") - transformer.lora_scheduling_enabled = True - if lora_strength == 0: - log.warning(f"LoRA {lora_path} has strength 0, skipping...") - continue - lora_sd = load_torch_file(lora_path, safe_load=True) - if "dwpose_embedding.0.weight" in lora_sd: #unianimate - from .unianimate.nodes import update_transformer - log.info("Unianimate LoRA detected, patching model...") - transformer = update_transformer(transformer, lora_sd) - - lora_sd = standardize_lora_key_format(lora_sd) - - if l["blocks"]: - lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"], l.get("layer_filter", [])) - - # 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} - - #spacepxl's control LoRA patch - # for key in lora_sd.keys(): - # print(key) - - if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd: - log.info("Control-LoRA detected, patching model...") - if not merge_loras: - log.warning("Control-LoRA patching is only supported with merge_loras=True, setting it to True") - merge_loras = True - control_lora = True - - in_cls = transformer.patch_embedding.__class__ # nn.Conv3d - old_in_dim = transformer.in_dim # 16 - new_in_dim = lora_sd["diffusion_model.patch_embedding.lora_A.weight"].shape[1] - assert new_in_dim == 32 - - new_in = in_cls( - new_in_dim, - transformer.patch_embedding.out_channels, - transformer.patch_embedding.kernel_size, - 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 - - 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(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength) - block.self_attn.k_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength) - block.self_attn.v_loras = LoRALinearLayer(dim, dim, rank=128, device=transformer_load_device, dtype=base_dtype, strength=lora_strength) - for lora in [block.self_attn.q_loras, block.self_attn.k_loras, block.self_attn.v_loras]: - for param in lora.parameters(): - param.requires_grad = False - for name, param in transformer.named_parameters(): - if "lora" in name: - param.data.copy_(lora_sd["diffusion_model." + name].to(param.device, dtype=param.dtype)) - else: - patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0) - - del lora_sd - - if not gguf and merge_loras: - log.info("Patching LoRA to the model...") - patcher = apply_lora( - patcher, device, transformer_load_device, - params_to_keep=params_to_keep, dtype=dtype, base_dtype=base_dtype, state_dict=sd, - low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights) - scale_weights.clear() - patcher.patches.clear() - if gguf: - #from diffusers.quantizers.gguf.utils import _replace_with_gguf_linear, GGUFParameter - from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter - log.info("Using GGUF to load and assign model weights to device...") - param_count = sum(1 for _ in transformer.named_parameters()) - - out_features = sd["blocks.0.self_attn.k.weight"].shape[1] + scale_weights = {} + if "fp8" in quantization: + for k, v in sd.items(): + if k.endswith(".scale_weight"): + scale_weights[k] = v - patcher.model.diffusion_model = _replace_with_gguf_linear(patcher.model.diffusion_model, base_dtype, sd, patches=patcher.patches) - pbar = ProgressBar(param_count) - cnt = 0 - for name, param in tqdm(patcher.model.diffusion_model.named_parameters(), - desc=f"Loading transformer parameters to {transformer_load_device}", - total=param_count, - leave=True): - if "loras" in name: - continue - #print(name, param.dtype, param.device, param.shape) - if isinstance(param, GGUFParameter): - dtype_to_use = torch.uint8 - elif "patch_embedding" in name: - dtype_to_use = torch.float32 - else: - dtype_to_use = base_dtype - set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) - cnt += 1 - if cnt % 100 == 0: - pbar.update(100) - - #for name, param in transformer.named_parameters(): - # print(name, param.dtype, param.device, param.shape) - #patcher.load(device, full_load=True) - pbar.update_absolute(param_count) - - patcher.model.is_patched = True + if "fp8_e4m3fn" in quantization: + weight_dtype = torch.float8_e4m3fn + elif "fp8_e5m2" in quantization: + 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"} + control_lora = False patch_linear = (True if "scaled" in quantization or (lora is not None and not merge_loras) else False) - + + if not merge_loras and control_lora: + log.warning("Control-LoRA patching is only supported with merge_loras=True") + + if lora is not None: + patcher, control_lora = add_lora_weights(patcher, lora, base_dtype, merge_loras=merge_loras) + + if not gguf: + if merge_loras and not patch_linear: + 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 + + 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, + low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights,) + if not control_lora: + scale_weights.clear() + patcher.patches.clear() + else: + from .custom_linear import _replace_linear + transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights) + if "fast" in quantization: if lora is not None and not merge_loras: raise NotImplementedError("fp8_fast is not supported with unmerged LoRAs") @@ -1234,7 +1287,7 @@ class WanVideoModelLoader: convert_fp8_linear(transformer, base_dtype, params_to_keep, scale_weight_keys=scale_weights) patch_linear = False - del sd + #del sd if multitalk_model is not None: transformer.audio_proj = multitalk_model["proj_model"] @@ -1263,18 +1316,18 @@ class WanVideoModelLoader: WanRMSNorm: AutoWrappedModule, }, module_config = dict( - offload_dtype=dtype, + offload_dtype=weight_dtype, offload_device=offload_device, - onload_dtype=dtype, + onload_dtype=weight_dtype, onload_device=device, computation_dtype=base_dtype, computation_device=device, ), max_num_param=params_to_keep, overflow_module_config = dict( - offload_dtype=dtype, + offload_dtype=weight_dtype, offload_device=offload_device, - onload_dtype=dtype, + onload_dtype=weight_dtype, onload_device=offload_device, computation_dtype=base_dtype, computation_device=device, @@ -1282,13 +1335,14 @@ class WanVideoModelLoader: compile_args = compile_args, ) - if load_device == "offload_device" and patcher.model.diffusion_model.device != offload_device: + if merge_loras: log.info(f"Moving diffusion model from {patcher.model.diffusion_model.device} to {offload_device}") patcher.model.diffusion_model.to(offload_device) gc.collect() mm.soft_empty_cache() - patcher.model["dtype"] = base_dtype + patcher.model["base_dtype"] = base_dtype + patcher.model["weight_dtype"] = weight_dtype patcher.model["base_path"] = model_path patcher.model["model_name"] = model patcher.model["manual_offloading"] = manual_offloading @@ -1296,9 +1350,11 @@ class WanVideoModelLoader: patcher.model["auto_cpu_offload"] = True if vram_management_args is not None else False patcher.model["control_lora"] = control_lora patcher.model["compile_args"] = compile_args - patcher.model["gguf"] = gguf + patcher.model["gguf_reader"] = gguf_reader patcher.model["fp8_matmul"] = "fast" in quantization patcher.model["scale_weights"] = scale_weights + patcher.model["sd"] = sd + patcher.model["lora"] = lora if 'transformer_options' not in patcher.model_options: patcher.model_options['transformer_options'] = {} diff --git a/utils.py b/utils.py index 95430f7..71782f7 100644 --- a/utils.py +++ b/utils.py @@ -173,7 +173,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d to_load.append((n, m, params)) to_load.sort(reverse=True) - pbar = ProgressBar(len(to_load)) + #pbar = ProgressBar(len(to_load)) for x in tqdm(to_load, desc="Loading model and applying LoRA weights:", leave=True): name = x[0] m = x[1] @@ -207,7 +207,7 @@ def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, d except: continue m.comfy_patched_weights = True - pbar.update(1) + #pbar.update(1) model.current_weight_patches_uuid = model.patches_uuid if low_mem_load: diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 7c06523..3c5a6bf 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1377,7 +1377,10 @@ class WanModel(torch.nn.Module): return block_mask def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None, prefetch_blocks=0, block_swap_debug=False): - log.info(f"Swapping {blocks_to_swap + 1} transformer blocks") + # Clamp blocks_to_swap to valid range + blocks_to_swap = max(0, min(blocks_to_swap, len(self.blocks))) + + log.info(f"Swapping {blocks_to_swap} transformer blocks") self.blocks_to_swap = blocks_to_swap self.prefetch_blocks = prefetch_blocks self.block_swap_debug = block_swap_debug @@ -1387,11 +1390,14 @@ class WanModel(torch.nn.Module): total_offload_memory = 0 total_main_memory = 0 + + # Calculate the index where swapping starts + swap_start_idx = len(self.blocks) - blocks_to_swap for b, block in tqdm(enumerate(self.blocks), total=len(self.blocks), desc="Initializing block swap"): block_memory = get_module_memory_mb(block) - if b > self.blocks_to_swap: + if b < swap_start_idx: block.to(self.main_device) total_main_memory += block_memory else: @@ -1402,12 +1408,17 @@ class WanModel(torch.nn.Module): vace_blocks_to_swap = 1 if vace_blocks_to_swap > 0 and self.vace_layers is not None: + # Clamp vace_blocks_to_swap to valid range + vace_blocks_to_swap = max(0, min(vace_blocks_to_swap, len(self.vace_blocks))) self.vace_blocks_to_swap = vace_blocks_to_swap + + # Calculate the index where VACE swapping starts + vace_swap_start_idx = len(self.vace_blocks) - vace_blocks_to_swap for b, block in tqdm(enumerate(self.vace_blocks), total=len(self.vace_blocks), desc="Initializing vace block swap"): block_memory = get_module_memory_mb(block) - if b > self.vace_blocks_to_swap: + if b < vace_swap_start_idx: block.to(self.main_device) total_main_memory += block_memory else: @@ -1447,9 +1458,10 @@ class WanModel(torch.nn.Module): hints = [] current_c = c + vace_swap_start_idx = len(self.vace_blocks) - self.vace_blocks_to_swap if self.vace_blocks_to_swap > 0 else len(self.vace_blocks) for b, block in enumerate(self.vace_blocks): - if b <= self.vace_blocks_to_swap and self.vace_blocks_to_swap >= 0: + if b >= vace_swap_start_idx and self.vace_blocks_to_swap > 0: block.to(self.main_device) if b == 0: @@ -1462,13 +1474,13 @@ class WanModel(torch.nn.Module): # Store skip connection c_skip = block.after_proj(c_processed) hints.append(c_skip.to( - self.offload_device if self.vace_blocks_to_swap != -1 else self.main_device, + self.offload_device if self.vace_blocks_to_swap > 0 else self.main_device, non_blocking=self.use_non_blocking )) current_c = c_processed - if b <= self.vace_blocks_to_swap and self.vace_blocks_to_swap >= 0: + if b >= vace_swap_start_idx and self.vace_blocks_to_swap > 0: block.to(self.offload_device, non_blocking=self.use_non_blocking) return hints @@ -1509,7 +1521,8 @@ class WanModel(torch.nn.Module): ref_target_masks=None, inner_t=None, standin_input=None, - fantasy_portrait_input=None + fantasy_portrait_input=None, + phantom_ref=None ): r""" Forward pass through the diffusion model @@ -1622,6 +1635,15 @@ class WanModel(torch.nn.Module): F += 1 x = [torch.concat([_fun_ref.unsqueeze(0), u], dim=1) for _fun_ref, u in zip(fun_ref, x)] + if phantom_ref is not None: + phantom_ref_frames = phantom_ref.size(1) + phantom_ref = self.original_patch_embedding(phantom_ref.unsqueeze(0).to(torch.float32)).flatten(2).transpose(1, 2).to(x[0].dtype) + grid_sizes = torch.stack([torch.tensor([u[0] + phantom_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) + phantom_ref_seq_len = phantom_ref.size(1) + seq_len += phantom_ref_seq_len + F += phantom_ref_frames + x = [torch.concat([u, phantom_ref.unsqueeze(0)], dim=1) for phantom_ref, u in zip(phantom_ref, x)] + seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long) assert seq_lens.max() <= seq_len x = torch.cat([ @@ -2005,20 +2027,21 @@ class WanModel(torch.nn.Module): # Asynchronous block offloading with CUDA streams and events cuda_stream = mm.get_offload_stream(device) events = [torch.cuda.Event() for _ in self.blocks] + swap_start_idx = len(self.blocks) - self.blocks_to_swap if self.blocks_to_swap > 0 else len(self.blocks) for b, block in enumerate(self.blocks): # Prefetch blocks if enabled if self.prefetch_blocks > 0: for prefetch_offset in range(1, self.prefetch_blocks + 1): prefetch_idx = b + prefetch_offset - if prefetch_idx < len(self.blocks) and self.blocks_to_swap >= 0 and prefetch_idx <= self.blocks_to_swap: + if prefetch_idx < len(self.blocks) and self.blocks_to_swap > 0 and prefetch_idx >= swap_start_idx: with torch.cuda.stream(cuda_stream): self.blocks[prefetch_idx].to(self.main_device, non_blocking=self.use_non_blocking) events[prefetch_idx].record(cuda_stream) if self.block_swap_debug: transfer_start = time.perf_counter() # Wait for block to be ready - if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: + if b >= swap_start_idx and self.blocks_to_swap > 0: if self.prefetch_blocks > 0: if not events[b].query(): events[b].synchronize() @@ -2037,7 +2060,7 @@ class WanModel(torch.nn.Module): compute_end = time.perf_counter() compute_time = compute_end - compute_start to_cpu_transfer_start = time.perf_counter() - if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: + if b >= swap_start_idx and self.blocks_to_swap > 0: block.to(self.offload_device, non_blocking=self.use_non_blocking) if self.block_swap_debug: to_cpu_transfer_end = time.perf_counter() @@ -2051,9 +2074,6 @@ class WanModel(torch.nn.Module): if (controlnet is not None) and (b % controlnet["controlnet_stride"] == 0) and (b // controlnet["controlnet_stride"] < len(controlnet["controlnet_states"])): x[:, :x_len] += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"] - if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: - block.to(self.offload_device, non_blocking=self.use_non_blocking) - if self.enable_teacache and (self.teacache_start_step <= current_step <= self.teacache_end_step) and pred_id is not None: self.teacache_state.update( pred_id, @@ -2076,9 +2096,14 @@ class WanModel(torch.nn.Module): ) if self.ref_conv is not None and fun_ref is not None: - full_ref_length = fun_ref.size(1) - x = x[:, full_ref_length:] + fun_ref_length = fun_ref.size(1) + x = x[:, fun_ref_length:] grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) + + if phantom_ref is not None: + phantom_ref_length = phantom_ref.size(1) + x = x[:, :-phantom_ref_length] + grid_sizes = torch.stack([torch.tensor([u[0] - phantom_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) if attn_cond is not None: x = x[:, :x_len] diff --git a/wanvideo/wan_video_vae.py b/wanvideo/wan_video_vae.py index 320a262..7ac9a05 100644 --- a/wanvideo/wan_video_vae.py +++ b/wanvideo/wan_video_vae.py @@ -1037,7 +1037,7 @@ class VideoVAE_(nn.Module): feat_cache=self._enc_feat_map, feat_idx=self._enc_conv_idx) out = torch.cat([out, out_], 2) - pbar.update(1) + pbar.update(iter_) mu = self.conv1(out).chunk(2, dim=1)[0] mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)