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",),
|
"vae": ("WANVAE",),
|
||||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||||
"memory_images": ("IMAGE",),
|
"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"
|
FUNCTION = "add"
|
||||||
CATEGORY = "WanVideoWrapper"
|
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)
|
updated = dict(embeds)
|
||||||
story_mem_latents, = WanVideoEncodeLatentBatch().encode(vae, memory_images)
|
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["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,)
|
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,
|
"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
|
"scail_input": scail_data_in, # SCAIL input
|
||||||
"dual_control_input": dual_control_in, # LongVie2 dual control 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
|
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,
|
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,
|
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
|
patch_size = self.patch_size
|
||||||
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
|
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)
|
torch.arange(0, steps_t - longcat_num_ref_latents, dtype=dtype, device=device)
|
||||||
], dim=0)
|
], dim=0)
|
||||||
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + grid_t.reshape(-1, 1, 1)
|
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:
|
else:
|
||||||
# Standard temporal encoding
|
# 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)
|
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
|
scail_input=None, # SCAIL pose
|
||||||
dual_control_input=None, # LongVie2 dual controlnet
|
dual_control_input=None, # LongVie2 dual controlnet
|
||||||
transformer_options={},
|
transformer_options={},
|
||||||
|
rope_negative_offset=0,
|
||||||
|
num_memory_frames=0,
|
||||||
):
|
):
|
||||||
r"""
|
r"""
|
||||||
Forward pass through the diffusion model
|
Forward pass through the diffusion model
|
||||||
@@ -2657,6 +2668,8 @@ class WanModel(torch.nn.Module):
|
|||||||
self.rope_embedder.k,
|
self.rope_embedder.k,
|
||||||
tuple(ntk_alphas),
|
tuple(ntk_alphas),
|
||||||
longcat_num_ref_latents,
|
longcat_num_ref_latents,
|
||||||
|
rope_negative_offset,
|
||||||
|
num_memory_frames,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Check cache using key comparison
|
# Check cache using key comparison
|
||||||
@@ -2672,6 +2685,8 @@ class WanModel(torch.nn.Module):
|
|||||||
ref_frame_shape=ref_frame_shape,
|
ref_frame_shape=ref_frame_shape,
|
||||||
pose_frame_shape=pose_frame_shape,
|
pose_frame_shape=pose_frame_shape,
|
||||||
longcat_num_ref_latents=longcat_num_ref_latents,
|
longcat_num_ref_latents=longcat_num_ref_latents,
|
||||||
|
rope_negative_offset=rope_negative_offset,
|
||||||
|
num_memory_frames=num_memory_frames,
|
||||||
device=x.device,
|
device=x.device,
|
||||||
dtype=x.dtype
|
dtype=x.dtype
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user