Compare commits
18
Commits
taylorseer
...
develop
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3ce9640497 | ||
|
|
675baf7574 | ||
|
|
9a5d192807 | ||
|
|
14cc8ae0a5 | ||
|
|
0d9fe9b86a | ||
|
|
27e541b0a5 | ||
|
|
9b321d09c5 | ||
|
|
91aabfc857 | ||
|
|
6682a0905e | ||
|
|
f5a243a898 | ||
|
|
00c8864900 | ||
|
|
02ff79b3f4 | ||
|
|
e18afda414 | ||
|
|
64b5d31765 | ||
|
|
468b103d38 | ||
|
|
6ea4d31b41 | ||
|
|
83f0bbb869 | ||
|
|
75190b756a |
File diff suppressed because it is too large
Load Diff
+3
-2
@@ -7,8 +7,9 @@ def fp8_linear_forward(cls, original_dtype, input):
|
|||||||
weight_dtype = cls.weight.dtype
|
weight_dtype = cls.weight.dtype
|
||||||
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
|
if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
|
||||||
if len(input.shape) == 3:
|
if len(input.shape) == 3:
|
||||||
target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn
|
#target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn
|
||||||
inn = input.reshape(-1, input.shape[2]).to(target_dtype)
|
#inn = input.reshape(-1, input.shape[2]).to(target_dtype)
|
||||||
|
inn = input.reshape(-1, input.shape[2]).to(weight_dtype)
|
||||||
w = cls.weight.t()
|
w = cls.weight.t()
|
||||||
|
|
||||||
scale = torch.ones((1), device=input.device, dtype=torch.float32)
|
scale = torch.ones((1), device=input.device, dtype=torch.float32)
|
||||||
|
|||||||
@@ -33,14 +33,15 @@ from diffusers.schedulers import DPMSolverMultistepScheduler
|
|||||||
|
|
||||||
from ...modules import HYVideoDiffusionTransformer
|
from ...modules import HYVideoDiffusionTransformer
|
||||||
from comfy.utils import ProgressBar
|
from comfy.utils import ProgressBar
|
||||||
|
import math
|
||||||
|
from ....utils import optimized_scale, fourier_filter
|
||||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||||
|
|
||||||
EXAMPLE_DOC_STRING = """"""
|
EXAMPLE_DOC_STRING = """"""
|
||||||
from ...modules.posemb_layers import get_nd_rotary_pos_embed
|
from ...modules.posemb_layers import get_nd_rotary_pos_embed, get_nd_rotary_pos_embed_new
|
||||||
from ....enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight
|
from ....enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight
|
||||||
|
|
||||||
def get_rotary_pos_embed(transformer, latent_video_length, height, width, k=0):
|
def get_rotary_pos_embed(transformer, latent_video_length, height, width, k=0, rope_func=get_nd_rotary_pos_embed):
|
||||||
target_ndim = 3
|
target_ndim = 3
|
||||||
ndim = 5 - 2
|
ndim = 5 - 2
|
||||||
rope_theta = 225
|
rope_theta = 225
|
||||||
@@ -79,7 +80,7 @@ def get_rotary_pos_embed(transformer, latent_video_length, height, width, k=0):
|
|||||||
assert (
|
assert (
|
||||||
sum(rope_dim_list) == head_dim
|
sum(rope_dim_list) == head_dim
|
||||||
), "sum(rope_dim_list) should equal to head_dim of attention layer"
|
), "sum(rope_dim_list) should equal to head_dim of attention layer"
|
||||||
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
|
freqs_cos, freqs_sin = rope_func(
|
||||||
rope_dim_list,
|
rope_dim_list,
|
||||||
rope_sizes,
|
rope_sizes,
|
||||||
theta=rope_theta,
|
theta=rope_theta,
|
||||||
@@ -254,6 +255,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
)
|
)
|
||||||
if latents is not None:
|
if latents is not None:
|
||||||
latents = latents.to(device)
|
latents = latents.to(device)
|
||||||
|
else:
|
||||||
|
original_latents = None
|
||||||
noise = randn_tensor(shape, generator=generator, device=device, dtype=self.base_dtype)
|
noise = randn_tensor(shape, generator=generator, device=device, dtype=self.base_dtype)
|
||||||
|
|
||||||
if freenoise:
|
if freenoise:
|
||||||
@@ -319,6 +322,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
latents = latents[:, :, :frames_needed, :, :]
|
latents = latents[:, :, :frames_needed, :, :]
|
||||||
logger.info(f"Frames needed less than current frames, cutting down to {frames_needed}")
|
logger.info(f"Frames needed less than current frames, cutting down to {frames_needed}")
|
||||||
|
|
||||||
|
original_latents = latents.clone()
|
||||||
latents = latents * (1 - latent_timestep / 1000) + latent_timestep / 1000 * noise
|
latents = latents * (1 - latent_timestep / 1000) + latent_timestep / 1000 * noise
|
||||||
print("latents shape:", latents.shape)
|
print("latents shape:", latents.shape)
|
||||||
|
|
||||||
@@ -338,7 +342,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
if hasattr(self.scheduler, "init_noise_sigma"):
|
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||||
# scale the initial noise by the standard deviation required by the scheduler
|
# scale the initial noise by the standard deviation required by the scheduler
|
||||||
latents = latents * self.scheduler.init_noise_sigma
|
latents = latents * self.scheduler.init_noise_sigma
|
||||||
return latents.to(device), timesteps, i2v_mask, image_cond_latents
|
return latents.to(device), timesteps, i2v_mask, image_cond_latents, noise, original_latents
|
||||||
|
|
||||||
# Copied from diffusers.pipelines.latent_consistency_models.pipeline_latent_consistency_text2img.LatentConsistencyModelPipeline.get_guidance_scale_embedding
|
# Copied from diffusers.pipelines.latent_consistency_models.pipeline_latent_consistency_text2img.LatentConsistencyModelPipeline.get_guidance_scale_embedding
|
||||||
def get_guidance_scale_embedding(
|
def get_guidance_scale_embedding(
|
||||||
@@ -423,6 +427,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
timesteps: List[int] = None,
|
timesteps: List[int] = None,
|
||||||
sigmas: List[float] = None,
|
sigmas: List[float] = None,
|
||||||
guidance_scale: float = 1.0,
|
guidance_scale: float = 1.0,
|
||||||
|
use_cfg_zero_star: bool = False,
|
||||||
|
fresca_args: Optional[Dict[str, Any]] = None,
|
||||||
|
slg_args: Optional[Dict[str, Any]] = None,
|
||||||
cfg_start_percent: float = 0.0,
|
cfg_start_percent: float = 0.0,
|
||||||
cfg_end_percent: float = 1.0,
|
cfg_end_percent: float = 1.0,
|
||||||
batched_cfg: bool = True,
|
batched_cfg: bool = True,
|
||||||
@@ -431,6 +438,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
denoise_strength: float = 1.0,
|
denoise_strength: float = 1.0,
|
||||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||||
latents: Optional[torch.Tensor] = None,
|
latents: Optional[torch.Tensor] = None,
|
||||||
|
mask_latents: Optional[torch.Tensor] = None,
|
||||||
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||||
guidance_rescale: float = 0.0,
|
guidance_rescale: float = 0.0,
|
||||||
clip_skip: Optional[int] = None,
|
clip_skip: Optional[int] = None,
|
||||||
@@ -452,8 +460,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
feta_args: Optional[Dict] = None,
|
feta_args: Optional[Dict] = None,
|
||||||
leapfusion_img2vid: Optional[bool] = False,
|
leapfusion_img2vid: Optional[bool] = False,
|
||||||
image_cond_latents: Optional[torch.Tensor] = None,
|
image_cond_latents: Optional[torch.Tensor] = None,
|
||||||
|
neg_image_cond_latents: Optional[torch.Tensor] = None,
|
||||||
riflex_freq_index: Optional[int] = None,
|
riflex_freq_index: Optional[int] = None,
|
||||||
i2v_stability=True,
|
i2v_stability=True,
|
||||||
|
loop_args: Optional[Dict] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
r"""
|
r"""
|
||||||
@@ -539,6 +549,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
# 2. Define call parameters
|
# 2. Define call parameters
|
||||||
|
|
||||||
batch_size = 1
|
batch_size = 1
|
||||||
|
ref_latents = None
|
||||||
device = self._execution_device
|
device = self._execution_device
|
||||||
|
|
||||||
prompt_embeds = prompt_embed_dict.get("prompt_embeds", None)
|
prompt_embeds = prompt_embed_dict.get("prompt_embeds", None)
|
||||||
@@ -633,14 +644,25 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
use_context_schedule = True
|
use_context_schedule = True
|
||||||
from ....context import get_context_scheduler
|
from ....context import get_context_scheduler
|
||||||
context = get_context_scheduler(context_schedule)
|
context = get_context_scheduler(context_schedule)
|
||||||
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
if i2v_condition_type == "reference":
|
||||||
self.transformer, context_frames, height, width
|
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
||||||
)
|
self.transformer, context_frames, height, width, rope_func=get_nd_rotary_pos_embed_new
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
||||||
|
self.transformer, context_frames, height, width
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
# rotary embeddings
|
# rotary embeddings
|
||||||
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
if i2v_condition_type == "reference":
|
||||||
self.transformer, latent_video_length, height, width, k=riflex_freq_index
|
print("Using reference condition")
|
||||||
)
|
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
||||||
|
self.transformer, latent_video_length, height, width, rope_func=get_nd_rotary_pos_embed_new
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
||||||
|
self.transformer, latent_video_length, height, width, k=riflex_freq_index
|
||||||
|
)
|
||||||
if not self.transformer.upcast_rope:
|
if not self.transformer.upcast_rope:
|
||||||
freqs_cos = freqs_cos.to(self.base_dtype).to(device)
|
freqs_cos = freqs_cos.to(self.base_dtype).to(device)
|
||||||
freqs_sin = freqs_sin.to(self.base_dtype).to(device)
|
freqs_sin = freqs_sin.to(self.base_dtype).to(device)
|
||||||
@@ -651,10 +673,11 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
if leapfusion_img2vid:
|
if leapfusion_img2vid:
|
||||||
logger.info("Single input latent frame detected, LeapFusion img2vid enabled")
|
logger.info("Single input latent frame detected, LeapFusion img2vid enabled")
|
||||||
original_latents = latents
|
original_latents = latents
|
||||||
|
|
||||||
# 5. Prepare latent variables
|
# 5. Prepare latent variables
|
||||||
#num_channels_latents = self.transformer.config.in_channels
|
#num_channels_latents = self.transformer.config.in_channels
|
||||||
num_channels_latents = 16
|
num_channels_latents = 16
|
||||||
latents, timesteps, i2v_mask, image_cond_latents = self.prepare_latents(
|
latents, timesteps, i2v_mask, image_cond_latents, noise, original_latents = self.prepare_latents(
|
||||||
batch_size * num_videos_per_prompt,
|
batch_size * num_videos_per_prompt,
|
||||||
num_channels_latents,
|
num_channels_latents,
|
||||||
num_inference_steps,
|
num_inference_steps,
|
||||||
@@ -690,6 +713,37 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
|
|
||||||
#print(self.scheduler.sigmas)
|
#print(self.scheduler.sigmas)
|
||||||
|
|
||||||
|
latent_shift_loop = False
|
||||||
|
if loop_args is not None:
|
||||||
|
latent_shift_loop = True
|
||||||
|
is_looped = True
|
||||||
|
latent_skip = loop_args["shift_skip"]
|
||||||
|
latent_shift_start_percent = loop_args["start_percent"]
|
||||||
|
latent_shift_end_percent = loop_args["end_percent"]
|
||||||
|
shift_idx = 0
|
||||||
|
|
||||||
|
if mask_latents is not None:
|
||||||
|
mask_latents_model_input = (
|
||||||
|
torch.cat([mask_latents] * 2)
|
||||||
|
if not math.isclose(self.guidance_scale, 1.0)
|
||||||
|
else mask_latents
|
||||||
|
)
|
||||||
|
print(f'mask_latents_model_input={mask_latents_model_input.shape} ')
|
||||||
|
|
||||||
|
if fresca_args is not None:
|
||||||
|
fresca_scale_low = fresca_args.get("fresca_scale_low", 1.0)
|
||||||
|
fresca_scale_high = fresca_args.get("fresca_scale_high", 1.25)
|
||||||
|
fresca_freq_cutoff = fresca_args.get("fresca_freq_cutoff", 20)
|
||||||
|
|
||||||
|
if slg_args is not None:
|
||||||
|
assert batched_cfg is not None, "Batched cfg is not supported with SLG"
|
||||||
|
self.transformer.slg_single_blocks = slg_args["single_blocks"]
|
||||||
|
self.transformer.slg_double_blocks = slg_args["double_blocks"]
|
||||||
|
self.transformer.slg_start_percent = slg_args["start_percent"]
|
||||||
|
self.transformer.slg_end_percent = slg_args["end_percent"]
|
||||||
|
else:
|
||||||
|
self.transformer.slg_single_blocks = self.transformer.slg_double_blocks = None
|
||||||
|
|
||||||
logger.info(f"Sampling {video_length} frames in {latents.shape[2]} latents at {width}x{height} with {len(timesteps)} inference steps")
|
logger.info(f"Sampling {video_length} frames in {latents.shape[2]} latents at {width}x{height} with {len(timesteps)} inference steps")
|
||||||
|
|
||||||
comfy_pbar = ProgressBar(len(timesteps))
|
comfy_pbar = ProgressBar(len(timesteps))
|
||||||
@@ -698,8 +752,16 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
if self.interrupt:
|
if self.interrupt:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
current_step_percentage = i / len(timesteps)
|
||||||
|
|
||||||
if image_cond_latents is not None and i2v_condition_type == "token_replace":
|
if image_cond_latents is not None and i2v_condition_type == "token_replace":
|
||||||
latents = torch.concat([original_image_latents, latents[:, :, 1:, :, :]], dim=2)
|
latents = torch.concat([original_image_latents, latents[:, :, 1:, :, :]], dim=2)
|
||||||
|
elif image_cond_latents is not None and i2v_condition_type == "reference":
|
||||||
|
ref_latents = image_cond_latents
|
||||||
|
if neg_image_cond_latents is not None:
|
||||||
|
uncond_ref_latents = neg_image_cond_latents
|
||||||
|
else:
|
||||||
|
uncond_ref_latents = image_cond_latents
|
||||||
|
|
||||||
latent_model_input = latents
|
latent_model_input = latents
|
||||||
input_prompt_embeds = prompt_embeds
|
input_prompt_embeds = prompt_embeds
|
||||||
@@ -708,7 +770,12 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
cfg_enabled = False
|
cfg_enabled = False
|
||||||
stg_enabled = False
|
stg_enabled = False
|
||||||
|
|
||||||
current_step_percentage = i / len(timesteps)
|
### latent shift
|
||||||
|
if latent_shift_loop:
|
||||||
|
if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent:
|
||||||
|
latent_model_input = torch.cat([latent_model_input[:, :, shift_idx:]] + [latent_model_input[:, :, :shift_idx]], dim=2)
|
||||||
|
|
||||||
|
|
||||||
if self.do_spatio_temporal_guidance:
|
if self.do_spatio_temporal_guidance:
|
||||||
if stg_start_percent <= current_step_percentage <= stg_end_percent:
|
if stg_start_percent <= current_step_percentage <= stg_end_percent:
|
||||||
stg_enabled = True
|
stg_enabled = True
|
||||||
@@ -745,6 +812,16 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
|
|
||||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||||
|
|
||||||
|
if mask_latents is not None:
|
||||||
|
original_latents_noise = original_latents * (1 - t / 1000.0) + t / 1000.0 * noise
|
||||||
|
original_latent_noise_model_input = (
|
||||||
|
torch.cat([original_latents_noise] * 2)
|
||||||
|
if self.do_classifier_free_guidance
|
||||||
|
else original_latents_noise
|
||||||
|
)
|
||||||
|
original_latent_noise_model_input = self.scheduler.scale_model_input(original_latent_noise_model_input, t)
|
||||||
|
latent_model_input = mask_latents_model_input * latent_model_input + (1 - mask_latents_model_input) * original_latent_noise_model_input
|
||||||
|
|
||||||
t_expand = t.repeat(latent_model_input.shape[0])
|
t_expand = t.repeat(latent_model_input.shape[0])
|
||||||
|
|
||||||
if leapfusion_img2vid:
|
if leapfusion_img2vid:
|
||||||
@@ -815,6 +892,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
stg_block_idx=stg_block_idx,
|
stg_block_idx=stg_block_idx,
|
||||||
stg_mode=stg_mode,
|
stg_mode=stg_mode,
|
||||||
return_dict=True,
|
return_dict=True,
|
||||||
|
ref_latents=ref_latents
|
||||||
)["x"]
|
)["x"]
|
||||||
window_mask = torch.ones_like(noise_pred_context)
|
window_mask = torch.ones_like(noise_pred_context)
|
||||||
|
|
||||||
@@ -851,6 +929,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
stg_block_idx=stg_block_idx,
|
stg_block_idx=stg_block_idx,
|
||||||
stg_mode=stg_mode,
|
stg_mode=stg_mode,
|
||||||
return_dict=True,
|
return_dict=True,
|
||||||
|
ref_latents=ref_latents,
|
||||||
|
is_uncond = False,
|
||||||
|
current_step = i,
|
||||||
|
current_step_percentage = current_step_percentage
|
||||||
)["x"]
|
)["x"]
|
||||||
else:
|
else:
|
||||||
uncond = self.transformer(
|
uncond = self.transformer(
|
||||||
@@ -861,10 +943,14 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
text_states_2=input_prompt_embeds_2[0].unsqueeze(0),
|
text_states_2=input_prompt_embeds_2[0].unsqueeze(0),
|
||||||
freqs_cos=freqs_cos,
|
freqs_cos=freqs_cos,
|
||||||
freqs_sin=freqs_sin,
|
freqs_sin=freqs_sin,
|
||||||
guidance=guidance_expand[0].unsqueeze(0),
|
guidance=guidance_expand[0].unsqueeze(0) if guidance_expand is not None else None,
|
||||||
stg_block_idx=stg_block_idx,
|
stg_block_idx=stg_block_idx,
|
||||||
stg_mode=stg_mode,
|
stg_mode=stg_mode,
|
||||||
return_dict=True,
|
return_dict=True,
|
||||||
|
ref_latents=uncond_ref_latents,
|
||||||
|
is_uncond = True,
|
||||||
|
current_step = i,
|
||||||
|
current_step_percentage = current_step_percentage
|
||||||
)["x"]
|
)["x"]
|
||||||
cond = self.transformer(
|
cond = self.transformer(
|
||||||
latent_model_input[1].unsqueeze(0),
|
latent_model_input[1].unsqueeze(0),
|
||||||
@@ -874,22 +960,40 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
text_states_2=input_prompt_embeds_2[1].unsqueeze(0),
|
text_states_2=input_prompt_embeds_2[1].unsqueeze(0),
|
||||||
freqs_cos=freqs_cos,
|
freqs_cos=freqs_cos,
|
||||||
freqs_sin=freqs_sin,
|
freqs_sin=freqs_sin,
|
||||||
guidance=guidance_expand[1].unsqueeze(0),
|
guidance=guidance_expand[1].unsqueeze(0) if guidance_expand is not None else None,
|
||||||
stg_block_idx=stg_block_idx,
|
stg_block_idx=stg_block_idx,
|
||||||
stg_mode=stg_mode,
|
stg_mode=stg_mode,
|
||||||
return_dict=True,
|
return_dict=True,
|
||||||
|
ref_latents=ref_latents,
|
||||||
|
is_uncond = False,
|
||||||
|
current_step = i,
|
||||||
|
current_step_percentage = current_step_percentage
|
||||||
)["x"]
|
)["x"]
|
||||||
|
|
||||||
# perform guidance
|
# perform guidance
|
||||||
if cfg_enabled and not self.do_spatio_temporal_guidance:
|
if cfg_enabled and not self.do_spatio_temporal_guidance:
|
||||||
if batched_cfg:
|
if batched_cfg:
|
||||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
uncond, cond = noise_pred.chunk(2)
|
||||||
noise_pred = noise_pred_uncond + self.guidance_scale * (
|
|
||||||
noise_pred_text - noise_pred_uncond
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
noise_pred = uncond + self.guidance_scale * (cond - uncond)
|
|
||||||
|
|
||||||
|
#https://github.com/WeichenFan/CFG-Zero-star/
|
||||||
|
if use_cfg_zero_star:
|
||||||
|
alpha = optimized_scale(
|
||||||
|
cond.view(batch_size, -1),
|
||||||
|
uncond.view(batch_size, -1)
|
||||||
|
).view(batch_size, 1, 1, 1)
|
||||||
|
else:
|
||||||
|
alpha = 1.0
|
||||||
|
#https://github.com/WikiChao/FreSca
|
||||||
|
if fresca_args is not None:
|
||||||
|
filtered_cond = fourier_filter(
|
||||||
|
cond - uncond,
|
||||||
|
scale_low=fresca_scale_low,
|
||||||
|
scale_high=fresca_scale_high,
|
||||||
|
freq_cutoff=fresca_freq_cutoff,
|
||||||
|
)
|
||||||
|
noise_pred = uncond * alpha + self.guidance_scale * filtered_cond * alpha
|
||||||
|
else:
|
||||||
|
noise_pred = uncond * alpha + self.guidance_scale * (cond - uncond * alpha)
|
||||||
|
|
||||||
elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance:
|
elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance:
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
@@ -904,6 +1008,19 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
noise_pred = noise_pred_text + self._stg_scale * (
|
noise_pred = noise_pred_text + self._stg_scale * (
|
||||||
noise_pred_text - noise_pred_perturb
|
noise_pred_text - noise_pred_perturb
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
if fresca_args is not None:
|
||||||
|
noise_pred = fourier_filter(
|
||||||
|
noise_pred,
|
||||||
|
scale_low=fresca_scale_low,
|
||||||
|
scale_high=fresca_scale_high,
|
||||||
|
freq_cutoff=fresca_freq_cutoff,
|
||||||
|
)
|
||||||
|
if latent_shift_loop:
|
||||||
|
#reverse latent shift
|
||||||
|
if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent:
|
||||||
|
noise_pred = torch.cat([noise_pred[:, :, latent_video_length - shift_idx:]] + [noise_pred[:, :, :latent_video_length - shift_idx]], dim=2)
|
||||||
|
shift_idx = (shift_idx + latent_skip) % latent_video_length
|
||||||
|
|
||||||
# compute the previous noisy sample x_t -> x_t-1
|
# compute the previous noisy sample x_t -> x_t-1
|
||||||
if image_cond_latents is not None and i2v_condition_type == "token_replace":
|
if image_cond_latents is not None and i2v_condition_type == "token_replace":
|
||||||
@@ -950,6 +1067,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
|||||||
else:
|
else:
|
||||||
comfy_pbar.update(1)
|
comfy_pbar.update(1)
|
||||||
|
|
||||||
|
if mask_latents is not None:
|
||||||
|
latents = mask_latents * latents + (1 - mask_latents) * original_latents
|
||||||
|
|
||||||
if image_cond_latents is not None:
|
if image_cond_latents is not None:
|
||||||
if leapfusion_img2vid or i2v_condition_type == "latent_concat":
|
if leapfusion_img2vid or i2v_condition_type == "latent_concat":
|
||||||
latents = latents[:, :, 1:, :, :]
|
latents = latents[:, :, 1:, :, :]
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import torch.nn as nn
|
|||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
from comfy.utils import load_torch_file
|
from comfy.utils import load_torch_file
|
||||||
|
|
||||||
|
@torch.compiler.disable()
|
||||||
def get_fp_maxval(bits=8, mantissa_bit=3, sign_bits=1):
|
def get_fp_maxval(bits=8, mantissa_bit=3, sign_bits=1):
|
||||||
_bits = torch.tensor(bits)
|
_bits = torch.tensor(bits)
|
||||||
_mantissa_bit = torch.tensor(mantissa_bit)
|
_mantissa_bit = torch.tensor(mantissa_bit)
|
||||||
@@ -17,6 +18,7 @@ def get_fp_maxval(bits=8, mantissa_bit=3, sign_bits=1):
|
|||||||
maxval = mantissa * 2 ** (2**E - 1 - bias)
|
maxval = mantissa * 2 ** (2**E - 1 - bias)
|
||||||
return maxval
|
return maxval
|
||||||
|
|
||||||
|
@torch.compiler.disable()
|
||||||
def quantize_to_fp8(x, bits=8, mantissa_bit=3, sign_bits=1):
|
def quantize_to_fp8(x, bits=8, mantissa_bit=3, sign_bits=1):
|
||||||
"""
|
"""
|
||||||
Default is E4M3.
|
Default is E4M3.
|
||||||
@@ -40,6 +42,7 @@ def quantize_to_fp8(x, bits=8, mantissa_bit=3, sign_bits=1):
|
|||||||
qdq_out = torch.round(input_clamp / log_scales) * log_scales
|
qdq_out = torch.round(input_clamp / log_scales) * log_scales
|
||||||
return qdq_out, log_scales
|
return qdq_out, log_scales
|
||||||
|
|
||||||
|
@torch.compiler.disable()
|
||||||
def fp8_tensor_quant(x, scale, bits=8, mantissa_bit=3, sign_bits=1):
|
def fp8_tensor_quant(x, scale, bits=8, mantissa_bit=3, sign_bits=1):
|
||||||
for i in range(len(x.shape) - 1):
|
for i in range(len(x.shape) - 1):
|
||||||
scale = scale.unsqueeze(-1)
|
scale = scale.unsqueeze(-1)
|
||||||
@@ -47,10 +50,10 @@ def fp8_tensor_quant(x, scale, bits=8, mantissa_bit=3, sign_bits=1):
|
|||||||
quant_dequant_x, log_scales = quantize_to_fp8(new_x, bits=bits, mantissa_bit=mantissa_bit, sign_bits=sign_bits)
|
quant_dequant_x, log_scales = quantize_to_fp8(new_x, bits=bits, mantissa_bit=mantissa_bit, sign_bits=sign_bits)
|
||||||
return quant_dequant_x, scale, log_scales
|
return quant_dequant_x, scale, log_scales
|
||||||
|
|
||||||
def fp8_activation_dequant(qdq_out, scale, dtype):
|
@torch.compiler.disable()
|
||||||
|
def fp8_activation_dequant(qdq_out, dtype):
|
||||||
qdq_out = qdq_out.type(dtype)
|
qdq_out = qdq_out.type(dtype)
|
||||||
quant_dequant_x = qdq_out * scale.to(dtype)
|
return qdq_out
|
||||||
return quant_dequant_x
|
|
||||||
|
|
||||||
def fp8_linear_forward(cls, original_dtype, input):
|
def fp8_linear_forward(cls, original_dtype, input):
|
||||||
weight_dtype = cls.weight.dtype
|
weight_dtype = cls.weight.dtype
|
||||||
@@ -62,33 +65,33 @@ def fp8_linear_forward(cls, original_dtype, input):
|
|||||||
linear_weight = linear_weight.to(torch.float8_e4m3fn)
|
linear_weight = linear_weight.to(torch.float8_e4m3fn)
|
||||||
weight_dtype = linear_weight.dtype
|
weight_dtype = linear_weight.dtype
|
||||||
else:
|
else:
|
||||||
scale = cls.fp8_scale.to(cls.weight.device)
|
scale = cls.fp8_scale#.to(cls.weight.device)
|
||||||
linear_weight = cls.weight
|
linear_weight = cls.weight
|
||||||
#####
|
#####
|
||||||
|
|
||||||
if weight_dtype == torch.float8_e4m3fn and cls.weight.sum() != 0:
|
#if weight_dtype == torch.float8_e4m3fn and cls.weight.sum() != 0:
|
||||||
if True or len(input.shape) == 3:
|
if weight_dtype == torch.float8_e4m3fn:
|
||||||
cls_dequant = fp8_activation_dequant(linear_weight, scale, original_dtype)
|
qdq_out = fp8_activation_dequant(linear_weight, original_dtype)
|
||||||
if cls.bias != None:
|
cls_dequant = qdq_out * scale
|
||||||
output = F.linear(input, cls_dequant, cls.bias)
|
if cls.bias != None:
|
||||||
else:
|
output = F.linear(input, cls_dequant, cls.bias)
|
||||||
output = F.linear(input, cls_dequant)
|
|
||||||
return output
|
|
||||||
else:
|
else:
|
||||||
return cls.original_forward(input.to(original_dtype))
|
output = F.linear(input, cls_dequant)
|
||||||
|
return output
|
||||||
else:
|
else:
|
||||||
return cls.original_forward(input)
|
return cls.original_forward(input)
|
||||||
|
|
||||||
def convert_fp8_linear(module, original_dtype):
|
def convert_fp8_linear(module, original_dtype, device, fp8_scale_map={}):
|
||||||
setattr(module, "fp8_matmul_enabled", True)
|
setattr(module, "fp8_matmul_enabled", True)
|
||||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
# loading fp8 mapping file
|
# loading fp8 mapping file
|
||||||
fp8_map_path = os.path.join(script_directory,"fp8_map.safetensors")
|
if not fp8_scale_map:
|
||||||
if os.path.exists(fp8_map_path):
|
fp8_map_path = os.path.join(script_directory,"fp8_map.safetensors")
|
||||||
fp8_map = load_torch_file(fp8_map_path, safe_load=True)
|
if os.path.exists(fp8_map_path):
|
||||||
else:
|
fp8_map = load_torch_file(fp8_map_path, safe_load=True)
|
||||||
raise ValueError(f"Invalid fp8_map path: {fp8_map_path}.")
|
else:
|
||||||
|
raise ValueError(f"Invalid fp8_map path: {fp8_map_path}.")
|
||||||
|
|
||||||
#fp8_layers = []
|
#fp8_layers = []
|
||||||
for key, layer in module.named_modules():
|
for key, layer in module.named_modules():
|
||||||
@@ -96,6 +99,6 @@ def convert_fp8_linear(module, original_dtype):
|
|||||||
#fp8_layers.append(key)
|
#fp8_layers.append(key)
|
||||||
original_forward = layer.forward
|
original_forward = layer.forward
|
||||||
#layer.weight = torch.nn.Parameter(layer.weight.to(torch.float8_e4m3fn))
|
#layer.weight = torch.nn.Parameter(layer.weight.to(torch.float8_e4m3fn))
|
||||||
setattr(layer, "fp8_scale", fp8_map[key].to(dtype=original_dtype))
|
setattr(layer, "fp8_scale", fp8_map[key].to(device=device, dtype=original_dtype))
|
||||||
setattr(layer, "original_forward", original_forward)
|
setattr(layer, "original_forward", original_forward)
|
||||||
setattr(layer, "forward", lambda input, m=layer: fp8_linear_forward(m, original_dtype, input))
|
setattr(layer, "forward", lambda input, m=layer: fp8_linear_forward(m, original_dtype, input))
|
||||||
|
|||||||
+78
-18
@@ -752,15 +752,24 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
|||||||
self.enable_teacache = False
|
self.enable_teacache = False
|
||||||
self.cnt = 0
|
self.cnt = 0
|
||||||
self.num_steps = 0
|
self.num_steps = 0
|
||||||
self.teacache_skipped_steps = 0
|
self.teacache_skipped_steps_cond = 0
|
||||||
|
self.teacache_skipped_steps_uncond = 0
|
||||||
|
self.teacache_start_step = 0
|
||||||
|
self.teacache_end_step = 100
|
||||||
self.rel_l1_thresh = 0.15
|
self.rel_l1_thresh = 0.15
|
||||||
self.accumulated_rel_l1_distance = 0
|
self.accumulated_rel_l1_distance = 0
|
||||||
self.previous_modulated_input = None
|
self.previous_modulated_input_cond = None
|
||||||
self.previous_residual = None
|
self.previous_modulated_input_uncond = None
|
||||||
self.last_dimensions = None
|
self.previous_residual_cond = None
|
||||||
self.last_frame_count = None
|
self.previous_residual_uncond = None
|
||||||
self.teacache_device = None
|
self.teacache_device = None
|
||||||
|
|
||||||
|
#slg
|
||||||
|
self.slg_single_blocks = None
|
||||||
|
self.slg_double_blocks = None
|
||||||
|
self.slg_start_percent = 0.0
|
||||||
|
self.slg_end_percent = 1.0
|
||||||
|
|
||||||
# thanks @2kpr for the initial block swap code!
|
# thanks @2kpr for the initial block swap code!
|
||||||
def block_swap(self, double_blocks_to_swap, single_blocks_to_swap, offload_txt_in=False, offload_img_in=False):
|
def block_swap(self, double_blocks_to_swap, single_blocks_to_swap, offload_txt_in=False, offload_img_in=False):
|
||||||
print(f"Swapping {double_blocks_to_swap + 1} double blocks and {single_blocks_to_swap + 1} single blocks")
|
print(f"Swapping {double_blocks_to_swap + 1} double blocks and {single_blocks_to_swap + 1} single blocks")
|
||||||
@@ -950,10 +959,19 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
|||||||
stg_mode: str = None,
|
stg_mode: str = None,
|
||||||
stg_block_idx: int = -1,
|
stg_block_idx: int = -1,
|
||||||
return_dict: bool = True,
|
return_dict: bool = True,
|
||||||
|
ref_latents: torch.Tensor = None,
|
||||||
|
is_uncond = False,
|
||||||
|
current_step: int = 0,
|
||||||
|
current_step_percentage: float = 0,
|
||||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||||
|
|
||||||
def _process_double_blocks(img, txt, vec, block_args):
|
def _process_double_blocks(img, txt, vec, block_args):
|
||||||
for b, block in enumerate(self.double_blocks):
|
for b, block in enumerate(self.double_blocks):
|
||||||
|
if self.slg_double_blocks is not None:
|
||||||
|
if b in self.slg_double_blocks and is_uncond:
|
||||||
|
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
|
||||||
|
print(f"Skipping double block {b}")
|
||||||
|
continue
|
||||||
if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0:
|
if b <= self.double_blocks_to_swap and self.double_blocks_to_swap >= 0:
|
||||||
block.to(self.main_device)
|
block.to(self.main_device)
|
||||||
|
|
||||||
@@ -965,6 +983,11 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
|||||||
|
|
||||||
def _process_single_blocks(x, vec, txt_seq_len, block_args, stg_mode=None, stg_block_idx=None):
|
def _process_single_blocks(x, vec, txt_seq_len, block_args, stg_mode=None, stg_block_idx=None):
|
||||||
for b, block in enumerate(self.single_blocks):
|
for b, block in enumerate(self.single_blocks):
|
||||||
|
if self.slg_single_blocks is not None:
|
||||||
|
if b in self.slg_single_blocks and is_uncond:
|
||||||
|
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
|
||||||
|
print(f"Skipping single block {b}")
|
||||||
|
continue
|
||||||
if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0:
|
if b <= self.single_blocks_to_swap and self.single_blocks_to_swap >= 0:
|
||||||
block.to(self.main_device)
|
block.to(self.main_device)
|
||||||
|
|
||||||
@@ -1030,6 +1053,12 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
|||||||
self.img_in.to(self.main_device)
|
self.img_in.to(self.main_device)
|
||||||
|
|
||||||
img = self.img_in(img)
|
img = self.img_in(img)
|
||||||
|
|
||||||
|
if ref_latents is not None:
|
||||||
|
ref_latents = self.img_in(ref_latents)
|
||||||
|
ref_length = ref_latents.shape[-2]
|
||||||
|
img = torch.cat([ref_latents, img], dim=-2) # t c
|
||||||
|
|
||||||
if self.text_projection == "linear":
|
if self.text_projection == "linear":
|
||||||
txt = self.txt_in(txt)
|
txt = self.txt_in(txt)
|
||||||
elif self.text_projection == "single_refiner":
|
elif self.text_projection == "single_refiner":
|
||||||
@@ -1070,7 +1099,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
|||||||
]
|
]
|
||||||
|
|
||||||
#tea_cache
|
#tea_cache
|
||||||
if self.enable_teacache:
|
if self.enable_teacache and self.teacache_start_step <= current_step <= self.teacache_end_step:
|
||||||
inp = img.clone()
|
inp = img.clone()
|
||||||
vec_ = vec.clone()
|
vec_ = vec.clone()
|
||||||
txt_ = txt.clone()
|
txt_ = txt.clone()
|
||||||
@@ -1088,29 +1117,51 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
|||||||
normed_inp, shift=img_mod1_shift, scale=img_mod1_scale
|
normed_inp, shift=img_mod1_shift, scale=img_mod1_scale
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Choose the appropriate cache based on whether this is a conditional or unconditional pass
|
||||||
|
previous_modulated_input = self.previous_modulated_input_uncond if is_uncond else self.previous_modulated_input_cond
|
||||||
|
previous_residual = self.previous_residual_uncond if is_uncond else self.previous_residual_cond
|
||||||
|
accumulated_rel_l1_distance = self.accumulated_rel_l1_distance_uncond if is_uncond else self.accumulated_rel_l1_distance_cond
|
||||||
|
|
||||||
if self.cnt == 0 or self.cnt == self.num_steps-1:
|
if self.cnt == 0 or self.cnt == self.num_steps-1:
|
||||||
should_calc = True
|
should_calc = True
|
||||||
self.accumulated_rel_l1_distance = 0
|
accumulated_rel_l1_distance = 0
|
||||||
self.previous_modulated_input = modulated_inp.clone()
|
previous_modulated_input = modulated_inp.clone()
|
||||||
else:
|
else:
|
||||||
coefficients = [7.33226126e+02, -4.01131952e+02, 6.75869174e+01, -3.14987800e+00, 9.61237896e-02]
|
coefficients = [7.33226126e+02, -4.01131952e+02, 6.75869174e+01, -3.14987800e+00, 9.61237896e-02]
|
||||||
rescale_func = np.poly1d(coefficients)
|
rescale_func = np.poly1d(coefficients)
|
||||||
self.accumulated_rel_l1_distance += rescale_func(((modulated_inp-self.previous_modulated_input).abs().mean() / self.previous_modulated_input.abs().mean()).cpu().item())
|
if previous_modulated_input is not None:
|
||||||
if self.accumulated_rel_l1_distance < self.rel_l1_thresh:
|
accumulated_rel_l1_distance += rescale_func(((modulated_inp-previous_modulated_input).abs().mean() / previous_modulated_input.abs().mean()).cpu().item())
|
||||||
should_calc = False
|
if accumulated_rel_l1_distance < self.rel_l1_thresh:
|
||||||
|
should_calc = False
|
||||||
|
else:
|
||||||
|
should_calc = True
|
||||||
|
accumulated_rel_l1_distance = 0
|
||||||
else:
|
else:
|
||||||
should_calc = True
|
should_calc = True
|
||||||
self.accumulated_rel_l1_distance = 0
|
accumulated_rel_l1_distance = 0
|
||||||
self.previous_modulated_input = modulated_inp.clone()
|
|
||||||
|
# Store back the appropriate values
|
||||||
|
if is_uncond:
|
||||||
|
self.previous_modulated_input_uncond = modulated_inp.clone()
|
||||||
|
self.accumulated_rel_l1_distance_uncond = accumulated_rel_l1_distance
|
||||||
|
else:
|
||||||
|
self.previous_modulated_input_cond = modulated_inp.clone()
|
||||||
|
self.accumulated_rel_l1_distance_cond = accumulated_rel_l1_distance
|
||||||
|
|
||||||
self.cnt += 1
|
self.cnt += 1
|
||||||
if self.cnt == self.num_steps:
|
if self.cnt == self.num_steps:
|
||||||
self.cnt = 0
|
self.cnt = 0
|
||||||
|
|
||||||
if not should_calc and self.previous_residual is not None:
|
if not should_calc and previous_residual is not None:
|
||||||
self.teacache_skipped_steps += 1
|
# Increment the appropriate skipped steps counter
|
||||||
|
if is_uncond:
|
||||||
|
self.teacache_skipped_steps_uncond += 1
|
||||||
|
else:
|
||||||
|
self.teacache_skipped_steps_cond += 1
|
||||||
|
|
||||||
# Verify tensor dimensions match before adding
|
# Verify tensor dimensions match before adding
|
||||||
if img.shape == self.previous_residual.shape:
|
if img.shape == previous_residual.shape:
|
||||||
img = img + self.previous_residual.to(img.device)
|
img = img + previous_residual.to(img.device)
|
||||||
else:
|
else:
|
||||||
should_calc = True # Force recalculation if dimensions don't match
|
should_calc = True # Force recalculation if dimensions don't match
|
||||||
|
|
||||||
@@ -1123,7 +1174,13 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
|||||||
x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode, stg_block_idx)
|
x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode, stg_block_idx)
|
||||||
|
|
||||||
img = x[:, :img_seq_len, ...]
|
img = x[:, :img_seq_len, ...]
|
||||||
self.previous_residual = (img - ori_img).to(self.teacache_device)
|
new_residual = (img - ori_img).to(self.teacache_device)
|
||||||
|
|
||||||
|
# Store the new residual in the appropriate cache
|
||||||
|
if is_uncond:
|
||||||
|
self.previous_residual_uncond = new_residual
|
||||||
|
else:
|
||||||
|
self.previous_residual_cond = new_residual
|
||||||
else:
|
else:
|
||||||
# Pass through DiT blocks
|
# Pass through DiT blocks
|
||||||
img, txt = _process_double_blocks(img, txt, vec, block_args)
|
img, txt = _process_double_blocks(img, txt, vec, block_args)
|
||||||
@@ -1132,6 +1189,9 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
|||||||
x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode, stg_block_idx)
|
x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode, stg_block_idx)
|
||||||
img = x[:, :img_seq_len, ...]
|
img = x[:, :img_seq_len, ...]
|
||||||
|
|
||||||
|
if ref_latents is not None:
|
||||||
|
img = img[:, ref_length:]
|
||||||
|
|
||||||
# ---------------------------- Final layer ------------------------------
|
# ---------------------------- Final layer ------------------------------
|
||||||
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
|
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
|
||||||
|
|
||||||
|
|||||||
@@ -302,3 +302,52 @@ def get_1d_rotary_pos_embed_riflex(
|
|||||||
torch.ones_like(freqs), freqs
|
torch.ones_like(freqs), freqs
|
||||||
) # complex64 # [S, D/2]
|
) # complex64 # [S, D/2]
|
||||||
return freqs_cis
|
return freqs_cis
|
||||||
|
|
||||||
|
def get_nd_rotary_pos_embed_new(rope_dim_list, start, *args, theta=10000., use_real=False,
|
||||||
|
theta_rescale_factor: Union[float, List[float]]=1.0,
|
||||||
|
interpolation_factor: Union[float, List[float]]=1.0,
|
||||||
|
concat_dict = {'mode': 'timecat-w', 'bias': -1}, num_frames: int = 129, k: int = 0,
|
||||||
|
):
|
||||||
|
|
||||||
|
grid = get_meshgrid_nd(start, *args, dim=len(rope_dim_list)) # [3, W, H, D] / [2, W, H]
|
||||||
|
if len(concat_dict)<1:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
if concat_dict['mode']=='timecat':
|
||||||
|
bias = grid[:,:1].clone()
|
||||||
|
bias[0] = concat_dict['bias']*torch.ones_like(bias[0])
|
||||||
|
grid = torch.cat([bias, grid], dim=1)
|
||||||
|
|
||||||
|
elif concat_dict['mode']=='timecat-w':
|
||||||
|
bias = grid[:,:1].clone()
|
||||||
|
bias[0] = concat_dict['bias']*torch.ones_like(bias[0])
|
||||||
|
bias[2] += start[-1] ## ref https://github.com/Yuanshi9815/OminiControl/blob/main/src/generate.py#L178
|
||||||
|
grid = torch.cat([bias, grid], dim=1)
|
||||||
|
if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):
|
||||||
|
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
|
||||||
|
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
|
||||||
|
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
|
||||||
|
assert len(theta_rescale_factor) == len(rope_dim_list), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
|
||||||
|
|
||||||
|
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
|
||||||
|
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
|
||||||
|
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
|
||||||
|
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
|
||||||
|
assert len(interpolation_factor) == len(rope_dim_list), "len(interpolation_factor) should equal to len(rope_dim_list)"
|
||||||
|
|
||||||
|
# use 1/ndim of dimensions to encode grid_axis
|
||||||
|
embs = []
|
||||||
|
for i in range(len(rope_dim_list)):
|
||||||
|
emb = get_1d_rotary_pos_embed(rope_dim_list[i], grid[i].reshape(-1), theta, use_real=use_real,
|
||||||
|
theta_rescale_factor=theta_rescale_factor[i],
|
||||||
|
interpolation_factor=interpolation_factor[i]) # 2 x [WHD, rope_dim_list[i]]
|
||||||
|
|
||||||
|
embs.append(emb)
|
||||||
|
|
||||||
|
if use_real:
|
||||||
|
cos = torch.cat([emb[0] for emb in embs], dim=1) # (WHD, D/2)
|
||||||
|
sin = torch.cat([emb[1] for emb in embs], dim=1) # (WHD, D/2)
|
||||||
|
return cos, sin
|
||||||
|
else:
|
||||||
|
emb = torch.cat(embs, dim=1) # (WHD, D/2)
|
||||||
|
return emb
|
||||||
@@ -2,7 +2,8 @@ import os
|
|||||||
import torch
|
import torch
|
||||||
import json
|
import json
|
||||||
import gc
|
import gc
|
||||||
from .utils import log, print_memory
|
from tqdm import tqdm
|
||||||
|
from .utils import log, print_memory, optimized_scale
|
||||||
from diffusers.video_processor import VideoProcessor
|
from diffusers.video_processor import VideoProcessor
|
||||||
from typing import List, Dict, Any, Tuple
|
from typing import List, Dict, Any, Tuple
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -220,6 +221,8 @@ class HyVideoTeaCache:
|
|||||||
"rel_l1_thresh": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.01,
|
"rel_l1_thresh": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.01,
|
||||||
"tooltip": "Higher values will make TeaCache more aggressive, faster, but may cause artifacts"}),
|
"tooltip": "Higher values will make TeaCache more aggressive, faster, but may cause artifacts"}),
|
||||||
"cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
|
"cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}),
|
||||||
|
"start_step": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1, "tooltip": "Start step to apply TeaCache"}),
|
||||||
|
"end_step": ("INT", {"default": -1, "min": -1, "max": 100, "step": 1, "tooltip": "End step to apply TeaCache"}),
|
||||||
|
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -229,14 +232,16 @@ class HyVideoTeaCache:
|
|||||||
CATEGORY = "HunyuanVideoWrapper"
|
CATEGORY = "HunyuanVideoWrapper"
|
||||||
DESCRIPTION = "TeaCache settings for HunyuanVideo to speed up inference"
|
DESCRIPTION = "TeaCache settings for HunyuanVideo to speed up inference"
|
||||||
|
|
||||||
def process(self, rel_l1_thresh, cache_device):
|
def process(self, rel_l1_thresh, cache_device, start_step, end_step):
|
||||||
if cache_device == "main_device":
|
if cache_device == "main_device":
|
||||||
teacache_device = mm.get_torch_device()
|
teacache_device = mm.get_torch_device()
|
||||||
else:
|
else:
|
||||||
teacache_device = mm.unet_offload_device()
|
teacache_device = mm.unet_offload_device()
|
||||||
teacache_args = {
|
teacache_args = {
|
||||||
"rel_l1_thresh": rel_l1_thresh,
|
"rel_l1_thresh": rel_l1_thresh,
|
||||||
"cache_device": teacache_device
|
"cache_device": teacache_device,
|
||||||
|
"start_step": start_step,
|
||||||
|
"end_step": end_step
|
||||||
}
|
}
|
||||||
return (teacache_args,)
|
return (teacache_args,)
|
||||||
|
|
||||||
@@ -276,7 +281,7 @@ class HyVideoModelLoader:
|
|||||||
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}),
|
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}),
|
||||||
|
|
||||||
"base_precision": (["fp32", "bf16"], {"default": "bf16"}),
|
"base_precision": (["fp32", "bf16"], {"default": "bf16"}),
|
||||||
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_scaled', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6", "torchao_int4", "torchao_int8"], {"default": 'disabled', "tooltip": "optional quantization method"}),
|
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_scaled'], {"default": 'disabled', "tooltip": "optional quantization method"}),
|
||||||
"load_device": (["main_device", "offload_device"], {"default": "main_device"}),
|
"load_device": (["main_device", "offload_device"], {"default": "main_device"}),
|
||||||
},
|
},
|
||||||
"optional": {
|
"optional": {
|
||||||
@@ -325,6 +330,8 @@ class HyVideoModelLoader:
|
|||||||
in_channels = sd["img_in.proj.weight"].shape[1]
|
in_channels = sd["img_in.proj.weight"].shape[1]
|
||||||
if in_channels == 16 and "i2v" in model.lower():
|
if in_channels == 16 and "i2v" in model.lower():
|
||||||
i2v_condition_type = "token_replace"
|
i2v_condition_type = "token_replace"
|
||||||
|
elif in_channels == 16 and "custom" in model.lower():
|
||||||
|
i2v_condition_type = "reference"
|
||||||
else:
|
else:
|
||||||
i2v_condition_type = "latent_concat"
|
i2v_condition_type = "latent_concat"
|
||||||
log.info(f"Condition type: {i2v_condition_type}")
|
log.info(f"Condition type: {i2v_condition_type}")
|
||||||
@@ -380,166 +387,98 @@ class HyVideoModelLoader:
|
|||||||
comfy_model=comfy_model,
|
comfy_model=comfy_model,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not "torchao" in quantization:
|
log.info("Using accelerate to load and assign model weights to device...")
|
||||||
log.info("Using accelerate to load and assign model weights to device...")
|
if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast" or quantization == "fp8_scaled":
|
||||||
if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast" or quantization == "fp8_scaled":
|
fp8_scale_map = {}
|
||||||
dtype = torch.float8_e4m3fn
|
if "fp8_scale" in sd:
|
||||||
elif quantization == "fp8_e5m2":
|
for k, v in sd.items():
|
||||||
dtype = torch.float8_e5m2
|
if k.endswith(".fp8_scale"):
|
||||||
else:
|
fp8_scale_map[k] = v
|
||||||
dtype = base_dtype
|
dtype = torch.float8_e4m3fn
|
||||||
params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"}
|
elif quantization == "fp8_e5m2":
|
||||||
for name, param in transformer.named_parameters():
|
dtype = torch.float8_e5m2
|
||||||
#print("Assigning Parameter name: ", name)
|
else:
|
||||||
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
|
dtype = base_dtype
|
||||||
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
|
params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"}
|
||||||
|
param_count = sum(1 for _ in transformer.named_parameters())
|
||||||
|
for name, param in tqdm(transformer.named_parameters(),
|
||||||
|
desc=f"Loading transformer parameters to {transformer_load_device}",
|
||||||
|
total=param_count,
|
||||||
|
leave=True):
|
||||||
|
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
|
||||||
|
set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
|
||||||
|
|
||||||
comfy_model.diffusion_model = transformer
|
comfy_model.diffusion_model = transformer
|
||||||
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
||||||
pipe.comfy_model = patcher
|
pipe.comfy_model = patcher
|
||||||
|
|
||||||
del sd
|
del sd
|
||||||
gc.collect()
|
gc.collect()
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
if lora is not None:
|
if lora is not None:
|
||||||
from comfy.sd import load_lora_for_models
|
from comfy.sd import load_lora_for_models
|
||||||
for l in lora:
|
for l in lora:
|
||||||
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
|
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
|
||||||
lora_path = l["path"]
|
lora_path = l["path"]
|
||||||
lora_strength = l["strength"]
|
lora_strength = l["strength"]
|
||||||
lora_sd = load_torch_file(lora_path, safe_load=True)
|
lora_sd = load_torch_file(lora_path, safe_load=True)
|
||||||
lora_sd = standardize_lora_key_format(lora_sd)
|
lora_sd = standardize_lora_key_format(lora_sd)
|
||||||
if l["blocks"]:
|
if l["blocks"]:
|
||||||
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"])
|
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"])
|
||||||
|
|
||||||
# patch in channels for keyframe LoRA
|
# patch in channels for keyframe LoRA
|
||||||
if "diffusion_model.img_in.proj.lora_A.weight" in lora_sd:
|
if "diffusion_model.img_in.proj.lora_A.weight" in lora_sd:
|
||||||
from .hyvideo.modules.embed_layers import PatchEmbed
|
from .hyvideo.modules.embed_layers import PatchEmbed
|
||||||
if lora_sd["diffusion_model.img_in.proj.lora_A.weight"].shape[1] != in_channels:
|
if lora_sd["diffusion_model.img_in.proj.lora_A.weight"].shape[1] != in_channels:
|
||||||
log.info(f"Different in_channels {lora_sd['diffusion_model.img_in.proj.lora_A.weight'].shape[1]} vs {in_channels}, patching...")
|
log.info(f"Different in_channels {lora_sd['diffusion_model.img_in.proj.lora_A.weight'].shape[1]} vs {in_channels}, patching...")
|
||||||
new_img_in = PatchEmbed(
|
new_img_in = PatchEmbed(
|
||||||
patch_size=patcher.model.diffusion_model.patch_size,
|
patch_size=patcher.model.diffusion_model.patch_size,
|
||||||
in_chans=32,
|
in_chans=32,
|
||||||
embed_dim=patcher.model.diffusion_model.hidden_size,
|
embed_dim=patcher.model.diffusion_model.hidden_size,
|
||||||
).to(patcher.model.diffusion_model.device, dtype=patcher.model.diffusion_model.dtype)
|
).to(patcher.model.diffusion_model.device, dtype=patcher.model.diffusion_model.dtype)
|
||||||
new_img_in.proj.weight.zero_()
|
new_img_in.proj.weight.zero_()
|
||||||
new_img_in.proj.weight[:, :in_channels].copy_(patcher.model.diffusion_model.img_in.proj.weight)
|
new_img_in.proj.weight[:, :in_channels].copy_(patcher.model.diffusion_model.img_in.proj.weight)
|
||||||
|
|
||||||
if patcher.model.diffusion_model.img_in.proj.bias is not None:
|
if patcher.model.diffusion_model.img_in.proj.bias is not None:
|
||||||
new_img_in.proj.bias.copy_(patcher.model.diffusion_model.img_in.proj.bias)
|
new_img_in.proj.bias.copy_(patcher.model.diffusion_model.img_in.proj.bias)
|
||||||
|
|
||||||
patcher.model.diffusion_model.img_in = new_img_in
|
patcher.model.diffusion_model.img_in = new_img_in
|
||||||
|
|
||||||
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
||||||
|
|
||||||
comfy.model_management.load_models_gpu([patcher])
|
comfy.model_management.load_models_gpu([patcher])
|
||||||
if load_device == "offload_device":
|
if load_device == "offload_device":
|
||||||
patcher.model.diffusion_model.to(offload_device)
|
patcher.model.diffusion_model.to(offload_device)
|
||||||
|
|
||||||
if quantization == "fp8_e4m3fn_fast":
|
if quantization == "fp8_e4m3fn_fast":
|
||||||
from .fp8_optimization import convert_fp8_linear
|
from .fp8_optimization import convert_fp8_linear
|
||||||
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep)
|
params_to_keep.update({"mlp", "modulation", "mod"})
|
||||||
elif quantization == "fp8_scaled":
|
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep)
|
||||||
from .hyvideo.modules.fp8_optimization import convert_fp8_linear
|
elif quantization == "fp8_scaled":
|
||||||
convert_fp8_linear(patcher.model.diffusion_model, base_dtype)
|
from .hyvideo.modules.fp8_optimization import convert_fp8_linear
|
||||||
|
convert_fp8_linear(patcher.model.diffusion_model, base_dtype, device, fp8_scale_map=fp8_scale_map)
|
||||||
|
|
||||||
if auto_cpu_offload:
|
if auto_cpu_offload:
|
||||||
transformer.enable_auto_offload(dtype=dtype, device=device)
|
if quantization == "fp8_scaled":
|
||||||
|
raise ValueError("Auto CPU offload and fp8 scaled quantization are not compatible.")
|
||||||
|
transformer.enable_auto_offload(dtype=dtype, device=device)
|
||||||
|
|
||||||
#compile
|
#compile
|
||||||
if compile_args is not None:
|
if compile_args is not None:
|
||||||
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
|
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
|
||||||
if compile_args["compile_single_blocks"]:
|
if compile_args["compile_single_blocks"]:
|
||||||
for i, block in enumerate(patcher.model.diffusion_model.single_blocks):
|
for i, block in enumerate(patcher.model.diffusion_model.single_blocks):
|
||||||
patcher.model.diffusion_model.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
|
||||||
if compile_args["compile_double_blocks"]:
|
|
||||||
for i, block in enumerate(patcher.model.diffusion_model.double_blocks):
|
|
||||||
patcher.model.diffusion_model.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
|
||||||
if compile_args["compile_txt_in"]:
|
|
||||||
patcher.model.diffusion_model.txt_in = torch.compile(patcher.model.diffusion_model.txt_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
|
||||||
if compile_args["compile_vector_in"]:
|
|
||||||
patcher.model.diffusion_model.vector_in = torch.compile(patcher.model.diffusion_model.vector_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
|
||||||
if compile_args["compile_final_layer"]:
|
|
||||||
patcher.model.diffusion_model.final_layer = torch.compile(patcher.model.diffusion_model.final_layer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
|
||||||
elif "torchao" in quantization:
|
|
||||||
try:
|
|
||||||
from torchao.quantization import (
|
|
||||||
quantize_,
|
|
||||||
fpx_weight_only,
|
|
||||||
float8_dynamic_activation_float8_weight,
|
|
||||||
int8_dynamic_activation_int8_weight,
|
|
||||||
int8_weight_only,
|
|
||||||
int4_weight_only
|
|
||||||
)
|
|
||||||
except:
|
|
||||||
raise ImportError("torchao is not installed")
|
|
||||||
|
|
||||||
# def filter_fn(module: nn.Module, fqn: str) -> bool:
|
|
||||||
# target_submodules = {'attn1', 'ff'} # avoid norm layers, 1.5 at least won't work with quantized norm1 #todo: test other models
|
|
||||||
# if any(sub in fqn for sub in target_submodules):
|
|
||||||
# return isinstance(module, nn.Linear)
|
|
||||||
# return False
|
|
||||||
|
|
||||||
if "fp6" in quantization:
|
|
||||||
quant_func = fpx_weight_only(3, 2)
|
|
||||||
elif "int4" in quantization:
|
|
||||||
quant_func = int4_weight_only()
|
|
||||||
elif "int8" in quantization:
|
|
||||||
quant_func = int8_weight_only()
|
|
||||||
elif "fp8dq" in quantization:
|
|
||||||
quant_func = float8_dynamic_activation_float8_weight()
|
|
||||||
elif 'fp8dqrow' in quantization:
|
|
||||||
from torchao.quantization.quant_api import PerRow
|
|
||||||
quant_func = float8_dynamic_activation_float8_weight(granularity=PerRow())
|
|
||||||
elif 'int8dq' in quantization:
|
|
||||||
quant_func = int8_dynamic_activation_int8_weight()
|
|
||||||
|
|
||||||
log.info(f"Quantizing model with {quant_func}")
|
|
||||||
comfy_model.diffusion_model = transformer
|
|
||||||
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
|
||||||
|
|
||||||
if lora is not None:
|
|
||||||
from comfy.sd import load_lora_for_models
|
|
||||||
for l in lora:
|
|
||||||
lora_path = l["path"]
|
|
||||||
lora_strength = l["strength"]
|
|
||||||
lora_sd = load_torch_file(lora_path, safe_load=True)
|
|
||||||
lora_sd = standardize_lora_key_format(lora_sd)
|
|
||||||
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
|
|
||||||
|
|
||||||
comfy.model_management.load_models_gpu([patcher])
|
|
||||||
|
|
||||||
for i, block in enumerate(patcher.model.diffusion_model.single_blocks):
|
|
||||||
log.info(f"Quantizing single_block {i}")
|
|
||||||
for name, _ in block.named_parameters(prefix=f"single_blocks.{i}"):
|
|
||||||
#print(f"Parameter name: {name}")
|
|
||||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name])
|
|
||||||
if compile_args is not None:
|
|
||||||
patcher.model.diffusion_model.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
patcher.model.diffusion_model.single_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||||
quantize_(block, quant_func)
|
if compile_args["compile_double_blocks"]:
|
||||||
print(block)
|
for i, block in enumerate(patcher.model.diffusion_model.double_blocks):
|
||||||
block.to(offload_device)
|
|
||||||
for i, block in enumerate(patcher.model.diffusion_model.double_blocks):
|
|
||||||
log.info(f"Quantizing double_block {i}")
|
|
||||||
for name, _ in block.named_parameters(prefix=f"double_blocks.{i}"):
|
|
||||||
#print(f"Parameter name: {name}")
|
|
||||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name])
|
|
||||||
if compile_args is not None:
|
|
||||||
patcher.model.diffusion_model.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
patcher.model.diffusion_model.double_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||||
quantize_(block, quant_func)
|
if compile_args["compile_txt_in"]:
|
||||||
for name, param in patcher.model.diffusion_model.named_parameters():
|
patcher.model.diffusion_model.txt_in = torch.compile(patcher.model.diffusion_model.txt_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||||
if "single_blocks" not in name and "double_blocks" not in name:
|
if compile_args["compile_vector_in"]:
|
||||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name])
|
patcher.model.diffusion_model.vector_in = torch.compile(patcher.model.diffusion_model.vector_in, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||||
|
if compile_args["compile_final_layer"]:
|
||||||
manual_offloading = False # to disable manual .to(device) calls
|
patcher.model.diffusion_model.final_layer = torch.compile(patcher.model.diffusion_model.final_layer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||||
log.info(f"Quantized transformer blocks to {quantization}")
|
|
||||||
for name, param in patcher.model.diffusion_model.named_parameters():
|
|
||||||
print(name, param.dtype)
|
|
||||||
#param.data = param.data.to(self.vae_dtype).to(device)
|
|
||||||
|
|
||||||
del sd
|
|
||||||
mm.soft_empty_cache()
|
|
||||||
|
|
||||||
patcher.model["pipe"] = pipe
|
patcher.model["pipe"] = pipe
|
||||||
patcher.model["dtype"] = base_dtype
|
patcher.model["dtype"] = base_dtype
|
||||||
@@ -661,10 +600,14 @@ class HyVideoTextEmbedBridge:
|
|||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {"required": {
|
return {"required": {
|
||||||
"positive": ("CONDITIONING", ),
|
"positive": ("CONDITIONING", ),
|
||||||
|
"cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "guidance scale"} ),
|
||||||
|
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ),
|
||||||
|
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ),
|
||||||
|
"batched_cfg": ("BOOLEAN", {"default": False, "tooltip": "Calculate cond and uncond as a batch, increases memory usage but can be faster"}),
|
||||||
|
"use_cfg_zero_star": ("BOOLEAN", {"default": True, "tooltip": "Use CFG zero star"}),
|
||||||
},
|
},
|
||||||
"optional": {
|
"optional": {
|
||||||
"negative": ("CONDITIONING", ),
|
"negative": ("CONDITIONING", ),
|
||||||
"hyvid_cfg": ("HYVID_CFG", {"tooltip": "The prompt from the cfg node is not used, only the settings"}),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
RETURN_TYPES = ("HYVIDEMBEDS",)
|
RETURN_TYPES = ("HYVIDEMBEDS",)
|
||||||
@@ -673,7 +616,7 @@ class HyVideoTextEmbedBridge:
|
|||||||
CATEGORY = "HunyuanVideoWrapper"
|
CATEGORY = "HunyuanVideoWrapper"
|
||||||
DESCRIPTION = "Acts as a bridge between the native ComfyUI conditioning and the HunyuanVideoWrapper embeds"
|
DESCRIPTION = "Acts as a bridge between the native ComfyUI conditioning and the HunyuanVideoWrapper embeds"
|
||||||
|
|
||||||
def convert(self, positive, negative=None, hyvid_cfg=None):
|
def convert(self, positive, cfg, start_percent, end_percent, batched_cfg, use_cfg_zero_star, negative=None):
|
||||||
positive_cond = positive[0][0]
|
positive_cond = positive[0][0]
|
||||||
positive_pooled = positive[0][1]["pooled_output"]
|
positive_pooled = positive[0][1]["pooled_output"]
|
||||||
positive_attention_mask = torch.ones(positive_cond.shape[1], dtype=torch.bool, device=positive_cond.device).unsqueeze(0)
|
positive_attention_mask = torch.ones(positive_cond.shape[1], dtype=torch.bool, device=positive_cond.device).unsqueeze(0)
|
||||||
@@ -689,10 +632,11 @@ class HyVideoTextEmbedBridge:
|
|||||||
"negative_attention_mask": negative_attention_mask,
|
"negative_attention_mask": negative_attention_mask,
|
||||||
"prompt_embeds_2": positive_pooled,
|
"prompt_embeds_2": positive_pooled,
|
||||||
"negative_prompt_embeds_2": negative_pooled,
|
"negative_prompt_embeds_2": negative_pooled,
|
||||||
"cfg": torch.tensor(hyvid_cfg["cfg"]) if hyvid_cfg is not None else None,
|
"cfg": torch.tensor(cfg),
|
||||||
"start_percent": torch.tensor(hyvid_cfg["start_percent"]) if hyvid_cfg is not None else None,
|
"start_percent": torch.tensor(start_percent),
|
||||||
"end_percent": torch.tensor(hyvid_cfg["end_percent"]) if hyvid_cfg is not None else None,
|
"end_percent": torch.tensor(end_percent),
|
||||||
"batched_cfg": torch.tensor(hyvid_cfg["batched_cfg"]) if hyvid_cfg is not None else None,
|
"batched_cfg": torch.tensor(batched_cfg),
|
||||||
|
"use_cfg_zero_star": torch.tensor(use_cfg_zero_star),
|
||||||
}
|
}
|
||||||
return (prompt_embeds_dict,)
|
return (prompt_embeds_dict,)
|
||||||
|
|
||||||
@@ -1139,7 +1083,8 @@ class HyVideoCFG:
|
|||||||
"cfg": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "guidance scale"} ),
|
"cfg": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "guidance scale"} ),
|
||||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ),
|
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ),
|
||||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ),
|
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the steps to apply CFG, rest of the steps use guidance_embeds"} ),
|
||||||
"batched_cfg": ("BOOLEAN", {"default": True, "tooltip": "Calculate cond and uncond as a batch, increases memory usage but can be faster"}),
|
"batched_cfg": ("BOOLEAN", {"default": False, "tooltip": "Calculate cond and uncond as a batch, increases memory usage but can be faster"}),
|
||||||
|
"use_cfg_zero_star": ("BOOLEAN", {"default": False, "tooltip": "Use CFG zero star"}),
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1149,13 +1094,14 @@ class HyVideoCFG:
|
|||||||
CATEGORY = "HunyuanVideoWrapper"
|
CATEGORY = "HunyuanVideoWrapper"
|
||||||
DESCRIPTION = "To use CFG with HunyuanVideo"
|
DESCRIPTION = "To use CFG with HunyuanVideo"
|
||||||
|
|
||||||
def process(self, negative_prompt, cfg, start_percent, end_percent, batched_cfg):
|
def process(self, negative_prompt, cfg, start_percent, end_percent, batched_cfg, use_cfg_zero_star):
|
||||||
cfg_dict = {
|
cfg_dict = {
|
||||||
"negative_prompt": negative_prompt,
|
"negative_prompt": negative_prompt,
|
||||||
"cfg": cfg,
|
"cfg": cfg,
|
||||||
"start_percent": start_percent,
|
"start_percent": start_percent,
|
||||||
"end_percent": end_percent,
|
"end_percent": end_percent,
|
||||||
"batched_cfg": batched_cfg
|
"batched_cfg": batched_cfg,
|
||||||
|
"use_cfg_zero_start": use_cfg_zero_star,
|
||||||
}
|
}
|
||||||
|
|
||||||
return (cfg_dict,)
|
return (cfg_dict,)
|
||||||
@@ -1234,6 +1180,7 @@ class HyVideoTextEmbedsLoad:
|
|||||||
"start_percent": loaded_tensors.get("start_percent", None),
|
"start_percent": loaded_tensors.get("start_percent", None),
|
||||||
"end_percent": loaded_tensors.get("end_percent", None),
|
"end_percent": loaded_tensors.get("end_percent", None),
|
||||||
"batched_cfg": loaded_tensors.get("batched_cfg", None),
|
"batched_cfg": loaded_tensors.get("batched_cfg", None),
|
||||||
|
"use_cfg_zero_star": loaded_tensors.get("use_cfg_zero_star", None),
|
||||||
}
|
}
|
||||||
|
|
||||||
return (prompt_embeds_dict,)
|
return (prompt_embeds_dict,)
|
||||||
@@ -1266,6 +1213,75 @@ class HyVideoContextOptions:
|
|||||||
}
|
}
|
||||||
|
|
||||||
return (context_options,)
|
return (context_options,)
|
||||||
|
|
||||||
|
class HyVideoLoopArgs:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"shift_skip": ("INT", {"default": 6, "min": 0, "tooltip": "Skip step of latent shift"}),
|
||||||
|
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the looping effect"}),
|
||||||
|
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the looping effect"}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("LOOPARGS", )
|
||||||
|
RETURN_NAMES = ("loop_args",)
|
||||||
|
FUNCTION = "process"
|
||||||
|
CATEGORY = "HunyuanVideoWrapper"
|
||||||
|
DESCRIPTION = "Looping through latent shift as shown in https://github.com/YisuiTT/Mobius/"
|
||||||
|
|
||||||
|
def process(self, **kwargs):
|
||||||
|
return (kwargs,)
|
||||||
|
|
||||||
|
class HunyuanVideoFresca:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"fresca_scale_low": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||||
|
"fresca_scale_high": ("FLOAT", {"default": 1.25, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||||
|
"fresca_freq_cutoff": ("INT", {"default": 20, "min": 0, "max": 10000, "step": 1}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("FRESCA_ARGS", )
|
||||||
|
RETURN_NAMES = ("fresca_args",)
|
||||||
|
FUNCTION = "process"
|
||||||
|
CATEGORY = "HunyuanVideoWrapper"
|
||||||
|
DESCRIPTION = "https://github.com/WikiChao/FreSca"
|
||||||
|
|
||||||
|
def process(self, **kwargs):
|
||||||
|
return (kwargs,)
|
||||||
|
|
||||||
|
class HunyuanVideoSLG:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"double_blocks": ("STRING", {"default": "", "tooltip": "Blocks to skip uncond on, separated by comma, index starts from 0"}),
|
||||||
|
"single_blocks": ("STRING", {"default": "20", "tooltip": "Blocks to skip uncond on, separated by comma, index starts from 0"}),
|
||||||
|
"start_percent": ("FLOAT", {"default": 0.4, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of SLG signal"}),
|
||||||
|
"end_percent": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of SLG signal"}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("SLGARGS", )
|
||||||
|
RETURN_NAMES = ("slg_args",)
|
||||||
|
FUNCTION = "process"
|
||||||
|
CATEGORY = "HunyuanVideoWrapper"
|
||||||
|
DESCRIPTION = "Skips uncond on the selected blocks"
|
||||||
|
|
||||||
|
def process(self, double_blocks, single_blocks, start_percent, end_percent):
|
||||||
|
|
||||||
|
slg_double_block_list = [int(x.strip()) for x in double_blocks.split(",")] if double_blocks else None
|
||||||
|
slg_single_block_list = [int(x.strip()) for x in single_blocks.split(",")] if single_blocks else None
|
||||||
|
|
||||||
|
slg_args = {
|
||||||
|
"double_blocks": slg_double_block_list,
|
||||||
|
"single_blocks": slg_single_block_list,
|
||||||
|
"start_percent": start_percent,
|
||||||
|
"end_percent": end_percent,
|
||||||
|
}
|
||||||
|
return (slg_args,)
|
||||||
|
|
||||||
#region Sampler
|
#region Sampler
|
||||||
class HyVideoSampler:
|
class HyVideoSampler:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -1287,6 +1303,7 @@ class HyVideoSampler:
|
|||||||
"optional": {
|
"optional": {
|
||||||
"samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ),
|
"samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ),
|
||||||
"image_cond_latents": ("LATENT", {"tooltip": "init Latents to use for image2video process"} ),
|
"image_cond_latents": ("LATENT", {"tooltip": "init Latents to use for image2video process"} ),
|
||||||
|
#"neg_image_cond_latents": ("LATENT", {"tooltip": "init Latents to use for image2video process"} ),
|
||||||
"denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
"denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||||
"stg_args": ("STGARGS", ),
|
"stg_args": ("STGARGS", ),
|
||||||
"context_options": ("HYVIDCONTEXT", ),
|
"context_options": ("HYVIDCONTEXT", ),
|
||||||
@@ -1298,6 +1315,10 @@ class HyVideoSampler:
|
|||||||
}),
|
}),
|
||||||
"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"}),
|
"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"}),
|
||||||
"i2v_mode": (["stability", "dynamic"], {"default": "dynamic", "tooltip": "I2V mode for image2video process"}),
|
"i2v_mode": (["stability", "dynamic"], {"default": "dynamic", "tooltip": "I2V mode for image2video process"}),
|
||||||
|
"loop_args": ("LOOPARGS", ),
|
||||||
|
"fresca_args": ("FRESCA_ARGS", ),
|
||||||
|
"slg_args": ("SLGARGS", ),
|
||||||
|
"mask": ("MASK", ),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1308,7 +1329,7 @@ class HyVideoSampler:
|
|||||||
|
|
||||||
def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, width, height, num_frames,
|
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,
|
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, i2v_mode="stability"):
|
teacache_args=None, scheduler=None, image_cond_latents=None, neg_image_cond_latents=None, riflex_freq_index=0, i2v_mode="stability", loop_args=None, fresca_args=None, slg_args=None, mask=None):
|
||||||
model = model.model
|
model = model.model
|
||||||
|
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
@@ -1331,11 +1352,13 @@ class HyVideoSampler:
|
|||||||
cfg_start_percent = float(hyvid_embeds.get("start_percent", 0.0))
|
cfg_start_percent = float(hyvid_embeds.get("start_percent", 0.0))
|
||||||
cfg_end_percent = float(hyvid_embeds.get("end_percent", 1.0))
|
cfg_end_percent = float(hyvid_embeds.get("end_percent", 1.0))
|
||||||
batched_cfg = hyvid_embeds.get("batched_cfg", True)
|
batched_cfg = hyvid_embeds.get("batched_cfg", True)
|
||||||
|
use_cfg_zero_star = hyvid_embeds.get("use_cfg_zero_star", True)
|
||||||
else:
|
else:
|
||||||
cfg = 1.0
|
cfg = 1.0
|
||||||
cfg_start_percent = 0.0
|
cfg_start_percent = 0.0
|
||||||
cfg_end_percent = 1.0
|
cfg_end_percent = 1.0
|
||||||
batched_cfg = False
|
batched_cfg = False
|
||||||
|
use_cfg_zero_star = False
|
||||||
|
|
||||||
if embedded_guidance_scale == 0.0:
|
if embedded_guidance_scale == 0.0:
|
||||||
embedded_guidance_scale = None
|
embedded_guidance_scale = None
|
||||||
@@ -1396,24 +1419,21 @@ class HyVideoSampler:
|
|||||||
|
|
||||||
# Initialize TeaCache if enabled
|
# Initialize TeaCache if enabled
|
||||||
if teacache_args is not None:
|
if teacache_args is not None:
|
||||||
# Check if dimensions have changed since last run
|
|
||||||
if (not hasattr(transformer, 'last_dimensions') or
|
|
||||||
transformer.last_dimensions != (height, width, num_frames) or
|
|
||||||
not hasattr(transformer, 'last_frame_count') or
|
|
||||||
transformer.last_frame_count != num_frames):
|
|
||||||
# Reset TeaCache state on dimension change
|
|
||||||
transformer.cnt = 0
|
|
||||||
transformer.teacache_skipped_steps = 0
|
|
||||||
transformer.accumulated_rel_l1_distance = 0
|
|
||||||
transformer.previous_modulated_input = None
|
|
||||||
transformer.previous_residual = None
|
|
||||||
transformer.last_dimensions = (height, width, num_frames)
|
|
||||||
transformer.last_frame_count = num_frames
|
|
||||||
transformer.teacache_device = device
|
|
||||||
|
|
||||||
transformer.enable_teacache = True
|
transformer.enable_teacache = True
|
||||||
|
transformer.cnt = 0
|
||||||
|
transformer.accumulated_rel_l1_distance = 0
|
||||||
|
transformer.teacache_skipped_steps_cond = transformer.teacache_skipped_steps_uncond =0
|
||||||
|
transformer.previous_modulated_input_cond = transformer.previous_modulated_input_uncond = None
|
||||||
|
transformer.previous_residual_cond = transformer.previous_residual_uncond = None
|
||||||
|
transformer.accumulated_rel_l1_distance_cond = transformer.accumulated_rel_l1_distance_uncond = 0
|
||||||
|
transformer.teacache_device = device
|
||||||
transformer.num_steps = steps
|
transformer.num_steps = steps
|
||||||
transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"]
|
transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"]
|
||||||
|
transformer.teacache_start_step = teacache_args["start_step"]
|
||||||
|
teacache_end_step = teacache_args["end_step"]
|
||||||
|
if teacache_end_step < 0:
|
||||||
|
teacache_end_step = steps - 1
|
||||||
|
transformer.teacache_end_step = teacache_end_step
|
||||||
else:
|
else:
|
||||||
transformer.enable_teacache = False
|
transformer.enable_teacache = False
|
||||||
|
|
||||||
@@ -1437,6 +1457,24 @@ class HyVideoSampler:
|
|||||||
if denoise_strength < 1.0:
|
if denoise_strength < 1.0:
|
||||||
input_latents *= VAE_SCALING_FACTOR
|
input_latents *= VAE_SCALING_FACTOR
|
||||||
|
|
||||||
|
mask_latents = None
|
||||||
|
if mask is not None:
|
||||||
|
from einops import rearrange
|
||||||
|
target_video_length = mask.shape[0]
|
||||||
|
target_height = mask.shape[1]
|
||||||
|
target_width = mask.shape[2]
|
||||||
|
|
||||||
|
mask_length = (target_video_length - 1) // 4 + 1
|
||||||
|
mask_height = target_height // 8
|
||||||
|
mask_width = target_width // 8
|
||||||
|
|
||||||
|
mask = mask.unsqueeze(-1).unsqueeze(0)
|
||||||
|
mask = rearrange(mask, "b t h w c -> b c t h w")
|
||||||
|
print("mask shape", mask.shape)
|
||||||
|
|
||||||
|
mask_latents = torch.nn.functional.interpolate(mask, size=(mask_length, mask_height, mask_width))
|
||||||
|
mask_latents = mask_latents.to(device)
|
||||||
|
|
||||||
out_latents = model["pipe"](
|
out_latents = model["pipe"](
|
||||||
num_inference_steps=steps,
|
num_inference_steps=steps,
|
||||||
height = target_height,
|
height = target_height,
|
||||||
@@ -1446,8 +1484,12 @@ class HyVideoSampler:
|
|||||||
cfg_start_percent=cfg_start_percent,
|
cfg_start_percent=cfg_start_percent,
|
||||||
cfg_end_percent=cfg_end_percent,
|
cfg_end_percent=cfg_end_percent,
|
||||||
batched_cfg=batched_cfg,
|
batched_cfg=batched_cfg,
|
||||||
|
use_cfg_zero_star=use_cfg_zero_star,
|
||||||
|
fresca_args=fresca_args,
|
||||||
|
slg_args=slg_args,
|
||||||
embedded_guidance_scale=embedded_guidance_scale,
|
embedded_guidance_scale=embedded_guidance_scale,
|
||||||
latents=input_latents,
|
latents=input_latents,
|
||||||
|
mask_latents=mask_latents,
|
||||||
denoise_strength=denoise_strength,
|
denoise_strength=denoise_strength,
|
||||||
prompt_embed_dict=hyvid_embeds,
|
prompt_embed_dict=hyvid_embeds,
|
||||||
generator=generator,
|
generator=generator,
|
||||||
@@ -1460,8 +1502,10 @@ class HyVideoSampler:
|
|||||||
feta_args=feta_args,
|
feta_args=feta_args,
|
||||||
leapfusion_img2vid = leapfusion_img2vid,
|
leapfusion_img2vid = leapfusion_img2vid,
|
||||||
image_cond_latents = image_cond_latents["samples"] * VAE_SCALING_FACTOR if image_cond_latents is not None else None,
|
image_cond_latents = image_cond_latents["samples"] * VAE_SCALING_FACTOR if image_cond_latents is not None else None,
|
||||||
|
neg_image_cond_latents = neg_image_cond_latents["samples"] * VAE_SCALING_FACTOR if neg_image_cond_latents is not None else None,
|
||||||
riflex_freq_index = riflex_freq_index,
|
riflex_freq_index = riflex_freq_index,
|
||||||
i2v_stability = i2v_stability,
|
i2v_stability = i2v_stability,
|
||||||
|
loop_args = loop_args,
|
||||||
)
|
)
|
||||||
|
|
||||||
print_memory(device)
|
print_memory(device)
|
||||||
@@ -1471,8 +1515,10 @@ class HyVideoSampler:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
if teacache_args is not None:
|
if teacache_args is not None:
|
||||||
log.info(f"TeaCache skipped {transformer.teacache_skipped_steps} steps")
|
|
||||||
transformer.teacache_skipped_steps = 0
|
log.info(f"TeaCache skipped {transformer.teacache_skipped_steps_cond} cond steps")
|
||||||
|
if transformer.teacache_skipped_steps_uncond > 0:
|
||||||
|
log.info(f"TeaCache skipped {transformer.teacache_skipped_steps_uncond} uncond steps")
|
||||||
|
|
||||||
if force_offload:
|
if force_offload:
|
||||||
if model["manual_offloading"]:
|
if model["manual_offloading"]:
|
||||||
@@ -1879,6 +1925,9 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"HyVideoI2VEncode": HyVideoI2VEncode,
|
"HyVideoI2VEncode": HyVideoI2VEncode,
|
||||||
"HyVideoEncodeKeyframes": HyVideoEncodeKeyframes,
|
"HyVideoEncodeKeyframes": HyVideoEncodeKeyframes,
|
||||||
"HyVideoTextEmbedBridge": HyVideoTextEmbedBridge,
|
"HyVideoTextEmbedBridge": HyVideoTextEmbedBridge,
|
||||||
|
"HyVideoLoopArgs": HyVideoLoopArgs,
|
||||||
|
"HunyuanVideoFresca": HunyuanVideoFresca,
|
||||||
|
"HunyuanVideoSLG": HunyuanVideoSLG
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"HyVideoSampler": "HunyuanVideo Sampler",
|
"HyVideoSampler": "HunyuanVideo Sampler",
|
||||||
@@ -1906,4 +1955,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
|||||||
"HyVideoI2VEncode": "HyVideo I2V Encode",
|
"HyVideoI2VEncode": "HyVideo I2V Encode",
|
||||||
"HyVideoEncodeKeyframes": "HyVideo Encode Keyframes",
|
"HyVideoEncodeKeyframes": "HyVideo Encode Keyframes",
|
||||||
"HyVideoTextEmbedBridge": "HyVideo TextEmbed Bridge",
|
"HyVideoTextEmbedBridge": "HyVideo TextEmbed Bridge",
|
||||||
|
"HyVideoLoopArgs": "HyVideo Loop Args",
|
||||||
|
"HunyuanVideoFresca": "HunyuanVideo Fresca",
|
||||||
|
"HunyuanVideoSLG": "HunyuanVideo SLG",
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -24,3 +24,69 @@ def print_memory(device):
|
|||||||
log.info(f"-------------------------------")
|
log.info(f"-------------------------------")
|
||||||
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
|
#memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False)
|
||||||
#log.info(f"Memory Summary:\n{memory_summary}")
|
#log.info(f"Memory Summary:\n{memory_summary}")
|
||||||
|
|
||||||
|
def optimized_scale(positive_flat, negative_flat):
|
||||||
|
|
||||||
|
# Calculate dot production
|
||||||
|
dot_product = torch.sum(positive_flat * negative_flat, dim=1, keepdim=True)
|
||||||
|
|
||||||
|
# Squared norm of uncondition
|
||||||
|
squared_norm = torch.sum(negative_flat ** 2, dim=1, keepdim=True) + 1e-8
|
||||||
|
|
||||||
|
# st_star = v_cond^T * v_uncond / ||v_uncond||^2
|
||||||
|
st_star = dot_product / squared_norm
|
||||||
|
|
||||||
|
return st_star
|
||||||
|
|
||||||
|
# Code based on https://github.com/WikiChao/FreSca (MIT License)
|
||||||
|
import torch
|
||||||
|
import torch.fft as fft
|
||||||
|
|
||||||
|
def fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
|
||||||
|
"""
|
||||||
|
Apply frequency-dependent scaling to an image tensor using Fourier transforms.
|
||||||
|
|
||||||
|
Parameters:
|
||||||
|
x: Input tensor of shape (B, C, H, W)
|
||||||
|
scale_low: Scaling factor for low-frequency components (default: 1.0)
|
||||||
|
scale_high: Scaling factor for high-frequency components (default: 1.5)
|
||||||
|
freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
x_filtered: Filtered version of x in spatial domain with frequency-specific scaling applied.
|
||||||
|
"""
|
||||||
|
# Preserve input dtype and device
|
||||||
|
dtype, device = x.dtype, x.device
|
||||||
|
|
||||||
|
# Convert to float32 for FFT computations
|
||||||
|
x = x.to(torch.float32)
|
||||||
|
|
||||||
|
# 1) Apply FFT and shift low frequencies to center
|
||||||
|
x_freq = fft.fftn(x, dim=(-2, -1))
|
||||||
|
x_freq = fft.fftshift(x_freq, dim=(-2, -1))
|
||||||
|
|
||||||
|
# 2) Create a mask to scale frequencies differently
|
||||||
|
B, C, T, H, W = x_freq.shape
|
||||||
|
crow, ccol = H // 2, W // 2
|
||||||
|
|
||||||
|
# Initialize mask with high-frequency scaling factor
|
||||||
|
mask = torch.ones((B, C, T, H, W), device=device) * scale_high
|
||||||
|
|
||||||
|
# Apply low-frequency scaling factor to center region
|
||||||
|
mask[
|
||||||
|
...,
|
||||||
|
crow - freq_cutoff : crow + freq_cutoff,
|
||||||
|
ccol - freq_cutoff : ccol + freq_cutoff,
|
||||||
|
] = scale_low
|
||||||
|
|
||||||
|
# 3) Apply frequency-specific scaling
|
||||||
|
x_freq = x_freq * mask
|
||||||
|
|
||||||
|
# 4) Convert back to spatial domain
|
||||||
|
x_freq = fft.ifftshift(x_freq, dim=(-2, -1))
|
||||||
|
x_filtered = fft.ifftn(x_freq, dim=(-2, -1)).real
|
||||||
|
|
||||||
|
# 5) Restore original dtype
|
||||||
|
x_filtered = x_filtered.to(dtype)
|
||||||
|
|
||||||
|
return x_filtered
|
||||||
Reference in New Issue
Block a user