From b6c1705f9f9a4dc118c32e07b5e1c96dd9eebe3a Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 19 Dec 2024 11:56:53 +0200 Subject: [PATCH] simple context windows with freenoise shuffling --- context.py | 184 ++++++++++++++ .../pipelines/pipeline_hunyuan_video.py | 230 ++++++++++++++---- nodes.py | 95 +++----- nodes_rf_inversion.py | 9 +- 4 files changed, 408 insertions(+), 110 deletions(-) create mode 100644 context.py diff --git a/context.py b/context.py new file mode 100644 index 0000000..6a30fed --- /dev/null +++ b/context.py @@ -0,0 +1,184 @@ +import numpy as np +from typing import Callable, Optional, List + + +def ordered_halving(val): + bin_str = f"{val:064b}" + bin_flip = bin_str[::-1] + as_int = int(bin_flip, 2) + + return as_int / (1 << 64) + +def does_window_roll_over(window: list[int], num_frames: int) -> tuple[bool, int]: + prev_val = -1 + for i, val in enumerate(window): + val = val % num_frames + if val < prev_val: + return True, i + prev_val = val + return False, -1 + +def shift_window_to_start(window: list[int], num_frames: int): + start_val = window[0] + for i in range(len(window)): + # 1) subtract each element by start_val to move vals relative to the start of all frames + # 2) add num_frames and take modulus to get adjusted vals + window[i] = ((window[i] - start_val) + num_frames) % num_frames + +def shift_window_to_end(window: list[int], num_frames: int): + # 1) shift window to start + shift_window_to_start(window, num_frames) + end_val = window[-1] + end_delta = num_frames - end_val - 1 + for i in range(len(window)): + # 2) add end_delta to each val to slide windows to end + window[i] = window[i] + end_delta + +def get_missing_indexes(windows: list[list[int]], num_frames: int) -> list[int]: + all_indexes = list(range(num_frames)) + for w in windows: + for val in w: + try: + all_indexes.remove(val) + except ValueError: + pass + return all_indexes + +def uniform_looped( + step: int = ..., + num_steps: Optional[int] = None, + num_frames: int = ..., + context_size: Optional[int] = None, + context_stride: int = 3, + context_overlap: int = 4, + closed_loop: bool = True, +): + if num_frames <= context_size: + yield list(range(num_frames)) + return + + context_stride = min(context_stride, int(np.ceil(np.log2(num_frames / context_size))) + 1) + + for context_step in 1 << np.arange(context_stride): + pad = int(round(num_frames * ordered_halving(step))) + for j in range( + int(ordered_halving(step) * context_step) + pad, + num_frames + pad + (0 if closed_loop else -context_overlap), + (context_size * context_step - context_overlap), + ): + yield [e % num_frames for e in range(j, j + context_size * context_step, context_step)] + +#from AnimateDiff-Evolved by Kosinkadink (https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved) +def uniform_standard( + step: int = ..., + num_steps: Optional[int] = None, + num_frames: int = ..., + context_size: Optional[int] = None, + context_stride: int = 3, + context_overlap: int = 4, + closed_loop: bool = True, +): + windows = [] + if num_frames <= context_size: + windows.append(list(range(num_frames))) + return windows + + context_stride = min(context_stride, int(np.ceil(np.log2(num_frames / context_size))) + 1) + + for context_step in 1 << np.arange(context_stride): + pad = int(round(num_frames * ordered_halving(step))) + for j in range( + int(ordered_halving(step) * context_step) + pad, + num_frames + pad + (0 if closed_loop else -context_overlap), + (context_size * context_step - context_overlap), + ): + windows.append([e % num_frames for e in range(j, j + context_size * context_step, context_step)]) + + # now that windows are created, shift any windows that loop, and delete duplicate windows + delete_idxs = [] + win_i = 0 + while win_i < len(windows): + # if window is rolls over itself, need to shift it + is_roll, roll_idx = does_window_roll_over(windows[win_i], num_frames) + if is_roll: + roll_val = windows[win_i][roll_idx] # roll_val might not be 0 for windows of higher strides + shift_window_to_end(windows[win_i], num_frames=num_frames) + # check if next window (cyclical) is missing roll_val + if roll_val not in windows[(win_i+1) % len(windows)]: + # need to insert new window here - just insert window starting at roll_val + windows.insert(win_i+1, list(range(roll_val, roll_val + context_size))) + # delete window if it's not unique + for pre_i in range(0, win_i): + if windows[win_i] == windows[pre_i]: + delete_idxs.append(win_i) + break + win_i += 1 + + # reverse delete_idxs so that they will be deleted in an order that doesn't break idx correlation + delete_idxs.reverse() + for i in delete_idxs: + windows.pop(i) + return windows + +def static_standard( + step: int = ..., + num_steps: Optional[int] = None, + num_frames: int = ..., + context_size: Optional[int] = None, + context_stride: int = 3, + context_overlap: int = 4, + closed_loop: bool = True, +): + windows = [] + if num_frames <= context_size: + windows.append(list(range(num_frames))) + return windows + # always return the same set of windows + delta = context_size - context_overlap + for start_idx in range(0, num_frames, delta): + # if past the end of frames, move start_idx back to allow same context_length + ending = start_idx + context_size + if ending >= num_frames: + final_delta = ending - num_frames + final_start_idx = start_idx - final_delta + windows.append(list(range(final_start_idx, final_start_idx + context_size))) + break + windows.append(list(range(start_idx, start_idx + context_size))) + return windows + +def get_context_scheduler(name: str) -> Callable: + if name == "uniform_looped": + return uniform_looped + elif name == "uniform_standard": + return uniform_standard + elif name == "static_standard": + return static_standard + else: + raise ValueError(f"Unknown context_overlap policy {name}") + + +def get_total_steps( + scheduler, + timesteps: List[int], + num_steps: Optional[int] = None, + num_frames: int = ..., + context_size: Optional[int] = None, + context_stride: int = 3, + context_overlap: int = 4, + closed_loop: bool = True, +): + return sum( + len( + list( + scheduler( + i, + num_steps, + num_frames, + context_size, + context_stride, + context_overlap, + ) + ) + ) + for i in range(len(timesteps)) + ) diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index 86e84c4..8d8572a 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -36,7 +36,55 @@ from comfy.utils import ProgressBar logger = logging.get_logger(__name__) # pylint: disable=invalid-name EXAMPLE_DOC_STRING = """""" +from ...modules.posemb_layers import get_nd_rotary_pos_embed +def get_rotary_pos_embed(transformer, latent_video_length, height, width): + target_ndim = 3 + ndim = 5 - 2 + rope_theta = 225 + patch_size = transformer.patch_size + rope_dim_list = transformer.rope_dim_list + hidden_size = transformer.hidden_size + heads_num = transformer.heads_num + head_dim = hidden_size // heads_num + + # 884 + latents_size = [latent_video_length, height // 8, width // 8] + + if isinstance(patch_size, int): + assert all(s % patch_size == 0 for s in latents_size), ( + f"Latent size(last {ndim} dimensions) should be divisible by patch size({patch_size}), " + f"but got {latents_size}." + ) + rope_sizes = [s // patch_size for s in latents_size] + elif isinstance(patch_size, list): + assert all( + s % patch_size[idx] == 0 + for idx, s in enumerate(latents_size) + ), ( + f"Latent size(last {ndim} dimensions) should be divisible by patch size({patch_size}), " + f"but got {latents_size}." + ) + rope_sizes = [ + s // patch_size[idx] for idx, s in enumerate(latents_size) + ] + + if len(rope_sizes) != target_ndim: + rope_sizes = [1] * (target_ndim - len(rope_sizes)) + rope_sizes # time axis + + if rope_dim_list is None: + rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)] + assert ( + sum(rope_dim_list) == head_dim + ), "sum(rope_dim_list) should equal to head_dim of attention layer" + freqs_cos, freqs_sin = get_nd_rotary_pos_embed( + rope_dim_list, + rope_sizes, + theta=rope_theta, + use_real=True, + theta_rescale_factor=1, + ) + return freqs_cos, freqs_sin def retrieve_timesteps( scheduler, num_inference_steps: Optional[int] = None, @@ -183,6 +231,9 @@ class HunyuanVideoPipeline(DiffusionPipeline): generator, latents=None, denoise_strength=1.0, + freenoise=False, + context_size=None, + context_overlap=None ): shape = ( batch_size, @@ -197,6 +248,40 @@ class HunyuanVideoPipeline(DiffusionPipeline): f" size of {batch_size}. Make sure the batch size matches the length of the generators." ) noise = randn_tensor(shape, generator=generator, device=device, dtype=self.base_dtype) + if freenoise: + logger.info("Applying FreeNoise") + # code and comments from AnimateDiff-Evolved by Kosinkadink (https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved) + #video_length = video_length // 4 + delta = context_size - context_overlap + for start_idx in range(0, video_length-context_size, delta): + # start_idx corresponds to the beginning of a context window + # goal: place shuffled in the delta region right after the end of the context window + # if space after context window is not enough to place the noise, adjust and finish + place_idx = start_idx + context_size + # if place_idx is outside the valid indexes, we are already finished + if place_idx >= video_length: + break + end_idx = place_idx - 1 + #print("video_length:", video_length, "start_idx:", start_idx, "end_idx:", end_idx, "place_idx:", place_idx, "delta:", delta) + + # if there is not enough room to copy delta amount of indexes, copy limited amount and finish + if end_idx + delta >= video_length: + final_delta = video_length - place_idx + # generate list of indexes in final delta region + list_idx = torch.tensor(list(range(start_idx,start_idx+final_delta)), device=torch.device("cpu"), dtype=torch.long) + # shuffle list + list_idx = list_idx[torch.randperm(final_delta, generator=generator)] + # apply shuffled indexes + noise[:, :, place_idx:place_idx + final_delta, :, :] = noise[:, :, list_idx, :, :] + break + # otherwise, do normal behavior + # generate list of indexes in delta region + list_idx = torch.tensor(list(range(start_idx,start_idx+delta)), device=torch.device("cpu"), dtype=torch.long) + # shuffle list + list_idx = list_idx[torch.randperm(delta, generator=generator)] + # apply shuffled indexes + #print("place_idx:", place_idx, "delta:", delta, "list_idx:", list_idx) + noise[:, :, place_idx:place_idx + delta, :, :] = noise[:, :, list_idx, :, :] if latents is None: latents = noise else: @@ -313,7 +398,6 @@ class HunyuanVideoPipeline(DiffusionPipeline): denoise_strength: float = 1.0, generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, latents: Optional[torch.Tensor] = None, - cross_attention_kwargs: Optional[Dict[str, Any]] = None, guidance_rescale: float = 0.0, clip_skip: Optional[int] = None, @@ -325,14 +409,13 @@ class HunyuanVideoPipeline(DiffusionPipeline): ] ] = None, callback_on_step_end_tensor_inputs: List[str] = ["latents"], - freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None, - n_tokens: Optional[int] = None, embedded_guidance_scale: Optional[float] = None, stg_mode: Optional[str] = None, stg_block_idx: Optional[int] = -1, stg_scale: Optional[float] = 0.0, stg_start_percent: Optional[float] = 0.0, stg_end_percent: Optional[float] = 1.0, + context_options: Optional[Dict[str, Any]] = None, **kwargs, ): r""" @@ -460,7 +543,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): # 4. Prepare timesteps extra_set_timesteps_kwargs = self.prepare_extra_func_kwargs( - self.scheduler.set_timesteps, {"n_tokens": n_tokens} + self.scheduler.set_timesteps, {} ) if hasattr(self.scheduler, "set_begin_index") and denoise_strength == 1.0: self.scheduler.set_begin_index(begin_index=0) @@ -477,6 +560,35 @@ class HunyuanVideoPipeline(DiffusionPipeline): latent_video_length = (video_length - 1) // 4 + 1 # elif "888" in vae_ver: # video_length = (video_length - 1) // 8 + 1 + + # context windows + use_context_schedule = False + freenoise = False + context_stride = 1 + context_overlap = 1 + context_frames = 65 + if context_options is not None: + context_schedule = context_options["context_schedule"] + context_frames = (context_options["context_frames"] - 1) // 4 + 1 + context_stride = context_options["context_stride"] // 4 + context_overlap = context_options["context_overlap"] // 4 + freenoise = context_options["freenoise"] + + logger.info(f"Context schedule enabled: {context_frames} frames, {context_stride} stride, {context_overlap} overlap") + use_context_schedule = True + from ....context import get_context_scheduler + context = get_context_scheduler(context_schedule) + freqs_cos, freqs_sin = get_rotary_pos_embed( + self.transformer, context_frames, height, width + ) + else: + # rotary embeddings + freqs_cos, freqs_sin = get_rotary_pos_embed( + self.transformer, latent_video_length, height, width + ) + + freqs_cos = freqs_cos.to(self.base_dtype).to(device) + freqs_sin = freqs_sin.to(self.base_dtype).to(device) # 5. Prepare latent variables @@ -493,6 +605,9 @@ class HunyuanVideoPipeline(DiffusionPipeline): generator, latents, denoise_strength=denoise_strength, + freenoise=freenoise, + context_size=context_frames, + context_overlap=context_overlap ) # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline @@ -553,9 +668,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): input_prompt_mask = prompt_mask[1].unsqueeze(0) input_prompt_embeds_2 = prompt_embeds_2[1].unsqueeze(0) - latent_model_input = self.scheduler.scale_model_input( - latent_model_input, t - ) + latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) t_expand = t.repeat(latent_model_input.shape[0]) if embedded_guidance_scale is not None and not cfg_enabled: @@ -568,44 +681,73 @@ class HunyuanVideoPipeline(DiffusionPipeline): ) else: guidance_expand = None - - # predict the noise residual - with torch.autocast( - device_type="cuda", dtype=self.base_dtype, enabled=True - ): - noise_pred = self.transformer( # For an input image (129, 192, 336) (1, 256, 256) - latent_model_input, # [2, 16, 33, 24, 42] - t_expand, # [2] - text_states=input_prompt_embeds, # [2, 256, 4096] - text_mask=input_prompt_mask, # [2, 256] - text_states_2=input_prompt_embeds_2, # [2, 768] - freqs_cos=freqs_cis[0], # [seqlen, head_dim] - freqs_sin=freqs_cis[1], # [seqlen, head_dim] - guidance=guidance_expand, - stg_block_idx=stg_block_idx, - stg_mode=stg_mode, - return_dict=True, - )["x"] - # perform guidance - if cfg_enabled and not self.do_spatio_temporal_guidance: - noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) - noise_pred = noise_pred_uncond + self.guidance_scale * ( - noise_pred_text - noise_pred_uncond - ) - elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance: - raise NotImplementedError - noise_pred_uncond, noise_pred_text, noise_pred_perturb = noise_pred.chunk(3) - noise_pred = noise_pred_uncond + self.guidance_scale * ( - noise_pred_text - noise_pred_uncond - ) + self._stg_scale * ( - noise_pred_text - noise_pred_perturb - ) - elif self.do_spatio_temporal_guidance and stg_enabled: - noise_pred_text, noise_pred_perturb = noise_pred.chunk(2) - noise_pred = noise_pred_text + self._stg_scale * ( - noise_pred_text - noise_pred_perturb - ) + if use_context_schedule: + counter = torch.zeros_like(latent_model_input) + noise_pred = torch.zeros_like(latent_model_input) + context_queue = list(context( + i, num_inference_steps, latents.shape[2], context_frames, context_stride, context_overlap, + )) + for c in context_queue: + partial_latent_model_input = latent_model_input[:, :, c, :, :] + print("partial_latent_model_input", partial_latent_model_input.shape) + with torch.autocast( + device_type="cuda", dtype=self.base_dtype, enabled=True): + noise_pred[:, :, c, :, :] += self.transformer( # For an input image (129, 192, 336) (1, 256, 256) + partial_latent_model_input, # [2, 16, 33, 24, 42] + t_expand, # [2] + text_states=input_prompt_embeds, # [2, 256, 4096] + text_mask=input_prompt_mask, # [2, 256] + text_states_2=input_prompt_embeds_2, # [2, 768] + freqs_cos=freqs_cos, # [seqlen, head_dim] + freqs_sin=freqs_sin, # [seqlen, head_dim] + guidance=guidance_expand, + stg_block_idx=stg_block_idx, + stg_mode=stg_mode, + return_dict=True, + )["x"] + + counter[:, :, c, :, :] += 1 + noise_pred = noise_pred.float() + noise_pred /= counter + else: + # predict the noise residual + with torch.autocast( + device_type="cuda", dtype=self.base_dtype, enabled=True + ): + noise_pred = self.transformer( # For an input image (129, 192, 336) (1, 256, 256) + latent_model_input, # [2, 16, 33, 24, 42] + t_expand, # [2] + text_states=input_prompt_embeds, # [2, 256, 4096] + text_mask=input_prompt_mask, # [2, 256] + text_states_2=input_prompt_embeds_2, # [2, 768] + freqs_cos=freqs_cos, # [seqlen, head_dim] + freqs_sin=freqs_sin, # [seqlen, head_dim] + guidance=guidance_expand, + stg_block_idx=stg_block_idx, + stg_mode=stg_mode, + return_dict=True, + )["x"] + + # perform guidance + if cfg_enabled and not self.do_spatio_temporal_guidance: + noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) + noise_pred = noise_pred_uncond + self.guidance_scale * ( + noise_pred_text - noise_pred_uncond + ) + elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance: + raise NotImplementedError + noise_pred_uncond, noise_pred_text, noise_pred_perturb = noise_pred.chunk(3) + noise_pred = noise_pred_uncond + self.guidance_scale * ( + noise_pred_text - noise_pred_uncond + ) + self._stg_scale * ( + noise_pred_text - noise_pred_perturb + ) + elif self.do_spatio_temporal_guidance and stg_enabled: + noise_pred_text, noise_pred_perturb = noise_pred.chunk(2) + noise_pred = noise_pred_text + self._stg_scale * ( + noise_pred_text - noise_pred_perturb + ) # compute the previous noisy sample x_t -> x_t-1 latents = self.scheduler.step( diff --git a/nodes.py b/nodes.py index 9a55d54..97e118a 100644 --- a/nodes.py +++ b/nodes.py @@ -9,7 +9,6 @@ from typing import List, Dict, Any, Tuple from .hyvideo.constants import PROMPT_TEMPLATE from .hyvideo.text_encoder import TextEncoder from .hyvideo.utils.data_utils import align_to -from .hyvideo.modules.posemb_layers import get_nd_rotary_pos_embed from .hyvideo.diffusion.schedulers import FlowMatchDiscreteScheduler from .hyvideo.diffusion.pipelines import HunyuanVideoPipeline from .hyvideo.vae.autoencoder_kl_causal_3d import AutoencoderKLCausal3D @@ -27,54 +26,6 @@ import comfy.latent_formats script_directory = os.path.dirname(os.path.abspath(__file__)) -def get_rotary_pos_embed(transformer, video_length, height, width): - target_ndim = 3 - ndim = 5 - 2 - rope_theta = 225 - patch_size = transformer.patch_size - rope_dim_list = transformer.rope_dim_list - hidden_size = transformer.hidden_size - heads_num = transformer.heads_num - head_dim = hidden_size // heads_num - - # 884 - latents_size = [(video_length - 1) // 4 + 1, height // 8, width // 8] - - if isinstance(patch_size, int): - assert all(s % patch_size == 0 for s in latents_size), ( - f"Latent size(last {ndim} dimensions) should be divisible by patch size({patch_size}), " - f"but got {latents_size}." - ) - rope_sizes = [s // patch_size for s in latents_size] - elif isinstance(patch_size, list): - assert all( - s % patch_size[idx] == 0 - for idx, s in enumerate(latents_size) - ), ( - f"Latent size(last {ndim} dimensions) should be divisible by patch size({patch_size}), " - f"but got {latents_size}." - ) - rope_sizes = [ - s // patch_size[idx] for idx, s in enumerate(latents_size) - ] - - if len(rope_sizes) != target_ndim: - rope_sizes = [1] * (target_ndim - len(rope_sizes)) + rope_sizes # time axis - - if rope_dim_list is None: - rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)] - assert ( - sum(rope_dim_list) == head_dim - ), "sum(rope_dim_list) should equal to head_dim of attention layer" - freqs_cos, freqs_sin = get_nd_rotary_pos_embed( - rope_dim_list, - rope_sizes, - theta=rope_theta, - use_real=True, - theta_rescale_factor=1, - ) - return freqs_cos, freqs_sin - def filter_state_dict_by_blocks(state_dict, blocks_mapping): filtered_dict = {} @@ -1027,7 +978,35 @@ class HyVideoTextEmbedsLoad: } return (prompt_embeds_dict,) + +class HyVideoContextOptions: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "context_schedule": (["uniform_standard", "uniform_looped", "static_standard"],), + "context_frames": ("INT", {"default": 65, "min": 2, "max": 1000, "step": 1, "tooltip": "Number of pixel frames in the context, NOTE: the latent space has 4 frames in 1"} ), + "context_stride": ("INT", {"default": 4, "min": 4, "max": 100, "step": 1, "tooltip": "Context stride as pixel frames, NOTE: the latent space has 4 frames in 1"} ), + "context_overlap": ("INT", {"default": 4, "min": 4, "max": 100, "step": 1, "tooltip": "Context overlap as pixel frames, NOTE: the latent space has 4 frames in 1"} ), + "freenoise": ("BOOLEAN", {"default": True, "tooltip": "Shuffle the noise"}), + } + } + RETURN_TYPES = ("COGCONTEXT", ) + RETURN_NAMES = ("context_options",) + FUNCTION = "process" + CATEGORY = "CogVideoWrapper" + DESCRIPTION = "Context options for HunyuanVideo, allows splitting the video into context windows and attemps blending them for longer generations than the model and memory otherwise would allow." + + def process(self, context_schedule, context_frames, context_stride, context_overlap, freenoise): + context_options = { + "context_schedule":context_schedule, + "context_frames":context_frames, + "context_stride":context_stride, + "context_overlap":context_overlap, + "freenoise":freenoise + } + + return (context_options,) #region Sampler class HyVideoSampler: @classmethod @@ -1050,6 +1029,7 @@ class HyVideoSampler: "samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ), "denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), "stg_args": ("STGARGS", ), + "context_options": ("COGCONTEXT", ), } } @@ -1058,7 +1038,8 @@ class HyVideoSampler: FUNCTION = "process" 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): + 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): model = model.model device = mm.get_torch_device() @@ -1100,15 +1081,6 @@ class HyVideoSampler: target_height = align_to(height, 16) target_width = align_to(width, 16) - freqs_cos, freqs_sin = get_rotary_pos_embed( - transformer, num_frames, target_height, target_width - ) - n_tokens = freqs_cos.shape[0] - freqs_cos = freqs_cos.to(dtype).to(device) - freqs_sin = freqs_sin.to(dtype).to(device) - - - model["pipe"].scheduler.shift = flow_shift if model["block_swap_args"] is not None: @@ -1150,13 +1122,12 @@ class HyVideoSampler: denoise_strength=denoise_strength, prompt_embed_dict=hyvid_embeds, generator=generator, - freqs_cis=(freqs_cos, freqs_sin), - n_tokens=n_tokens, stg_mode=stg_args["stg_mode"] if stg_args is not None else None, stg_block_idx=stg_args["stg_block_idx"] if stg_args is not None else -1, stg_scale=stg_args["stg_scale"] if stg_args is not None else 0.0, stg_start_percent=stg_args["stg_start_percent"] if stg_args is not None else 0.0, stg_end_percent=stg_args["stg_end_percent"] if stg_args is not None else 1.0, + context_options=context_options, ) print_memory(device) @@ -1403,6 +1374,7 @@ NODE_CLASS_MAPPINGS = { "HyVideoLoraBlockEdit": HyVideoLoraBlockEdit, "HyVideoTextEmbedsSave": HyVideoTextEmbedsSave, "HyVideoTextEmbedsLoad": HyVideoTextEmbedsLoad, + "HyVideoContextOptions": HyVideoContextOptions, } NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoSampler": "HunyuanVideo Sampler", @@ -1423,4 +1395,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoLoraBlockEdit": "HunyuanVideo Lora Block Edit", "HyVideoTextEmbedsSave": "HunyuanVideo TextEmbeds Save", "HyVideoTextEmbedsLoad": "HunyuanVideo TextEmbeds Load", + "HyVideoContextOptions": "HunyuanVideo Context Options", } diff --git a/nodes_rf_inversion.py b/nodes_rf_inversion.py index 4a6982b..0eec2a1 100644 --- a/nodes_rf_inversion.py +++ b/nodes_rf_inversion.py @@ -4,10 +4,9 @@ import gc import os from .utils import log, print_memory -from .hyvideo.utils.data_utils import align_to from diffusers.utils.torch_utils import randn_tensor import comfy.model_management as mm -from .nodes import get_rotary_pos_embed +from .hyvideo.diffusion.pipelines.pipeline_hunyuan_video import get_rotary_pos_embed script_directory = os.path.dirname(os.path.abspath(__file__)) @@ -117,7 +116,7 @@ class HyVideoInverseSampler: f"Input (height, width, video_length) = ({height}, {width}, {num_frames})" ) - freqs_cos, freqs_sin = get_rotary_pos_embed(transformer, num_frames, height, width) + freqs_cos, freqs_sin = get_rotary_pos_embed(transformer, latent_num_frames, height, width) pipeline.scheduler.shift = flow_shift @@ -327,7 +326,7 @@ class HyVideoReSampler: f"Input (height, width, video_length) = ({height}, {width}, {num_frames})" ) - freqs_cos, freqs_sin = get_rotary_pos_embed(transformer, num_frames, height, width) + freqs_cos, freqs_sin = get_rotary_pos_embed(transformer, latent_num_frames, height, width) pipeline.scheduler.shift = flow_shift @@ -505,7 +504,7 @@ class HyVideoPromptMixSampler: f"Input (height, width, video_length) = ({height}, {width}, {num_frames})" ) latent_video_length = (num_frames - 1) // 4 + 1 - freqs_cos, freqs_sin = get_rotary_pos_embed(transformer, num_frames, height, width) + freqs_cos, freqs_sin = get_rotary_pos_embed(transformer, latent_video_length, height, width) pipeline.scheduler.shift = flow_shift