From d903ea7d3dddaadddccbe8a1881bc75ac8db998a Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 6 Jul 2025 23:55:37 +0300 Subject: [PATCH] small fixes, add comfy pbar for model loading --- nodes_model_loading.py | 15 ++++++++++++--- uni3c/controlnet.py | 12 ++++++++---- uni3c/nodes.py | 6 ++++-- wanvideo/modules/model.py | 13 +++++++------ 4 files changed, 31 insertions(+), 15 deletions(-) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index baae7fd..a7d77aa 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -10,7 +10,7 @@ from accelerate.utils import set_module_tensor_to_device import folder_paths import comfy.model_management as mm -from comfy.utils import load_torch_file +from comfy.utils import load_torch_file, ProgressBar import comfy.model_base from comfy.sd import load_lora_for_models @@ -712,6 +712,7 @@ class WanVideoModelLoader: if not lora_low_mem_load: log.info("Using accelerate to load and assign model weights to device...") param_count = sum(1 for _ in transformer.named_parameters()) + pbar = ProgressBar(param_count) for name, param in tqdm(transformer.named_parameters(), desc=f"Loading transformer parameters to {transformer_load_device}", total=param_count, @@ -719,7 +720,8 @@ class WanVideoModelLoader: dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype if "patch_embedding" in name: dtype_to_use = torch.float32 - set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) + set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) + pbar.update(1) comfy_model.diffusion_model = transformer comfy_model.load_device = transformer_load_device @@ -786,10 +788,16 @@ class WanVideoModelLoader: #from diffusers.quantizers.gguf.utils import _replace_with_gguf_linear, GGUFParameter from .gguf.gguf import _replace_with_gguf_linear, GGUFParameter log.info("Using GGUF to load and assign model weights to device...") + param_count = sum(1 for _ in transformer.named_parameters()) + out_features = sd["blocks.0.self_attn.k.weight"].shape[1] patcher.model.diffusion_model = _replace_with_gguf_linear(patcher.model.diffusion_model, base_dtype, sd, patches=patcher.patches) - for name, param in patcher.model.diffusion_model.named_parameters(): + pbar = ProgressBar(param_count) + for name, param in tqdm(patcher.model.diffusion_model.named_parameters(), + desc=f"Loading transformer parameters to {transformer_load_device}", + total=param_count, + leave=True): #print(name, param.dtype, param.device, param.shape) if isinstance(param, GGUFParameter): dtype_to_use = torch.uint8 @@ -798,6 +806,7 @@ class WanVideoModelLoader: else: dtype_to_use = base_dtype set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) + pbar.update(1) #for name, param in transformer.named_parameters(): # print(name, param.dtype, param.device, param.shape) #patcher.load(device, full_load=True) diff --git a/uni3c/controlnet.py b/uni3c/controlnet.py index 0517e74..3b0e84b 100644 --- a/uni3c/controlnet.py +++ b/uni3c/controlnet.py @@ -189,6 +189,8 @@ class WanControlNet(ModelMixin): self.in_channels = controlnet_cfg["in_channels"] self.dim = controlnet_cfg["dim"] self.num_heads = controlnet_cfg["num_heads"] + self.quantized = controlnet_cfg["quantized"] + self.base_dtype = controlnet_cfg["base_dtype"] if controlnet_cfg["conv_out_dim"] != controlnet_cfg["dim"]: self.proj_in = nn.Linear(controlnet_cfg["conv_out_dim"], controlnet_cfg["dim"]) @@ -230,11 +232,13 @@ class WanControlNet(ModelMixin): self.controlnet_mask_embedding = MaskCamEmbed(controlnet_cfg) - def forward(self, render_latent, render_mask, camera_embedding, temb, device): + def forward(self, render_latent, render_mask, camera_embedding, temb, device): controlnet_rotary_emb = self.controlnet_rope(render_latent) - - controlnet_inputs = self.controlnet_patch_embedding(render_latent.to(torch.float32)).to(render_latent.dtype) - controlnet_inputs = controlnet_inputs.to(render_latent.dtype) + controlnet_inputs = self.controlnet_patch_embedding(render_latent.to(torch.float32)) + if not self.quantized: + controlnet_inputs = controlnet_inputs.to(render_latent.dtype) + else: + controlnet_inputs = controlnet_inputs.to(self.base_dtype) controlnet_inputs = controlnet_inputs.flatten(2).transpose(1, 2) diff --git a/uni3c/nodes.py b/uni3c/nodes.py index efc2b29..3720356 100644 --- a/uni3c/nodes.py +++ b/uni3c/nodes.py @@ -69,7 +69,9 @@ class WanVideoUni3C_ControlnetLoader: "num_layers": 20, "add_channels": 7, "mid_channels": 256, - "attention_mode": attention_mode + "attention_mode": attention_mode, + "quantized": True if quantization != "disabled" else False, + "base_dtype": base_dtype } from .controlnet import WanControlNet @@ -94,7 +96,7 @@ class WanVideoUni3C_ControlnetLoader: dtype = torch.float8_e5m2 else: dtype = base_dtype - params_to_keep = {"norm", "head", "time_in", "vector_in", "controlnet_patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter"} + params_to_keep = {"norm", "head", "time_in", "vector_in", "controlnet_patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "proj_in"} log.info("Using accelerate to load and assign controlnet model weights to device...") param_count = sum(1 for _ in controlnet.named_parameters()) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 683e1d8..a359793 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1631,12 +1631,13 @@ class WanModel(ModelMixin, ConfigMixin): if (pcd_data["start"] <= current_step_percentage <= pcd_data["end"]) or \ (pcd_data["end"] > 0 and current_step == 0 and current_step_percentage >= pcd_data["start"]): self.controlnet.to(self.main_device) - pdc_controlnet_states = self.controlnet( - render_latent=render_latent.to(self.main_device, self.controlnet.dtype), - render_mask=pcd_data["render_mask"], - camera_embedding=pcd_data["camera_embedding"], - temb=e.to(self.main_device), - device=self.offload_device) + with torch.autocast(device_type=mm.get_autocast_device(device), dtype=x.dtype, enabled=True): + pdc_controlnet_states = self.controlnet( + render_latent=render_latent.to(self.main_device, self.controlnet.dtype), + render_mask=pcd_data["render_mask"], + camera_embedding=pcd_data["camera_embedding"], + temb=e.to(self.main_device), + device=self.offload_device) self.controlnet.to(self.offload_device) for b, block in enumerate(self.blocks):