Merge a3b9b3c1c3646b589362d29eec000dbdce07483e into fd5dfb812cfc9e53ff4e83534a49468b72509661

This commit is contained in:
wl2018 2024-12-12 19:36:30 +07:00 committed by GitHub
commit 2d6912e83a
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -421,11 +421,19 @@ class VAE:
return comfy.utils.tiled_scale_multidim(samples, encode_fn, tile=(tile_x,), overlap=overlap, upscale_amount=(1/self.downscale_ratio), out_channels=self.latent_channels, output_device=self.output_device)
def decode(self, samples_in):
predicted_oom = False
samples = None
out = None
pixel_samples = None
try:
memory_used = self.memory_used_decode(samples_in.shape, self.vae_dtype)
model_management.load_models_gpu([self.patcher], memory_required=memory_used)
free_memory = model_management.get_free_memory(self.device)
logging.debug(f"Free memory: {free_memory} bytes, predicted memory useage of one batch: {memory_used} bytes")
if free_memory < memory_used:
logging.warning("Warning: Out of memory is predicted for regular VAE decoding, directly switch to tiled VAE decoding.")
predicted_oom = True
raise model_management.OOM_EXCEPTION
batch_number = int(free_memory / memory_used)
batch_number = max(1, batch_number)
@ -436,7 +444,11 @@ class VAE:
pixel_samples = torch.empty((samples_in.shape[0],) + tuple(out.shape[1:]), device=self.output_device)
pixel_samples[x:x+batch_number] = out
except model_management.OOM_EXCEPTION as e:
logging.warning("Warning: Ran out of memory when regular VAE decoding, retrying with tiled VAE decoding.")
samples = None
out = None
pixel_samples = None
if not predicted_oom:
logging.warning("Warning: Ran out of memory when regular VAE decoding, retrying with tiled VAE decoding.")
dims = samples_in.ndim - 2
if dims == 1:
pixel_samples = self.decode_tiled_1d(samples_in)