Merge branch 'kijai:main' into feat/encode_progress

This commit is contained in:
Mel Massadian
2025-08-10 15:37:02 +02:00
committed by GitHub
7 changed files with 206 additions and 75 deletions
@@ -1059,10 +1059,10 @@
480,
832,
"lanczos",
"pad",
"crop",
"0,0,0",
"top",
2,
"center",
16,
"cpu"
]
},
@@ -1587,7 +1587,7 @@
"crop",
"0, 0, 0",
"center",
2,
16,
"cpu"
]
},
+37 -30
View File
@@ -1,31 +1,31 @@
#based on ComfyUI's and MinusZoneAI's fp8_linear optimization
import torch
import torch.nn as nn
from .utils import log
def fp8_linear_forward(cls, original_dtype, input):
#based on ComfyUI's and MinusZoneAI's fp8_linear optimization
def fp8_linear_forward(cls, base_dtype, input):
weight_dtype = cls.weight.dtype
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
if len(input.shape) == 3:
#target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn
inn = input.reshape(-1, input.shape[2]).to(weight_dtype)
w = cls.weight.t()
scale = torch.ones((1), device=input.device, dtype=torch.float32)
bias = cls.bias.to(original_dtype) if cls.bias is not None else None
if bias is not None:
o = torch._scaled_mm(inn, w, out_dtype=original_dtype, bias=bias, scale_a=scale, scale_b=scale)
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:
o = torch._scaled_mm(inn, w, out_dtype=original_dtype, scale_a=scale, scale_b=scale)
scale_weight = scale_weight.to(input.device)
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
if isinstance(o, tuple):
o = o[0]
bias = cls.bias.to(base_dtype) if cls.bias is not None else None
return o.reshape((-1, input.shape[1], cls.weight.shape[0]))
o = torch._scaled_mm(inn, cls.weight.t(), out_dtype=base_dtype, bias=bias, scale_a=scale_input, scale_b=scale_weight)
return o.reshape((-1, input_shape[1], cls.weight.shape[0]))
else:
return cls.original_forward(input.to(original_dtype))
return cls.original_forward(input.to(base_dtype))
else:
return cls.original_forward(input)
@@ -67,20 +67,23 @@ def linear_with_lora_and_scale_forward(cls, input):
weight = apply_lora(weight, lora, cls.step).to(input.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)
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():
if not any(keyword in name for keyword in params_to_keep):
if isinstance(submodule, nn.Linear):
if scale_weight_keys is not None:
scale_key = f"{name}.scale_weight"
if scale_key in scale_weight_keys:
print("Setting scale_weight for", name)
setattr(submodule, "scale_weight", scale_weight_keys[scale_key])
original_forward = submodule.forward
setattr(submodule, "original_forward", original_forward)
setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, original_dtype, input))
setattr(submodule, "forward", lambda input, m=submodule: fp8_linear_forward(m, base_dtype, input))
def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=None, params_to_keep={}):
log.info("Patching Linear layers...")
for name, submodule in module.named_modules():
if not any(keyword in name for keyword in params_to_keep):
# Set scale_weight if present
@@ -90,9 +93,13 @@ def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=N
setattr(submodule, "scale_weight", scale_weight_keys[scale_key])
# Set LoRA if present
if hasattr(submodule, "lora"):
#print(f"removing old LoRA in {name}" )
delattr(submodule, "lora")
if patches is not None:
patch_key = f"diffusion_model.{name}.weight"
patch = patches.get(patch_key, [])
patch_key1 = f"diffusion_model.{name}.weight"
patch_key_compiled = f"diffusion_model.{name.replace('_orig_mod.', '')}.weight"
patch = patches.get(patch_key1, []) or patches.get(patch_key_compiled, [])
if len(patch) != 0:
lora_diffs = []
for p in patch:
@@ -108,17 +115,17 @@ def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=N
lora_strengths = [p[0] for p in patch]
lora = (lora_diffs, lora_strengths)
setattr(submodule, "lora", lora)
#print(f"Added LoRA to {name} with {len(lora_diffs)} diffs and strengths {lora_strengths}")
# Set forward if Linear and has either scale or lora
if isinstance(submodule, nn.Linear):
has_scale = hasattr(submodule, "scale_weight")
has_lora = hasattr(submodule, "lora")
if not hasattr(submodule, "original_forward"):
setattr(submodule, "original_forward", submodule.forward)
if has_scale or has_lora:
original_forward_ = submodule.forward
setattr(submodule, "original_forward_", original_forward_)
setattr(submodule, "forward", lambda input, m=submodule: linear_with_lora_and_scale_forward(m, input))
setattr(submodule, "step", 0) # Initialize step for LoRA if needed
setattr(submodule, "step", 0) # Initialize step for LoRA scheduling
def remove_lora_from_module(module):
unloaded = False
+113
View File
@@ -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
+19 -11
View File
@@ -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
@@ -1499,19 +1499,26 @@ class WanVideoSampler:
transformer = model.diffusion_model
dtype = model["dtype"]
fp8_matmul = model["fp8_matmul"]
gguf = model["gguf"]
scale_weights = model["scale_weights"]
control_lora = model["control_lora"]
transformer_options = patcher.model_options.get("transformer_options", None)
merge_loras = transformer_options["merge_loras"]
is_5b = transformer.out_dim == 48
vae_upscale_factor = 16 if is_5b else 8
if len(patcher.patches) != 0 and transformer_options.get("linear_with_lora", False) is True:
patch_linear = transformer_options.get("patch_linear", False)
if gguf:
set_lora_params(transformer, patcher.patches)
elif len(patcher.patches) != 0 and patch_linear:
log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model")
if not gguf:
convert_linear_with_lora_and_scale(transformer, patches=patcher.patches)
else:
set_lora_params(transformer, patcher.patches)
if not merge_loras and fp8_matmul:
raise NotImplementedError("FP8 matmul with unmerged LoRAs is not supported")
set_lora_params_fp8(transformer, patcher.patches)
else:
remove_lora_from_module(transformer)
@@ -1924,10 +1931,11 @@ class WanVideoSampler:
all_indices = []
for entry in extra_latents:
add_index = entry["index"]
noise[:, add_index] = entry["samples"].squeeze(0).squeeze(1).to(noise)
log.info(f"Adding extra samples to latent index {add_index}")
all_indices.append(add_index)
num_extra_frames = entry["samples"].shape[2]
noise[:, add_index:add_index+num_extra_frames] = entry["samples"].to(noise)
log.info(f"Adding extra samples to latent indices {add_index} to {add_index+num_extra_frames-1}")
all_indices.extend(range(add_index, add_index+num_extra_frames))
latent = noise.to(device)
@@ -3172,7 +3180,7 @@ class WanVideoSampler:
"drop_last": drop_last,
"generator_state": seed_g.get_state(),
},{
"samples": callback_latent.unsqueeze(0).cpu(),
"samples": callback_latent.unsqueeze(0).cpu() if callback is not None else None,
})
#region VideoDecode
+31 -28
View File
@@ -7,12 +7,11 @@ from tqdm import tqdm
from .wanvideo.modules.model import WanModel
from .wanvideo.modules.t5 import T5EncoderModel
from .wanvideo.modules.clip import CLIPModel
from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
from accelerate import init_empty_weights
from .utils import set_module_tensor_to_device
from .fp8_optimization import convert_linear_with_lora_and_scale
import folder_paths
import comfy.model_management as mm
from comfy.utils import load_torch_file, ProgressBar
@@ -687,6 +686,10 @@ class WanVideoSetLoRAs:
lora_sd = standardize_lora_key_format(lora_sd)
if l["blocks"]:
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"], l.get("layer_filter", []))
# Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys
if not any('img' in k for k in model.model.diffusion_model.state_dict().keys()):
lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k}
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
raise NotImplementedError("Control LoRA patching is not implemented in this node.")
@@ -698,7 +701,7 @@ class WanVideoSetLoRAs:
if 'transformer_options' not in patcher.model_options:
patcher.model_options['transformer_options'] = {}
patcher.model_options['transformer_options']["linear_with_lora"] = True
patcher.model_options['transformer_options']["patch_linear"] = True
return (patcher,)
@@ -711,7 +714,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", "fp8_e4m3fn_scaled", "fp8_e5m2_scaled"], {"default": "disabled", "tooltip": "optional quantization method"}),
"quantization": (["disabled", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e4m3fn_scaled", "fp8_e4m3fn_scaled_fast", "fp8_e5m2", "fp8_e5m2_fast", "fp8_e5m2_scaled", "fp8_e5m2_scaled_fast"], {"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": {
@@ -1027,16 +1030,16 @@ class WanVideoModelLoader:
model_type=comfy.model_base.ModelType.FLOW,
device=device,
)
scale_weights = {}
if not gguf:
scale_weights = {}
if "scaled" in quantization:
scale_weights = {}
if "fp8" in quantization:
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:
@@ -1094,9 +1097,14 @@ class WanVideoModelLoader:
transformer = update_transformer(transformer, lora_sd)
lora_sd = standardize_lora_key_format(lora_sd)
if l["blocks"]:
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"], l.get("layer_filter", []))
# Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys
if not any('img' in k for k in sd.keys()):
lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k}
#spacepxl's control LoRA patch
# for key in lora_sd.keys():
# print(key)
@@ -1140,6 +1148,8 @@ class WanVideoModelLoader:
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, control_lora=control_lora, scale_weights=scale_weights)
scale_weights.clear()
patcher.patches.clear()
if gguf:
#from diffusers.quantizers.gguf.utils import _replace_with_gguf_linear, GGUFParameter
@@ -1175,22 +1185,14 @@ class WanVideoModelLoader:
patcher.model.is_patched = True
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)
patch_linear = (True if "scaled" in quantization or (lora is not None and not merge_loras) else False)
if "scaled" in quantization and not merge_loras:
log.info("Using FP8 scaled linear quantization")
convert_linear_with_lora_and_scale(patcher.model.diffusion_model, scale_weights, patches=patcher.patches)
elif lora is not None and not merge_loras and not gguf:
log.info("LoRAs will be applied at runtime")
convert_linear_with_lora_and_scale(patcher.model.diffusion_model, patches=patcher.patches)
if "fast" in quantization:
if lora is not None and not merge_loras:
raise NotImplementedError("fp8_fast is not supported with unmerged LoRAs")
from .fp8_optimization import convert_fp8_linear
convert_fp8_linear(transformer, base_dtype, params_to_keep, scale_weight_keys=scale_weights)
patch_linear = False
del sd
@@ -1255,11 +1257,14 @@ class WanVideoModelLoader:
patcher.model["control_lora"] = control_lora
patcher.model["compile_args"] = compile_args
patcher.model["gguf"] = gguf
patcher.model["fp8_matmul"] = "fast" in quantization
patcher.model["scale_weights"] = scale_weights
if 'transformer_options' not in patcher.model_options:
patcher.model_options['transformer_options'] = {}
patcher.model_options["transformer_options"]["block_swap_args"] = block_swap_args
patcher.model_options["transformer_options"]["linear_with_lora"] = True if not merge_loras else False
patcher.model_options["transformer_options"]["patch_linear"] = patch_linear
patcher.model_options["transformer_options"]["merge_loras"] = merge_loras
for model in mm.current_loaded_models:
if model._model() == patcher:
@@ -1323,8 +1328,6 @@ class WanVideoVAELoader:
DESCRIPTION = "Loads Wan VAE model from 'ComfyUI/models/vae'"
def loadmodel(self, model_name, precision, compile_args=None):
from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
#with open(os.path.join(script_directory, 'configs', 'hy_vae_config.json')) as f:
# vae_config = json.load(f)
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "ComfyUI-WanVideoWrapper"
description = "ComfyUI wrapper nodes for WanVideo"
version = "1.2.6"
version = "1.2.7"
license = {file = "LICENSE"}
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.15.0", "ftfy", "gguf >= 0.14.0", "pyloudnorm"]
+1 -1
View File
@@ -149,7 +149,7 @@ class WanVideoDiffusionForcingSampler:
gguf = model["gguf"]
transformer_options = patcher.model_options.get("transformer_options", None)
if len(patcher.patches) != 0 and transformer_options.get("linear_with_lora", False) is True:
if len(patcher.patches) != 0 and transformer_options.get("linear_patched", False) is True:
log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model")
if not gguf:
convert_linear_with_lora_and_scale(transformer, patches=patcher.patches)