Use torch custom_ops to avoid graph breaks with torch.compile
Hopefully finally fixes the torch.compile VRAM issues...
This commit is contained in:
+129
-29
@@ -1,14 +1,52 @@
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from accelerate import init_empty_weights
|
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
|
#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())
|
has_children = list(model.children())
|
||||||
if not has_children:
|
if not has_children:
|
||||||
return
|
return
|
||||||
|
|
||||||
allow_compile = False
|
allow_compile = False
|
||||||
|
|
||||||
for name, module in model.named_children():
|
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)
|
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
|
||||||
module_prefix = prefix + name + "."
|
module_prefix = prefix + name + "."
|
||||||
module_prefix = module_prefix.replace("_orig_mod.", "")
|
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:
|
if isinstance(module, nn.Linear) and "loras" not in module_prefix and name not in modules_to_not_convert:
|
||||||
in_features = state_dict[module_prefix + "weight"].shape[1]
|
weight_key = module_prefix + "weight"
|
||||||
out_features = state_dict[module_prefix + "weight"].shape[0]
|
if weight_key not in state_dict:
|
||||||
if scale_weights is not None:
|
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_key = f"{module_prefix}scale_weight"
|
||||||
|
scale_weight = scale_weights.get(scale_key)
|
||||||
|
|
||||||
with init_empty_weights():
|
with init_empty_weights():
|
||||||
model._modules[name] = CustomLinear(
|
model._modules[name] = CustomLinear(
|
||||||
@@ -30,8 +77,9 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
|
|||||||
out_features,
|
out_features,
|
||||||
module.bias is not None,
|
module.bias is not None,
|
||||||
compute_dtype=compute_dtype,
|
compute_dtype=compute_dtype,
|
||||||
scale_weight=scale_weights.get(scale_key) if scale_weights else None,
|
scale_weight=scale_weight,
|
||||||
allow_compile=allow_compile
|
allow_compile=allow_compile,
|
||||||
|
is_gguf=is_gguf
|
||||||
)
|
)
|
||||||
model._modules[name].source_cls = type(module)
|
model._modules[name].source_cls = type(module)
|
||||||
model._modules[name].requires_grad_(False)
|
model._modules[name].requires_grad_(False)
|
||||||
@@ -84,7 +132,8 @@ class CustomLinear(nn.Linear):
|
|||||||
compute_dtype=None,
|
compute_dtype=None,
|
||||||
device=None,
|
device=None,
|
||||||
scale_weight=None,
|
scale_weight=None,
|
||||||
allow_compile=False
|
allow_compile=False,
|
||||||
|
is_gguf=False
|
||||||
) -> None:
|
) -> None:
|
||||||
super().__init__(in_features, out_features, bias, device)
|
super().__init__(in_features, out_features, bias, device)
|
||||||
self.compute_dtype = compute_dtype
|
self.compute_dtype = compute_dtype
|
||||||
@@ -93,11 +142,51 @@ class CustomLinear(nn.Linear):
|
|||||||
self.scale_weight = scale_weight
|
self.scale_weight = scale_weight
|
||||||
self.lora_strengths = []
|
self.lora_strengths = []
|
||||||
self.allow_compile = allow_compile
|
self.allow_compile = allow_compile
|
||||||
|
self.is_gguf = is_gguf
|
||||||
|
|
||||||
if not allow_compile:
|
if not allow_compile:
|
||||||
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
|
# Disable compilation for both methods
|
||||||
self.forward = torch.compiler.disable()(self.forward)
|
#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")):
|
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
|
||||||
self.lora_diffs = []
|
self.lora_diffs = []
|
||||||
for i, diff in enumerate(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")
|
self.lora_diffs.append(f"lora_diff_{i}_0")
|
||||||
|
|
||||||
def _get_weight_with_lora(self, weight):
|
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"):
|
if not hasattr(self, "lora_diff_0_0"):
|
||||||
return weight
|
return weight
|
||||||
|
|
||||||
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
|
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
|
||||||
if isinstance(lora_strength, list):
|
if isinstance(lora_strength, list):
|
||||||
lora_strength = lora_strength[self.step]
|
lora_strength = lora_strength[self.step]
|
||||||
@@ -122,39 +211,50 @@ class CustomLinear(nn.Linear):
|
|||||||
continue
|
continue
|
||||||
elif lora_strength == 0.0:
|
elif lora_strength == 0.0:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if isinstance(lora_diff_names, tuple):
|
if isinstance(lora_diff_names, tuple):
|
||||||
lora_diff_0 = getattr(self, lora_diff_names[0])
|
lora_diff_0 = getattr(self, lora_diff_names[0])
|
||||||
lora_diff_1 = getattr(self, lora_diff_names[1])
|
lora_diff_1 = getattr(self, lora_diff_names[1])
|
||||||
lora_diff_2 = getattr(self, lora_diff_names[2])
|
lora_diff_2 = getattr(self, lora_diff_names[2])
|
||||||
patch_diff = torch.mm(
|
|
||||||
lora_diff_0.flatten(start_dim=1),
|
weight = self._apply_lora_impl(
|
||||||
lora_diff_1.flatten(start_dim=1)
|
weight, lora_diff_0, lora_diff_1,
|
||||||
).reshape(weight.shape) + 0
|
float(lora_diff_2) if lora_diff_2 is not None else 0.0,
|
||||||
alpha = lora_diff_2 / lora_diff_1.shape[0] if lora_diff_2 is not None else 1.0
|
float(lora_strength)
|
||||||
scale = lora_strength * alpha
|
)
|
||||||
weight = weight.add(patch_diff, alpha=scale)
|
|
||||||
else:
|
else:
|
||||||
lora_diff = getattr(self, lora_diff_names)
|
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
|
return weight
|
||||||
|
|
||||||
def forward(self, input):
|
def forward(self, input):
|
||||||
|
weight = self._prepare_weight(input)
|
||||||
|
|
||||||
if self.bias is not None:
|
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:
|
else:
|
||||||
bias = None
|
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():
|
if weight.numel() < input.numel():
|
||||||
weight = weight * self.scale_weight
|
weight = weight * self.scale_weight
|
||||||
else:
|
else:
|
||||||
input = input * self.scale_weight
|
input = input * self.scale_weight
|
||||||
|
|
||||||
weight = self._get_weight_with_lora(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):
|
def remove_lora_from_module(module):
|
||||||
for name, submodule in module.named_modules():
|
for name, submodule in module.named_modules():
|
||||||
if hasattr(submodule, "lora_diffs"):
|
if hasattr(submodule, "lora_diffs"):
|
||||||
|
|||||||
+7
-143
@@ -1,15 +1,11 @@
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import gguf
|
import gguf
|
||||||
from accelerate import init_empty_weights
|
|
||||||
|
|
||||||
from .gguf_utils import GGUFParameter, dequantize_gguf_tensor
|
from .gguf_utils import GGUFParameter
|
||||||
from ..utils import log
|
|
||||||
|
|
||||||
def load_gguf(model_path):
|
def load_gguf(model_path):
|
||||||
from gguf import GGUFReader
|
reader = gguf.GGUFReader(model_path)
|
||||||
reader = GGUFReader(model_path)
|
|
||||||
parsed_parameters = {}
|
parsed_parameters = {}
|
||||||
for tensor in reader.tensors:
|
for tensor in reader.tensors:
|
||||||
# if the tensor is a torch supported dtype do not use GGUFParameter
|
# 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
|
parsed_parameters[tensor.name] = GGUFParameter(meta_tensor, quant_type=tensor.tensor_type) if is_gguf_quant else meta_tensor
|
||||||
return parsed_parameters, reader
|
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 _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):
|
return _replace_linear(model, compute_dtype, state_dict, prefix, patches, None, compile_args, modules_to_not_convert)
|
||||||
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
|
|
||||||
|
|
||||||
def set_lora_params_gguf(module, patches, module_prefix="", device=torch.device("cpu")):
|
def set_lora_params_gguf(module, patches, module_prefix="", device=torch.device("cpu")):
|
||||||
# Recursively set lora_diffs and lora_strengths for all GGUFLinear layers
|
return set_lora_params(module, patches, module_prefix, device)
|
||||||
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
|
|
||||||
|
|
||||||
|
GGUFLinear = CustomLinear
|
||||||
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
|
|
||||||
@@ -18,14 +18,21 @@ except Exception as e:
|
|||||||
# Sage Attention imports
|
# Sage Attention imports
|
||||||
try:
|
try:
|
||||||
from sageattention import sageattn
|
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):
|
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)
|
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:
|
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)
|
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:
|
else:
|
||||||
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout)
|
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"):
|
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):
|
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")
|
log.warning("sageattention DLL loading error, sageattention will not be available")
|
||||||
sageattn_func = None
|
sageattn_func = None
|
||||||
|
|
||||||
try:
|
|
||||||
from sageattn3 import sageattn3_blackwell as sageattn_blackwell
|
|
||||||
except:
|
|
||||||
try:
|
|
||||||
from sageattn import sageattn_blackwell
|
|
||||||
except:
|
|
||||||
SAGE3_AVAILABLE = False
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from sageattention import sageattn_varlen
|
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_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)
|
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):
|
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)
|
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:
|
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)
|
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:
|
except:
|
||||||
sageattn_varlen_func = None
|
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__ = [
|
__all__ = [
|
||||||
'flash_attention',
|
'flash_attention',
|
||||||
'attention',
|
'attention',
|
||||||
@@ -228,7 +243,7 @@ def attention(
|
|||||||
per_block_mean=False #seems necessary for reasonable VRAM usage, not sure of other implications
|
per_block_mean=False #seems necessary for reasonable VRAM usage, not sure of other implications
|
||||||
).transpose(1,2).contiguous()
|
).transpose(1,2).contiguous()
|
||||||
elif attention_mode == 'sageattn_varlen':
|
elif attention_mode == 'sageattn_varlen':
|
||||||
return sageattn_varlen_func(
|
return torch.ops.wanvideo.sageattn_varlen(
|
||||||
q,k,v,
|
q,k,v,
|
||||||
q_lens=q_lens,
|
q_lens=q_lens,
|
||||||
k_lens=k_lens,
|
k_lens=k_lens,
|
||||||
@@ -238,4 +253,4 @@ def attention(
|
|||||||
elif attention_mode == 'sageattn_compiled':
|
elif attention_mode == 'sageattn_compiled':
|
||||||
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
|
return sageattn_func_compiled(q, k, v, tensor_layout="NHD").contiguous()
|
||||||
else:
|
else:
|
||||||
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
|
return torch.ops.wanvideo.sageattn(q, k, v, tensor_layout="NHD").contiguous()
|
||||||
|
|||||||
Reference in New Issue
Block a user