simple context windows with freenoise shuffling

This commit is contained in:
kijai
2024-12-19 11:56:53 +02:00
parent ec0fdf1fc4
commit b6c1705f9f
4 changed files with 408 additions and 110 deletions
+184
View File
@@ -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))
)
@@ -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(
+34 -61
View File
@@ -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",
}
+4 -5
View File
@@ -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