From 6ea4d31b41cbd5cc23982278b7b1e357a17ce308 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 9 May 2025 09:16:02 +0300 Subject: [PATCH] initial hunyuan custom support --- .../pipelines/pipeline_hunyuan_video.py | 98 ++++-- hyvideo/modules/fp8_optimization.py | 43 +-- hyvideo/modules/models.py | 77 ++++- hyvideo/modules/posemb_layers.py | 49 +++ nodes.py | 296 ++++++++---------- utils.py | 15 +- 6 files changed, 361 insertions(+), 217 deletions(-) diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index f612ad1..65234e2 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -33,14 +33,15 @@ from diffusers.schedulers import DPMSolverMultistepScheduler from ...modules import HYVideoDiffusionTransformer from comfy.utils import ProgressBar - +import math +from ....utils import optimized_scale logger = logging.get_logger(__name__) # pylint: disable=invalid-name 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 -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 ndim = 5 - 2 rope_theta = 225 @@ -79,7 +80,7 @@ def get_rotary_pos_embed(transformer, latent_video_length, height, width, k=0): 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( + freqs_cos, freqs_sin = rope_func( rope_dim_list, rope_sizes, theta=rope_theta, @@ -254,6 +255,8 @@ class HunyuanVideoPipeline(DiffusionPipeline): ) if latents is not None: latents = latents.to(device) + else: + original_latents = None noise = randn_tensor(shape, generator=generator, device=device, dtype=self.base_dtype) if freenoise: @@ -318,7 +321,8 @@ class HunyuanVideoPipeline(DiffusionPipeline): elif frames_needed < current_frames: latents = latents[:, :, :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 print("latents shape:", latents.shape) @@ -338,7 +342,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): if hasattr(self.scheduler, "init_noise_sigma"): # scale the initial noise by the standard deviation required by the scheduler 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 def get_guidance_scale_embedding( @@ -423,6 +427,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): timesteps: List[int] = None, sigmas: List[float] = None, guidance_scale: float = 1.0, + use_cfg_zero_star: bool = False, cfg_start_percent: float = 0.0, cfg_end_percent: float = 1.0, batched_cfg: bool = True, @@ -431,6 +436,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): denoise_strength: float = 1.0, generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, latents: Optional[torch.Tensor] = None, + mask_latents: Optional[torch.Tensor] = None, cross_attention_kwargs: Optional[Dict[str, Any]] = None, guidance_rescale: float = 0.0, clip_skip: Optional[int] = None, @@ -452,6 +458,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): feta_args: Optional[Dict] = None, leapfusion_img2vid: Optional[bool] = False, image_cond_latents: Optional[torch.Tensor] = None, + neg_image_cond_latents: Optional[torch.Tensor] = None, riflex_freq_index: Optional[int] = None, i2v_stability=True, loop_args: Optional[Dict] = None, @@ -540,6 +547,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): # 2. Define call parameters batch_size = 1 + ref_latents = None device = self._execution_device prompt_embeds = prompt_embed_dict.get("prompt_embeds", None) @@ -634,14 +642,25 @@ class HunyuanVideoPipeline(DiffusionPipeline): 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 - ) + if i2v_condition_type == "reference": + 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: # rotary embeddings - freqs_cos, freqs_sin = get_rotary_pos_embed( - self.transformer, latent_video_length, height, width, k=riflex_freq_index - ) + if i2v_condition_type == "reference": + 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: freqs_cos = freqs_cos.to(self.base_dtype).to(device) freqs_sin = freqs_sin.to(self.base_dtype).to(device) @@ -652,10 +671,11 @@ class HunyuanVideoPipeline(DiffusionPipeline): if leapfusion_img2vid: logger.info("Single input latent frame detected, LeapFusion img2vid enabled") original_latents = latents + # 5. Prepare latent variables #num_channels_latents = self.transformer.config.in_channels 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, num_channels_latents, num_inference_steps, @@ -699,6 +719,14 @@ class HunyuanVideoPipeline(DiffusionPipeline): 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} ') logger.info(f"Sampling {video_length} frames in {latents.shape[2]} latents at {width}x{height} with {len(timesteps)} inference steps") @@ -712,6 +740,12 @@ class HunyuanVideoPipeline(DiffusionPipeline): if image_cond_latents is not None and i2v_condition_type == "token_replace": 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 input_prompt_embeds = prompt_embeds @@ -762,6 +796,16 @@ class HunyuanVideoPipeline(DiffusionPipeline): 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]) if leapfusion_img2vid: @@ -868,6 +912,8 @@ class HunyuanVideoPipeline(DiffusionPipeline): stg_block_idx=stg_block_idx, stg_mode=stg_mode, return_dict=True, + ref_latents=ref_latents, + is_uncond = False )["x"] else: uncond = self.transformer( @@ -878,10 +924,12 @@ class HunyuanVideoPipeline(DiffusionPipeline): text_states_2=input_prompt_embeds_2[0].unsqueeze(0), freqs_cos=freqs_cos, 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_mode=stg_mode, return_dict=True, + ref_latents=uncond_ref_latents, + is_uncond = True )["x"] cond = self.transformer( latent_model_input[1].unsqueeze(0), @@ -891,21 +939,28 @@ class HunyuanVideoPipeline(DiffusionPipeline): text_states_2=input_prompt_embeds_2[1].unsqueeze(0), freqs_cos=freqs_cos, 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_mode=stg_mode, return_dict=True, + ref_latents=ref_latents, + is_uncond = False )["x"] # perform guidance if cfg_enabled and not self.do_spatio_temporal_guidance: if batched_cfg: - noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) - noise_pred = noise_pred_uncond + self.guidance_scale * ( - noise_pred_text - noise_pred_uncond - ) + uncond, cond = noise_pred.chunk(2) + + #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: - noise_pred = uncond + self.guidance_scale * (cond - uncond) + alpha = 1.0 + noise_pred = uncond * alpha + self.guidance_scale * (cond - uncond * alpha) elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance: @@ -972,6 +1027,9 @@ class HunyuanVideoPipeline(DiffusionPipeline): else: 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 leapfusion_img2vid or i2v_condition_type == "latent_concat": latents = latents[:, :, 1:, :, :] diff --git a/hyvideo/modules/fp8_optimization.py b/hyvideo/modules/fp8_optimization.py index 86b7af7..a67e713 100644 --- a/hyvideo/modules/fp8_optimization.py +++ b/hyvideo/modules/fp8_optimization.py @@ -4,6 +4,7 @@ import torch.nn as nn from torch.nn import functional as F from comfy.utils import load_torch_file +@torch.compiler.disable() def get_fp_maxval(bits=8, mantissa_bit=3, sign_bits=1): _bits = torch.tensor(bits) _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) return maxval +@torch.compiler.disable() def quantize_to_fp8(x, bits=8, mantissa_bit=3, sign_bits=1): """ 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 return qdq_out, log_scales +@torch.compiler.disable() def fp8_tensor_quant(x, scale, bits=8, mantissa_bit=3, sign_bits=1): for i in range(len(x.shape) - 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) 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) - quant_dequant_x = qdq_out * scale.to(dtype) - return quant_dequant_x + return qdq_out def fp8_linear_forward(cls, original_dtype, input): 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) weight_dtype = linear_weight.dtype else: - scale = cls.fp8_scale.to(cls.weight.device) + scale = cls.fp8_scale#.to(cls.weight.device) linear_weight = cls.weight ##### - if weight_dtype == torch.float8_e4m3fn and cls.weight.sum() != 0: - if True or len(input.shape) == 3: - cls_dequant = fp8_activation_dequant(linear_weight, scale, original_dtype) - if cls.bias != None: - output = F.linear(input, cls_dequant, cls.bias) - else: - output = F.linear(input, cls_dequant) - return output + #if weight_dtype == torch.float8_e4m3fn and cls.weight.sum() != 0: + if weight_dtype == torch.float8_e4m3fn: + qdq_out = fp8_activation_dequant(linear_weight, original_dtype) + cls_dequant = qdq_out * scale + if cls.bias != None: + output = F.linear(input, cls_dequant, cls.bias) else: - return cls.original_forward(input.to(original_dtype)) + output = F.linear(input, cls_dequant) + return output else: 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) script_directory = os.path.dirname(os.path.abspath(__file__)) # loading fp8 mapping file - fp8_map_path = os.path.join(script_directory,"fp8_map.safetensors") - if os.path.exists(fp8_map_path): - fp8_map = load_torch_file(fp8_map_path, safe_load=True) - else: - raise ValueError(f"Invalid fp8_map path: {fp8_map_path}.") + if not fp8_scale_map: + fp8_map_path = os.path.join(script_directory,"fp8_map.safetensors") + if os.path.exists(fp8_map_path): + fp8_map = load_torch_file(fp8_map_path, safe_load=True) + else: + raise ValueError(f"Invalid fp8_map path: {fp8_map_path}.") #fp8_layers = [] for key, layer in module.named_modules(): @@ -96,6 +99,6 @@ def convert_fp8_linear(module, original_dtype): #fp8_layers.append(key) original_forward = layer.forward #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, "forward", lambda input, m=layer: fp8_linear_forward(m, original_dtype, input)) diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index 795b129..d1ee739 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -752,7 +752,8 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): self.enable_teacache = False self.cnt = 0 self.num_steps = 0 - self.teacache_skipped_steps = 0 + self.teacache_skipped_steps_cond = 0 + self.teacache_skipped_steps_uncond = 0 self.rel_l1_thresh = 0.15 self.accumulated_rel_l1_distance = 0 self.previous_modulated_input = None @@ -950,6 +951,8 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): stg_mode: str = None, stg_block_idx: int = -1, return_dict: bool = True, + ref_latents: torch.Tensor = None, + is_uncond = False, ) -> Union[torch.Tensor, Dict[str, torch.Tensor]]: def _process_double_blocks(img, txt, vec, block_args): @@ -1030,6 +1033,12 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): self.img_in.to(self.main_device) 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": txt = self.txt_in(txt) elif self.text_projection == "single_refiner": @@ -1088,29 +1097,62 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): normed_inp, shift=img_mod1_shift, scale=img_mod1_scale ) + # Use separate variables for conditional and unconditional passes + if not hasattr(self, 'previous_modulated_input_cond'): + self.previous_modulated_input_cond = None + self.previous_modulated_input_uncond = None + self.previous_residual_cond = None + self.previous_residual_uncond = None + self.accumulated_rel_l1_distance_cond = 0 + self.accumulated_rel_l1_distance_uncond = 0 + self.teacache_skipped_steps_cond = 0 + self.teacache_skipped_steps_uncond = 0 + + # 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: should_calc = True - self.accumulated_rel_l1_distance = 0 - self.previous_modulated_input = modulated_inp.clone() + accumulated_rel_l1_distance = 0 + previous_modulated_input = modulated_inp.clone() else: coefficients = [7.33226126e+02, -4.01131952e+02, 6.75869174e+01, -3.14987800e+00, 9.61237896e-02] 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 self.accumulated_rel_l1_distance < self.rel_l1_thresh: - should_calc = False + if previous_modulated_input is not None: + accumulated_rel_l1_distance += rescale_func(((modulated_inp-previous_modulated_input).abs().mean() / previous_modulated_input.abs().mean()).cpu().item()) + if accumulated_rel_l1_distance < self.rel_l1_thresh: + should_calc = False + else: + should_calc = True + accumulated_rel_l1_distance = 0 else: should_calc = True - self.accumulated_rel_l1_distance = 0 - self.previous_modulated_input = modulated_inp.clone() + accumulated_rel_l1_distance = 0 + + # 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 if self.cnt == self.num_steps: self.cnt = 0 - if not should_calc and self.previous_residual is not None: - self.teacache_skipped_steps += 1 + if not should_calc and previous_residual is not None: + # 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 - if img.shape == self.previous_residual.shape: - img = img + self.previous_residual.to(img.device) + if img.shape == previous_residual.shape: + img = img + previous_residual.to(img.device) else: should_calc = True # Force recalculation if dimensions don't match @@ -1123,7 +1165,13 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode, stg_block_idx) 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: # Pass through DiT blocks img, txt = _process_double_blocks(img, txt, vec, block_args) @@ -1132,6 +1180,9 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): x = _process_single_blocks(x, vec, txt.shape[1], block_args, stg_mode, stg_block_idx) img = x[:, :img_seq_len, ...] + if ref_latents is not None: + img = img[:, ref_length:] + # ---------------------------- Final layer ------------------------------ img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels) diff --git a/hyvideo/modules/posemb_layers.py b/hyvideo/modules/posemb_layers.py index 278e6b5..c3698d0 100644 --- a/hyvideo/modules/posemb_layers.py +++ b/hyvideo/modules/posemb_layers.py @@ -302,3 +302,52 @@ def get_1d_rotary_pos_embed_riflex( torch.ones_like(freqs), freqs ) # complex64 # [S, D/2] 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 \ No newline at end of file diff --git a/nodes.py b/nodes.py index 7683ecd..8e5fb2b 100644 --- a/nodes.py +++ b/nodes.py @@ -2,7 +2,8 @@ import os import torch import json 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 typing import List, Dict, Any, Tuple import numpy as np @@ -276,7 +277,7 @@ class HyVideoModelLoader: "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"}), - "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"}), }, "optional": { @@ -325,6 +326,8 @@ class HyVideoModelLoader: in_channels = sd["img_in.proj.weight"].shape[1] if in_channels == 16 and "i2v" in model.lower(): i2v_condition_type = "token_replace" + elif in_channels == 16 and not "i2v" in model.lower(): + i2v_condition_type = "reference" else: i2v_condition_type = "latent_concat" log.info(f"Condition type: {i2v_condition_type}") @@ -380,166 +383,97 @@ class HyVideoModelLoader: comfy_model=comfy_model, ) - if not "torchao" in quantization: - 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": - dtype = torch.float8_e4m3fn - elif quantization == "fp8_e5m2": - dtype = torch.float8_e5m2 - else: - dtype = base_dtype - params_to_keep = {"norm", "bias", "time_in", "vector_in", "guidance_in", "txt_in", "img_in"} - for name, param in transformer.named_parameters(): - #print("Assigning Parameter name: ", name) - 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]) + 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": + fp8_scale_map = {} + if "fp8_scale" in sd: + for k, v in sd.items(): + if k.endswith(".fp8_scale"): + fp8_scale_map[k] = v + dtype = torch.float8_e4m3fn + elif quantization == "fp8_e5m2": + dtype = torch.float8_e5m2 + else: + dtype = base_dtype + 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 - patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) - pipe.comfy_model = patcher + comfy_model.diffusion_model = transformer + patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) + pipe.comfy_model = patcher - del sd - gc.collect() - mm.soft_empty_cache() + del sd + gc.collect() + mm.soft_empty_cache() - if lora is not None: - from comfy.sd import load_lora_for_models - for l in lora: - log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}") - 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) - if l["blocks"]: - lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"]) - - # patch in channels for keyframe LoRA - if "diffusion_model.img_in.proj.lora_A.weight" in lora_sd: - from .hyvideo.modules.embed_layers import PatchEmbed - 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...") - new_img_in = PatchEmbed( - patch_size=patcher.model.diffusion_model.patch_size, - in_chans=32, - embed_dim=patcher.model.diffusion_model.hidden_size, - ).to(patcher.model.diffusion_model.device, dtype=patcher.model.diffusion_model.dtype) - new_img_in.proj.weight.zero_() - new_img_in.proj.weight[:, :in_channels].copy_(patcher.model.diffusion_model.img_in.proj.weight) + if lora is not None: + from comfy.sd import load_lora_for_models + for l in lora: + log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}") + 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) + if l["blocks"]: + lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"]) + + # patch in channels for keyframe LoRA + if "diffusion_model.img_in.proj.lora_A.weight" in lora_sd: + from .hyvideo.modules.embed_layers import PatchEmbed + 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...") + new_img_in = PatchEmbed( + patch_size=patcher.model.diffusion_model.patch_size, + in_chans=32, + embed_dim=patcher.model.diffusion_model.hidden_size, + ).to(patcher.model.diffusion_model.device, dtype=patcher.model.diffusion_model.dtype) + new_img_in.proj.weight.zero_() + 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: - new_img_in.proj.bias.copy_(patcher.model.diffusion_model.img_in.proj.bias) + 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) - 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]) - if load_device == "offload_device": - patcher.model.diffusion_model.to(offload_device) + comfy.model_management.load_models_gpu([patcher]) + if load_device == "offload_device": + patcher.model.diffusion_model.to(offload_device) - if quantization == "fp8_e4m3fn_fast": - from .fp8_optimization import convert_fp8_linear - convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep) - elif quantization == "fp8_scaled": - from .hyvideo.modules.fp8_optimization import convert_fp8_linear - convert_fp8_linear(patcher.model.diffusion_model, base_dtype) + if quantization == "fp8_e4m3fn_fast": + from .fp8_optimization import convert_fp8_linear + convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep) + elif quantization == "fp8_scaled": + 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: - transformer.enable_auto_offload(dtype=dtype, device=device) + if auto_cpu_offload: + 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 - if compile_args is not None: - torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"] - if compile_args["compile_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: + #compile + if compile_args is not None: + torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"] + if compile_args["compile_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"]) - quantize_(block, quant_func) - print(block) - 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: + 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"]) - quantize_(block, quant_func) - for name, param in patcher.model.diffusion_model.named_parameters(): - if "single_blocks" not in name and "double_blocks" not in name: - set_module_tensor_to_device(patcher.model.diffusion_model, name, device=patcher.model.diffusion_model_load_device, dtype=base_dtype, value=sd[name]) - - manual_offloading = False # to disable manual .to(device) calls - 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() + 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"]) patcher.model["pipe"] = pipe patcher.model["dtype"] = base_dtype @@ -661,10 +595,14 @@ class HyVideoTextEmbedBridge: def INPUT_TYPES(s): return {"required": { "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": { "negative": ("CONDITIONING", ), - "hyvid_cfg": ("HYVID_CFG", {"tooltip": "The prompt from the cfg node is not used, only the settings"}), } } RETURN_TYPES = ("HYVIDEMBEDS",) @@ -673,7 +611,7 @@ class HyVideoTextEmbedBridge: CATEGORY = "HunyuanVideoWrapper" 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_pooled = positive[0][1]["pooled_output"] positive_attention_mask = torch.ones(positive_cond.shape[1], dtype=torch.bool, device=positive_cond.device).unsqueeze(0) @@ -689,10 +627,11 @@ class HyVideoTextEmbedBridge: "negative_attention_mask": negative_attention_mask, "prompt_embeds_2": positive_pooled, "negative_prompt_embeds_2": negative_pooled, - "cfg": torch.tensor(hyvid_cfg["cfg"]) if hyvid_cfg is not None else None, - "start_percent": torch.tensor(hyvid_cfg["start_percent"]) if hyvid_cfg is not None else None, - "end_percent": torch.tensor(hyvid_cfg["end_percent"]) if hyvid_cfg is not None else None, - "batched_cfg": torch.tensor(hyvid_cfg["batched_cfg"]) if hyvid_cfg is not None else None, + "cfg": torch.tensor(cfg), + "start_percent": torch.tensor(start_percent), + "end_percent": torch.tensor(end_percent), + "batched_cfg": torch.tensor(batched_cfg), + "use_cfg_zero_star": torch.tensor(use_cfg_zero_star), } return (prompt_embeds_dict,) @@ -1139,7 +1078,8 @@ class HyVideoCFG: "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"} ), "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 +1089,14 @@ class HyVideoCFG: CATEGORY = "HunyuanVideoWrapper" 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 = { "negative_prompt": negative_prompt, "cfg": cfg, "start_percent": start_percent, "end_percent": end_percent, - "batched_cfg": batched_cfg + "batched_cfg": batched_cfg, + "use_cfg_zero_start": use_cfg_zero_star, } return (cfg_dict,) @@ -1234,6 +1175,7 @@ class HyVideoTextEmbedsLoad: "start_percent": loaded_tensors.get("start_percent", None), "end_percent": loaded_tensors.get("end_percent", 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,) @@ -1307,6 +1249,7 @@ class HyVideoSampler: "optional": { "samples": ("LATENT", {"tooltip": "init Latents to use for video2video 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}), "stg_args": ("STGARGS", ), "context_options": ("HYVIDCONTEXT", ), @@ -1319,6 +1262,7 @@ 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"}), "i2v_mode": (["stability", "dynamic"], {"default": "dynamic", "tooltip": "I2V mode for image2video process"}), "loop_args": ("LOOPARGS", ), + "mask": ("MASK", ), } } @@ -1329,7 +1273,7 @@ class HyVideoSampler: def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, width, height, num_frames, samples=None, denoise_strength=1.0, force_offload=True, stg_args=None, context_options=None, feta_args=None, - teacache_args=None, scheduler=None, image_cond_latents=None, riflex_freq_index=0, i2v_mode="stability", loop_args=None): + teacache_args=None, scheduler=None, image_cond_latents=None, neg_image_cond_latents=None, riflex_freq_index=0, i2v_mode="stability", loop_args=None, mask=None): model = model.model device = mm.get_torch_device() @@ -1352,11 +1296,13 @@ class HyVideoSampler: cfg_start_percent = float(hyvid_embeds.get("start_percent", 0.0)) cfg_end_percent = float(hyvid_embeds.get("end_percent", 1.0)) batched_cfg = hyvid_embeds.get("batched_cfg", True) + use_cfg_zero_star = hyvid_embeds.get("use_cfg_zero_star", True) else: cfg = 1.0 cfg_start_percent = 0.0 cfg_end_percent = 1.0 batched_cfg = False + use_cfg_zero_star = False if embedded_guidance_scale == 0.0: embedded_guidance_scale = None @@ -1424,7 +1370,8 @@ class HyVideoSampler: transformer.last_frame_count != num_frames): # Reset TeaCache state on dimension change transformer.cnt = 0 - transformer.teacache_skipped_steps = 0 + transformer.teacache_skipped_steps_cond = 0 + transformer.teacache_skipped_steps_uncond = 0 transformer.accumulated_rel_l1_distance = 0 transformer.previous_modulated_input = None transformer.previous_residual = None @@ -1458,6 +1405,24 @@ class HyVideoSampler: if denoise_strength < 1.0: 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"]( num_inference_steps=steps, height = target_height, @@ -1467,8 +1432,10 @@ class HyVideoSampler: cfg_start_percent=cfg_start_percent, cfg_end_percent=cfg_end_percent, batched_cfg=batched_cfg, + use_cfg_zero_star=use_cfg_zero_star, embedded_guidance_scale=embedded_guidance_scale, latents=input_latents, + mask_latents=mask_latents, denoise_strength=denoise_strength, prompt_embed_dict=hyvid_embeds, generator=generator, @@ -1481,6 +1448,7 @@ class HyVideoSampler: feta_args=feta_args, leapfusion_img2vid = leapfusion_img2vid, image_cond_latents = image_cond_latents["samples"] * VAE_SCALING_FACTOR if image_cond_latents is not None else None, + 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, i2v_stability = i2v_stability, loop_args = loop_args, @@ -1493,8 +1461,10 @@ class HyVideoSampler: pass 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 model["manual_offloading"]: diff --git a/utils.py b/utils.py index bce6e2b..4e0ae46 100644 --- a/utils.py +++ b/utils.py @@ -23,4 +23,17 @@ def print_memory(device): log.info(f"Max reserved memory: {max_reserved=:.3f} GB") log.info(f"-------------------------------") #memory_summary = torch.cuda.memory_summary(device=device, abbreviated=False) - #log.info(f"Memory Summary:\n{memory_summary}") \ No newline at end of file + #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 \ No newline at end of file