possible mps fixes

This commit is contained in:
kijai
2024-06-03 20:26:36 +03:00
parent fbe8648001
commit 03eea22840
5 changed files with 33 additions and 16 deletions
+3 -1
View File
@@ -33,6 +33,8 @@ from ...lvdm.models.autoencoder_dualref import VideoDecoder
__conditioning_keys__ = {'concat': 'c_concat',
'crossattn': 'c_crossattn',
'adm': 'y'}
import comfy.model_management as mm
device = mm.get_torch_device()
class DDPM(pl.LightningModule):
# classic DDPM with Gaussian diffusion, in image space
@@ -528,7 +530,7 @@ class LatentDiffusion(DDPM):
n_samples = default(self.en_and_decode_n_samples_a_time, self.temporal_length)
n_rounds = math.ceil(z.shape[0] / n_samples)
with torch.autocast("cuda", enabled=True):
with torch.autocast(mm.get_autocast_device(device), enabled=True):
for n in range(n_rounds):
if isinstance(self.first_stage_model.decoder, VideoDecoder):
kwargs.update({"timesteps": len(z[n * n_samples : (n + 1) * n_samples])})
+5 -2
View File
@@ -6,6 +6,9 @@ from ....lvdm.common import noise_like
from ....lvdm.common import extract_into_tensor
import copy
import comfy.utils
import comfy.model_management as mm
device = mm.get_torch_device()
class DDIMSampler(object):
def __init__(self, model, schedule="linear", **kwargs):
@@ -17,8 +20,8 @@ class DDIMSampler(object):
def register_buffer(self, name, attr):
if type(attr) == torch.Tensor:
if attr.device != torch.device("cuda"):
attr = attr.to(torch.device("cuda"))
if attr.device != torch.device(device):
attr = attr.to(torch.device(device))
setattr(self, name, attr)
def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., verbose=True):
+11 -4
View File
@@ -4,6 +4,8 @@ import torch
import torch.nn.functional as F
from einops import repeat
import comfy.model_management as mm
device = mm.get_torch_device()
def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False):
"""
@@ -29,14 +31,19 @@ def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False):
def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3):
if mm.is_device_mps(device):
dtype = torch.float32
else:
dtype = torch.float64
if schedule == "linear":
betas = (
torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=torch.float64) ** 2
torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=dtype) ** 2
)
elif schedule == "cosine":
timesteps = (
torch.arange(n_timestep + 1, dtype=torch.float64) / n_timestep + cosine_s
torch.arange(n_timestep + 1, dtype=dtype) / n_timestep + cosine_s
)
alphas = timesteps / (1 + cosine_s) * np.pi / 2
alphas = torch.cos(alphas).pow(2)
@@ -45,9 +52,9 @@ def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2,
betas = np.clip(betas, a_min=0, a_max=0.999)
elif schedule == "sqrt_linear":
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64)
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=dtype)
elif schedule == "sqrt":
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64) ** 0.5
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=dtype) ** 0.5
else:
raise ValueError(f"schedule '{schedule}' unknown.")
return betas.numpy()
+10 -8
View File
@@ -7,6 +7,8 @@ from transformers import T5Tokenizer, T5EncoderModel, CLIPTokenizer, CLIPTextMod
from ....lvdm.common import autocast
from ....utils.utils import count_params
import comfy.model_management as mm
device = mm.get_torch_device()
class AbstractEncoder(nn.Module):
def __init__(self):
@@ -41,7 +43,7 @@ class ClassEmbedder(nn.Module):
c = self.embedding(c)
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)
uc = torch.ones((bs,), device=device) * uc_class
uc = {self.key: uc}
@@ -57,7 +59,7 @@ def disabled_train(self, mode=True):
class FrozenT5Embedder(AbstractEncoder):
"""Uses the T5 transformer encoder for text"""
def __init__(self, version="google/t5-v1_1-large", device="cuda", max_length=77,
def __init__(self, version="google/t5-v1_1-large", 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)
@@ -94,7 +96,7 @@ class FrozenCLIPEmbedder(AbstractEncoder):
"hidden"
]
def __init__(self, version="openai/clip-vit-large-patch14", device="cuda", max_length=77,
def __init__(self, version="openai/clip-vit-large-patch14", device=device, max_length=77,
freeze=True, layer="last", layer_idx=None): # clip-vit-base-patch32
super().__init__()
assert layer in self.LAYERS
@@ -138,7 +140,7 @@ class ClipImageEmbedder(nn.Module):
self,
model,
jit=False,
device='cuda' if torch.cuda.is_available() else 'cpu',
device=device,
antialias=True,
ucg_rate=0.
):
@@ -181,7 +183,7 @@ class FrozenOpenCLIPEmbedder(AbstractEncoder):
"penultimate"
]
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77,
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device=device, max_length=77,
freeze=True, layer="last"):
super().__init__()
assert layer in self.LAYERS
@@ -239,7 +241,7 @@ class FrozenOpenCLIPImageEmbedder(AbstractEncoder):
Uses the OpenCLIP vision transformer encoder for images
"""
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77,
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device=device, max_length=77,
freeze=True, layer="pooled", antialias=True, ucg_rate=0.):
super().__init__()
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'),
@@ -297,7 +299,7 @@ class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder):
Uses the OpenCLIP vision transformer encoder for images
"""
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda",
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device=device,
freeze=True, layer="pooled", antialias=True):
super().__init__()
return
@@ -373,7 +375,7 @@ class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder):
return x
class FrozenCLIPT5Encoder(AbstractEncoder):
def __init__(self, clip_version="openai/clip-vit-large-patch14", t5_version="google/t5-v1_1-xl", device="cuda",
def __init__(self, clip_version="openai/clip-vit-large-patch14", t5_version="google/t5-v1_1-xl", device=device,
clip_max_length=77, t5_max_length=77):
super().__init__()
self.clip_encoder = FrozenCLIPEmbedder(clip_version, device, max_length=clip_max_length)
+4 -1
View File
@@ -603,7 +603,10 @@ class ToonCrafterInterpolation:
imtext_cond = torch.cat([text_emb, img_embeds], dim=1)
del cond_images, img_emb, img_emb2, text_emb
fs = torch.tensor([fs], dtype=torch.long, device=self.model.device)
if comfy.model_management.is_device_mps(device):
fs = torch.tensor([fs], dtype=torch.float32, device=self.model.device)
else:
fs = torch.tensor([fs], dtype=torch.float64, device=self.model.device)
cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat]}
if noise_shape[-1] == 32: