Add RifleX, ability to generate more frames without looping
https://github.com/thu-ml/RIFLEx
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user