Support fp8_scaled models and allow running LoRAs unmerged on other models as well
This commit is contained in:
+105
-5
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
|
||||
+45
-22
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user