Less dependant on cuda

This commit is contained in:
kijai
2024-03-04 13:40:45 +02:00
parent a3794f3a38
commit 565644f746
4 changed files with 42 additions and 33 deletions
+19 -17
View File
@@ -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)
+4 -2
View File
@@ -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
+5 -2
View File
@@ -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),
+14 -12
View File
@@ -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,
):