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:
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user