From 18b68bb828337a6439247c140e6aaf7ace0a662c Mon Sep 17 00:00:00 2001 From: shiimizu Date: Sun, 1 Dec 2024 13:48:41 -0800 Subject: [PATCH] Use the VAE's device. Fixes #37 --- tiled_vae.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/tiled_vae.py b/tiled_vae.py index 033b4e9..53a71d8 100644 --- a/tiled_vae.py +++ b/tiled_vae.py @@ -352,7 +352,7 @@ def perfcount(fn): def wrapper(*args, **kwargs): ts = time() - if torch.cuda.is_available(): + if torch.cuda.is_available() and devices.device not in ['cpu', devices.cpu]: torch.cuda.reset_peak_memory_stats(devices.device) devices.torch_gc() gc.collect() @@ -361,7 +361,7 @@ def perfcount(fn): devices.torch_gc() gc.collect() - if torch.cuda.is_available(): + if torch.cuda.is_available() and devices.device not in ['cpu', devices.cpu]: vram = torch.cuda.max_memory_allocated(devices.device) / 2**20 print(f'[Tiled VAE]: Done in {time() - ts:.3f}s, max VRAM alloc {vram:.3f} MB') else: @@ -417,8 +417,6 @@ class GroupNormParam: pixels = torch.tensor(self.pixel_list, dtype=torch.float32, device=devices.device) / max_value sum_pixels = torch.sum(pixels) pixels = pixels.unsqueeze(1) / sum_pixels - # var = torch.sum(var * pixels.to(var.device), dim=0) - # mean = torch.sum(mean * pixels.to(var.device), dim=0) var = torch.sum(var * pixels, dim=0) mean = torch.sum(mean * pixels, dim=0) return lambda x: custom_group_norm(x, 32, mean, var, self.weight, self.bias) @@ -761,6 +759,7 @@ class TiledVAE: fast = kwargs['fast'] if 'fast' in kwargs else args[3] color_fix = kwargs['color_fix'] if 'color_fix' in kwargs else False is_decoder = self.is_decoder + devices.device = _vae.device # for shorthand vae = _vae.first_stage_model