partial Uni3C implementation

This commit is contained in:
kijai
2025-05-26 12:35:36 +03:00
parent d9ca90c1f0
commit ef1ed29178
7 changed files with 855 additions and 4 deletions
+9
View File
@@ -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"]
+24 -2
View File
@@ -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
@@ -2327,6 +2328,15 @@ class WanVideoSampler:
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)
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
+89
View File
@@ -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
+263
View File
@@ -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
+237
View File
@@ -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",
}
+206
View File
@@ -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
+25
View File
@@ -1161,6 +1161,7 @@ class WanModel(ModelMixin, ConfigMixin):
audio_proj=None,
audio_context_lens=None,
audio_scale=1.0,
pcd_data=None,
):
r"""
@@ -1207,6 +1208,12 @@ class WanModel(ModelMixin, ConfigMixin):
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:
self.expanded_patch_embedding.to(device)
@@ -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)