From ea61900a4cdc83f91518614af952ba11d824b5da Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 20 Jun 2024 18:03:15 -0400 Subject: [PATCH] fix bfloat and pag --- modules/pag/pipe_sdxl.py | 7 +++---- modules/sd_samplers_common.py | 1 - modules/sd_vae_approx.py | 2 +- 3 files changed, 4 insertions(+), 6 deletions(-) diff --git a/modules/pag/pipe_sdxl.py b/modules/pag/pipe_sdxl.py index 89a1caa76..429384ea3 100644 --- a/modules/pag/pipe_sdxl.py +++ b/modules/pag/pipe_sdxl.py @@ -461,10 +461,9 @@ class StableDiffusionXLPAGPipeline( image_encoder=image_encoder, feature_extractor=feature_extractor, ) - # if 'force_zeros_for_empty_prompt' in self.config: - # self.register_to_config(force_zeros_for_empty_prompt=force_zeros_for_empty_prompt) - # if 'requires_aesthetics_score' in self.config: - # self.register_to_config(requires_aesthetics_score=requires_aesthetics_score) + if 'requires_aesthetics_score' in self.config: + self.register_to_config(requires_aesthetics_score=requires_aesthetics_score) + self.register_to_config(force_zeros_for_empty_prompt=force_zeros_for_empty_prompt) self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor) self.default_sample_size = self.unet.config.sample_size diff --git a/modules/sd_samplers_common.py b/modules/sd_samplers_common.py index 8d6694f5c..54a38cf55 100644 --- a/modules/sd_samplers_common.py +++ b/modules/sd_samplers_common.py @@ -51,7 +51,6 @@ def single_sample_to_image(sample, approximation=None): sd_cascade = True if len(sample.shape) == 4 and sample.shape[0]: # likely animatediff latent sample = sample.permute(1, 0, 2, 3)[0] - if shared.native: # [-x,x] to [-5,5] sample_max = torch.max(sample) if sample_max > 5: diff --git a/modules/sd_vae_approx.py b/modules/sd_vae_approx.py index 1e5984145..2b4399edb 100644 --- a/modules/sd_vae_approx.py +++ b/modules/sd_vae_approx.py @@ -51,7 +51,7 @@ def nn_approximation(sample): # Approximate NN in_sample = sample.to(device, dtype).unsqueeze(0) sd_vae_approx_model.to(device, dtype) x_sample = sd_vae_approx_model(in_sample) - x_sample = x_sample[0].detach().cpu() + x_sample = x_sample[0].to(torch.float32).detach().cpu() return x_sample except Exception as e: shared.log.error(f'VAE decode approximate: {e}')