diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 407c39f..dc1dfec 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -390,28 +390,28 @@ class AnimateDiffModel(nn.Module): self.motion_embedding = FancyVideoCondEmbedding(in_channels=in_channels, cond_embed_dim=cond_embed_dim) self.motion_embedding.apply(initialize_weights_to_zero) - def get_fancyvideo_emb_patches(self, dtype, device, fps=16, motion_score=1.0): + def get_fancyvideo_emb_patches(self, dtype, device, fps=25, motion_score=3.0): patches = [] if self.fps_embedding is not None: if fps is not None: - def fps_emb_patch(x: Tensor, emb: Tensor, model_channels: int, transformer_options: dict[str]): + def fps_emb_patch(emb: Tensor, model_channels: int, transformer_options: dict[str]): nonlocal fps if fps is None: return emb fps = torch.tensor(fps).to(dtype=emb.dtype, device=emb.device) - fps = fps.expand(x.shape[0]) + fps = fps.expand(emb.shape[0]) fps_emb = timestep_embedding(fps, model_channels, repeat_only=False).to(dtype=emb.dtype) fps_emb = self.fps_embedding(fps_emb) return emb + fps_emb patches.append(fps_emb_patch) if self.motion_embedding is not None: if motion_score is not None: - def motion_emb_patch(x: Tensor, emb: Tensor, model_channels: int, transformer_options: dict[str]): + def motion_emb_patch(emb: Tensor, model_channels: int, transformer_options: dict[str]): nonlocal motion_score if motion_score is None: return emb motion_score = torch.tensor(motion_score).to(dtype=emb.dtype, device=emb.device) - motion_score = motion_score.expand(x.shape[0]) + motion_score = motion_score.expand(emb.shape[0]) motion_emb = timestep_embedding(motion_score, model_channels, repeat_only=False).to(dtype=emb.dtype) motion_emb = self.motion_embedding(motion_emb) return emb + motion_emb diff --git a/animatediff/nodes_sigma_schedule.py b/animatediff/nodes_sigma_schedule.py index e09e217..369c3f1 100644 --- a/animatediff/nodes_sigma_schedule.py +++ b/animatediff/nodes_sigma_schedule.py @@ -59,7 +59,12 @@ class RawSigmaScheduleNode: sampling: str, lcm_original_timesteps: int, zsnr: bool, lcm_zsnr: bool=None): if lcm_zsnr is not None: zsnr = lcm_zsnr - new_config = ModelSamplingConfig(beta_schedule=raw_beta_schedule, linear_start=linear_start, linear_end=linear_end) + # from pathlib import Path + # log_name = 'enforce_zero_terminal_snr_betas' + # betas_file = Path(__file__).parent.parent / rf"{log_name}.pt" + # given_betas = torch.load(betas_file, weights_only=True) + # given_betas[-1] = 0.0 + new_config = ModelSamplingConfig(beta_schedule=raw_beta_schedule, linear_start=linear_start, linear_end=linear_end)#, given_betas=given_betas) if sampling != ModelSamplingType.LCM: lcm_original_timesteps=None model_type = ModelSamplingType.from_alias(sampling) diff --git a/animatediff/utils_model.py b/animatediff/utils_model.py index 930cdc9..009aff8 100644 --- a/animatediff/utils_model.py +++ b/animatediff/utils_model.py @@ -75,13 +75,16 @@ def vae_decode_raw_batched(vae: VAE, latents: Tensor, per_batch=16, show_pbar=Fa class ModelSamplingConfig: - def __init__(self, beta_schedule: str, linear_start: float=None, linear_end: float=None): + def __init__(self, beta_schedule: str, linear_start: float=None, linear_end: float=None, given_betas: Tensor=None, timesteps: int=None): self.sampling_settings = {"beta_schedule": beta_schedule} if linear_start is not None: self.sampling_settings["linear_start"] = linear_start if linear_end is not None: self.sampling_settings["linear_end"] = linear_end - self.beta_schedule = beta_schedule # keeping this for backwards compatibility + if given_betas is not None: + self.sampling_settings["given_betas"] = given_betas + if timesteps is not None: + self.sampling_settings["timesteps"] = timesteps class ModelSamplingType: @@ -112,7 +115,7 @@ def factory_model_sampling_discrete_distilled(original_timesteps=50): # based on code in comfy_extras/nodes_model_advanced.py -def evolved_model_sampling(model_config: ModelSamplingConfig, model_type: ModelType, alias: str, original_timesteps: int=None): +def evolved_model_sampling(model_config: ModelSamplingConfig, model_type: ModelType, alias: str, original_timesteps: Union[int, None]=None): # if LCM, need to handle manually if BetaSchedules.is_lcm(alias) or original_timesteps is not None: sampling_type = comfy_extras.nodes_model_advanced.LCM @@ -129,7 +132,16 @@ def evolved_model_sampling(model_config: ModelSamplingConfig, model_type: ModelT # NOTE: if I want to support zsnr, this is where I would add that code return ModelSamplingAdvancedEvolved(model_config) # otherwise, use vanilla model_sampling function - return model_sampling(model_config, model_type) + ms = model_sampling(model_config, model_type) + if "given_betas" in model_config.sampling_settings: + beta_schedule = model_config.sampling_settings.get("beta_schedule", "linear") + linear_start = model_config.sampling_settings.get("linear_start", 0.00085) + linear_end = model_config.sampling_settings.get("linear_end", 0.012) + timesteps = model_config.sampling_settings.get("timesteps", 1000) + given_betas = model_config.sampling_settings.get("given_betas", None) + ms._register_schedule(given_betas=given_betas, beta_schedule=beta_schedule, + timesteps=timesteps, linear_start=linear_start, linear_end=linear_end) + return ms class BetaSchedules: