simple context windows with freenoise shuffling
This commit is contained in:
+184
@@ -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(
|
||||
|
||||
@@ -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,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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user