1111
1212from __future__ import annotations
1313
14+ import inspect
1415import math
1516import warnings
1617from abc import ABC , abstractmethod
@@ -862,6 +863,96 @@ def __init__(self, scheduler: Scheduler) -> None: # type: ignore[override]
862863
863864 self .scheduler = scheduler
864865
866+ @staticmethod
867+ def _scheduler_step_supports_kwarg (scheduler : Scheduler , kwarg : str ) -> bool :
868+ try :
869+ return kwarg in inspect .signature (scheduler .step ).parameters
870+ except (TypeError , ValueError ):
871+ return False
872+
873+ @staticmethod
874+ def _get_previous_sample_from_step_output (step_output : Any ) -> torch .Tensor :
875+ if isinstance (step_output , tuple ):
876+ return step_output [0 ]
877+ if isinstance (step_output , Mapping ):
878+ return step_output ["prev_sample" ]
879+ if hasattr (step_output , "prev_sample" ):
880+ return step_output .prev_sample
881+ raise TypeError ("Unsupported scheduler.step output. Expected a tuple or an object with `prev_sample`." )
882+
883+ @staticmethod
884+ def _get_scheduler_name (scheduler : Scheduler ) -> str :
885+ if hasattr (scheduler , "_get_name" ):
886+ return scheduler ._get_name ()
887+ return scheduler .__class__ .__name__
888+
889+ @staticmethod
890+ def _get_scheduler_config_value (scheduler : Scheduler , name : str , default : Any = None ) -> Any :
891+ config = getattr (scheduler , "config" , None )
892+ if isinstance (config , Mapping ):
893+ if name in config :
894+ return config [name ]
895+ elif config is not None and hasattr (config , name ):
896+ return getattr (config , name )
897+
898+ if hasattr (scheduler , name ):
899+ return getattr (scheduler , name )
900+ return default
901+
902+ @staticmethod
903+ def _get_posterior_mean (
904+ scheduler : Scheduler , timestep : int | torch .Tensor , x_0 : torch .Tensor , x_t : torch .Tensor
905+ ) -> torch .Tensor :
906+ alpha_t = scheduler .alphas [timestep ]
907+ alpha_prod_t = scheduler .alphas_cumprod [timestep ]
908+ alpha_prod_t_prev = scheduler .alphas_cumprod [timestep - 1 ] if timestep > 0 else scheduler .one
909+
910+ x_0_coefficient = alpha_prod_t_prev .sqrt () * scheduler .betas [timestep ] / (1 - alpha_prod_t )
911+ x_t_coefficient = alpha_t .sqrt () * (1 - alpha_prod_t_prev ) / (1 - alpha_prod_t )
912+
913+ return x_0_coefficient * x_0 + x_t_coefficient * x_t
914+
915+ def _get_posterior_variance (
916+ self , scheduler : Scheduler , timestep : int | torch .Tensor , predicted_variance : torch .Tensor | None = None
917+ ) -> torch .Tensor :
918+ alpha_prod_t = scheduler .alphas_cumprod [timestep ]
919+ alpha_prod_t_prev = scheduler .alphas_cumprod [timestep - 1 ] if timestep > 0 else scheduler .one
920+ variance = (1 - alpha_prod_t_prev ) / (1 - alpha_prod_t ) * scheduler .betas [timestep ]
921+ variance_type = self ._get_scheduler_config_value (scheduler , "variance_type" )
922+
923+ if variance_type == "fixed_small" :
924+ variance = torch .clamp (variance , min = 1e-20 )
925+ elif variance_type == "fixed_large" :
926+ variance = scheduler .betas [timestep ]
927+ elif variance_type == "learned" and predicted_variance is not None :
928+ return predicted_variance
929+ elif variance_type == "learned_range" and predicted_variance is not None :
930+ min_log = variance
931+ max_log = scheduler .betas [timestep ]
932+ frac = (predicted_variance + 1 ) / 2
933+ variance = frac * max_log + (1 - frac ) * min_log
934+
935+ return variance
936+
937+ def _scheduler_step (
938+ self ,
939+ scheduler : Scheduler ,
940+ model_output : torch .Tensor ,
941+ timestep : int | torch .Tensor ,
942+ sample : torch .Tensor ,
943+ next_timestep : int | torch .Tensor | None = None ,
944+ ) -> torch .Tensor :
945+ step_kwargs = {}
946+ if self ._scheduler_step_supports_kwarg (scheduler , "return_dict" ):
947+ step_kwargs ["return_dict" ] = False
948+
949+ if isinstance (scheduler , RFlowScheduler ):
950+ step_output = scheduler .step (model_output , timestep , sample , next_timestep , ** step_kwargs ) # type: ignore
951+ else :
952+ step_output = scheduler .step (model_output , timestep , sample , ** step_kwargs ) # type: ignore
953+
954+ return self ._get_previous_sample_from_step_output (step_output )
955+
865956 def __call__ ( # type: ignore[override]
866957 self ,
867958 inputs : torch .Tensor ,
@@ -941,7 +1032,12 @@ def sample(
9411032 scheduler = self .scheduler
9421033 image = input_noise
9431034
944- all_next_timesteps = torch .cat ((scheduler .timesteps [1 :], torch .tensor ([0 ], dtype = scheduler .timesteps .dtype )))
1035+ all_next_timesteps = torch .cat (
1036+ (
1037+ scheduler .timesteps [1 :],
1038+ torch .tensor ([0 ], dtype = scheduler .timesteps .dtype , device = scheduler .timesteps .device ),
1039+ )
1040+ )
9451041 if verbose and has_tqdm :
9461042 progress_bar = tqdm (
9471043 zip (scheduler .timesteps , all_next_timesteps ),
@@ -985,10 +1081,9 @@ def sample(
9851081 model_output = model_output_uncond + cfg * (model_output_cond - model_output_uncond )
9861082
9871083 # 2. compute previous image: x_t -> x_t-1
988- if not isinstance (scheduler , RFlowScheduler ):
989- image , _ = scheduler .step (model_output , t , image ) # type: ignore
990- else :
991- image , _ = scheduler .step (model_output , t , image , next_t ) # type: ignore
1084+ image = self ._scheduler_step (
1085+ scheduler = scheduler , model_output = model_output , timestep = t , sample = image , next_timestep = next_t
1086+ )
9921087 if save_intermediates and t % intermediate_steps == 0 :
9931088 intermediates .append (image )
9941089
@@ -1029,10 +1124,10 @@ def get_likelihood(
10291124
10301125 if not scheduler :
10311126 scheduler = self .scheduler
1032- if scheduler ._get_name () != "DDPMScheduler" :
1127+ scheduler_name = self ._get_scheduler_name (scheduler )
1128+ if scheduler_name != "DDPMScheduler" :
10331129 raise NotImplementedError (
1034- f"Likelihood computation is only compatible with DDPMScheduler,"
1035- f" you are using { scheduler ._get_name ()} "
1130+ f"Likelihood computation is only compatible with DDPMScheduler, you are using { scheduler_name } ,"
10361131 )
10371132 if mode not in ["crossattn" , "concat" ]:
10381133 raise NotImplementedError (f"{ mode } condition is not supported" )
@@ -1047,7 +1142,7 @@ def get_likelihood(
10471142 total_kl = torch .zeros (inputs .shape [0 ]).to (inputs .device )
10481143 for t in progress_bar :
10491144 timesteps = torch .full (inputs .shape [:1 ], t , device = inputs .device ).long ()
1050- noisy_image = self . scheduler .add_noise (original_samples = inputs , noise = noise , timesteps = timesteps )
1145+ noisy_image = scheduler .add_noise (original_samples = inputs , noise = noise , timesteps = timesteps )
10511146 diffusion_model = (
10521147 partial (diffusion_model , seg = seg )
10531148 if isinstance (diffusion_model , SPADEDiffusionModelUNet )
@@ -1060,7 +1155,8 @@ def get_likelihood(
10601155 model_output = diffusion_model (x = noisy_image , timesteps = timesteps , context = conditioning )
10611156
10621157 # get the model's predicted mean, and variance if it is predicted
1063- if model_output .shape [1 ] == inputs .shape [1 ] * 2 and scheduler .variance_type in ["learned" , "learned_range" ]:
1158+ variance_type = self ._get_scheduler_config_value (scheduler , "variance_type" )
1159+ if model_output .shape [1 ] == inputs .shape [1 ] * 2 and variance_type in ["learned" , "learned_range" ]:
10641160 model_output , predicted_variance = torch .split (model_output , inputs .shape [1 ], dim = 1 )
10651161 else :
10661162 predicted_variance = None
@@ -1073,15 +1169,17 @@ def get_likelihood(
10731169
10741170 # 2. compute predicted original sample from predicted noise also called
10751171 # "predicted x_0" of formula (15) from https://arxiv.org/pdf/2006.11239.pdf
1076- if scheduler .prediction_type == "epsilon" :
1172+ prediction_type = self ._get_scheduler_config_value (scheduler , "prediction_type" )
1173+ if prediction_type == "epsilon" :
10771174 pred_original_sample = (noisy_image - beta_prod_t ** (0.5 ) * model_output ) / alpha_prod_t ** (0.5 )
1078- elif scheduler . prediction_type == "sample" :
1175+ elif prediction_type == "sample" :
10791176 pred_original_sample = model_output
1080- elif scheduler . prediction_type == "v_prediction" :
1177+ elif prediction_type == "v_prediction" :
10811178 pred_original_sample = (alpha_prod_t ** 0.5 ) * noisy_image - (beta_prod_t ** 0.5 ) * model_output
10821179 # 3. Clip "predicted x_0"
1083- if scheduler .clip_sample :
1084- pred_original_sample = torch .clamp (pred_original_sample , - 1 , 1 )
1180+ if self ._get_scheduler_config_value (scheduler , "clip_sample" ):
1181+ clip_sample_range = self ._get_scheduler_config_value (scheduler , "clip_sample_range" , 1.0 )
1182+ pred_original_sample = torch .clamp (pred_original_sample , - clip_sample_range , clip_sample_range )
10851183
10861184 # 4. Compute coefficients for pred_original_sample x_0 and current sample x_t
10871185 # See formula (7) from https://arxiv.org/pdf/2006.11239.pdf
@@ -1093,11 +1191,15 @@ def get_likelihood(
10931191 predicted_mean = pred_original_sample_coeff * pred_original_sample + current_sample_coeff * noisy_image
10941192
10951193 # get the posterior mean and variance
1096- posterior_mean = scheduler ._get_mean (timestep = t , x_0 = inputs , x_t = noisy_image ) # type: ignore[operator]
1097- posterior_variance = scheduler ._get_variance (timestep = t , predicted_variance = predicted_variance ) # type: ignore[operator]
1194+ posterior_mean = self ._get_posterior_mean (scheduler = scheduler , timestep = t , x_0 = inputs , x_t = noisy_image )
1195+ posterior_variance = self ._get_posterior_variance (
1196+ scheduler = scheduler , timestep = t , predicted_variance = predicted_variance
1197+ )
10981198
10991199 log_posterior_variance = torch .log (posterior_variance )
1100- log_predicted_variance = torch .log (predicted_variance ) if predicted_variance else log_posterior_variance
1200+ log_predicted_variance = (
1201+ torch .log (predicted_variance ) if predicted_variance is not None else log_posterior_variance
1202+ )
11011203
11021204 if t == 0 :
11031205 # compute -log p(x_0|x_1)
@@ -1510,7 +1612,12 @@ def sample( # type: ignore[override]
15101612 scheduler = self .scheduler
15111613 image = input_noise
15121614
1513- all_next_timesteps = torch .cat ((scheduler .timesteps [1 :], torch .tensor ([0 ], dtype = scheduler .timesteps .dtype )))
1615+ all_next_timesteps = torch .cat (
1616+ (
1617+ scheduler .timesteps [1 :],
1618+ torch .tensor ([0 ], dtype = scheduler .timesteps .dtype , device = scheduler .timesteps .device ),
1619+ )
1620+ )
15141621 if verbose and has_tqdm :
15151622 progress_bar = tqdm (
15161623 zip (scheduler .timesteps , all_next_timesteps ),
@@ -1584,10 +1691,9 @@ def sample( # type: ignore[override]
15841691 model_output = model_output_uncond + cfg * (model_output_cond - model_output_uncond )
15851692
15861693 # 3. compute previous image: x_t -> x_t-1
1587- if not isinstance (scheduler , RFlowScheduler ):
1588- image , _ = scheduler .step (model_output , t , image ) # type: ignore
1589- else :
1590- image , _ = scheduler .step (model_output , t , image , next_t ) # type: ignore
1694+ image = self ._scheduler_step (
1695+ scheduler = scheduler , model_output = model_output , timestep = t , sample = image , next_timestep = next_t
1696+ )
15911697
15921698 if save_intermediates and t % intermediate_steps == 0 :
15931699 intermediates .append (image )
@@ -1632,10 +1738,10 @@ def get_likelihood( # type: ignore[override]
16321738
16331739 if not scheduler :
16341740 scheduler = self .scheduler
1635- if scheduler ._get_name () != "DDPMScheduler" :
1741+ scheduler_name = self ._get_scheduler_name (scheduler )
1742+ if scheduler_name != "DDPMScheduler" :
16361743 raise NotImplementedError (
1637- f"Likelihood computation is only compatible with DDPMScheduler,"
1638- f" you are using { scheduler ._get_name ()} "
1744+ f"Likelihood computation is only compatible with DDPMScheduler, you are using { scheduler_name } ."
16391745 )
16401746 if mode not in ["crossattn" , "concat" ]:
16411747 raise NotImplementedError (f"{ mode } condition is not supported" )
@@ -1648,7 +1754,7 @@ def get_likelihood( # type: ignore[override]
16481754 total_kl = torch .zeros (inputs .shape [0 ]).to (inputs .device )
16491755 for t in progress_bar :
16501756 timesteps = torch .full (inputs .shape [:1 ], t , device = inputs .device ).long ()
1651- noisy_image = self . scheduler .add_noise (original_samples = inputs , noise = noise , timesteps = timesteps )
1757+ noisy_image = scheduler .add_noise (original_samples = inputs , noise = noise , timesteps = timesteps )
16521758
16531759 diffuse = diffusion_model
16541760 if isinstance (diffusion_model , SPADEDiffusionModelUNet ):
@@ -1681,7 +1787,8 @@ def get_likelihood( # type: ignore[override]
16811787 mid_block_additional_residual = mid_block_res_sample ,
16821788 )
16831789 # get the model's predicted mean, and variance if it is predicted
1684- if model_output .shape [1 ] == inputs .shape [1 ] * 2 and scheduler .variance_type in ["learned" , "learned_range" ]:
1790+ variance_type = self ._get_scheduler_config_value (scheduler , "variance_type" )
1791+ if model_output .shape [1 ] == inputs .shape [1 ] * 2 and variance_type in ["learned" , "learned_range" ]:
16851792 model_output , predicted_variance = torch .split (model_output , inputs .shape [1 ], dim = 1 )
16861793 else :
16871794 predicted_variance = None
@@ -1694,15 +1801,17 @@ def get_likelihood( # type: ignore[override]
16941801
16951802 # 2. compute predicted original sample from predicted noise also called
16961803 # "predicted x_0" of formula (15) from https://arxiv.org/pdf/2006.11239.pdf
1697- if scheduler .prediction_type == "epsilon" :
1804+ prediction_type = self ._get_scheduler_config_value (scheduler , "prediction_type" )
1805+ if prediction_type == "epsilon" :
16981806 pred_original_sample = (noisy_image - beta_prod_t ** (0.5 ) * model_output ) / alpha_prod_t ** (0.5 )
1699- elif scheduler . prediction_type == "sample" :
1807+ elif prediction_type == "sample" :
17001808 pred_original_sample = model_output
1701- elif scheduler . prediction_type == "v_prediction" :
1809+ elif prediction_type == "v_prediction" :
17021810 pred_original_sample = (alpha_prod_t ** 0.5 ) * noisy_image - (beta_prod_t ** 0.5 ) * model_output
17031811 # 3. Clip "predicted x_0"
1704- if scheduler .clip_sample :
1705- pred_original_sample = torch .clamp (pred_original_sample , - 1 , 1 )
1812+ if self ._get_scheduler_config_value (scheduler , "clip_sample" ):
1813+ clip_sample_range = self ._get_scheduler_config_value (scheduler , "clip_sample_range" , 1.0 )
1814+ pred_original_sample = torch .clamp (pred_original_sample , - clip_sample_range , clip_sample_range )
17061815
17071816 # 4. Compute coefficients for pred_original_sample x_0 and current sample x_t
17081817 # See formula (7) from https://arxiv.org/pdf/2006.11239.pdf
@@ -1714,11 +1823,15 @@ def get_likelihood( # type: ignore[override]
17141823 predicted_mean = pred_original_sample_coeff * pred_original_sample + current_sample_coeff * noisy_image
17151824
17161825 # get the posterior mean and variance
1717- posterior_mean = scheduler ._get_mean (timestep = t , x_0 = inputs , x_t = noisy_image ) # type: ignore[operator]
1718- posterior_variance = scheduler ._get_variance (timestep = t , predicted_variance = predicted_variance ) # type: ignore[operator]
1826+ posterior_mean = self ._get_posterior_mean (scheduler = scheduler , timestep = t , x_0 = inputs , x_t = noisy_image )
1827+ posterior_variance = self ._get_posterior_variance (
1828+ scheduler = scheduler , timestep = t , predicted_variance = predicted_variance
1829+ )
17191830
17201831 log_posterior_variance = torch .log (posterior_variance )
1721- log_predicted_variance = torch .log (predicted_variance ) if predicted_variance else log_posterior_variance
1832+ log_predicted_variance = (
1833+ torch .log (predicted_variance ) if predicted_variance is not None else log_posterior_variance
1834+ )
17221835
17231836 if t == 0 :
17241837 # compute -log p(x_0|x_1)
0 commit comments