diff --git a/__init__.py b/__init__.py index 5634205..e69827e 100644 --- a/__init__.py +++ b/__init__.py @@ -4,17 +4,26 @@ from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPI from .skyreels.nodes import NODE_CLASS_MAPPINGS as SKYREELS_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SKYREELS_NODE_DISPLAY_NAME_MAPPINGS from .fantasytalking.nodes import NODE_CLASS_MAPPINGS as FANTASYTALKING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS from .fun_camera.nodes import NODE_CLASS_MAPPINGS as FUN_CAMERA_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS +from .uni3c.nodes import NODE_CLASS_MAPPINGS as UNI3C_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNI3C_NODE_DISPLAY_NAME_MAPPINGS + +#from .causvid.nodes import NODE_CLASS_MAPPINGS as CAUSVID_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as CAUSVID_NODE_DISPLAY_NAME_MAPPINGS NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(FANTASYTALKING_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(FUN_CAMERA_NODE_CLASS_MAPPINGS) +NODE_CLASS_MAPPINGS.update(UNI3C_NODE_CLASS_MAPPINGS) + +#NODE_CLASS_MAPPINGS.update(CAUSVID_NODE_CLASS_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(SKYREELS_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS) +NODE_DISPLAY_NAME_MAPPINGS.update(UNI3C_NODE_DISPLAY_NAME_MAPPINGS) + +#NODE_DISPLAY_NAME_MAPPINGS.update(CAUSVID_NODE_DISPLAY_NAME_MAPPINGS) __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/nodes.py b/nodes.py index 949d810..9b63708 100644 --- a/nodes.py +++ b/nodes.py @@ -2264,7 +2264,7 @@ class WanVideoSampler: "shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), "force_offload": ("BOOLEAN", {"default": True, "tooltip": "Moves the model to the offload device after sampling"}), - "scheduler": (["unipc", "unipc/beta", "dpm++", "dpm++/beta","dpm++_sde", "dpm++_sde/beta", "euler", "euler/beta", "deis", "lcm", "lcm/beta", "flowmatch_causvid"], + "scheduler": (["unipc", "unipc/beta", "dpm++", "dpm++/beta","dpm++_sde", "dpm++_sde/beta", "euler", "euler/beta", "euler/accvideo", "deis", "lcm", "lcm/beta", "flowmatch_causvid"], { "default": 'unipc' }), @@ -2287,6 +2287,7 @@ class WanVideoSampler: "sigmas": ("SIGMAS", ), "unianimate_poses": ("UNIANIMATE_POSE", ), "fantasytalking_embeds": ("FANTASYTALKING_EMBEDS", ), + "uni3c_embeds": ("UNI3C_EMBEDS", ), } } @@ -2298,7 +2299,7 @@ class WanVideoSampler: def process(self, model, text_embeds, image_embeds, shift, steps, cfg, seed, scheduler, riflex_freq_index, force_offload=True, samples=None, feta_args=None, denoise_strength=1.0, context_options=None, teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default", loop_args=None, - experimental_args=None, sigmas=None, unianimate_poses=None, fantasytalking_embeds=None): + experimental_args=None, sigmas=None, unianimate_poses=None, fantasytalking_embeds=None, uni3c_embeds=None): #assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache." patcher = model model = model.model @@ -2326,7 +2327,16 @@ class WanVideoSampler: if flowedit_args: #seems to work better timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift)) else: - sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) + sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) + elif scheduler in ['euler/accvideo']: + if steps != 50: + raise Exception("Steps must be set to 50 for accvideo scheduler, 10 actual steps are used") + sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta')) + sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) + start_latent_list = [0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50] + sample_scheduler.sigmas = sample_scheduler.sigmas[start_latent_list] + num_inference_steps = len(start_latent_list) - 1 + sample_scheduler.timesteps = timesteps = sample_scheduler.timesteps[start_latent_list[:num_inference_steps]] elif 'dpm++' in scheduler: if 'sde' in scheduler: algorithm_type = "sde-dpmsolver++" @@ -2696,6 +2706,17 @@ class WanVideoSampler: module.onload() elif model["manual_offloading"]: transformer.to(device) + + #uni3c + pcd_data = None + if uni3c_embeds is not None: + transformer.controlnet = uni3c_embeds["controlnet"] + pcd_data = { + "render_latent": uni3c_embeds["render_latent"], + "render_mask": uni3c_embeds["render_mask"], + "camera_embedding": uni3c_embeds["camera_embedding"], + } + #feta if feta_args is not None and latent_video_length > 1: set_enhance_weight(feta_args["weight"]) @@ -2872,6 +2893,7 @@ class WanVideoSampler: 'audio_proj': audio_proj if fantasytalking_embeds is not None else None, 'audio_context_lens': audio_context_lens if fantasytalking_embeds is not None else None, 'audio_scale': audio_scale if fantasytalking_embeds is not None else None, + "pcd_data": pcd_data } batch_size = 1 diff --git a/uni3c/camera.py b/uni3c/camera.py new file mode 100644 index 0000000..77d55bb --- /dev/null +++ b/uni3c/camera.py @@ -0,0 +1,89 @@ +import einops +import torch +import torch.nn.functional as F + + +@torch.amp.autocast("cuda", enabled=False) +def batch_sample_rays(intrinsic, extrinsic, image_h=None, image_w=None): + ''' get rays + Args: + intrinsic: [BF, 3, 3], + extrinsic: [BF, 4, 4], + h, w: int + # normalize: let the first camera R=I + Returns: + rays_o, rays_d: [BF, N, 3] + ''' + + # FIXME: PPU does not support inverse in GPU + device = intrinsic.device + B = intrinsic.shape[0] + + c2w = torch.inverse(extrinsic)[:, :3, :4].to(device) # [BF,3,4] + x = torch.arange(image_w, device=device).float() - 0.5 + y = torch.arange(image_h, device=device).float() + 0.5 + points = torch.stack(torch.meshgrid(x, y, indexing='ij'), -1) + points = einops.repeat(points, 'w h c -> b (h w) c', b=B) + points = torch.cat([points, torch.ones_like(points)[:, :, 0:1]], dim=-1) + directions = points @ intrinsic.inverse().to(device).transpose(-1, -2) * 1 # depth is 1 + + rays_d = F.normalize(directions @ c2w[:, :3, :3].transpose(-1, -2), dim=-1) # [BF,N,3] + rays_o = c2w[..., :3, 3] # [BF, 3] + + rays_o = rays_o[:, None, :].expand_as(rays_d) # [BF, N, 3] + + return rays_o, rays_d + + +@torch.amp.autocast("cuda", enabled=False) +def embed_rays(rays_o, rays_d, nframe): + if len(rays_o.shape) == 4: # [b,f,n,3] + rays_o = einops.rearrange(rays_o, "b f n c -> (b f) n c") + rays_d = einops.rearrange(rays_d, "b f n c -> (b f) n c") + cross_od = torch.cross(rays_o, rays_d, dim=-1) + cam_emb = torch.cat([rays_d, cross_od], dim=-1) + cam_emb = einops.rearrange(cam_emb, "(b f) n c -> b f n c", f=nframe) + return cam_emb + + +@torch.amp.autocast("cuda", enabled=False) +def camera_center_normalization(w2c, nframe, camera_scale=2.0): + # copy from SEVA, w2c: [BF, 4, 4] + # ensure the first view is eye matrix + c2w_view0 = w2c[::nframe].inverse() # [B,4,4] + c2w_view0 = c2w_view0.repeat_interleave(nframe, dim=0) # [BF,4,4] + w2c = c2w_view0 @ w2c + + # camera centering + c2w = torch.linalg.inv(w2c) + camera_dist_2med = torch.norm(c2w[:, :3, 3] - c2w[:, :3, 3].median(0, keepdim=True).values, dim=-1) + valid_mask = camera_dist_2med <= torch.clamp(torch.quantile(camera_dist_2med, 0.97) * 10, max=1e6) + c2w[:, :3, 3] -= c2w[valid_mask, :3, 3].mean(0, keepdim=True) + w2c = torch.linalg.inv(c2w) + + # camera normalization + camera_dists = c2w[:, :3, 3].clone() + translation_scaling_factor = ( + camera_scale + if torch.isclose( + torch.norm(camera_dists[0]), + torch.zeros(1, dtype=camera_dists.dtype, device=camera_dists.device), + atol=1e-5, + ).any() + else (camera_scale / torch.norm(camera_dists[0])) + ) + w2c[:, :3, 3] *= translation_scaling_factor + c2w[:, :3, 3] *= translation_scaling_factor + + return w2c + + +def get_camera_embedding(intrinsic, extrinsic, f, h, w, normalize=True): + if normalize: + extrinsic = camera_center_normalization(extrinsic, nframe=f) + + rays_o, rays_d = batch_sample_rays(intrinsic, extrinsic, image_h=h, image_w=w) + camera_embedding = embed_rays(rays_o, rays_d, nframe=f) + camera_embedding = einops.rearrange(camera_embedding, "b f (h w) c -> b c f h w", h=h, w=w) + + return camera_embedding diff --git a/uni3c/controlnet.py b/uni3c/controlnet.py new file mode 100644 index 0000000..0517e74 --- /dev/null +++ b/uni3c/controlnet.py @@ -0,0 +1,263 @@ +import torch +import torch.nn as nn +from diffusers.models import ModelMixin +from typing import Optional +import torch.nn.functional as F +from diffusers.models.attention_processor import Attention +from diffusers.models.transformers.transformer_wan import WanRotaryPosEmbed +from einops import rearrange + +from ..wanvideo.modules.attention import sageattn_func + +def zero_module(module): + # Zero out the parameters of a module and return it. + for p in module.parameters(): + p.detach().zero_() + return module + + +class SimpleAttnProcessor2_0: + def __init__(self, attention_mode): + self.attention_mode = attention_mode + def __call__( + self, + attn: Attention, + hidden_states: torch.Tensor, + attention_mask: Optional[torch.Tensor] = None, + rotary_emb: Optional[torch.Tensor] = None, + **kwargs + ) -> torch.Tensor: + + query = attn.to_q(hidden_states) + key = attn.to_k(hidden_states) + value = attn.to_v(hidden_states) + + if attn.norm_q is not None: + query = attn.norm_q(query) + if attn.norm_k is not None: + key = attn.norm_k(key) + + query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2) + key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2) + value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2) # [b,head,l,c] + + if rotary_emb is not None: + def apply_rotary_emb(hidden_states: torch.Tensor, freqs: torch.Tensor): + x_rotated = torch.view_as_complex(hidden_states.to(torch.float64).unflatten(3, (-1, 2))) + x_out = torch.view_as_real(x_rotated * freqs).flatten(3, 4) + return x_out.type_as(hidden_states) + + query = apply_rotary_emb(query, rotary_emb) + key = apply_rotary_emb(key, rotary_emb) + + if self.attention_mode == 'sdpa': + hidden_states = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + elif self.attention_mode == 'sageattn': + hidden_states = sageattn_func( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + hidden_states = hidden_states.transpose(1, 2).flatten(2, 3) + hidden_states = hidden_states.type_as(query) + + hidden_states = attn.to_out[0](hidden_states) + hidden_states = attn.to_out[1](hidden_states) + return hidden_states + + +class SimpleCogVideoXLayerNormZero(nn.Module): + def __init__( + self, + conditioning_dim: int, + embedding_dim: int, + elementwise_affine: bool = True, + eps: float = 1e-5, + bias: bool = True, + ) -> None: + super().__init__() + + self.silu = nn.SiLU() + self.linear = nn.Linear(conditioning_dim, 3 * embedding_dim, bias=bias) + self.norm = nn.LayerNorm(embedding_dim, eps=eps, elementwise_affine=elementwise_affine) + + def forward(self, hidden_states: torch.Tensor, temb: torch.Tensor): + shift, scale, gate = self.linear(self.silu(temb)).chunk(3, dim=1) + hidden_states = self.norm(hidden_states) * (1 + scale)[:, None, :] + shift[:, None, :] + return hidden_states, gate[:, None, :] + + +class SingleAttentionBlock(nn.Module): + + def __init__( + self, + dim, + ffn_dim, + num_heads, + time_embed_dim=512, + qk_norm="rms_norm_across_heads", + eps=1e-6, + attention_mode="sdpa", + ): + super().__init__() + self.dim = dim + self.ffn_dim = ffn_dim + self.num_heads = num_heads + self.qk_norm = qk_norm + self.eps = eps + + # layers + self.norm1 = SimpleCogVideoXLayerNormZero( + time_embed_dim, dim, elementwise_affine=True, eps=1e-5, bias=True + ) + self.self_attn = Attention( + query_dim=dim, + heads=num_heads, + kv_heads=num_heads, + dim_head=dim // num_heads, + qk_norm=qk_norm, + eps=eps, + bias=True, + cross_attention_dim=None, + out_bias=True, + processor=SimpleAttnProcessor2_0(attention_mode), + ) + self.norm2 = SimpleCogVideoXLayerNormZero( + time_embed_dim, dim, elementwise_affine=True, eps=1e-5, bias=True + ) + self.ffn = nn.Sequential( + nn.Linear(dim, ffn_dim), + nn.GELU(approximate='tanh'), + nn.Linear(ffn_dim, dim) + ) + + def forward( + self, + hidden_states, + temb, + rotary_emb, + ): + # norm & modulate + norm_hidden_states, gate_msa = self.norm1(hidden_states, temb) + + # attention + attn_hidden_states = self.self_attn(hidden_states=norm_hidden_states, + rotary_emb=rotary_emb) + + hidden_states = hidden_states + gate_msa * attn_hidden_states + + # norm & modulate + norm_hidden_states, gate_ff = self.norm2(hidden_states, temb) + + # feed-forward + ff_output = self.ffn(norm_hidden_states) + + hidden_states = hidden_states + gate_ff * ff_output + + return hidden_states + +class MaskCamEmbed(nn.Module): + def __init__(self, controlnet_cfg) -> None: + super().__init__() + + # padding bug fixed + if controlnet_cfg.get("interp", False): + self.mask_padding = [0, 0, 0, 0, 3, 3] # 左右上下前后, I2V-interp,首尾帧 + else: + self.mask_padding = [0, 0, 0, 0, 3, 0] # 左右上下前后, I2V + add_channels = controlnet_cfg.get("add_channels", 1) + mid_channels = controlnet_cfg.get("mid_channels", 64) + self.mask_proj = nn.Sequential(nn.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8)), + nn.GroupNorm(mid_channels // 8, mid_channels), nn.SiLU()) + self.mask_zero_proj = zero_module(nn.Conv3d(mid_channels, controlnet_cfg["conv_out_dim"], kernel_size=(1, 2, 2), stride=(1, 2, 2))) + + def forward(self, add_inputs: torch.Tensor): + # render_mask.shape [b,c,f,h,w] + warp_add_pad = F.pad(add_inputs, self.mask_padding, mode="constant", value=0) + add_embeds = self.mask_proj(warp_add_pad) # [B,C,F,H,W] + add_embeds = self.mask_zero_proj(add_embeds) + add_embeds = rearrange(add_embeds, "b c f h w -> b (f h w) c") + + return add_embeds + +class WanControlNet(ModelMixin): + def __init__(self, controlnet_cfg): + super().__init__() + + self.rope_max_seq_len = 1024 + self.patch_size = (1, 2, 2) + self.in_channels = controlnet_cfg["in_channels"] + self.dim = controlnet_cfg["dim"] + self.num_heads = controlnet_cfg["num_heads"] + + if controlnet_cfg["conv_out_dim"] != controlnet_cfg["dim"]: + self.proj_in = nn.Linear(controlnet_cfg["conv_out_dim"], controlnet_cfg["dim"]) + else: + self.proj_in = nn.Identity() + + self.controlnet_blocks = nn.ModuleList( + [ + SingleAttentionBlock( + dim=self.dim, + ffn_dim=controlnet_cfg["ffn_dim"], + num_heads=self.num_heads, + time_embed_dim=controlnet_cfg["time_embed_dim"], + qk_norm="rms_norm_across_heads", + attention_mode=controlnet_cfg["attention_mode"], + ) + for _ in range(controlnet_cfg["num_layers"]) + ] + ) + self.proj_out = nn.ModuleList( + [ + zero_module(nn.Linear(self.dim, 5120)) + for _ in range(controlnet_cfg["num_layers"]) + ] + ) + + self.gradient_checkpointing = False + + self.controlnet_rope = WanRotaryPosEmbed(self.dim // self.num_heads, + self.patch_size, self.rope_max_seq_len) + + self.controlnet_patch_embedding = nn.Conv3d( + self.in_channels, + controlnet_cfg["conv_out_dim"], + kernel_size=self.patch_size, + stride=self.patch_size, + dtype=torch.float32 + ) + + self.controlnet_mask_embedding = MaskCamEmbed(controlnet_cfg) + + 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 = controlnet_inputs.flatten(2).transpose(1, 2) + + # additional inputs (mask, camera embedding) + add_inputs = None + if camera_embedding is not None and render_mask is not None: + add_inputs = torch.cat([render_mask, camera_embedding], dim=1) + elif render_mask is not None: + add_inputs = render_mask + + 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 = [] + for i, block in enumerate(self.controlnet_blocks): + hidden_states = block( + hidden_states=hidden_states, + temb=temb, + rotary_emb=controlnet_rotary_emb + ) + controlnet_states.append(self.proj_out[i](hidden_states).to(device)) + + return controlnet_states diff --git a/uni3c/nodes.py b/uni3c/nodes.py new file mode 100644 index 0000000..0494b69 --- /dev/null +++ b/uni3c/nodes.py @@ -0,0 +1,237 @@ + +import torch +from ..utils import log +import comfy.model_management as mm +from comfy.utils import ProgressBar, load_torch_file +from tqdm import tqdm +import gc + +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 + def INPUT_TYPES(s): + return { + "required": { + "model": (folder_paths.get_filename_list("controlnet"), {"tooltip": "These models are loaded from the 'ComfyUI/models/controlnet' -folder",}), + + "base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}), + "quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_e4m3fn_fast_no_ffn'], {"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"}), + "attention_mode": ([ + "sdpa", + "sageattn", + ], {"default": "sdpa"}), + }, + "optional": { + "compile_args": ("WANCOMPILEARGS", ), + #"block_swap_args": ("BLOCKSWAPARGS", ), + } + } + + RETURN_TYPES = ("WANVIDEOCONTROLNET",) + RETURN_NAMES = ("controlnet", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + + def loadmodel(self, model, base_precision, load_device, quantization, attention_mode, compile_args=None): + + device = mm.get_torch_device() + 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] + + controlnet_cfg = { + "in_channels": in_channels, + "conv_out_dim": 5120, + "time_embed_dim": 5120, + "dim": 1024, + "ffn_dim": ffn_dim, + "num_heads": 16, + "num_layers": 20, + "add_channels": 7, + "mid_channels": 256, + "attention_mode": attention_mode + } + + from .controlnet import WanControlNet + + with init_empty_weights(): + controlnet = WanControlNet(controlnet_cfg) + controlnet.eval() + + if quantization == "disabled": + for k, v in sd.items(): + if isinstance(v, torch.Tensor): + if v.dtype == torch.float8_e4m3fn: + quantization = "fp8_e4m3fn" + break + elif v.dtype == torch.float8_e5m2: + quantization = "fp8_e5m2" + break + + if "fp8_e4m3fn" in quantization: + dtype = torch.float8_e4m3fn + elif quantization == "fp8_e5m2": + 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"} + + 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}", + 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: + torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"] + try: + if hasattr(torch, '_dynamo') and hasattr(torch._dynamo, 'config'): + torch._dynamo.config.recompile_limit = compile_args["dynamo_recompile_limit"] + except Exception as e: + log.warning(f"Could not set recompile_limit: {e}") + if compile_args["compile_transformer_blocks_only"]: + 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"]) + + + if load_device == "offload_device" and controlnet.device != offload_device: + log.info(f"Moving controlnet model from {controlnet.device} to {offload_device}") + controlnet.to(offload_device) + gc.collect() + mm.soft_empty_cache() + + return (controlnet,) + +class WanVideoUni3C_embeds: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "controlnet": ("WANVIDEOCONTROLNET",), + "render_latent": ("LATENT",), + # "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}), + # "vace_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply VACE"}), + # "vace_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply VACE"}), + }, + "optional": { + "render_mask": ("MASK",), + }, + } + + RETURN_TYPES = ("UNI3C_EMBEDS", ) + RETURN_NAMES = ("uni3c_embeds",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + + def process(self, controlnet, render_latent, render_mask=None): + + device = mm.get_torch_device() + + latent_mask = None + + latents = render_latent["samples"] + nframe = latents.shape[2] * 4 + height = latents.shape[3] * 8 + width = latents.shape[4] * 8 + + if render_mask is not None: + 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")) + # w2cs = torch.tensor(np.array(cam_info["extrinsic"]), dtype=torch.float32, device=device) + # intrinsic = torch.tensor(np.array(cam_info["intrinsic"]), dtype=torch.float32, device=device) + # intrinsic[0, :] = intrinsic[0, :] / cam_info["width"] * width + # intrinsic[1, :] = intrinsic[1, :] / cam_info["height"] * height + # intrinsic = intrinsic[None].repeat(nframe, 1, 1) + + # from .utils import build_cameras, set_initial_camera, traj_map + + # focal_length = 1.0 + # start_elevation = 5.0 + # depth_avg = 0.5 + # traj_type = "orbit" + # cam_traj, x_offset, y_offset, z_offset, d_theta, d_phi, d_r = traj_map(traj_type) + # focallength_px = focal_length * width + + # K = torch.tensor([[focallength_px, 0, width / 2], + # [0, focallength_px, height / 2], + # [0, 0, 1]], dtype=torch.float32) + # 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, + # c2w_0=c2w_0, + # intrinsic=intrinsic, + # nframe=nframe, + # focal_length=focal_length, + # d_theta=d_theta, + # d_phi=d_phi, + # d_r=d_r, + # radius=depth_avg, + # x_offset=x_offset, + # 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) + + uni3c_embeds = { + "controlnet": controlnet, + "render_latent": latents.to(device), + "render_mask": latent_mask, + "camera_embedding": None + } + + return (uni3c_embeds,) + +NODE_CLASS_MAPPINGS = { + "WanVideoUni3C_ControlnetLoader": WanVideoUni3C_ControlnetLoader, + "WanVideoUni3C_embeds": WanVideoUni3C_embeds, + } +NODE_DISPLAY_NAME_MAPPINGS = { + "WanVideoUni3C_ControlnetLoader": "WanVideo Uni3C Controlnet Loader", + "WanVideoUni3C_embeds": "WanVideo Uni3C Embeds", + } + + \ No newline at end of file diff --git a/uni3c/utils.py b/uni3c/utils.py new file mode 100644 index 0000000..7ba7f2c --- /dev/null +++ b/uni3c/utils.py @@ -0,0 +1,206 @@ +import imageio +import numpy as np +import torch +from PIL import Image +from scipy.interpolate import UnivariateSpline +from scipy.interpolate import interp1d + + +def load_video(video_path): + reader = imageio.get_reader(video_path) + total_frames = reader.count_frames() + frames = [] + for i in range(total_frames): + frame = reader.get_data(i) + frames.append(Image.fromarray(frame)) + + reader.close() + + return frames + + +def points_padding(points): + padding = torch.ones_like(points)[..., 0:1] + points = torch.cat([points, padding], dim=-1) + return points + + +def np_points_padding(points): + padding = np.ones_like(points)[..., 0:1] + points = np.concatenate([points, padding], axis=-1) + return points + + +def txt_interpolation(input_list, n, mode='smooth'): + x = np.linspace(0, 1, len(input_list)) + if mode == 'smooth': + f = UnivariateSpline(x, input_list, k=3) + elif mode == 'linear': + f = interp1d(x, input_list) + else: + raise KeyError(f"Invalid txt interpolation mode: {mode}") + xnew = np.linspace(0, 1, n) + ynew = f(xnew) + return ynew + + +def traj_map(traj_type): + # pre-defined trajectories + if traj_type == "free1": # Zoom out and rotate to the upper left + cam_traj = "free" + x_offset = 0.0 + y_offset = 0.0 + z_offset = 0.0 + d_theta = -15.0 + d_phi = 45.0 + d_r = 1.6 + elif traj_type == "free2": # Rotate to the right horizontally + cam_traj = "free" + x_offset = -0.05 + y_offset = 0.0 + z_offset = 0.0 + d_theta = 0.0 + d_phi = -60.0 + d_r = 1.0 + elif traj_type == "free3": # Move back to the left + cam_traj = "free" + x_offset = -0.25 + y_offset = 0.0 + z_offset = 0.0 + d_theta = 0.0 + d_phi = 0.0 + d_r = 1.7 + elif traj_type == "free4": # Rotate and approach to the upper right + cam_traj = "free" + x_offset = 0.0 + y_offset = 0.0 + z_offset = 0.0 + d_theta = -15.0 + d_phi = -60.0 + d_r = 0.75 + elif traj_type == "free5": # Large-angle camera movement to the upper right + cam_traj = "free" + x_offset = 0.0 + y_offset = 0.0 + z_offset = 0.0 + d_theta = -15.0 + d_phi = -120.0 + d_r = 1.6 + elif traj_type == "swing1": # Swing shot 1 + cam_traj = "swing1" + x_offset = 0.0 + y_offset = 0.0 + z_offset = 0.0 + d_theta = 0.0 + d_phi = 0.0 + d_r = 1.0 + elif traj_type == "swing2": # Swing shot 2 + cam_traj = "swing2" + x_offset = 0.0 + y_offset = 0.0 + z_offset = 0.0 + d_theta = 0.0 + d_phi = 0.0 + d_r = 1.0 + elif traj_type == "orbit": # 360-degree counterclockwise rotation + cam_traj = "free" + x_offset = 0.0 + y_offset = 0.0 + z_offset = 0.0 + d_theta = 0.0 + d_phi = -360.0 + d_r = 1.0 + else: + raise NotImplementedError + return cam_traj, x_offset, y_offset, z_offset, d_theta, d_phi, d_r + + +def set_initial_camera(start_elevation, radius): + c2w_0 = torch.tensor([[1, 0, 0, 0], + [0, 1, 0, 0], + [0, 0, 1, -radius], + [0, 0, 0, 1]], dtype=torch.float32) + elevation_rad = np.deg2rad(start_elevation) + R_elevation = torch.tensor([[1, 0, 0, 0], + [0, np.cos(-elevation_rad), -np.sin(-elevation_rad), 0], + [0, np.sin(-elevation_rad), np.cos(-elevation_rad), 0], + [0, 0, 0, 1]], dtype=torch.float32) + c2w_0 = R_elevation @ c2w_0 + w2c_0 = c2w_0.inverse() + + return w2c_0, c2w_0 + + +def build_cameras(cam_traj, w2c_0, c2w_0, intrinsic, nframe, focal_length, + d_theta, d_phi, d_r, radius, x_offset, y_offset, z_offset): + # build camera viewpoints according to d_theta,d_phi, d_r + # return: w2cs:[V,4,4], c2ws:[V,4,4], intrinsic:[V,3,3] + if intrinsic.ndim == 2: + intrinsic = intrinsic[None].repeat(nframe, 1, 1) + + c2ws = [c2w_0] + w2cs = [w2c_0] + d_thetas, d_phis, d_rs = [], [], [] + x_offsets, y_offsets, z_offsets = [], [], [] + if cam_traj == "free": + for i in range(nframe - 1): + coef = (i + 1) / (nframe - 1) + d_thetas.append(d_theta * coef) + d_phis.append(d_phi * coef) + d_rs.append(coef * d_r + (1 - coef) * 1.0) + x_offsets.append(radius * x_offset * ((i + 1) / nframe)) + y_offsets.append(radius * y_offset * ((i + 1) / nframe)) + z_offsets.append(radius * z_offset * ((i + 1) / nframe)) + elif cam_traj == "swing1": + phis__ = [0, -5, -25, -30, -20, -8, 0] + thetas__ = [0, -8, -12, -20, -17, -12, -5, -2, 1, 5, 3, 1, 0] + rs__ = [0, 0.2] + d_phis = txt_interpolation(phis__, nframe, mode='smooth') + d_phis[0] = phis__[0] + d_phis[-1] = phis__[-1] + d_thetas = txt_interpolation(thetas__, nframe, mode='smooth') + d_thetas[0] = thetas__[0] + d_thetas[-1] = thetas__[-1] + d_rs = txt_interpolation(rs__, nframe, mode='linear') + d_rs = 1.0 + d_rs + elif cam_traj == "swing2": + phis__ = [0, 5, 25, 30, 20, 10, 0] + thetas__ = [0, -5, -14, -11, 0, 1, 5, 3, 0] + rs__ = [0, -0.03, -0.1, -0.2, -0.17, -0.1, 0] + d_phis = txt_interpolation(phis__, nframe, mode='smooth') + d_phis[0] = phis__[0] + d_phis[-1] = phis__[-1] + d_thetas = txt_interpolation(thetas__, nframe, mode='smooth') + d_thetas[0] = thetas__[0] + d_thetas[-1] = thetas__[-1] + d_rs = txt_interpolation(rs__, nframe, mode='smooth') + d_rs = 1.0 + d_rs + else: + raise NotImplementedError("Unknown trajectory type...") + + for i in range(nframe - 1): + d_theta_rad = np.deg2rad(d_thetas[i]) + R_theta = torch.tensor([[1, 0, 0, 0], + [0, np.cos(d_theta_rad), -np.sin(d_theta_rad), 0], + [0, np.sin(d_theta_rad), np.cos(d_theta_rad), 0], + [0, 0, 0, 1]], dtype=torch.float32) + d_phi_rad = np.deg2rad(d_phis[i]) + R_phi = torch.tensor([[np.cos(d_phi_rad), 0, np.sin(d_phi_rad), 0], + [0, 1, 0, 0], + [-np.sin(d_phi_rad), 0, np.cos(d_phi_rad), 0], + [0, 0, 0, 1]], dtype=torch.float32) + c2w_1 = R_phi @ R_theta @ c2w_0 + if i < len(x_offsets) and i < len(y_offsets) and i < len(z_offsets): + c2w_1[:3, -1] += torch.tensor([x_offsets[i], y_offsets[i], z_offsets[i]]) + c2w_1[:3, -1] *= d_rs[i] + w2c_1 = c2w_1.inverse() + c2ws.append(c2w_1) + w2cs.append(w2c_1) + + intrinsic[i + 1, :2, :2] = intrinsic[i + 1, :2, :2] * focal_length * ((i + 1) / nframe) + \ + intrinsic[i + 1, :2, :2] * ((nframe - (i + 1)) / nframe) + + w2cs = torch.stack(w2cs, dim=0) + c2ws = torch.stack(c2ws, dim=0) + + return w2cs, c2ws, intrinsic diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 3dd1710..8ad322b 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1161,6 +1161,7 @@ class WanModel(ModelMixin, ConfigMixin): audio_proj=None, audio_context_lens=None, audio_scale=1.0, + pcd_data=None, ): r""" @@ -1190,7 +1191,7 @@ class WanModel(ModelMixin, ConfigMixin): freqs = freqs.to(device) _, F, H, W = x[0].shape - + # Construct blockwise causal attn mask if self.attention_mode == 'flex_attention' and current_step == 0: self.block_mask = self._prepare_blockwise_causal_attn_mask( @@ -1206,6 +1207,12 @@ class WanModel(ModelMixin, ConfigMixin): if random_ref_emb is not None: y[0] = y[0] + random_ref_emb * unianim_data["strength"] x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)] + + #uni3c controlnet + + if pcd_data is not None: + hidden_states = x[0].unsqueeze(0).clone().float() + render_latent = torch.cat([hidden_states[:, :20], pcd_data["render_latent"]], dim=1) # embeddings if control_lora_enabled: @@ -1414,6 +1421,18 @@ class WanModel(ModelMixin, ConfigMixin): kwargs['vace_hints'] = vace_hint_list kwargs['vace_context_scale'] = vace_scale_list + #uni3c controlnet + if pcd_data is not None: + self.controlnet.to(self.main_device) + controlnet_states = self.controlnet( + render_latent=render_latent.to(self.main_device), + 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): if self.slg_blocks is not None: if b in self.slg_blocks and is_uncond: @@ -1422,6 +1441,12 @@ class WanModel(ModelMixin, ConfigMixin): if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: block.to(self.main_device) x = block(x, **kwargs) + + #uni3c controlnet + if pcd_data is not None: + if b < len(controlnet_states): + x += controlnet_states[b].to(x.device) + if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: block.to(self.offload_device, non_blocking=self.use_non_blocking)