add option to do cfg sequentially
This commit is contained in:
@@ -402,6 +402,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
guidance_scale: float = 1.0,
|
||||
cfg_start_percent: float = 0.0,
|
||||
cfg_end_percent: float = 1.0,
|
||||
batched_cfg: bool = True,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
eta: float = 0.0,
|
||||
denoise_strength: float = 1.0,
|
||||
@@ -664,8 +665,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
|
||||
#print(self.scheduler.sigmas)
|
||||
|
||||
|
||||
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))
|
||||
with self.progress_bar(total=len(timesteps)) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
@@ -771,6 +772,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
with torch.autocast(
|
||||
device_type="cuda", dtype=self.base_dtype, enabled=True
|
||||
):
|
||||
if batched_cfg or not cfg_enabled:
|
||||
noise_pred = self.transformer( # For an input image (129, 192, 336) (1, 256, 256)
|
||||
latent_model_input, # [2, 16, 33, 24, 42]
|
||||
t_expand, # [2]
|
||||
@@ -784,13 +786,45 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
stg_mode=stg_mode,
|
||||
return_dict=True,
|
||||
)["x"]
|
||||
else:
|
||||
uncond = self.transformer(
|
||||
latent_model_input[0].unsqueeze(0),
|
||||
t_expand[0].unsqueeze(0),
|
||||
text_states=input_prompt_embeds[0].unsqueeze(0),
|
||||
text_mask=input_prompt_mask[0].unsqueeze(0),
|
||||
text_states_2=input_prompt_embeds_2[0].unsqueeze(0),
|
||||
freqs_cos=freqs_cos,
|
||||
freqs_sin=freqs_sin,
|
||||
guidance=guidance_expand[0].unsqueeze(0),
|
||||
stg_block_idx=stg_block_idx,
|
||||
stg_mode=stg_mode,
|
||||
return_dict=True,
|
||||
)["x"]
|
||||
cond = self.transformer(
|
||||
latent_model_input[1].unsqueeze(0),
|
||||
t_expand[1].unsqueeze(0),
|
||||
text_states=input_prompt_embeds[1].unsqueeze(0),
|
||||
text_mask=input_prompt_mask[1].unsqueeze(0),
|
||||
text_states_2=input_prompt_embeds_2[1].unsqueeze(0),
|
||||
freqs_cos=freqs_cos,
|
||||
freqs_sin=freqs_sin,
|
||||
guidance=guidance_expand[1].unsqueeze(0),
|
||||
stg_block_idx=stg_block_idx,
|
||||
stg_mode=stg_mode,
|
||||
return_dict=True,
|
||||
)["x"]
|
||||
|
||||
# perform guidance
|
||||
if cfg_enabled and not self.do_spatio_temporal_guidance:
|
||||
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
|
||||
)
|
||||
else:
|
||||
noise_pred = uncond + self.guidance_scale * (cond - uncond)
|
||||
|
||||
|
||||
elif self.do_classifier_free_guidance and self.do_spatio_temporal_guidance:
|
||||
raise NotImplementedError
|
||||
noise_pred_uncond, noise_pred_text, noise_pred_perturb = noise_pred.chunk(3)
|
||||
|
||||
@@ -965,6 +965,7 @@ class HyVideoTextEncode:
|
||||
"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,
|
||||
}
|
||||
return (prompt_embeds_dict,)
|
||||
|
||||
@@ -1003,6 +1004,7 @@ 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"}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -1012,12 +1014,13 @@ class HyVideoCFG:
|
||||
CATEGORY = "HunyuanVideoWrapper"
|
||||
DESCRIPTION = "To use CFG with HunyuanVideo"
|
||||
|
||||
def process(self, negative_prompt, cfg, start_percent, end_percent):
|
||||
def process(self, negative_prompt, cfg, start_percent, end_percent, batched_cfg):
|
||||
cfg_dict = {
|
||||
"negative_prompt": negative_prompt,
|
||||
"cfg": cfg,
|
||||
"start_percent": start_percent,
|
||||
"end_percent": end_percent,
|
||||
"batched_cfg": batched_cfg
|
||||
}
|
||||
|
||||
return (cfg_dict,)
|
||||
@@ -1095,6 +1098,7 @@ class HyVideoTextEmbedsLoad:
|
||||
"cfg": loaded_tensors.get("cfg", None),
|
||||
"start_percent": loaded_tensors.get("start_percent", None),
|
||||
"end_percent": loaded_tensors.get("end_percent", None),
|
||||
"batched_cfg": loaded_tensors.get("batched_cfg", None),
|
||||
}
|
||||
|
||||
return (prompt_embeds_dict,)
|
||||
@@ -1185,10 +1189,12 @@ class HyVideoSampler:
|
||||
cfg = float(hyvid_embeds.get("cfg", 1.0))
|
||||
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)
|
||||
else:
|
||||
cfg = 1.0
|
||||
cfg_start_percent = 0.0
|
||||
cfg_end_percent = 1.0
|
||||
batched_cfg = False
|
||||
|
||||
if embedded_guidance_scale == 0.0:
|
||||
embedded_guidance_scale = None
|
||||
@@ -1291,6 +1297,7 @@ class HyVideoSampler:
|
||||
guidance_scale=cfg,
|
||||
cfg_start_percent=cfg_start_percent,
|
||||
cfg_end_percent=cfg_end_percent,
|
||||
batched_cfg=batched_cfg,
|
||||
embedded_guidance_scale=embedded_guidance_scale,
|
||||
latents=input_latents,
|
||||
denoise_strength=denoise_strength,
|
||||
|
||||
Reference in New Issue
Block a user