diff --git a/SUPIR/modules/SUPIR_v0.py b/SUPIR/modules/SUPIR_v0.py index b57e5c6..f97f117 100644 --- a/SUPIR/modules/SUPIR_v0.py +++ b/SUPIR/modules/SUPIR_v0.py @@ -24,6 +24,8 @@ import re import torch from functools import partial +import comfy.model_management +device = comfy.model_management.get_torch_device() try: import xformers @@ -680,14 +682,14 @@ if __name__ == '__main__': # # model = instantiate_from_config(opt.model.params.control_stage_config) # hint = model(torch.randn([1, 4, 64, 64]), torch.randn([1]), torch.randn([1, 4, 64, 64])) - # hint = [h.cuda() for h in hint] + # hint = [h.device for h in hint] # print(sum(map(lambda hint: hint.numel(), model.parameters()))) # # unet = instantiate_from_config(opt.model.params.network_config) - # unet = unet.cuda() + # unet = unet.device # - # _output = unet(torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1]).cuda(), torch.randn([1, 77, 1280]).cuda(), - # torch.randn([1, 2560]).cuda(), hint) + # _output = unet(torch.randn([1, 4, 64, 64]).device, torch.randn([1]).device, torch.randn([1, 77, 1280]).device, + # torch.randn([1, 2560]).device, hint) # print(sum(map(lambda _output: _output.numel(), unet.parameters()))) # base @@ -695,31 +697,31 @@ if __name__ == '__main__': opt = OmegaConf.load('../../options/dev/SUPIR_tmp.yaml') model = instantiate_from_config(opt.model.params.control_stage_config) - model = model.cuda() + model = model.to(device) - hint = model(torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1]).cuda(), torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1, 77, 2048]).cuda(), - torch.randn([1, 2816]).cuda()) + hint = model(torch.randn([1, 4, 64, 64]).device, torch.randn([1]).device, torch.randn([1, 4, 64, 64]).device, torch.randn([1, 77, 2048]).device, + torch.randn([1, 2816]).device) #for h in hint: # print(h.shape) # unet = instantiate_from_config(opt.model.params.network_config) - unet = unet.cuda() - _output = unet(torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1]).cuda(), torch.randn([1, 77, 2048]).cuda(), - torch.randn([1, 2816]).cuda(), hint) + unet = unet.to(device) + _output = unet(torch.randn([1, 4, 64, 64]).device, torch.randn([1]).device, torch.randn([1, 77, 2048]).device, + torch.randn([1, 2816]).device, hint) # model = instantiate_from_config(opt.model.params.control_stage_config) - # model = model.cuda() + # model = model.device # # hint = model(torch.randn([1, 4, 64, 64]), torch.randn([1]), torch.randn([1, 4, 64, 64])) - # hint = model(torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1]).cuda(), torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1, 77, 1280]).cuda(), - # torch.randn([1, 2560]).cuda()) - # # hint = [h.cuda() for h in hint] + # hint = model(torch.randn([1, 4, 64, 64]).device, torch.randn([1]).device, torch.randn([1, 4, 64, 64]).device, torch.randn([1, 77, 1280]).device, + # torch.randn([1, 2560]).device) + # # hint = [h.device for h in hint] # # for h in hint: # print(h.shape) # # unet = instantiate_from_config(opt.model.params.network_config) - # unet = unet.cuda() - # _output = unet(torch.randn([1, 4, 64, 64]).cuda(), torch.randn([1]).cuda(), torch.randn([1, 77, 1280]).cuda(), - # torch.randn([1, 2560]).cuda(), hint) + # unet = unet.device + # _output = unet(torch.randn([1, 4, 64, 64]).device, torch.randn([1]).device, torch.randn([1, 77, 1280]).device, + # torch.randn([1, 2560]).device, hint) diff --git a/sgm/models/diffusion.py b/sgm/models/diffusion.py index 9d7c561..2e23d36 100644 --- a/sgm/models/diffusion.py +++ b/sgm/models/diffusion.py @@ -17,7 +17,9 @@ from ..util import ( instantiate_from_config, log_txt_as_img, ) +import comfy.model_management +device = comfy.model_management.get_torch_device() class DiffusionEngine(pl.LightningModule): def __init__( @@ -117,13 +119,13 @@ class DiffusionEngine(pl.LightningModule): @torch.no_grad() def decode_first_stage(self, z): z = 1.0 / self.scale_factor * z - with torch.autocast("cuda", enabled=not self.disable_first_stage_autocast): + with torch.autocast(device, enabled=not self.disable_first_stage_autocast): out = self.first_stage_model.decode(z) return out @torch.no_grad() def encode_first_stage(self, x): - with torch.autocast("cuda", enabled=not self.disable_first_stage_autocast): + with torch.autocast(device, enabled=not self.disable_first_stage_autocast): z = self.first_stage_model.encode(x) z = self.scale_factor * z return z diff --git a/sgm/modules/diffusionmodules/wrappers.py b/sgm/modules/diffusionmodules/wrappers.py index 64571fd..a899e8e 100644 --- a/sgm/modules/diffusionmodules/wrappers.py +++ b/sgm/modules/diffusionmodules/wrappers.py @@ -6,7 +6,9 @@ from packaging import version # torch._dynamo.config.cache_size_limit = 512 OPENAIUNETWRAPPER = ".sgm.modules.diffusionmodules.wrappers.OpenAIWrapper" - +import comfy.model_management +from contextlib import nullcontext +device = comfy.model_management.get_torch_device() class IdentityWrapper(nn.Module): def __init__(self, diffusion_model, compile_model: bool = False): @@ -84,7 +86,8 @@ class ControlWrapper(nn.Module): def forward( self, x: torch.Tensor, t: torch.Tensor, c: dict, control_scale=1, **kwargs ) -> torch.Tensor: - with torch.autocast("cuda", dtype=self.dtype): + autocast_condition = (self.dtype == torch.float16 or self.dtype == torch.bfloat16) and not comfy.model_management.is_device_mps(device) + with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.dtype) if autocast_condition else nullcontext(): control = self.control_model(x=c.get("control", None), timesteps=t, xt=x, control_vector=c.get("control_vector", None), mask_x=c.get("mask_x", None), diff --git a/sgm/modules/encoders/modules.py b/sgm/modules/encoders/modules.py index 1a874ab..a550162 100644 --- a/sgm/modules/encoders/modules.py +++ b/sgm/modules/encoders/modules.py @@ -33,6 +33,8 @@ from ...util import ( ) from ....CKPT_PTH import SDXL_CLIP1_PATH, SDXL_CLIP2_CKPT_PTH +import comfy.model_management +device = comfy.model_management.get_torch_device() class Conv2d(torch.nn.Conv2d): def reset_parameters(self): @@ -342,7 +344,7 @@ class ClassEmbedder(AbstractEmbModel): c = c[:, None, :] return c - def get_unconditional_conditioning(self, bs, device="cuda"): + def get_unconditional_conditioning(self, bs, device=device): uc_class = ( self.n_classes - 1 ) # 1000 classes --> 0 ... 999, one extra class for ucg (class 1000) @@ -367,7 +369,7 @@ class FrozenT5Embedder(AbstractEmbModel): """Uses the T5 transformer encoder for text""" def __init__( - self, version="google/t5-v1_1-xxl", device="cuda", max_length=77, freeze=True + self, version="google/t5-v1_1-xxl", device=device, max_length=77, freeze=True ): # others are google/t5-v1_1-xl and google/t5-v1_1-xxl super().__init__() self.tokenizer = T5Tokenizer.from_pretrained(version) @@ -395,7 +397,7 @@ class FrozenT5Embedder(AbstractEmbModel): return_tensors="pt", ) tokens = batch_encoding["input_ids"].to(self.device) - with torch.autocast("cuda", enabled=False): + with torch.autocast(device, enabled=False): outputs = self.transformer(input_ids=tokens) z = outputs.last_hidden_state return z @@ -410,7 +412,7 @@ class FrozenByT5Embedder(AbstractEmbModel): """ def __init__( - self, version="google/byt5-base", device="cuda", max_length=77, freeze=True + self, version="google/byt5-base", device=device, max_length=77, freeze=True ): # others are google/t5-v1_1-xl and google/t5-v1_1-xxl super().__init__() self.tokenizer = ByT5Tokenizer.from_pretrained(version) @@ -437,7 +439,7 @@ class FrozenByT5Embedder(AbstractEmbModel): return_tensors="pt", ) tokens = batch_encoding["input_ids"].to(self.device) - with torch.autocast("cuda", enabled=False): + with torch.autocast(device, enabled=False): outputs = self.transformer(input_ids=tokens) z = outputs.last_hidden_state return z @@ -454,7 +456,7 @@ class FrozenCLIPEmbedder(AbstractEmbModel): def __init__( self, version="openai/clip-vit-large-patch14", - device="cuda", + device=device, max_length=77, freeze=True, layer="last", @@ -522,7 +524,7 @@ class FrozenOpenCLIPEmbedder2(AbstractEmbModel): self, arch="ViT-H-14", version="laion2b_s32b_b79k", - device="cuda", + device=device, max_length=77, freeze=True, layer="last", @@ -558,7 +560,7 @@ class FrozenOpenCLIPEmbedder2(AbstractEmbModel): for param in self.parameters(): param.requires_grad = False - @autocast + #@autocast def forward(self, text): tokens = open_clip.tokenize(text) z = self.encode_with_transformer(tokens.to(self.device)) @@ -624,7 +626,7 @@ class FrozenOpenCLIPEmbedder(AbstractEmbModel): self, arch="ViT-H-14", version="laion2b_s32b_b79k", - device="cuda", + device=device, max_length=77, freeze=True, layer="last", @@ -694,7 +696,7 @@ class FrozenOpenCLIPImageEmbedder(AbstractEmbModel): self, arch="ViT-H-14", version="laion2b_s32b_b79k", - device="cuda", + device=device, max_length=77, freeze=True, antialias=True, @@ -753,7 +755,7 @@ class FrozenOpenCLIPImageEmbedder(AbstractEmbModel): for param in self.parameters(): param.requires_grad = False - @autocast + # @autocast def forward(self, image, no_dropout=False): z = self.encode_with_vision_transformer(image) tokens = None @@ -850,7 +852,7 @@ class FrozenCLIPT5Encoder(AbstractEmbModel): self, clip_version="openai/clip-vit-large-patch14", t5_version="google/t5-v1_1-xl", - device="cuda", + device=device, clip_max_length=77, t5_max_length=77, ):