diff --git a/nodes.py b/nodes.py index 23ff341..722bc77 100644 --- a/nodes.py +++ b/nodes.py @@ -895,6 +895,8 @@ class WanVideoAddStoryMemLatents: "vae": ("WANVAE",), "embeds": ("WANVIDIMAGE_EMBEDS",), "memory_images": ("IMAGE",), + "rope_negative_offset": ("BOOLEAN", {"default": False, "tooltip": "Use positive RoPE frequency offset for the memory latents"}), + "rope_negative_offset_frames": ("INT", {"default": 5, "min": 0, "max": 100, "step": 1, "tooltip": "RoPE frequency offset for the memory latents"}), } } @@ -903,10 +905,11 @@ class WanVideoAddStoryMemLatents: FUNCTION = "add" CATEGORY = "WanVideoWrapper" - def add(self, vae, embeds, memory_images): + def add(self, vae, embeds, memory_images, rope_negative_offset, rope_negative_offset_frames): updated = dict(embeds) story_mem_latents, = WanVideoEncodeLatentBatch().encode(vae, memory_images) updated["story_mem_latents"] = story_mem_latents["samples"].squeeze(2).permute(1, 0, 2, 3) # [C, T, H, W] + updated["rope_negative_offset_frames"] = rope_negative_offset_frames if rope_negative_offset else 0 return (updated,) diff --git a/nodes_sampler.py b/nodes_sampler.py index 2e8f7c7..0e6b7a8 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -1488,7 +1488,9 @@ class WanVideoSampler: "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 - "transformer_options": transformer_options + "transformer_options": transformer_options, + "rope_negative_offset": image_embeds.get("rope_negative_offset_frames", 0), # StoryMem rope negative offset + "num_memory_frames": story_mem_latents.shape[1] if story_mem_latents is not None else 0, # StoryMem memory frames } batch_size = 1 diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index a10d4b5..ed0a228 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -2201,7 +2201,7 @@ class WanModel(torch.nn.Module): def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, ref_frame_shape=None, pose_frame_shape=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None, - ref_frame_index=10, longcat_num_ref_latents=0): + ref_frame_index=10, longcat_num_ref_latents=0, num_memory_frames=3, rope_negative_offset=5): patch_size = self.patch_size t_len = ((t + (patch_size[0] // 2)) // patch_size[0]) @@ -2225,6 +2225,15 @@ class WanModel(torch.nn.Module): torch.arange(0, steps_t - longcat_num_ref_latents, dtype=dtype, device=device) ], dim=0) img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + grid_t.reshape(-1, 1, 1) + elif num_memory_frames > 0 and rope_negative_offset > 0: + # Negative RoPE shift for memory frames + # Memory frames get negative indices: {-f_m*S, -(f_m-1)*S, ..., -S} + # Current video frames start from 0: {0, 1, ..., f-1} + memory_indices = torch.arange(-num_memory_frames * rope_negative_offset, 0, rope_negative_offset, dtype=dtype, device=device) + current_indices = torch.arange(0, steps_t - num_memory_frames, dtype=dtype, device=device) + grid_t = torch.cat([memory_indices, current_indices], dim=0) + log.info(f"{num_memory_frames} memory frames, temporal rope positions: {grid_t}") + img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + grid_t.reshape(-1, 1, 1) else: # Standard temporal encoding img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start+freq_offset + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1) @@ -2334,6 +2343,8 @@ class WanModel(torch.nn.Module): scail_input=None, # SCAIL pose dual_control_input=None, # LongVie2 dual controlnet transformer_options={}, + rope_negative_offset=0, + num_memory_frames=0, ): r""" Forward pass through the diffusion model @@ -2657,6 +2668,8 @@ class WanModel(torch.nn.Module): self.rope_embedder.k, tuple(ntk_alphas), longcat_num_ref_latents, + rope_negative_offset, + num_memory_frames, ) # Check cache using key comparison @@ -2672,6 +2685,8 @@ class WanModel(torch.nn.Module): ref_frame_shape=ref_frame_shape, pose_frame_shape=pose_frame_shape, longcat_num_ref_latents=longcat_num_ref_latents, + rope_negative_offset=rope_negative_offset, + num_memory_frames=num_memory_frames, device=x.device, dtype=x.dtype )