RoPE frequency offset option for storymem

I'm not 100% sure on this, I initially tested this when I noticed the original code doesn't, but it's described in the paper... now I see the original code has added it too so it seems to be the intended way to use it.
This commit is contained in:
kijai
2026-01-23 14:31:04 +02:00
parent 8640bfad52
commit 2c2a6e1889
3 changed files with 23 additions and 3 deletions
+4 -1
View File
@@ -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,)
+3 -1
View File
@@ -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
+16 -1
View File
@@ -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
)