Less dependant on cuda
This commit is contained in:
+19
-17
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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,
|
||||
):
|
||||
|
||||
Reference in New Issue
Block a user