Add RifleX, ability to generate more frames without looping

https://github.com/thu-ml/RIFLEx
This commit is contained in:
kijai
2025-02-24 23:05:09 +02:00
parent 28ba2dc638
commit 848168b8e4
3 changed files with 21 additions and 3 deletions
@@ -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)
+12
View File
@@ -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]
+4 -1
View File
@@ -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)