Spaces:
Running on Zero
Running on Zero
File size: 3,238 Bytes
d3ea518 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 | # 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)
|