From 689797bdb7c01ed154bdd342753512dbb2c02a13 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 17 Aug 2025 15:02:45 +0300 Subject: [PATCH 1/3] fix --- nodes.py | 52 ++++++++++++++++++++++++++-------------------------- 1 file changed, 26 insertions(+), 26 deletions(-) diff --git a/nodes.py b/nodes.py index 594b9c3..ac83acf 100644 --- a/nodes.py +++ b/nodes.py @@ -2128,39 +2128,16 @@ class WanVideoSampler: rope_function = "default" #echoshot does not support comfy rope function log.info(f"Number of shots in prompt: {shot_num}, Shot token lengths: {shot_len}") - #region transformer settings - #rope - freqs = None - transformer.rope_embedder.k = None - transformer.rope_embedder.num_frames = None - if "default" in rope_function or bidirectional_sampling: - d = transformer.dim // transformer.num_heads - freqs = torch.cat([ - rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index), - rope_params(1024, 2 * (d // 6)), - rope_params(1024, 2 * (d // 6)) - ], - dim=1) - elif "comfy" in rope_function: - transformer.rope_embedder.k = riflex_freq_index - transformer.rope_embedder.num_frames = latent_video_length - - transformer.rope_func = rope_function - for block in transformer.blocks: - block.rope_func = rope_function - if transformer.vace_layers is not None: - for block in transformer.vace_blocks: - block.rope_func = rope_function - - #blockswap init mm.unload_all_models() mm.soft_empty_cache() gc.collect() - + + #region transformer settings if transformer_options is not None: block_swap_args = transformer_options.get("block_swap_args", None) + #blockswap init if block_swap_args is not None: transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False) for name, param in transformer.named_parameters(): @@ -2286,6 +2263,29 @@ class WanVideoSampler: import copy sample_scheduler_flipped = copy.deepcopy(sample_scheduler) + #rope + freqs = None + transformer.rope_embedder.k = None + transformer.rope_embedder.num_frames = None + if "default" in rope_function or bidirectional_sampling: + d = transformer.dim // transformer.num_heads + freqs = torch.cat([ + rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index), + rope_params(1024, 2 * (d // 6)), + rope_params(1024, 2 * (d // 6)) + ], + dim=1) + elif "comfy" in rope_function: + transformer.rope_embedder.k = riflex_freq_index + transformer.rope_embedder.num_frames = latent_video_length + + transformer.rope_func = rope_function + for block in transformer.blocks: + block.rope_func = rope_function + if transformer.vace_layers is not None: + for block in transformer.vace_blocks: + block.rope_func = rope_function + #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, From 1d7e0848e06894722d23bbf208b3d9c19918d4c6 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 17 Aug 2025 16:46:50 +0300 Subject: [PATCH 2/3] Allow Phantom to work with MultiTalk --- multitalk/multitalk.py | 23 +++++++++++++++++++---- 1 file changed, 19 insertions(+), 4 deletions(-) diff --git a/multitalk/multitalk.py b/multitalk/multitalk.py index 610b829..581b43f 100644 --- a/multitalk/multitalk.py +++ b/multitalk/multitalk.py @@ -253,8 +253,14 @@ class SingleStreamAttention(nn.Module): self.attention_mode = attention_mode def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None, enable_sp=False, kv_seq=None) -> torch.Tensor: - N_t, N_h, N_w = shape + + x_extra = None + if x.shape[0] != encoder_hidden_states.shape[0]: + x_extra = x[:, -N_h * N_w:, :] + x = x[:, :-N_h * N_w, :] + N_t = N_t - 1 + x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) # get q for hidden_state @@ -282,17 +288,18 @@ class SingleStreamAttention(nn.Module): encoder_v.transpose(1, 2), attention_mode=self.attention_mode ) - #x = torch.nn.functional.scaled_dot_product_attention( - # q, encoder_k, encoder_v, attn_mask=None, is_causal=False, dropout_p=0.0) # linear transform x_output_shape = (B, N, C) #x = x.transpose(1, 2) x = x.reshape(x_output_shape) x = self.proj(x) - x = self.proj_drop(x) + x = self.proj_drop(x) x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t) + + if x_extra is not None: + x = torch.cat([x, torch.zeros_like(x_extra)], dim=1) return x @@ -363,6 +370,12 @@ class SingleStreamMultiAttention(SingleStreamAttention): N_t, _, _ = shape x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) + x_extra = None + if x.shape[0] != encoder_hidden_states.shape[0]: + x_extra = x[:, -N_h * N_w:, :] + x = x[:, :-N_h * N_w, :] + N_t = N_t - 1 + # Query projection B, N, C = x.shape q = self.q_linear(x) @@ -473,5 +486,7 @@ class SingleStreamMultiAttention(SingleStreamAttention): # Restore original layout x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t) + if x_extra is not None: + x = torch.cat([x, torch.zeros_like(x_extra)], dim=1) return x \ No newline at end of file From 427e7be6c2c3a2be6563a631b0275443f69b0d94 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 17 Aug 2025 18:24:48 +0300 Subject: [PATCH 3/3] fix multitalk --- multitalk/multitalk.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/multitalk/multitalk.py b/multitalk/multitalk.py index 581b43f..5922ac5 100644 --- a/multitalk/multitalk.py +++ b/multitalk/multitalk.py @@ -256,12 +256,13 @@ class SingleStreamAttention(nn.Module): N_t, N_h, N_w = shape x_extra = None - if x.shape[0] != encoder_hidden_states.shape[0]: + try: + x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) + except: x_extra = x[:, -N_h * N_w:, :] x = x[:, :-N_h * N_w, :] N_t = N_t - 1 - - x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) + x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) # get q for hidden_state B, N, C = x.shape @@ -367,7 +368,7 @@ class SingleStreamMultiAttention(SingleStreamAttention): if human_num is None or human_num <= 1: return super().forward(x, encoder_hidden_states, shape) - N_t, _, _ = shape + N_t, N_h, N_w = shape x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) x_extra = None