diff --git a/fp8_optimization.py b/fp8_optimization.py index f32eae3..33d11d1 100644 --- a/fp8_optimization.py +++ b/fp8_optimization.py @@ -27,13 +27,113 @@ def fp8_linear_forward(cls, original_dtype, input): return cls.original_forward(input.to(original_dtype)) else: return cls.original_forward(input) + +def fp8_scaled_linear_forward(cls, original_dtype, input): + weight = cls.weight.to(original_dtype) + scale_weight = cls.scale_weight.to(input.device) + bias = cls.bias.to(original_dtype) if cls.bias is not None else None + + if weight.numel() < input.numel(): + weight = weight * scale_weight + else: + input = input * scale_weight + + lora = getattr(cls, "lora", None) + if lora is not None: + for lora_diff, lora_strength in zip(lora[0], lora[1]): + 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).to(original_dtype) + + return torch.nn.functional.linear(input, weight, bias) + +def linear_with_lora_forward(cls, original_dtype, input): + weight = cls.weight.to(original_dtype) + bias = cls.bias.to(original_dtype) if cls.bias is not None else None + + lora = getattr(cls, "lora", None) + if lora is not None: + for lora_diff, lora_strength in zip(lora[0], lora[1]): + 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).to(original_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) - for name, module in module.named_modules(): + for name, submodule in module.named_modules(): if not any(keyword in name for keyword in params_to_keep): - if isinstance(module, nn.Linear): - original_forward = module.forward - setattr(module, "original_forward", original_forward) - setattr(module, "forward", lambda input, m=module: fp8_linear_forward(m, original_dtype, input)) + if isinstance(submodule, nn.Linear): + original_forward = submodule.forward + setattr(submodule, "original_forward", original_forward) + setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, original_dtype, input)) + +def convert_fp8_scaled_linear(module, sd, original_dtype, params_to_keep={}, patches=None): + setattr(module, "fp8_scaled_enabled", True) + + for name, submodule in module.named_modules(): + if not any(keyword in name for keyword in params_to_keep): + scale_key = f"{name}.scale_weight" + has_scale = scale_key in sd + weight = getattr(submodule, 'weight', None) + has_fp8_weight = weight is not None and weight.dtype in [torch.float8_e4m3fn, torch.float8_e5m2] + if has_scale: + setattr(submodule, "scale_weight", sd[scale_key]) + + if patches is not None: + patch_key = f"diffusion_model.{name}.weight" + patch = patches.get(patch_key, []) + #print("Patches for", patch_key, ":", patch) + if len(patch) != 0: + lora_diffs = [] + for p in patch: + lora_obj = p[1] + if 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] + lora = (lora_diffs, lora_strengths) + setattr(submodule, "lora", lora) + + if isinstance(submodule, nn.Linear) and (has_scale or has_fp8_weight): + original_forward = submodule.forward + setattr(submodule, "original_forward", original_forward) + setattr(submodule, "forward", lambda input, m=submodule: fp8_scaled_linear_forward(m, original_dtype, input)) + +def convert_linear_with_lora(module, original_dtype, patches=None): + for name, submodule in module.named_modules(): + if isinstance(submodule, nn.Linear): + patch_key = f"diffusion_model.{name}.weight" + patch = patches.get(patch_key, []) + if len(patch) != 0: + lora_diffs = [] + for p in patch: + lora_obj = p[1] + if 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] + lora = (lora_diffs, lora_strengths) + setattr(submodule, "lora", lora) + # original_forward = submodule.forward + # setattr(submodule, "original_forward", original_forward) + setattr(submodule, "forward", lambda input, m=submodule: linear_with_lora_forward(m, original_dtype, input)) diff --git a/gguf/gguf.py b/gguf/gguf.py index e930867..aa6288c 100644 --- a/gguf/gguf.py +++ b/gguf/gguf.py @@ -71,13 +71,11 @@ class GGUFLinear(nn.Linear): device=None, lora_diffs=None, lora_strengths=None, - lora_alphas=None ) -> None: super().__init__(in_features, out_features, bias, device) self.compute_dtype = compute_dtype self.lora_diffs = lora_diffs self.lora_strengths = lora_strengths - self.lora_alphas = lora_alphas def forward(self, inputs): weight = dequantize_without_compile(self.weight) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 37e542d..2ffbc7e 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -339,7 +339,8 @@ class WanVideoLoraSelect: "optional": { "prev_lora":("WANVIDLORA", {"default": None, "tooltip": "For loading multiple LoRAs"}), "blocks":("SELECTEDBLOCKS", ), - "low_mem_load": ("BOOLEAN", {"default": False, "tooltip": "Load the LORA model with less VRAM usage, slower loading"}), + "low_mem_load": ("BOOLEAN", {"default": False, "tooltip": "Load the LORA model with less VRAM usage, slower loading. This affects ALL LoRAs, not just the current one"}), + "merge_loras": ("BOOLEAN", {"default": True, "tooltip": "Merge LoRAs into the model, otherwise they are loaded on the fly. Always enabled for GGUF and scaled fp8 models. This affects ALL LoRAs, not just the current one"}), }, "hidden": { "unique_id": "UNIQUE_ID", @@ -352,7 +353,7 @@ class WanVideoLoraSelect: CATEGORY = "WanVideoWrapper" DESCRIPTION = "Select a LoRA model from ComfyUI/models/loras" - def getlorapath(self, lora, strength, unique_id, blocks={}, prev_lora=None, low_mem_load=False): + def getlorapath(self, lora, strength, unique_id, blocks={}, prev_lora=None, low_mem_load=False, merge_loras=True): loras_list = [] strength = round(strength, 4) @@ -408,6 +409,7 @@ class WanVideoLoraSelect: "blocks": blocks.get("selected_blocks", {}), "layer_filter": blocks.get("layer_filter", ""), "low_mem_load": low_mem_load, + "merge_loras": merge_loras, } if prev_lora is not None: loras_list.extend(prev_lora) @@ -437,6 +439,8 @@ class WanVideoLoraSelectMulti: "prev_lora":("WANVIDLORA", {"default": None, "tooltip": "For loading multiple LoRAs"}), "blocks":("SELECTEDBLOCKS", ), "low_mem_load": ("BOOLEAN", {"default": False, "tooltip": "Load the LORA model with less VRAM usage, slower loading"}), + "merge_loras": ("BOOLEAN", {"default": True, "tooltip": "Merge LoRAs into the model, otherwise they are loaded on the fly. Always enabled for GGUF and scaled fp8 models. This affects ALL LoRAs, not just the current one"}), + } } @@ -448,7 +452,7 @@ class WanVideoLoraSelectMulti: 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): + low_mem_load=False, merge_loras=True): loras_list = [] strength_0 = round(strength_0, 4) @@ -481,6 +485,7 @@ class WanVideoLoraSelectMulti: "blocks": blocks.get("selected_blocks", {}), "layer_filter": blocks.get("layer_filter", ""), "low_mem_load": low_mem_load, + "merge_loras": merge_loras, } loras_list.append(lora) @@ -547,7 +552,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'], {"default": 'disabled', "tooltip": "optional quantization method"}), + "quantization": (["disabled", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2", "fp8_e4m3fn_fast_no_ffn", "fp8_e4m3fn_scaled"], {"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": { @@ -577,10 +582,12 @@ class WanVideoModelLoader: def loadmodel(self, model, base_precision, load_device, quantization, compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, vram_management_args=None, vace_model=None, fantasytalking_model=None, multitalk_model=None): assert not (vram_management_args is not None and block_swap_args is not None), "Can't use both block_swap_args and vram_management_args at the same time" - lora_low_mem_load = False + + lora_low_mem_load = merge_loras = False if lora is not None: for l in lora: - lora_low_mem_load = l.get("low_mem_load") if lora is not None else False + lora_low_mem_load = l.get("low_mem_load", False) + merge_loras = l.get("merge_loras", True) transformer = None mm.unload_all_models() @@ -602,6 +609,7 @@ class WanVideoModelLoader: raise ValueError("Quantization should be disabled when loading GGUF models.") quantization = "gguf" gguf = True + merge_loras = False manual_offloading = True @@ -629,11 +637,22 @@ class WanVideoModelLoader: else: from diffusers.models.model_loading_utils import load_gguf_checkpoint sd = load_gguf_checkpoint(model_path) - # for k, v in sd.items(): - # if isinstance(v, torch.Tensor): - # print(f"{k}: {v.shape} {v.dtype}") - # else: - # print(f"{k}: {type(v)}") + + if quantization == "disabled": + if "scaled_fp8" in sd: + quantization = "fp8_e4m3fn_scaled" + else: + for k, v in sd.items(): + if isinstance(v, torch.Tensor): + if v.dtype == torch.float8_e4m3fn: + quantization = "fp8_e4m3fn" + break + elif v.dtype == torch.float8_e5m2: + quantization = "fp8_e5m2" + break + + if merge_loras and "scaled" in quantization: + raise ValueError("scaled models currently do not support merging LoRAs, please disable merging or use a non-scaled model") if "vace_blocks.0.after_proj.weight" in sd and not "patch_embedding.weight" in sd: @@ -834,16 +853,6 @@ class WanVideoModelLoader: device=device, ) - if quantization == "disabled": - for k, v in sd.items(): - if isinstance(v, torch.Tensor): - if v.dtype == torch.float8_e4m3fn: - quantization = "fp8_e4m3fn" - break - elif v.dtype == torch.float8_e5m2: - quantization = "fp8_e5m2" - break - if not gguf: if "fp8_e4m3fn" in quantization: dtype = torch.float8_e4m3fn @@ -852,6 +861,8 @@ class WanVideoModelLoader: else: dtype = base_dtype params_to_keep = {"norm", "head", "bias", "time_in", "vector_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add"} + if "scaled" in quantization: + params_to_keep = {"patch_embedding", "modulation","norm", "bias"} #if lora is not None: # transformer_load_device = device if not lora_low_mem_load: @@ -926,7 +937,8 @@ class WanVideoModelLoader: del lora_sd - if not gguf: + if not gguf and not "scaled" in quantization 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) if gguf: @@ -961,11 +973,22 @@ class WanVideoModelLoader: 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) + + if "scaled" in quantization: + log.info("Using FP8 scaled linear quantization") + from .fp8_optimization import convert_fp8_scaled_linear + convert_fp8_scaled_linear(patcher.model.diffusion_model, sd, base_dtype, params_to_keep=params_to_keep, patches=patcher.patches) + elif not merge_loras and not gguf: + log.info("LoRAs will be applied at runtime") + from .fp8_optimization import convert_linear_with_lora + convert_linear_with_lora(patcher.model.diffusion_model, base_dtype, patches=patcher.patches) del sd