diff --git a/custom_linear.py b/custom_linear.py index 5b047b1..712e769 100644 --- a/custom_linear.py +++ b/custom_linear.py @@ -1,14 +1,52 @@ import torch import torch.nn as nn from accelerate import init_empty_weights +from .gguf.gguf_utils import GGUFParameter, dequantize_gguf_tensor + +@torch.library.custom_op("wanvideo::apply_lora", mutates_args=()) +def apply_lora(weight: torch.Tensor, lora_diff_0: torch.Tensor, lora_diff_1: torch.Tensor, lora_diff_2: float, lora_strength: float) -> torch.Tensor: + patch_diff = torch.mm( + lora_diff_0.flatten(start_dim=1), + lora_diff_1.flatten(start_dim=1) + ).reshape(weight.shape) + + alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 != 0.0 else 1.0 + scale = lora_strength * alpha + + return weight.add(patch_diff, alpha=scale) + +@apply_lora.register_fake +def _(weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength): + # Return weight with same metadata + return weight.clone() + +@torch.library.custom_op("wanvideo::apply_single_lora", mutates_args=()) +def apply_single_lora(weight: torch.Tensor, lora_diff: torch.Tensor, lora_strength: float) -> torch.Tensor: + return weight.add(lora_diff, alpha=lora_strength) + +@apply_single_lora.register_fake +def _(weight, lora_diff, lora_strength): + # Return weight with same metadata + return weight.clone() + +@torch.library.custom_op("wanvideo::linear_forward", mutates_args=()) +def linear_forward(input: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor | None) -> torch.Tensor: + return torch.nn.functional.linear(input, weight, bias) + +@linear_forward.register_fake +def _(input, weight, bias): + # Calculate output shape: (..., out_features) + out_features = weight.shape[0] + output_shape = list(input.shape[:-1]) + [out_features] + return input.new_empty(output_shape) #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, compile_args=None): - +def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None, compile_args=None, modules_to_not_convert=[]): + has_children = list(model.children()) if not has_children: return - + allow_compile = False for name, module in model.named_children(): @@ -16,13 +54,22 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s allow_compile = compile_args.get("allow_unmerged_lora_compile", False) module_prefix = prefix + name + "." module_prefix = module_prefix.replace("_orig_mod.", "") - _replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args) + _replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args, modules_to_not_convert) - 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: + if isinstance(module, nn.Linear) and "loras" not in module_prefix and name not in modules_to_not_convert: + weight_key = module_prefix + "weight" + if weight_key not in state_dict: + continue + + in_features = state_dict[weight_key].shape[1] + out_features = state_dict[weight_key].shape[0] + + is_gguf = isinstance(state_dict[weight_key], GGUFParameter) + + scale_weight = None + if not is_gguf and scale_weights is not None: scale_key = f"{module_prefix}scale_weight" + scale_weight = scale_weights.get(scale_key) with init_empty_weights(): model._modules[name] = CustomLinear( @@ -30,8 +77,9 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s out_features, module.bias is not None, compute_dtype=compute_dtype, - scale_weight=scale_weights.get(scale_key) if scale_weights else None, - allow_compile=allow_compile + scale_weight=scale_weight, + allow_compile=allow_compile, + is_gguf=is_gguf ) model._modules[name].source_cls = type(module) model._modules[name].requires_grad_(False) @@ -84,7 +132,8 @@ class CustomLinear(nn.Linear): compute_dtype=None, device=None, scale_weight=None, - allow_compile=False + allow_compile=False, + is_gguf=False ) -> None: super().__init__(in_features, out_features, bias, device) self.compute_dtype = compute_dtype @@ -93,11 +142,51 @@ class CustomLinear(nn.Linear): self.scale_weight = scale_weight self.lora_strengths = [] self.allow_compile = allow_compile + self.is_gguf = is_gguf if not allow_compile: - self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora) - self.forward = torch.compiler.disable()(self.forward) - + # Disable compilation for both methods + #self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora) + #self.forward = torch.compiler.disable()(self.forward) + # Use regular implementations instead of custom ops + self._apply_lora_impl = self._apply_lora_custom_op + self._apply_single_lora_impl = self._apply_single_lora_custom_op + self._linear_forward_impl = self._linear_forward_custom_op + else: + self._apply_lora_impl = self._apply_lora_direct + self._apply_single_lora_impl = self._apply_single_lora_direct + self._linear_forward_impl = self._linear_forward_direct + + + # Direct implementations (no custom ops) + def _apply_lora_direct(self, weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength): + patch_diff = torch.mm( + lora_diff_0.flatten(start_dim=1), + lora_diff_1.flatten(start_dim=1) + ).reshape(weight.shape) + alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 != 0.0 else 1.0 + scale = lora_strength * alpha + return weight.add(patch_diff, alpha=scale) + + def _apply_single_lora_direct(self, weight, lora_diff, lora_strength): + return weight.add(lora_diff, alpha=lora_strength) + + def _linear_forward_direct(self, input, weight, bias): + return torch.nn.functional.linear(input, weight, bias) + + # Custom op implementations + def _apply_lora_custom_op(self, weight, lora_diff_0, lora_diff_1, lora_diff_2, lora_strength): + return torch.ops.wanvideo.apply_lora(weight, lora_diff_0, lora_diff_1, + float(lora_diff_2) if lora_diff_2 is not None else 0.0, + float(lora_strength) + ) + + def _apply_single_lora_custom_op(self, weight, lora_diff, lora_strength): + return torch.ops.wanvideo.apply_single_lora(weight, lora_diff, float(lora_strength)) + + def _linear_forward_custom_op(self, input, weight, bias): + return torch.ops.wanvideo.linear_forward(input, weight, bias) + def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")): self.lora_diffs = [] for i, diff in enumerate(lora_diffs): @@ -111,10 +200,10 @@ class CustomLinear(nn.Linear): self.lora_diffs.append(f"lora_diff_{i}_0") def _get_weight_with_lora(self, weight): - """Apply LoRA outside compiled region""" + """Apply LoRA using custom ops to avoid graph breaks""" if not hasattr(self, "lora_diff_0_0"): return weight - + for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths): if isinstance(lora_strength, list): lora_strength = lora_strength[self.step] @@ -122,39 +211,50 @@ class CustomLinear(nn.Linear): continue elif lora_strength == 0.0: continue + 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) + + weight = self._apply_lora_impl( + weight, lora_diff_0, lora_diff_1, + float(lora_diff_2) if lora_diff_2 is not None else 0.0, + float(lora_strength) + ) else: lora_diff = getattr(self, lora_diff_names) - weight = weight.add(lora_diff, alpha=lora_strength) + weight = self._apply_single_lora_impl(weight, lora_diff,float(lora_strength)) + return weight + + def _prepare_weight(self, input): + """Prepare weight tensor - handles both regular and GGUF weights""" + if self.is_gguf: + weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype) + else: + weight = self.weight.to(input) return weight def forward(self, input): + weight = self._prepare_weight(input) + if self.bias is not None: - bias = self.bias.to(input) + bias = self.bias.to(input if not self.is_gguf else self.compute_dtype) else: bias = None - weight = self.weight.to(input) - if self.scale_weight is not None: + # Only apply scale_weight for non-GGUF models + if not self.is_gguf and self.scale_weight is not None: if weight.numel() < input.numel(): weight = weight * self.scale_weight else: input = input * self.scale_weight weight = self._get_weight_with_lora(weight) + out = self._linear_forward_impl(input, weight, bias) + del weight, input, bias + return out - return torch.nn.functional.linear(input, weight, bias) - def remove_lora_from_module(module): for name, submodule in module.named_modules(): if hasattr(submodule, "lora_diffs"): diff --git a/gguf/gguf.py b/gguf/gguf.py index 76dd87e..4b168aa 100644 --- a/gguf/gguf.py +++ b/gguf/gguf.py @@ -1,15 +1,11 @@ import torch -import torch.nn as nn import numpy as np import gguf -from accelerate import init_empty_weights -from .gguf_utils import GGUFParameter, dequantize_gguf_tensor -from ..utils import log +from .gguf_utils import GGUFParameter def load_gguf(model_path): - from gguf import GGUFReader - reader = GGUFReader(model_path) + reader = gguf.GGUFReader(model_path) parsed_parameters = {} for tensor in reader.tensors: # if the tensor is a torch supported dtype do not use GGUFParameter @@ -18,144 +14,12 @@ def load_gguf(model_path): 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 +from ..custom_linear import _replace_linear, set_lora_params, CustomLinear + def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None, compile_args=None): - def _should_convert_to_gguf(state_dict, prefix): - weight_key = prefix + "weight" - return weight_key in state_dict and isinstance(state_dict[weight_key], GGUFParameter) - - has_children = list(model.children()) - if not has_children: - return - - allow_compile = False - - for name, module in model.named_children(): - if compile_args is not None: - allow_compile = compile_args.get("allow_unmerged_lora_compile", False) - module_prefix = prefix + name + "." - _replace_with_gguf_linear(module, compute_dtype, state_dict, module_prefix, modules_to_not_convert, patches, compile_args) - - if ( - isinstance(module, nn.Linear) - and not isinstance(module, GGUFLinear) - and _should_convert_to_gguf(state_dict, module_prefix) - and name not in modules_to_not_convert - ): - in_features = state_dict[module_prefix + "weight"].shape[1] - out_features = state_dict[module_prefix + "weight"].shape[0] - - with init_empty_weights(): - model._modules[name] = GGUFLinear( - in_features, - out_features, - module.bias is not None, - compute_dtype=compute_dtype, - allow_compile=allow_compile - ) - - model._modules[name].source_cls = type(module) - model._modules[name].requires_grad_(False) - return model + return _replace_linear(model, compute_dtype, state_dict, prefix, patches, None, compile_args, modules_to_not_convert) 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, 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: - 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 - 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 + return set_lora_params(module, patches, module_prefix, device) - -class GGUFLinear(nn.Linear): - def __init__( - self, - in_features, - out_features, - bias=False, - compute_dtype=None, - device=None, - allow_compile=False - ) -> None: - super().__init__(in_features, out_features, bias, device) - self.compute_dtype = compute_dtype - self.lora_diffs = [] - self.lora_strengths = [] - self.step = 0 - self.allow_compile = allow_compile - - if not allow_compile: - self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora) - - def forward(self, inputs): - 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 - - weight = self._get_weight_with_lora(weight)#.to(self.compute_dtype) - - return torch.nn.functional.linear(inputs, weight, bias) - - def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")): - self.lora_diffs = [] - for i, diff in enumerate(lora_diffs): - if len(diff) > 1: - self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype)) - self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device, self.compute_dtype)) - 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}_0", diff[0].to(device, self.compute_dtype)) - self.lora_diffs.append(f"lora_diff_{i}_0") - - def _get_weight_with_lora(self, weight): - """Apply LoRA outside compiled region""" - if not hasattr(self, "lora_diff_0_0"): - return weight - - for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths): - if isinstance(lora_strength, list): - lora_strength = lora_strength[self.step] - if lora_strength == 0.0: - continue - elif lora_strength == 0.0: - continue - 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 +GGUFLinear = CustomLinear \ No newline at end of file diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index d25a4d4..1f4e4d0 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -18,14 +18,21 @@ except Exception as e: # Sage Attention imports try: from sageattention import sageattn - @torch.compiler.disable() - def sageattn_func(q, k, v, attn_mask=None, dropout_p=0, is_causal=False, tensor_layout="HND"): + + @torch.library.custom_op("wanvideo::sageattn", mutates_args=()) + def sageattn_func(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, tensor_layout: str = "HND" + ) -> torch.Tensor: if not (q.dtype == k.dtype == v.dtype): return sageattn(q, k.to(q.dtype), v.to(q.dtype), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout) elif q.dtype == torch.float32: return sageattn(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout).to(torch.float32) else: return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout) + + @sageattn_func.register_fake + def _(q, k, v, attn_mask=None, dropout_p=0.0, is_causal=False, tensor_layout="HND"): + # Return tensor with same shape as q + return q.clone() def sageattn_func_compiled(q, k, v, attn_mask=None, dropout_p=0, is_causal=False, tensor_layout="HND"): if not (q.dtype == k.dtype == v.dtype): @@ -42,18 +49,12 @@ except Exception as e: log.warning("sageattention DLL loading error, sageattention will not be available") sageattn_func = None -try: - from sageattn3 import sageattn3_blackwell as sageattn_blackwell -except: - try: - from sageattn import sageattn_blackwell - except: - SAGE3_AVAILABLE = False - try: from sageattention import sageattn_varlen - @torch.compiler.disable() - def sageattn_varlen_func(q, k, v, q_lens, k_lens, max_seqlen_q, max_seqlen_k, dropout_p=0, is_causal=False): + + @torch.library.custom_op("wanvideo::sageattn_varlen", mutates_args=()) + def sageattn_varlen_func(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, q_lens: list, k_lens: list, max_seqlen_q: int, max_seqlen_k: int, dropout_p: float = 0.0, is_causal: bool = False + ) -> torch.Tensor: cu_seqlens_q = torch.tensor([0] + list(torch.cumsum(torch.tensor(q_lens), dim=0)), device=q.device, dtype=torch.int32) cu_seqlens_k = torch.tensor([0] + list(torch.cumsum(torch.tensor(k_lens), dim=0)), device=q.device, dtype=torch.int32) if not (q.dtype == k.dtype == v.dtype): @@ -62,9 +63,23 @@ try: return sageattn_varlen(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal).to(torch.float32) else: return sageattn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, dropout_p=dropout_p, is_causal=is_causal) + + @sageattn_varlen_func.register_fake + def _(q, k, v, q_lens, k_lens, max_seqlen_q, max_seqlen_k, dropout_p=0.0, is_causal=False): + # Return tensor with same shape as q + return q.clone() except: sageattn_varlen_func = None +try: + from sageattn3 import sageattn3_blackwell as sageattn_blackwell +except: + try: + from sageattn import sageattn_blackwell + except: + SAGE3_AVAILABLE = False + + __all__ = [ 'flash_attention', 'attention', @@ -228,7 +243,7 @@ def attention( per_block_mean=False #seems necessary for reasonable VRAM usage, not sure of other implications ).transpose(1,2).contiguous() elif attention_mode == 'sageattn_varlen': - return sageattn_varlen_func( + return torch.ops.wanvideo.sageattn_varlen( q,k,v, q_lens=q_lens, k_lens=k_lens, @@ -238,4 +253,4 @@ def attention( elif attention_mode == 'sageattn_compiled': return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous() else: - return sageattn_func(q, k, v, tensor_layout="NHD").contiguous() + return torch.ops.wanvideo.sageattn(q, k, v, tensor_layout="NHD").contiguous()