Add SLG
For testing, no clue what blocks to use at this point
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user