diff --git a/fp8_optimization_v2.py b/custom_linear.py similarity index 95% rename from fp8_optimization_v2.py rename to custom_linear.py index 03b3e7a..ae5e81d 100644 --- a/fp8_optimization_v2.py +++ b/custom_linear.py @@ -19,7 +19,7 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s scale_key = f"{module_prefix}scale_weight" with init_empty_weights(): - model._modules[name] = Fp8Linear( + model._modules[name] = CustomLinear( in_features, out_features, module.bias is not None, @@ -34,11 +34,11 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s return model def set_lora_params(module, patches, module_prefix=""): - # Recursively set lora_diffs and lora_strengths for all Fp8Linear layers + # Recursively set lora_diffs and lora_strengths for all CustomLinear 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): + if isinstance(module, CustomLinear): key = f"diffusion_model.{module_prefix}weight" patch = patches.get(key, []) #print(f"Processing LoRA patches for {key}: {len(patch)} patches found") @@ -59,7 +59,7 @@ def set_lora_params(module, patches, module_prefix=""): module.step = 0 # Initialize step for LoRA scheduling -class Fp8Linear(nn.Linear): +class CustomLinear(nn.Linear): def __init__( self, in_features, diff --git a/gguf/gguf.py b/gguf/gguf.py index 1c0fc8a..44948ad 100644 --- a/gguf/gguf.py +++ b/gguf/gguf.py @@ -38,18 +38,18 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul module.bias is not None, compute_dtype=compute_dtype ) - #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=""): +def set_lora_params_gguf(module, patches, module_prefix=""): # Recursively set lora_diffs and lora_strengths for all GGUFLinear layers for name, child in module.named_children(): child_prefix = (f"{module_prefix}{name}.") - set_lora_params(child, patches, child_prefix) + set_lora_params_gguf(child, patches, child_prefix) if isinstance(module, GGUFLinear): key = f"diffusion_model.{module_prefix}weight" patch = patches.get(key, []) diff --git a/nodes.py b/nodes.py index 211bc03..ed72325 100644 --- a/nodes.py +++ b/nodes.py @@ -8,9 +8,9 @@ import hashlib from diffusers.schedulers import FlowMatchEulerDiscreteScheduler from .wanvideo.modules.model import rope_params -from .fp8_optimization_v2 import remove_lora_from_module, set_lora_params as set_lora_params_fp8 +from .custom_linear import remove_lora_from_module, set_lora_params from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps, scheduler_list -from .gguf.gguf import set_lora_params +from .gguf.gguf import set_lora_params_gguf from .multitalk.multitalk import timestep_transform, add_noise from .utils import(log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, add_noise_to_reference_video, optimized_scale, setup_radial_attention, @@ -1501,7 +1501,6 @@ class WanVideoSampler: 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) @@ -1513,12 +1512,12 @@ class WanVideoSampler: patch_linear = transformer_options.get("patch_linear", False) if gguf: - set_lora_params(transformer, patcher.patches) + 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") if not merge_loras and fp8_matmul: raise NotImplementedError("FP8 matmul with unmerged LoRAs is not supported") - set_lora_params_fp8(transformer, patcher.patches) + set_lora_params(transformer, patcher.patches) else: remove_lora_from_module(transformer) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 19b9679..0e7b7cc 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1037,7 +1037,7 @@ class WanVideoModelLoader: if k.endswith(".scale_weight"): scale_weights[k] = v if not merge_loras: - from .fp8_optimization_v2 import _replace_linear + from .custom_linear import _replace_linear transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights) if "fp8_e4m3fn" in quantization: diff --git a/skyreels/nodes.py b/skyreels/nodes.py index 134c1dc..7fe8f15 100644 --- a/skyreels/nodes.py +++ b/skyreels/nodes.py @@ -8,9 +8,9 @@ from tqdm import tqdm from ..wanvideo.modules.model import rope_params from ..wanvideo.schedulers.fm_solvers_unipc import FlowUniPCMultistepScheduler from diffusers.schedulers import FlowMatchEulerDiscreteScheduler -from ..fp8_optimization import convert_linear_with_lora_and_scale, remove_lora_from_module +from ..custom_linear import remove_lora_from_module, set_lora_params from ..wanvideo.schedulers.scheduling_flow_match_lcm import FlowMatchLCMScheduler -from ..gguf.gguf import set_lora_params +from ..gguf.gguf import set_lora_params_gguf from einops import rearrange from ..enhance_a_video.globals import disable_enhance @@ -146,19 +146,25 @@ class WanVideoDiffusionForcingSampler: device = mm.get_torch_device() offload_device = mm.unet_offload_device() + fp8_matmul = model["fp8_matmul"] gguf = model["gguf"] + merge_loras = transformer_options["merge_loras"] transformer_options = patcher.model_options.get("transformer_options", None) - if len(patcher.patches) != 0 and transformer_options.get("linear_patched", False) is True: + patch_linear = transformer_options.get("patch_linear", False) + + if gguf: + 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") - 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(transformer, patcher.patches) else: - log.info("Unloading all LoRAs") remove_lora_from_module(transformer) + transformer.lora_scheduling_enabled = transformer_options.get("lora_scheduling_enabled", False) + #torch.compile if model["auto_cpu_offload"] is False: transformer = compile_model(transformer, model["compile_args"])