homa / hyavatar /diffusion /__init__.py
multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
d3ea518 verified
Raw
History Blame Contribute Delete
3.24 kB
# Inference-only diffusion pipeline loaders.
from diffusers.schedulers import DDPMScheduler
from .pipelines import PipelineAllInOne, PipelineAllInOneLong
from .schedulers import FlowMatchDiscreteScheduler
def _build_scheduler(args, data_type='video'):
""" Build the denoising scheduler for inference. """
if args.denoise_type in ("ddpm", "diffusers_ddpm"):
rescale_betas_zero_snr = False
if args.enforce_zero_terminal_snr:
if args.predict_type == "v_prediction":
rescale_betas_zero_snr = True
else:
raise ValueError("We only support enforcing terminal SNR to 0 when predict_type==v_prediction.")
return DDPMScheduler(
beta_start=args.ddpm_beta_start,
beta_end=args.ddpm_beta_end,
beta_schedule=args.ddpm_noise_schedule,
variance_type='learned_range' if args.ddpm_learn_sigma else 'fixed_small',
prediction_type=args.ddpm_predict_type,
steps_offset=1,
clip_sample=False,
rescale_betas_zero_snr=rescale_betas_zero_snr,
)
elif args.denoise_type == "flow":
return FlowMatchDiscreteScheduler(
shift=args.flow_shift_eval_video if data_type == 'video' else args.flow_shift_eval,
reverse=args.flow_reverse,
solver=args.flow_solver,
use_flux_shift=args.use_flux_shift,
flux_base_shift=args.flux_base_shift,
flux_max_shift=args.flux_max_shift,
flux_base_token=args.flux_base_token,
flux_max_token=args.flux_max_token,
flux_shift_factor=args.flux_shift_factor,
)
raise ValueError(f"Invalid denoise type {args.denoise_type}")
def _load_pipeline(pipeline_cls, args, rank, vae, text_encoder, text_encoder_2, model,
scheduler=None, device=None, progress_bar_config=None, data_type='video'):
if scheduler is None:
scheduler = _build_scheduler(args, data_type=data_type)
# Only enable progress bar for rank 0
progress_bar_config = progress_bar_config or {'leave': True, 'disable': rank != 0}
pipeline = pipeline_cls(
vae=vae,
text_encoder=text_encoder,
text_encoder_2=text_encoder_2,
unet=model,
scheduler=scheduler,
safety_checker=None,
feature_extractor=None,
requires_safety_checker=False,
progress_bar_config=progress_bar_config,
args=args,
)
return pipeline.to(device)
def load_diffusion_pipeline_all_in_one(
args, rank, vae, text_encoder, text_encoder_2, model, scheduler=None,
device=None, progress_bar_config=None, data_type='video'
):
return _load_pipeline(PipelineAllInOne, args, rank, vae, text_encoder, text_encoder_2,
model, scheduler, device, progress_bar_config, data_type)
def load_diffusion_pipeline_all_in_one_long(
args, rank, vae, text_encoder, text_encoder_2, model, scheduler=None,
device=None, progress_bar_config=None, data_type='video'
):
return _load_pipeline(PipelineAllInOneLong, args, rank, vae, text_encoder, text_encoder_2,
model, scheduler, device, progress_bar_config, data_type)