possible mps fixes
This commit is contained in:
@@ -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])})
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user