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)