diff --git a/LongVie2/modules.py b/LongVie2/modules.py new file mode 100644 index 0000000..7d89789 --- /dev/null +++ b/LongVie2/modules.py @@ -0,0 +1,199 @@ +import torch +import torch.nn as nn +from einops import rearrange + +from ..wanvideo.modules.attention import attention + +def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor): + return (x * (1 + scale) + shift) + + +def sinusoidal_embedding_1d(dim, position): + sinusoid = torch.outer(position.type(torch.float64), torch.pow( + 10000, -torch.arange(dim//2, dtype=torch.float64, device=position.device).div(dim//2))) + x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1) + return x.to(position.dtype) + + +def precompute_freqs_cis_3d(dim: int, end: int = 1024, theta: float = 10000.0): + # 3d rope precompute + f_freqs_cis = precompute_freqs_cis(dim - 2 * (dim // 3), end, theta) + h_freqs_cis = precompute_freqs_cis(dim // 3, end, theta) + w_freqs_cis = precompute_freqs_cis(dim // 3, end, theta) + return f_freqs_cis, h_freqs_cis, w_freqs_cis + + +def precompute_freqs_cis(dim: int, end: int = 1024, theta: float = 10000.0): + # 1d rope precompute + freqs = 1.0 / (theta ** (torch.arange(0, dim, 2) + [: (dim // 2)].double() / dim)) + freqs = torch.outer(torch.arange(end, device=freqs.device), freqs) + freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 + return freqs_cis + + +def rope_apply(x, freqs, num_heads): + x = rearrange(x, "b s (n d) -> b s n d", n=num_heads) + x_out = torch.view_as_complex(x.to(torch.float64).reshape( + x.shape[0], x.shape[1], x.shape[2], -1, 2)) + x_out = torch.view_as_real(x_out * freqs).flatten(2) + return x_out.to(x.dtype) + + +class RMSNorm(nn.Module): + def __init__(self, dim, eps=1e-5): + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim)) + + def norm(self, x): + return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps) + + def forward(self, x): + dtype = x.dtype + return self.norm(x.float()).to(dtype) * self.weight + + +class AttentionModule(nn.Module): + def __init__(self, num_heads, head_dim): + super().__init__() + self.num_heads = num_heads + self.head_dim = head_dim + + def forward(self, q, k, v): + b, n, d = q.size(0), self.num_heads, self.head_dim + x = attention( + q.view(b, -1, n, d), + k.view(b, -1, n, d), + v.view(b, -1, n, d) + ) + return x.flatten(2) + + +class SelfAttention(nn.Module): + def __init__(self, dim: int, num_heads: int, eps: float = 1e-6): + super().__init__() + self.dim = dim + self.num_heads = num_heads + self.head_dim = dim // num_heads + + self.q = nn.Linear(dim, dim) + self.k = nn.Linear(dim, dim) + self.v = nn.Linear(dim, dim) + self.o = nn.Linear(dim, dim) + self.norm_q = RMSNorm(dim, eps=eps) + self.norm_k = RMSNorm(dim, eps=eps) + + self.attn = AttentionModule(self.num_heads, self.head_dim) + + def forward(self, x, freqs): + q = self.norm_q(self.q(x)) + k = self.norm_k(self.k(x)) + v = self.v(x) + q = rope_apply(q, freqs, self.num_heads) + k = rope_apply(k, freqs, self.num_heads) + x = self.attn(q, k, v) + return self.o(x) + + +class CrossAttention(nn.Module): + def __init__(self, dim: int, num_heads: int, eps: float = 1e-6, clip_fea: torch.Tensor = None): + super().__init__() + self.dim = dim + self.num_heads = num_heads + self.head_dim = dim // num_heads + + self.q = nn.Linear(dim, dim) + self.k = nn.Linear(dim, dim) + self.v = nn.Linear(dim, dim) + self.o = nn.Linear(dim, dim) + self.norm_q = RMSNorm(dim, eps=eps) + self.norm_k = RMSNorm(dim, eps=eps) + + + self.k_img = nn.Linear(dim, dim) + self.v_img = nn.Linear(dim, dim) + self.norm_k_img = RMSNorm(dim, eps=eps) + + self.attn = AttentionModule(self.num_heads, self.head_dim) + + def forward(self, x: torch.Tensor, y: torch.Tensor, clip_fea: torch.Tensor = None): + ctx = y + q = self.norm_q(self.q(x)) + k = self.norm_k(self.k(ctx)) + v = self.v(ctx) + x = self.attn(q, k, v) + if clip_fea is not None: + k_img = self.norm_k_img(self.k_img(clip_fea)) + v_img = self.v_img(clip_fea) + y = self.attn(q, k_img, v_img) + x = x + y + return self.o(x) + + +class GateModule(nn.Module): + def __init__(self,): + super().__init__() + + def forward(self, x, gate, residual): + return x + gate * residual + +class DiTBlock(nn.Module): + def __init__(self, dim: int, num_heads: int, ffn_dim: int, eps: float = 1e-6): + super().__init__() + self.dim = dim + self.num_heads = num_heads + self.ffn_dim = ffn_dim + + self.self_attn = SelfAttention(dim, num_heads, eps) + self.cross_attn = CrossAttention(dim, num_heads, eps) + self.norm1 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False) + self.norm2 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False) + self.norm3 = nn.LayerNorm(dim, eps=eps) + self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU( + approximate='tanh'), nn.Linear(ffn_dim, dim)) + self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) + self.gate = GateModule() + + def forward(self, x, context, t_mod, freqs, clip_fea=None): + has_seq = len(t_mod.shape) == 4 + chunk_dim = 2 if has_seq else 1 + # msa: multi-head self-attention mlp: multi-layer perceptron + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( + self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod).chunk(6, dim=chunk_dim) + if has_seq: + shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = ( + shift_msa.squeeze(2), scale_msa.squeeze(2), gate_msa.squeeze(2), + shift_mlp.squeeze(2), scale_mlp.squeeze(2), gate_mlp.squeeze(2), + ) + input_x = modulate(self.norm1(x), shift_msa, scale_msa) + x = self.gate(x, gate_msa, self.self_attn(input_x, freqs)) + x = x + self.cross_attn(self.norm3(x), context, clip_fea=clip_fea) + input_x = modulate(self.norm2(x), shift_mlp, scale_mlp) + x = self.gate(x, gate_mlp, self.ffn(input_x)) + return x + + +class WanModelDualControl(torch.nn.Module): + def __init__(self, dim: int, ffn_dim: int, eps: float, num_heads: int, control_layers = 12): + super().__init__() + self.control_layers = control_layers + self.control_blocks_dense = nn.ModuleList([ + DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps) + for _ in range(self.control_layers) + ]) + + self.control_blocks_sparse = nn.ModuleList([ + DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps) + for _ in range(self.control_layers) + ]) + + self.control_initial_combine_linear_dense = torch.nn.Linear(dim, dim//2) + self.control_initial_combine_linear_sparse = torch.nn.Linear(dim, dim//2) + + self.control_text_linear = torch.nn.Linear(dim, dim//2) + self.control_t_mod = torch.nn.Linear(dim, dim//2) + + self.control_combine_linears = torch.nn.ModuleList([torch.nn.Linear(dim//2, dim) for _ in range(self.control_layers)]) + head_dim = dim // num_heads + self.freqs = precompute_freqs_cis_3d(head_dim) diff --git a/LongVie2/nodes.py b/LongVie2/nodes.py new file mode 100644 index 0000000..6383d76 --- /dev/null +++ b/LongVie2/nodes.py @@ -0,0 +1,88 @@ +import torch +from ..utils import log +import comfy.model_management as mm + +device = mm.get_torch_device() +offload_device = mm.unet_offload_device() + +class WanVideoAddDualControlEmbeds: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "embeds": ("WANVIDIMAGE_EMBEDS",), + "vae": ("WANVAE", {"tooltip": "VAE model"}), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}), + "first_frame_noise_level": ("FLOAT", {"default": 0.925926, "min": 0.0, "max": 1.0, "step": 0.000001, "tooltip": "Noise level for the first frame when using previous frames"}), + }, + "optional": { + "dense": ("IMAGE", {"tooltip": "Dense control signal (depth) video input"}), + "sparse": ("IMAGE", {"tooltip": "Sparse control signal (tracks) video input"}), + "prev_images": ("IMAGE", {"tooltip": "Previous frames for temporal consistency, default is 8 frames"}), + } + } + + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",) + RETURN_NAMES = ("image_embeds",) + FUNCTION = "add" + CATEGORY = "WanVideoWrapper" + + def add(self, embeds, vae, strength, start_percent, end_percent, first_frame_noise_level, dense=None, sparse=None, prev_images=None): + updated = dict(embeds) + updated.setdefault("dual_control", {}) + + if dense is None and sparse is None: + raise ValueError("At least one of dense or sparse inputs must be provided.") + + num_frames = dense.shape[0] if dense is not None else sparse.shape[0] + height = dense.shape[1] if dense is not None else sparse.shape[1] + width = dense.shape[2] if dense is not None else sparse.shape[2] + msk = torch.ones(1, num_frames, height//8, width//8, device=device) + msk[:, 1:] = 0 + msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1) + msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8) + msk = msk.transpose(1, 2) + + dense_input_latent = sparse_input_latent = None + + vae.to(device) + if dense is not None: + dense_images = dense[..., :3].permute(3, 0, 1, 2) * 2 - 1 + dense_images = 1 - dense_images # Invert colors for depth to match the usual range in comfy + dense_video_latent = vae.encode([dense_images.to(device, vae.dtype)], device, tiled=False) + dense_first = (dense_images[:, :1]).to(device, vae.dtype) + vae_input_dense = torch.cat([dense_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1) + dense_concat_latent = vae.encode([vae_input_dense], device, tiled=False) + dense_concat_latent = torch.cat([msk, dense_concat_latent], dim=1) + dense_input_latent = torch.cat([dense_video_latent, dense_concat_latent],dim=1) + if sparse is not None: + sparse_images = sparse[..., :3].permute(3, 0, 1, 2) * 2 - 1 + sparse_video_latent = vae.encode([sparse_images.to(device, vae.dtype)], device, tiled=False) + sparse_first = (sparse_images[:, :1]).to(device, vae.dtype) + vae_input_sparse = torch.cat([sparse_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1) + sparse_concat_latent = vae.encode([vae_input_sparse], device, tiled=False) + sparse_concat_latent = torch.cat([msk, sparse_concat_latent], dim=1) + sparse_input_latent = torch.cat([sparse_video_latent, sparse_concat_latent],dim=1) + + if prev_images is not None: + prev_images = prev_images[..., :3].permute(3, 0, 1, 2) * 2 - 1 + prev_video_latent = vae.encode([prev_images.to(device, vae.dtype)], device, tiled=False) + updated["dual_control"]["prev_latent"] = prev_video_latent[0] + + vae.to(offload_device) + updated["dual_control"]["dense_input_latent"] = dense_input_latent + updated["dual_control"]["sparse_input_latent"] = sparse_input_latent + updated["dual_control"]["strength"] = strength + updated["dual_control"]["start_percent"] = start_percent + updated["dual_control"]["end_percent"] = end_percent + updated["dual_control"]["first_frame_noise_level"] = first_frame_noise_level + return (updated,) + + +NODE_CLASS_MAPPINGS = { + "WanVideoAddDualControlEmbeds": WanVideoAddDualControlEmbeds, + } +NODE_DISPLAY_NAME_MAPPINGS = { + "WanVideoAddDualControlEmbeds": "WanVideo Add Dual Control Embeds", + } diff --git a/__init__.py b/__init__.py index 4efb858..f075e69 100644 --- a/__init__.py +++ b/__init__.py @@ -49,6 +49,7 @@ OPTIONAL_MODULES = [ (".WanMove.nodes", "WanMove"), (".SCAIL.nodes", "SCAIL"), (".LongCat.nodes", "LongCat"), + (".LongVie2.nodes", "LongVie2"), ] def register_nodes(module_path: str, name: str, optional: bool) -> None: @@ -71,4 +72,4 @@ for module_path, name in REQUIRED_MODULES: for module_path, name in OPTIONAL_MODULES: register_nodes(module_path, name, optional=True) -__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/custom_linear.py b/custom_linear.py index b80be20..9fe4762 100644 --- a/custom_linear.py +++ b/custom_linear.py @@ -56,7 +56,7 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s module_prefix = module_prefix.replace("_orig_mod.", "") _replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args, modules_to_not_convert) - if isinstance(module, nn.Linear) and "loras" not in module_prefix and name not in modules_to_not_convert: + if isinstance(module, nn.Linear) and "loras" not in module_prefix and "dual_controller" not in module_prefix and name not in modules_to_not_convert: weight_key = module_prefix + "weight" if weight_key not in state_dict: continue diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 68239da..6c45e7a 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1605,6 +1605,24 @@ class WanVideoModelLoader: block.ref_attn_v_img = nn.Linear(in_features, out_features) block.ref_attn_norm_k_img = WanRMSNorm(out_features, eps=1e-6) + if "blocks.0.control_blocks_dense.cross_attn.k.weight" in sd: + log.info("LongVie2 model detected, patching model...") + from .LongVie2.modules import WanModelDualControl + control_layers = 12 + with init_empty_weights(): + dual_controller = WanModelDualControl(dim=5120, ffn_dim=13824, eps=1e-06, num_heads=40, control_layers=control_layers) + for b in range(control_layers): + transformer.blocks[b].control_blocks_dense = dual_controller.control_blocks_dense[b] + transformer.blocks[b].control_blocks_sparse = dual_controller.control_blocks_sparse[b] + transformer.blocks[b].control_combine_linears = dual_controller.control_combine_linears[b] + transformer.dual_controller = nn.Module() + transformer.dual_controller.control_initial_combine_linear_dense = dual_controller.control_initial_combine_linear_dense + transformer.dual_controller.control_initial_combine_linear_sparse = dual_controller.control_initial_combine_linear_sparse + transformer.dual_controller.control_t_mod = dual_controller.control_t_mod + transformer.dual_controller.control_text_linear = dual_controller.control_text_linear + transformer.dual_controller_freqs = dual_controller.freqs + + comfy_model.diffusion_model = transformer comfy_model.load_device = transformer_load_device patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) diff --git a/nodes_sampler.py b/nodes_sampler.py index a1ec60f..eeb00ad 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -1259,6 +1259,18 @@ class WanVideoSampler: if context_options is None: image_cond = replace_feature(image_cond.unsqueeze(0).clone(), track_pos.unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0] + # LongVie2 dual control + dual_control_embeds = image_embeds.get("dual_control", None) + if dual_control_embeds is not None and context_options is None: + dual_control_input = dict_to_device(dual_control_embeds.copy(), device, dtype) if dual_control_embeds is not None else None + prev_latents = dual_control_input.get("prev_latent", None) + if prev_latents is not None: + _sigma = dual_control_embeds.get("first_frame_noise_level", 0.925926) + log.info(f"Using dual control previous latents with first frame noise level: {_sigma}") + latent[:, :1] = (1 - _sigma) * prev_latents[:, -1:].to(latent) + _sigma * noise[:, :1] + prev_ones = torch.ones(20, *prev_latents.shape[1:], device=device, dtype=dtype) + dual_control_input["prev_latent"] = torch.cat([prev_ones, prev_latents]).unsqueeze(0) + #region model pred def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None, @@ -1513,6 +1525,19 @@ class WanVideoSampler: if wanmove_embeds is not None and context_window is not None: image_cond_input = replace_feature(image_cond_input.unsqueeze(0), track_pos[:, context_window].unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0] + dual_control_in = None + if dual_control_embeds is not None: + if context_window is not None: + dual_control_in = dual_control_embeds.copy() + dense_input_latent = dual_control_embeds.get("dense_input_latent", None) + if dense_input_latent is not None: + dual_control_in["dense_input_latent"] = dual_control_embeds["dense_input_latent"][:, :, context_window] + sparse_input_latent = dual_control_embeds.get("sparse_input_latent", None) + if sparse_input_latent is not None: + dual_control_in["sparse_input_latent"] = dual_control_embeds["sparse_input_latent"][:, :, context_window] + else: + dual_control_in = dual_control_input + base_params = { 'x': [z], # latent 'y': [image_cond_input] if image_cond_input is not None else None, # image cond @@ -1575,6 +1600,7 @@ class WanVideoSampler: "one_to_all_input": one_to_all_data, # One-to-All input "one_to_all_controlnet_strength": one_to_all_data["controlnet_strength"] if one_to_all_data is not None else 0.0, "scail_input": scail_data_in, # SCAIL input + "dual_control_input": dual_control_in, # LongVie2 dual control input } batch_size = 1 diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 89f180f..dd3da2a 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -2318,6 +2318,7 @@ class WanModel(torch.nn.Module): sdancer_input=None, # SteadyDancer one_to_all_input=None, one_to_all_controlnet_strength=0.0, # One-to-All scail_input=None, # SCAIL pose + dual_control_input=None, # LongVie2 dual controlnet ): r""" Forward pass through the diffusion model @@ -2340,6 +2341,7 @@ class WanModel(torch.nn.Module): List[Tensor]: List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8] """ + print("input x shape:", x[0].shape) # Stand-In only used on first positive pass, then cached in kv_cache if is_uncond or current_step > 0: standin_input = None @@ -2544,6 +2546,20 @@ class WanModel(torch.nn.Module): x = [u.flatten(2).transpose(1, 2) for u in x] self.original_seq_len = x[0].shape[1] + prev_latent = None + if dual_control_input is not None: + prev_latent = dual_control_input.get("prev_latent", None) + if prev_latent is not None: + F += prev_latent.shape[2] + print("Using Dual ControlNet prev latent shape:", prev_latent.shape) + prev_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in prev_latent] + prev_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in prev_x] + print("Dual ControlNet prev_x shape:", prev_x[0].shape) + print("x shape before concat:", x[0].shape) + seq_len += prev_x[0].shape[1] + print("F updated to:", F) + x = [torch.cat([u, v], dim=1) for u, v in zip(prev_x, x)] + # SCAIL pose if scail_input is not None: scail_pose_latents = scail_input.get("pose_latent", None) @@ -2834,6 +2850,44 @@ class WanModel(torch.nn.Module): chunked_self_attention = False seq_chunks = 0 + # dual control + if dual_control_input is not None and dual_control_input["start_percent"] <= current_step_percentage <= dual_control_input["end_percent"]: + dense_latent = dual_control_input["dense_input_latent"] + print("dense_latent shape:", dense_latent.shape) + sparse_latent = dual_control_input["sparse_input_latent"] + if dense_latent is None and sparse_latent is None: + raise ValueError("At least one of dense_input_latent or sparse_input_latent must be provided in dual_control_input") + + if dense_latent is not None: + dense_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in dense_latent] + dense_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in dense_x] + dense = self.dual_controller.control_initial_combine_linear_dense(dense_x[0]) + + if sparse_latent is not None: + sparse_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in sparse_latent] + sparse_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in sparse_x] + sparse = self.dual_controller.control_initial_combine_linear_sparse(sparse_x[0]) + + if dense_latent is None: + dense = torch.zeros_like(sparse) + elif sparse_latent is None: + sparse = torch.zeros_like(dense) + + control_context = clip_fea_control = None + if context != []: + control_context = self.dual_controller.control_text_linear(context) + if clip_embed is not None: + clip_fea_control = self.dual_controller.control_text_linear(clip_embed) + control_t_mod = self.dual_controller.control_t_mod(e0) + + control_freqs = torch.cat([ + self.dual_controller_freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1), + self.dual_controller_freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1), + self.dual_controller_freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1) + ], dim=-1).reshape(f * h * w, 1, -1).to(x.device) + else: + dual_control_input = None + # MultiTalk if multitalk_audio is not None: self.multitalk_audio_proj.to(self.main_device) @@ -3173,6 +3227,18 @@ class WanModel(torch.nn.Module): x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, x_onetoall_ref=x_onetoall_ref, onetoall_freqs=onetoall_freqs, **kwargs) # ---post block----# + # dual controlnet + if dual_control_input is not None and (hasattr(block, "control_blocks_dense") or hasattr(block, "control_blocks_sparse")): + if dense_latent is not None and hasattr(block, "control_blocks_dense"): + dense = block.control_blocks_dense(dense, control_context, control_t_mod, control_freqs, clip_fea=clip_fea_control) + if sparse_latent is not None and hasattr(block, "control_blocks_sparse"): + sparse = block.control_blocks_sparse(sparse, control_context, control_t_mod, control_freqs, clip_fea=clip_fea_control) + + if prev_latent is not None: + x[:, -self.original_seq_len:] += block.control_combine_linears(dense + sparse) * dual_control_input["strength"] + else: + x += block.control_combine_linears(dense + sparse) * dual_control_input["strength"] + if self.audio_injector is not None and s2v_audio_input is not None: x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v if block.has_face_fuser_block and motion_vec is not None: @@ -3275,8 +3341,10 @@ class WanModel(torch.nn.Module): # x = x[:, :self.original_seq_len] #grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) - - x = x[:, :self.original_seq_len] + if prev_latent is not None: + x = x[:, -self.original_seq_len:] + else: + x = x[:, :self.original_seq_len] x = self.head(x, e.to(x.device), temp_length=F, e_tr=e_token_replace.to(x.device) if use_token_replace else None, tr_start=token_replace_start, tr_num=replace_token_num)