# 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)