Updated emb_patch, some upgrades for ModelSamplingConfig

This commit is contained in:
Jedrzej Kosinski
2024-09-07 03:06:53 -05:00
parent 1a9d70ee4c
commit a77460cfaf
3 changed files with 27 additions and 10 deletions
+5 -5
View File
@@ -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
+6 -1
View File
@@ -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)
+16 -4
View File
@@ -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: