From 6682a0905e3d4bbeac8f9d3e8de28d5bb1389f0f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 9 May 2025 14:10:55 +0300 Subject: [PATCH] Add SLG For testing, no clue what blocks to use at this point --- .../pipelines/pipeline_hunyuan_video.py | 17 +++++++- hyvideo/modules/models.py | 16 ++++++++ nodes.py | 40 +++++++++++++++++-- 3 files changed, 68 insertions(+), 5 deletions(-) diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index d8e5b2b..c05627c 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -429,6 +429,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): 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_end_percent: float = 1.0, batched_cfg: bool = True, @@ -734,6 +735,15 @@ class HunyuanVideoPipeline(DiffusionPipeline): 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_blocks = None + 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)) @@ -922,6 +932,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): ref_latents=ref_latents, is_uncond = False, current_step = i, + current_step_percentage = current_step_percentage )["x"] else: uncond = self.transformer( @@ -938,7 +949,8 @@ class HunyuanVideoPipeline(DiffusionPipeline): return_dict=True, ref_latents=uncond_ref_latents, is_uncond = True, - current_step = i + current_step = i, + current_step_percentage = current_step_percentage )["x"] cond = self.transformer( latent_model_input[1].unsqueeze(0), @@ -954,7 +966,8 @@ class HunyuanVideoPipeline(DiffusionPipeline): return_dict=True, ref_latents=ref_latents, is_uncond = False, - current_step = i + current_step = i, + current_step_percentage = current_step_percentage )["x"] # perform guidance diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index c49fb31..903bd38 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -764,6 +764,11 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): self.last_frame_count = None self.teacache_device = None + 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! 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") @@ -956,10 +961,16 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): ref_latents: torch.Tensor = None, is_uncond = False, current_step: int = 0, + current_step_percentage: float = 0, ) -> Union[torch.Tensor, Dict[str, torch.Tensor]]: def _process_double_blocks(img, txt, vec, block_args): 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: block.to(self.main_device) @@ -971,6 +982,11 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): 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): + 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: block.to(self.main_device) diff --git a/nodes.py b/nodes.py index 8cc485c..8e52884 100644 --- a/nodes.py +++ b/nodes.py @@ -1251,6 +1251,36 @@ class HunyuanVideoFresca: 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": "10", "tooltip": "Blocks to skip uncond on, separated by comma, index starts from 0"}), + "start_percent": ("FLOAT", {"default": 0.1, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of SLG signal"}), + "end_percent": ("FLOAT", {"default": 1.0, "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 class HyVideoSampler: @@ -1287,6 +1317,7 @@ class HyVideoSampler: "i2v_mode": (["stability", "dynamic"], {"default": "dynamic", "tooltip": "I2V mode for image2video process"}), "loop_args": ("LOOPARGS", ), "fresca_args": ("FRESCA_ARGS", ), + "slg_args": ("SLGARGS", ), "mask": ("MASK", ), } } @@ -1298,7 +1329,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, neg_image_cond_latents=None, riflex_freq_index=0, i2v_mode="stability", loop_args=None, fresca_args=None, mask=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, fresca_args=None, slg_args=None, mask=None): model = model.model device = mm.get_torch_device() @@ -1466,6 +1497,7 @@ class HyVideoSampler: 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, latents=input_latents, mask_latents=mask_latents, @@ -1905,7 +1937,8 @@ NODE_CLASS_MAPPINGS = { "HyVideoEncodeKeyframes": HyVideoEncodeKeyframes, "HyVideoTextEmbedBridge": HyVideoTextEmbedBridge, "HyVideoLoopArgs": HyVideoLoopArgs, - "HunyuanVideoFresca": HunyuanVideoFresca + "HunyuanVideoFresca": HunyuanVideoFresca, + "HunyuanVideoSLG": HunyuanVideoSLG } NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoSampler": "HunyuanVideo Sampler", @@ -1934,5 +1967,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "HyVideoEncodeKeyframes": "HyVideo Encode Keyframes", "HyVideoTextEmbedBridge": "HyVideo TextEmbed Bridge", "HyVideoLoopArgs": "HyVideo Loop Args", - "HunyuanVideoFresca": "HunyuanVideo Fresca" + "HunyuanVideoFresca": "HunyuanVideo Fresca", + "HunyuanVideoSLG": "HunyuanVideo SLG", }