diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index 389f05d..dc5a693 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -40,7 +40,7 @@ EXAMPLE_DOC_STRING = """""" from ...modules.posemb_layers import get_nd_rotary_pos_embed from ....enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight -def get_rotary_pos_embed(transformer, latent_video_length, height, width): +def get_rotary_pos_embed(transformer, latent_video_length, height, width, k=0): target_ndim = 3 ndim = 5 - 2 rope_theta = 225 @@ -85,6 +85,8 @@ def get_rotary_pos_embed(transformer, latent_video_length, height, width): theta=rope_theta, use_real=True, theta_rescale_factor=1, + num_frames=latent_video_length, + k=k, ) return freqs_cos, freqs_sin def retrieve_timesteps( @@ -424,6 +426,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): feta_args: Optional[Dict] = None, leapfusion_img2vid: Optional[bool] = False, image_cond_latents: Optional[torch.Tensor] = None, + riflex_freq_index: Optional[int] = None, **kwargs, ): r""" @@ -611,7 +614,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): else: # rotary embeddings freqs_cos, freqs_sin = get_rotary_pos_embed( - self.transformer, latent_video_length, height, width + self.transformer, latent_video_length, height, width, k=riflex_freq_index ) if not self.transformer.upcast_rope: freqs_cos = freqs_cos.to(self.base_dtype).to(device) diff --git a/hyvideo/modules/posemb_layers.py b/hyvideo/modules/posemb_layers.py index 2ac4ce7..0fbbada 100644 --- a/hyvideo/modules/posemb_layers.py +++ b/hyvideo/modules/posemb_layers.py @@ -113,6 +113,8 @@ def get_nd_rotary_pos_embed( use_real=False, theta_rescale_factor: Union[float, List[float]] = 1.0, interpolation_factor: Union[float, List[float]] = 1.0, + num_frames: int = 129, + k: int = 0, ): """ This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure. @@ -163,6 +165,8 @@ def get_nd_rotary_pos_embed( use_real=use_real, theta_rescale_factor=theta_rescale_factor[i], interpolation_factor=interpolation_factor[i], + L_test=num_frames, + k=k, ) # 2 x [WHD, rope_dim_list[i]] embs.append(emb) @@ -182,6 +186,8 @@ def get_1d_rotary_pos_embed( use_real: bool = False, theta_rescale_factor: float = 1.0, interpolation_factor: float = 1.0, + L_test: int = 100, + k: int = 0, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: """ Precompute the frequency tensor for complex exponential (cis) with given dimensions. @@ -215,6 +221,12 @@ def get_1d_rotary_pos_embed( theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim) ) # [D/2] # assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}" + + #RIFLEx https://github.com/thu-ml/RIFLEx + if k > 0: + freqs[k-1] = 0.9 * 2 * torch.pi / L_test + + freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2] if use_real: freqs_cos = freqs.cos().repeat_interleave(2, dim=1) # [S, D] diff --git a/nodes.py b/nodes.py index 24eb0fa..37f94d9 100644 --- a/nodes.py +++ b/nodes.py @@ -1155,6 +1155,7 @@ class HyVideoSampler: { "default": 'FlowMatchDiscreteScheduler' }), + "riflex_freq_index": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Frequency index for RIFLEX, disabled when 0, default 4. Allows for new frames to be generated after 129 without looping"}), } } @@ -1164,7 +1165,8 @@ class HyVideoSampler: CATEGORY = "HunyuanVideoWrapper" def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, width, height, num_frames, - samples=None, denoise_strength=1.0, force_offload=True, stg_args=None, context_options=None, feta_args=None, teacache_args=None, scheduler=None, image_cond_latents=None): + samples=None, denoise_strength=1.0, force_offload=True, stg_args=None, context_options=None, feta_args=None, + teacache_args=None, scheduler=None, image_cond_latents=None, riflex_freq_index=0): model = model.model device = mm.get_torch_device() @@ -1306,6 +1308,7 @@ class HyVideoSampler: feta_args=feta_args, leapfusion_img2vid = leapfusion_img2vid, image_cond_latents = image_cond_latents["samples"] * VAE_SCALING_FACTOR if image_cond_latents is not None else None, + riflex_freq_index = riflex_freq_index ) print_memory(device)