diff --git a/example_workflows/wanvideo_Fun_2_2_control_example_01.json b/example_workflows/wanvideo_Fun_2_2_control_example_01.json index 8a45901..dd166a6 100644 --- a/example_workflows/wanvideo_Fun_2_2_control_example_01.json +++ b/example_workflows/wanvideo_Fun_2_2_control_example_01.json @@ -1059,10 +1059,10 @@ 480, 832, "lanczos", - "pad", + "crop", "0,0,0", - "top", - 2, + "center", + 16, "cpu" ] }, @@ -1587,7 +1587,7 @@ "crop", "0, 0, 0", "center", - 2, + 16, "cpu" ] }, diff --git a/fp8_optimization.py b/fp8_optimization.py index b63e91d..dac0d3d 100644 --- a/fp8_optimization.py +++ b/fp8_optimization.py @@ -1,31 +1,31 @@ -#based on ComfyUI's and MinusZoneAI's fp8_linear optimization - import torch import torch.nn as nn from .utils import log -def fp8_linear_forward(cls, original_dtype, input): +#based on ComfyUI's and MinusZoneAI's fp8_linear optimization +def fp8_linear_forward(cls, base_dtype, input): weight_dtype = cls.weight.dtype if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: if len(input.shape) == 3: - #target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn - inn = input.reshape(-1, input.shape[2]).to(weight_dtype) - w = cls.weight.t() - - scale = torch.ones((1), device=input.device, dtype=torch.float32) - bias = cls.bias.to(original_dtype) if cls.bias is not None else None - - if bias is not None: - o = torch._scaled_mm(inn, w, out_dtype=original_dtype, bias=bias, scale_a=scale, scale_b=scale) + input_shape = input.shape + + scale_weight = getattr(cls, 'scale_weight', None) + if scale_weight is None: + scale_weight = torch.ones((), device=input.device, dtype=torch.float32) else: - o = torch._scaled_mm(inn, w, out_dtype=original_dtype, scale_a=scale, scale_b=scale) + scale_weight = scale_weight.to(input.device) + + scale_input = torch.ones((), device=input.device, dtype=torch.float32) + input = torch.clamp(input, min=-448, max=448, out=input) + inn = input.reshape(-1, input_shape[2]).to(torch.float8_e4m3fn).contiguous() #always e4m3fn because e5m2 * e5m2 is not supported - if isinstance(o, tuple): - o = o[0] + bias = cls.bias.to(base_dtype) if cls.bias is not None else None - return o.reshape((-1, input.shape[1], cls.weight.shape[0])) + o = torch._scaled_mm(inn, cls.weight.t(), out_dtype=base_dtype, bias=bias, scale_a=scale_input, scale_b=scale_weight) + + return o.reshape((-1, input_shape[1], cls.weight.shape[0])) else: - return cls.original_forward(input.to(original_dtype)) + return cls.original_forward(input.to(base_dtype)) else: return cls.original_forward(input) @@ -67,20 +67,23 @@ def linear_with_lora_and_scale_forward(cls, input): weight = apply_lora(weight, lora, cls.step).to(input.dtype) return torch.nn.functional.linear(input, weight, bias) - -def convert_fp8_linear(module, original_dtype, params_to_keep={}): - setattr(module, "fp8_matmul_enabled", True) - +def convert_fp8_linear(module, base_dtype, params_to_keep={}, scale_weight_keys=None): + log.info("FP8 matmul enabled") for name, submodule in module.named_modules(): if not any(keyword in name for keyword in params_to_keep): if isinstance(submodule, nn.Linear): + if scale_weight_keys is not None: + scale_key = f"{name}.scale_weight" + if scale_key in scale_weight_keys: + print("Setting scale_weight for", name) + setattr(submodule, "scale_weight", scale_weight_keys[scale_key]) original_forward = submodule.forward setattr(submodule, "original_forward", original_forward) - setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, original_dtype, input)) - - + setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, base_dtype, input)) + def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=None, params_to_keep={}): + log.info("Patching Linear layers...") for name, submodule in module.named_modules(): if not any(keyword in name for keyword in params_to_keep): # Set scale_weight if present @@ -90,9 +93,13 @@ def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=N setattr(submodule, "scale_weight", scale_weight_keys[scale_key]) # Set LoRA if present + if hasattr(submodule, "lora"): + #print(f"removing old LoRA in {name}" ) + delattr(submodule, "lora") if patches is not None: - patch_key = f"diffusion_model.{name}.weight" - patch = patches.get(patch_key, []) + patch_key1 = f"diffusion_model.{name}.weight" + patch_key_compiled = f"diffusion_model.{name.replace('_orig_mod.', '')}.weight" + patch = patches.get(patch_key1, []) or patches.get(patch_key_compiled, []) if len(patch) != 0: lora_diffs = [] for p in patch: @@ -108,17 +115,17 @@ def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=N lora_strengths = [p[0] for p in patch] lora = (lora_diffs, lora_strengths) setattr(submodule, "lora", lora) + #print(f"Added LoRA to {name} with {len(lora_diffs)} diffs and strengths {lora_strengths}") # Set forward if Linear and has either scale or lora if isinstance(submodule, nn.Linear): has_scale = hasattr(submodule, "scale_weight") has_lora = hasattr(submodule, "lora") + if not hasattr(submodule, "original_forward"): + setattr(submodule, "original_forward", submodule.forward) if has_scale or has_lora: - original_forward_ = submodule.forward - setattr(submodule, "original_forward_", original_forward_) setattr(submodule, "forward", lambda input, m=submodule: linear_with_lora_and_scale_forward(m, input)) - setattr(submodule, "step", 0) # Initialize step for LoRA if needed - + setattr(submodule, "step", 0) # Initialize step for LoRA scheduling def remove_lora_from_module(module): unloaded = False diff --git a/fp8_optimization_v2.py b/fp8_optimization_v2.py new file mode 100644 index 0000000..03b3e7a --- /dev/null +++ b/fp8_optimization_v2.py @@ -0,0 +1,113 @@ +import torch +import torch.nn as nn +from accelerate import init_empty_weights + +#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py +def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None): + + has_children = list(model.children()) + if not has_children: + return + for name, module in model.named_children(): + module_prefix = prefix + name + "." + _replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights) + + if isinstance(module, nn.Linear): + in_features = state_dict[module_prefix + "weight"].shape[1] + out_features = state_dict[module_prefix + "weight"].shape[0] + if scale_weights is not None: + scale_key = f"{module_prefix}scale_weight" + + with init_empty_weights(): + model._modules[name] = Fp8Linear( + in_features, + out_features, + module.bias is not None, + compute_dtype=compute_dtype, + scale_weight=scale_weights.get(scale_key) if scale_weights else None + ) + #set_lora_params(model._modules[name], patches, module_prefix) + model._modules[name].source_cls = type(module) + # Force requires_grad to False to avoid unexpected errors + model._modules[name].requires_grad_(False) + + return model + +def set_lora_params(module, patches, module_prefix=""): + # Recursively set lora_diffs and lora_strengths for all Fp8Linear layers + for name, child in module.named_children(): + child_prefix = (f"{module_prefix}{name}.") + set_lora_params(child, patches, child_prefix) + if isinstance(module, Fp8Linear): + key = f"diffusion_model.{module_prefix}weight" + patch = patches.get(key, []) + #print(f"Processing LoRA patches for {key}: {len(patch)} patches found") + if len(patch) != 0: + lora_diffs = [] + for p in patch: + lora_obj = p[1] + if "head" in key: + continue # For now skip LoRA for head layers + elif hasattr(lora_obj, "weights"): + lora_diffs.append(lora_obj.weights) + elif isinstance(lora_obj, tuple) and lora_obj[0] == "diff": + lora_diffs.append(lora_obj[1]) + else: + continue + lora_strengths = [p[0] for p in patch] + module.lora = (lora_diffs, lora_strengths) + module.step = 0 # Initialize step for LoRA scheduling + + +class Fp8Linear(nn.Linear): + def __init__( + self, + in_features, + out_features, + bias=False, + compute_dtype=None, + device=None, + scale_weight=None + ) -> None: + super().__init__(in_features, out_features, bias, device) + self.compute_dtype = compute_dtype + self.lora = None + self.step = 0 + self.scale_weight = scale_weight + + def forward(self, input): + weight = self.weight.to(input.dtype) + bias = self.bias.to(input.dtype) if self.bias is not None else None + if self.scale_weight is not None: + scale_weight = self.scale_weight.to(input.device) + if weight.numel() < input.numel(): + weight = weight * scale_weight + else: + input = input * scale_weight + + if self.lora is not None: + weight = self.apply_lora(weight).to(input.dtype) + + return torch.nn.functional.linear(input, weight, bias) + + @torch.compiler.disable() + def apply_lora(self, weight): + for lora_diff, lora_strength in zip(self.lora[0], self.lora[1]): + if isinstance(lora_strength, list): + lora_strength = lora_strength[self.step] + if lora_strength == 0.0: + continue + elif lora_strength == 0.0: + continue + patch_diff = torch.mm( + lora_diff[0].flatten(start_dim=1).to(weight.device), + lora_diff[1].flatten(start_dim=1).to(weight.device) + ).reshape(weight.shape) + alpha = lora_diff[2] / lora_diff[1].shape[0] if lora_diff[2] is not None else 1.0 + scale = lora_strength * alpha + weight = weight.add(patch_diff, alpha=scale) + return weight + +def remove_lora_from_module(module): + for name, submodule in module.named_modules(): + submodule.lora = None \ No newline at end of file diff --git a/nodes.py b/nodes.py index 6ed1e77..24d52bf 100644 --- a/nodes.py +++ b/nodes.py @@ -8,7 +8,7 @@ import hashlib from diffusers.schedulers import FlowMatchEulerDiscreteScheduler from .wanvideo.modules.model import rope_params -from .fp8_optimization import convert_linear_with_lora_and_scale, remove_lora_from_module +from .fp8_optimization_v2 import remove_lora_from_module, set_lora_params as set_lora_params_fp8 from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps, scheduler_list from .gguf.gguf import set_lora_params from .multitalk.multitalk import timestep_transform, add_noise @@ -1499,19 +1499,26 @@ class WanVideoSampler: transformer = model.diffusion_model dtype = model["dtype"] + fp8_matmul = model["fp8_matmul"] gguf = model["gguf"] + scale_weights = model["scale_weights"] control_lora = model["control_lora"] + transformer_options = patcher.model_options.get("transformer_options", None) + merge_loras = transformer_options["merge_loras"] is_5b = transformer.out_dim == 48 vae_upscale_factor = 16 if is_5b else 8 - if len(patcher.patches) != 0 and transformer_options.get("linear_with_lora", False) is True: + patch_linear = transformer_options.get("patch_linear", False) + + if gguf: + set_lora_params(transformer, patcher.patches) + elif len(patcher.patches) != 0 and patch_linear: log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model") - if not gguf: - convert_linear_with_lora_and_scale(transformer, patches=patcher.patches) - else: - set_lora_params(transformer, patcher.patches) + if not merge_loras and fp8_matmul: + raise NotImplementedError("FP8 matmul with unmerged LoRAs is not supported") + set_lora_params_fp8(transformer, patcher.patches) else: remove_lora_from_module(transformer) @@ -1924,10 +1931,11 @@ class WanVideoSampler: all_indices = [] for entry in extra_latents: add_index = entry["index"] - noise[:, add_index] = entry["samples"].squeeze(0).squeeze(1).to(noise) - log.info(f"Adding extra samples to latent index {add_index}") - all_indices.append(add_index) - + num_extra_frames = entry["samples"].shape[2] + noise[:, add_index:add_index+num_extra_frames] = entry["samples"].to(noise) + log.info(f"Adding extra samples to latent indices {add_index} to {add_index+num_extra_frames-1}") + all_indices.extend(range(add_index, add_index+num_extra_frames)) + latent = noise.to(device) @@ -3172,7 +3180,7 @@ class WanVideoSampler: "drop_last": drop_last, "generator_state": seed_g.get_state(), },{ - "samples": callback_latent.unsqueeze(0).cpu(), + "samples": callback_latent.unsqueeze(0).cpu() if callback is not None else None, }) #region VideoDecode diff --git a/nodes_model_loading.py b/nodes_model_loading.py index a4a42e9..a2b064e 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -7,12 +7,11 @@ from tqdm import tqdm from .wanvideo.modules.model import WanModel from .wanvideo.modules.t5 import T5EncoderModel from .wanvideo.modules.clip import CLIPModel +from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38 from accelerate import init_empty_weights from .utils import set_module_tensor_to_device -from .fp8_optimization import convert_linear_with_lora_and_scale - import folder_paths import comfy.model_management as mm from comfy.utils import load_torch_file, ProgressBar @@ -687,6 +686,10 @@ class WanVideoSetLoRAs: 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 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.") @@ -698,7 +701,7 @@ class WanVideoSetLoRAs: if 'transformer_options' not in patcher.model_options: patcher.model_options['transformer_options'] = {} - patcher.model_options['transformer_options']["linear_with_lora"] = True + patcher.model_options['transformer_options']["patch_linear"] = True return (patcher,) @@ -711,7 +714,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_e5m2", "fp8_e4m3fn_fast_no_ffn", "fp8_e4m3fn_scaled", "fp8_e5m2_scaled"], {"default": "disabled", "tooltip": "optional quantization method"}), + "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"}), "load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}), }, "optional": { @@ -1027,16 +1030,16 @@ class WanVideoModelLoader: model_type=comfy.model_base.ModelType.FLOW, device=device, ) - + scale_weights = {} if not gguf: - - scale_weights = {} - if "scaled" in quantization: - scale_weights = {} + if "fp8" in quantization: for k, v in sd.items(): if k.endswith(".scale_weight"): scale_weights[k] = v - + if not merge_loras: + from .fp8_optimization_v2 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: @@ -1094,9 +1097,14 @@ class WanVideoModelLoader: 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) @@ -1140,6 +1148,8 @@ class WanVideoModelLoader: 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 @@ -1175,22 +1185,14 @@ class WanVideoModelLoader: patcher.model.is_patched = True - - if "fast" in quantization: - if not merge_loras: - raise ValueError("FP8 fast quantization requires LoRAs to be merged into the model, please set merge_loras=True in the LoRA input") - from .fp8_optimization import convert_fp8_linear - if quantization == "fp8_e4m3fn_fast_no_ffn": - params_to_keep.update({"ffn"}) - print(params_to_keep) - convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep) + patch_linear = (True if "scaled" in quantization or (lora is not None and not merge_loras) else False) - if "scaled" in quantization and not merge_loras: - log.info("Using FP8 scaled linear quantization") - convert_linear_with_lora_and_scale(patcher.model.diffusion_model, scale_weights, patches=patcher.patches) - elif lora is not None and not merge_loras and not gguf: - log.info("LoRAs will be applied at runtime") - convert_linear_with_lora_and_scale(patcher.model.diffusion_model, patches=patcher.patches) + if "fast" in quantization: + if lora is not None and not merge_loras: + raise NotImplementedError("fp8_fast is not supported with unmerged LoRAs") + from .fp8_optimization import convert_fp8_linear + convert_fp8_linear(transformer, base_dtype, params_to_keep, scale_weight_keys=scale_weights) + patch_linear = False del sd @@ -1255,11 +1257,14 @@ class WanVideoModelLoader: patcher.model["control_lora"] = control_lora patcher.model["compile_args"] = compile_args patcher.model["gguf"] = gguf + patcher.model["fp8_matmul"] = "fast" in quantization + patcher.model["scale_weights"] = scale_weights if 'transformer_options' not in patcher.model_options: patcher.model_options['transformer_options'] = {} patcher.model_options["transformer_options"]["block_swap_args"] = block_swap_args - patcher.model_options["transformer_options"]["linear_with_lora"] = True if not merge_loras else False + patcher.model_options["transformer_options"]["patch_linear"] = patch_linear + patcher.model_options["transformer_options"]["merge_loras"] = merge_loras for model in mm.current_loaded_models: if model._model() == patcher: @@ -1323,8 +1328,6 @@ class WanVideoVAELoader: DESCRIPTION = "Loads Wan VAE model from 'ComfyUI/models/vae'" def loadmodel(self, model_name, precision, compile_args=None): - from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38 - dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] #with open(os.path.join(script_directory, 'configs', 'hy_vae_config.json')) as f: # vae_config = json.load(f) diff --git a/pyproject.toml b/pyproject.toml index e02605a..eb67e91 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "ComfyUI-WanVideoWrapper" description = "ComfyUI wrapper nodes for WanVideo" -version = "1.2.6" +version = "1.2.7" license = {file = "LICENSE"} dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.15.0", "ftfy", "gguf >= 0.14.0", "pyloudnorm"] diff --git a/skyreels/nodes.py b/skyreels/nodes.py index e0da49b..134c1dc 100644 --- a/skyreels/nodes.py +++ b/skyreels/nodes.py @@ -149,7 +149,7 @@ class WanVideoDiffusionForcingSampler: gguf = model["gguf"] transformer_options = patcher.model_options.get("transformer_options", None) - if len(patcher.patches) != 0 and transformer_options.get("linear_with_lora", False) is True: + if len(patcher.patches) != 0 and transformer_options.get("linear_patched", False) is True: log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model") if not gguf: convert_linear_with_lora_and_scale(transformer, patches=patcher.patches)