support using comfy native VAE decoding
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user