@@ -832,12 +832,12 @@ def text_model_fn(x_t, ts, transformer_out, **kwargs):
832832
833833 # 1. Sample gaussian noise
834834 batch_size = 2 # second image is empty for classifier-free guidance
835- image = self . text_noise_scheduler . sample_noise (
836- (batch_size , self .text_unet .in_channels , 64 , 64 ), device = torch_device , generator = generator
837- )
835+ image = torch . randn (
836+ (batch_size , self .text_unet .in_channels , 64 , 64 ), generator = generator
837+ ). to ( torch_device )
838838
839839 # 2. Encode tokens
840- # an empty input is needed to guide the model away from (
840+ # an empty input is needed to guide the model away from it
841841 inputs = self .tokenizer ([prompt , "" ], padding = "max_length" , max_length = 128 , return_tensors = "pt" )
842842 input_ids = inputs ["input_ids" ].to (torch_device )
843843 attention_mask = inputs ["attention_mask" ].to (torch_device )
@@ -850,7 +850,7 @@ def text_model_fn(x_t, ts, transformer_out, **kwargs):
850850 mean , variance , log_variance , pred_xstart = self .p_mean_variance (
851851 text_model_fn , self .text_noise_scheduler , image , t , transformer_out = transformer_out
852852 )
853- noise = self . text_noise_scheduler . sample_noise (image .shape , device = torch_device , generator = generator )
853+ noise = torch . randn (image .shape , generator = generator ). to ( torch_device )
854854 nonzero_mask = (t != 0 ).float ().view (- 1 , * ([1 ] * (len (image .shape ) - 1 ))) # no noise when t == 0
855855 image = mean + nonzero_mask * torch .exp (0.5 * log_variance ) * noise
856856
@@ -873,8 +873,8 @@ def text_model_fn(x_t, ts, transformer_out, **kwargs):
873873 self .upscale_unet .resolution ,
874874 ),
875875 generator = generator ,
876- )
877- image = image . to ( torch_device ) * upsample_temp
876+ ). to ( torch_device )
877+ image = image * upsample_temp
878878
879879 num_trained_timesteps = self .upscale_noise_scheduler .timesteps
880880 inference_step_times = range (0 , num_trained_timesteps , num_trained_timesteps // num_inference_steps_upscale )
@@ -896,7 +896,7 @@ def text_model_fn(x_t, ts, transformer_out, **kwargs):
896896 # 3. optionally sample variance
897897 variance = 0
898898 if eta > 0 :
899- noise = torch .randn (image .shape , generator = generator ).to (image . device )
899+ noise = torch .randn (image .shape , generator = generator ).to (torch_device )
900900 variance = (
901901 self .upscale_noise_scheduler .get_variance (t , num_inference_steps_upscale ).sqrt () * eta * noise
902902 )
0 commit comments