This commit is contained in:
kijai
2025-02-26 02:42:03 +02:00
parent 3b4b879117
commit 2d1afe1c3d
3 changed files with 42 additions and 55 deletions
+13 -29
View File
@@ -132,7 +132,7 @@ class WanVideoModelLoader:
"required": {
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}),
"base_precision": (["fp32", "bf16"], {"default": "bf16"}),
"base_precision": (["fp32", "bf16", "fp16"], {"default": "bf16"}),
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_scaled', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6", "torchao_int4", "torchao_int8"], {"default": 'disabled', "tooltip": "optional quantization method"}),
"load_device": (["main_device", "offload_device"], {"default": "main_device"}),
},
@@ -239,7 +239,8 @@ class WanVideoModelLoader:
if quantization == "fp8_e4m3fn_fast":
from .fp8_optimization import convert_fp8_linear
params_to_keep.update({"ff"})
#params_to_keep.update({"ffn"})
print(params_to_keep)
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep)
#compile
@@ -285,38 +286,19 @@ class WanVideoModelLoader:
comfy_model.diffusion_model = transformer
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
if lora is not None:
from comfy.sd import load_lora_for_models
for l in lora:
lora_path = l["path"]
lora_strength = l["strength"]
lora_sd = load_torch_file(lora_path, safe_load=True)
lora_sd = standardize_lora_key_format(lora_sd)
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
comfy.model_management.load_models_gpu([patcher])
for i, block in enumerate(patcher.model.diffusion_model.single_blocks):
log.info(f"Quantizing single_block {i}")
for name, _ in block.named_parameters(prefix=f"single_blocks.{i}"):
for i, block in enumerate(patcher.model.diffusion_model.blocks):
log.info(f"Quantizing block {i}")
for name, _ in block.named_parameters(prefix=f"blocks.{i}"):
#print(f"Parameter name: {name}")
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name])
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=base_dtype, value=sd[name])
if compile_args is not None:
patcher.model.diffusion_model.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
patcher.model.diffusion_model.blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
quantize_(block, quant_func)
print(block)
block.to(offload_device)
for i, block in enumerate(patcher.model.diffusion_model.double_blocks):
log.info(f"Quantizing double_block {i}")
for name, _ in block.named_parameters(prefix=f"double_blocks.{i}"):
#print(f"Parameter name: {name}")
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name])
if compile_args is not None:
patcher.model.diffusion_model.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
quantize_(block, quant_func)
#block.to(offload_device)
for name, param in patcher.model.diffusion_model.named_parameters():
if "single_blocks" not in name and "double_blocks" not in name:
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name])
if "blocks" not in name:
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=base_dtype, value=sd[name])
manual_offloading = False # to disable manual .to(device) calls
log.info(f"Quantized transformer blocks to {quantization}")
@@ -729,6 +711,7 @@ class WanVideoSampler:
model["block_swap_args"]["blocks_to_swap"] - 1 ,
)
else:
if model["manual_offloading"]:
transformer.to(device)
@@ -898,6 +881,7 @@ class WanVideoSampler:
del latent_model_input, timestep
if force_offload:
if model["manual_offloading"]:
transformer.to(offload_device)
mm.soft_empty_cache()
gc.collect()
+5 -2
View File
@@ -15,8 +15,11 @@ except ModuleNotFoundError:
try:
from sageattention import sageattn
@torch.compiler.disable()
def sageattn_func(q, k, v, attn_mask=None, dropout_p=0, is_causal=False):
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal)
except ModuleNotFoundError:
sageattn = None
sageattn_func = None
import warnings
__all__ = [
@@ -191,7 +194,7 @@ def attention(
k = k.transpose(1, 2).to(dtype)
v = v.transpose(1, 2).to(dtype)
out = sageattn(
out = sageattn_func(
q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p)
out = out.transpose(1, 2).contiguous()
+20 -20
View File
@@ -626,26 +626,26 @@ class WanModel(ModelMixin, ConfigMixin):
out.append(u)
return out
def init_weights(self):
r"""
Initialize model parameters using Xavier initialization.
"""
# def init_weights(self):
# r"""
# Initialize model parameters using Xavier initialization.
# """
# basic init
for m in self.modules():
if isinstance(m, nn.Linear):
nn.init.xavier_uniform_(m.weight)
if m.bias is not None:
nn.init.zeros_(m.bias)
# # basic init
# for m in self.modules():
# if isinstance(m, nn.Linear):
# nn.init.xavier_uniform_(m.weight)
# if m.bias is not None:
# nn.init.zeros_(m.bias)
# init embeddings
nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
for m in self.text_embedding.modules():
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=.02)
for m in self.time_embedding.modules():
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=.02)
# # init embeddings
# nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
# for m in self.text_embedding.modules():
# if isinstance(m, nn.Linear):
# nn.init.normal_(m.weight, std=.02)
# for m in self.time_embedding.modules():
# if isinstance(m, nn.Linear):
# nn.init.normal_(m.weight, std=.02)
# init output layer
nn.init.zeros_(self.head.head.weight)
# # init output layer
# nn.init.zeros_(self.head.head.weight)