From cc9bf1e4f53b156f35e947f182653f25a9a587a9 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 30 Oct 2025 16:44:03 +0200 Subject: [PATCH] Store lora diffs in buffers for GGUF as well --- fp8_optimization.py | 25 +-- gguf/gguf.py | 86 ++++++---- gguf/gguf_utils.py | 338 ++++++++++++++++++++++++++++++++++++++ wanvideo/modules/model.py | 62 ++++--- 4 files changed, 419 insertions(+), 92 deletions(-) create mode 100644 gguf/gguf_utils.py diff --git a/fp8_optimization.py b/fp8_optimization.py index ae5164c..217aa7b 100644 --- a/fp8_optimization.py +++ b/fp8_optimization.py @@ -8,15 +8,15 @@ def fp8_linear_forward(cls, base_dtype, input): if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: if len(input.shape) == 3: 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: scale_weight = scale_weight.to(input.device).squeeze() - + 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 @@ -31,24 +31,6 @@ def fp8_linear_forward(cls, base_dtype, input): return cls.original_forward(input) -@torch.compiler.disable() -def apply_lora(weight, lora, step=None): - for lora_diff, lora_strength in zip(lora[0], lora[1]): - if isinstance(lora_strength, list): - lora_strength = lora_strength[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 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(): @@ -61,4 +43,3 @@ def convert_fp8_linear(module, base_dtype, params_to_keep={}, scale_weight_keys= original_forward = submodule.forward setattr(submodule, "original_forward", original_forward) setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, base_dtype, input)) - diff --git a/gguf/gguf.py b/gguf/gguf.py index fa24023..a7d860b 100644 --- a/gguf/gguf.py +++ b/gguf/gguf.py @@ -1,13 +1,11 @@ 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 accelerate import init_empty_weights + +from .gguf_utils import GGUFParameter, dequantize_gguf_tensor from ..utils import log -if is_accelerate_available(): - from accelerate import init_empty_weights def load_gguf(model_path): from gguf import GGUFReader @@ -43,8 +41,7 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul in_features = state_dict[module_prefix + "weight"].shape[1] out_features = state_dict[module_prefix + "weight"].shape[0] - ctx = init_empty_weights if is_accelerate_available() else nullcontext - with ctx(): + with init_empty_weights(): model._modules[name] = GGUFLinear( in_features, out_features, @@ -53,19 +50,26 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul ) 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_gguf(module, patches, module_prefix=""): +def set_lora_params_gguf(module, patches, module_prefix="", device=torch.device("cpu")): # Recursively set lora_diffs and lora_strengths for all GGUFLinear layers for name, child in module.named_children(): + params = list(child.parameters()) + if params: + device = params[0].device + else: + device = torch.device("cpu") child_prefix = (f"{module_prefix}{name}.") - set_lora_params_gguf(child, patches, child_prefix) + set_lora_params_gguf(child, patches, child_prefix, device) if isinstance(module, GGUFLinear): 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: + key = key.replace("_orig_mod.", "") + patch = patches.get(key, []) if len(patch) != 0: lora_diffs = [] for p in patch: @@ -78,8 +82,8 @@ def set_lora_params_gguf(module, patches, module_prefix=""): lora_diffs.append(lora_obj[1]) else: continue - lora_strengths = [p[0] for p in patch] - module.lora = (lora_diffs, lora_strengths) + module.lora_strengths = [p[0] for p in patch] + module.set_lora_diffs(lora_diffs, device=device) module.step = 0 # Initialize step for LoRA scheduling @@ -94,41 +98,51 @@ class GGUFLinear(nn.Linear): ) -> None: super().__init__(in_features, out_features, bias, device) self.compute_dtype = compute_dtype - self.lora = None + self.lora_diffs = [] + self.lora_strengths = [] self.step = 0 def forward(self, inputs): - weight = self.dequantize_without_compile() - weight = weight.to(self.compute_dtype) + weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype) bias = self.bias.to(self.compute_dtype) if self.bias is not None else None - if hasattr(self, "lora") and self.lora is not None: - weight = self.apply_lora(weight, self.step).to(self.compute_dtype) + if hasattr(self, f"lora_diff_0_0"): + weight = self.apply_lora(weight).to(self.compute_dtype) - output = torch.nn.functional.linear(inputs, weight, bias) - return output + return torch.nn.functional.linear(inputs, weight, bias) - @torch.compiler.disable() - def dequantize_without_compile(self): - return dequantize_gguf_tensor(self.weight) + def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")): + self.lora_diffs = [] + for i, diff in enumerate(lora_diffs): + if isinstance(diff, tuple): + self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device)) + self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device)) + setattr(self, f"lora_diff_{i}_2", diff[2]) + self.lora_diffs.append((f"lora_diff_{i}_0", f"lora_diff_{i}_1", f"lora_diff_{i}_2")) + else: + self.register_buffer(f"lora_diff_{i}", diff.to(device)) + self.lora_diffs.append(f"lora_diff_{i}") - @torch.compiler.disable() - def apply_lora(self, weight, step=None): - for lora_diff, lora_strength in zip(self.lora[0], self.lora[1]): + def apply_lora(self, weight): + for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths): if isinstance(lora_strength, list): - lora_strength = lora_strength[step] + lora_strength = lora_strength[self.step] if lora_strength == 0.0: continue elif lora_strength == 0.0: continue - if len(lora_diff) == 1: - weight = weight.add(lora_diff[0].to(weight.device), alpha=lora_strength) - 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) + if isinstance(lora_diff_names, tuple): + lora_diff_0 = getattr(self, lora_diff_names[0]) + lora_diff_1 = getattr(self, lora_diff_names[1]) + lora_diff_2 = getattr(self, lora_diff_names[2]) + patch_diff = torch.mm( + lora_diff_0.flatten(start_dim=1), + lora_diff_1.flatten(start_dim=1) + ).reshape(weight.shape) + 0 + 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) + else: + lora_diff = getattr(self, lora_diff_names) + weight = weight.add(lora_diff, alpha=lora_strength) return weight \ No newline at end of file diff --git a/gguf/gguf_utils.py b/gguf/gguf_utils.py new file mode 100644 index 0000000..97436b4 --- /dev/null +++ b/gguf/gguf_utils.py @@ -0,0 +1,338 @@ +# Copyright 2024 The HuggingFace Team and City96. All rights reserved. +# # +# # Licensed under the Apache License, Version 2.0 (the "License"); +# # you may not use this file except in compliance with the License. +# # You may obtain a copy of the License at +# # +# # http://www.apache.org/licenses/LICENSE-2.0 +# # +# # Unless required by applicable law or agreed to in writing, software +# # distributed under the License is distributed on an "AS IS" BASIS, +# # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# # See the License for the specific language governing permissions and +# # limitations under the License. + + +import gguf +import torch +import torch.nn as nn + +# dequantize operations based on torch ports of GGUF dequantize_functions +# from City96 +# more info: https://github.com/city96/ComfyUI-GGUF/blob/main/dequant.py + + +QK_K = 256 +K_SCALE_SIZE = 12 + + +def to_uint32(x): + x = x.view(torch.uint8).to(torch.int32) + return (x[:, 0] | x[:, 1] << 8 | x[:, 2] << 16 | x[:, 3] << 24).unsqueeze(1) + + +def split_block_dims(blocks, *args): + n_max = blocks.shape[1] + dims = list(args) + [n_max - sum(args)] + return torch.split(blocks, dims, dim=1) + + +def get_scale_min(scales): + n_blocks = scales.shape[0] + scales = scales.view(torch.uint8) + scales = scales.reshape((n_blocks, 3, 4)) + + d, m, m_d = torch.split(scales, scales.shape[-2] // 3, dim=-2) + + sc = torch.cat([d & 0x3F, (m_d & 0x0F) | ((d >> 2) & 0x30)], dim=-1) + min = torch.cat([m & 0x3F, (m_d >> 4) | ((m >> 2) & 0x30)], dim=-1) + + return (sc.reshape((n_blocks, 8)), min.reshape((n_blocks, 8))) + + +def dequantize_blocks_Q8_0(blocks, block_size, type_size, dtype=None): + d, x = split_block_dims(blocks, 2) + d = d.view(torch.float16).to(dtype) + x = x.view(torch.int8) + return d * x + + +def dequantize_blocks_Q5_1(blocks, block_size, type_size, dtype=None): + n_blocks = blocks.shape[0] + + d, m, qh, qs = split_block_dims(blocks, 2, 2, 4) + d = d.view(torch.float16).to(dtype) + m = m.view(torch.float16).to(dtype) + qh = to_uint32(qh) + + qh = qh.reshape((n_blocks, 1)) >> torch.arange(32, device=d.device, dtype=torch.int32).reshape(1, 32) + ql = qs.reshape((n_blocks, -1, 1, block_size // 2)) >> torch.tensor( + [0, 4], device=d.device, dtype=torch.uint8 + ).reshape(1, 1, 2, 1) + qh = (qh & 1).to(torch.uint8) + ql = (ql & 0x0F).reshape((n_blocks, -1)) + + qs = ql | (qh << 4) + return (d * qs) + m + + +def dequantize_blocks_Q5_0(blocks, block_size, type_size, dtype=None): + n_blocks = blocks.shape[0] + + d, qh, qs = split_block_dims(blocks, 2, 4) + d = d.view(torch.float16).to(dtype) + qh = to_uint32(qh) + + qh = qh.reshape(n_blocks, 1) >> torch.arange(32, device=d.device, dtype=torch.int32).reshape(1, 32) + ql = qs.reshape(n_blocks, -1, 1, block_size // 2) >> torch.tensor( + [0, 4], device=d.device, dtype=torch.uint8 + ).reshape(1, 1, 2, 1) + + qh = (qh & 1).to(torch.uint8) + ql = (ql & 0x0F).reshape(n_blocks, -1) + + qs = (ql | (qh << 4)).to(torch.int8) - 16 + return d * qs + + +def dequantize_blocks_Q4_1(blocks, block_size, type_size, dtype=None): + n_blocks = blocks.shape[0] + + d, m, qs = split_block_dims(blocks, 2, 2) + d = d.view(torch.float16).to(dtype) + m = m.view(torch.float16).to(dtype) + + qs = qs.reshape((n_blocks, -1, 1, block_size // 2)) >> torch.tensor( + [0, 4], device=d.device, dtype=torch.uint8 + ).reshape(1, 1, 2, 1) + qs = (qs & 0x0F).reshape(n_blocks, -1) + + return (d * qs) + m + + +def dequantize_blocks_Q4_0(blocks, block_size, type_size, dtype=None): + n_blocks = blocks.shape[0] + + d, qs = split_block_dims(blocks, 2) + d = d.view(torch.float16).to(dtype) + + qs = qs.reshape((n_blocks, -1, 1, block_size // 2)) >> torch.tensor( + [0, 4], device=d.device, dtype=torch.uint8 + ).reshape((1, 1, 2, 1)) + qs = (qs & 0x0F).reshape((n_blocks, -1)).to(torch.int8) - 8 + return d * qs + + +def dequantize_blocks_Q6_K(blocks, block_size, type_size, dtype=None): + n_blocks = blocks.shape[0] + + ( + ql, + qh, + scales, + d, + ) = split_block_dims(blocks, QK_K // 2, QK_K // 4, QK_K // 16) + + scales = scales.view(torch.int8).to(dtype) + d = d.view(torch.float16).to(dtype) + d = (d * scales).reshape((n_blocks, QK_K // 16, 1)) + + ql = ql.reshape((n_blocks, -1, 1, 64)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape( + (1, 1, 2, 1) + ) + ql = (ql & 0x0F).reshape((n_blocks, -1, 32)) + qh = qh.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 2, 4, 6], device=d.device, dtype=torch.uint8).reshape( + (1, 1, 4, 1) + ) + qh = (qh & 0x03).reshape((n_blocks, -1, 32)) + q = (ql | (qh << 4)).to(torch.int8) - 32 + q = q.reshape((n_blocks, QK_K // 16, -1)) + + return (d * q).reshape((n_blocks, QK_K)) + + +def dequantize_blocks_Q5_K(blocks, block_size, type_size, dtype=None): + n_blocks = blocks.shape[0] + + d, dmin, scales, qh, qs = split_block_dims(blocks, 2, 2, K_SCALE_SIZE, QK_K // 8) + + d = d.view(torch.float16).to(dtype) + dmin = dmin.view(torch.float16).to(dtype) + + sc, m = get_scale_min(scales) + + d = (d * sc).reshape((n_blocks, -1, 1)) + dm = (dmin * m).reshape((n_blocks, -1, 1)) + + ql = qs.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape( + (1, 1, 2, 1) + ) + qh = qh.reshape((n_blocks, -1, 1, 32)) >> torch.arange(0, 8, device=d.device, dtype=torch.uint8).reshape( + (1, 1, 8, 1) + ) + ql = (ql & 0x0F).reshape((n_blocks, -1, 32)) + qh = (qh & 0x01).reshape((n_blocks, -1, 32)) + q = ql | (qh << 4) + + return (d * q - dm).reshape((n_blocks, QK_K)) + + +def dequantize_blocks_Q4_K(blocks, block_size, type_size, dtype=None): + n_blocks = blocks.shape[0] + + d, dmin, scales, qs = split_block_dims(blocks, 2, 2, K_SCALE_SIZE) + d = d.view(torch.float16).to(dtype) + dmin = dmin.view(torch.float16).to(dtype) + + sc, m = get_scale_min(scales) + + d = (d * sc).reshape((n_blocks, -1, 1)) + dm = (dmin * m).reshape((n_blocks, -1, 1)) + + qs = qs.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape( + (1, 1, 2, 1) + ) + qs = (qs & 0x0F).reshape((n_blocks, -1, 32)) + + return (d * qs - dm).reshape((n_blocks, QK_K)) + + +def dequantize_blocks_Q3_K(blocks, block_size, type_size, dtype=None): + n_blocks = blocks.shape[0] + + hmask, qs, scales, d = split_block_dims(blocks, QK_K // 8, QK_K // 4, 12) + d = d.view(torch.float16).to(dtype) + + lscales, hscales = scales[:, :8], scales[:, 8:] + lscales = lscales.reshape((n_blocks, 1, 8)) >> torch.tensor([0, 4], device=d.device, dtype=torch.uint8).reshape( + (1, 2, 1) + ) + lscales = lscales.reshape((n_blocks, 16)) + hscales = hscales.reshape((n_blocks, 1, 4)) >> torch.tensor( + [0, 2, 4, 6], device=d.device, dtype=torch.uint8 + ).reshape((1, 4, 1)) + hscales = hscales.reshape((n_blocks, 16)) + scales = (lscales & 0x0F) | ((hscales & 0x03) << 4) + scales = scales.to(torch.int8) - 32 + + dl = (d * scales).reshape((n_blocks, 16, 1)) + + ql = qs.reshape((n_blocks, -1, 1, 32)) >> torch.tensor([0, 2, 4, 6], device=d.device, dtype=torch.uint8).reshape( + (1, 1, 4, 1) + ) + qh = hmask.reshape(n_blocks, -1, 1, 32) >> torch.arange(0, 8, device=d.device, dtype=torch.uint8).reshape( + (1, 1, 8, 1) + ) + ql = ql.reshape((n_blocks, 16, QK_K // 16)) & 3 + qh = (qh.reshape((n_blocks, 16, QK_K // 16)) & 1) ^ 1 + q = ql.to(torch.int8) - (qh << 2).to(torch.int8) + + return (dl * q).reshape((n_blocks, QK_K)) + + +def dequantize_blocks_Q2_K(blocks, block_size, type_size, dtype=None): + n_blocks = blocks.shape[0] + + scales, qs, d, dmin = split_block_dims(blocks, QK_K // 16, QK_K // 4, 2) + d = d.view(torch.float16).to(dtype) + dmin = dmin.view(torch.float16).to(dtype) + + # (n_blocks, 16, 1) + dl = (d * (scales & 0xF)).reshape((n_blocks, QK_K // 16, 1)) + ml = (dmin * (scales >> 4)).reshape((n_blocks, QK_K // 16, 1)) + + shift = torch.tensor([0, 2, 4, 6], device=d.device, dtype=torch.uint8).reshape((1, 1, 4, 1)) + + qs = (qs.reshape((n_blocks, -1, 1, 32)) >> shift) & 3 + qs = qs.reshape((n_blocks, QK_K // 16, 16)) + qs = dl * qs - ml + + return qs.reshape((n_blocks, -1)) + + +def dequantize_blocks_BF16(blocks, block_size, type_size, dtype=None): + return (blocks.view(torch.int16).to(torch.int32) << 16).view(torch.float32) + + +GGML_QUANT_SIZES = gguf.GGML_QUANT_SIZES +dequantize_functions = { + gguf.GGMLQuantizationType.BF16: dequantize_blocks_BF16, + gguf.GGMLQuantizationType.Q8_0: dequantize_blocks_Q8_0, + gguf.GGMLQuantizationType.Q5_1: dequantize_blocks_Q5_1, + gguf.GGMLQuantizationType.Q5_0: dequantize_blocks_Q5_0, + gguf.GGMLQuantizationType.Q4_1: dequantize_blocks_Q4_1, + gguf.GGMLQuantizationType.Q4_0: dequantize_blocks_Q4_0, + gguf.GGMLQuantizationType.Q6_K: dequantize_blocks_Q6_K, + gguf.GGMLQuantizationType.Q5_K: dequantize_blocks_Q5_K, + gguf.GGMLQuantizationType.Q4_K: dequantize_blocks_Q4_K, + gguf.GGMLQuantizationType.Q3_K: dequantize_blocks_Q3_K, + gguf.GGMLQuantizationType.Q2_K: dequantize_blocks_Q2_K, +} +SUPPORTED_GGUF_QUANT_TYPES = list(dequantize_functions.keys()) + + +def _quant_shape_from_byte_shape(shape, type_size, block_size): + return (*shape[:-1], shape[-1] // type_size * block_size) + + +def dequantize_gguf_tensor(tensor): + if not hasattr(tensor, "quant_type"): + return tensor + + quant_type = tensor.quant_type + dequant_fn = dequantize_functions[quant_type] + + block_size, type_size = GGML_QUANT_SIZES[quant_type] + + tensor = tensor.view(torch.uint8) + shape = _quant_shape_from_byte_shape(tensor.shape, type_size, block_size) + + n_blocks = tensor.numel() // type_size + blocks = tensor.reshape((n_blocks, type_size)) + + dequant = dequant_fn(blocks, block_size, type_size) + dequant = dequant.reshape(shape) + + return dequant.as_tensor() + + +class GGUFParameter(torch.nn.Parameter): + def __new__(cls, data, requires_grad=False, quant_type=None): + data = data if data is not None else torch.empty(0) + self = torch.Tensor._make_subclass(cls, data, requires_grad) + self.quant_type = quant_type + block_size, type_size = GGML_QUANT_SIZES[quant_type] + self.quant_shape = _quant_shape_from_byte_shape(self.shape, type_size, block_size) + + return self + + def as_tensor(self): + return torch.Tensor._make_subclass(torch.Tensor, self, self.requires_grad) + + @classmethod + def __torch_function__(cls, func, types, args=(), kwargs=None): + if kwargs is None: + kwargs = {} + + result = super().__torch_function__(func, types, args, kwargs) + + # When converting from original format checkpoints we often use splits, cats etc on tensors + # this method ensures that the returned tensor type from those operations remains GGUFParameter + # so that we preserve quant_type information + quant_type = None + for arg in args: + if isinstance(arg, list) and isinstance(arg[0], GGUFParameter): + quant_type = arg[0].quant_type + break + if isinstance(arg, GGUFParameter): + quant_type = arg.quant_type + break + if isinstance(result, torch.Tensor): + return cls(result, quant_type=quant_type) + # Handle tuples and lists + elif isinstance(result, (tuple, list)): + # Preserve the original type (tuple or list) + wrapped = [cls(x, quant_type=quant_type) if isinstance(x, torch.Tensor) else x for x in result] + return type(result)(wrapped) + else: + return result diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index c602fe4..fc03124 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -24,6 +24,28 @@ from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot from ...MTV.mtv import apply_rotary_emb +from comfy import model_management as mm + +__all__ = ['WanModel'] + +class AdaLayerNorm(nn.Module): + def __init__(self, embedding_dim, output_dim=None, norm_elementwise_affine=False, norm_eps=1e-5, dtype=None, device=None, operations=None): + super().__init__() + + output_dim = output_dim or embedding_dim * 2 + + self.silu = nn.SiLU() + self.linear = operations.Linear(embedding_dim, output_dim, dtype=dtype, device=device) + self.norm = operations.LayerNorm(output_dim // 2, norm_eps, norm_elementwise_affine, dtype=dtype, device=device) + + def forward(self, x, temb): + temb = self.linear(self.silu(temb)) + shift, scale = temb.chunk(2, dim=1) + shift = shift[:, None, :] + scale = scale[:, None, :] + x = self.norm(x) * (1 + scale) + shift + return x + class FramePackMotioner(nn.Module):#from comfy.ldm.wan.model def __init__( self, @@ -77,22 +99,11 @@ class FramePackMotioner(nn.Module):#from comfy.ldm.wan.model rope = torch.cat([rope_post, rope_2x, rope_4x], dim=1) return motion_lat, rope -from diffusers.models.attention import AdaLayerNorm - -__all__ = ['WanModel'] - -from comfy import model_management as mm - - def zero_module(module): - """ - Zero out the parameters of a module and return it. - """ for p in module.parameters(): p.detach().zero_() return module - def torch_dfs(model: nn.Module, parent_name='root'): module_names, modules = [], [] current_name = parent_name if parent_name else 'root' @@ -404,7 +415,7 @@ class WanLayerNorm(nn.LayerNorm): """ return super().forward(x) - +#region selfattn class WanSelfAttention(nn.Module): def __init__(self, @@ -883,29 +894,12 @@ WAN_CROSSATTENTION_CLASSES = { class WanAttentionBlock(nn.Module): def __init__(self, - cross_attn_type, - in_features, - out_features, - ffn_dim, - ffn2_dim, - num_heads, - qk_norm=True, - cross_attn_norm=False, - eps=1e-6, - attention_mode="sdpa", - rope_func="comfy", - rms_norm_function="default", - use_motion_attn=False, - use_humo_audio_attn=False, - face_fuser_block=False, - lynx_ip_layers=None, - lynx_ref_layers=None, - block_idx=0, - # long cat - is_longcat = False, - ): + cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads, + qk_norm=True, cross_attn_norm=False, eps=1e-6, attention_mode="sdpa", rope_func="comfy", rms_norm_function="default", + use_motion_attn=False, use_humo_audio_attn=False, face_fuser_block=False, lynx_ip_layers=None, lynx_ref_layers=None, + block_idx=0, is_longcat=False): super().__init__() - self.dim = min(out_features, in_features) + self.dim = out_features self.ffn_dim = ffn_dim self.num_heads = num_heads self.head_dim = out_features // num_heads