diff --git a/nodes.py b/nodes.py index d35576b..91db929 100644 --- a/nodes.py +++ b/nodes.py @@ -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,7 +711,8 @@ class WanVideoSampler: model["block_swap_args"]["blocks_to_swap"] - 1 , ) else: - transformer.to(device) + if model["manual_offloading"]: + transformer.to(device) # # Initialize TeaCache if enabled @@ -898,9 +881,10 @@ class WanVideoSampler: del latent_model_input, timestep if force_offload: - transformer.to(offload_device) - mm.soft_empty_cache() - gc.collect() + if model["manual_offloading"]: + transformer.to(offload_device) + mm.soft_empty_cache() + gc.collect() print_memory(device) try: diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index 8e7d6c2..0735663 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -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() diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 9a6da1a..047fea7 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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)