Better fp8 linear layer patching with torch.compile
This commit is contained in:
@@ -0,0 +1,113 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from accelerate import init_empty_weights
|
||||
|
||||
#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):
|
||||
|
||||
has_children = list(model.children())
|
||||
if not has_children:
|
||||
return
|
||||
for name, module in model.named_children():
|
||||
module_prefix = prefix + name + "."
|
||||
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights)
|
||||
|
||||
if isinstance(module, nn.Linear):
|
||||
in_features = state_dict[module_prefix + "weight"].shape[1]
|
||||
out_features = state_dict[module_prefix + "weight"].shape[0]
|
||||
if scale_weights is not None:
|
||||
scale_key = f"{module_prefix}scale_weight"
|
||||
|
||||
with init_empty_weights():
|
||||
model._modules[name] = Fp8Linear(
|
||||
in_features,
|
||||
out_features,
|
||||
module.bias is not None,
|
||||
compute_dtype=compute_dtype,
|
||||
scale_weight=scale_weights.get(scale_key) if scale_weights else None
|
||||
)
|
||||
#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=""):
|
||||
# Recursively set lora_diffs and lora_strengths for all Fp8Linear 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):
|
||||
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:
|
||||
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
|
||||
lora_strengths = [p[0] for p in patch]
|
||||
module.lora = (lora_diffs, lora_strengths)
|
||||
module.step = 0 # Initialize step for LoRA scheduling
|
||||
|
||||
|
||||
class Fp8Linear(nn.Linear):
|
||||
def __init__(
|
||||
self,
|
||||
in_features,
|
||||
out_features,
|
||||
bias=False,
|
||||
compute_dtype=None,
|
||||
device=None,
|
||||
scale_weight=None
|
||||
) -> None:
|
||||
super().__init__(in_features, out_features, bias, device)
|
||||
self.compute_dtype = compute_dtype
|
||||
self.lora = None
|
||||
self.step = 0
|
||||
self.scale_weight = scale_weight
|
||||
|
||||
def forward(self, input):
|
||||
weight = self.weight.to(input.dtype)
|
||||
bias = self.bias.to(input.dtype) if self.bias is not None else None
|
||||
if self.scale_weight is not None:
|
||||
scale_weight = self.scale_weight.to(input.device)
|
||||
if weight.numel() < input.numel():
|
||||
weight = weight * scale_weight
|
||||
else:
|
||||
input = input * scale_weight
|
||||
|
||||
if self.lora is not None:
|
||||
weight = self.apply_lora(weight).to(input.dtype)
|
||||
|
||||
return torch.nn.functional.linear(input, weight, bias)
|
||||
|
||||
@torch.compiler.disable()
|
||||
def apply_lora(self, weight):
|
||||
for lora_diff, lora_strength in zip(self.lora[0], self.lora[1]):
|
||||
if isinstance(lora_strength, list):
|
||||
lora_strength = lora_strength[self.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 remove_lora_from_module(module):
|
||||
for name, submodule in module.named_modules():
|
||||
submodule.lora = None
|
||||
@@ -8,7 +8,7 @@ import hashlib
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
|
||||
from .wanvideo.modules.model import rope_params
|
||||
from .fp8_optimization import convert_linear_with_lora_and_scale, remove_lora_from_module
|
||||
from .fp8_optimization_v2 import remove_lora_from_module, set_lora_params as set_lora_params_fp8
|
||||
from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps, scheduler_list
|
||||
from .gguf.gguf import set_lora_params
|
||||
from .multitalk.multitalk import timestep_transform, add_noise
|
||||
@@ -1513,9 +1513,7 @@ class WanVideoSampler:
|
||||
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")
|
||||
convert_linear_with_lora_and_scale(transformer, patches=patcher.patches, scale_weight_keys=scale_weights)
|
||||
elif patch_linear:
|
||||
convert_linear_with_lora_and_scale(transformer, scale_weight_keys=scale_weights)
|
||||
set_lora_params_fp8(transformer, patcher.patches)
|
||||
else:
|
||||
remove_lora_from_module(transformer)
|
||||
|
||||
|
||||
@@ -1036,7 +1036,10 @@ class WanVideoModelLoader:
|
||||
for k, v in sd.items():
|
||||
if k.endswith(".scale_weight"):
|
||||
scale_weights[k] = v
|
||||
|
||||
if not merge_loras:
|
||||
from .fp8_optimization_v2 import _replace_linear
|
||||
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
|
||||
|
||||
if "fp8_e4m3fn" in quantization:
|
||||
dtype = torch.float8_e4m3fn
|
||||
elif "fp8_e5m2" in quantization:
|
||||
@@ -1182,7 +1185,7 @@ class WanVideoModelLoader:
|
||||
|
||||
patcher.model.is_patched = True
|
||||
|
||||
patch_linear = (True if "scaled" in quantization or not merge_loras else False)
|
||||
patch_linear = (True if "scaled" in quantization or (lora is not None and not merge_loras) else False)
|
||||
|
||||
if "fast" in quantization:
|
||||
if lora is not None and not merge_loras:
|
||||
|
||||
Reference in New Issue
Block a user