Updated emb_patch, some upgrades for ModelSamplingConfig
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user