From 93f7af6dc8559e2f6815caa67fd0982e4d8940dd Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 17 Dec 2025 16:22:37 +0200 Subject: [PATCH] Make uni3c offloading optional --- nodes_sampler.py | 10 ++----- uni3c/controlnet.py | 6 ++-- uni3c/nodes.py | 62 ++++++++++++++++++--------------------- wanvideo/modules/model.py | 14 +++++---- 4 files changed, 42 insertions(+), 50 deletions(-) diff --git a/nodes_sampler.py b/nodes_sampler.py index 2ac249a..cdbda4c 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -888,16 +888,10 @@ class WanVideoSampler: if uni3c_embeds is not None: transformer.uni3c_controlnet = uni3c_embeds["controlnet"] render_latent = uni3c_embeds["render_latent"].to(device) + uni3c_data = uni3c_embeds.copy() if render_latent.shape != noise.shape: render_latent = torch.nn.functional.interpolate(render_latent, size=(noise.shape[1], noise.shape[2], noise.shape[3]), mode='trilinear', align_corners=False) - uni3c_data = { - "render_latent": render_latent, - "render_mask": uni3c_embeds["render_mask"], - "camera_embedding": uni3c_embeds["camera_embedding"], - "controlnet_weight": uni3c_embeds["controlnet_weight"], - "start": uni3c_embeds["start"], - "end": uni3c_embeds["end"], - } + uni3c_data["render_latent"] = render_latent # Enhance-a-video (feta) if feta_args is not None and latent_video_length > 1: diff --git a/uni3c/controlnet.py b/uni3c/controlnet.py index dd2e713..2c8707d 100644 --- a/uni3c/controlnet.py +++ b/uni3c/controlnet.py @@ -334,7 +334,7 @@ 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, out_device): controlnet_rotary_emb = self.controlnet_rope(render_latent) controlnet_inputs = self.controlnet_patch_embedding(render_latent.to(torch.float32)) if not self.quantized: @@ -354,7 +354,7 @@ class WanControlNet(ModelMixin): if add_inputs is not None: add_inputs = self.controlnet_mask_embedding(add_inputs) controlnet_inputs = controlnet_inputs + add_inputs - + hidden_states = self.proj_in(controlnet_inputs) controlnet_states = [] @@ -364,6 +364,6 @@ class WanControlNet(ModelMixin): temb=temb, rotary_emb=controlnet_rotary_emb ) - controlnet_states.append(self.proj_out[i](hidden_states).to(device)) + controlnet_states.append(self.proj_out[i](hidden_states).to(out_device)) return controlnet_states diff --git a/uni3c/nodes.py b/uni3c/nodes.py index 2e5d114..844f553 100644 --- a/uni3c/nodes.py +++ b/uni3c/nodes.py @@ -10,8 +10,6 @@ from accelerate import init_empty_weights from accelerate.utils import set_module_tensor_to_device import folder_paths -import json -import numpy as np class WanVideoUni3C_ControlnetLoader: @classmethod @@ -22,7 +20,7 @@ class WanVideoUni3C_ControlnetLoader: "base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}), "quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e5m2'], {"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"}), + "load_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}), "attention_mode": ([ "sdpa", "sageattn", @@ -45,17 +43,17 @@ class WanVideoUni3C_ControlnetLoader: offload_device = mm.unet_offload_device() transformer_load_device = device if load_device == "main_device" else offload_device - + base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision] - + model_path = folder_paths.get_full_path_or_raise("controlnet", model) - + sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True) if not "controlnet_patch_embedding.weight" in sd: raise ValueError("Invalid ControlNet model") - + in_channels = sd["controlnet_patch_embedding.weight"].shape[1] ffn_dim = sd["controlnet_blocks.0.ffn.0.bias"].shape[0] @@ -79,7 +77,7 @@ class WanVideoUni3C_ControlnetLoader: with init_empty_weights(): controlnet = WanControlNet(controlnet_cfg) controlnet.eval() - + if quantization == "disabled": for k, v in sd.items(): if isinstance(v, torch.Tensor): @@ -97,18 +95,18 @@ class WanVideoUni3C_ControlnetLoader: else: dtype = base_dtype 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()) - for name, param in tqdm(controlnet.named_parameters(), - desc=f"Loading transformer parameters to {transformer_load_device}", + for name, param in tqdm(controlnet.named_parameters(), + desc=f"Loading transformer parameters to {transformer_load_device}", total=param_count, leave=True): dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype if "controlnet_patch_embedding" in name: dtype_to_use = torch.float32 set_module_tensor_to_device(controlnet, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) - + del sd if compile_args is not None: @@ -123,8 +121,8 @@ class WanVideoUni3C_ControlnetLoader: for i, block in enumerate(controlnet.controlnet_blocks): controlnet.controlnet_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) else: - controlnet = torch.compile(controlnet, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) - + controlnet = torch.compile(controlnet, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + if load_device == "offload_device" and controlnet.device != offload_device: log.info(f"Moving controlnet model from {controlnet.device} to {offload_device}") @@ -146,6 +144,7 @@ class WanVideoUni3C_embeds: "optional": { "render_latent": ("LATENT",), "render_mask": ("MASK", {"tooltip": "NOT IMPLEMENTED!"}), + "offload": ("BOOLEAN", {"default": True, "tooltip": "If enabled, the controlnet model will be offloaded before main model block processing to save VRAM."}), }, } @@ -154,9 +153,7 @@ class WanVideoUni3C_embeds: FUNCTION = "process" CATEGORY = "WanVideoWrapper" - def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None): - - device = mm.get_torch_device() + def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None, offload=True): latent_mask = latents = None if render_latent is not None: @@ -164,17 +161,17 @@ class WanVideoUni3C_embeds: # nframe = latents.shape[2] * 4 # height = latents.shape[3] * 8 # width = latents.shape[4] * 8 - + if render_mask is not None: raise NotImplementedError("render_mask is not implemented at this time") - mask = torch.nn.functional.interpolate( - render_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W] - size=(nframe, height, width), - mode='trilinear', - align_corners=False - ).squeeze(0) - latent_mask = mask.unsqueeze(0).to(device) - log.info(f"latent mask shape {latent_mask.shape}") + # mask = torch.nn.functional.interpolate( + # render_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W] + # size=(nframe, height, width), + # mode='trilinear', + # align_corners=False + # ).squeeze(0) + # latent_mask = mask.unsqueeze(0).to(device) + # log.info(f"latent mask shape {latent_mask.shape}") # # load camera # cam_info = json.load(open(f"{render_path}/cam_info.json")) @@ -199,7 +196,7 @@ class WanVideoUni3C_embeds: # K_inv = K.inverse() # intrinsic = K[None].repeat(nframe, 1, 1) - + # w2c_0, c2w_0 = set_initial_camera(start_elevation, depth_avg) # w2cs, c2ws, intrinsic = build_cameras(cam_traj=cam_traj, # w2c_0=w2c_0, @@ -215,7 +212,7 @@ class WanVideoUni3C_embeds: # y_offset=y_offset, # z_offset=z_offset) - + # from .camera import get_camera_embedding # camera_embedding = get_camera_embedding(intrinsic, w2cs, nframe, height, width, normalize=True) #print("camera embedding shape", camera_embedding.shape) @@ -227,11 +224,12 @@ class WanVideoUni3C_embeds: "end": end_percent, "render_latent": latents, "render_mask": latent_mask, - "camera_embedding": None + "camera_embedding": None, + "offload": offload, } - + return (uni3c_embeds,) - + NODE_CLASS_MAPPINGS = { "WanVideoUni3C_ControlnetLoader": WanVideoUni3C_ControlnetLoader, "WanVideoUni3C_embeds": WanVideoUni3C_embeds, @@ -240,5 +238,3 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoUni3C_ControlnetLoader": "WanVideo Uni3C Controlnet Loader", "WanVideoUni3C_embeds": "WanVideo Uni3C Embeds", } - - \ No newline at end of file diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 86e72a4..b2d19f8 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -3013,15 +3013,17 @@ class WanModel(torch.nn.Module): if uni3c_data is not None: if (uni3c_data["start"] <= current_step_percentage <= uni3c_data["end"]) or \ (uni3c_data["end"] > 0 and current_step == 0 and current_step_percentage >= uni3c_data["start"]): - self.uni3c_controlnet.to(self.main_device) - with torch.autocast(device_type=mm.get_autocast_device(device), dtype=self.base_dtype, enabled=True): + if uni3c_data["offload"] or self.uni3c_controlnet.device != self.main_device: + self.uni3c_controlnet.to(self.main_device) + with torch.autocast(device_type=mm.get_autocast_device(device), dtype=self.base_dtype, enabled=self.uni3c_controlnet.quantized): uni3c_controlnet_states = self.uni3c_controlnet( render_latent=render_latent.to(self.main_device, self.uni3c_controlnet.dtype), - render_mask=uni3c_data["render_mask"], - camera_embedding=uni3c_data["camera_embedding"], + render_mask=uni3c_data["render_mask"], + camera_embedding=uni3c_data["camera_embedding"], temb=e.to(self.main_device), - device=self.offload_device) - self.uni3c_controlnet.to(self.offload_device) + out_device=self.offload_device if uni3c_data["offload"] else device) + if uni3c_data["offload"]: + self.uni3c_controlnet.to(self.offload_device) # Asynchronous block offloading with CUDA streams and events if torch.cuda.is_available():