support using comfy native VAE decoding

This commit is contained in:
kijai
2025-01-06 14:58:05 +02:00
parent e3a2fa7ca7
commit e8e0ac4e56
3 changed files with 25 additions and 13 deletions
+7 -7
View File
@@ -48,6 +48,8 @@ import comfy.latent_formats
script_directory = os.path.dirname(os.path.abspath(__file__))
VAE_SCALING_FACTOR = 0.476986
def filter_state_dict_by_blocks(state_dict, blocks_mapping):
filtered_dict = {}
@@ -1164,11 +1166,10 @@ class HyVideoSampler:
model["scheduler_config"]["flow_shift"] = flow_shift
model["scheduler_config"]["algorithm_type"] = "sde-dpmsolver++"
#model["scheduler_config"]["use_beta_sigmas"] = True
#model["scheduler_config"]["use_beta_flow_sigmas"] = True
noise_scheduler = scheduler_mapping[scheduler].from_config(model["scheduler_config"])
model["pipe"].scheduler = noise_scheduler
#model["pipe"].scheduler.flow_shift = flow_shift
if model["block_swap_args"] is not None:
for name, param in transformer.named_parameters():
@@ -1229,7 +1230,7 @@ class HyVideoSampler:
cfg_start_percent=cfg_start_percent,
cfg_end_percent=cfg_end_percent,
embedded_guidance_scale=embedded_guidance_scale,
latents=samples["samples"] if samples is not None else None,
latents=samples["samples"] * VAE_SCALING_FACTOR if samples is not None else None,
denoise_strength=denoise_strength,
prompt_embed_dict=hyvid_embeds,
generator=generator,
@@ -1255,7 +1256,7 @@ class HyVideoSampler:
gc.collect()
return ({
"samples": out_latents
"samples": out_latents.cpu() / VAE_SCALING_FACTOR
},)
#region VideoDecode
@@ -1310,8 +1311,7 @@ class HyVideoDecode:
raise ValueError(
f"Only support latents with shape (b, c, h, w) or (b, c, f, h, w), but got {latents.shape}."
)
latents = latents / vae.config.scaling_factor
#latents = latents / vae.config.scaling_factor
latents = latents.to(vae.dtype).to(device)
if enable_vae_tiling:
@@ -1385,7 +1385,7 @@ class HyVideoEncode:
if enable_vae_tiling:
vae.enable_tiling()
latents = vae.encode(image).latent_dist.sample(generator)
latents = latents * vae.config.scaling_factor
#latents = latents * vae.config.scaling_factor
vae.to(offload_device)
print("encoded latents shape",latents.shape)
+8 -6
View File
@@ -11,6 +11,8 @@ from .enhance_a_video.globals import enable_enhance, disable_enhance, set_enhanc
script_directory = os.path.dirname(os.path.abspath(__file__))
VAE_SCALING_FACTOR = 0.476986
def generate_eta_values(
timesteps,
start_step,
@@ -97,7 +99,7 @@ class HyVideoInverseSampler:
generator = torch.Generator(device=torch.device("cpu")).manual_seed(seed)
latents = samples["samples"] if samples is not None else None
latents = samples["samples"] * VAE_SCALING_FACTOR if samples is not None else None
batch_size, num_channels_latents, latent_num_frames, latent_height, latent_width = latents.shape
height = latent_height * pipeline.vae_scale_factor
width = latent_width * pipeline.vae_scale_factor
@@ -270,7 +272,7 @@ class HyVideoInverseSampler:
gc.collect()
return ({
"samples": latents
"samples": latents / VAE_SCALING_FACTOR
},)
class HyVideoReSampler:
@@ -312,7 +314,7 @@ class HyVideoReSampler:
transformer = model["pipe"].transformer
pipeline = model["pipe"]
target_latents = samples["samples"]
target_latents = samples["samples"] * VAE_SCALING_FACTOR
batch_size, num_channels_latents, latent_num_frames, latent_height, latent_width = target_latents.shape
height = latent_height * pipeline.vae_scale_factor
@@ -370,7 +372,7 @@ class HyVideoReSampler:
target_latents = target_latents.to(device)
latents = inversed_latents["samples"]
latents = inversed_latents["samples"] * VAE_SCALING_FACTOR
# 7. Denoising loop
self._num_timesteps = len(timesteps)
@@ -475,7 +477,7 @@ class HyVideoReSampler:
gc.collect()
return ({
"samples": latents
"samples": latents / VAE_SCALING_FACTOR
},)
class HyVideoPromptMixSampler:
@@ -712,7 +714,7 @@ class HyVideoPromptMixSampler:
gc.collect()
return ({
"samples": latents
"samples": latents / VAE_SCALING_FACTOR
},)
NODE_CLASS_MAPPINGS = {
+10
View File
@@ -219,6 +219,7 @@ class DPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin):
use_beta_sigmas: Optional[bool] = False,
use_lu_lambdas: Optional[bool] = False,
use_flow_sigmas: Optional[bool] = False,
use_beta_flow_sigmas: Optional[bool] = False,
flow_shift: Optional[float] = 1.0,
final_sigmas_type: Optional[str] = "zero", # "zero", "sigma_min"
lambda_min_clipped: float = -float("inf"),
@@ -409,6 +410,15 @@ class DPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin):
sigmas = np.flip(sigmas).copy()
sigmas = self._convert_to_beta(in_sigmas=sigmas, num_inference_steps=num_inference_steps)
timesteps = np.array([self._sigma_to_t(sigma, log_sigmas) for sigma in sigmas])
elif self.config.use_beta_flow_sigmas:
alphas = np.linspace(1, 1 / self.config.num_train_timesteps, num_inference_steps + 1)
flow_sigmas = 1.0 - alphas
flow_sigmas = np.flip(self.config.flow_shift * flow_sigmas /
(1 + (self.config.flow_shift - 1) * flow_sigmas))[:-1]
sigmas = self._convert_to_beta(in_sigmas=flow_sigmas,
num_inference_steps=num_inference_steps)
#timesteps = np.array([self._sigma_to_t(sigma, log_sigmas) for sigma in sigmas])
timesteps = (sigmas * self.config.num_train_timesteps).copy()
elif self.config.use_flow_sigmas:
alphas = np.linspace(1, 1 / self.config.num_train_timesteps, num_inference_steps + 1)
sigmas = 1.0 - alphas