diff --git a/nodes.py b/nodes.py index 378eb80..b6bf6d1 100644 --- a/nodes.py +++ b/nodes.py @@ -1592,6 +1592,24 @@ class WanVideoLoopArgs: def process(self, **kwargs): return (kwargs,) + +class WanVideoExperimentalArgs: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "video_attention_split_steps": ("STRING", {"default": "2, 3", "tooltip": "Steps to split self attention when using multiple prompts"}), + }, + } + + RETURN_TYPES = ("EXPERIMENTALARGS", ) + RETURN_NAMES = ("exp_args",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Experimental stuff" + EXPERIMENTAL = True + + def process(self, **kwargs): + return (kwargs,) class WanVideoSampler: @classmethod @@ -1625,6 +1643,7 @@ class WanVideoSampler: "slg_args": ("SLGARGS", ), "rope_function": (["default", "comfy"], {"default": "default", "tooltip": "!EXPERIMENTAL! Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile"}), "loop_args": ("LOOPARGS", ), + "experimental_args": ("EXPERIMENTALARGS", ), } } @@ -1635,7 +1654,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): + teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default", loop_args=None, experimental_args=None): #assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache." patcher = model model = model.model @@ -1956,6 +1975,11 @@ class WanVideoSampler: drift_timesteps = torch.cat([drift_timesteps, torch.tensor([0]).to(drift_timesteps.device)]).to(drift_timesteps.device) timesteps[-drift_steps:] = drift_timesteps[-drift_steps:] + if experimental_args is not None: + video_attention_split_steps = experimental_args.get("video_attention_split_steps", []) + if video_attention_split_steps: + transformer.video_attention_split_steps = [int(x.strip()) for x in video_attention_split_steps.split(",")] + def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, teacache_state=None): with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True): nonlocal patcher @@ -2628,6 +2652,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoLoopArgs": WanVideoLoopArgs, "WanVideoImageResizeToClosest": WanVideoImageResizeToClosest, "WanVideoSetBlockSwap": WanVideoSetBlockSwap, + "WanVideoExperimentalArgs": WanVideoExperimentalArgs, } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -2660,4 +2685,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoLoopArgs": "WanVideo Loop Args", "WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest", "WanVideoSetBlockSwap": "WanVideo Set BlockSwap", + "WanVideoExperimentalArgs": "WanVideo Experimental Args", } diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 32329c0..bf0c642 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -221,7 +221,7 @@ class WanSelfAttention(nn.Module): self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity() self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity() - def forward(self, x, seq_lens, grid_sizes, freqs, rope_func = "default"): + def forward(self, x, seq_lens, grid_sizes, freqs, seq_chunks=1,current_step=0, video_attention_split_steps = [], rope_func = "default"): r""" Args: x(Tensor): Shape [B, L, num_heads, C / num_heads] @@ -240,35 +240,84 @@ class WanSelfAttention(nn.Module): q, k, v = qkv_fn(x) - if self.attention_mode == 'spargeattn_tune' or self.attention_mode == 'spargeattn': - tune_mode = False - if self.attention_mode == 'spargeattn_tune': - tune_mode = True + # if self.attention_mode == 'spargeattn_tune' or self.attention_mode == 'spargeattn': + # tune_mode = False + # if self.attention_mode == 'spargeattn_tune': + # tune_mode = True - if hasattr(self, 'inner_attention'): - #print("has inner attention") - q=rope_apply(q, grid_sizes, freqs) - k=rope_apply(k, grid_sizes, freqs) - q = q.permute(0, 2, 1, 3) - k = k.permute(0, 2, 1, 3) - v = v.permute(0, 2, 1, 3) - x = self.inner_attention( - q=q, - k=k, - v=v, - is_causal=False, - tune_mode=tune_mode - ).permute(0, 2, 1, 3) - #print("inner attention", x.shape) #inner attention torch.Size([1, 12, 32760, 128]) + # if hasattr(self, 'inner_attention'): + # #print("has inner attention") + # q=rope_apply(q, grid_sizes, freqs) + # k=rope_apply(k, grid_sizes, freqs) + # q = q.permute(0, 2, 1, 3) + # k = k.permute(0, 2, 1, 3) + # v = v.permute(0, 2, 1, 3) + # x = self.inner_attention( + # q=q, + # k=k, + # v=v, + # is_causal=False, + # tune_mode=tune_mode + # ).permute(0, 2, 1, 3) + # #print("inner attention", x.shape) #inner attention torch.Size([1, 12, 32760, 128]) + #else: + if rope_func == "comfy": + q, k = apply_rope_comfy(q, k, freqs) else: - if rope_func == "comfy": - q, k = apply_rope_comfy(q, k, freqs) - else: - q=rope_apply(q, grid_sizes, freqs) - k=rope_apply(k, grid_sizes, freqs) - if is_enhance_enabled(): - feta_scores = get_feta_scores(q, k) + q=rope_apply(q, grid_sizes, freqs) + k=rope_apply(k, grid_sizes, freqs) + if is_enhance_enabled(): + feta_scores = get_feta_scores(q, k) + # Split by frames + if seq_chunks > 1 and current_step in video_attention_split_steps: + outputs = [] + # Extract frame, height, width from grid_sizes - force to CPU scalars + frames = grid_sizes[0][0].item() + height = grid_sizes[0][1].item() + width = grid_sizes[0][2].item() + tokens_per_frame = height * width + + actual_chunks = min(seq_chunks, frames) + if isinstance(actual_chunks, torch.Tensor): + actual_chunks = actual_chunks.item() + + frame_chunks = [] # Pre-calculate all chunk boundaries + start_frame = 0 + base_frames_per_chunk = frames // actual_chunks + extra_frames = frames % actual_chunks + + # Pre-calculate all chunks + for i in range(actual_chunks): + chunk_size = base_frames_per_chunk + (1 if i < extra_frames else 0) + end_frame = start_frame + chunk_size + frame_chunks.append((start_frame, end_frame)) + start_frame = end_frame + + # Process each chunk using the pre-calculated boundaries + for start_frame, end_frame in frame_chunks: + # Convert to token indices + start_idx = int(start_frame * tokens_per_frame) + end_idx = int(end_frame * tokens_per_frame) + + chunk_q = q[:, start_idx:end_idx, :, :] + chunk_k = k[:, start_idx:end_idx, :, :] + chunk_v = v[:, start_idx:end_idx, :, :] + + chunk_out = attention( + q=chunk_q, + k=chunk_k, + v=chunk_v, + k_lens=seq_lens, + window_size=self.window_size, + attention_mode=self.attention_mode) + + outputs.append(chunk_out) + + # Concatenate outputs along the sequence dimension + x = torch.cat(outputs, dim=1) + else: + # Original attention computation x = attention( q=q, k=k, @@ -416,8 +465,11 @@ class WanAttentionBlock(nn.Module): freqs, context, context_lens, + current_step, + video_attention_split_steps=[], rope_func = "default", clip_embed=None, + ): r""" Args: @@ -433,29 +485,42 @@ class WanAttentionBlock(nn.Module): # self-attention y = self.self_attn( - self.norm1(x).float() * (1 + e[1]) + e[0], seq_lens, grid_sizes, - freqs, rope_func=rope_func) + self.norm1(x).float() * (1 + e[1]) + e[0], + seq_lens, grid_sizes, + freqs, rope_func=rope_func, + seq_chunks=max(context.shape[0], clip_embed.shape[0] if clip_embed is not None else 0), + current_step=current_step, + video_attention_split_steps=video_attention_split_steps + ) + x = x.to(torch.float32) + (y.to(torch.float32) * e[2].to(torch.float32)) # cross-attention & ffn function - def cross_attn_ffn(x, context, context_lens, e, clip_embed=None): + def cross_attn_ffn(x, context, context_lens, e, clip_embed=None, grid_sizes=None): if context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1): # Get number of prompts num_prompts = context.shape[0] num_clip_embeds = 0 if clip_embed is None else clip_embed.shape[0] num_segments = max(num_prompts, num_clip_embeds) - # split the sequence dimension - seq_len = x.shape[1] - segment_length = seq_len // num_prompts + # Extract spatial dimensions + frames, height, width = grid_sizes[0] # Assuming batch size 1 + tokens_per_frame = height * width + + # Distribute frames across prompts + frames_per_segment = max(1, frames // num_segments) # Process each prompt segment x_combined = torch.zeros_like(x) for i in range(num_segments): - # Calculate indices for this segment - start_idx = i * segment_length - end_idx = (i+1) * segment_length if i < num_segments-1 else seq_len + # Calculate frame boundaries for this segment + start_frame = i * frames_per_segment + end_frame = min((i+1) * frames_per_segment, frames) if i < num_segments-1 else frames + + # Convert frame indices to token indices + start_idx = start_frame * tokens_per_frame + end_idx = end_frame * tokens_per_frame segment_indices = torch.arange(start_idx, end_idx, device=x.device, dtype=torch.long) # Get prompt segment (cycle through available prompts if needed) @@ -494,7 +559,7 @@ class WanAttentionBlock(nn.Module): x = x.to(torch.float32) + (y.to(torch.float32) * e[5].to(torch.float32)) return x - x = cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed) + x = cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes) return x @@ -654,6 +719,8 @@ class WanModel(ModelMixin, ConfigMixin): self.use_non_blocking = True + self.video_attention_split_steps = [] + # embeddings self.patch_embedding = nn.Conv3d( in_dim, dim, kernel_size=patch_size, stride=patch_size) @@ -900,7 +967,9 @@ class WanModel(ModelMixin, ConfigMixin): context=context, context_lens=context_lens, clip_embed=clip_embed, - rope_func=rope_func + rope_func=rope_func, + current_step=current_step, + video_attention_split_steps=self.video_attention_split_steps, ) for b, block in enumerate(self.blocks):