Spaces:
Running on Zero
Running on Zero
| # 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) | |