diff --git a/lvdm/models/ddpm3d.py b/lvdm/models/ddpm3d.py index 77e2d25..fc89595 100644 --- a/lvdm/models/ddpm3d.py +++ b/lvdm/models/ddpm3d.py @@ -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])}) diff --git a/lvdm/models/samplers/ddim.py b/lvdm/models/samplers/ddim.py index 7e8237a..55a16af 100644 --- a/lvdm/models/samplers/ddim.py +++ b/lvdm/models/samplers/ddim.py @@ -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): diff --git a/lvdm/models/utils_diffusion.py b/lvdm/models/utils_diffusion.py index 30043d3..db02a80 100644 --- a/lvdm/models/utils_diffusion.py +++ b/lvdm/models/utils_diffusion.py @@ -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() diff --git a/lvdm/modules/encoders/condition.py b/lvdm/modules/encoders/condition.py index d2a55ce..2cb8108 100644 --- a/lvdm/modules/encoders/condition.py +++ b/lvdm/modules/encoders/condition.py @@ -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) diff --git a/nodes.py b/nodes.py index c90b62f..03f20e4 100644 --- a/nodes.py +++ b/nodes.py @@ -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: