@@ -57,7 +57,7 @@ def __init__(
5757 img_size = img_size , device = device , schedule_name = schedule_name , latent = latent ,
5858 latent_channel = latent_channel , autoencoder = autoencoder )
5959
60- def sample (
60+ def _sample_loop (
6161 self ,
6262 model : nn .Module ,
6363 x : Optional [torch .Tensor ] = None ,
@@ -66,7 +66,7 @@ def sample(
6666 cfg_scale : Optional [float ] = None
6767 ) -> torch .Tensor :
6868 """
69- PLMS sample method
69+ PLMS sample loop method
7070 :param model: Model
7171 :param x: Input image tensor, if provided, will be used as the starting point for sampling
7272 :param n: Number of sample images, x priority is greater than n
@@ -75,62 +75,57 @@ def sample(
7575 Avoiding the posterior collapse problem, Reference paper: 'Classifier-Free Diffusion Guidance'
7676 :return: Sample images
7777 """
78- # Input dim: [n, img_channel, img_size_h, img_size_w]
79- x , n = self ._get_input_image (n = n , x = x )
80- logger .info (msg = f"PLMS Sampling { n } new images...." )
81- model .eval ()
82- with torch .no_grad ():
83- # Old eps
84- old_eps = []
85- # The list of current time and previous time
86- for i , p_i in tqdm (self .time_step ):
87- # Time step, creating a tensor of size n
88- t = (torch .ones (n ) * i ).long ().to (self .device )
89- # Previous time step, creating a tensor of size n
90- p_t = (torch .ones (n ) * p_i ).long ().to (self .device )
91- # Expand to a 4-dimensional tensor, and get the value according to the time step t
92- alpha_t = self .alpha_hat [t ][:, None , None , None ]
93- alpha_prev = self .alpha_hat [p_t ][:, None , None , None ]
94- noise = torch .randn_like (x ) if i > 1 else torch .zeros_like (x )
78+ # Old eps
79+ old_eps = []
80+ # The list of current time and previous time
81+ for i , p_i in tqdm (self .time_step ):
82+ # Time step, creating a tensor of size n
83+ t = (torch .ones (n ) * i ).long ().to (self .device )
84+ # Previous time step, creating a tensor of size n
85+ p_t = (torch .ones (n ) * p_i ).long ().to (self .device )
86+ # Expand to a 4-dimensional tensor, and get the value according to the time step t
87+ alpha_t = self .alpha_hat [t ][:, None , None , None ]
88+ alpha_prev = self .alpha_hat [p_t ][:, None , None , None ]
89+ noise = torch .randn_like (x ) if i > 1 else torch .zeros_like (x )
9590
96- # Predict noise
97- predicted_noise = self ._get_predicted_noise (model , x , t , labels , cfg_scale )
91+ # Predict noise
92+ predicted_noise = self ._get_predicted_noise (model , x , t , labels , cfg_scale )
9893
99- # Calculation formula
100- if len (old_eps ) == 0 :
101- # Pseudo Improved Euler (2nd order)
102- x0_t = torch .clamp ((x - (predicted_noise * torch .sqrt ((1 - alpha_t )))) / torch .sqrt (alpha_t ), - 1 , 1 )
103- c1 = self .eta * torch .sqrt ((1 - alpha_t / alpha_prev ) * (1 - alpha_prev ) / (1 - alpha_t ))
104- c2 = torch .sqrt ((1 - alpha_prev ) - c1 ** 2 )
105- p_x = torch .sqrt (alpha_prev ) * x0_t + c2 * predicted_noise + c1 * noise
106- if labels is None and cfg_scale is None :
107- # Images and time steps input into the model
108- predicted_noise_next = model (p_x , p_t )
109- else :
110- predicted_noise_next = model (p_x , p_t , labels )
111- predicted_noise_prime = (predicted_noise + predicted_noise_next ) / 2
112- elif len (old_eps ) == 1 :
113- # 2nd order Pseudo Linear Multistep (Adams-Bashforth)
114- predicted_noise_prime = (3 * predicted_noise - old_eps [- 1 ]) / 2
115- elif len (old_eps ) == 2 :
116- # 3rd order Pseudo Linear Multistep (Adams-Bashforth)
117- predicted_noise_prime = (23 * predicted_noise - 16 * old_eps [- 1 ] + 5 * old_eps [- 2 ]) / 12
118- elif len (old_eps ) >= 3 :
119- # 4th order Pseudo Linear Multistep (Adams-Bashforth)
120- predicted_noise_prime = (55 * predicted_noise - 59 * old_eps [- 1 ] + 37 * old_eps [- 2 ] -
121- 9 * old_eps [- 3 ]) / 24
122-
123- x0_t = torch .clamp ((x - (predicted_noise_prime * torch .sqrt ((1 - alpha_t )))) / torch .sqrt (alpha_t ),
124- - 1 , 1 )
94+ # Calculation formula
95+ if len (old_eps ) == 0 :
96+ # Pseudo Improved Euler (2nd order)
97+ x0_t = torch .clamp ((x - (predicted_noise * torch .sqrt ((1 - alpha_t )))) / torch .sqrt (alpha_t ), - 1 , 1 )
12598 c1 = self .eta * torch .sqrt ((1 - alpha_t / alpha_prev ) * (1 - alpha_prev ) / (1 - alpha_t ))
12699 c2 = torch .sqrt ((1 - alpha_prev ) - c1 ** 2 )
127- x = torch .sqrt (alpha_prev ) * x0_t + c2 * predicted_noise_prime + c1 * noise
128- # Save old predicted_noise
129- old_eps .append (predicted_noise )
130- # Only the last 3 historical values are retained to save memory
131- if len (old_eps ) > 3 :
132- old_eps .pop (0 )
133- # Post process
134- x = self .post_process (x = x )
135- model .train ()
100+ p_x = torch .sqrt (alpha_prev ) * x0_t + c2 * predicted_noise + c1 * noise
101+ if labels is None and cfg_scale is None :
102+ # Images and time steps input into the model
103+ predicted_noise_next = model (p_x , p_t )
104+ else :
105+ predicted_noise_next = model (p_x , p_t , labels )
106+ predicted_noise_prime = (predicted_noise + predicted_noise_next ) / 2
107+ elif len (old_eps ) == 1 :
108+ # 2nd order Pseudo Linear Multistep (Adams-Bashforth)
109+ predicted_noise_prime = (3 * predicted_noise - old_eps [- 1 ]) / 2
110+ elif len (old_eps ) == 2 :
111+ # 3rd order Pseudo Linear Multistep (Adams-Bashforth)
112+ predicted_noise_prime = (23 * predicted_noise - 16 * old_eps [- 1 ] + 5 * old_eps [- 2 ]) / 12
113+ elif len (old_eps ) >= 3 :
114+ # 4th order Pseudo Linear Multistep (Adams-Bashforth)
115+ predicted_noise_prime = (55 * predicted_noise - 59 * old_eps [- 1 ] + 37 * old_eps [- 2 ] -
116+ 9 * old_eps [- 3 ]) / 24
117+ else :
118+ raise ValueError (f"Unexpected number of old_eps: { len (old_eps )} " )
119+
120+ x0_t = torch .clamp ((x - (predicted_noise_prime * torch .sqrt ((1 - alpha_t )))) / torch .sqrt (alpha_t ),
121+ - 1 , 1 )
122+ c1 = self .eta * torch .sqrt ((1 - alpha_t / alpha_prev ) * (1 - alpha_prev ) / (1 - alpha_t ))
123+ c2 = torch .sqrt ((1 - alpha_prev ) - c1 ** 2 )
124+ x = torch .sqrt (alpha_prev ) * x0_t + c2 * predicted_noise_prime + c1 * noise
125+ # Save old predicted_noise
126+ old_eps .append (predicted_noise )
127+ # Only the last 3 historical values are retained to save memory
128+ if len (old_eps ) > 3 :
129+ old_eps .pop (0 )
130+
136131 return x
0 commit comments