diff --git a/README.ZH_CN.md b/README.ZH_CN.md index 021451c..d2b5ee9 100644 --- a/README.ZH_CN.md +++ b/README.ZH_CN.md @@ -52,6 +52,10 @@ git clone https://github.com/yolain/ComfyUI-Easy-Use ## 📜 更新日志 +**v1.2.7** + +- 使用一种新的方式在 loader 中显示模型缩略图(支持 diffusion_models、lors、checkpoints) + **v1.2.6** - 修复了在缺少自定义节点时缺少 “红色框框” 样式的问题。 diff --git a/README.md b/README.md index 1ab1128..8196644 100644 --- a/README.md +++ b/README.md @@ -47,6 +47,10 @@ Double-click install.bat to install the required dependencies ## 📜 Changelog +`**v1.2.7** + +- Using a new way to display the models thumbnails in the loaders (supported diffusion_models、lors、checkpoints) +` **v1.2.6** - Fix missing the "Red Rect" styles when you are missing custom nodes. diff --git a/__init__.py b/__init__.py index d64af85..679ee01 100644 --- a/__init__.py +++ b/__init__.py @@ -11,7 +11,8 @@ node_list = [ "api", "easyNodes", "image", - "logic" + "logic", + "deprecated", ] NODE_CLASS_MAPPINGS = {} diff --git a/py/deprecated.py b/py/deprecated.py new file mode 100644 index 0000000..9985f14 --- /dev/null +++ b/py/deprecated.py @@ -0,0 +1,360 @@ +import torch +import comfy +import comfy.model_management +from .libs.log import log_node_info, log_node_warn +from .libs.adv_encode import advanced_encode +from nodes import ConditioningSetMask, RepeatLatentBatch +from comfy_extras.nodes_mask import LatentCompositeMasked +from .libs.utils import AlwaysEqualProxy +any_type = AlwaysEqualProxy("*") + + +class If: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "any": (any_type,), + "if": (any_type,), + "else": (any_type,), + }, + } + + RETURN_TYPES = (any_type,) + RETURN_NAMES = ("?",) + FUNCTION = "execute" + CATEGORY = "EasyUse/🚫 Deprecated" + DEPRECATED = True + + def execute(self, *args, **kwargs): + return (kwargs['if'] if kwargs['any'] else kwargs['else'],) + + +class poseEditor: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image": ("STRING", {"default": ""}) + }} + + FUNCTION = "output_pose" + CATEGORY = "EasyUse/🚫 Deprecated" + DEPRECATED = True + RETURN_TYPES = () + RETURN_NAMES = () + + def output_pose(self, image): + return () + + +class imageToMask: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image": ("IMAGE",), + "channel": (['red', 'green', 'blue'],), + } + } + + RETURN_TYPES = ("MASK",) + FUNCTION = "convert" + CATEGORY = "EasyUse/🚫 Deprecated" + DEPRECATED = True + + def convert_to_single_channel(self, image, channel='red'): + from PIL import Image + # Convert to RGB mode to access individual channels + image = image.convert('RGB') + + # Extract the desired channel and convert to greyscale + if channel == 'red': + channel_img = image.split()[0].convert('L') + elif channel == 'green': + channel_img = image.split()[1].convert('L') + elif channel == 'blue': + channel_img = image.split()[2].convert('L') + else: + raise ValueError( + "Invalid channel option. Please choose 'red', 'green', or 'blue'.") + + # Convert the greyscale channel back to RGB mode + channel_img = Image.merge( + 'RGB', (channel_img, channel_img, channel_img)) + + return channel_img + + def convert(self, image, channel='red'): + from .libs.image import pil2tensor, tensor2pil + image = self.convert_to_single_channel(tensor2pil(image), channel) + image = pil2tensor(image) + return (image.squeeze().mean(2),) + +# 显示推理时间 +class showSpentTime: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "pipe": ("PIPE_LINE",), + "spent_time": ("INFO", {"default": 'Time will be displayed when reasoning is complete', "forceInput": False}), + }, + "hidden": { + "unique_id": "UNIQUE_ID", + "extra_pnginfo": "EXTRA_PNGINFO", + }, + } + + FUNCTION = "notify" + OUTPUT_NODE = True + CATEGORY = "EasyUse/🚫 Deprecated" + DEPRECATED = True + RETURN_TYPES = () + RETURN_NAMES = () + + def notify(self, pipe, spent_time=None, unique_id=None, extra_pnginfo=None): + if unique_id and extra_pnginfo and "workflow" in extra_pnginfo: + workflow = extra_pnginfo["workflow"] + node = next((x for x in workflow["nodes"] if str(x["id"]) == unique_id), None) + if node: + spent_time = pipe['loader_settings']['spent_time'] if 'spent_time' in pipe['loader_settings'] else '' + node["widgets_values"] = [spent_time] + + return {"ui": {"text": spent_time}, "result": {}} + + +# 潜空间sigma相乘 +class latentNoisy: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "sampler_name": (comfy.samplers.KSampler.SAMPLERS,), + "scheduler": (comfy.samplers.KSampler.SCHEDULERS,), + "steps": ("INT", {"default": 10000, "min": 0, "max": 10000}), + "start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}), + "end_at_step": ("INT", {"default": 10000, "min": 1, "max": 10000}), + "source": (["CPU", "GPU"],), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), + }, + "optional": { + "pipe": ("PIPE_LINE",), + "optional_model": ("MODEL",), + "optional_latent": ("LATENT",) + }} + + RETURN_TYPES = ("PIPE_LINE", "LATENT", "FLOAT",) + RETURN_NAMES = ("pipe", "latent", "sigma",) + FUNCTION = "run" + DEPRECATED = True + + CATEGORY = "EasyUse/🚫 Deprecated" + + def run(self, sampler_name, scheduler, steps, start_at_step, end_at_step, source, seed, pipe=None, optional_model=None, optional_latent=None): + model = optional_model if optional_model is not None else pipe["model"] + batch_size = pipe["loader_settings"]["batch_size"] + empty_latent_height = pipe["loader_settings"]["empty_latent_height"] + empty_latent_width = pipe["loader_settings"]["empty_latent_width"] + + if optional_latent is not None: + samples = optional_latent + else: + torch.manual_seed(seed) + if source == "CPU": + device = "cpu" + else: + device = comfy.model_management.get_torch_device() + noise = torch.randn((batch_size, 4, empty_latent_height // 8, empty_latent_width // 8), dtype=torch.float32, + device=device).cpu() + + samples = {"samples": noise} + + device = comfy.model_management.get_torch_device() + end_at_step = min(steps, end_at_step) + start_at_step = min(start_at_step, end_at_step) + comfy.model_management.load_model_gpu(model) + model_patcher = comfy.model_patcher.ModelPatcher(model.model, load_device=device, offload_device=comfy.model_management.unet_offload_device()) + sampler = comfy.samplers.KSampler(model_patcher, steps=steps, device=device, sampler=sampler_name, + scheduler=scheduler, denoise=1.0, model_options=model.model_options) + sigmas = sampler.sigmas + sigma = sigmas[start_at_step] - sigmas[end_at_step] + sigma /= model.model.latent_format.scale_factor + sigma = sigma.cpu().numpy() + + samples_out = samples.copy() + + s1 = samples["samples"] + samples_out["samples"] = s1 * sigma + + if pipe is None: + pipe = {} + new_pipe = { + **pipe, + "samples": samples_out + } + del pipe + + return (new_pipe, samples_out, sigma) + +# Latent遮罩复合 +class latentCompositeMaskedWithCond: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "pipe": ("PIPE_LINE",), + "text_combine": ("LIST",), + "source_latent": ("LATENT",), + "source_mask": ("MASK",), + "destination_mask": ("MASK",), + "text_combine_mode": (["add", "replace", "cover"], {"default": "add"}), + "replace_text": ("STRING", {"default": ""}) + }, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO", "my_unique_id": "UNIQUE_ID"}, + } + + OUTPUT_IS_LIST = (False, False, True) + + RETURN_TYPES = ("PIPE_LINE", "LATENT", "CONDITIONING") + RETURN_NAMES = ("pipe", "latent", "conditioning",) + FUNCTION = "run" + + CATEGORY = "EasyUse/🚫 Deprecated" + DEPRECATED = True + + def run(self, pipe, text_combine, source_latent, source_mask, destination_mask, text_combine_mode, replace_text, prompt=None, extra_pnginfo=None, my_unique_id=None): + positive = None + clip = pipe["clip"] + destination_latent = pipe["samples"] + + conds = [] + + for text in text_combine: + if text_combine_mode == 'cover': + positive = text + elif text_combine_mode == 'replace' and replace_text != '': + positive = pipe["loader_settings"]["positive"].replace(replace_text, text) + else: + positive = pipe["loader_settings"]["positive"] + ',' + text + positive_token_normalization = pipe["loader_settings"]["positive_token_normalization"] + positive_weight_interpretation = pipe["loader_settings"]["positive_weight_interpretation"] + a1111_prompt_style = pipe["loader_settings"]["a1111_prompt_style"] + positive_cond = pipe["positive"] + + log_node_warn("Positive encoding...") + steps = pipe["loader_settings"]["steps"] if "steps" in pipe["loader_settings"] else 1 + positive_embeddings_final = advanced_encode(clip, positive, + positive_token_normalization, + positive_weight_interpretation, w_max=1.0, + apply_to_pooled='enable', a1111_prompt_style=a1111_prompt_style, steps=steps) + + # source cond + (cond_1,) = ConditioningSetMask().append(positive_cond, source_mask, "default", 1) + (cond_2,) = ConditioningSetMask().append(positive_embeddings_final, destination_mask, "default", 1) + positive_cond = cond_1 + cond_2 + + conds.append(positive_cond) + # latent composite masked + (samples,) = LatentCompositeMasked().composite(destination_latent, source_latent, 0, 0, False) + + new_pipe = { + **pipe, + "samples": samples, + "loader_settings": { + **pipe["loader_settings"], + "positive": positive, + } + } + + del pipe + + return (new_pipe, samples, conds) + +# 噪声注入到潜空间 +class injectNoiseToLatent: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "strength": ("FLOAT", {"default": 0.1, "min": 0.0, "max": 200.0, "step": 0.0001}), + "normalize": ("BOOLEAN", {"default": False}), + "average": ("BOOLEAN", {"default": False}), + }, + "optional": { + "pipe_to_noise": ("PIPE_LINE",), + "image_to_latent": ("IMAGE",), + "latent": ("LATENT",), + "noise": ("LATENT",), + "mask": ("MASK",), + "mix_randn_amount": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1000.0, "step": 0.001}), + "seed": ("INT", {"default": 123, "min": 0, "max": 0xffffffffffffffff, "step": 1}), + } + } + + RETURN_TYPES = ("LATENT",) + FUNCTION = "inject" + CATEGORY = "EasyUse/🚫 Deprecated" + DEPRECATED = True + + + def inject(self,strength, normalize, average, pipe_to_noise=None, noise=None, image_to_latent=None, latent=None, mix_randn_amount=0, mask=None, seed=None): + + vae = pipe_to_noise["vae"] if pipe_to_noise is not None else pipe_to_noise["vae"] + batch_size = pipe_to_noise["loader_settings"]["batch_size"] if pipe_to_noise is not None and "batch_size" in pipe_to_noise["loader_settings"] else 1 + if noise is None and pipe_to_noise is not None: + noise = pipe_to_noise["samples"] + elif noise is None: + raise Exception("InjectNoiseToLatent: No noise provided") + + if image_to_latent is not None and vae is not None: + samples = {"samples": vae.encode(image_to_latent[:, :, :, :3])} + latents = RepeatLatentBatch().repeat(samples, batch_size)[0] + elif latent is not None: + latents = latent + else: + latents = {"samples": noise["samples"].clone()} + + samples = latents.copy() + if latents["samples"].shape != noise["samples"].shape: + raise ValueError("InjectNoiseToLatent: Latent and noise must have the same shape") + if average: + noised = (samples["samples"].clone() + noise["samples"].clone()) / 2 + else: + noised = samples["samples"].clone() + noise["samples"].clone() * strength + if normalize: + noised = noised / noised.std() + if mask is not None: + mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), + size=(noised.shape[2], noised.shape[3]), mode="bilinear") + mask = mask.expand((-1, noised.shape[1], -1, -1)) + if mask.shape[0] < noised.shape[0]: + mask = mask.repeat((noised.shape[0] - 1) // mask.shape[0] + 1, 1, 1, 1)[:noised.shape[0]] + noised = mask * noised + (1 - mask) * latents["samples"] + if mix_randn_amount > 0: + if seed is not None: + torch.manual_seed(seed) + rand_noise = torch.randn_like(noised) + noised = ((1 - mix_randn_amount) * noised + mix_randn_amount * + rand_noise) / ((mix_randn_amount ** 2 + (1 - mix_randn_amount) ** 2) ** 0.5) + samples["samples"] = noised + return (samples,) + + +NODE_CLASS_MAPPINGS = { + "easy if": If, + "easy poseEditor": poseEditor, + "easy imageToMask": imageToMask, + "easy showSpentTime": showSpentTime, + # latent 潜空间 + "easy latentNoisy": latentNoisy, + "easy latentCompositeMaskedWithCond": latentCompositeMaskedWithCond, + "easy injectNoiseToLatent": injectNoiseToLatent, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "easy if": "If (🚫Deprecated)", + "easy poseEditor": "PoseEditor (🚫Deprecated)", + "easy imageToMask": "ImageToMask (🚫Deprecated)", + "easy showSpentTime": "Show Spent Time (🚫Deprecated)", + # latent 潜空间 + "easy latentNoisy": "LatentNoisy (🚫Deprecated)", + "easy latentCompositeMaskedWithCond": "LatentCompositeMaskedWithCond (🚫Deprecated)", + "easy injectNoiseToLatent": "InjectNoiseToLatent (🚫Deprecated)", +} \ No newline at end of file diff --git a/py/dynamiCrafter/__init__.py b/py/dynamiCrafter/__init__.py deleted file mode 100644 index 84023ba..0000000 --- a/py/dynamiCrafter/__init__.py +++ /dev/null @@ -1,334 +0,0 @@ -#credit to ExponentialML for this module -#from https://github.com/ExponentialML/ComfyUI_Native_DynamiCrafter -import os -import torch -import comfy - -from einops import rearrange -from comfy import model_base, model_management -from .lvdm.modules.networks.openaimodel3d import UNetModel as DynamiCrafterUNetModel - -from .utils.model_utils import DynamiCrafterBase, DYNAMICRAFTER_CONFIG, load_image_proj_dict, load_dynamicrafter_dict, get_image_proj_model - -class DynamiCrafter: - - def __init__(self): - self.model_patcher = None - - # There is probably a better way to do this, but with the apply_model callback, this seems necessary. - # The model gets wrapped around a CFG Denoiser class, and handles the conditioning parts there. - # We cannot access it, so we must find the conditioning according to how ComfyUI handles it. - def get_conditioning_pair(self, c_crossattn, use_cfg: bool): - if not use_cfg: - return c_crossattn - - conditioning_group = [] - - for i in range(c_crossattn.shape[0]): - # Get the positive and negative conditioning. - positive_idx = i + 1 - negative_idx = i - - if positive_idx >= c_crossattn.shape[0]: - break - - if not torch.equal(c_crossattn[[positive_idx]], c_crossattn[[negative_idx]]): - conditioning_group = [ - c_crossattn[[positive_idx]], - c_crossattn[[negative_idx]] - ] - break - - if len(conditioning_group) == 0: - raise ValueError("Could not get the appropriate conditioning group.") - - return torch.cat(conditioning_group) - - # apply_model, {"input": input_x, "timestep": timestep_, "c": c, "cond_or_uncond": cond_or_uncond} - def _forward(self, *args): - transformer_options = self.model_patcher.model_options['transformer_options'] - conditioning = transformer_options['conditioning'] - - apply_model = args[0] - - # forward_dict - fd = args[1] - - x, t, model_in_kwargs, _ = fd['input'], fd['timestep'], fd['c'], fd['cond_or_uncond'] - - c_crossattn = model_in_kwargs.pop("c_crossattn") - c_concat = conditioning['c_concat'] - num_video_frames = conditioning['num_video_frames'] - fs = conditioning['fs'] - - original_num_frames = num_video_frames - - # Better way to determine if we're using CFG - # The cond batch will always be num_frames >= 2 since we're doing video, - # so we need get this condition differently here. - if x.shape[0] > num_video_frames: - num_video_frames *= 2 - batch_size = 2 - use_cfg = True - else: - use_cfg = False - batch_size = 1 - - if use_cfg: - c_concat = torch.cat([c_concat] * 2) - - self.validate_forwardable_latent(x, c_concat, num_video_frames, use_cfg) - - x_in, c_concat = map(lambda xc: rearrange(xc, '(b t) c h w -> b c t h w', b=batch_size), (x, c_concat)) - - # We always assume video, so there will always be batched conditionings. - c_crossattn = self.get_conditioning_pair(c_crossattn, use_cfg) - c_crossattn = c_crossattn[:2] if use_cfg else c_crossattn[:1] - context_in = c_crossattn - - img_embs = conditioning['image_emb'] - - if use_cfg: - img_emb_uncond = conditioning['image_emb_uncond'] - img_embs = torch.cat([img_embs, img_emb_uncond]) - - fs = torch.cat([fs] * x_in.shape[0]) - - outs = [] - for i in range(batch_size): - model_in_kwargs['transformer_options']['cond_idx'] = i - x_out = apply_model( - x_in[[i]], - t=torch.cat([t[:1]]), - context_in=context_in[[i]], - c_crossattn=c_crossattn, - cc_concat=c_concat[[i]], # "cc" is to handle naming conflict with apply_model wrapper. - # We want to handle this in the UNet forward. - num_video_frames=num_video_frames // 2 if batch_size > 1 else num_video_frames, - img_emb=img_embs[[i]], - fs=fs[[i]], - **model_in_kwargs - ) - outs.append(x_out) - - x_out = torch.cat(list(reversed(outs))) - x_out = rearrange(x_out, 'b c t h w -> (b t) c h w') - - return x_out - - def assign_forward_args( - self, - model, - c_concat, - image_emb, - image_emb_uncond, - fs, - frames, - ): - model.model_options['transformer_options']['conditioning'] = { - "c_concat": c_concat, - "image_emb": image_emb, - 'image_emb_uncond': image_emb_uncond, - "fs": fs, - "num_video_frames": frames, - } - - def validate_forwardable_latent(self, latent, c_concat, num_video_frames, use_cfg): - check_no_cfg = latent.shape[0] != num_video_frames - check_with_cfg = latent.shape[0] != (num_video_frames * 2) - - latent_batch_size = latent.shape[0] if not use_cfg else latent.shape[0] // 2 - num_frames = num_video_frames if not use_cfg else num_video_frames // 2 - - if all([check_no_cfg, check_with_cfg]): - raise ValueError( - "Please make sure your latent inputs match the number of frames in the DynamiCrafter Processor." - f"Got a latent batch size of ({latent_batch_size}) with number of frames being ({num_frames})." - ) - - latent_h, latent_w = latent.shape[-2:] - c_concat_h, c_concat_w = c_concat.shape[-2:] - - if not all([latent_h == c_concat_h, latent_w == c_concat_w]): - raise ValueError( - "Please make sure that your input latent and image frames are the same height and width.", - f"Image Size: {c_concat_w * 8}, {c_concat_h * 8}, Latent Size: {latent_h * 8}, {latent_w * 8}" - ) - - def process_image_conditioning( - self, - model, - clip_vision, - vae, - image_proj_model, - images, - use_interpolate, - fps: int, - frames: int, - scale_latents: bool - ): - self.model_patcher = model - encoded_latent = vae.encode(images[:, :, :, :3]) - - encoded_image = clip_vision.encode_image(images[:1])['last_hidden_state'] - image_emb = image_proj_model(encoded_image) - - encoded_image_uncond = clip_vision.encode_image(torch.zeros_like(images)[:1])['last_hidden_state'] - image_emb_uncond = image_proj_model(encoded_image_uncond) - - c_concat = encoded_latent - - if scale_latents: - vae_process_input = vae.process_input - vae.process_input = lambda image: (image - .5) * 2 - c_concat = vae.encode(images[:, :, :, :3]) - vae.process_input = vae_process_input - c_concat = model.model.process_latent_in(c_concat) * 1.3 - else: - c_concat = model.model.process_latent_in(c_concat) - - fs = torch.tensor([fps], dtype=torch.long, device=model_management.intermediate_device()) - - model.set_model_unet_function_wrapper(self._forward) - - used_interpolate_processing = False - - if use_interpolate and frames > 16: - raise ValueError( - "When using interpolation mode, the maximum amount of frames are 16." - "If you're doing long video generation, consider using the last frame\ - from the first generation for the next one (autoregressive)." - ) - if encoded_latent.shape[0] == 1: - c_concat = torch.cat([c_concat] * frames, dim=0)[:frames] - - if use_interpolate: - mask = torch.zeros_like(c_concat) - mask[:1] = c_concat[:1] - c_concat = mask - - used_interpolate_processing = True - else: - if use_interpolate and c_concat.shape[0] in [2, 3]: - input_frame_count = c_concat.shape[0] - - # We're just padding to the same type an size of the concat - masked_frames = torch.zeros_like(torch.cat([c_concat[:1]] * frames))[:frames] - - # Start frame - masked_frames[:1] = c_concat[:1] - - end_frame_idx = -1 - - # TODO - speed = 1.0 - if speed < 1.0: - possible_speeds = list(torch.linspace(0, 1.0, c_concat.shape[0])) - speed_from_frames = enumerate(possible_speeds) - speed_idx = min(speed_from_frames, key=lambda n: n[1] - speed)[0] - end_frame_idx = speed_idx - - # End frame - masked_frames[-1:] = c_concat[[end_frame_idx]] - - # Possible middle frame, but not working at the moment. - if input_frame_count == 3: - middle_idx = masked_frames.shape[0] // 2 - middle_idx_frame = c_concat.shape[0] // 2 - masked_frames[[middle_idx]] = c_concat[[middle_idx_frame]] - - c_concat = masked_frames - used_interpolate_processing = True - - print(f"Using interpolation mode with {input_frame_count} frames.") - - if c_concat.shape[0] < frames and not used_interpolate_processing: - print( - "Multiple images found, but interpolation mode is unset. Using the first frame as condition.", - ) - c_concat = torch.cat([c_concat[:1]] * frames) - - c_concat = c_concat[:frames] - - if encoded_latent.shape[0] == 1: - encoded_latent = torch.cat([encoded_latent] * frames)[:frames] - - if encoded_latent.shape[0] < frames and encoded_latent.shape[0] != 1: - encoded_latent = torch.cat( - [encoded_latent] + [encoded_latent[-1:]] * abs(encoded_latent.shape[0] - frames) - )[:frames] - - # We could store this as a state in this Node Class Instance, but to prevent any weird edge cases, - # this should always be passed through the 'stateless' way, and let ComfyUI handle the transformer_options state. - self.assign_forward_args(model, c_concat, image_emb, image_emb_uncond, fs, frames) - - return (model, {"samples": torch.zeros_like(c_concat)}, {"samples": encoded_latent},) - - - # Loader for the DynamiCrafter model. - def load_model_sicts(self, model_path: str): - model_state_dict = comfy.utils.load_torch_file(model_path) - dynamicrafter_dict = load_dynamicrafter_dict(model_state_dict) - image_proj_dict = load_image_proj_dict(model_state_dict) - - return dynamicrafter_dict, image_proj_dict - - def get_prediction_type(self, is_eps: bool, model_config): - if not is_eps and "image_cross_attention_scale_learnable" in model_config.unet_config.keys(): - model_config.unet_config["image_cross_attention_scale_learnable"] = False - - return model_base.ModelType.EPS if is_eps else model_base.ModelType.V_PREDICTION - - def handle_model_management(self, dynamicrafter_dict: dict, model_config): - parameters = comfy.utils.calculate_parameters(dynamicrafter_dict, "model.diffusion_model.") - load_device = model_management.get_torch_device() - unet_dtype = model_management.unet_dtype( - model_params=parameters, - supported_dtypes=model_config.supported_inference_dtypes - ) - manual_cast_dtype = model_management.unet_manual_cast( - unet_dtype, - load_device, - model_config.supported_inference_dtypes - ) - model_config.set_inference_dtype(unet_dtype, manual_cast_dtype) - inital_load_device = model_management.unet_inital_load_device(parameters, unet_dtype) - offload_device = model_management.unet_offload_device() - - return load_device, inital_load_device - - def check_leftover_keys(self, state_dict: dict): - left_over = state_dict.keys() - if len(left_over) > 0: - print("left over keys:", left_over) - - def load_dynamicrafter(self, model_path): - - if os.path.exists(model_path): - dynamicrafter_dict, image_proj_dict = self.load_model_sicts(model_path) - model_config = DynamiCrafterBase(DYNAMICRAFTER_CONFIG) - - dynamicrafter_dict, is_eps = model_config.process_dict_version(state_dict=dynamicrafter_dict) - - MODEL_TYPE = self.get_prediction_type(is_eps, model_config) - load_device, inital_load_device = self.handle_model_management(dynamicrafter_dict, model_config) - - model = model_base.BaseModel( - model_config, - model_type=MODEL_TYPE, - device=inital_load_device, - unet_model=DynamiCrafterUNetModel - ) - - image_proj_model = get_image_proj_model(image_proj_dict) - model.load_model_weights(dynamicrafter_dict, "model.diffusion_model.") - self.check_leftover_keys(dynamicrafter_dict) - - model_patcher = comfy.model_patcher.ModelPatcher( - model, - load_device=load_device, - offload_device=model_management.unet_offload_device(), - current_device=inital_load_device - ) - - return (model_patcher, image_proj_model,) \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/__init__.py b/py/dynamiCrafter/lvdm/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/py/dynamiCrafter/lvdm/basics.py b/py/dynamiCrafter/lvdm/basics.py deleted file mode 100644 index 3fd14ad..0000000 --- a/py/dynamiCrafter/lvdm/basics.py +++ /dev/null @@ -1,102 +0,0 @@ -# adopted from -# https://github.com/openai/improved-diffusion/blob/main/improved_diffusion/gaussian_diffusion.py -# and -# https://github.com/lucidrains/denoising-diffusion-pytorch/blob/7706bdfc6f527f58d33f84b7b522e61e6e3164b3/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py -# and -# https://github.com/openai/guided-diffusion/blob/0ba878e517b276c45d1195eb29f6f5f72659a05b/guided_diffusion/nn.py -# -# thanks! - -import torch.nn as nn -import comfy.ops -ops = comfy.ops.disable_weight_init - -from ..utils.utils import instantiate_from_config - -def disabled_train(self, mode=True): - """Overwrite model.train with this function to make sure train/eval mode - does not change anymore.""" - return self - -def zero_module(module): - """ - Zero out the parameters of a module and return it. - """ - for p in module.parameters(): - p.detach().zero_() - return module - -def scale_module(module, scale): - """ - Scale the parameters of a module and return it. - """ - for p in module.parameters(): - p.detach().mul_(scale) - return module - - -def conv_nd(dims, *args, **kwargs): - """ - Create a 1D, 2D, or 3D convolution module. - """ - if dims == 1: - return nn.Conv1d(*args, **kwargs) - elif dims == 2: - return ops.Conv2d(*args, **kwargs) - elif dims == 3: - return ops.Conv3d(*args, **kwargs) - raise ValueError(f"unsupported dimensions: {dims}") - - -def linear(*args, **kwargs): - """ - Create a linear module. - """ - return ops.Linear(*args, **kwargs) - - -def avg_pool_nd(dims, *args, **kwargs): - """ - Create a 1D, 2D, or 3D average pooling module. - """ - if dims == 1: - return nn.AvgPool1d(*args, **kwargs) - elif dims == 2: - return nn.AvgPool2d(*args, **kwargs) - elif dims == 3: - return nn.AvgPool3d(*args, **kwargs) - raise ValueError(f"unsupported dimensions: {dims}") - - -def nonlinearity(type='silu'): - if type == 'silu': - return nn.SiLU() - elif type == 'leaky_relu': - return nn.LeakyReLU() - - -class GroupNormSpecific(ops.GroupNorm): - def forward(self, x): - return super().forward(x.float()).type(x.dtype) - - -def normalization(channels, num_groups=32, dtype=None, device=None): - """ - Make a standard normalization layer. - :param channels: number of input channels. - :return: an nn.Module for normalization. - """ - return GroupNormSpecific(num_groups, channels, dtype=dtype, device=device) - - -class HybridConditioner(nn.Module): - - def __init__(self, c_concat_config, c_crossattn_config): - super().__init__() - self.concat_conditioner = instantiate_from_config(c_concat_config) - self.crossattn_conditioner = instantiate_from_config(c_crossattn_config) - - def forward(self, c_concat, c_crossattn): - c_concat = self.concat_conditioner(c_concat) - c_crossattn = self.crossattn_conditioner(c_crossattn) - return {'c_concat': [c_concat], 'c_crossattn': [c_crossattn]} \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/common.py b/py/dynamiCrafter/lvdm/common.py deleted file mode 100644 index 55a150b..0000000 --- a/py/dynamiCrafter/lvdm/common.py +++ /dev/null @@ -1,94 +0,0 @@ -import math -from inspect import isfunction -import torch -from torch import nn -import torch.distributed as dist - - -def gather_data(data, return_np=True): - ''' gather data from multiple processes to one list ''' - data_list = [torch.zeros_like(data) for _ in range(dist.get_world_size())] - dist.all_gather(data_list, data) # gather not supported with NCCL - if return_np: - data_list = [data.cpu().numpy() for data in data_list] - return data_list - -def autocast(f): - def do_autocast(*args, **kwargs): - with torch.cuda.amp.autocast(enabled=True, - dtype=torch.get_autocast_gpu_dtype(), - cache_enabled=torch.is_autocast_cache_enabled()): - return f(*args, **kwargs) - return do_autocast - - -def extract_into_tensor(a, t, x_shape): - b, *_ = t.shape - out = a.gather(-1, t) - return out.reshape(b, *((1,) * (len(x_shape) - 1))) - - -def noise_like(shape, device, repeat=False): - repeat_noise = lambda: torch.randn((1, *shape[1:]), device=device).repeat(shape[0], *((1,) * (len(shape) - 1))) - noise = lambda: torch.randn(shape, device=device) - return repeat_noise() if repeat else noise() - - -def default(val, d): - if exists(val): - return val - return d() if isfunction(d) else d - -def exists(val): - return val is not None - -def identity(*args, **kwargs): - return nn.Identity() - -def uniq(arr): - return{el: True for el in arr}.keys() - -def mean_flat(tensor): - """ - Take the mean over all non-batch dimensions. - """ - return tensor.mean(dim=list(range(1, len(tensor.shape)))) - -def ismap(x): - if not isinstance(x, torch.Tensor): - return False - return (len(x.shape) == 4) and (x.shape[1] > 3) - -def isimage(x): - if not isinstance(x,torch.Tensor): - return False - return (len(x.shape) == 4) and (x.shape[1] == 3 or x.shape[1] == 1) - -def max_neg_value(t): - return -torch.finfo(t.dtype).max - -def shape_to_str(x): - shape_str = "x".join([str(x) for x in x.shape]) - return shape_str - -def init_(tensor): - dim = tensor.shape[-1] - std = 1 / math.sqrt(dim) - tensor.uniform_(-std, std) - return tensor - -ckpt = torch.utils.checkpoint.checkpoint -def checkpoint(func, inputs, params, flag): - """ - Evaluate a function without caching intermediate activations, allowing for - reduced memory at the expense of extra compute in the backward pass. - :param func: the function to evaluate. - :param inputs: the argument sequence to pass to `func`. - :param params: a sequence of parameters `func` depends on but does not - explicitly take as arguments. - :param flag: if False, disable gradient checkpointing. - """ - if flag: - return ckpt(func, *inputs, use_reentrant=False) - else: - return func(*inputs) \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/distributions.py b/py/dynamiCrafter/lvdm/distributions.py deleted file mode 100644 index 9a2a82e..0000000 --- a/py/dynamiCrafter/lvdm/distributions.py +++ /dev/null @@ -1,95 +0,0 @@ -import torch -import numpy as np - - -class AbstractDistribution: - def sample(self): - raise NotImplementedError() - - def mode(self): - raise NotImplementedError() - - -class DiracDistribution(AbstractDistribution): - def __init__(self, value): - self.value = value - - def sample(self): - return self.value - - def mode(self): - return self.value - - -class DiagonalGaussianDistribution(object): - def __init__(self, parameters, deterministic=False): - self.parameters = parameters - self.mean, self.logvar = torch.chunk(parameters, 2, dim=1) - self.logvar = torch.clamp(self.logvar, -30.0, 20.0) - self.deterministic = deterministic - self.std = torch.exp(0.5 * self.logvar) - self.var = torch.exp(self.logvar) - if self.deterministic: - self.var = self.std = torch.zeros_like(self.mean).to(device=self.parameters.device) - - def sample(self, noise=None): - if noise is None: - noise = torch.randn(self.mean.shape) - - x = self.mean + self.std * noise.to(device=self.parameters.device) - return x - - def kl(self, other=None): - if self.deterministic: - return torch.Tensor([0.]) - else: - if other is None: - return 0.5 * torch.sum(torch.pow(self.mean, 2) - + self.var - 1.0 - self.logvar, - dim=[1, 2, 3]) - else: - return 0.5 * torch.sum( - torch.pow(self.mean - other.mean, 2) / other.var - + self.var / other.var - 1.0 - self.logvar + other.logvar, - dim=[1, 2, 3]) - - def nll(self, sample, dims=[1,2,3]): - if self.deterministic: - return torch.Tensor([0.]) - logtwopi = np.log(2.0 * np.pi) - return 0.5 * torch.sum( - logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var, - dim=dims) - - def mode(self): - return self.mean - - -def normal_kl(mean1, logvar1, mean2, logvar2): - """ - source: https://github.com/openai/guided-diffusion/blob/27c20a8fab9cb472df5d6bdd6c8d11c8f430b924/guided_diffusion/losses.py#L12 - Compute the KL divergence between two gaussians. - Shapes are automatically broadcasted, so batches can be compared to - scalars, among other use cases. - """ - tensor = None - for obj in (mean1, logvar1, mean2, logvar2): - if isinstance(obj, torch.Tensor): - tensor = obj - break - assert tensor is not None, "at least one argument must be a Tensor" - - # Force variances to be Tensors. Broadcasting helps convert scalars to - # Tensors, but it does not work for torch.exp(). - logvar1, logvar2 = [ - x if isinstance(x, torch.Tensor) else torch.tensor(x).to(tensor) - for x in (logvar1, logvar2) - ] - - return 0.5 * ( - -1.0 - + logvar2 - - logvar1 - + torch.exp(logvar1 - logvar2) - + ((mean1 - mean2) ** 2) * torch.exp(-logvar2) - ) \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/ema.py b/py/dynamiCrafter/lvdm/ema.py deleted file mode 100644 index cd2f8e3..0000000 --- a/py/dynamiCrafter/lvdm/ema.py +++ /dev/null @@ -1,76 +0,0 @@ -import torch -from torch import nn - - -class LitEma(nn.Module): - def __init__(self, model, decay=0.9999, use_num_upates=True): - super().__init__() - if decay < 0.0 or decay > 1.0: - raise ValueError('Decay must be between 0 and 1') - - self.m_name2s_name = {} - self.register_buffer('decay', torch.tensor(decay, dtype=torch.float32)) - self.register_buffer('num_updates', torch.tensor(0,dtype=torch.int) if use_num_upates - else torch.tensor(-1,dtype=torch.int)) - - for name, p in model.named_parameters(): - if p.requires_grad: - #remove as '.'-character is not allowed in buffers - s_name = name.replace('.','') - self.m_name2s_name.update({name:s_name}) - self.register_buffer(s_name,p.clone().detach().data) - - self.collected_params = [] - - def forward(self,model): - decay = self.decay - - if self.num_updates >= 0: - self.num_updates += 1 - decay = min(self.decay,(1 + self.num_updates) / (10 + self.num_updates)) - - one_minus_decay = 1.0 - decay - - with torch.no_grad(): - m_param = dict(model.named_parameters()) - shadow_params = dict(self.named_buffers()) - - for key in m_param: - if m_param[key].requires_grad: - sname = self.m_name2s_name[key] - shadow_params[sname] = shadow_params[sname].type_as(m_param[key]) - shadow_params[sname].sub_(one_minus_decay * (shadow_params[sname] - m_param[key])) - else: - assert not key in self.m_name2s_name - - def copy_to(self, model): - m_param = dict(model.named_parameters()) - shadow_params = dict(self.named_buffers()) - for key in m_param: - if m_param[key].requires_grad: - m_param[key].data.copy_(shadow_params[self.m_name2s_name[key]].data) - else: - assert not key in self.m_name2s_name - - def store(self, parameters): - """ - Save the current parameters for restoring later. - Args: - parameters: Iterable of `torch.nn.Parameter`; the parameters to be - temporarily stored. - """ - self.collected_params = [param.clone() for param in parameters] - - def restore(self, parameters): - """ - Restore the parameters stored with the `store` method. - Useful to validate the model with EMA parameters without affecting the - original optimization process. Store the parameters before the - `copy_to` method. After validation (or model saving), use this to - restore the former parameters. - Args: - parameters: Iterable of `torch.nn.Parameter`; the parameters to be - updated with the stored parameters. - """ - for c_param, param in zip(self.collected_params, parameters): - param.data.copy_(c_param.data) \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/models/autoencoder.py b/py/dynamiCrafter/lvdm/models/autoencoder.py deleted file mode 100644 index cfa86e9..0000000 --- a/py/dynamiCrafter/lvdm/models/autoencoder.py +++ /dev/null @@ -1,219 +0,0 @@ -import os -from contextlib import contextmanager -import torch -import numpy as np -from einops import rearrange -import torch.nn.functional as F -import pytorch_lightning as pl -from ...modules.networks.ae_modules import Encoder, Decoder -from ...distributions import DiagonalGaussianDistribution -from utils.utils import instantiate_from_config - - -class AutoencoderKL(pl.LightningModule): - def __init__(self, - ddconfig, - lossconfig, - embed_dim, - ckpt_path=None, - ignore_keys=[], - image_key="image", - colorize_nlabels=None, - monitor=None, - test=False, - logdir=None, - input_dim=4, - test_args=None, - ): - super().__init__() - self.image_key = image_key - self.encoder = Encoder(**ddconfig) - self.decoder = Decoder(**ddconfig) - self.loss = instantiate_from_config(lossconfig) - assert ddconfig["double_z"] - self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1) - self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1) - self.embed_dim = embed_dim - self.input_dim = input_dim - self.test = test - self.test_args = test_args - self.logdir = logdir - if colorize_nlabels is not None: - assert type(colorize_nlabels)==int - self.register_buffer("colorize", torch.randn(3, colorize_nlabels, 1, 1)) - if monitor is not None: - self.monitor = monitor - if ckpt_path is not None: - self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys) - if self.test: - self.init_test() - - def init_test(self,): - self.test = True - save_dir = os.path.join(self.logdir, "test") - if 'ckpt' in self.test_args: - ckpt_name = os.path.basename(self.test_args.ckpt).split('.ckpt')[0] + f'_epoch{self._cur_epoch}' - self.root = os.path.join(save_dir, ckpt_name) - else: - self.root = save_dir - if 'test_subdir' in self.test_args: - self.root = os.path.join(save_dir, self.test_args.test_subdir) - - self.root_zs = os.path.join(self.root, "zs") - self.root_dec = os.path.join(self.root, "reconstructions") - self.root_inputs = os.path.join(self.root, "inputs") - os.makedirs(self.root, exist_ok=True) - - if self.test_args.save_z: - os.makedirs(self.root_zs, exist_ok=True) - if self.test_args.save_reconstruction: - os.makedirs(self.root_dec, exist_ok=True) - if self.test_args.save_input: - os.makedirs(self.root_inputs, exist_ok=True) - assert(self.test_args is not None) - self.test_maximum = getattr(self.test_args, 'test_maximum', None) - self.count = 0 - self.eval_metrics = {} - self.decodes = [] - self.save_decode_samples = 2048 - - def init_from_ckpt(self, path, ignore_keys=list()): - sd = torch.load(path, map_location="cpu") - try: - self._cur_epoch = sd['epoch'] - sd = sd["state_dict"] - except: - self._cur_epoch = 'null' - keys = list(sd.keys()) - for k in keys: - for ik in ignore_keys: - if k.startswith(ik): - print("Deleting key {} from state_dict.".format(k)) - del sd[k] - self.load_state_dict(sd, strict=False) - # self.load_state_dict(sd, strict=True) - print(f"Restored from {path}") - - def encode(self, x, **kwargs): - - h = self.encoder(x) - moments = self.quant_conv(h) - posterior = DiagonalGaussianDistribution(moments) - return posterior - - def decode(self, z, **kwargs): - z = self.post_quant_conv(z) - dec = self.decoder(z) - return dec - - def forward(self, input, sample_posterior=True): - posterior = self.encode(input) - if sample_posterior: - z = posterior.sample() - else: - z = posterior.mode() - dec = self.decode(z) - return dec, posterior - - def get_input(self, batch, k): - x = batch[k] - if x.dim() == 5 and self.input_dim == 4: - b,c,t,h,w = x.shape - self.b = b - self.t = t - x = rearrange(x, 'b c t h w -> (b t) c h w') - - return x - - def training_step(self, batch, batch_idx, optimizer_idx): - inputs = self.get_input(batch, self.image_key) - reconstructions, posterior = self(inputs) - - if optimizer_idx == 0: - # train encoder+decoder+logvar - aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step, - last_layer=self.get_last_layer(), split="train") - self.log("aeloss", aeloss, prog_bar=True, logger=True, on_step=True, on_epoch=True) - self.log_dict(log_dict_ae, prog_bar=False, logger=True, on_step=True, on_epoch=False) - return aeloss - - if optimizer_idx == 1: - # train the discriminator - discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, optimizer_idx, self.global_step, - last_layer=self.get_last_layer(), split="train") - - self.log("discloss", discloss, prog_bar=True, logger=True, on_step=True, on_epoch=True) - self.log_dict(log_dict_disc, prog_bar=False, logger=True, on_step=True, on_epoch=False) - return discloss - - def validation_step(self, batch, batch_idx): - inputs = self.get_input(batch, self.image_key) - reconstructions, posterior = self(inputs) - aeloss, log_dict_ae = self.loss(inputs, reconstructions, posterior, 0, self.global_step, - last_layer=self.get_last_layer(), split="val") - - discloss, log_dict_disc = self.loss(inputs, reconstructions, posterior, 1, self.global_step, - last_layer=self.get_last_layer(), split="val") - - self.log("val/rec_loss", log_dict_ae["val/rec_loss"]) - self.log_dict(log_dict_ae) - self.log_dict(log_dict_disc) - return self.log_dict - - def configure_optimizers(self): - lr = self.learning_rate - opt_ae = torch.optim.Adam(list(self.encoder.parameters())+ - list(self.decoder.parameters())+ - list(self.quant_conv.parameters())+ - list(self.post_quant_conv.parameters()), - lr=lr, betas=(0.5, 0.9)) - opt_disc = torch.optim.Adam(self.loss.discriminator.parameters(), - lr=lr, betas=(0.5, 0.9)) - return [opt_ae, opt_disc], [] - - def get_last_layer(self): - return self.decoder.conv_out.weight - - @torch.no_grad() - def log_images(self, batch, only_inputs=False, **kwargs): - log = dict() - x = self.get_input(batch, self.image_key) - x = x.to(self.device) - if not only_inputs: - xrec, posterior = self(x) - if x.shape[1] > 3: - # colorize with random projection - assert xrec.shape[1] > 3 - x = self.to_rgb(x) - xrec = self.to_rgb(xrec) - log["samples"] = self.decode(torch.randn_like(posterior.sample())) - log["reconstructions"] = xrec - log["inputs"] = x - return log - - def to_rgb(self, x): - assert self.image_key == "segmentation" - if not hasattr(self, "colorize"): - self.register_buffer("colorize", torch.randn(3, x.shape[1], 1, 1).to(x)) - x = F.conv2d(x, weight=self.colorize) - x = 2.*(x-x.min())/(x.max()-x.min()) - 1. - return x - -class IdentityFirstStage(torch.nn.Module): - def __init__(self, *args, vq_interface=False, **kwargs): - self.vq_interface = vq_interface # TODO: Should be true by default but check to not break older stuff - super().__init__() - - def encode(self, x, *args, **kwargs): - return x - - def decode(self, x, *args, **kwargs): - return x - - def quantize(self, x, *args, **kwargs): - if self.vq_interface: - return x, None, [None, None, None] - return x - - def forward(self, x, *args, **kwargs): - return x \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/models/ddpm3d.py b/py/dynamiCrafter/lvdm/models/ddpm3d.py deleted file mode 100644 index a126ed9..0000000 --- a/py/dynamiCrafter/lvdm/models/ddpm3d.py +++ /dev/null @@ -1,762 +0,0 @@ -""" -wild mixture of -https://github.com/openai/improved-diffusion/blob/e94489283bb876ac1477d5dd7709bbbd2d9902ce/improved_diffusion/gaussian_diffusion.py -https://github.com/lucidrains/denoising-diffusion-pytorch/blob/7706bdfc6f527f58d33f84b7b522e61e6e3164b3/denoising_diffusion_pytorch/denoising_diffusion_pytorch.py -https://github.com/CompVis/taming-transformers --- merci -""" - -from functools import partial -from contextlib import contextmanager -import numpy as np -from tqdm import tqdm -from einops import rearrange, repeat -import logging -mainlogger = logging.getLogger('mainlogger') -import torch -import torch.nn as nn -from torchvision.utils import make_grid - -from ...utils.utils import instantiate_from_config -from ..ema import LitEma -from ..distributions import DiagonalGaussianDistribution -from ..models.utils_diffusion import make_beta_schedule, rescale_zero_terminal_snr -from ..basics import disabled_train -from ..common import ( - extract_into_tensor, - noise_like, - exists, - default -) - -__conditioning_keys__ = {'concat': 'c_concat', - 'crossattn': 'c_crossattn', - 'adm': 'y'} - -class DDPM(nn.Module): - # classic DDPM with Gaussian diffusion, in image space - def __init__(self, - unet_config, - timesteps=1000, - beta_schedule="linear", - loss_type="l2", - ckpt_path=None, - ignore_keys=[], - load_only_unet=False, - monitor=None, - use_ema=True, - first_stage_key="image", - image_size=256, - channels=3, - log_every_t=100, - clip_denoised=True, - linear_start=1e-4, - linear_end=2e-2, - cosine_s=8e-3, - given_betas=None, - original_elbo_weight=0., - v_posterior=0., # weight for choosing posterior variance as sigma = (1-v) * beta_tilde + v * beta - l_simple_weight=1., - conditioning_key=None, - parameterization="eps", # all assuming fixed variance schedules - scheduler_config=None, - use_positional_encodings=False, - learn_logvar=False, - logvar_init=0., - rescale_betas_zero_snr=False, - ): - super().__init__() - assert parameterization in ["eps", "x0", "v"], 'currently only supporting "eps" and "x0" and "v"' - self.parameterization = parameterization - mainlogger.info(f"{self.__class__.__name__}: Running in {self.parameterization}-prediction mode") - self.cond_stage_model = None - self.clip_denoised = clip_denoised - self.log_every_t = log_every_t - self.first_stage_key = first_stage_key - self.channels = channels - self.temporal_length = unet_config.params.temporal_length - self.image_size = image_size # try conv? - if isinstance(self.image_size, int): - self.image_size = [self.image_size, self.image_size] - self.use_positional_encodings = use_positional_encodings - self.model = DiffusionWrapper(unet_config, conditioning_key) - #count_params(self.model, verbose=True) - self.use_ema = use_ema - self.rescale_betas_zero_snr = rescale_betas_zero_snr - if self.use_ema: - self.model_ema = LitEma(self.model) - mainlogger.info(f"Keeping EMAs of {len(list(self.model_ema.buffers()))}.") - - self.use_scheduler = scheduler_config is not None - if self.use_scheduler: - self.scheduler_config = scheduler_config - - self.v_posterior = v_posterior - self.original_elbo_weight = original_elbo_weight - self.l_simple_weight = l_simple_weight - - if monitor is not None: - self.monitor = monitor - if ckpt_path is not None: - self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys, only_model=load_only_unet) - - self.register_schedule(given_betas=given_betas, beta_schedule=beta_schedule, timesteps=timesteps, - linear_start=linear_start, linear_end=linear_end, cosine_s=cosine_s) - - self.loss_type = loss_type - - self.learn_logvar = learn_logvar - self.logvar = torch.full(fill_value=logvar_init, size=(self.num_timesteps,)) - if self.learn_logvar: - self.logvar = nn.Parameter(self.logvar, requires_grad=True) - - def register_schedule(self, given_betas=None, beta_schedule="linear", timesteps=1000, - linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3): - if exists(given_betas): - betas = given_betas - else: - betas = make_beta_schedule(beta_schedule, timesteps, linear_start=linear_start, linear_end=linear_end, - cosine_s=cosine_s) - if self.rescale_betas_zero_snr: - betas = rescale_zero_terminal_snr(betas) - - alphas = 1. - betas - alphas_cumprod = np.cumprod(alphas, axis=0) - alphas_cumprod_prev = np.append(1., alphas_cumprod[:-1]) - - timesteps, = betas.shape - self.num_timesteps = int(timesteps) - self.linear_start = linear_start - self.linear_end = linear_end - assert alphas_cumprod.shape[0] == self.num_timesteps, 'alphas have to be defined for each timestep' - - to_torch = partial(torch.tensor, dtype=torch.float32) - - self.register_buffer('betas', to_torch(betas)) - self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod)) - self.register_buffer('alphas_cumprod_prev', to_torch(alphas_cumprod_prev)) - - # calculations for diffusion q(x_t | x_{t-1}) and others - self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod))) - self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod))) - self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod))) - - if self.parameterization != 'v': - self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod))) - self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod - 1))) - else: - self.register_buffer('sqrt_recip_alphas_cumprod', torch.zeros_like(to_torch(alphas_cumprod))) - self.register_buffer('sqrt_recipm1_alphas_cumprod', torch.zeros_like(to_torch(alphas_cumprod))) - - # calculations for posterior q(x_{t-1} | x_t, x_0) - posterior_variance = (1 - self.v_posterior) * betas * (1. - alphas_cumprod_prev) / ( - 1. - alphas_cumprod) + self.v_posterior * betas - # above: equal to 1. / (1. / (1. - alpha_cumprod_tm1) + alpha_t / beta_t) - self.register_buffer('posterior_variance', to_torch(posterior_variance)) - # below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain - self.register_buffer('posterior_log_variance_clipped', to_torch(np.log(np.maximum(posterior_variance, 1e-20)))) - self.register_buffer('posterior_mean_coef1', to_torch( - betas * np.sqrt(alphas_cumprod_prev) / (1. - alphas_cumprod))) - self.register_buffer('posterior_mean_coef2', to_torch( - (1. - alphas_cumprod_prev) * np.sqrt(alphas) / (1. - alphas_cumprod))) - - if self.parameterization == "eps": - lvlb_weights = self.betas ** 2 / ( - 2 * self.posterior_variance * to_torch(alphas) * (1 - self.alphas_cumprod)) - elif self.parameterization == "x0": - lvlb_weights = 0.5 * np.sqrt(torch.Tensor(alphas_cumprod)) / (2. * 1 - torch.Tensor(alphas_cumprod)) - elif self.parameterization == "v": - lvlb_weights = torch.ones_like(self.betas ** 2 / ( - 2 * self.posterior_variance * to_torch(alphas) * (1 - self.alphas_cumprod))) - else: - raise NotImplementedError("mu not supported") - # TODO how to choose this term - lvlb_weights[0] = lvlb_weights[1] - self.register_buffer('lvlb_weights', lvlb_weights, persistent=False) - assert not torch.isnan(self.lvlb_weights).all() - - @contextmanager - def ema_scope(self, context=None): - if self.use_ema: - self.model_ema.store(self.model.parameters()) - self.model_ema.copy_to(self.model) - if context is not None: - mainlogger.info(f"{context}: Switched to EMA weights") - try: - yield None - finally: - if self.use_ema: - self.model_ema.restore(self.model.parameters()) - if context is not None: - mainlogger.info(f"{context}: Restored training weights") - - def init_from_ckpt(self, path, ignore_keys=list(), only_model=False): - sd = torch.load(path, map_location="cpu") - if "state_dict" in list(sd.keys()): - sd = sd["state_dict"] - keys = list(sd.keys()) - for k in keys: - for ik in ignore_keys: - if k.startswith(ik): - mainlogger.info("Deleting key {} from state_dict.".format(k)) - del sd[k] - missing, unexpected = self.load_state_dict(sd, strict=False) if not only_model else self.model.load_state_dict( - sd, strict=False) - mainlogger.info(f"Restored from {path} with {len(missing)} missing and {len(unexpected)} unexpected keys") - if len(missing) > 0: - mainlogger.info(f"Missing Keys: {missing}") - if len(unexpected) > 0: - mainlogger.info(f"Unexpected Keys: {unexpected}") - - def q_mean_variance(self, x_start, t): - """ - Get the distribution q(x_t | x_0). - :param x_start: the [N x C x ...] tensor of noiseless inputs. - :param t: the number of diffusion steps (minus 1). Here, 0 means one step. - :return: A tuple (mean, variance, log_variance), all of x_start's shape. - """ - mean = (extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start) - variance = extract_into_tensor(1.0 - self.alphas_cumprod, t, x_start.shape) - log_variance = extract_into_tensor(self.log_one_minus_alphas_cumprod, t, x_start.shape) - return mean, variance, log_variance - - def predict_start_from_noise(self, x_t, t, noise): - return ( - extract_into_tensor(self.sqrt_recip_alphas_cumprod, t, x_t.shape) * x_t - - extract_into_tensor(self.sqrt_recipm1_alphas_cumprod, t, x_t.shape) * noise - ) - - def predict_start_from_z_and_v(self, x_t, t, v): - # self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod))) - # self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod))) - return ( - extract_into_tensor(self.sqrt_alphas_cumprod, t, x_t.shape) * x_t - - extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_t.shape) * v - ) - - def predict_eps_from_z_and_v(self, x_t, t, v): - return ( - extract_into_tensor(self.sqrt_alphas_cumprod, t, x_t.shape) * v + - extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_t.shape) * x_t - ) - - def q_posterior(self, x_start, x_t, t): - posterior_mean = ( - extract_into_tensor(self.posterior_mean_coef1, t, x_t.shape) * x_start + - extract_into_tensor(self.posterior_mean_coef2, t, x_t.shape) * x_t - ) - posterior_variance = extract_into_tensor(self.posterior_variance, t, x_t.shape) - posterior_log_variance_clipped = extract_into_tensor(self.posterior_log_variance_clipped, t, x_t.shape) - return posterior_mean, posterior_variance, posterior_log_variance_clipped - - def p_mean_variance(self, x, t, clip_denoised: bool): - model_out = self.model(x, t) - if self.parameterization == "eps": - x_recon = self.predict_start_from_noise(x, t=t, noise=model_out) - elif self.parameterization == "x0": - x_recon = model_out - if clip_denoised: - x_recon.clamp_(-1., 1.) - - model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start=x_recon, x_t=x, t=t) - return model_mean, posterior_variance, posterior_log_variance - - @torch.no_grad() - def p_sample(self, x, t, clip_denoised=True, repeat_noise=False): - b, *_, device = *x.shape, x.device - model_mean, _, model_log_variance = self.p_mean_variance(x=x, t=t, clip_denoised=clip_denoised) - noise = noise_like(x.shape, device, repeat_noise) - # no noise when t == 0 - nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1))) - return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise - - @torch.no_grad() - def p_sample_loop(self, shape, return_intermediates=False): - device = self.betas.device - b = shape[0] - img = torch.randn(shape, device=device) - intermediates = [img] - for i in tqdm(reversed(range(0, self.num_timesteps)), desc='Sampling t', total=self.num_timesteps): - img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long), - clip_denoised=self.clip_denoised) - if i % self.log_every_t == 0 or i == self.num_timesteps - 1: - intermediates.append(img) - if return_intermediates: - return img, intermediates - return img - - @torch.no_grad() - def sample(self, batch_size=16, return_intermediates=False): - image_size = self.image_size - channels = self.channels - return self.p_sample_loop((batch_size, channels, image_size, image_size), - return_intermediates=return_intermediates) - - def q_sample(self, x_start, t, noise=None): - noise = default(noise, lambda: torch.randn_like(x_start)) - return (extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) * x_start + - extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise) - - def get_v(self, x, noise, t): - return ( - extract_into_tensor(self.sqrt_alphas_cumprod, t, x.shape) * noise - - extract_into_tensor(self.sqrt_one_minus_alphas_cumprod, t, x.shape) * x - ) - - def get_input(self, batch, k): - x = batch[k] - x = x.to(memory_format=torch.contiguous_format).float() - return x - - def _get_rows_from_list(self, samples): - n_imgs_per_row = len(samples) - denoise_grid = rearrange(samples, 'n b c h w -> b n c h w') - denoise_grid = rearrange(denoise_grid, 'b n c h w -> (b n) c h w') - denoise_grid = make_grid(denoise_grid, nrow=n_imgs_per_row) - return denoise_grid - - @torch.no_grad() - def log_images(self, batch, N=8, n_row=2, sample=True, return_keys=None, **kwargs): - log = dict() - x = self.get_input(batch, self.first_stage_key) - N = min(x.shape[0], N) - n_row = min(x.shape[0], n_row) - x = x.to(self.device)[:N] - log["inputs"] = x - - # get diffusion row - diffusion_row = list() - x_start = x[:n_row] - - for t in range(self.num_timesteps): - if t % self.log_every_t == 0 or t == self.num_timesteps - 1: - t = repeat(torch.tensor([t]), '1 -> b', b=n_row) - t = t.to(self.device).long() - noise = torch.randn_like(x_start) - x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise) - diffusion_row.append(x_noisy) - - log["diffusion_row"] = self._get_rows_from_list(diffusion_row) - - if sample: - # get denoise row - with self.ema_scope("Plotting"): - samples, denoise_row = self.sample(batch_size=N, return_intermediates=True) - - log["samples"] = samples - log["denoise_row"] = self._get_rows_from_list(denoise_row) - - if return_keys: - if np.intersect1d(list(log.keys()), return_keys).shape[0] == 0: - return log - else: - return {key: log[key] for key in return_keys} - return log - - -class LatentDiffusion(DDPM): - """main class""" - def __init__(self, - first_stage_config, - cond_stage_config, - num_timesteps_cond=None, - cond_stage_key="caption", - cond_stage_trainable=False, - cond_stage_forward=None, - conditioning_key=None, - uncond_prob=0.2, - uncond_type="empty_seq", - scale_factor=1.0, - scale_by_std=False, - encoder_type="2d", - only_model=False, - noise_strength=0, - use_dynamic_rescale=False, - base_scale=0.7, - turning_step=400, - loop_video=False, - fps_condition_type='fs', - perframe_ae=False, - *args, **kwargs): - self.num_timesteps_cond = default(num_timesteps_cond, 1) - self.scale_by_std = scale_by_std - assert self.num_timesteps_cond <= kwargs['timesteps'] - # for backwards compatibility after implementation of DiffusionWrapper - ckpt_path = kwargs.pop("ckpt_path", None) - ignore_keys = kwargs.pop("ignore_keys", []) - conditioning_key = default(conditioning_key, 'crossattn') - super().__init__(conditioning_key=conditioning_key, *args, **kwargs) - - self.cond_stage_trainable = cond_stage_trainable - self.cond_stage_key = cond_stage_key - self.noise_strength = noise_strength - self.use_dynamic_rescale = use_dynamic_rescale - self.loop_video = loop_video - self.fps_condition_type = fps_condition_type - self.perframe_ae = perframe_ae - try: - self.num_downs = len(first_stage_config.params.ddconfig.ch_mult) - 1 - except: - self.num_downs = 0 - if not scale_by_std: - self.scale_factor = scale_factor - else: - self.register_buffer('scale_factor', torch.tensor(scale_factor)) - - if use_dynamic_rescale: - scale_arr1 = np.linspace(1.0, base_scale, turning_step) - scale_arr2 = np.full(self.num_timesteps, base_scale) - scale_arr = np.concatenate((scale_arr1, scale_arr2)) - to_torch = partial(torch.tensor, dtype=torch.float32) - self.register_buffer('scale_arr', to_torch(scale_arr)) - - self.instantiate_first_stage(first_stage_config) - self.instantiate_cond_stage(cond_stage_config) - self.first_stage_config = first_stage_config - self.cond_stage_config = cond_stage_config - self.clip_denoised = False - - self.cond_stage_forward = cond_stage_forward - self.encoder_type = encoder_type - assert(encoder_type in ["2d", "3d"]) - self.uncond_prob = uncond_prob - self.classifier_free_guidance = True if uncond_prob > 0 else False - assert(uncond_type in ["zero_embed", "empty_seq"]) - self.uncond_type = uncond_type - - self.restarted_from_ckpt = False - if ckpt_path is not None: - self.init_from_ckpt(ckpt_path, ignore_keys, only_model=only_model) - self.restarted_from_ckpt = True - - - def make_cond_schedule(self, ): - self.cond_ids = torch.full(size=(self.num_timesteps,), fill_value=self.num_timesteps - 1, dtype=torch.long) - ids = torch.round(torch.linspace(0, self.num_timesteps - 1, self.num_timesteps_cond)).long() - self.cond_ids[:self.num_timesteps_cond] = ids - - def instantiate_first_stage(self, config): - model = instantiate_from_config(config) - self.first_stage_model = model.eval() - self.first_stage_model.train = disabled_train - for param in self.first_stage_model.parameters(): - param.requires_grad = False - - def instantiate_cond_stage(self, config): - if not self.cond_stage_trainable: - model = instantiate_from_config(config) - self.cond_stage_model = model.eval() - self.cond_stage_model.train = disabled_train - for param in self.cond_stage_model.parameters(): - param.requires_grad = False - else: - model = instantiate_from_config(config) - self.cond_stage_model = model - - def get_learned_conditioning(self, c): - if self.cond_stage_forward is None: - if hasattr(self.cond_stage_model, 'encode') and callable(self.cond_stage_model.encode): - c = self.cond_stage_model.encode(c) - if isinstance(c, DiagonalGaussianDistribution): - c = c.mode() - else: - c = self.cond_stage_model(c) - else: - assert hasattr(self.cond_stage_model, self.cond_stage_forward) - c = getattr(self.cond_stage_model, self.cond_stage_forward)(c) - return c - - def get_first_stage_encoding(self, encoder_posterior, noise=None): - if isinstance(encoder_posterior, DiagonalGaussianDistribution): - z = encoder_posterior.sample(noise=noise) - elif isinstance(encoder_posterior, torch.Tensor): - z = encoder_posterior - else: - raise NotImplementedError(f"encoder_posterior of type '{type(encoder_posterior)}' not yet implemented") - return self.scale_factor * z - - @torch.no_grad() - def encode_first_stage(self, x): - if self.encoder_type == "2d" and x.dim() == 5: - b, _, t, _, _ = x.shape - x = rearrange(x, 'b c t h w -> (b t) c h w') - reshape_back = True - else: - reshape_back = False - - ## consume more GPU memory but faster - if not self.perframe_ae: - encoder_posterior = self.first_stage_model.encode(x) - results = self.get_first_stage_encoding(encoder_posterior).detach() - else: ## consume less GPU memory but slower - results = [] - for index in range(x.shape[0]): - frame_batch = self.first_stage_model.encode(x[index:index+1,:,:,:]) - frame_result = self.get_first_stage_encoding(frame_batch).detach() - results.append(frame_result) - results = torch.cat(results, dim=0) - - if reshape_back: - results = rearrange(results, '(b t) c h w -> b c t h w', b=b,t=t) - - return results - - def decode_core(self, z, **kwargs): - if self.encoder_type == "2d" and z.dim() == 5: - b, _, t, _, _ = z.shape - z = rearrange(z, 'b c t h w -> (b t) c h w') - reshape_back = True - else: - reshape_back = False - - if not self.perframe_ae: - z = 1. / self.scale_factor * z - results = self.first_stage_model.decode(z, **kwargs) - else: - results = [] - for index in range(z.shape[0]): - frame_z = 1. / self.scale_factor * z[index:index+1,:,:,:] - frame_result = self.first_stage_model.decode(frame_z, **kwargs) - results.append(frame_result) - results = torch.cat(results, dim=0) - - if reshape_back: - results = rearrange(results, '(b t) c h w -> b c t h w', b=b,t=t) - return results - - @torch.no_grad() - def decode_first_stage(self, z, **kwargs): - return self.decode_core(z, **kwargs) - - # same as above but without decorator - def differentiable_decode_first_stage(self, z, **kwargs): - return self.decode_core(z, **kwargs) - - def forward(self, x, c, **kwargs): - t = torch.randint(0, self.num_timesteps, (x.shape[0],), device=self.device).long() - if self.use_dynamic_rescale: - x = x * extract_into_tensor(self.scale_arr, t, x.shape) - return self.p_losses(x, c, t, **kwargs) - - def apply_model(self, x_noisy, t, cond, **kwargs): - if isinstance(cond, dict): - # hybrid case, cond is exptected to be a dict - pass - else: - if not isinstance(cond, list): - cond = [cond] - key = 'c_concat' if self.model.conditioning_key == 'concat' else 'c_crossattn' - cond = {key: cond} - - x_recon = self.model(x_noisy, t, **cond, **kwargs) - - if isinstance(x_recon, tuple): - return x_recon[0] - else: - return x_recon - - def _get_denoise_row_from_list(self, samples, desc=''): - denoise_row = [] - for zd in tqdm(samples, desc=desc): - denoise_row.append(self.decode_first_stage(zd.to(self.device))) - n_log_timesteps = len(denoise_row) - - denoise_row = torch.stack(denoise_row) # n_log_timesteps, b, C, H, W - - if denoise_row.dim() == 5: - denoise_grid = rearrange(denoise_row, 'n b c h w -> b n c h w') - denoise_grid = rearrange(denoise_grid, 'b n c h w -> (b n) c h w') - denoise_grid = make_grid(denoise_grid, nrow=n_log_timesteps) - elif denoise_row.dim() == 6: - # video, grid_size=[n_log_timesteps*bs, t] - video_length = denoise_row.shape[3] - denoise_grid = rearrange(denoise_row, 'n b c t h w -> b n c t h w') - denoise_grid = rearrange(denoise_grid, 'b n c t h w -> (b n) c t h w') - denoise_grid = rearrange(denoise_grid, 'n c t h w -> (n t) c h w') - denoise_grid = make_grid(denoise_grid, nrow=video_length) - else: - raise ValueError - - return denoise_grid - - - def p_mean_variance(self, x, c, t, clip_denoised: bool, return_x0=False, score_corrector=None, corrector_kwargs=None, **kwargs): - t_in = t - model_out = self.apply_model(x, t_in, c, **kwargs) - - if score_corrector is not None: - assert self.parameterization == "eps" - model_out = score_corrector.modify_score(self, model_out, x, t, c, **corrector_kwargs) - - if self.parameterization == "eps": - x_recon = self.predict_start_from_noise(x, t=t, noise=model_out) - elif self.parameterization == "x0": - x_recon = model_out - else: - raise NotImplementedError() - - if clip_denoised: - x_recon.clamp_(-1., 1.) - - model_mean, posterior_variance, posterior_log_variance = self.q_posterior(x_start=x_recon, x_t=x, t=t) - - if return_x0: - return model_mean, posterior_variance, posterior_log_variance, x_recon - else: - return model_mean, posterior_variance, posterior_log_variance - - @torch.no_grad() - def p_sample(self, x, c, t, clip_denoised=False, repeat_noise=False, return_x0=False, \ - temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, **kwargs): - b, *_, device = *x.shape, x.device - outputs = self.p_mean_variance(x=x, c=c, t=t, clip_denoised=clip_denoised, return_x0=return_x0, \ - score_corrector=score_corrector, corrector_kwargs=corrector_kwargs, **kwargs) - if return_x0: - model_mean, _, model_log_variance, x0 = outputs - else: - model_mean, _, model_log_variance = outputs - - noise = noise_like(x.shape, device, repeat_noise) * temperature - if noise_dropout > 0.: - noise = torch.nn.functional.dropout(noise, p=noise_dropout) - # no noise when t == 0 - nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1))) - - if return_x0: - return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise, x0 - else: - return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise - - @torch.no_grad() - def p_sample_loop(self, cond, shape, return_intermediates=False, x_T=None, verbose=True, callback=None, \ - timesteps=None, mask=None, x0=None, img_callback=None, start_T=None, log_every_t=None, **kwargs): - - if not log_every_t: - log_every_t = self.log_every_t - device = self.betas.device - b = shape[0] - # sample an initial noise - if x_T is None: - img = torch.randn(shape, device=device) - else: - img = x_T - - intermediates = [img] - if timesteps is None: - timesteps = self.num_timesteps - if start_T is not None: - timesteps = min(timesteps, start_T) - - iterator = tqdm(reversed(range(0, timesteps)), desc='Sampling t', total=timesteps) if verbose else reversed(range(0, timesteps)) - - if mask is not None: - assert x0 is not None - assert x0.shape[2:3] == mask.shape[2:3] # spatial size has to match - - for i in iterator: - ts = torch.full((b,), i, device=device, dtype=torch.long) - if self.shorten_cond_schedule: - assert self.model.conditioning_key != 'hybrid' - tc = self.cond_ids[ts].to(cond.device) - cond = self.q_sample(x_start=cond, t=tc, noise=torch.randn_like(cond)) - - img = self.p_sample(img, cond, ts, clip_denoised=self.clip_denoised, **kwargs) - if mask is not None: - img_orig = self.q_sample(x0, ts) - img = img_orig * mask + (1. - mask) * img - - if i % log_every_t == 0 or i == timesteps - 1: - intermediates.append(img) - if callback: callback(i) - if img_callback: img_callback(img, i) - - if return_intermediates: - return img, intermediates - return img - - -class LatentVisualDiffusion(LatentDiffusion): - def __init__(self, img_cond_stage_config, image_proj_stage_config, freeze_embedder=True, *args, **kwargs): - super().__init__(*args, **kwargs) - self._init_embedder(img_cond_stage_config, freeze_embedder) - self.image_proj_model = instantiate_from_config(image_proj_stage_config) - - def _init_embedder(self, config, freeze=True): - embedder = instantiate_from_config(config) - if freeze: - self.embedder = embedder.eval() - self.embedder.train = disabled_train - for param in self.embedder.parameters(): - param.requires_grad = False - - -class DiffusionWrapper(nn.Module): - def __init__(self, diff_model_config, conditioning_key): - super().__init__() - self.diffusion_model = instantiate_from_config(diff_model_config) - self.conditioning_key = conditioning_key - - def forward(self, x, t, c_concat: list = None, c_crossattn: list = None, - c_adm=None, s=None, mask=None, **kwargs): - # temporal_context = fps is foNone - if self.conditioning_key is None: - out = self.diffusion_model(x, t) - elif self.conditioning_key == 'concat': - xc = torch.cat([x] + c_concat, dim=1) - out = self.diffusion_model(xc, t, **kwargs) - elif self.conditioning_key == 'crossattn': - cc = torch.cat(c_crossattn, 1) - out = self.diffusion_model(x, t, context=cc, **kwargs) - elif self.conditioning_key == 'hybrid': - ## it is just right [b,c,t,h,w]: concatenate in channel dim - xc = torch.cat([x] + c_concat, dim=1) - cc = torch.cat(c_crossattn, 1) - out = self.diffusion_model(xc, t, context=cc, **kwargs) - elif self.conditioning_key == 'resblockcond': - cc = c_crossattn[0] - out = self.diffusion_model(x, t, context=cc) - elif self.conditioning_key == 'adm': - cc = c_crossattn[0] - out = self.diffusion_model(x, t, y=cc) - elif self.conditioning_key == 'hybrid-adm': - assert c_adm is not None - xc = torch.cat([x] + c_concat, dim=1) - cc = torch.cat(c_crossattn, 1) - out = self.diffusion_model(xc, t, context=cc, y=c_adm, **kwargs) - elif self.conditioning_key == 'hybrid-time': - assert s is not None - xc = torch.cat([x] + c_concat, dim=1) - cc = torch.cat(c_crossattn, 1) - out = self.diffusion_model(xc, t, context=cc, s=s) - elif self.conditioning_key == 'concat-time-mask': - # assert s is not None - xc = torch.cat([x] + c_concat, dim=1) - out = self.diffusion_model(xc, t, context=None, s=s, mask=mask) - elif self.conditioning_key == 'concat-adm-mask': - # assert s is not None - if c_concat is not None: - xc = torch.cat([x] + c_concat, dim=1) - else: - xc = x - out = self.diffusion_model(xc, t, context=None, y=s, mask=mask) - elif self.conditioning_key == 'hybrid-adm-mask': - cc = torch.cat(c_crossattn, 1) - if c_concat is not None: - xc = torch.cat([x] + c_concat, dim=1) - else: - xc = x - out = self.diffusion_model(xc, t, context=cc, y=s, mask=mask) - elif self.conditioning_key == 'hybrid-time-adm': # adm means y, e.g., class index - # assert s is not None - assert c_adm is not None - xc = torch.cat([x] + c_concat, dim=1) - cc = torch.cat(c_crossattn, 1) - out = self.diffusion_model(xc, t, context=cc, s=s, y=c_adm) - elif self.conditioning_key == 'crossattn-adm': - assert c_adm is not None - cc = torch.cat(c_crossattn, 1) - out = self.diffusion_model(x, t, context=cc, y=c_adm) - else: - raise NotImplementedError() - - return out \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/models/samplers/ddim.py b/py/dynamiCrafter/lvdm/models/samplers/ddim.py deleted file mode 100644 index a3270a0..0000000 --- a/py/dynamiCrafter/lvdm/models/samplers/ddim.py +++ /dev/null @@ -1,317 +0,0 @@ -import numpy as np -from tqdm import tqdm -import torch -from ..models.utils_diffusion import make_ddim_sampling_parameters, make_ddim_timesteps, rescale_noise_cfg -from ..common import noise_like -from ..common import extract_into_tensor -import copy - - -class DDIMSampler(object): - def __init__(self, model, schedule="linear", **kwargs): - super().__init__() - self.model = model - self.ddpm_num_timesteps = model.num_timesteps - self.schedule = schedule - self.counter = 0 - - def register_buffer(self, name, attr): - if type(attr) == torch.Tensor: - if attr.device != torch.device("cuda"): - attr = attr.to(torch.device("cuda")) - setattr(self, name, attr) - - def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., verbose=True): - self.ddim_timesteps = make_ddim_timesteps(ddim_discr_method=ddim_discretize, num_ddim_timesteps=ddim_num_steps, - num_ddpm_timesteps=self.ddpm_num_timesteps,verbose=verbose) - alphas_cumprod = self.model.alphas_cumprod - assert alphas_cumprod.shape[0] == self.ddpm_num_timesteps, 'alphas have to be defined for each timestep' - to_torch = lambda x: x.clone().detach().to(torch.float32).to(self.model.device) - - if self.model.use_dynamic_rescale: - self.ddim_scale_arr = self.model.scale_arr[self.ddim_timesteps] - self.ddim_scale_arr_prev = torch.cat([self.ddim_scale_arr[0:1], self.ddim_scale_arr[:-1]]) - - self.register_buffer('betas', to_torch(self.model.betas)) - self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod)) - self.register_buffer('alphas_cumprod_prev', to_torch(self.model.alphas_cumprod_prev)) - - # calculations for diffusion q(x_t | x_{t-1}) and others - self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod.cpu()))) - self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod.cpu()))) - self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod.cpu()))) - self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu()))) - self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu() - 1))) - - # ddim sampling parameters - ddim_sigmas, ddim_alphas, ddim_alphas_prev = make_ddim_sampling_parameters(alphacums=alphas_cumprod.cpu(), - ddim_timesteps=self.ddim_timesteps, - eta=ddim_eta,verbose=verbose) - self.register_buffer('ddim_sigmas', ddim_sigmas) - self.register_buffer('ddim_alphas', ddim_alphas) - self.register_buffer('ddim_alphas_prev', ddim_alphas_prev) - self.register_buffer('ddim_sqrt_one_minus_alphas', np.sqrt(1. - ddim_alphas)) - sigmas_for_original_sampling_steps = ddim_eta * torch.sqrt( - (1 - self.alphas_cumprod_prev) / (1 - self.alphas_cumprod) * ( - 1 - self.alphas_cumprod / self.alphas_cumprod_prev)) - self.register_buffer('ddim_sigmas_for_original_num_steps', sigmas_for_original_sampling_steps) - - @torch.no_grad() - def sample(self, - S, - batch_size, - shape, - conditioning=None, - callback=None, - normals_sequence=None, - img_callback=None, - quantize_x0=False, - eta=0., - mask=None, - x0=None, - temperature=1., - noise_dropout=0., - score_corrector=None, - corrector_kwargs=None, - verbose=True, - schedule_verbose=False, - x_T=None, - log_every_t=100, - unconditional_guidance_scale=1., - unconditional_conditioning=None, - precision=None, - fs=None, - timestep_spacing='uniform', #uniform_trailing for starting from last timestep - guidance_rescale=0.0, - **kwargs - ): - - # check condition bs - if conditioning is not None: - if isinstance(conditioning, dict): - try: - cbs = conditioning[list(conditioning.keys())[0]].shape[0] - except: - cbs = conditioning[list(conditioning.keys())[0]][0].shape[0] - - if cbs != batch_size: - print(f"Warning: Got {cbs} conditionings but batch-size is {batch_size}") - else: - if conditioning.shape[0] != batch_size: - print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}") - - self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, verbose=schedule_verbose) - - # make shape - if len(shape) == 3: - C, H, W = shape - size = (batch_size, C, H, W) - elif len(shape) == 4: - C, T, H, W = shape - size = (batch_size, C, T, H, W) - - samples, intermediates = self.ddim_sampling(conditioning, size, - callback=callback, - img_callback=img_callback, - quantize_denoised=quantize_x0, - mask=mask, x0=x0, - ddim_use_original_steps=False, - noise_dropout=noise_dropout, - temperature=temperature, - score_corrector=score_corrector, - corrector_kwargs=corrector_kwargs, - x_T=x_T, - log_every_t=log_every_t, - unconditional_guidance_scale=unconditional_guidance_scale, - unconditional_conditioning=unconditional_conditioning, - verbose=verbose, - precision=precision, - fs=fs, - guidance_rescale=guidance_rescale, - **kwargs) - return samples, intermediates - - @torch.no_grad() - def ddim_sampling(self, cond, shape, - x_T=None, ddim_use_original_steps=False, - callback=None, timesteps=None, quantize_denoised=False, - mask=None, x0=None, img_callback=None, log_every_t=100, - temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, - unconditional_guidance_scale=1., unconditional_conditioning=None, verbose=True,precision=None,fs=None,guidance_rescale=0.0, - **kwargs): - device = self.model.betas.device - b = shape[0] - if x_T is None: - img = torch.randn(shape, device=device) - else: - img = x_T - if precision is not None: - if precision == 16: - img = img.to(dtype=torch.float16) - - if timesteps is None: - timesteps = self.ddpm_num_timesteps if ddim_use_original_steps else self.ddim_timesteps - elif timesteps is not None and not ddim_use_original_steps: - subset_end = int(min(timesteps / self.ddim_timesteps.shape[0], 1) * self.ddim_timesteps.shape[0]) - 1 - timesteps = self.ddim_timesteps[:subset_end] - - intermediates = {'x_inter': [img], 'pred_x0': [img]} - time_range = reversed(range(0,timesteps)) if ddim_use_original_steps else np.flip(timesteps) - total_steps = timesteps if ddim_use_original_steps else timesteps.shape[0] - if verbose: - iterator = tqdm(time_range, desc='DDIM Sampler', total=total_steps) - else: - iterator = time_range - - clean_cond = kwargs.pop("clean_cond", False) - - # cond_copy, unconditional_conditioning_copy = copy.deepcopy(cond), copy.deepcopy(unconditional_conditioning) - for i, step in enumerate(iterator): - index = total_steps - i - 1 - ts = torch.full((b,), step, device=device, dtype=torch.long) - - ## use mask to blend noised original latent (img_orig) & new sampled latent (img) - if mask is not None: - assert x0 is not None - if clean_cond: - img_orig = x0 - else: - img_orig = self.model.q_sample(x0, ts) # TODO: deterministic forward pass? - img = img_orig * mask + (1. - mask) * img # keep original & modify use img - - - - - outs = self.p_sample_ddim(img, cond, ts, index=index, use_original_steps=ddim_use_original_steps, - quantize_denoised=quantize_denoised, temperature=temperature, - noise_dropout=noise_dropout, score_corrector=score_corrector, - corrector_kwargs=corrector_kwargs, - unconditional_guidance_scale=unconditional_guidance_scale, - unconditional_conditioning=unconditional_conditioning, - mask=mask,x0=x0,fs=fs,guidance_rescale=guidance_rescale, - **kwargs) - - - img, pred_x0 = outs - if callback: callback(i) - if img_callback: img_callback(pred_x0, i) - - if index % log_every_t == 0 or index == total_steps - 1: - intermediates['x_inter'].append(img) - intermediates['pred_x0'].append(pred_x0) - - return img, intermediates - - @torch.no_grad() - def p_sample_ddim(self, x, c, t, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False, - temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, - unconditional_guidance_scale=1., unconditional_conditioning=None, - uc_type=None, conditional_guidance_scale_temporal=None,mask=None,x0=None,guidance_rescale=0.0,**kwargs): - b, *_, device = *x.shape, x.device - if x.dim() == 5: - is_video = True - else: - is_video = False - - if unconditional_conditioning is None or unconditional_guidance_scale == 1.: - model_output = self.model.apply_model(x, t, c, **kwargs) # unet denoiser - else: - ### do_classifier_free_guidance - if isinstance(c, torch.Tensor) or isinstance(c, dict): - e_t_cond = self.model.apply_model(x, t, c, **kwargs) - e_t_uncond = self.model.apply_model(x, t, unconditional_conditioning, **kwargs) - else: - raise NotImplementedError - - model_output = e_t_uncond + unconditional_guidance_scale * (e_t_cond - e_t_uncond) - - if guidance_rescale > 0.0: - model_output = rescale_noise_cfg(model_output, e_t_cond, guidance_rescale=guidance_rescale) - - if self.model.parameterization == "v": - e_t = self.model.predict_eps_from_z_and_v(x, t, model_output) - else: - e_t = model_output - - if score_corrector is not None: - assert self.model.parameterization == "eps", 'not implemented' - e_t = score_corrector.modify_score(self.model, e_t, x, t, c, **corrector_kwargs) - - alphas = self.model.alphas_cumprod if use_original_steps else self.ddim_alphas - alphas_prev = self.model.alphas_cumprod_prev if use_original_steps else self.ddim_alphas_prev - sqrt_one_minus_alphas = self.model.sqrt_one_minus_alphas_cumprod if use_original_steps else self.ddim_sqrt_one_minus_alphas - # sigmas = self.model.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas - sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas - # select parameters corresponding to the currently considered timestep - - if is_video: - size = (b, 1, 1, 1, 1) - else: - size = (b, 1, 1, 1) - a_t = torch.full(size, alphas[index], device=device) - a_prev = torch.full(size, alphas_prev[index], device=device) - sigma_t = torch.full(size, sigmas[index], device=device) - sqrt_one_minus_at = torch.full(size, sqrt_one_minus_alphas[index],device=device) - - # current prediction for x_0 - if self.model.parameterization != "v": - pred_x0 = (x - sqrt_one_minus_at * e_t) / a_t.sqrt() - else: - pred_x0 = self.model.predict_start_from_z_and_v(x, t, model_output) - - if self.model.use_dynamic_rescale: - scale_t = torch.full(size, self.ddim_scale_arr[index], device=device) - prev_scale_t = torch.full(size, self.ddim_scale_arr_prev[index], device=device) - rescale = (prev_scale_t / scale_t) - pred_x0 *= rescale - - if quantize_denoised: - pred_x0, _, *_ = self.model.first_stage_model.quantize(pred_x0) - # direction pointing to x_t - dir_xt = (1. - a_prev - sigma_t**2).sqrt() * e_t - - noise = sigma_t * noise_like(x.shape, device, repeat_noise) * temperature - if noise_dropout > 0.: - noise = torch.nn.functional.dropout(noise, p=noise_dropout) - - x_prev = a_prev.sqrt() * pred_x0 + dir_xt + noise - - return x_prev, pred_x0 - - @torch.no_grad() - def decode(self, x_latent, cond, t_start, unconditional_guidance_scale=1.0, unconditional_conditioning=None, - use_original_steps=False, callback=None): - - timesteps = np.arange(self.ddpm_num_timesteps) if use_original_steps else self.ddim_timesteps - timesteps = timesteps[:t_start] - - time_range = np.flip(timesteps) - total_steps = timesteps.shape[0] - print(f"Running DDIM Sampling with {total_steps} timesteps") - - iterator = tqdm(time_range, desc='Decoding image', total=total_steps) - x_dec = x_latent - for i, step in enumerate(iterator): - index = total_steps - i - 1 - ts = torch.full((x_latent.shape[0],), step, device=x_latent.device, dtype=torch.long) - x_dec, _ = self.p_sample_ddim(x_dec, cond, ts, index=index, use_original_steps=use_original_steps, - unconditional_guidance_scale=unconditional_guidance_scale, - unconditional_conditioning=unconditional_conditioning) - if callback: callback(i) - return x_dec - - @torch.no_grad() - def stochastic_encode(self, x0, t, use_original_steps=False, noise=None): - # fast, but does not allow for exact reconstruction - # t serves as an index to gather the correct alphas - if use_original_steps: - sqrt_alphas_cumprod = self.sqrt_alphas_cumprod - sqrt_one_minus_alphas_cumprod = self.sqrt_one_minus_alphas_cumprod - else: - sqrt_alphas_cumprod = torch.sqrt(self.ddim_alphas) - sqrt_one_minus_alphas_cumprod = self.ddim_sqrt_one_minus_alphas - - if noise is None: - noise = torch.randn_like(x0) - return (extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 + - extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) * noise) diff --git a/py/dynamiCrafter/lvdm/models/samplers/ddim_multiplecond.py b/py/dynamiCrafter/lvdm/models/samplers/ddim_multiplecond.py deleted file mode 100644 index 1548a23..0000000 --- a/py/dynamiCrafter/lvdm/models/samplers/ddim_multiplecond.py +++ /dev/null @@ -1,323 +0,0 @@ -import numpy as np -from tqdm import tqdm -import torch -from ...models.utils_diffusion import make_ddim_sampling_parameters, make_ddim_timesteps, rescale_noise_cfg -from ..common import noise_like -from ..common import extract_into_tensor -import copy - - -class DDIMSampler(object): - def __init__(self, model, schedule="linear", **kwargs): - super().__init__() - self.model = model - self.ddpm_num_timesteps = model.num_timesteps - self.schedule = schedule - self.counter = 0 - - def register_buffer(self, name, attr): - if type(attr) == torch.Tensor: - if attr.device != torch.device("cuda"): - attr = attr.to(torch.device("cuda")) - setattr(self, name, attr) - - def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., verbose=True): - self.ddim_timesteps = make_ddim_timesteps(ddim_discr_method=ddim_discretize, num_ddim_timesteps=ddim_num_steps, - num_ddpm_timesteps=self.ddpm_num_timesteps,verbose=verbose) - alphas_cumprod = self.model.alphas_cumprod - assert alphas_cumprod.shape[0] == self.ddpm_num_timesteps, 'alphas have to be defined for each timestep' - to_torch = lambda x: x.clone().detach().to(torch.float32).to(self.model.device) - - if self.model.use_dynamic_rescale: - self.ddim_scale_arr = self.model.scale_arr[self.ddim_timesteps] - self.ddim_scale_arr_prev = torch.cat([self.ddim_scale_arr[0:1], self.ddim_scale_arr[:-1]]) - - self.register_buffer('betas', to_torch(self.model.betas)) - self.register_buffer('alphas_cumprod', to_torch(alphas_cumprod)) - self.register_buffer('alphas_cumprod_prev', to_torch(self.model.alphas_cumprod_prev)) - - # calculations for diffusion q(x_t | x_{t-1}) and others - self.register_buffer('sqrt_alphas_cumprod', to_torch(np.sqrt(alphas_cumprod.cpu()))) - self.register_buffer('sqrt_one_minus_alphas_cumprod', to_torch(np.sqrt(1. - alphas_cumprod.cpu()))) - self.register_buffer('log_one_minus_alphas_cumprod', to_torch(np.log(1. - alphas_cumprod.cpu()))) - self.register_buffer('sqrt_recip_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu()))) - self.register_buffer('sqrt_recipm1_alphas_cumprod', to_torch(np.sqrt(1. / alphas_cumprod.cpu() - 1))) - - # ddim sampling parameters - ddim_sigmas, ddim_alphas, ddim_alphas_prev = make_ddim_sampling_parameters(alphacums=alphas_cumprod.cpu(), - ddim_timesteps=self.ddim_timesteps, - eta=ddim_eta,verbose=verbose) - self.register_buffer('ddim_sigmas', ddim_sigmas) - self.register_buffer('ddim_alphas', ddim_alphas) - self.register_buffer('ddim_alphas_prev', ddim_alphas_prev) - self.register_buffer('ddim_sqrt_one_minus_alphas', np.sqrt(1. - ddim_alphas)) - sigmas_for_original_sampling_steps = ddim_eta * torch.sqrt( - (1 - self.alphas_cumprod_prev) / (1 - self.alphas_cumprod) * ( - 1 - self.alphas_cumprod / self.alphas_cumprod_prev)) - self.register_buffer('ddim_sigmas_for_original_num_steps', sigmas_for_original_sampling_steps) - - @torch.no_grad() - def sample(self, - S, - batch_size, - shape, - conditioning=None, - callback=None, - normals_sequence=None, - img_callback=None, - quantize_x0=False, - eta=0., - mask=None, - x0=None, - temperature=1., - noise_dropout=0., - score_corrector=None, - corrector_kwargs=None, - verbose=True, - schedule_verbose=False, - x_T=None, - log_every_t=100, - unconditional_guidance_scale=1., - unconditional_conditioning=None, - precision=None, - fs=None, - timestep_spacing='uniform', #uniform_trailing for starting from last timestep - guidance_rescale=0.0, - # this has to come in the same format as the conditioning, # e.g. as encoded tokens, ... - **kwargs - ): - - # check condition bs - if conditioning is not None: - if isinstance(conditioning, dict): - try: - cbs = conditioning[list(conditioning.keys())[0]].shape[0] - except: - cbs = conditioning[list(conditioning.keys())[0]][0].shape[0] - - if cbs != batch_size: - print(f"Warning: Got {cbs} conditionings but batch-size is {batch_size}") - else: - if conditioning.shape[0] != batch_size: - print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}") - - # print('==> timestep_spacing: ', timestep_spacing, guidance_rescale) - self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, verbose=schedule_verbose) - - # make shape - if len(shape) == 3: - C, H, W = shape - size = (batch_size, C, H, W) - elif len(shape) == 4: - C, T, H, W = shape - size = (batch_size, C, T, H, W) - # print(f'Data shape for DDIM sampling is {size}, eta {eta}') - - samples, intermediates = self.ddim_sampling(conditioning, size, - callback=callback, - img_callback=img_callback, - quantize_denoised=quantize_x0, - mask=mask, x0=x0, - ddim_use_original_steps=False, - noise_dropout=noise_dropout, - temperature=temperature, - score_corrector=score_corrector, - corrector_kwargs=corrector_kwargs, - x_T=x_T, - log_every_t=log_every_t, - unconditional_guidance_scale=unconditional_guidance_scale, - unconditional_conditioning=unconditional_conditioning, - verbose=verbose, - precision=precision, - fs=fs, - guidance_rescale=guidance_rescale, - **kwargs) - return samples, intermediates - - @torch.no_grad() - def ddim_sampling(self, cond, shape, - x_T=None, ddim_use_original_steps=False, - callback=None, timesteps=None, quantize_denoised=False, - mask=None, x0=None, img_callback=None, log_every_t=100, - temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, - unconditional_guidance_scale=1., unconditional_conditioning=None, verbose=True,precision=None,fs=None,guidance_rescale=0.0, - **kwargs): - device = self.model.betas.device - b = shape[0] - if x_T is None: - img = torch.randn(shape, device=device) - else: - img = x_T - if precision is not None: - if precision == 16: - img = img.to(dtype=torch.float16) - - - if timesteps is None: - timesteps = self.ddpm_num_timesteps if ddim_use_original_steps else self.ddim_timesteps - elif timesteps is not None and not ddim_use_original_steps: - subset_end = int(min(timesteps / self.ddim_timesteps.shape[0], 1) * self.ddim_timesteps.shape[0]) - 1 - timesteps = self.ddim_timesteps[:subset_end] - - intermediates = {'x_inter': [img], 'pred_x0': [img]} - time_range = reversed(range(0,timesteps)) if ddim_use_original_steps else np.flip(timesteps) - total_steps = timesteps if ddim_use_original_steps else timesteps.shape[0] - if verbose: - iterator = tqdm(time_range, desc='DDIM Sampler', total=total_steps) - else: - iterator = time_range - - clean_cond = kwargs.pop("clean_cond", False) - - # cond_copy, unconditional_conditioning_copy = copy.deepcopy(cond), copy.deepcopy(unconditional_conditioning) - for i, step in enumerate(iterator): - index = total_steps - i - 1 - ts = torch.full((b,), step, device=device, dtype=torch.long) - - ## use mask to blend noised original latent (img_orig) & new sampled latent (img) - if mask is not None: - assert x0 is not None - if clean_cond: - img_orig = x0 - else: - img_orig = self.model.q_sample(x0, ts) # TODO: deterministic forward pass? - img = img_orig * mask + (1. - mask) * img # keep original & modify use img - - - - - outs = self.p_sample_ddim(img, cond, ts, index=index, use_original_steps=ddim_use_original_steps, - quantize_denoised=quantize_denoised, temperature=temperature, - noise_dropout=noise_dropout, score_corrector=score_corrector, - corrector_kwargs=corrector_kwargs, - unconditional_guidance_scale=unconditional_guidance_scale, - unconditional_conditioning=unconditional_conditioning, - mask=mask,x0=x0,fs=fs,guidance_rescale=guidance_rescale, - **kwargs) - - - - img, pred_x0 = outs - if callback: callback(i) - if img_callback: img_callback(pred_x0, i) - - if index % log_every_t == 0 or index == total_steps - 1: - intermediates['x_inter'].append(img) - intermediates['pred_x0'].append(pred_x0) - - return img, intermediates - - @torch.no_grad() - def p_sample_ddim(self, x, c, t, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False, - temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, - unconditional_guidance_scale=1., unconditional_conditioning=None, - uc_type=None, cfg_img=None,mask=None,x0=None,guidance_rescale=0.0, **kwargs): - b, *_, device = *x.shape, x.device - if x.dim() == 5: - is_video = True - else: - is_video = False - if cfg_img is None: - cfg_img = unconditional_guidance_scale - - unconditional_conditioning_img_nonetext = kwargs['unconditional_conditioning_img_nonetext'] - - - if unconditional_conditioning is None or unconditional_guidance_scale == 1.: - model_output = self.model.apply_model(x, t, c, **kwargs) # unet denoiser - else: - ### with unconditional condition - e_t_cond = self.model.apply_model(x, t, c, **kwargs) - e_t_uncond = self.model.apply_model(x, t, unconditional_conditioning, **kwargs) - e_t_uncond_img = self.model.apply_model(x, t, unconditional_conditioning_img_nonetext, **kwargs) - # text cfg - model_output = e_t_uncond + cfg_img * (e_t_uncond_img - e_t_uncond) + unconditional_guidance_scale * (e_t_cond - e_t_uncond_img) - if guidance_rescale > 0.0: - model_output = rescale_noise_cfg(model_output, e_t_cond, guidance_rescale=guidance_rescale) - - if self.model.parameterization == "v": - e_t = self.model.predict_eps_from_z_and_v(x, t, model_output) - else: - e_t = model_output - - if score_corrector is not None: - assert self.model.parameterization == "eps", 'not implemented' - e_t = score_corrector.modify_score(self.model, e_t, x, t, c, **corrector_kwargs) - - alphas = self.model.alphas_cumprod if use_original_steps else self.ddim_alphas - alphas_prev = self.model.alphas_cumprod_prev if use_original_steps else self.ddim_alphas_prev - sqrt_one_minus_alphas = self.model.sqrt_one_minus_alphas_cumprod if use_original_steps else self.ddim_sqrt_one_minus_alphas - sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas - # select parameters corresponding to the currently considered timestep - - if is_video: - size = (b, 1, 1, 1, 1) - else: - size = (b, 1, 1, 1) - a_t = torch.full(size, alphas[index], device=device) - a_prev = torch.full(size, alphas_prev[index], device=device) - sigma_t = torch.full(size, sigmas[index], device=device) - sqrt_one_minus_at = torch.full(size, sqrt_one_minus_alphas[index],device=device) - - # current prediction for x_0 - if self.model.parameterization != "v": - pred_x0 = (x - sqrt_one_minus_at * e_t) / a_t.sqrt() - else: - pred_x0 = self.model.predict_start_from_z_and_v(x, t, model_output) - - if self.model.use_dynamic_rescale: - scale_t = torch.full(size, self.ddim_scale_arr[index], device=device) - prev_scale_t = torch.full(size, self.ddim_scale_arr_prev[index], device=device) - rescale = (prev_scale_t / scale_t) - pred_x0 *= rescale - - if quantize_denoised: - pred_x0, _, *_ = self.model.first_stage_model.quantize(pred_x0) - # direction pointing to x_t - dir_xt = (1. - a_prev - sigma_t**2).sqrt() * e_t - - noise = sigma_t * noise_like(x.shape, device, repeat_noise) * temperature - if noise_dropout > 0.: - noise = torch.nn.functional.dropout(noise, p=noise_dropout) - - x_prev = a_prev.sqrt() * pred_x0 + dir_xt + noise - - return x_prev, pred_x0 - - @torch.no_grad() - def decode(self, x_latent, cond, t_start, unconditional_guidance_scale=1.0, unconditional_conditioning=None, - use_original_steps=False, callback=None): - - timesteps = np.arange(self.ddpm_num_timesteps) if use_original_steps else self.ddim_timesteps - timesteps = timesteps[:t_start] - - time_range = np.flip(timesteps) - total_steps = timesteps.shape[0] - print(f"Running DDIM Sampling with {total_steps} timesteps") - - iterator = tqdm(time_range, desc='Decoding image', total=total_steps) - x_dec = x_latent - for i, step in enumerate(iterator): - index = total_steps - i - 1 - ts = torch.full((x_latent.shape[0],), step, device=x_latent.device, dtype=torch.long) - x_dec, _ = self.p_sample_ddim(x_dec, cond, ts, index=index, use_original_steps=use_original_steps, - unconditional_guidance_scale=unconditional_guidance_scale, - unconditional_conditioning=unconditional_conditioning) - if callback: callback(i) - return x_dec - - @torch.no_grad() - def stochastic_encode(self, x0, t, use_original_steps=False, noise=None): - # fast, but does not allow for exact reconstruction - # t serves as an index to gather the correct alphas - if use_original_steps: - sqrt_alphas_cumprod = self.sqrt_alphas_cumprod - sqrt_one_minus_alphas_cumprod = self.sqrt_one_minus_alphas_cumprod - else: - sqrt_alphas_cumprod = torch.sqrt(self.ddim_alphas) - sqrt_one_minus_alphas_cumprod = self.ddim_sqrt_one_minus_alphas - - if noise is None: - noise = torch.randn_like(x0) - return (extract_into_tensor(sqrt_alphas_cumprod, t, x0.shape) * x0 + - extract_into_tensor(sqrt_one_minus_alphas_cumprod, t, x0.shape) * noise) \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/models/samplers/unipc/__init__.py b/py/dynamiCrafter/lvdm/models/samplers/unipc/__init__.py deleted file mode 100644 index cdf30d0..0000000 --- a/py/dynamiCrafter/lvdm/models/samplers/unipc/__init__.py +++ /dev/null @@ -1 +0,0 @@ -from .sampler import UniPCSampler \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/models/samplers/unipc/sampler.py b/py/dynamiCrafter/lvdm/models/samplers/unipc/sampler.py deleted file mode 100644 index b68ffff..0000000 --- a/py/dynamiCrafter/lvdm/models/samplers/unipc/sampler.py +++ /dev/null @@ -1,79 +0,0 @@ -"""SAMPLING ONLY.""" - -import torch - -from .uni_pc import NoiseScheduleVP, model_wrapper, UniPC - -class UniPCSampler(object): - def __init__(self, model, **kwargs): - super().__init__() - self.model = model - to_torch = lambda x: x.clone().detach().to(torch.float32).to(model.device) - self.register_buffer('alphas_cumprod', to_torch(model.alphas_cumprod)) - - def register_buffer(self, name, attr): - if type(attr) == torch.Tensor: - if attr.device != torch.device("cuda"): - attr = attr.to(torch.device("cuda")) - setattr(self, name, attr) - - @torch.no_grad() - def sample(self, - S, - batch_size, - shape, - conditioning=None, - callback=None, - normals_sequence=None, - img_callback=None, - quantize_x0=False, - eta=0., - mask=None, - x0=None, - temperature=1., - noise_dropout=0., - score_corrector=None, - corrector_kwargs=None, - verbose=True, - x_T=None, - log_every_t=100, - unconditional_guidance_scale=1., - unconditional_conditioning=None, - # this has to come in the same format as the conditioning, # e.g. as encoded tokens, ... - **kwargs - ): - if conditioning is not None: - if isinstance(conditioning, dict): - cbs = conditioning[list(conditioning.keys())[0]].shape[0] - if cbs != batch_size: - print(f"Warning: Got {cbs} conditionings but batch-size is {batch_size}") - else: - if conditioning.shape[0] != batch_size: - print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}") - - # sampling - C, F, H, W = shape - size = (batch_size, C, H, W) - - device = self.model.betas.device - if x_T is None: - img = torch.randn(size, device=device) - else: - img = x_T - - ns = NoiseScheduleVP('discrete', alphas_cumprod=self.alphas_cumprod) - - model_fn = model_wrapper( - lambda x, t, c: self.model.apply_model(x, t, c), - ns, - model_type="noise", - guidance_type="classifier-free", - condition=conditioning, - unconditional_condition=unconditional_conditioning, - guidance_scale=unconditional_guidance_scale, - ) - - uni_pc = UniPC(model_fn, ns, predict_x0=True, thresholding=False) - x = uni_pc.sample(img, steps=S, skip_type="time_uniform", method="multistep", order=3, lower_order_final=True) - - return x.to(device), None \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/models/samplers/unipc/uni_pc.py b/py/dynamiCrafter/lvdm/models/samplers/unipc/uni_pc.py deleted file mode 100644 index 9a42069..0000000 --- a/py/dynamiCrafter/lvdm/models/samplers/unipc/uni_pc.py +++ /dev/null @@ -1,808 +0,0 @@ -import torch -import torch.nn.functional as F -import math - - -class NoiseScheduleVP: - def __init__( - self, - schedule='discrete', - betas=None, - alphas_cumprod=None, - continuous_beta_0=0.1, - continuous_beta_1=20., - ): - """Create a wrapper class for the forward SDE (VP type). - - *** - Update: We support discrete-time diffusion models by implementing a picewise linear interpolation for log_alpha_t. - We recommend to use schedule='discrete' for the discrete-time diffusion models, especially for high-resolution images. - *** - - The forward SDE ensures that the condition distribution q_{t|0}(x_t | x_0) = N ( alpha_t * x_0, sigma_t^2 * I ). - We further define lambda_t = log(alpha_t) - log(sigma_t), which is the half-logSNR (described in the DPM-Solver paper). - Therefore, we implement the functions for computing alpha_t, sigma_t and lambda_t. For t in [0, T], we have: - - log_alpha_t = self.marginal_log_mean_coeff(t) - sigma_t = self.marginal_std(t) - lambda_t = self.marginal_lambda(t) - - Moreover, as lambda(t) is an invertible function, we also support its inverse function: - - t = self.inverse_lambda(lambda_t) - - =============================================================== - - We support both discrete-time DPMs (trained on n = 0, 1, ..., N-1) and continuous-time DPMs (trained on t in [t_0, T]). - - 1. For discrete-time DPMs: - - For discrete-time DPMs trained on n = 0, 1, ..., N-1, we convert the discrete steps to continuous time steps by: - t_i = (i + 1) / N - e.g. for N = 1000, we have t_0 = 1e-3 and T = t_{N-1} = 1. - We solve the corresponding diffusion ODE from time T = 1 to time t_0 = 1e-3. - - Args: - betas: A `torch.Tensor`. The beta array for the discrete-time DPM. (See the original DDPM paper for details) - alphas_cumprod: A `torch.Tensor`. The cumprod alphas for the discrete-time DPM. (See the original DDPM paper for details) - - Note that we always have alphas_cumprod = cumprod(betas). Therefore, we only need to set one of `betas` and `alphas_cumprod`. - - **Important**: Please pay special attention for the args for `alphas_cumprod`: - The `alphas_cumprod` is the \hat{alpha_n} arrays in the notations of DDPM. Specifically, DDPMs assume that - q_{t_n | 0}(x_{t_n} | x_0) = N ( \sqrt{\hat{alpha_n}} * x_0, (1 - \hat{alpha_n}) * I ). - Therefore, the notation \hat{alpha_n} is different from the notation alpha_t in DPM-Solver. In fact, we have - alpha_{t_n} = \sqrt{\hat{alpha_n}}, - and - log(alpha_{t_n}) = 0.5 * log(\hat{alpha_n}). - - - 2. For continuous-time DPMs: - - We support two types of VPSDEs: linear (DDPM) and cosine (improved-DDPM). The hyperparameters for the noise - schedule are the default settings in DDPM and improved-DDPM: - - Args: - beta_min: A `float` number. The smallest beta for the linear schedule. - beta_max: A `float` number. The largest beta for the linear schedule. - cosine_s: A `float` number. The hyperparameter in the cosine schedule. - cosine_beta_max: A `float` number. The hyperparameter in the cosine schedule. - T: A `float` number. The ending time of the forward process. - - =============================================================== - - Args: - schedule: A `str`. The noise schedule of the forward SDE. 'discrete' for discrete-time DPMs, - 'linear' or 'cosine' for continuous-time DPMs. - Returns: - A wrapper object of the forward SDE (VP type). - - =============================================================== - - Example: - - # For discrete-time DPMs, given betas (the beta array for n = 0, 1, ..., N - 1): - >>> ns = NoiseScheduleVP('discrete', betas=betas) - - # For discrete-time DPMs, given alphas_cumprod (the \hat{alpha_n} array for n = 0, 1, ..., N - 1): - >>> ns = NoiseScheduleVP('discrete', alphas_cumprod=alphas_cumprod) - - # For continuous-time DPMs (VPSDE), linear schedule: - >>> ns = NoiseScheduleVP('linear', continuous_beta_0=0.1, continuous_beta_1=20.) - - """ - - if schedule not in ['discrete', 'linear', 'cosine']: - raise ValueError("Unsupported noise schedule {}. The schedule needs to be 'discrete' or 'linear' or 'cosine'".format(schedule)) - - self.schedule = schedule - if schedule == 'discrete': - if betas is not None: - log_alphas = 0.5 * torch.log(1 - betas).cumsum(dim=0) - else: - assert alphas_cumprod is not None - log_alphas = 0.5 * torch.log(alphas_cumprod) - self.total_N = len(log_alphas) - self.T = 1. - self.t_array = torch.linspace(0., 1., self.total_N + 1)[1:].reshape((1, -1)) - self.log_alpha_array = log_alphas.reshape((1, -1,)) - else: - self.total_N = 1000 - self.beta_0 = continuous_beta_0 - self.beta_1 = continuous_beta_1 - self.cosine_s = 0.008 - self.cosine_beta_max = 999. - self.cosine_t_max = math.atan(self.cosine_beta_max * (1. + self.cosine_s) / math.pi) * 2. * (1. + self.cosine_s) / math.pi - self.cosine_s - self.cosine_log_alpha_0 = math.log(math.cos(self.cosine_s / (1. + self.cosine_s) * math.pi / 2.)) - self.schedule = schedule - if schedule == 'cosine': - # For the cosine schedule, T = 1 will have numerical issues. So we manually set the ending time T. - # Note that T = 0.9946 may be not the optimal setting. However, we find it works well. - self.T = 0.9946 - else: - self.T = 1. - - def marginal_log_mean_coeff(self, t): - """ - Compute log(alpha_t) of a given continuous-time label t in [0, T]. - """ - if self.schedule == 'discrete': - return interpolate_fn(t.reshape((-1, 1)), self.t_array.to(t.device), self.log_alpha_array.to(t.device)).reshape((-1)) - elif self.schedule == 'linear': - return -0.25 * t ** 2 * (self.beta_1 - self.beta_0) - 0.5 * t * self.beta_0 - elif self.schedule == 'cosine': - log_alpha_fn = lambda s: torch.log(torch.cos((s + self.cosine_s) / (1. + self.cosine_s) * math.pi / 2.)) - log_alpha_t = log_alpha_fn(t) - self.cosine_log_alpha_0 - return log_alpha_t - - def marginal_alpha(self, t): - """ - Compute alpha_t of a given continuous-time label t in [0, T]. - """ - return torch.exp(self.marginal_log_mean_coeff(t)) - - def marginal_std(self, t): - """ - Compute sigma_t of a given continuous-time label t in [0, T]. - """ - return torch.sqrt(1. - torch.exp(2. * self.marginal_log_mean_coeff(t))) - - def marginal_lambda(self, t): - """ - Compute lambda_t = log(alpha_t) - log(sigma_t) of a given continuous-time label t in [0, T]. - """ - log_mean_coeff = self.marginal_log_mean_coeff(t) - log_std = 0.5 * torch.log(1. - torch.exp(2. * log_mean_coeff)) - return log_mean_coeff - log_std - - def inverse_lambda(self, lamb): - """ - Compute the continuous-time label t in [0, T] of a given half-logSNR lambda_t. - """ - if self.schedule == 'linear': - tmp = 2. * (self.beta_1 - self.beta_0) * torch.logaddexp(-2. * lamb, torch.zeros((1,)).to(lamb)) - Delta = self.beta_0**2 + tmp - return tmp / (torch.sqrt(Delta) + self.beta_0) / (self.beta_1 - self.beta_0) - elif self.schedule == 'discrete': - log_alpha = -0.5 * torch.logaddexp(torch.zeros((1,)).to(lamb.device), -2. * lamb) - t = interpolate_fn(log_alpha.reshape((-1, 1)), torch.flip(self.log_alpha_array.to(lamb.device), [1]), torch.flip(self.t_array.to(lamb.device), [1])) - return t.reshape((-1,)) - else: - log_alpha = -0.5 * torch.logaddexp(-2. * lamb, torch.zeros((1,)).to(lamb)) - t_fn = lambda log_alpha_t: torch.arccos(torch.exp(log_alpha_t + self.cosine_log_alpha_0)) * 2. * (1. + self.cosine_s) / math.pi - self.cosine_s - t = t_fn(log_alpha) - return t - - -def model_wrapper( - model, - noise_schedule, - model_type="noise", - model_kwargs={}, - guidance_type="uncond", - condition=None, - unconditional_condition=None, - guidance_scale=1., - classifier_fn=None, - classifier_kwargs={}, -): - """Create a wrapper function for the noise prediction model. - - DPM-Solver needs to solve the continuous-time diffusion ODEs. For DPMs trained on discrete-time labels, we need to - firstly wrap the model function to a noise prediction model that accepts the continuous time as the input. - - We support four types of the diffusion model by setting `model_type`: - - 1. "noise": noise prediction model. (Trained by predicting noise). - - 2. "x_start": data prediction model. (Trained by predicting the data x_0 at time 0). - - 3. "v": velocity prediction model. (Trained by predicting the velocity). - The "v" prediction is derivation detailed in Appendix D of [1], and is used in Imagen-Video [2]. - - [1] Salimans, Tim, and Jonathan Ho. "Progressive distillation for fast sampling of diffusion models." - arXiv preprint arXiv:2202.00512 (2022). - [2] Ho, Jonathan, et al. "Imagen Video: High Definition Video Generation with Diffusion Models." - arXiv preprint arXiv:2210.02303 (2022). - - 4. "score": marginal score function. (Trained by denoising score matching). - Note that the score function and the noise prediction model follows a simple relationship: - ``` - noise(x_t, t) = -sigma_t * score(x_t, t) - ``` - - We support three types of guided sampling by DPMs by setting `guidance_type`: - 1. "uncond": unconditional sampling by DPMs. - The input `model` has the following format: - `` - model(x, t_input, **model_kwargs) -> noise | x_start | v | score - `` - - 2. "classifier": classifier guidance sampling [3] by DPMs and another classifier. - The input `model` has the following format: - `` - model(x, t_input, **model_kwargs) -> noise | x_start | v | score - `` - - The input `classifier_fn` has the following format: - `` - classifier_fn(x, t_input, cond, **classifier_kwargs) -> logits(x, t_input, cond) - `` - - [3] P. Dhariwal and A. Q. Nichol, "Diffusion models beat GANs on image synthesis," - in Advances in Neural Information Processing Systems, vol. 34, 2021, pp. 8780-8794. - - 3. "classifier-free": classifier-free guidance sampling by conditional DPMs. - The input `model` has the following format: - `` - model(x, t_input, cond, **model_kwargs) -> noise | x_start | v | score - `` - And if cond == `unconditional_condition`, the model output is the unconditional DPM output. - - [4] Ho, Jonathan, and Tim Salimans. "Classifier-free diffusion guidance." - arXiv preprint arXiv:2207.12598 (2022). - - - The `t_input` is the time label of the model, which may be discrete-time labels (i.e. 0 to 999) - or continuous-time labels (i.e. epsilon to T). - - We wrap the model function to accept only `x` and `t_continuous` as inputs, and outputs the predicted noise: - `` - def model_fn(x, t_continuous) -> noise: - t_input = get_model_input_time(t_continuous) - return noise_pred(model, x, t_input, **model_kwargs) - `` - where `t_continuous` is the continuous time labels (i.e. epsilon to T). And we use `model_fn` for DPM-Solver. - - =============================================================== - - Args: - model: A diffusion model with the corresponding format described above. - noise_schedule: A noise schedule object, such as NoiseScheduleVP. - model_type: A `str`. The parameterization type of the diffusion model. - "noise" or "x_start" or "v" or "score". - model_kwargs: A `dict`. A dict for the other inputs of the model function. - guidance_type: A `str`. The type of the guidance for sampling. - "uncond" or "classifier" or "classifier-free". - condition: A pytorch tensor. The condition for the guided sampling. - Only used for "classifier" or "classifier-free" guidance type. - unconditional_condition: A pytorch tensor. The condition for the unconditional sampling. - Only used for "classifier-free" guidance type. - guidance_scale: A `float`. The scale for the guided sampling. - classifier_fn: A classifier function. Only used for the classifier guidance. - classifier_kwargs: A `dict`. A dict for the other inputs of the classifier function. - Returns: - A noise prediction model that accepts the noised data and the continuous time as the inputs. - """ - - def get_model_input_time(t_continuous): - """ - Convert the continuous-time `t_continuous` (in [epsilon, T]) to the model input time. - For discrete-time DPMs, we convert `t_continuous` in [1 / N, 1] to `t_input` in [0, 1000 * (N - 1) / N]. - For continuous-time DPMs, we just use `t_continuous`. - """ - if noise_schedule.schedule == 'discrete': - return (t_continuous - 1. / noise_schedule.total_N) * 1000. - else: - return t_continuous - - def noise_pred_fn(x, t_continuous, cond=None): - if t_continuous.reshape((-1,)).shape[0] == 1: - t_continuous = t_continuous.expand((x.shape[0])) - t_input = get_model_input_time(t_continuous) - if cond is None: - output = model(x, t_input, None, **model_kwargs) - else: - output = model(x, t_input, cond, **model_kwargs) - if model_type == "noise": - return output - elif model_type == "x_start": - alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous) - dims = x.dim() - return (x - expand_dims(alpha_t, dims) * output) / expand_dims(sigma_t, dims) - elif model_type == "v": - alpha_t, sigma_t = noise_schedule.marginal_alpha(t_continuous), noise_schedule.marginal_std(t_continuous) - dims = x.dim() - return expand_dims(alpha_t, dims) * output + expand_dims(sigma_t, dims) * x - elif model_type == "score": - sigma_t = noise_schedule.marginal_std(t_continuous) - dims = x.dim() - return -expand_dims(sigma_t, dims) * output - - def cond_grad_fn(x, t_input): - """ - Compute the gradient of the classifier, i.e. nabla_{x} log p_t(cond | x_t). - """ - with torch.enable_grad(): - x_in = x.detach().requires_grad_(True) - log_prob = classifier_fn(x_in, t_input, condition, **classifier_kwargs) - return torch.autograd.grad(log_prob.sum(), x_in)[0] - - def model_fn(x, t_continuous): - """ - The noise predicition model function that is used for DPM-Solver. - """ - if t_continuous.reshape((-1,)).shape[0] == 1: - t_continuous = t_continuous.expand((x.shape[0])) - if guidance_type == "uncond": - return noise_pred_fn(x, t_continuous) - elif guidance_type == "classifier": - assert classifier_fn is not None - t_input = get_model_input_time(t_continuous) - cond_grad = cond_grad_fn(x, t_input) - sigma_t = noise_schedule.marginal_std(t_continuous) - noise = noise_pred_fn(x, t_continuous) - return noise - guidance_scale * expand_dims(sigma_t, dims=cond_grad.dim()) * cond_grad - elif guidance_type == "classifier-free": - if guidance_scale == 1. or unconditional_condition is None: - return noise_pred_fn(x, t_continuous, cond=condition) - else: - x_in = torch.cat([x] * 2) - t_in = torch.cat([t_continuous] * 2) - c_in = torch.cat([unconditional_condition, condition]) - noise_uncond, noise = noise_pred_fn(x_in, t_in, cond=c_in).chunk(2) - return noise_uncond + guidance_scale * (noise - noise_uncond) - - assert model_type in ["noise", "x_start", "v"] - assert guidance_type in ["uncond", "classifier", "classifier-free"] - return model_fn - - -class UniPC: - def __init__( - self, - model_fn, - noise_schedule, - predict_x0=True, - thresholding=False, - max_val=1., - variant='bh1' - ): - """Construct a UniPC. - - We support both data_prediction and noise_prediction. - """ - self.model = model_fn - self.noise_schedule = noise_schedule - self.variant = variant - self.predict_x0 = predict_x0 - self.thresholding = thresholding - self.max_val = max_val - - def dynamic_thresholding_fn(self, x0, t=None): - """ - The dynamic thresholding method. - """ - dims = x0.dim() - p = self.dynamic_thresholding_ratio - s = torch.quantile(torch.abs(x0).reshape((x0.shape[0], -1)), p, dim=1) - s = expand_dims(torch.maximum(s, self.thresholding_max_val * torch.ones_like(s).to(s.device)), dims) - x0 = torch.clamp(x0, -s, s) / s - return x0 - - def noise_prediction_fn(self, x, t): - """ - Return the noise prediction model. - """ - return self.model(x, t) - - def data_prediction_fn(self, x, t): - """ - Return the data prediction model (with thresholding). - """ - noise = self.noise_prediction_fn(x, t) - dims = x.dim() - alpha_t, sigma_t = self.noise_schedule.marginal_alpha(t), self.noise_schedule.marginal_std(t) - x0 = (x - expand_dims(sigma_t, dims) * noise) / expand_dims(alpha_t, dims) - if self.thresholding: - p = 0.995 # A hyperparameter in the paper of "Imagen" [1]. - s = torch.quantile(torch.abs(x0).reshape((x0.shape[0], -1)), p, dim=1) - s = expand_dims(torch.maximum(s, self.max_val * torch.ones_like(s).to(s.device)), dims) - x0 = torch.clamp(x0, -s, s) / s - return x0 - - def model_fn(self, x, t): - """ - Convert the model to the noise prediction model or the data prediction model. - """ - if self.predict_x0: - return self.data_prediction_fn(x, t) - else: - return self.noise_prediction_fn(x, t) - - def get_time_steps(self, skip_type, t_T, t_0, N, device): - """Compute the intermediate time steps for sampling. - """ - if skip_type == 'logSNR': - lambda_T = self.noise_schedule.marginal_lambda(torch.tensor(t_T).to(device)) - lambda_0 = self.noise_schedule.marginal_lambda(torch.tensor(t_0).to(device)) - logSNR_steps = torch.linspace(lambda_T.cpu().item(), lambda_0.cpu().item(), N + 1).to(device) - return self.noise_schedule.inverse_lambda(logSNR_steps) - elif skip_type == 'time_uniform': - return torch.linspace(t_T, t_0, N + 1).to(device) - elif skip_type == 'time_quadratic': - t_order = 2 - t = torch.linspace(t_T**(1. / t_order), t_0**(1. / t_order), N + 1).pow(t_order).to(device) - return t - else: - raise ValueError("Unsupported skip_type {}, need to be 'logSNR' or 'time_uniform' or 'time_quadratic'".format(skip_type)) - - def get_orders_and_timesteps_for_singlestep_solver(self, steps, order, skip_type, t_T, t_0, device): - """ - Get the order of each step for sampling by the singlestep DPM-Solver. - """ - if order == 3: - K = steps // 3 + 1 - if steps % 3 == 0: - orders = [3,] * (K - 2) + [2, 1] - elif steps % 3 == 1: - orders = [3,] * (K - 1) + [1] - else: - orders = [3,] * (K - 1) + [2] - elif order == 2: - if steps % 2 == 0: - K = steps // 2 - orders = [2,] * K - else: - K = steps // 2 + 1 - orders = [2,] * (K - 1) + [1] - elif order == 1: - K = steps - orders = [1,] * steps - else: - raise ValueError("'order' must be '1' or '2' or '3'.") - if skip_type == 'logSNR': - # To reproduce the results in DPM-Solver paper - timesteps_outer = self.get_time_steps(skip_type, t_T, t_0, K, device) - else: - timesteps_outer = self.get_time_steps(skip_type, t_T, t_0, steps, device)[torch.cumsum(torch.tensor([0,] + orders), 0).to(device)] - return timesteps_outer, orders - - def denoise_to_zero_fn(self, x, s): - """ - Denoise at the final step, which is equivalent to solve the ODE from lambda_s to infty by first-order discretization. - """ - return self.data_prediction_fn(x, s) - - def multistep_uni_pc_update(self, x, model_prev_list, t_prev_list, t, order, **kwargs): - if len(t.shape) == 0: - t = t.view(-1) - if 'bh' in self.variant: - return self.multistep_uni_pc_bh_update(x, model_prev_list, t_prev_list, t, order, **kwargs) - else: - assert self.variant == 'vary_coeff' - return self.multistep_uni_pc_vary_update(x, model_prev_list, t_prev_list, t, order, **kwargs) - - def multistep_uni_pc_vary_update(self, x, model_prev_list, t_prev_list, t, order, use_corrector=True): - print(f'using unified predictor-corrector with order {order} (solver type: vary coeff)') - ns = self.noise_schedule - assert order <= len(model_prev_list) - - # first compute rks - t_prev_0 = t_prev_list[-1] - lambda_prev_0 = ns.marginal_lambda(t_prev_0) - lambda_t = ns.marginal_lambda(t) - model_prev_0 = model_prev_list[-1] - sigma_prev_0, sigma_t = ns.marginal_std(t_prev_0), ns.marginal_std(t) - log_alpha_t = ns.marginal_log_mean_coeff(t) - alpha_t = torch.exp(log_alpha_t) - - h = lambda_t - lambda_prev_0 - - rks = [] - D1s = [] - for i in range(1, order): - t_prev_i = t_prev_list[-(i + 1)] - model_prev_i = model_prev_list[-(i + 1)] - lambda_prev_i = ns.marginal_lambda(t_prev_i) - rk = (lambda_prev_i - lambda_prev_0) / h - rks.append(rk) - D1s.append((model_prev_i - model_prev_0) / rk) - - rks.append(1.) - rks = torch.tensor(rks, device=x.device) - - K = len(rks) - # build C matrix - C = [] - - col = torch.ones_like(rks) - for k in range(1, K + 1): - C.append(col) - col = col * rks / (k + 1) - C = torch.stack(C, dim=1) - - if len(D1s) > 0: - D1s = torch.stack(D1s, dim=1) # (B, K) - C_inv_p = torch.linalg.inv(C[:-1, :-1]) - A_p = C_inv_p - - if use_corrector: - print('using corrector') - C_inv = torch.linalg.inv(C) - A_c = C_inv - - hh = -h if self.predict_x0 else h - h_phi_1 = torch.expm1(hh) - h_phi_ks = [] - factorial_k = 1 - h_phi_k = h_phi_1 - for k in range(1, K + 2): - h_phi_ks.append(h_phi_k) - h_phi_k = h_phi_k / hh - 1 / factorial_k - factorial_k *= (k + 1) - - model_t = None - if self.predict_x0: - x_t_ = ( - sigma_t / sigma_prev_0 * x - - alpha_t * h_phi_1 * model_prev_0 - ) - # now predictor - x_t = x_t_ - if len(D1s) > 0: - # compute the residuals for predictor - for k in range(K - 1): - x_t = x_t - alpha_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_p[k]) - # now corrector - if use_corrector: - model_t = self.model_fn(x_t, t) - D1_t = (model_t - model_prev_0) - x_t = x_t_ - k = 0 - for k in range(K - 1): - x_t = x_t - alpha_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_c[k][:-1]) - x_t = x_t - alpha_t * h_phi_ks[K] * (D1_t * A_c[k][-1]) - else: - log_alpha_prev_0, log_alpha_t = ns.marginal_log_mean_coeff(t_prev_0), ns.marginal_log_mean_coeff(t) - x_t_ = ( - (torch.exp(log_alpha_t - log_alpha_prev_0)) * x - - (sigma_t * h_phi_1) * model_prev_0 - ) - # now predictor - x_t = x_t_ - if len(D1s) > 0: - # compute the residuals for predictor - for k in range(K - 1): - x_t = x_t - sigma_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_p[k]) - # now corrector - if use_corrector: - model_t = self.model_fn(x_t, t) - D1_t = (model_t - model_prev_0) - x_t = x_t_ - k = 0 - for k in range(K - 1): - x_t = x_t - sigma_t * h_phi_ks[k + 1] * torch.einsum('bkchw,k->bchw', D1s, A_c[k][:-1]) - x_t = x_t - sigma_t * h_phi_ks[K] * (D1_t * A_c[k][-1]) - return x_t, model_t - - def multistep_uni_pc_bh_update(self, x, model_prev_list, t_prev_list, t, order, x_t=None, use_corrector=True): - print(f'using unified predictor-corrector with order {order} (solver type: B(h))') - ns = self.noise_schedule - assert order <= len(model_prev_list) - dims = x.dim() - - # first compute rks - t_prev_0 = t_prev_list[-1] - lambda_prev_0 = ns.marginal_lambda(t_prev_0) - lambda_t = ns.marginal_lambda(t) - model_prev_0 = model_prev_list[-1] - sigma_prev_0, sigma_t = ns.marginal_std(t_prev_0), ns.marginal_std(t) - log_alpha_prev_0, log_alpha_t = ns.marginal_log_mean_coeff(t_prev_0), ns.marginal_log_mean_coeff(t) - alpha_t = torch.exp(log_alpha_t) - - h = lambda_t - lambda_prev_0 - - rks = [] - D1s = [] - for i in range(1, order): - t_prev_i = t_prev_list[-(i + 1)] - model_prev_i = model_prev_list[-(i + 1)] - lambda_prev_i = ns.marginal_lambda(t_prev_i) - rk = ((lambda_prev_i - lambda_prev_0) / h)[0] - rks.append(rk) - D1s.append((model_prev_i - model_prev_0) / rk) - - rks.append(1.) - rks = torch.tensor(rks, device=x.device) - - R = [] - b = [] - - hh = -h[0] if self.predict_x0 else h[0] - h_phi_1 = torch.expm1(hh) # h\phi_1(h) = e^h - 1 - h_phi_k = h_phi_1 / hh - 1 - - factorial_i = 1 - - if self.variant == 'bh1': - B_h = hh - elif self.variant == 'bh2': - B_h = torch.expm1(hh) - else: - raise NotImplementedError() - - for i in range(1, order + 1): - R.append(torch.pow(rks, i - 1)) - b.append(h_phi_k * factorial_i / B_h) - factorial_i *= (i + 1) - h_phi_k = h_phi_k / hh - 1 / factorial_i - - R = torch.stack(R) - b = torch.tensor(b, device=x.device) - - # now predictor - use_predictor = len(D1s) > 0 and x_t is None - if len(D1s) > 0: - D1s = torch.stack(D1s, dim=1) # (B, K) - if x_t is None: - # for order 2, we use a simplified version - if order == 2: - rhos_p = torch.tensor([0.5], device=b.device) - else: - rhos_p = torch.linalg.solve(R[:-1, :-1], b[:-1]) - else: - D1s = None - - if use_corrector: - print('using corrector') - # for order 1, we use a simplified version - if order == 1: - rhos_c = torch.tensor([0.5], device=b.device) - else: - rhos_c = torch.linalg.solve(R, b) - - model_t = None - if self.predict_x0: - x_t_ = ( - expand_dims(sigma_t / sigma_prev_0, dims) * x - - expand_dims(alpha_t * h_phi_1, dims)* model_prev_0 - ) - - if x_t is None: - if use_predictor: - pred_res = torch.einsum('k,bkchw->bchw', rhos_p, D1s) - else: - pred_res = 0 - x_t = x_t_ - expand_dims(alpha_t * B_h, dims) * pred_res - - if use_corrector: - model_t = self.model_fn(x_t, t) - if D1s is not None: - corr_res = torch.einsum('k,bkchw->bchw', rhos_c[:-1], D1s) - else: - corr_res = 0 - D1_t = (model_t - model_prev_0) - x_t = x_t_ - expand_dims(alpha_t * B_h, dims) * (corr_res + rhos_c[-1] * D1_t) - else: - x_t_ = ( - expand_dims(torch.exp(log_alpha_t - log_alpha_prev_0), dims) * x - - expand_dims(sigma_t * h_phi_1, dims) * model_prev_0 - ) - if x_t is None: - if use_predictor: - pred_res = torch.einsum('k,bkchw->bchw', rhos_p, D1s) - else: - pred_res = 0 - x_t = x_t_ - expand_dims(sigma_t * B_h, dims) * pred_res - - if use_corrector: - model_t = self.model_fn(x_t, t) - if D1s is not None: - corr_res = torch.einsum('k,bkchw->bchw', rhos_c[:-1], D1s) - else: - corr_res = 0 - D1_t = (model_t - model_prev_0) - x_t = x_t_ - expand_dims(sigma_t * B_h, dims) * (corr_res + rhos_c[-1] * D1_t) - return x_t, model_t - - - def sample(self, x, steps=20, t_start=None, t_end=None, order=3, skip_type='time_uniform', - method='singlestep', lower_order_final=True, denoise_to_zero=False, solver_type='dpm_solver', - atol=0.0078, rtol=0.05, corrector=False, - ): - t_0 = 1. / self.noise_schedule.total_N if t_end is None else t_end - t_T = self.noise_schedule.T if t_start is None else t_start - device = x.device - if method == 'multistep': - assert steps >= order - timesteps = self.get_time_steps(skip_type=skip_type, t_T=t_T, t_0=t_0, N=steps, device=device) - assert timesteps.shape[0] - 1 == steps - with torch.no_grad(): - vec_t = timesteps[0].expand((x.shape[0])) - model_prev_list = [self.model_fn(x, vec_t)] - t_prev_list = [vec_t] - # Init the first `order` values by lower order multistep DPM-Solver. - for init_order in range(1, order): - vec_t = timesteps[init_order].expand(x.shape[0]) - x, model_x = self.multistep_uni_pc_update(x, model_prev_list, t_prev_list, vec_t, init_order, use_corrector=True) - if model_x is None: - model_x = self.model_fn(x, vec_t) - model_prev_list.append(model_x) - t_prev_list.append(vec_t) - for step in range(order, steps + 1): - vec_t = timesteps[step].expand(x.shape[0]) - if lower_order_final: - step_order = min(order, steps + 1 - step) - else: - step_order = order - print('this step order:', step_order) - if step == steps: - print('do not run corrector at the last step') - use_corrector = False - else: - use_corrector = True - x, model_x = self.multistep_uni_pc_update(x, model_prev_list, t_prev_list, vec_t, step_order, use_corrector=use_corrector) - for i in range(order - 1): - t_prev_list[i] = t_prev_list[i + 1] - model_prev_list[i] = model_prev_list[i + 1] - t_prev_list[-1] = vec_t - # We do not need to evaluate the final model value. - if step < steps: - if model_x is None: - model_x = self.model_fn(x, vec_t) - model_prev_list[-1] = model_x - else: - raise NotImplementedError() - if denoise_to_zero: - x = self.denoise_to_zero_fn(x, torch.ones((x.shape[0],)).to(device) * t_0) - return x - - -############################################################# -# other utility functions -############################################################# - -def interpolate_fn(x, xp, yp): - """ - A piecewise linear function y = f(x), using xp and yp as keypoints. - We implement f(x) in a differentiable way (i.e. applicable for autograd). - The function f(x) is well-defined for all x-axis. (For x beyond the bounds of xp, we use the outmost points of xp to define the linear function.) - - Args: - x: PyTorch tensor with shape [N, C], where N is the batch size, C is the number of channels (we use C = 1 for DPM-Solver). - xp: PyTorch tensor with shape [C, K], where K is the number of keypoints. - yp: PyTorch tensor with shape [C, K]. - Returns: - The function values f(x), with shape [N, C]. - """ - N, K = x.shape[0], xp.shape[1] - all_x = torch.cat([x.unsqueeze(2), xp.unsqueeze(0).repeat((N, 1, 1))], dim=2) - sorted_all_x, x_indices = torch.sort(all_x, dim=2) - x_idx = torch.argmin(x_indices, dim=2) - cand_start_idx = x_idx - 1 - start_idx = torch.where( - torch.eq(x_idx, 0), - torch.tensor(1, device=x.device), - torch.where( - torch.eq(x_idx, K), torch.tensor(K - 2, device=x.device), cand_start_idx, - ), - ) - end_idx = torch.where(torch.eq(start_idx, cand_start_idx), start_idx + 2, start_idx + 1) - start_x = torch.gather(sorted_all_x, dim=2, index=start_idx.unsqueeze(2)).squeeze(2) - end_x = torch.gather(sorted_all_x, dim=2, index=end_idx.unsqueeze(2)).squeeze(2) - start_idx2 = torch.where( - torch.eq(x_idx, 0), - torch.tensor(0, device=x.device), - torch.where( - torch.eq(x_idx, K), torch.tensor(K - 2, device=x.device), cand_start_idx, - ), - ) - y_positions_expanded = yp.unsqueeze(0).expand(N, -1, -1) - start_y = torch.gather(y_positions_expanded, dim=2, index=start_idx2.unsqueeze(2)).squeeze(2) - end_y = torch.gather(y_positions_expanded, dim=2, index=(start_idx2 + 1).unsqueeze(2)).squeeze(2) - cand = start_y + (x - start_x) * (end_y - start_y) / (end_x - start_x) - return cand - - -def expand_dims(v, dims): - """ - Expand the tensor `v` to the dim `dims`. - - Args: - `v`: a PyTorch tensor with shape [N]. - `dim`: a `int`. - Returns: - a PyTorch tensor with shape [N, 1, 1, ..., 1] and the total dimension is `dims`. - """ - return v[(...,) + (None,)*(dims - 1)] \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/models/utils_diffusion.py b/py/dynamiCrafter/lvdm/models/utils_diffusion.py deleted file mode 100644 index 403b7b3..0000000 --- a/py/dynamiCrafter/lvdm/models/utils_diffusion.py +++ /dev/null @@ -1,158 +0,0 @@ -import math -import numpy as np -import torch -import torch.nn.functional as F -from einops import repeat - - -def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False, dtype=None): - """ - Create sinusoidal timestep embeddings. - :param timesteps: a 1-D Tensor of N indices, one per batch element. - These may be fractional. - :param dim: the dimension of the output. - :param max_period: controls the minimum frequency of the embeddings. - :return: an [N x dim] Tensor of positional embeddings. - """ - if not repeat_only: - half = dim // 2 - freqs = torch.exp( - -math.log(max_period) * torch.arange(start=0, end=half, dtype=dtype) / half - ).to(device=timesteps.device) - args = timesteps[:, None].float() * freqs[None] - embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) - if dim % 2: - embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) - else: - embedding = repeat(timesteps, 'b -> b d', d=dim) - return embedding.to(dtype) - - -def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3): - if schedule == "linear": - betas = ( - torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=torch.float64) ** 2 - ) - - elif schedule == "cosine": - timesteps = ( - torch.arange(n_timestep + 1, dtype=torch.float64) / n_timestep + cosine_s - ) - alphas = timesteps / (1 + cosine_s) * np.pi / 2 - alphas = torch.cos(alphas).pow(2) - alphas = alphas / alphas[0] - betas = 1 - alphas[1:] / alphas[:-1] - 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) - elif schedule == "sqrt": - betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64) ** 0.5 - else: - raise ValueError(f"schedule '{schedule}' unknown.") - return betas.numpy() - - -def make_ddim_timesteps(ddim_discr_method, num_ddim_timesteps, num_ddpm_timesteps, verbose=True): - if ddim_discr_method == 'uniform': - c = num_ddpm_timesteps // num_ddim_timesteps - ddim_timesteps = np.asarray(list(range(0, num_ddpm_timesteps, c))) - steps_out = ddim_timesteps + 1 - elif ddim_discr_method == 'uniform_trailing': - c = num_ddpm_timesteps / num_ddim_timesteps - ddim_timesteps = np.flip(np.round(np.arange(num_ddpm_timesteps, 0, -c))).astype(np.int64) - steps_out = ddim_timesteps - 1 - elif ddim_discr_method == 'quad': - ddim_timesteps = ((np.linspace(0, np.sqrt(num_ddpm_timesteps * .8), num_ddim_timesteps)) ** 2).astype(int) - steps_out = ddim_timesteps + 1 - else: - raise NotImplementedError(f'There is no ddim discretization method called "{ddim_discr_method}"') - - # assert ddim_timesteps.shape[0] == num_ddim_timesteps - # add one to get the final alpha values right (the ones from first scale to data during sampling) - # steps_out = ddim_timesteps + 1 - if verbose: - print(f'Selected timesteps for ddim sampler: {steps_out}') - return steps_out - - -def make_ddim_sampling_parameters(alphacums, ddim_timesteps, eta, verbose=True): - # select alphas for computing the variance schedule - # print(f'ddim_timesteps={ddim_timesteps}, len_alphacums={len(alphacums)}') - alphas = alphacums[ddim_timesteps] - alphas_prev = np.asarray([alphacums[0]] + alphacums[ddim_timesteps[:-1]].tolist()) - - # according the the formula provided in https://arxiv.org/abs/2010.02502 - sigmas = eta * np.sqrt((1 - alphas_prev) / (1 - alphas) * (1 - alphas / alphas_prev)) - if verbose: - print(f'Selected alphas for ddim sampler: a_t: {alphas}; a_(t-1): {alphas_prev}') - print(f'For the chosen value of eta, which is {eta}, ' - f'this results in the following sigma_t schedule for ddim sampler {sigmas}') - return sigmas, alphas, alphas_prev - - -def betas_for_alpha_bar(num_diffusion_timesteps, alpha_bar, max_beta=0.999): - """ - Create a beta schedule that discretizes the given alpha_t_bar function, - which defines the cumulative product of (1-beta) over time from t = [0,1]. - :param num_diffusion_timesteps: the number of betas to produce. - :param alpha_bar: a lambda that takes an argument t from 0 to 1 and - produces the cumulative product of (1-beta) up to that - part of the diffusion process. - :param max_beta: the maximum beta to use; use values lower than 1 to - prevent singularities. - """ - betas = [] - for i in range(num_diffusion_timesteps): - t1 = i / num_diffusion_timesteps - t2 = (i + 1) / num_diffusion_timesteps - betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta)) - return np.array(betas) - -def rescale_zero_terminal_snr(betas): - """ - Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1) - - Args: - betas (`numpy.ndarray`): - the betas that the scheduler is being initialized with. - - Returns: - `numpy.ndarray`: rescaled betas with zero terminal SNR - """ - # Convert betas to alphas_bar_sqrt - alphas = 1.0 - betas - alphas_cumprod = np.cumprod(alphas, axis=0) - alphas_bar_sqrt = np.sqrt(alphas_cumprod) - - # Store old values. - alphas_bar_sqrt_0 = alphas_bar_sqrt[0].copy() - alphas_bar_sqrt_T = alphas_bar_sqrt[-1].copy() - - # Shift so the last timestep is zero. - alphas_bar_sqrt -= alphas_bar_sqrt_T - - # Scale so the first timestep is back to the old value. - alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T) - - # Convert alphas_bar_sqrt to betas - alphas_bar = alphas_bar_sqrt**2 # Revert sqrt - alphas = alphas_bar[1:] / alphas_bar[:-1] # Revert cumprod - alphas = np.concatenate([alphas_bar[0:1], alphas]) - betas = 1 - alphas - - return betas - - -def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0): - """ - Rescale `noise_cfg` according to `guidance_rescale`. Based on findings of [Common Diffusion Noise Schedules and - Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). See Section 3.4 - """ - std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True) - std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True) - # rescale the results from guidance (fixes overexposure) - noise_pred_rescaled = noise_cfg * (std_text / std_cfg) - # mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images - noise_cfg = guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg - return noise_cfg \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/modules/attention.py b/py/dynamiCrafter/lvdm/modules/attention.py deleted file mode 100644 index a136029..0000000 --- a/py/dynamiCrafter/lvdm/modules/attention.py +++ /dev/null @@ -1,809 +0,0 @@ -import torch -from torch import nn, einsum -import torch.nn.functional as F -from einops import rearrange, repeat -from functools import partial -from ..common import ( - checkpoint, - exists, - default, -) -from ..basics import zero_module -import comfy.ops -ops = comfy.ops.disable_weight_init -from comfy import model_management -from comfy.ldm.modules.attention import optimized_attention, optimized_attention_masked - -if model_management.xformers_enabled(): - import xformers - import xformers.ops - XFORMERS_IS_AVAILBLE = True -else: - XFORMERS_IS_AVAILBLE = False - -class RelativePosition(nn.Module): - """ https://github.com/evelinehong/Transformer_Relative_Position_PyTorch/blob/master/relative_position.py """ - - def __init__(self, num_units, max_relative_position): - super().__init__() - self.num_units = num_units - self.max_relative_position = max_relative_position - self.embeddings_table = nn.Parameter(torch.Tensor(max_relative_position * 2 + 1, num_units)) - nn.init.xavier_uniform_(self.embeddings_table) - - def forward(self, length_q, length_k): - device = self.embeddings_table.device - range_vec_q = torch.arange(length_q, device=device) - range_vec_k = torch.arange(length_k, device=device) - distance_mat = range_vec_k[None, :] - range_vec_q[:, None] - distance_mat_clipped = torch.clamp(distance_mat, -self.max_relative_position, self.max_relative_position) - final_mat = distance_mat_clipped + self.max_relative_position - final_mat = final_mat.long() - embeddings = self.embeddings_table[final_mat] - return embeddings - - -# TODO Add native Comfy optimized attention. -class CrossAttention(nn.Module): - - def __init__( - self, - query_dim, - context_dim=None, - heads=8, - dim_head=64, - dropout=0., - relative_position=False, - temporal_length=None, - video_length=None, - image_cross_attention=False, - image_cross_attention_scale=1.0, - image_cross_attention_scale_learnable=False, - text_context_len=77, - device=None, - dtype=None, - operations=ops - ): - super().__init__() - inner_dim = dim_head * heads - context_dim = default(context_dim, query_dim) - self.scale = dim_head**-0.5 - self.heads = heads - self.dim_head = dim_head - self.to_q = operations.Linear(query_dim, inner_dim, bias=False, device=device, dtype=dtype) - self.to_k = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype) - self.to_v = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype) - - self.to_out = nn.Sequential( - operations.Linear(inner_dim, query_dim, device=device, dtype=dtype), - nn.Dropout(dropout) - ) - - self.relative_position = relative_position - if self.relative_position: - assert(temporal_length is not None) - self.relative_position_k = RelativePosition(num_units=dim_head, max_relative_position=temporal_length) - self.relative_position_v = RelativePosition(num_units=dim_head, max_relative_position=temporal_length) - else: - ## only used for spatial attention, while NOT for temporal attention - if XFORMERS_IS_AVAILBLE and temporal_length is None: - self.forward = self.efficient_forward - else: - self.forward = self.comfy_efficient_forward - - self.video_length = video_length - self.image_cross_attention = image_cross_attention - self.image_cross_attention_scale = image_cross_attention_scale - self.text_context_len = text_context_len - self.image_cross_attention_scale_learnable = image_cross_attention_scale_learnable - if self.image_cross_attention: - self.to_k_ip = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype) - self.to_v_ip = operations.Linear(context_dim, inner_dim, bias=False, device=device, dtype=dtype) - if image_cross_attention_scale_learnable: - self.register_parameter('alpha', nn.Parameter(torch.tensor(0.)) ) - - def comfy_efficient_forward(self, x, context=None, mask=None, *args, **kwargs): - spatial_self_attn = (context is None) - k_ip, v_ip, out_ip = None, None, None - - h = self.heads - q = self.to_q(x) - context = default(context, x) - - if self.image_cross_attention and not spatial_self_attn: - context, context_image = context[:,:self.text_context_len,:], context[:,self.text_context_len:,:] - k = self.to_k(context) - v = self.to_v(context) - k_ip = self.to_k_ip(context_image) - v_ip = self.to_v_ip(context_image) - else: - if not spatial_self_attn: - context = context[:,:self.text_context_len,:] - k = self.to_k(context) - v = self.to_v(context) - - out = optimized_attention(q, k, v, h) - - if exists(mask): - ## feasible for causal attention mask only - out = optimized_attention_masked(q, k, v, h) - - ## for image cross-attention - if k_ip is not None: - q = rearrange(q, 'b n (h d) -> (b h) n d', h=h) - k_ip, v_ip = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (k_ip, v_ip)) - sim_ip = torch.einsum('b i d, b j d -> b i j', q, k_ip) * self.scale - del k_ip - sim_ip = sim_ip.softmax(dim=-1) - out_ip = torch.einsum('b i j, b j d -> b i d', sim_ip, v_ip) - out_ip = rearrange(out_ip, '(b h) n d -> b n (h d)', h=h) - - if out_ip is not None: - if self.image_cross_attention_scale_learnable: - out = out + self.image_cross_attention_scale * out_ip * (torch.tanh(self.alpha)+1) - else: - out = out + self.image_cross_attention_scale * out_ip - - return self.to_out(out) - - def forward(self, x, context=None, mask=None): - spatial_self_attn = (context is None) - k_ip, v_ip, out_ip = None, None, None - - h = self.heads - q = self.to_q(x) - context = default(context, x) - - if self.image_cross_attention and not spatial_self_attn: - context, context_image = context[:,:self.text_context_len,:], context[:,self.text_context_len:,:] - k = self.to_k(context) - v = self.to_v(context) - k_ip = self.to_k_ip(context_image) - v_ip = self.to_v_ip(context_image) - else: - - # Assumed Spatial Attention (b c h w) - if not spatial_self_attn: - context = context[:,:self.text_context_len,:] - k = self.to_k(context) - v = self.to_v(context) - - - q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (q, k, v)) - - sim = torch.einsum('b i d, b j d -> b i j', q, k) * self.scale - if self.relative_position: - len_q, len_k, len_v = q.shape[1], k.shape[1], v.shape[1] - k2 = self.relative_position_k(len_q, len_k) - sim2 = einsum('b t d, t s d -> b t s', q, k2) * self.scale # TODO check - sim += sim2 - del k - - if exists(mask): - ## feasible for causal attention mask only - max_neg_value = -torch.finfo(sim.dtype).max - mask = repeat(mask, 'b i j -> (b h) i j', h=h) - sim.masked_fill_(~(mask>0.5), max_neg_value) - - # attention, what we cannot get enough of - sim = sim.softmax(dim=-1) - - out = torch.einsum('b i j, b j d -> b i d', sim, v) - if self.relative_position: - v2 = self.relative_position_v(len_q, len_v) - out2 = einsum('b t s, t s d -> b t d', sim, v2) # TODO check - out += out2 - out = rearrange(out, '(b h) n d -> b n (h d)', h=h) - - - ## for image cross-attention - if k_ip is not None: - k_ip, v_ip = map(lambda t: rearrange(t, 'b n (h d) -> (b h) n d', h=h), (k_ip, v_ip)) - sim_ip = torch.einsum('b i d, b j d -> b i j', q, k_ip) * self.scale - del k_ip - sim_ip = sim_ip.softmax(dim=-1) - out_ip = torch.einsum('b i j, b j d -> b i d', sim_ip, v_ip) - out_ip = rearrange(out_ip, '(b h) n d -> b n (h d)', h=h) - - - if out_ip is not None: - if self.image_cross_attention_scale_learnable: - out = out + self.image_cross_attention_scale * out_ip * (torch.tanh(self.alpha)+1) - else: - out = out + self.image_cross_attention_scale * out_ip - - return self.to_out(out) - - def efficient_forward(self, x, context=None, mask=None): - spatial_self_attn = (context is None) - k_ip, v_ip, out_ip = None, None, None - - q = self.to_q(x) - context = default(context, x) - - if self.image_cross_attention and not spatial_self_attn: - context, context_image = context[:,:self.text_context_len,:], context[:,self.text_context_len:,:] - k = self.to_k(context) - v = self.to_v(context) - k_ip = self.to_k_ip(context_image) - v_ip = self.to_v_ip(context_image) - else: - if not spatial_self_attn: - context = context[:,:self.text_context_len,:] - k = self.to_k(context) - v = self.to_v(context) - - b, _, _ = q.shape - q, k, v = map( - lambda t: t.unsqueeze(3) - .reshape(b, t.shape[1], self.heads, self.dim_head) - .permute(0, 2, 1, 3) - .reshape(b * self.heads, t.shape[1], self.dim_head) - .contiguous(), - (q, k, v), - ) - # actually compute the attention, what we cannot get enough of - out = xformers.ops.memory_efficient_attention(q, k, v, attn_bias=None, op=None) - - ## for image cross-attention - if k_ip is not None: - k_ip, v_ip = map( - lambda t: t.unsqueeze(3) - .reshape(b, t.shape[1], self.heads, self.dim_head) - .permute(0, 2, 1, 3) - .reshape(b * self.heads, t.shape[1], self.dim_head) - .contiguous(), - (k_ip, v_ip), - ) - out_ip = xformers.ops.memory_efficient_attention(q, k_ip, v_ip, attn_bias=None, op=None) - out_ip = ( - out_ip.unsqueeze(0) - .reshape(b, self.heads, out.shape[1], self.dim_head) - .permute(0, 2, 1, 3) - .reshape(b, out.shape[1], self.heads * self.dim_head) - ) - - if exists(mask): - raise NotImplementedError - out = ( - out.unsqueeze(0) - .reshape(b, self.heads, out.shape[1], self.dim_head) - .permute(0, 2, 1, 3) - .reshape(b, out.shape[1], self.heads * self.dim_head) - ) - if out_ip is not None: - if self.image_cross_attention_scale_learnable: - out = out + self.image_cross_attention_scale * out_ip * (torch.tanh(self.alpha)+1) - else: - out = out + self.image_cross_attention_scale * out_ip - - return self.to_out(out) - - -class BasicTransformerBlock(nn.Module): - - def __init__( - self, - dim, - n_heads, - d_head, - dropout=0., - context_dim=None, - gated_ff=True, - checkpoint=True, - disable_self_attn=False, - attention_cls=None, - video_length=None, - inner_dim=None, - image_cross_attention=False, - image_cross_attention_scale=1.0, - image_cross_attention_scale_learnable=False, - switch_temporal_ca_to_sa=False, - text_context_len=77, - ff_in=None, - device=None, - dtype=None, - operations=ops - ): - super().__init__() - attn_cls = CrossAttention if attention_cls is None else attention_cls - - self.ff_in = ff_in or inner_dim is not None - if self.ff_in: - self.norm_in = operations.LayerNorm(dim, dtype=dtype, device=device) - self.ff_in = FeedForward( - dim, - dim_out=inner_dim, - dropout=dropout, - glu=gated_ff, - dtype=dtype, - device=device, - operations=operations - ) - if inner_dim is None: - inner_dim = dim - - self.is_res = inner_dim == dim - self.disable_self_attn = disable_self_attn - self.attn1 = attn_cls(query_dim=dim, heads=n_heads, dim_head=d_head, dropout=dropout, - context_dim=None, device=device, dtype=dtype if self.disable_self_attn else None) - self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff, device=device, dtype=dtype) - self.attn2 = attn_cls( - query_dim=dim, - context_dim=context_dim, - heads=n_heads, - dim_head=d_head, - dropout=dropout, - video_length=video_length, - image_cross_attention=image_cross_attention, - image_cross_attention_scale=image_cross_attention_scale, - image_cross_attention_scale_learnable=image_cross_attention_scale_learnable, - text_context_len=text_context_len, - device=device, - dtype=dtype - ) - self.image_cross_attention = image_cross_attention - - self.norm1 = operations.LayerNorm(dim, device=device, dtype=dtype) - self.norm2 = operations.LayerNorm(dim, device=device, dtype=dtype) - self.norm3 = operations.LayerNorm(dim, device=device, dtype=dtype) - - self.n_heads = n_heads - self.d_head = d_head - self.checkpoint = checkpoint - self.switch_temporal_ca_to_sa = switch_temporal_ca_to_sa - - def forward(self, x, context=None, mask=None, **kwargs): - ## implementation tricks: because checkpointing doesn't support non-tensor (e.g. None or scalar) arguments - input_tuple = (x,) ## should not be (x), otherwise *input_tuple will decouple x into multiple arguments - if context is not None: - input_tuple = (x, context) - if mask is not None: - forward_mask = partial(self._forward, mask=mask) - return checkpoint(forward_mask, (x,), self.parameters(), self.checkpoint) - return checkpoint(self._forward, input_tuple, self.parameters(), self.checkpoint) - - - def _forward(self, x, context=None, mask=None, transformer_options={}): - extra_options = {} - block = transformer_options.get("block", None) - block_index = transformer_options.get("block_index", 0) - transformer_patches = {} - transformer_patches_replace = {} - - for k in transformer_options: - if k == "patches": - transformer_patches = transformer_options[k] - elif k == "patches_replace": - transformer_patches_replace = transformer_options[k] - else: - extra_options[k] = transformer_options[k] - - extra_options["n_heads"] = self.n_heads - extra_options["dim_head"] = self.d_head - - if self.ff_in: - x_skip = x - x = self.ff_in(self.norm_in(x)) - if self.is_res: - x += x_skip - - n = self.norm1(x) - if self.disable_self_attn: - context_attn1 = context - else: - context_attn1 = None - value_attn1 = None - - if "attn1_patch" in transformer_patches: - patch = transformer_patches["attn1_patch"] - if context_attn1 is None: - context_attn1 = n - value_attn1 = context_attn1 - for p in patch: - n, context_attn1, value_attn1 = p(n, context_attn1, value_attn1, extra_options) - - if block is not None: - transformer_block = (block[0], block[1], block_index) - else: - transformer_block = None - attn1_replace_patch = transformer_patches_replace.get("attn1", {}) - block_attn1 = transformer_block - if block_attn1 not in attn1_replace_patch: - block_attn1 = block - - if block_attn1 in attn1_replace_patch: - if context_attn1 is None: - context_attn1 = n - value_attn1 = n - n = self.attn1.to_q(n) - context_attn1 = self.attn1.to_k(context_attn1) - value_attn1 = self.attn1.to_v(value_attn1) - n = attn1_replace_patch[block_attn1](n, context_attn1, value_attn1, extra_options) - n = self.attn1.to_out(n) - else: - n = self.attn1(n, context=context_attn1, value=value_attn1) - - if "attn1_output_patch" in transformer_patches: - patch = transformer_patches["attn1_output_patch"] - for p in patch: - n = p(n, extra_options) - - x += n - if "middle_patch" in transformer_patches: - patch = transformer_patches["middle_patch"] - for p in patch: - x = p(x, extra_options) - - if self.attn2 is not None: - n = self.norm2(x) - if self.switch_temporal_ca_to_sa: - context_attn2 = n - else: - context_attn2 = context - value_attn2 = None - if "attn2_patch" in transformer_patches: - patch = transformer_patches["attn2_patch"] - value_attn2 = context_attn2 - for p in patch: - n, context_attn2, value_attn2 = p(n, context_attn2, value_attn2, extra_options) - - attn2_replace_patch = transformer_patches_replace.get("attn2", {}) - block_attn2 = transformer_block - if block_attn2 not in attn2_replace_patch: - block_attn2 = block - - if block_attn2 in attn2_replace_patch: - if value_attn2 is None: - value_attn2 = context_attn2 - n = self.attn2.to_q(n) - context_attn2 = self.attn2.to_k(context_attn2) - value_attn2 = self.attn2.to_v(value_attn2) - n = attn2_replace_patch[block_attn2](n, context_attn2, value_attn2, extra_options) - n = self.attn2.to_out(n) - else: - n = self.attn2(n, context=context_attn2, value=value_attn2) - - if "attn2_output_patch" in transformer_patches: - patch = transformer_patches["attn2_output_patch"] - for p in patch: - n = p(n, extra_options) - - x += n - if self.is_res: - x_skip = x - x = self.ff(self.norm3(x)) - if self.is_res: - x += x_skip - - return x - - -class SpatialTransformer(nn.Module): - """ - Transformer block for image-like data in spatial axis. - First, project the input (aka embedding) - and reshape to b, t, d. - Then apply standard transformer action. - Finally, reshape to image - NEW: use_linear for more efficiency instead of the 1x1 convs - """ - - def __init__( - self, - in_channels, - n_heads, - d_head, - depth=1, - dropout=0., - context_dim=None, - use_checkpoint=True, - disable_self_attn=False, - use_linear=False, - video_length=None, - image_cross_attention=False, - image_cross_attention_scale_learnable=False, - device=None, - dtype=None, - operations=ops - ): - super().__init__() - self.in_channels = in_channels - inner_dim = n_heads * d_head - self.norm = operations.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True, device=device, dtype=dtype) - if not use_linear: - self.proj_in = opeations.Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0, device=device, dtype=dtype) - else: - self.proj_in = operations.Linear(in_channels, inner_dim, device=device, dtype=dtype) - - attention_cls = None - self.transformer_blocks = nn.ModuleList([ - BasicTransformerBlock( - inner_dim, - n_heads, - d_head, - dropout=dropout, - context_dim=context_dim, - disable_self_attn=disable_self_attn, - checkpoint=use_checkpoint, - attention_cls=attention_cls, - video_length=video_length, - image_cross_attention=image_cross_attention, - image_cross_attention_scale_learnable=image_cross_attention_scale_learnable, - device=device, - dtype=dtype - ) for d in range(depth) - ]) - if not use_linear: - self.proj_out = zero_module(operations.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0, device=device, dtype=dtype)) - else: - self.proj_out = zero_module(operations.Linear(inner_dim, in_channels, device=device, dtype=dtype)) - self.use_linear = use_linear - - def forward(self, x, context=None, transformer_options={}, **kwargs): - b, c, h, w = x.shape - x_in = x - x = self.norm(x) - if not self.use_linear: - x = self.proj_in(x) - x = rearrange(x, 'b c h w -> b (h w) c').contiguous() - if self.use_linear: - x = self.proj_in(x) - for i, block in enumerate(self.transformer_blocks): - transformer_options['block_index'] = i - x = block(x, context=context, **kwargs) - if self.use_linear: - x = self.proj_out(x) - x = rearrange(x, 'b (h w) c -> b c h w', h=h, w=w).contiguous() - if not self.use_linear: - x = self.proj_out(x) - return x + x_in - - -class TemporalTransformer(nn.Module): - """ - Transformer block for image-like data in temporal axis. - First, reshape to b, t, d. - Then apply standard transformer action. - Finally, reshape to image - """ - def __init__( - self, - in_channels, - n_heads, - d_head, - depth=1, - dropout=0., - context_dim=None, - use_checkpoint=True, - use_linear=False, - only_self_att=True, - causal_attention=False, - causal_block_size=1, - relative_position=False, - temporal_length=None, - device=None, - dtype=None, - operations=ops - ): - super().__init__() - self.only_self_att = only_self_att - self.relative_position = relative_position - self.causal_attention = causal_attention - self.causal_block_size = causal_block_size - - if only_self_att: - context_dim = None - - self.in_channels = in_channels - inner_dim = n_heads * d_head - self.norm = operations.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True, device=device, dtype=dtype) - self.proj_in = nn.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0).to(device, dtype) - if not use_linear: - self.proj_in = nn.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0).to(device, dtype) - else: - self.proj_in = operations.Linear(in_channels, inner_dim, device=device, dtype=dtype) - - if relative_position: - assert(temporal_length is not None) - attention_cls = partial(CrossAttention, relative_position=True, temporal_length=temporal_length, device=device, dtype=dtype) - else: - attention_cls = partial(CrossAttention, temporal_length=temporal_length, device=device, dtype=dtype) - if self.causal_attention: - assert(temporal_length is not None) - self.mask = torch.tril(torch.ones([1, temporal_length, temporal_length])) - - if self.only_self_att: - context_dim = None - self.transformer_blocks = nn.ModuleList([ - BasicTransformerBlock( - inner_dim, - n_heads, - d_head, - dropout=dropout, - context_dim=context_dim, - attention_cls=attention_cls, - checkpoint=use_checkpoint, - device=device, - dtype=dtype - ) for d in range(depth) - ]) - if not use_linear: - self.proj_out = zero_module(nn.Conv1d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0).to(device, dtype)) - else: - self.proj_out = zero_module(operations.Linear(inner_dim, in_channels, device=device, dtype=dtype)) - self.use_linear = use_linear - - def forward(self, x, context=None): - b, c, t, h, w = x.shape - x_in = x - x = self.norm(x) - x = rearrange(x, 'b c t h w -> (b h w) c t').contiguous() - if not self.use_linear: - x = self.proj_in(x) - x = rearrange(x, 'bhw c t -> bhw t c').contiguous() - if self.use_linear: - x = self.proj_in(x) - - temp_mask = None - if self.causal_attention: - # slice the from mask map - temp_mask = self.mask[:,:t,:t].to(x.device) - - if temp_mask is not None: - mask = temp_mask.to(x.device) - mask = repeat(mask, 'l i j -> (l bhw) i j', bhw=b*h*w) - else: - mask = None - - if self.only_self_att: - ## note: if no context is given, cross-attention defaults to self-attention - for i, block in enumerate(self.transformer_blocks): - x = block(x, mask=mask) - x = rearrange(x, '(b hw) t c -> b hw t c', b=b).contiguous() - else: - x = rearrange(x, '(b hw) t c -> b hw t c', b=b).contiguous() - context = rearrange(context, '(b t) l con -> b t l con', t=t).contiguous() - for i, block in enumerate(self.transformer_blocks): - # calculate each batch one by one (since number in shape could not greater then 65,535 for some package) - for j in range(b): - context_j = repeat( - context[j], - 't l con -> (t r) l con', r=(h * w) // t, t=t).contiguous() - ## note: causal mask will not applied in cross-attention case - x[j] = block(x[j], context=context_j) - - if self.use_linear: - x = self.proj_out(x) - x = rearrange(x, 'b (h w) t c -> b c t h w', h=h, w=w).contiguous() - if not self.use_linear: - x = rearrange(x, 'b hw t c -> (b hw) c t').contiguous() - x = self.proj_out(x) - x = rearrange(x, '(b h w) c t -> b c t h w', b=b, h=h, w=w).contiguous() - - return x + x_in - - -class GEGLU(nn.Module): - def __init__(self, dim_in, dim_out, device=None, dtype=None, operations=ops): - super().__init__() - self.proj = operations.Linear(dim_in, dim_out * 2, device=device, dtype=dtype) - - def forward(self, x): - x, gate = self.proj(x).chunk(2, dim=-1) - return x * F.gelu(gate) - - -class FeedForward(nn.Module): - def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0., device=None, dtype=None, operations=ops): - super().__init__() - inner_dim = int(dim * mult) - dim_out = default(dim_out, dim) - project_in = nn.Sequential( - operations.Linear(dim, inner_dim, device=device, dtype=dtype), - nn.GELU() - ) if not glu else GEGLU(dim, inner_dim) - - self.net = nn.Sequential( - project_in, - nn.Dropout(dropout), - operations.Linear(inner_dim, dim_out, device=device, dtype=dtype) - ) - - def forward(self, x): - return self.net(x) - - -class LinearAttention(nn.Module): - def __init__(self, dim, heads=4, dim_head=32, device=None, dtype=None, operations=ops): - super().__init__() - self.heads = heads - hidden_dim = dim_head * heads - self.to_qkv = operations.Conv2d(dim, hidden_dim * 3, 1, bias = False, device=device, dtype=dtype) - self.to_out = operations.Conv2d(hidden_dim, dim, 1, device=device, dtype=dtype) - - def forward(self, x): - b, c, h, w = x.shape - qkv = self.to_qkv(x) - q, k, v = rearrange(qkv, 'b (qkv heads c) h w -> qkv b heads c (h w)', heads = self.heads, qkv=3) - k = k.softmax(dim=-1) - context = torch.einsum('bhdn,bhen->bhde', k, v) - out = torch.einsum('bhde,bhdn->bhen', context, q) - out = rearrange(out, 'b heads c (h w) -> b (heads c) h w', heads=self.heads, h=h, w=w) - return self.to_out(out) - - -class SpatialSelfAttention(nn.Module): - def __init__(self, in_channels, device=None, dtype=None, operations=ops): - super().__init__() - self.in_channels = in_channels - - self.norm = operations.GroupNorm( - num_groups=32, - num_channels=in_channels, - eps=1e-6, - affine=True, - device=device, - dtype=dtype - ) - self.q = operations.Conv2d( - in_channels, - in_channels, - kernel_size=1, - stride=1, - padding=0, - device=device, - dtype=dtype - ) - self.k = operations.Conv2d( - in_channels, - in_channels, - kernel_size=1, - stride=1, - padding=0, - device=device, - dtype=dtype - ) - self.v = operations.Conv2d( - in_channels, - in_channels, - kernel_size=1, - stride=1, - padding=0, - device=device, - dtype=dtype - ) - self.proj_out = operations.Conv2d( - in_channels, - in_channels, - kernel_size=1, - stride=1, - padding=0, - device=device, - dtype=dtype - ) - - def forward(self, x): - h_ = x - h_ = self.norm(h_) - q = self.q(h_) - k = self.k(h_) - v = self.v(h_) - - # compute attention - b,c,h,w = q.shape - q = rearrange(q, 'b c h w -> b (h w) c') - k = rearrange(k, 'b c h w -> b c (h w)') - w_ = torch.einsum('bij,bjk->bik', q, k) - - w_ = w_ * (int(c)**(-0.5)) - w_ = torch.nn.functional.softmax(w_, dim=2) - - # attend to values - v = rearrange(v, 'b c h w -> b c (h w)') - w_ = rearrange(w_, 'b i j -> b j i') - h_ = torch.einsum('bij,bjk->bik', v, w_) - h_ = rearrange(h_, 'b c (h w) -> b c h w', h=h) - h_ = self.proj_out(h_) - - return x+h_ diff --git a/py/dynamiCrafter/lvdm/modules/encoders/condition.py b/py/dynamiCrafter/lvdm/modules/encoders/condition.py deleted file mode 100644 index 610322b..0000000 --- a/py/dynamiCrafter/lvdm/modules/encoders/condition.py +++ /dev/null @@ -1,389 +0,0 @@ -import torch -import torch.nn as nn -import kornia -import open_clip -from torch.utils.checkpoint import checkpoint -from transformers import T5Tokenizer, T5EncoderModel, CLIPTokenizer, CLIPTextModel -from ..common import autocast -from utils.utils import count_params - - -class AbstractEncoder(nn.Module): - def __init__(self): - super().__init__() - - def encode(self, *args, **kwargs): - raise NotImplementedError - - -class IdentityEncoder(AbstractEncoder): - def encode(self, x): - return x - - -class ClassEmbedder(nn.Module): - def __init__(self, embed_dim, n_classes=1000, key='class', ucg_rate=0.1): - super().__init__() - self.key = key - self.embedding = nn.Embedding(n_classes, embed_dim) - self.n_classes = n_classes - self.ucg_rate = ucg_rate - - def forward(self, batch, key=None, disable_dropout=False): - if key is None: - key = self.key - # this is for use in crossattn - c = batch[key][:, None] - if self.ucg_rate > 0. and not disable_dropout: - mask = 1. - torch.bernoulli(torch.ones_like(c) * self.ucg_rate) - c = mask * c + (1 - mask) * torch.ones_like(c) * (self.n_classes - 1) - c = c.long() - c = self.embedding(c) - return c - - def get_unconditional_conditioning(self, bs, device="cuda"): - 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} - return uc - - -def disabled_train(self, mode=True): - """Overwrite model.train with this function to make sure train/eval mode - does not change anymore.""" - return self - - -class FrozenT5Embedder(AbstractEncoder): - """Uses the T5 transformer encoder for text""" - - def __init__(self, version="google/t5-v1_1-large", device="cuda", 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) - self.transformer = T5EncoderModel.from_pretrained(version) - self.device = device - self.max_length = max_length # TODO: typical value? - if freeze: - self.freeze() - - def freeze(self): - self.transformer = self.transformer.eval() - # self.train = disabled_train - for param in self.parameters(): - param.requires_grad = False - - def forward(self, text): - batch_encoding = self.tokenizer(text, truncation=True, max_length=self.max_length, return_length=True, - return_overflowing_tokens=False, padding="max_length", return_tensors="pt") - tokens = batch_encoding["input_ids"].to(self.device) - outputs = self.transformer(input_ids=tokens) - - z = outputs.last_hidden_state - return z - - def encode(self, text): - return self(text) - - -class FrozenCLIPEmbedder(AbstractEncoder): - """Uses the CLIP transformer encoder for text (from huggingface)""" - LAYERS = [ - "last", - "pooled", - "hidden" - ] - - def __init__(self, version="openai/clip-vit-large-patch14", device="cuda", max_length=77, - freeze=True, layer="last", layer_idx=None): # clip-vit-base-patch32 - super().__init__() - assert layer in self.LAYERS - self.tokenizer = CLIPTokenizer.from_pretrained(version) - self.transformer = CLIPTextModel.from_pretrained(version) - self.device = device - self.max_length = max_length - if freeze: - self.freeze() - self.layer = layer - self.layer_idx = layer_idx - if layer == "hidden": - assert layer_idx is not None - assert 0 <= abs(layer_idx) <= 12 - - def freeze(self): - self.transformer = self.transformer.eval() - # self.train = disabled_train - for param in self.parameters(): - param.requires_grad = False - - def forward(self, text): - batch_encoding = self.tokenizer(text, truncation=True, max_length=self.max_length, return_length=True, - return_overflowing_tokens=False, padding="max_length", return_tensors="pt") - tokens = batch_encoding["input_ids"].to(self.device) - outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden") - if self.layer == "last": - z = outputs.last_hidden_state - elif self.layer == "pooled": - z = outputs.pooler_output[:, None, :] - else: - z = outputs.hidden_states[self.layer_idx] - return z - - def encode(self, text): - return self(text) - - -class ClipImageEmbedder(nn.Module): - def __init__( - self, - model, - jit=False, - device='cuda' if torch.cuda.is_available() else 'cpu', - antialias=True, - ucg_rate=0. - ): - super().__init__() - from clip import load as load_clip - self.model, _ = load_clip(name=model, device=device, jit=jit) - - self.antialias = antialias - - self.register_buffer('mean', torch.Tensor([0.48145466, 0.4578275, 0.40821073]), persistent=False) - self.register_buffer('std', torch.Tensor([0.26862954, 0.26130258, 0.27577711]), persistent=False) - self.ucg_rate = ucg_rate - - def preprocess(self, x): - # normalize to [0,1] - x = kornia.geometry.resize(x, (224, 224), - interpolation='bicubic', align_corners=True, - antialias=self.antialias) - x = (x + 1.) / 2. - # re-normalize according to clip - x = kornia.enhance.normalize(x, self.mean, self.std) - return x - - def forward(self, x, no_dropout=False): - # x is assumed to be in range [-1,1] - out = self.model.encode_image(self.preprocess(x)) - out = out.to(x.dtype) - if self.ucg_rate > 0. and not no_dropout: - out = torch.bernoulli((1. - self.ucg_rate) * torch.ones(out.shape[0], device=out.device))[:, None] * out - return out - - -class FrozenOpenCLIPEmbedder(AbstractEncoder): - """ - Uses the OpenCLIP transformer encoder for text - """ - LAYERS = [ - # "pooled", - "last", - "penultimate" - ] - - def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77, - freeze=True, layer="last"): - super().__init__() - assert layer in self.LAYERS - model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), pretrained=version) - del model.visual - self.model = model - - self.device = device - self.max_length = max_length - if freeze: - self.freeze() - self.layer = layer - if self.layer == "last": - self.layer_idx = 0 - elif self.layer == "penultimate": - self.layer_idx = 1 - else: - raise NotImplementedError() - - def freeze(self): - self.model = self.model.eval() - for param in self.parameters(): - param.requires_grad = False - - def forward(self, text): - tokens = open_clip.tokenize(text) ## all clip models use 77 as context length - z = self.encode_with_transformer(tokens.to(self.device)) - return z - - def encode_with_transformer(self, text): - x = self.model.token_embedding(text) # [batch_size, n_ctx, d_model] - x = x + self.model.positional_embedding - x = x.permute(1, 0, 2) # NLD -> LND - x = self.text_transformer_forward(x, attn_mask=self.model.attn_mask) - x = x.permute(1, 0, 2) # LND -> NLD - x = self.model.ln_final(x) - return x - - def text_transformer_forward(self, x: torch.Tensor, attn_mask=None): - for i, r in enumerate(self.model.transformer.resblocks): - if i == len(self.model.transformer.resblocks) - self.layer_idx: - break - if self.model.transformer.grad_checkpointing and not torch.jit.is_scripting(): - x = checkpoint(r, x, attn_mask) - else: - x = r(x, attn_mask=attn_mask) - return x - - def encode(self, text): - return self(text) - - -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, - freeze=True, layer="pooled", antialias=True, ucg_rate=0.): - super().__init__() - model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), - pretrained=version, ) - del model.transformer - self.model = model - # self.mapper = torch.nn.Linear(1280, 1024) - self.device = device - self.max_length = max_length - if freeze: - self.freeze() - self.layer = layer - if self.layer == "penultimate": - raise NotImplementedError() - self.layer_idx = 1 - - self.antialias = antialias - - self.register_buffer('mean', torch.Tensor([0.48145466, 0.4578275, 0.40821073]), persistent=False) - self.register_buffer('std', torch.Tensor([0.26862954, 0.26130258, 0.27577711]), persistent=False) - self.ucg_rate = ucg_rate - - def preprocess(self, x): - # normalize to [0,1] - x = kornia.geometry.resize(x, (224, 224), - interpolation='bicubic', align_corners=True, - antialias=self.antialias) - x = (x + 1.) / 2. - # renormalize according to clip - x = kornia.enhance.normalize(x, self.mean, self.std) - return x - - def freeze(self): - self.model = self.model.eval() - for param in self.model.parameters(): - param.requires_grad = False - - @autocast - def forward(self, image, no_dropout=False): - z = self.encode_with_vision_transformer(image) - if self.ucg_rate > 0. and not no_dropout: - z = torch.bernoulli((1. - self.ucg_rate) * torch.ones(z.shape[0], device=z.device))[:, None] * z - return z - - def encode_with_vision_transformer(self, img): - img = self.preprocess(img) - x = self.model.visual(img) - return x - - def encode(self, text): - return self(text) - -class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder): - """ - Uses the OpenCLIP vision transformer encoder for images - """ - - def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", - freeze=True, layer="pooled", antialias=True): - super().__init__() - model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), - pretrained=version, ) - del model.transformer - self.model = model - self.device = device - - if freeze: - self.freeze() - self.layer = layer - if self.layer == "penultimate": - raise NotImplementedError() - self.layer_idx = 1 - - self.antialias = antialias - - self.register_buffer('mean', torch.Tensor([0.48145466, 0.4578275, 0.40821073]), persistent=False) - self.register_buffer('std', torch.Tensor([0.26862954, 0.26130258, 0.27577711]), persistent=False) - - - def preprocess(self, x): - # normalize to [0,1] - x = kornia.geometry.resize(x, (224, 224), - interpolation='bicubic', align_corners=True, - antialias=self.antialias) - x = (x + 1.) / 2. - # renormalize according to clip - x = kornia.enhance.normalize(x, self.mean, self.std) - return x - - def freeze(self): - self.model = self.model.eval() - for param in self.model.parameters(): - param.requires_grad = False - - def forward(self, image, no_dropout=False): - ## image: b c h w - z = self.encode_with_vision_transformer(image) - return z - - def encode_with_vision_transformer(self, x): - x = self.preprocess(x) - - # to patches - whether to use dual patchnorm - https://arxiv.org/abs/2302.01327v1 - if self.model.visual.input_patchnorm: - # einops - rearrange(x, 'b c (h p1) (w p2) -> b (h w) (c p1 p2)') - x = x.reshape(x.shape[0], x.shape[1], self.model.visual.grid_size[0], self.model.visual.patch_size[0], self.model.visual.grid_size[1], self.model.visual.patch_size[1]) - x = x.permute(0, 2, 4, 1, 3, 5) - x = x.reshape(x.shape[0], self.model.visual.grid_size[0] * self.model.visual.grid_size[1], -1) - x = self.model.visual.patchnorm_pre_ln(x) - x = self.model.visual.conv1(x) - else: - x = self.model.visual.conv1(x) # shape = [*, width, grid, grid] - x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2] - x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width] - - # class embeddings and positional embeddings - x = torch.cat( - [self.model.visual.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device), - x], dim=1) # shape = [*, grid ** 2 + 1, width] - x = x + self.model.visual.positional_embedding.to(x.dtype) - - # a patch_dropout of 0. would mean it is disabled and this function would do nothing but return what was passed in - x = self.model.visual.patch_dropout(x) - x = self.model.visual.ln_pre(x) - - x = x.permute(1, 0, 2) # NLD -> LND - x = self.model.visual.transformer(x) - x = x.permute(1, 0, 2) # LND -> NLD - - return x - -class FrozenCLIPT5Encoder(AbstractEncoder): - def __init__(self, clip_version="openai/clip-vit-large-patch14", t5_version="google/t5-v1_1-xl", device="cuda", - clip_max_length=77, t5_max_length=77): - super().__init__() - self.clip_encoder = FrozenCLIPEmbedder(clip_version, device, max_length=clip_max_length) - self.t5_encoder = FrozenT5Embedder(t5_version, device, max_length=t5_max_length) - print(f"{self.clip_encoder.__class__.__name__} has {count_params(self.clip_encoder) * 1.e-6:.2f} M parameters, " - f"{self.t5_encoder.__class__.__name__} comes with {count_params(self.t5_encoder) * 1.e-6:.2f} M params.") - - def encode(self, text): - return self(text) - - def forward(self, text): - clip_z = self.clip_encoder.encode(text) - t5_z = self.t5_encoder.encode(text) - return [clip_z, t5_z] diff --git a/py/dynamiCrafter/lvdm/modules/encoders/resampler.py b/py/dynamiCrafter/lvdm/modules/encoders/resampler.py deleted file mode 100644 index 0c30c58..0000000 --- a/py/dynamiCrafter/lvdm/modules/encoders/resampler.py +++ /dev/null @@ -1,145 +0,0 @@ -# modified from https://github.com/mlfoundations/open_flamingo/blob/main/open_flamingo/src/helpers.py -# and https://github.com/lucidrains/imagen-pytorch/blob/main/imagen_pytorch/imagen_pytorch.py -# and https://github.com/tencent-ailab/IP-Adapter/blob/main/ip_adapter/resampler.py -import math -import torch -import torch.nn as nn - - -class ImageProjModel(nn.Module): - """Projection Model""" - def __init__(self, cross_attention_dim=1024, clip_embeddings_dim=1024, clip_extra_context_tokens=4): - super().__init__() - self.cross_attention_dim = cross_attention_dim - self.clip_extra_context_tokens = clip_extra_context_tokens - self.proj = nn.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim) - self.norm = nn.LayerNorm(cross_attention_dim) - - def forward(self, image_embeds): - #embeds = image_embeds - embeds = image_embeds.type(list(self.proj.parameters())[0].dtype) - clip_extra_context_tokens = self.proj(embeds).reshape(-1, self.clip_extra_context_tokens, self.cross_attention_dim) - clip_extra_context_tokens = self.norm(clip_extra_context_tokens) - return clip_extra_context_tokens - - -# FFN -def FeedForward(dim, mult=4): - inner_dim = int(dim * mult) - return nn.Sequential( - nn.LayerNorm(dim), - nn.Linear(dim, inner_dim, bias=False), - nn.GELU(), - nn.Linear(inner_dim, dim, bias=False), - ) - - -def reshape_tensor(x, heads): - bs, length, width = x.shape - #(bs, length, width) --> (bs, length, n_heads, dim_per_head) - x = x.view(bs, length, heads, -1) - # (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head) - x = x.transpose(1, 2) - # (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head) - x = x.reshape(bs, heads, length, -1) - return x - - -class PerceiverAttention(nn.Module): - def __init__(self, *, dim, dim_head=64, heads=8): - super().__init__() - self.scale = dim_head**-0.5 - self.dim_head = dim_head - self.heads = heads - inner_dim = dim_head * heads - - self.norm1 = nn.LayerNorm(dim) - self.norm2 = nn.LayerNorm(dim) - - self.to_q = nn.Linear(dim, inner_dim, bias=False) - self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False) - self.to_out = nn.Linear(inner_dim, dim, bias=False) - - - def forward(self, x, latents): - """ - Args: - x (torch.Tensor): image features - shape (b, n1, D) - latent (torch.Tensor): latent features - shape (b, n2, D) - """ - x = self.norm1(x) - latents = self.norm2(latents) - - b, l, _ = latents.shape - - q = self.to_q(latents) - kv_input = torch.cat((x, latents), dim=-2) - k, v = self.to_kv(kv_input).chunk(2, dim=-1) - - q = reshape_tensor(q, self.heads) - k = reshape_tensor(k, self.heads) - v = reshape_tensor(v, self.heads) - - # attention - scale = 1 / math.sqrt(math.sqrt(self.dim_head)) - weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards - weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) - out = weight @ v - - out = out.permute(0, 2, 1, 3).reshape(b, l, -1) - - return self.to_out(out) - - -class Resampler(nn.Module): - def __init__( - self, - dim=1024, - depth=8, - dim_head=64, - heads=16, - num_queries=8, - embedding_dim=768, - output_dim=1024, - ff_mult=4, - video_length=None, # using frame-wise version or not - ): - super().__init__() - ## queries for a single frame / image - self.num_queries = num_queries - self.video_length = video_length - - ## queries for each frame - if video_length is not None: - num_queries = num_queries * video_length - - self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5) - self.proj_in = nn.Linear(embedding_dim, dim) - self.proj_out = nn.Linear(dim, output_dim) - self.norm_out = nn.LayerNorm(output_dim) - - self.layers = nn.ModuleList([]) - for _ in range(depth): - self.layers.append( - nn.ModuleList( - [ - PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads), - FeedForward(dim=dim, mult=ff_mult), - ] - ) - ) - - def forward(self, x): - latents = self.latents.repeat(x.size(0), 1, 1) ## B (T L) C - x = self.proj_in(x) - - for attn, ff in self.layers: - latents = attn(x, latents) + latents - latents = ff(latents) + latents - - latents = self.proj_out(latents) - latents = self.norm_out(latents) # B L C or B (T L) C - - return latents \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/modules/networks/ae_modules.py b/py/dynamiCrafter/lvdm/modules/networks/ae_modules.py deleted file mode 100644 index 35b6817..0000000 --- a/py/dynamiCrafter/lvdm/modules/networks/ae_modules.py +++ /dev/null @@ -1,1026 +0,0 @@ -# pytorch_diffusion + derived encoder decoder -import math -import torch -import numpy as np -import torch.nn as nn -from einops import rearrange -from utils.utils import instantiate_from_config -from ...modules.attention import LinearAttention - -import comfy.ops -ops = comfy.ops.disable_weight_init - -def nonlinearity(x): - # swish - return x*torch.sigmoid(x) - - -def Normalize(in_channels, num_groups=32, device=None, dtype=None, operations=ops): - return operations.GroupNorm( - num_groups=num_groups, - num_channels=in_channels, - eps=1e-6, - affine=True, - device=device, - dtype=dtype - ) - -class LinAttnBlock(LinearAttention): - """to match AttnBlock usage""" - def __init__(self, in_channels, device=None, dtype=None): - super().__init__(dim=in_channels, heads=1, dim_head=in_channels, device=device, dtype=dtype) - - -class AttnBlock(nn.Module): - def __init__(self, in_channels, device=None, dtype=None, operations=ops): - super().__init__() - self.in_channels = in_channels - - self.norm = Normalize(in_channels, device=device, dtype=dtype) - self.q = operations.Conv2d( - in_channels, - in_channels, - kernel_size=1, - stride=1, - padding=0, - device=device, - dtype=dtype - ) - self.k = operations.Conv2d( - in_channels, - in_channels, - kernel_size=1, - stride=1, - padding=0, - device=device, - dtype=dtype - ) - self.v = operations.Conv2d( - in_channels, - in_channels, - kernel_size=1, - stride=1, - padding=0, - device=device, - dtype=dtype - ) - self.proj_out = operations.Conv2d( - in_channels, - in_channels, - kernel_size=1, - stride=1, - padding=0, - device=device, - dtype=dtype - ) - - def forward(self, x): - h_ = x - h_ = self.norm(h_) - q = self.q(h_) - k = self.k(h_) - v = self.v(h_) - - # compute attention - b,c,h,w = q.shape - q = q.reshape(b,c,h*w) # bcl - q = q.permute(0,2,1) # bcl -> blc l=hw - k = k.reshape(b,c,h*w) # bcl - - w_ = torch.bmm(q,k) # b,hw,hw w[b,i,j]=sum_c q[b,i,c]k[b,c,j] - w_ = w_ * (int(c)**(-0.5)) - w_ = torch.nn.functional.softmax(w_, dim=2) - - # attend to values - v = v.reshape(b,c,h*w) - w_ = w_.permute(0,2,1) # b,hw,hw (first hw of k, second of q) - h_ = torch.bmm(v,w_) # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j] - h_ = h_.reshape(b,c,h,w) - - h_ = self.proj_out(h_) - - return x+h_ - -def make_attn(in_channels, attn_type="vanilla", device=None, dtype=None): - assert attn_type in ["vanilla", "linear", "none"], f'attn_type {attn_type} unknown' - #print(f"making attention of type '{attn_type}' with {in_channels} in_channels") - if attn_type == "vanilla": - return AttnBlock(in_channels, device=device, dtype=dtype) - elif attn_type == "none": - return nn.Identity(in_channels) - else: - return LinAttnBlock(in_channels, device=device, dtype=dtype) - -class Downsample(nn.Module): - def __init__(self, in_channels, with_conv, device=None, dtype=None, operations=ops): - super().__init__() - self.with_conv = with_conv - self.in_channels = in_channels - if self.with_conv: - # no asymmetric padding in torch conv, must do it ourselves - self.conv = operations.Conv2d( - in_channels, - in_channels, - kernel_size=3, - stride=2, - padding=0, - device=device, - dtype=dtype - ) - def forward(self, x): - if self.with_conv: - pad = (0,1,0,1) - x = torch.nn.functional.pad(x, pad, mode="constant", value=0) - x = self.conv(x) - else: - x = torch.nn.functional.avg_pool2d(x, kernel_size=2, stride=2) - return x - -class Upsample(nn.Module): - def __init__(self, in_channels, with_conv, device=None, dtype=None, operations=ops): - super().__init__() - self.with_conv = with_conv - self.in_channels = in_channels - if self.with_conv: - self.conv = operations.Conv2d( - in_channels, - in_channels, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype - ) - - def forward(self, x): - x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest") - if self.with_conv: - x = self.conv(x) - return x - -def get_timestep_embedding(timesteps, embedding_dim): - """ - This matches the implementation in Denoising Diffusion Probabilistic Models: - From Fairseq. - Build sinusoidal embeddings. - This matches the implementation in tensor2tensor, but differs slightly - from the description in Section 3.5 of "Attention Is All You Need". - """ - assert len(timesteps.shape) == 1 - - half_dim = embedding_dim // 2 - emb = math.log(10000) / (half_dim - 1) - emb = torch.exp(torch.arange(half_dim, dtype=torch.float32) * -emb) - emb = emb.to(device=timesteps.device) - emb = timesteps.float()[:, None] * emb[None, :] - emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1) - if embedding_dim % 2 == 1: # zero pad - emb = torch.nn.functional.pad(emb, (0,1,0,0)) - return emb - - - -class ResnetBlock(nn.Module): - def __init__( - self, - *, - in_channels, - out_channels=None, - conv_shortcut=False, - dropout, - temb_channels=512, - device=None, - dtype=None, - operations=ops - ): - super().__init__() - self.in_channels = in_channels - out_channels = in_channels if out_channels is None else out_channels - self.out_channels = out_channels - self.use_conv_shortcut = conv_shortcut - - self.norm1 = Normalize(in_channels, device=device, dtype=dtype) - self.conv1 = operations.Conv2d( - in_channels, - out_channels, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype - ) - if temb_channels > 0: - self.temb_proj = operations.Linear( - temb_channels, - out_channels, - device=device, - dtype=dtype - ) - self.norm2 = Normalize(out_channels, device=device, dtype=dtype) - self.dropout = torch.nn.Dropout(dropout) - self.conv2 = operations.Conv2d( - out_channels, - out_channels, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype - ) - if self.in_channels != self.out_channels: - if self.use_conv_shortcut: - self.conv_shortcut = operations.Conv2d( - in_channels, - out_channels, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=device - ) - else: - self.nin_shortcut = operations.Conv2d( - in_channels, - out_channels, - kernel_size=1, - stride=1, - padding=0, - device=device, - dtype=dtype - ) - - def forward(self, x, temb): - h = x - h = self.norm1(h) - h = nonlinearity(h) - h = self.conv1(h) - - if temb is not None: - h = h + self.temb_proj(nonlinearity(temb))[:,:,None,None] - - h = self.norm2(h) - h = nonlinearity(h) - h = self.dropout(h) - h = self.conv2(h) - - if self.in_channels != self.out_channels: - if self.use_conv_shortcut: - x = self.conv_shortcut(x) - else: - x = self.nin_shortcut(x) - - return x+h - -class Model(nn.Module): - def __init__( - self, - *, - ch, - out_ch, - ch_mult=(1,2,4,8), - num_res_blocks, - attn_resolutions, - dropout=0.0, - resamp_with_conv=True, - in_channels, - resolution, - use_timestep=True, - use_linear_attn=False, - attn_type="vanilla", - device=None, - dtype=None, - operations=ops - ): - super().__init__() - if use_linear_attn: attn_type = "linear" - self.ch = ch - self.temb_ch = self.ch*4 - self.num_resolutions = len(ch_mult) - self.num_res_blocks = num_res_blocks - self.resolution = resolution - self.in_channels = in_channels - - self.use_timestep = use_timestep - if self.use_timestep: - # timestep embedding - self.temb = nn.Module() - self.temb.dense = nn.ModuleList([ - operations.Linear( - self.ch, - self.temb_ch, - device=device, - dtype=dtype - ), - operations.Linear( - self.temb_ch, - self.temb_ch, - device=device, - dtype=dtype - ), - ]) - - # downsampling - self.conv_in = operations.Conv2d( - in_channels, - self.ch, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype - ) - - curr_res = resolution - in_ch_mult = (1,)+tuple(ch_mult) - self.down = nn.ModuleList() - for i_level in range(self.num_resolutions): - block = nn.ModuleList() - attn = nn.ModuleList() - block_in = ch*in_ch_mult[i_level] - block_out = ch*ch_mult[i_level] - for i_block in range(self.num_res_blocks): - block.append(ResnetBlock(in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - device=device, - dtype=dtype)) - block_in = block_out - if curr_res in attn_resolutions: - attn.append(make_attn(block_in, attn_type=attn_type, device=device, dtype=dtype)) - down = nn.Module() - down.block = block - down.attn = attn - if i_level != self.num_resolutions-1: - down.downsample = Downsample(block_in, resamp_with_conv, device=device, dtype=dtype) - curr_res = curr_res // 2 - self.down.append(down) - - # middle - self.mid = nn.Module() - self.mid.block_1 = ResnetBlock(in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - device=device, - dtype=dtype) - self.mid.attn_1 = make_attn(block_in, attn_type=attn_type, device=device, dtype=dtype) - self.mid.block_2 = ResnetBlock(in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - device=device, - dtype=dtype) - - # upsampling - self.up = nn.ModuleList() - for i_level in reversed(range(self.num_resolutions)): - block = nn.ModuleList() - attn = nn.ModuleList() - block_out = ch*ch_mult[i_level] - skip_in = ch*ch_mult[i_level] - for i_block in range(self.num_res_blocks+1): - if i_block == self.num_res_blocks: - skip_in = ch*in_ch_mult[i_level] - block.append(ResnetBlock(in_channels=block_in+skip_in, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - device=device, - dtype=dtype)) - block_in = block_out - if curr_res in attn_resolutions: - attn.append(make_attn(block_in, attn_type=attn_type, device=device, dtype=dtype)) - up = nn.Module() - up.block = block - up.attn = attn - if i_level != 0: - up.upsample = Upsample(block_in, resamp_with_conv, device=device, dtype=dtype) - curr_res = curr_res * 2 - self.up.insert(0, up) # prepend to get consistent order - - # end - self.norm_out = Normalize(block_in, device=device, dtype=device) - self.conv_out = torch.nn.Conv2d(block_in, - out_ch, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype) - - def forward(self, x, t=None, context=None): - #assert x.shape[2] == x.shape[3] == self.resolution - if context is not None: - # assume aligned context, cat along channel axis - x = torch.cat((x, context), dim=1) - if self.use_timestep: - # timestep embedding - assert t is not None - temb = get_timestep_embedding(t, self.ch) - temb = self.temb.dense[0](temb) - temb = nonlinearity(temb) - temb = self.temb.dense[1](temb) - else: - temb = None - - # downsampling - hs = [self.conv_in(x)] - for i_level in range(self.num_resolutions): - for i_block in range(self.num_res_blocks): - h = self.down[i_level].block[i_block](hs[-1], temb) - if len(self.down[i_level].attn) > 0: - h = self.down[i_level].attn[i_block](h) - hs.append(h) - if i_level != self.num_resolutions-1: - hs.append(self.down[i_level].downsample(hs[-1])) - - # middle - h = hs[-1] - h = self.mid.block_1(h, temb) - h = self.mid.attn_1(h) - h = self.mid.block_2(h, temb) - - # upsampling - for i_level in reversed(range(self.num_resolutions)): - for i_block in range(self.num_res_blocks+1): - h = self.up[i_level].block[i_block]( - torch.cat([h, hs.pop()], dim=1), temb) - if len(self.up[i_level].attn) > 0: - h = self.up[i_level].attn[i_block](h) - if i_level != 0: - h = self.up[i_level].upsample(h) - - # end - h = self.norm_out(h) - h = nonlinearity(h) - h = self.conv_out(h) - return h - - def get_last_layer(self): - return self.conv_out.weight - - -class Encoder(nn.Module): - def __init__( - self, - *, - ch, - out_ch, - ch_mult=(1,2,4,8), - num_res_blocks, - attn_resolutions, - dropout=0.0, - resamp_with_conv=True, - in_channels, - resolution, - z_channels, - double_z=True, - use_linear_attn=False, - attn_type="vanilla", - device=None, - dtype=None, - operations=ops, - **ignore_kwargs - ): - super().__init__() - if use_linear_attn: attn_type = "linear" - self.ch = ch - self.temb_ch = 0 - self.num_resolutions = len(ch_mult) - self.num_res_blocks = num_res_blocks - self.resolution = resolution - self.in_channels = in_channels - - # downsampling - self.conv_in = operations.Conv2d(in_channels, - self.ch, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype) - - curr_res = resolution - in_ch_mult = (1,)+tuple(ch_mult) - self.in_ch_mult = in_ch_mult - self.down = nn.ModuleList() - for i_level in range(self.num_resolutions): - block = nn.ModuleList() - attn = nn.ModuleList() - block_in = ch*in_ch_mult[i_level] - block_out = ch*ch_mult[i_level] - for i_block in range(self.num_res_blocks): - block.append(ResnetBlock(in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - device=device, - dtype=dtype)) - block_in = block_out - if curr_res in attn_resolutions: - attn.append(make_attn(block_in, attn_type=attn_type, device=device, dtype=dtype)) - down = nn.Module() - down.block = block - down.attn = attn - if i_level != self.num_resolutions-1: - down.downsample = Downsample(block_in, resamp_with_conv, device=device, dtype=dtype) - curr_res = curr_res // 2 - self.down.append(down) - - # middle - self.mid = nn.Module() - self.mid.block_1 = ResnetBlock(in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - device=device, - dtype=dtype) - self.mid.attn_1 = make_attn(block_in, attn_type=attn_type, device=device, dtype=dtype) - self.mid.block_2 = ResnetBlock(in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - device=device, - dtype=dtype) - - # end - self.norm_out = Normalize(block_in, device=device, dtype=dtype) - self.conv_out = operations.Conv2d(block_in, - 2*z_channels if double_z else z_channels, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype) - - def forward(self, x): - # timestep embedding - temb = None - - # print(f'encoder-input={x.shape}') - # downsampling - hs = [self.conv_in(x)] - # print(f'encoder-conv in feat={hs[0].shape}') - for i_level in range(self.num_resolutions): - for i_block in range(self.num_res_blocks): - h = self.down[i_level].block[i_block](hs[-1], temb) - # print(f'encoder-down feat={h.shape}') - if len(self.down[i_level].attn) > 0: - h = self.down[i_level].attn[i_block](h) - hs.append(h) - if i_level != self.num_resolutions-1: - # print(f'encoder-downsample (input)={hs[-1].shape}') - hs.append(self.down[i_level].downsample(hs[-1])) - # print(f'encoder-downsample (output)={hs[-1].shape}') - - # middle - h = hs[-1] - h = self.mid.block_1(h, temb) - # print(f'encoder-mid1 feat={h.shape}') - h = self.mid.attn_1(h) - h = self.mid.block_2(h, temb) - # print(f'encoder-mid2 feat={h.shape}') - - # end - h = self.norm_out(h) - h = nonlinearity(h) - h = self.conv_out(h) - # print(f'end feat={h.shape}') - return h - - -class Decoder(nn.Module): - def __init__( - self, - *, - ch, - out_ch, - ch_mult=(1,2,4,8), - num_res_blocks, - attn_resolutions, - dropout=0.0, - resamp_with_conv=True, - in_channels, - resolution, - z_channels, - give_pre_end=False, - tanh_out=False, - use_linear_attn=False, - attn_type="vanilla", - device=None, - dtype=None, - operations=ops, - **ignorekwargs - ): - super().__init__() - if use_linear_attn: attn_type = "linear" - self.ch = ch - self.temb_ch = 0 - self.num_resolutions = len(ch_mult) - self.num_res_blocks = num_res_blocks - self.resolution = resolution - self.in_channels = in_channels - self.give_pre_end = give_pre_end - self.tanh_out = tanh_out - - # compute in_ch_mult, block_in and curr_res at lowest res - in_ch_mult = (1,)+tuple(ch_mult) - block_in = ch*ch_mult[self.num_resolutions-1] - curr_res = resolution // 2**(self.num_resolutions-1) - self.z_shape = (1,z_channels,curr_res,curr_res) - print("AE working on z of shape {} = {} dimensions.".format( - self.z_shape, np.prod(self.z_shape))) - - # z to block_in - self.conv_in = torch.nn.Conv2d(z_channels, - block_in, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype) - - # middle - self.mid = nn.Module() - self.mid.block_1 = ResnetBlock(in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - device=device, - dtype=dtype) - self.mid.attn_1 = make_attn(block_in, attn_type=attn_type, device=device, dtype=dtype) - self.mid.block_2 = ResnetBlock(in_channels=block_in, - out_channels=block_in, - temb_channels=self.temb_ch, - dropout=dropout, - device=device, - dtype=dtype) - - # upsampling - self.up = nn.ModuleList() - for i_level in reversed(range(self.num_resolutions)): - block = nn.ModuleList() - attn = nn.ModuleList() - block_out = ch*ch_mult[i_level] - for i_block in range(self.num_res_blocks+1): - block.append(ResnetBlock(in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - device=device, - dtype=dtype)) - block_in = block_out - if curr_res in attn_resolutions: - attn.append(make_attn(block_in, attn_type=attn_type, device=device, dtype=dtype)) - up = nn.Module() - up.block = block - up.attn = attn - if i_level != 0: - up.upsample = Upsample(block_in, resamp_with_conv, device=device, dtype=dtype) - curr_res = curr_res * 2 - self.up.insert(0, up) # prepend to get consistent order - - # end - self.norm_out = Normalize(block_in, device=device, dtype=dtype) - self.conv_out = operations.Conv2d(block_in, - out_ch, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype) - - def forward(self, z): - #assert z.shape[1:] == self.z_shape[1:] - self.last_z_shape = z.shape - - # print(f'decoder-input={z.shape}') - # timestep embedding - temb = None - - # z to block_in - h = self.conv_in(z) - # print(f'decoder-conv in feat={h.shape}') - - # middle - h = self.mid.block_1(h, temb) - h = self.mid.attn_1(h) - h = self.mid.block_2(h, temb) - # print(f'decoder-mid feat={h.shape}') - - # upsampling - for i_level in reversed(range(self.num_resolutions)): - for i_block in range(self.num_res_blocks+1): - h = self.up[i_level].block[i_block](h, temb) - if len(self.up[i_level].attn) > 0: - h = self.up[i_level].attn[i_block](h) - # print(f'decoder-up feat={h.shape}') - if i_level != 0: - h = self.up[i_level].upsample(h) - # print(f'decoder-upsample feat={h.shape}') - - # end - if self.give_pre_end: - return h - - h = self.norm_out(h) - h = nonlinearity(h) - h = self.conv_out(h) - # print(f'decoder-conv_out feat={h.shape}') - if self.tanh_out: - h = torch.tanh(h) - return h - - -class SimpleDecoder(nn.Module): - def __init__(self, in_channels, out_channels, device=None, dtype=None, operations=ops, *args, **kwargs): - super().__init__() - self.model = nn.ModuleList([nn.Conv2d(in_channels, in_channels, 1), - ResnetBlock(in_channels=in_channels, - out_channels=2 * in_channels, - temb_channels=0, dropout=0.0, - device=device, - dtype=dtype), - ResnetBlock(in_channels=2 * in_channels, - out_channels=4 * in_channels, - temb_channels=0, dropout=0.0, - device=device, - dtype=dtype), - ResnetBlock(in_channels=4 * in_channels, - out_channels=2 * in_channels, - temb_channels=0, dropout=0.0, - device=device, - dtype=dtype), - operations.Conv2d(2*in_channels, in_channels, 1), - Upsample(in_channels, with_conv=True, device=device, dtype=dtype)]) - # end - self.norm_out = Normalize(in_channels, device=device, dtype=dtype) - self.conv_out = torch.nn.Conv2d(in_channels, - out_channels, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype) - - def forward(self, x): - for i, layer in enumerate(self.model): - if i in [1,2,3]: - x = layer(x, None) - else: - x = layer(x) - - h = self.norm_out(x) - h = nonlinearity(h) - x = self.conv_out(h) - return x - - -class UpsampleDecoder(nn.Module): - def __init__(self, in_channels, out_channels, ch, num_res_blocks, resolution, - ch_mult=(2,2), dropout=0.0, device=None, dtype=None, operations=ops): - super().__init__() - # upsampling - self.temb_ch = 0 - self.num_resolutions = len(ch_mult) - self.num_res_blocks = num_res_blocks - block_in = in_channels - curr_res = resolution // 2 ** (self.num_resolutions - 1) - self.res_blocks = nn.ModuleList() - self.upsample_blocks = nn.ModuleList() - for i_level in range(self.num_resolutions): - res_block = [] - block_out = ch * ch_mult[i_level] - for i_block in range(self.num_res_blocks + 1): - res_block.append(ResnetBlock(in_channels=block_in, - out_channels=block_out, - temb_channels=self.temb_ch, - dropout=dropout, - device=device, - dtype=dtype)) - block_in = block_out - self.res_blocks.append(nn.ModuleList(res_block)) - if i_level != self.num_resolutions - 1: - self.upsample_blocks.append(Upsample(block_in, True, device=device, dtype=dtype)) - curr_res = curr_res * 2 - - # end - self.norm_out = Normalize(block_in, device=device, dtype=dtype) - self.conv_out = torch.nn.Conv2d(block_in, - out_channels, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype) - - def forward(self, x): - # upsampling - h = x - for k, i_level in enumerate(range(self.num_resolutions)): - for i_block in range(self.num_res_blocks + 1): - h = self.res_blocks[i_level][i_block](h, None) - if i_level != self.num_resolutions - 1: - h = self.upsample_blocks[k](h) - h = self.norm_out(h) - h = nonlinearity(h) - h = self.conv_out(h) - return h - - -class LatentRescaler(nn.Module): - def __init__(self, factor, in_channels, mid_channels, out_channels, depth=2, device=None, dtype=None, operations=ops): - super().__init__() - # residual block, interpolate, residual block - self.factor = factor - self.conv_in = operations.Conv2d(in_channels, - mid_channels, - kernel_size=3, - stride=1, - padding=1, - device=device, - dtype=dtype) - self.res_block1 = nn.ModuleList([ResnetBlock(in_channels=mid_channels, - out_channels=mid_channels, - temb_channels=0, - dropout=0.0, - device=device, - dtype=dtype) for _ in range(depth)]) - self.attn = AttnBlock(mid_channels, device=device, dtype=dtype) - self.res_block2 = nn.ModuleList([ResnetBlock(in_channels=mid_channels, - out_channels=mid_channels, - temb_channels=0, - dropout=0.0, - device=device, - dtype=dtype) for _ in range(depth)]) - - self.conv_out = operations.Conv2d(mid_channels, - out_channels, - kernel_size=1, - device=device, - dtype=dtype - ) - - def forward(self, x): - x = self.conv_in(x) - for block in self.res_block1: - x = block(x, None) - x = torch.nn.functional.interpolate(x, size=(int(round(x.shape[2]*self.factor)), int(round(x.shape[3]*self.factor)))) - x = self.attn(x) - for block in self.res_block2: - x = block(x, None) - x = self.conv_out(x) - return x - - -class MergedRescaleEncoder(nn.Module): - def __init__(self, in_channels, ch, resolution, out_ch, num_res_blocks, - attn_resolutions, dropout=0.0, resamp_with_conv=True, - ch_mult=(1,2,4,8), rescale_factor=1.0, rescale_module_depth=1, device=None, dtype=None, operations=ops): - super().__init__() - intermediate_chn = ch * ch_mult[-1] - self.encoder = Encoder(in_channels=in_channels, num_res_blocks=num_res_blocks, ch=ch, ch_mult=ch_mult, - z_channels=intermediate_chn, double_z=False, resolution=resolution, - attn_resolutions=attn_resolutions, dropout=dropout, resamp_with_conv=resamp_with_conv, - out_ch=None, device=device, dtype=dtype) - self.rescaler = LatentRescaler(factor=rescale_factor, in_channels=intermediate_chn, - mid_channels=intermediate_chn, out_channels=out_ch, depth=rescale_module_depth, - device=device, dtype=dtype) - - def forward(self, x): - x = self.encoder(x) - x = self.rescaler(x) - return x - - -class MergedRescaleDecoder(nn.Module): - def __init__(self, z_channels, out_ch, resolution, num_res_blocks, attn_resolutions, ch, ch_mult=(1,2,4,8), - dropout=0.0, resamp_with_conv=True, rescale_factor=1.0, rescale_module_depth=1, - device=None, dtype=None, operations=ops): - super().__init__() - tmp_chn = z_channels*ch_mult[-1] - self.decoder = Decoder(out_ch=out_ch, z_channels=tmp_chn, attn_resolutions=attn_resolutions, dropout=dropout, - resamp_with_conv=resamp_with_conv, in_channels=None, num_res_blocks=num_res_blocks, - ch_mult=ch_mult, resolution=resolution, ch=ch, device=device, operations=ops) - self.rescaler = LatentRescaler(factor=rescale_factor, in_channels=z_channels, mid_channels=tmp_chn, - out_channels=tmp_chn, depth=rescale_module_depth, device=device, operations=ops) - - def forward(self, x): - x = self.rescaler(x) - x = self.decoder(x) - return x - - -class Upsampler(nn.Module): - def __init__(self, in_size, out_size, in_channels, out_channels, ch_mult=2, device=None, dtype=None, operations=ops): - super().__init__() - assert out_size >= in_size - num_blocks = int(np.log2(out_size//in_size))+1 - factor_up = 1.+ (out_size % in_size) - print(f"Building {self.__class__.__name__} with in_size: {in_size} --> out_size {out_size} and factor {factor_up}") - self.rescaler = LatentRescaler(factor=factor_up, in_channels=in_channels, mid_channels=2*in_channels, - out_channels=in_channels, device=device, dtype=dtype) - self.decoder = Decoder(out_ch=out_channels, resolution=out_size, z_channels=in_channels, num_res_blocks=2, - attn_resolutions=[], in_channels=None, ch=in_channels, device=device, dtype=dtype, - ch_mult=[ch_mult for _ in range(num_blocks)]) - - def forward(self, x): - x = self.rescaler(x) - x = self.decoder(x) - return x - - -class Resize(nn.Module): - def __init__(self, in_channels=None, learned=False, mode="bilinear", device=None, dtype=None, operations=ops): - super().__init__() - self.with_conv = learned - self.mode = mode - if self.with_conv: - print(f"Note: {self.__class__.__name} uses learned downsampling and will ignore the fixed {mode} mode") - raise NotImplementedError() - assert in_channels is not None - # no asymmetric padding in torch conv, must do it ourselves - self.conv = operations.Conv2d(in_channels, - in_channels, - kernel_size=4, - stride=2, - padding=1, - device=device, - dtype=dtype) - - def forward(self, x, scale_factor=1.0): - if scale_factor==1.0: - return x - else: - x = torch.nn.functional.interpolate(x, mode=self.mode, align_corners=False, scale_factor=scale_factor) - return x - -class FirstStagePostProcessor(nn.Module): - - def __init__(self, ch_mult:list, in_channels, - pretrained_model:nn.Module=None, - reshape=False, - n_channels=None, - dropout=0., - pretrained_config=None, - device=None, - dtype=None, - operations=ops): - super().__init__() - if pretrained_config is None: - assert pretrained_model is not None, 'Either "pretrained_model" or "pretrained_config" must not be None' - self.pretrained_model = pretrained_model - else: - assert pretrained_config is not None, 'Either "pretrained_model" or "pretrained_config" must not be None' - self.instantiate_pretrained(pretrained_config) - - self.do_reshape = reshape - - if n_channels is None: - n_channels = self.pretrained_model.encoder.ch - - self.proj_norm = Normalize(in_channels,num_groups=in_channels//2, device=device, dtype=dtype) - self.proj = nn.Conv2d(in_channels,n_channels,kernel_size=3, - stride=1,padding=1, device=device, dtype=dtype) - - blocks = [] - downs = [] - ch_in = n_channels - for m in ch_mult: - blocks.append(ResnetBlock(in_channels=ch_in,out_channels=m*n_channels,dropout=dropout, device=device, dtype=dtype)) - ch_in = m * n_channels - downs.append(Downsample(ch_in, with_conv=False, device=device, dtype=dtype)) - - self.model = nn.ModuleList(blocks) - self.downsampler = nn.ModuleList(downs) - - - def instantiate_pretrained(self, config): - model = instantiate_from_config(config) - self.pretrained_model = model.eval() - # self.pretrained_model.train = False - for param in self.pretrained_model.parameters(): - param.requires_grad = False - - - @torch.no_grad() - def encode_with_pretrained(self,x): - c = self.pretrained_model.encode(x) - if isinstance(c, DiagonalGaussianDistribution): - c = c.mode() - return c - - def forward(self,x): - z_fs = self.encode_with_pretrained(x) - z = self.proj_norm(z_fs) - z = self.proj(z) - z = nonlinearity(z) - - for submodel, downmodel in zip(self.model,self.downsampler): - z = submodel(z,temb=None) - z = downmodel(z) - - if self.do_reshape: - z = rearrange(z,'b c h w -> b (h w) c') - return z \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/modules/networks/openaimodel3d.py b/py/dynamiCrafter/lvdm/modules/networks/openaimodel3d.py deleted file mode 100644 index cd4e1c6..0000000 --- a/py/dynamiCrafter/lvdm/modules/networks/openaimodel3d.py +++ /dev/null @@ -1,822 +0,0 @@ -from functools import partial -from abc import abstractmethod -import torch -import torch.nn as nn -from einops import rearrange -import torch.nn.functional as F -from ...models.utils_diffusion import timestep_embedding -from ...common import checkpoint -from ...basics import ( - zero_module, - conv_nd, - linear, - avg_pool_nd, - normalization -) -from ...modules.attention import SpatialTransformer, TemporalTransformer -import comfy.ops -import logging - -ops = comfy.ops.disable_weight_init - -class TimestepBlock(nn.Module): - """ - Any module where forward() takes timestep embeddings as a second argument. - """ - @abstractmethod - def forward(self, x, emb): - """ - Apply the module to `x` given `emb` timestep embeddings. - """ - -#This is needed because accelerate makes a copy of transformer_options which breaks "transformer_index" -def forward_timestep_embed(ts, x, emb, context=None, batch_size=None, transformer_options={}): - for layer in ts: - if isinstance(layer, TimestepBlock): - x = layer(x, emb, batch_size=batch_size) - elif isinstance(layer, SpatialTransformer): - x = layer(x, context) - if "transformer_index" in transformer_options: - transformer_options["transformer_index"] += 1 - elif isinstance(layer, TemporalTransformer): - x = rearrange(x, '(b f) c h w -> b c f h w', b=batch_size) - x = layer(x, context) - if "transformer_index" in transformer_options: - transformer_options["transformer_index"] += 1 - x = rearrange(x, 'b c f h w -> (b f) c h w') - else: - x = layer(x) - return x - -class TimestepEmbedSequential(nn.Sequential, TimestepBlock): - """ - A sequential module that passes timestep embeddings to the children that - support it as an extra input. - """ - - def forward(self, *args, **kwargs): - return forward_timestep_embed(self, *args, **kwargs) - -class Downsample(nn.Module): - """ - A downsampling layer with an optional convolution. - :param channels: channels in the inputs and outputs. - :param use_conv: a bool determining if a convolution is applied. - :param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then - downsampling occurs in the inner-two dimensions. - """ - - def __init__(self, channels, use_conv, dims=2, out_channels=None, padding=1, dtype=None, device=None, operations=ops): - super().__init__() - self.channels = channels - self.out_channels = out_channels or channels - self.use_conv = use_conv - self.dims = dims - stride = 2 if dims != 3 else (1, 2, 2) - if use_conv: - self.op = operations.conv_nd( - dims, self.channels, self.out_channels, 3, stride=stride, padding=padding - ) - else: - assert self.channels == self.out_channels - self.op = avg_pool_nd(dims, kernel_size=stride, stride=stride) - - def forward(self, x): - assert x.shape[1] == self.channels - return self.op(x) - -class Upsample(nn.Module): - """ - An upsampling layer with an optional convolution. - :param channels: channels in the inputs and outputs. - :param use_conv: a bool determining if a convolution is applied. - :param dims: determines if the signal is 1D, 2D, or 3D. If 3D, then - upsampling occurs in the inner-two dimensions. - """ - - def __init__(self, channels, use_conv, dims=2, out_channels=None, padding=1, dtype=None, device=None, operations=ops): - super().__init__() - self.channels = channels - self.out_channels = out_channels or channels - self.use_conv = use_conv - self.dims = dims - if use_conv: - self.conv = operations.conv_nd(dims, self.channels, self.out_channels, 3, padding=padding, dtype=dtype, device=device) - - def forward(self, x): - assert x.shape[1] == self.channels - if self.dims == 3: - x = F.interpolate(x, (x.shape[2], x.shape[3] * 2, x.shape[4] * 2), mode='nearest') - else: - x = F.interpolate(x, scale_factor=2, mode='nearest') - if self.use_conv: - x = self.conv(x) - return x - -class ResBlock(TimestepBlock): - """ - A residual block that can optionally change the number of channels. - :param channels: the number of input channels. - :param emb_channels: the number of timestep embedding channels. - :param dropout: the rate of dropout. - :param out_channels: if specified, the number of out channels. - :param use_conv: if True and out_channels is specified, use a spatial - convolution instead of a smaller 1x1 convolution to change the - channels in the skip connection. - :param dims: determines if the signal is 1D, 2D, or 3D. - :param up: if True, use this block for upsampling. - :param down: if True, use this block for downsampling. - :param use_temporal_conv: if True, use the temporal convolution. - :param use_image_dataset: if True, the temporal parameters will not be optimized. - """ - - def __init__( - self, - channels, - emb_channels, - dropout, - out_channels=None, - use_scale_shift_norm=False, - dims=2, - use_checkpoint=False, - use_conv=False, - up=False, - down=False, - kernel_size=3, - use_temporal_conv=False, - tempspatial_aware=False, - dtype=None, - device=None, - operations=ops - ): - super().__init__() - self.channels = channels - self.emb_channels = emb_channels - self.dropout = dropout - self.out_channels = out_channels or channels - self.use_conv = use_conv - self.use_checkpoint = use_checkpoint - self.use_scale_shift_norm = use_scale_shift_norm - self.use_temporal_conv = use_temporal_conv - - if isinstance(kernel_size, list): - padding =[k // 2 for k in kernel_size] - else: - padding = kernel_size // 2 - - # operations used in normalization function - self.in_layers = nn.Sequential( - normalization(channels, dtype=dtype, device=device), - nn.SiLU(), - operations.conv_nd(dims, channels, self.out_channels, 3, padding=1, dtype=dtype, device=device), - ) - - self.updown = up or down - - if up: - self.h_upd = Upsample(channels, False, dims, dtype=dtype, device=device) - self.x_upd = Upsample(channels, False, dims, dtype=dtype, device=device) - elif down: - self.h_upd = Downsample(channels, False, dims, dtype=dtype, device=device) - self.x_upd = Downsample(channels, False, dims, dtype=dtype, device=device) - else: - self.h_upd = self.x_upd = nn.Identity() - - self.emb_layers = nn.Sequential( - nn.SiLU(), - operations.Linear( - emb_channels, - 2 * self.out_channels if use_scale_shift_norm else self.out_channels, - dtype=dtype, - device=device - ), - ) - self.out_layers = nn.Sequential( - normalization(self.out_channels, dtype=dtype, device=device), - nn.SiLU(), - nn.Dropout(p=dropout), - zero_module(operations.Conv2d(self.out_channels, self.out_channels, 3, padding=1, dtype=dtype, device=device)), - ) - - if self.out_channels == channels: - self.skip_connection = nn.Identity() - elif use_conv: - self.skip_connection = operations.conv_nd(dims, channels, self.out_channels, 3, padding=1, dtype=dtype, device=device) - else: - self.skip_connection = operations.conv_nd(dims, channels, self.out_channels, 1, dtype=dtype, device=device) - - if self.use_temporal_conv: - self.temopral_conv = TemporalConvBlock( - self.out_channels, - self.out_channels, - dropout=0.1, - spatial_aware=tempspatial_aware, - dtype=dtype, - device=device - ) - - def forward(self, x, emb, batch_size=None): - """ - Apply the block to a Tensor, conditioned on a timestep embedding. - :param x: an [N x C x ...] Tensor of features. - :param emb: an [N x emb_channels] Tensor of timestep embeddings. - :return: an [N x C x ...] Tensor of outputs. - """ - input_tuple = (x, emb) - if batch_size: - forward_batchsize = partial(self._forward, batch_size=batch_size) - return checkpoint(forward_batchsize, input_tuple, self.parameters(), self.use_checkpoint) - return checkpoint(self._forward, input_tuple, self.parameters(), self.use_checkpoint) - - def _forward(self, x, emb, batch_size=None): - if self.updown: - in_rest, in_conv = self.in_layers[:-1], self.in_layers[-1] - h = in_rest(x) - h = self.h_upd(h) - x = self.x_upd(x) - h = in_conv(h) - else: - h = self.in_layers(x) - emb_out = self.emb_layers(emb).type(h.dtype) - while len(emb_out.shape) < len(h.shape): - emb_out = emb_out[..., None] - if self.use_scale_shift_norm: - out_norm, out_rest = self.out_layers[0], self.out_layers[1:] - scale, shift = torch.chunk(emb_out, 2, dim=1) - h = out_norm(h) * (1 + scale) + shift - h = out_rest(h) - else: - h = h + emb_out - h = self.out_layers(h) - h = self.skip_connection(x) + h - - if self.use_temporal_conv and batch_size: - h = rearrange(h, '(b t) c h w -> b c t h w', b=batch_size) - h = self.temopral_conv(h) - h = rearrange(h, 'b c t h w -> (b t) c h w') - return h - -class TemporalConvBlock(nn.Module): - """ - Adapted from modelscope: https://github.com/modelscope/modelscope/blob/master/modelscope/models/multi_modal/video_synthesis/unet_sd.py - """ - def __init__( - self, - in_channels, - out_channels=None, - dropout=0.0, - spatial_aware=False, - dtype=None, - device=None, - operations=ops - ): - super(TemporalConvBlock, self).__init__() - if out_channels is None: - out_channels = in_channels - self.in_channels = in_channels - self.out_channels = out_channels - th_kernel_shape = (3, 1, 1) if not spatial_aware else (3, 3, 1) - th_padding_shape = (1, 0, 0) if not spatial_aware else (1, 1, 0) - tw_kernel_shape = (3, 1, 1) if not spatial_aware else (3, 1, 3) - tw_padding_shape = (1, 0, 0) if not spatial_aware else (1, 0, 1) - - # conv layers - self.conv1 = nn.Sequential( - operations.GroupNorm(32, in_channels, device=device, dtype=dtype), nn.SiLU(), - operations.Conv3d(in_channels, out_channels, th_kernel_shape, padding=th_padding_shape, device=device, dtype=dtype)) - self.conv2 = nn.Sequential( - operations.GroupNorm(32, out_channels, device=device, dtype=dtype), nn.SiLU(), nn.Dropout(dropout), - operations.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape, device=device, dtype=dtype)) - self.conv3 = nn.Sequential( - operations.GroupNorm(32, out_channels, device=device, dtype=dtype), nn.SiLU(), nn.Dropout(dropout), - operations.Conv3d(out_channels, in_channels, th_kernel_shape, padding=th_padding_shape, device=device, dtype=dtype)) - self.conv4 = nn.Sequential( - operations.GroupNorm(32, out_channels, device=device, dtype=dtype), nn.SiLU(), nn.Dropout(dropout), - operations.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape, device=device, dtype=dtype)) - - # zero out the last layer params,so the conv block is identity - nn.init.zeros_(self.conv4[-1].weight) - nn.init.zeros_(self.conv4[-1].bias) - - def forward(self, x): - identity = x - x = self.conv1(x) - x = self.conv2(x) - x = self.conv3(x) - x = self.conv4(x) - - return identity + x - -def context_processor(context, t, img_emb=None, temporal_size=16, concat_only=False, disable_concat=False): - if disable_concat: - return context - - ## repeat t times for context [(b t) 77 768] & time embedding - ## check if we use per-frame image conditioning - - if img_emb is not None: - context = torch.cat([context, img_emb.to(context.device, context.dtype)], dim=1) - - if concat_only: - return context - - b, l_context, _ = context.shape - if l_context == 77 + t * temporal_size: - context_text, context_img = context[:,:77,:], context[:,77:,:] - context_text = context_text.repeat_interleave(repeats=t, dim=0) - context_img = rearrange(context_img, 'b (t l) c -> (b t) l c', t=t) - context = torch.cat([context_text, context_img], dim=1) - else: - context = context.repeat_interleave(repeats=t, dim=0) - - return context - -def apply_control(h, control, name, cond_idx=None): - if control is not None and name in control and len(control[name]) > 0: - frames = h.shape[0] - ctrl = control[name].pop() - if ctrl is not None: - try: - if cond_idx is not None and ctrl.shape[0] > frames: - ctrl_frames_list = list(range(ctrl.shape[0])) - ctrl_frames = len(ctrl_frames_list) - - idxs = ( - ctrl_frames_list[ctrl_frames // 2:] if cond_idx == 0 else \ - ctrl_frames_list[:ctrl_frames // 2] - ) - - ctrl = ctrl[idxs] - - h += ctrl - except Exception as e: - if h.shape != ctrl.shape: - logging.warning( - "warning control could not be applied {} {}".format(h.shape, ctrl.shape) - ) - logging.warning(e) - return h - -class UNetModel(nn.Module): - """ - The full UNet model with attention and timestep embedding. - :param in_channels: in_channels in the input Tensor. - :param model_channels: base channel count for the model. - :param out_channels: channels in the output Tensor. - :param num_res_blocks: number of residual blocks per downsample. - :param attention_resolutions: a collection of downsample rates at which - attention will take place. May be a set, list, or tuple. - For example, if this contains 4, then at 4x downsampling, attention - will be used. - :param dropout: the dropout probability. - :param channel_mult: channel multiplier for each level of the UNet. - :param conv_resample: if True, use learned convolutions for upsampling and - downsampling. - :param dims: determines if the signal is 1D, 2D, or 3D. - :param num_classes: if specified (as an int), then this model will be - class-conditional with `num_classes` classes. - :param use_checkpoint: use gradient checkpointing to reduce memory usage. - :param num_heads: the number of attention heads in each attention layer. - :param num_heads_channels: if specified, ignore num_heads and instead use - a fixed channel width per attention head. - :param num_heads_upsample: works with num_heads to set a different number - of heads for upsampling. Deprecated. - :param use_scale_shift_norm: use a FiLM-like conditioning mechanism. - :param resblock_updown: use residual blocks for up/downsampling. - :param use_new_attention_order: use a different attention pattern for potentially - increased efficiency. - """ - - def __init__(self, - in_channels, - model_channels, - out_channels, - num_res_blocks, - attention_resolutions, - dropout=0.0, - channel_mult=(1, 2, 4, 8), - conv_resample=True, - dims=2, - context_dim=None, - use_scale_shift_norm=False, - resblock_updown=False, - num_heads=-1, - num_head_channels=-1, - transformer_depth=1, - use_linear=False, - use_checkpoint=False, - temporal_conv=False, - tempspatial_aware=False, - temporal_attention=True, - use_relative_position=True, - use_causal_attention=False, - temporal_length=None, - use_fp16=False, - addition_attention=False, - temporal_selfatt_only=True, - image_cross_attention=False, - image_cross_attention_scale_learnable=False, - default_fs=4, - fs_condition=False, - device=None, - dtype=torch.float16, - operations=ops - ): - super(UNetModel, self).__init__() - if num_heads == -1: - assert num_head_channels != -1, 'Either num_heads or num_head_channels has to be set' - if num_head_channels == -1: - assert num_heads != -1, 'Either num_heads or num_head_channels has to be set' - - self.in_channels = in_channels - self.model_channels = model_channels - self.out_channels = out_channels - self.num_res_blocks = num_res_blocks - self.attention_resolutions = attention_resolutions - self.dropout = dropout - self.channel_mult = channel_mult - self.conv_resample = conv_resample - self.temporal_attention = temporal_attention - time_embed_dim = model_channels * 4 - self.use_checkpoint = use_checkpoint - temporal_self_att_only = True - self.addition_attention = addition_attention - self.temporal_length = temporal_length - self.image_cross_attention = image_cross_attention - self.image_cross_attention_scale_learnable = image_cross_attention_scale_learnable - self.default_fs = default_fs - self.fs_condition = fs_condition - self.device = device - #self.dtype = dtype - self.dtype = torch.float32 - - ## Time embedding blocks - self.time_embed = nn.Sequential( - linear(model_channels, time_embed_dim, device=device, dtype=self.dtype), - nn.SiLU(), - linear(time_embed_dim, time_embed_dim, device=device, dtype=self.dtype), - ) - if fs_condition: - self.fps_embedding = nn.Sequential( - linear(model_channels, time_embed_dim, device=device, dtype=self.dtype), - nn.SiLU(), - linear(time_embed_dim, time_embed_dim, device=device, dtype=self.dtype), - ) - nn.init.zeros_(self.fps_embedding[-1].weight) - nn.init.zeros_(self.fps_embedding[-1].bias) - ## Input Block - self.input_blocks = nn.ModuleList( - [ - TimestepEmbedSequential( - operations.conv_nd( - dims, - in_channels, - model_channels, - 3, - padding=1, - device=device, - dtype=self.dtype - )) - ] - ) - if self.addition_attention: - self.init_attn=TimestepEmbedSequential( - TemporalTransformer( - model_channels, - n_heads=8, - d_head=num_head_channels, - depth=transformer_depth, - context_dim=context_dim, - use_checkpoint=use_checkpoint, only_self_att=temporal_selfatt_only, - causal_attention=False, relative_position=use_relative_position, - temporal_length=temporal_length, - device=device, - dtype=self.dtype - )) - - input_block_chans = [model_channels] - ch = model_channels - ds = 1 - for level, mult in enumerate(channel_mult): - for _ in range(num_res_blocks): - layers = [ - ResBlock(ch, time_embed_dim, dropout, - out_channels=mult * model_channels, dims=dims, use_checkpoint=use_checkpoint, - use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware, - use_temporal_conv=temporal_conv, - device=device, - dtype=self.dtype - ) - ] - ch = mult * model_channels - if ds in attention_resolutions: - if num_head_channels == -1: - dim_head = ch // num_heads - else: - num_heads = ch // num_head_channels - dim_head = num_head_channels - layers.append( - SpatialTransformer(ch, num_heads, dim_head, - depth=transformer_depth, context_dim=context_dim, use_linear=use_linear, - use_checkpoint=use_checkpoint, disable_self_attn=False, - video_length=temporal_length, image_cross_attention=self.image_cross_attention, - image_cross_attention_scale_learnable=self.image_cross_attention_scale_learnable, - device=device, - dtype=self.dtype - ) - ) - if self.temporal_attention: - layers.append( - TemporalTransformer(ch, num_heads, dim_head, - depth=transformer_depth, context_dim=context_dim, use_linear=use_linear, - use_checkpoint=use_checkpoint, only_self_att=temporal_self_att_only, - causal_attention=use_causal_attention, relative_position=use_relative_position, - temporal_length=temporal_length, - device=device, - dtype=self.dtype - ) - ) - self.input_blocks.append(TimestepEmbedSequential(*layers)) - input_block_chans.append(ch) - if level != len(channel_mult) - 1: - out_ch = ch - self.input_blocks.append( - TimestepEmbedSequential( - ResBlock(ch, time_embed_dim, dropout, - out_channels=out_ch, dims=dims, use_checkpoint=use_checkpoint, - use_scale_shift_norm=use_scale_shift_norm, - down=True, - device=device, - dtype=self.dtype - ) - if resblock_updown - else Downsample( - ch, - conv_resample, - dims=dims, - out_channels=out_ch, - device=device, - dtype=self.dtype - ) - ) - ) - ch = out_ch - input_block_chans.append(ch) - ds *= 2 - - if num_head_channels == -1: - dim_head = ch // num_heads - else: - num_heads = ch // num_head_channels - dim_head = num_head_channels - layers = [ - ResBlock(ch, time_embed_dim, dropout, - dims=dims, use_checkpoint=use_checkpoint, - use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware, - use_temporal_conv=temporal_conv, - device=device, - dtype=self.dtype - ), - SpatialTransformer(ch, num_heads, dim_head, - depth=transformer_depth, context_dim=context_dim, use_linear=use_linear, - use_checkpoint=use_checkpoint, disable_self_attn=False, video_length=temporal_length, - image_cross_attention=self.image_cross_attention,image_cross_attention_scale_learnable=self.image_cross_attention_scale_learnable, - device=device, - dtype=self.dtype - ) - ] - if self.temporal_attention: - layers.append( - TemporalTransformer(ch, num_heads, dim_head, - depth=transformer_depth, context_dim=context_dim, use_linear=use_linear, - use_checkpoint=use_checkpoint, only_self_att=temporal_self_att_only, - causal_attention=use_causal_attention, relative_position=use_relative_position, - temporal_length=temporal_length, - device=device, - dtype=self.dtype - ) - ) - layers.append( - ResBlock(ch, time_embed_dim, dropout, - dims=dims, use_checkpoint=use_checkpoint, - use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware, - use_temporal_conv=temporal_conv, - device=device, - dtype=self.dtype - ) - ) - - ## Middle Block - self.middle_block = TimestepEmbedSequential(*layers) - - ## Output Block - self.output_blocks = nn.ModuleList([]) - for level, mult in list(enumerate(channel_mult))[::-1]: - for i in range(num_res_blocks + 1): - ich = input_block_chans.pop() - layers = [ - ResBlock(ch + ich, time_embed_dim, dropout, - out_channels=mult * model_channels, dims=dims, use_checkpoint=use_checkpoint, - use_scale_shift_norm=use_scale_shift_norm, tempspatial_aware=tempspatial_aware, - use_temporal_conv=temporal_conv, - device=device, - dtype=self.dtype - ) - ] - ch = model_channels * mult - if ds in attention_resolutions: - if num_head_channels == -1: - dim_head = ch // num_heads - else: - num_heads = ch // num_head_channels - dim_head = num_head_channels - layers.append( - SpatialTransformer(ch, num_heads, dim_head, - depth=transformer_depth, context_dim=context_dim, use_linear=use_linear, - use_checkpoint=use_checkpoint, disable_self_attn=False, video_length=temporal_length, - image_cross_attention=self.image_cross_attention,image_cross_attention_scale_learnable=self.image_cross_attention_scale_learnable, - device=device, - dtype=self.dtype - ) - ) - if self.temporal_attention: - layers.append( - TemporalTransformer(ch, num_heads, dim_head, - depth=transformer_depth, context_dim=context_dim, use_linear=use_linear, - use_checkpoint=use_checkpoint, only_self_att=temporal_self_att_only, - causal_attention=use_causal_attention, relative_position=use_relative_position, - temporal_length=temporal_length, - device=device, - dtype=self.dtype - ) - ) - if level and i == num_res_blocks: - out_ch = ch - layers.append( - ResBlock(ch, time_embed_dim, dropout, - out_channels=out_ch, dims=dims, use_checkpoint=use_checkpoint, - use_scale_shift_norm=use_scale_shift_norm, - up=True, - device=device, - dtype=self.dtype - ) - if resblock_updown - else Upsample(ch, conv_resample, dims=dims, out_channels=out_ch) - ) - ds //= 2 - self.output_blocks.append(TimestepEmbedSequential(*layers)) - - self.out = nn.Sequential( - normalization(ch, device=device, dtype=self.dtype), - nn.SiLU(), - zero_module( - operations.conv_nd( - dims, - model_channels, - out_channels, - 3, - padding=1, - device=device, - dtype=self.dtype - ) - ), - ) - - # TODO Add Transformer options to leverage the usage of patches. - def forward( - self, - x, - timesteps, - context=None, - context_in=None, - cc_concat=None, - num_video_frames=16, - features_adapter=None, - fs=None, - img_emb=None, - control=None, - transformer_options={}, - cond_idx=None, - **kwargs - ): - - if any([fs is None, img_emb is None, cc_concat is None]): - raise ValueError("One or more of the required inputs for UNet Forward is None.") - - cond_idx = transformer_options.get("cond_idx", None) - transformer_options['original_shape'] = list(x.shape) - transformer_options['transformer_index'] = 0 - transformer_patches = transformer_options.get("patches", {}) - - # In ComfyUI, the frames are always with the batch, so we deconstruct it here. - # This is mandatory as this is a video based model. - # We usually denote "f" as frames, but will use "t" (time) to be consistent with DynamiCrafter. - b,_,t,_,_ = x.shape - - context = context_in - cc_concat = cc_concat.to(x.device, x.dtype) - x = torch.cat([x, cc_concat], dim=1) - - fs = fs.to(x.device, x.dtype) - - timestep = timesteps - context = context_processor(context, num_video_frames, img_emb=img_emb) - - t_emb = timestep_embedding(timestep, self.model_channels, repeat_only=False, dtype=self.dtype) - emb = self.time_embed(t_emb) - emb = emb.repeat_interleave(repeats=t, dim=0) - - ## always in shape (b t) c h w, except for temporal layer - x = rearrange(x, 'b c t h w -> (b t) c h w') - - ## combine emb - if self.fs_condition: - if fs is None: - fs = torch.tensor( - [self.default_fs] * b, dtype=torch.long, device=x.device) - fs_emb = timestep_embedding(fs, self.model_channels, repeat_only=False, dtype=self.dtype).type(x.dtype) - - fs_embed = self.fps_embedding(fs_emb) - fs_embed = fs_embed.repeat_interleave(repeats=t, dim=0) - - emb = emb + fs_embed - - h = x.type(self.dtype) - adapter_idx = 0 - hs = [] - - for id, module in enumerate(self.input_blocks): - transformer_options["block"] = ("input", id) - #h = module(h, emb, context=context, batch_size=b) - h = forward_timestep_embed( - module, - h, - emb, - context=context, - batch_size=b, - transformer_options=transformer_options - ) - h = apply_control(h, control, 'input', cond_idx) - - if "input_block_patch" in transformer_patches: - patch = transformer_patches["input_block_patch"] - for p in patch: - h = p(h, transformer_options) - - if id ==0 and self.addition_attention: - h = forward_timestep_embed( - self.init_attn, - h, - emb, - context=context, - batch_size=b, - transformer_options=transformer_options - ) - ## plug-in adapter features - if ((id+1)%3 == 0) and features_adapter is not None: - h = h + features_adapter[adapter_idx] - adapter_idx += 1 - hs.append(h) - if "input_block_patch_after_skip" in transformer_patches: - patch = transformer_patches["input_block_patch_after_skip"] - for p in patch: - h = p(h, transformer_options) - if features_adapter is not None: - assert len(features_adapter)==adapter_idx, 'Wrong features_adapter' - transformer_options["block"] = ("middle", 0) - h = forward_timestep_embed( - self.middle_block, - h, - emb, - context=context, - batch_size=b, - transformer_options=transformer_options - ) - h = apply_control(h, control, 'middle', cond_idx) - for id, module in enumerate(self.output_blocks): - transformer_options["block"] = ("output", id) - hsp = hs.pop() - hsp = apply_control(hsp, control, 'output', cond_idx) - - if "output_block_patch" in transformer_patches: - patch = transformer_patches["output_block_patch"] - for p in patch: - h, hsp = p(h, hsp, transformer_options) - - h = torch.cat([h, hsp], dim=1) - del hsp - h = forward_timestep_embed( - module, - h, - emb, - context=context, - batch_size=b, - transformer_options=transformer_options - ) - h = h.type(x.dtype) - h = self.out(h) - - # We output with the tensor unfolded framewise, then reshape them to batched using ComfyUI nodes. - h = rearrange(h, '(b t) c h w -> b c t h w', t=num_video_frames) - - return h \ No newline at end of file diff --git a/py/dynamiCrafter/lvdm/modules/x_transformer.py b/py/dynamiCrafter/lvdm/modules/x_transformer.py deleted file mode 100644 index 5321012..0000000 --- a/py/dynamiCrafter/lvdm/modules/x_transformer.py +++ /dev/null @@ -1,639 +0,0 @@ -"""shout-out to https://github.com/lucidrains/x-transformers/tree/main/x_transformers""" -from functools import partial -from inspect import isfunction -from collections import namedtuple -from einops import rearrange, repeat -import torch -from torch import nn, einsum -import torch.nn.functional as F - -# constants -DEFAULT_DIM_HEAD = 64 - -Intermediates = namedtuple('Intermediates', [ - 'pre_softmax_attn', - 'post_softmax_attn' -]) - -LayerIntermediates = namedtuple('Intermediates', [ - 'hiddens', - 'attn_intermediates' -]) - - -class AbsolutePositionalEmbedding(nn.Module): - def __init__(self, dim, max_seq_len): - super().__init__() - self.emb = nn.Embedding(max_seq_len, dim) - self.init_() - - def init_(self): - nn.init.normal_(self.emb.weight, std=0.02) - - def forward(self, x): - n = torch.arange(x.shape[1], device=x.device) - return self.emb(n)[None, :, :] - - -class FixedPositionalEmbedding(nn.Module): - def __init__(self, dim): - super().__init__() - inv_freq = 1. / (10000 ** (torch.arange(0, dim, 2).float() / dim)) - self.register_buffer('inv_freq', inv_freq) - - def forward(self, x, seq_dim=1, offset=0): - t = torch.arange(x.shape[seq_dim], device=x.device).type_as(self.inv_freq) + offset - sinusoid_inp = torch.einsum('i , j -> i j', t, self.inv_freq) - emb = torch.cat((sinusoid_inp.sin(), sinusoid_inp.cos()), dim=-1) - return emb[None, :, :] - - -# helpers - -def exists(val): - return val is not None - - -def default(val, d): - if exists(val): - return val - return d() if isfunction(d) else d - - -def always(val): - def inner(*args, **kwargs): - return val - return inner - - -def not_equals(val): - def inner(x): - return x != val - return inner - - -def equals(val): - def inner(x): - return x == val - return inner - - -def max_neg_value(tensor): - return -torch.finfo(tensor.dtype).max - - -# keyword argument helpers - -def pick_and_pop(keys, d): - values = list(map(lambda key: d.pop(key), keys)) - return dict(zip(keys, values)) - - -def group_dict_by_key(cond, d): - return_val = [dict(), dict()] - for key in d.keys(): - match = bool(cond(key)) - ind = int(not match) - return_val[ind][key] = d[key] - return (*return_val,) - - -def string_begins_with(prefix, str): - return str.startswith(prefix) - - -def group_by_key_prefix(prefix, d): - return group_dict_by_key(partial(string_begins_with, prefix), d) - - -def groupby_prefix_and_trim(prefix, d): - kwargs_with_prefix, kwargs = group_dict_by_key(partial(string_begins_with, prefix), d) - kwargs_without_prefix = dict(map(lambda x: (x[0][len(prefix):], x[1]), tuple(kwargs_with_prefix.items()))) - return kwargs_without_prefix, kwargs - - -# classes -class Scale(nn.Module): - def __init__(self, value, fn): - super().__init__() - self.value = value - self.fn = fn - - def forward(self, x, **kwargs): - x, *rest = self.fn(x, **kwargs) - return (x * self.value, *rest) - - -class Rezero(nn.Module): - def __init__(self, fn): - super().__init__() - self.fn = fn - self.g = nn.Parameter(torch.zeros(1)) - - def forward(self, x, **kwargs): - x, *rest = self.fn(x, **kwargs) - return (x * self.g, *rest) - - -class ScaleNorm(nn.Module): - def __init__(self, dim, eps=1e-5): - super().__init__() - self.scale = dim ** -0.5 - self.eps = eps - self.g = nn.Parameter(torch.ones(1)) - - def forward(self, x): - norm = torch.norm(x, dim=-1, keepdim=True) * self.scale - return x / norm.clamp(min=self.eps) * self.g - - -class RMSNorm(nn.Module): - def __init__(self, dim, eps=1e-8): - super().__init__() - self.scale = dim ** -0.5 - self.eps = eps - self.g = nn.Parameter(torch.ones(dim)) - - def forward(self, x): - norm = torch.norm(x, dim=-1, keepdim=True) * self.scale - return x / norm.clamp(min=self.eps) * self.g - - -class Residual(nn.Module): - def forward(self, x, residual): - return x + residual - - -class GRUGating(nn.Module): - def __init__(self, dim): - super().__init__() - self.gru = nn.GRUCell(dim, dim) - - def forward(self, x, residual): - gated_output = self.gru( - rearrange(x, 'b n d -> (b n) d'), - rearrange(residual, 'b n d -> (b n) d') - ) - - return gated_output.reshape_as(x) - - -# feedforward - -class GEGLU(nn.Module): - def __init__(self, dim_in, dim_out): - super().__init__() - self.proj = nn.Linear(dim_in, dim_out * 2) - - def forward(self, x): - x, gate = self.proj(x).chunk(2, dim=-1) - return x * F.gelu(gate) - - -class FeedForward(nn.Module): - def __init__(self, dim, dim_out=None, mult=4, glu=False, dropout=0.): - super().__init__() - inner_dim = int(dim * mult) - dim_out = default(dim_out, dim) - project_in = nn.Sequential( - nn.Linear(dim, inner_dim), - nn.GELU() - ) if not glu else GEGLU(dim, inner_dim) - - self.net = nn.Sequential( - project_in, - nn.Dropout(dropout), - nn.Linear(inner_dim, dim_out) - ) - - def forward(self, x): - return self.net(x) - - -# attention. -class Attention(nn.Module): - def __init__( - self, - dim, - dim_head=DEFAULT_DIM_HEAD, - heads=8, - causal=False, - mask=None, - talking_heads=False, - sparse_topk=None, - use_entmax15=False, - num_mem_kv=0, - dropout=0., - on_attn=False - ): - super().__init__() - if use_entmax15: - raise NotImplementedError("Check out entmax activation instead of softmax activation!") - self.scale = dim_head ** -0.5 - self.heads = heads - self.causal = causal - self.mask = mask - - inner_dim = dim_head * heads - - self.to_q = nn.Linear(dim, inner_dim, bias=False) - self.to_k = nn.Linear(dim, inner_dim, bias=False) - self.to_v = nn.Linear(dim, inner_dim, bias=False) - self.dropout = nn.Dropout(dropout) - - # talking heads - self.talking_heads = talking_heads - if talking_heads: - self.pre_softmax_proj = nn.Parameter(torch.randn(heads, heads)) - self.post_softmax_proj = nn.Parameter(torch.randn(heads, heads)) - - # explicit topk sparse attention - self.sparse_topk = sparse_topk - - # entmax - #self.attn_fn = entmax15 if use_entmax15 else F.softmax - self.attn_fn = F.softmax - - # add memory key / values - self.num_mem_kv = num_mem_kv - if num_mem_kv > 0: - self.mem_k = nn.Parameter(torch.randn(heads, num_mem_kv, dim_head)) - self.mem_v = nn.Parameter(torch.randn(heads, num_mem_kv, dim_head)) - - # attention on attention - self.attn_on_attn = on_attn - self.to_out = nn.Sequential(nn.Linear(inner_dim, dim * 2), nn.GLU()) if on_attn else nn.Linear(inner_dim, dim) - - def forward( - self, - x, - context=None, - mask=None, - context_mask=None, - rel_pos=None, - sinusoidal_emb=None, - prev_attn=None, - mem=None - ): - b, n, _, h, talking_heads, device = *x.shape, self.heads, self.talking_heads, x.device - kv_input = default(context, x) - - q_input = x - k_input = kv_input - v_input = kv_input - - if exists(mem): - k_input = torch.cat((mem, k_input), dim=-2) - v_input = torch.cat((mem, v_input), dim=-2) - - if exists(sinusoidal_emb): - # in shortformer, the query would start at a position offset depending on the past cached memory - offset = k_input.shape[-2] - q_input.shape[-2] - q_input = q_input + sinusoidal_emb(q_input, offset=offset) - k_input = k_input + sinusoidal_emb(k_input) - - q = self.to_q(q_input) - k = self.to_k(k_input) - v = self.to_v(v_input) - - q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h=h), (q, k, v)) - - input_mask = None - if any(map(exists, (mask, context_mask))): - q_mask = default(mask, lambda: torch.ones((b, n), device=device).bool()) - k_mask = q_mask if not exists(context) else context_mask - k_mask = default(k_mask, lambda: torch.ones((b, k.shape[-2]), device=device).bool()) - q_mask = rearrange(q_mask, 'b i -> b () i ()') - k_mask = rearrange(k_mask, 'b j -> b () () j') - input_mask = q_mask * k_mask - - if self.num_mem_kv > 0: - mem_k, mem_v = map(lambda t: repeat(t, 'h n d -> b h n d', b=b), (self.mem_k, self.mem_v)) - k = torch.cat((mem_k, k), dim=-2) - v = torch.cat((mem_v, v), dim=-2) - if exists(input_mask): - input_mask = F.pad(input_mask, (self.num_mem_kv, 0), value=True) - - dots = einsum('b h i d, b h j d -> b h i j', q, k) * self.scale - mask_value = max_neg_value(dots) - - if exists(prev_attn): - dots = dots + prev_attn - - pre_softmax_attn = dots - - if talking_heads: - dots = einsum('b h i j, h k -> b k i j', dots, self.pre_softmax_proj).contiguous() - - if exists(rel_pos): - dots = rel_pos(dots) - - if exists(input_mask): - dots.masked_fill_(~input_mask, mask_value) - del input_mask - - if self.causal: - i, j = dots.shape[-2:] - r = torch.arange(i, device=device) - mask = rearrange(r, 'i -> () () i ()') < rearrange(r, 'j -> () () () j') - mask = F.pad(mask, (j - i, 0), value=False) - dots.masked_fill_(mask, mask_value) - del mask - - if exists(self.sparse_topk) and self.sparse_topk < dots.shape[-1]: - top, _ = dots.topk(self.sparse_topk, dim=-1) - vk = top[..., -1].unsqueeze(-1).expand_as(dots) - mask = dots < vk - dots.masked_fill_(mask, mask_value) - del mask - - attn = self.attn_fn(dots, dim=-1) - post_softmax_attn = attn - - attn = self.dropout(attn) - - if talking_heads: - attn = einsum('b h i j, h k -> b k i j', attn, self.post_softmax_proj).contiguous() - - out = einsum('b h i j, b h j d -> b h i d', attn, v) - out = rearrange(out, 'b h n d -> b n (h d)') - - intermediates = Intermediates( - pre_softmax_attn=pre_softmax_attn, - post_softmax_attn=post_softmax_attn - ) - - return self.to_out(out), intermediates - - -class AttentionLayers(nn.Module): - def __init__( - self, - dim, - depth, - heads=8, - causal=False, - cross_attend=False, - only_cross=False, - use_scalenorm=False, - use_rmsnorm=False, - use_rezero=False, - rel_pos_num_buckets=32, - rel_pos_max_distance=128, - position_infused_attn=False, - custom_layers=None, - sandwich_coef=None, - par_ratio=None, - residual_attn=False, - cross_residual_attn=False, - macaron=False, - pre_norm=True, - gate_residual=False, - **kwargs - ): - super().__init__() - ff_kwargs, kwargs = groupby_prefix_and_trim('ff_', kwargs) - attn_kwargs, _ = groupby_prefix_and_trim('attn_', kwargs) - - dim_head = attn_kwargs.get('dim_head', DEFAULT_DIM_HEAD) - - self.dim = dim - self.depth = depth - self.layers = nn.ModuleList([]) - - self.has_pos_emb = position_infused_attn - self.pia_pos_emb = FixedPositionalEmbedding(dim) if position_infused_attn else None - self.rotary_pos_emb = always(None) - - assert rel_pos_num_buckets <= rel_pos_max_distance, 'number of relative position buckets must be less than the relative position max distance' - self.rel_pos = None - - self.pre_norm = pre_norm - - self.residual_attn = residual_attn - self.cross_residual_attn = cross_residual_attn - - norm_class = ScaleNorm if use_scalenorm else nn.LayerNorm - norm_class = RMSNorm if use_rmsnorm else norm_class - norm_fn = partial(norm_class, dim) - - norm_fn = nn.Identity if use_rezero else norm_fn - branch_fn = Rezero if use_rezero else None - - if cross_attend and not only_cross: - default_block = ('a', 'c', 'f') - elif cross_attend and only_cross: - default_block = ('c', 'f') - else: - default_block = ('a', 'f') - - if macaron: - default_block = ('f',) + default_block - - if exists(custom_layers): - layer_types = custom_layers - elif exists(par_ratio): - par_depth = depth * len(default_block) - assert 1 < par_ratio <= par_depth, 'par ratio out of range' - default_block = tuple(filter(not_equals('f'), default_block)) - par_attn = par_depth // par_ratio - depth_cut = par_depth * 2 // 3 # 2 / 3 attention layer cutoff suggested by PAR paper - par_width = (depth_cut + depth_cut // par_attn) // par_attn - assert len(default_block) <= par_width, 'default block is too large for par_ratio' - par_block = default_block + ('f',) * (par_width - len(default_block)) - par_head = par_block * par_attn - layer_types = par_head + ('f',) * (par_depth - len(par_head)) - elif exists(sandwich_coef): - assert sandwich_coef > 0 and sandwich_coef <= depth, 'sandwich coefficient should be less than the depth' - layer_types = ('a',) * sandwich_coef + default_block * (depth - sandwich_coef) + ('f',) * sandwich_coef - else: - layer_types = default_block * depth - - self.layer_types = layer_types - self.num_attn_layers = len(list(filter(equals('a'), layer_types))) - - for layer_type in self.layer_types: - if layer_type == 'a': - layer = Attention(dim, heads=heads, causal=causal, **attn_kwargs) - elif layer_type == 'c': - layer = Attention(dim, heads=heads, **attn_kwargs) - elif layer_type == 'f': - layer = FeedForward(dim, **ff_kwargs) - layer = layer if not macaron else Scale(0.5, layer) - else: - raise Exception(f'invalid layer type {layer_type}') - - if isinstance(layer, Attention) and exists(branch_fn): - layer = branch_fn(layer) - - if gate_residual: - residual_fn = GRUGating(dim) - else: - residual_fn = Residual() - - self.layers.append(nn.ModuleList([ - norm_fn(), - layer, - residual_fn - ])) - - def forward( - self, - x, - context=None, - mask=None, - context_mask=None, - mems=None, - return_hiddens=False - ): - hiddens = [] - intermediates = [] - prev_attn = None - prev_cross_attn = None - - mems = mems.copy() if exists(mems) else [None] * self.num_attn_layers - - for ind, (layer_type, (norm, block, residual_fn)) in enumerate(zip(self.layer_types, self.layers)): - is_last = ind == (len(self.layers) - 1) - - if layer_type == 'a': - hiddens.append(x) - layer_mem = mems.pop(0) - - residual = x - - if self.pre_norm: - x = norm(x) - - if layer_type == 'a': - out, inter = block(x, mask=mask, sinusoidal_emb=self.pia_pos_emb, rel_pos=self.rel_pos, - prev_attn=prev_attn, mem=layer_mem) - elif layer_type == 'c': - out, inter = block(x, context=context, mask=mask, context_mask=context_mask, prev_attn=prev_cross_attn) - elif layer_type == 'f': - out = block(x) - - x = residual_fn(out, residual) - - if layer_type in ('a', 'c'): - intermediates.append(inter) - - if layer_type == 'a' and self.residual_attn: - prev_attn = inter.pre_softmax_attn - elif layer_type == 'c' and self.cross_residual_attn: - prev_cross_attn = inter.pre_softmax_attn - - if not self.pre_norm and not is_last: - x = norm(x) - - if return_hiddens: - intermediates = LayerIntermediates( - hiddens=hiddens, - attn_intermediates=intermediates - ) - - return x, intermediates - - return x - - -class Encoder(AttentionLayers): - def __init__(self, **kwargs): - assert 'causal' not in kwargs, 'cannot set causality on encoder' - super().__init__(causal=False, **kwargs) - - - -class TransformerWrapper(nn.Module): - def __init__( - self, - *, - num_tokens, - max_seq_len, - attn_layers, - emb_dim=None, - max_mem_len=0., - emb_dropout=0., - num_memory_tokens=None, - tie_embedding=False, - use_pos_emb=True - ): - super().__init__() - assert isinstance(attn_layers, AttentionLayers), 'attention layers must be one of Encoder or Decoder' - - dim = attn_layers.dim - emb_dim = default(emb_dim, dim) - - self.max_seq_len = max_seq_len - self.max_mem_len = max_mem_len - self.num_tokens = num_tokens - - self.token_emb = nn.Embedding(num_tokens, emb_dim) - self.pos_emb = AbsolutePositionalEmbedding(emb_dim, max_seq_len) if ( - use_pos_emb and not attn_layers.has_pos_emb) else always(0) - self.emb_dropout = nn.Dropout(emb_dropout) - - self.project_emb = nn.Linear(emb_dim, dim) if emb_dim != dim else nn.Identity() - self.attn_layers = attn_layers - self.norm = nn.LayerNorm(dim) - - self.init_() - - self.to_logits = nn.Linear(dim, num_tokens) if not tie_embedding else lambda t: t @ self.token_emb.weight.t() - - # memory tokens (like [cls]) from Memory Transformers paper - num_memory_tokens = default(num_memory_tokens, 0) - self.num_memory_tokens = num_memory_tokens - if num_memory_tokens > 0: - self.memory_tokens = nn.Parameter(torch.randn(num_memory_tokens, dim)) - - # let funnel encoder know number of memory tokens, if specified - if hasattr(attn_layers, 'num_memory_tokens'): - attn_layers.num_memory_tokens = num_memory_tokens - - def init_(self): - nn.init.normal_(self.token_emb.weight, std=0.02) - - def forward( - self, - x, - return_embeddings=False, - mask=None, - return_mems=False, - return_attn=False, - mems=None, - **kwargs - ): - b, n, device, num_mem = *x.shape, x.device, self.num_memory_tokens - x = self.token_emb(x) - x += self.pos_emb(x) - x = self.emb_dropout(x) - - x = self.project_emb(x) - - if num_mem > 0: - mem = repeat(self.memory_tokens, 'n d -> b n d', b=b) - x = torch.cat((mem, x), dim=1) - - # auto-handle masking after appending memory tokens - if exists(mask): - mask = F.pad(mask, (num_mem, 0), value=True) - - x, intermediates = self.attn_layers(x, mask=mask, mems=mems, return_hiddens=True, **kwargs) - x = self.norm(x) - - mem, x = x[:, :num_mem], x[:, num_mem:] - - out = self.to_logits(x) if not return_embeddings else x - - if return_mems: - hiddens = intermediates.hiddens - new_mems = list(map(lambda pair: torch.cat(pair, dim=-2), zip(mems, hiddens))) if exists(mems) else hiddens - new_mems = list(map(lambda t: t[..., -self.max_mem_len:, :].detach(), new_mems)) - return out, new_mems - - if return_attn: - attn_maps = list(map(lambda t: t.post_softmax_attn, intermediates.attn_intermediates)) - return out, attn_maps - - return out \ No newline at end of file diff --git a/py/dynamiCrafter/utils/model_utils.py b/py/dynamiCrafter/utils/model_utils.py deleted file mode 100644 index 72d2c40..0000000 --- a/py/dynamiCrafter/utils/model_utils.py +++ /dev/null @@ -1,146 +0,0 @@ - -import torch - -from collections import OrderedDict - -from comfy import model_base -from comfy import utils -from comfy import diffusers_convert - -try: - import comfy.text_encoders.sd2_clip -except ImportError: - from comfy import sd2_clip - -from comfy import supported_models_base -from comfy import latent_formats - -from ..lvdm.modules.encoders.resampler import Resampler - -DYNAMICRAFTER_CONFIG = { - 'in_channels': 8, - 'out_channels': 4, - 'model_channels': 320, - 'attention_resolutions': [4, 2, 1], - 'num_res_blocks': 2, - 'channel_mult': [1, 2, 4, 4], - 'num_head_channels': 64, - 'transformer_depth': 1, - 'context_dim': 1024, - 'use_linear': True, - 'use_checkpoint': False, - 'temporal_conv': True, - 'temporal_attention': True, - 'temporal_selfatt_only': True, - 'use_relative_position': False, - 'use_causal_attention': False, - 'temporal_length': 16, - 'addition_attention': True, - 'image_cross_attention': True, - 'image_cross_attention_scale_learnable': True, - 'default_fs': 3, - 'fs_condition': True -} - -IMAGE_PROJ_CONFIG = { - "dim": 1024, - "depth": 4, - "dim_head": 64, - "heads": 12, - "num_queries": 16, - "embedding_dim": 1280, - "output_dim": 1024, - "ff_mult": 4, - "video_length": 16 -} - -def process_list_or_str(target_key_or_keys, k): - if isinstance(target_key_or_keys, list): - return any([list_k in k for list_k in target_key_or_keys]) - else: - return target_key_or_keys in k - -def simple_state_dict_loader(state_dict: dict, target_key: str, target_dict: dict = None): - out_dict = {} - - if target_dict is None: - for k, v in state_dict.items(): - if process_list_or_str(target_key, k): - out_dict[k] = v - else: - for k, v in target_dict.items(): - out_dict[k] = state_dict[k] - - return out_dict - -def load_image_proj_dict(state_dict: dict): - return simple_state_dict_loader(state_dict, 'image_proj') - -def load_dynamicrafter_dict(state_dict: dict): - return simple_state_dict_loader(state_dict, 'model.diffusion_model') - -def load_vae_dict(state_dict: dict): - return simple_state_dict_loader(state_dict, 'first_stage_model') - -def get_base_model(state_dict: dict, version_checker=False): - - is_256_model = False - - for k in state_dict.keys(): - if "framestride_embed" in k: - is_256_model = True - break - -def get_image_proj_model(state_dict: dict): - - state_dict = {k.replace('image_proj_model.', ''): v for k, v in state_dict.items()} - #target_dict = Resampler().state_dict() - - ImageProjModel = Resampler(**IMAGE_PROJ_CONFIG) - ImageProjModel.load_state_dict(state_dict) - - print("Image Projection Model loaded successfully") - #del target_dict - return ImageProjModel - -class DynamiCrafterBase(supported_models_base.BASE): - unet_config = {} - unet_extra_config = {} - - latent_format = latent_formats.SD15 - - def process_clip_state_dict(self, state_dict): - replace_prefix = {} - replace_prefix["conditioner.embedders.0.model."] = "clip_h." #SD2 in sgm format - replace_prefix["cond_stage_model.model."] = "clip_h." - state_dict = utils.state_dict_prefix_replace(state_dict, replace_prefix, filter_keys=True) - state_dict = utils.clip_text_transformers_convert(state_dict, "clip_h.", "clip_h.transformer.") - return state_dict - - def process_clip_state_dict_for_saving(self, state_dict): - replace_prefix = {} - replace_prefix["clip_h"] = "cond_stage_model.model" - state_dict = utils.state_dict_prefix_replace(state_dict, replace_prefix) - state_dict = diffusers_convert.convert_text_enc_state_dict_v20(state_dict) - return state_dict - - def clip_target(self): - return supported_models_base.ClipTarget(sd2_clip.SD2Tokenizer, sd2_clip.SD2ClipModel) - - def process_dict_version(self, state_dict: dict): - processed_dict = OrderedDict() - is_eps = False - - for k in list(state_dict.keys()): - if "framestride_embed" in k: - new_key = k.replace("framestride_embed", "fps_embedding") - processed_dict[new_key] = state_dict[k] - is_eps = True - continue - - processed_dict[k] = state_dict[k] - - return processed_dict, is_eps - - - diff --git a/py/dynamiCrafter/utils/utils.py b/py/dynamiCrafter/utils/utils.py deleted file mode 100644 index 0e23b25..0000000 --- a/py/dynamiCrafter/utils/utils.py +++ /dev/null @@ -1,82 +0,0 @@ -import importlib -import numpy as np -import cv2 -import torch -import torch.distributed as dist - -MODEL_EXTS = ['ckpt', 'safetensors', 'bin'] - -def get_models_directory(directory: list): - files_list = list(filter(lambda f: f.split(".")[-1] in MODEL_EXTS, directory)) - return files_list - -def count_params(model, verbose=False): - total_params = sum(p.numel() for p in model.parameters()) - if verbose: - print(f"{model.__class__.__name__} has {total_params*1.e-6:.2f} M params.") - return total_params - - -def check_istarget(name, para_list): - """ - name: full name of source para - para_list: partial name of target para - """ - istarget=False - for para in para_list: - if para in name: - return True - return istarget - - -def instantiate_from_config(config): - if not "target" in config: - if config == '__is_first_stage__': - return None - elif config == "__is_unconditional__": - return None - raise KeyError("Expected key `target` to instantiate.") - return get_obj_from_str(config["target"])(**config.get("params", dict())) - - -def get_obj_from_str(string, reload=False): - module, cls = string.rsplit(".", 1) - if reload: - module_imp = importlib.import_module(module) - importlib.reload(module_imp) - return getattr(importlib.import_module(module, package=None), cls) - - -def load_npz_from_dir(data_dir): - data = [np.load(os.path.join(data_dir, data_name))['arr_0'] for data_name in os.listdir(data_dir)] - data = np.concatenate(data, axis=0) - return data - - -def load_npz_from_paths(data_paths): - data = [np.load(data_path)['arr_0'] for data_path in data_paths] - data = np.concatenate(data, axis=0) - return data - - -def resize_numpy_image(image, max_resolution=512 * 512, resize_short_edge=None): - h, w = image.shape[:2] - if resize_short_edge is not None: - k = resize_short_edge / min(h, w) - else: - k = max_resolution / (h * w) - k = k**0.5 - h = int(np.round(h * k / 64)) * 64 - w = int(np.round(w * k / 64)) * 64 - image = cv2.resize(image, (w, h), interpolation=cv2.INTER_LANCZOS4) - return image - - -def setup_dist(args): - if dist.is_initialized(): - return - torch.cuda.set_device(args.local_rank) - torch.distributed.init_process_group( - 'nccl', - init_method='env://' - ) \ No newline at end of file diff --git a/py/easyNodes.py b/py/easyNodes.py index e7cbd07..44ae925 100644 --- a/py/easyNodes.py +++ b/py/easyNodes.py @@ -572,217 +572,6 @@ class portraitMaster: # ---------------------------------------------------------------提示词 结束----------------------------------------------------------------------# -# ---------------------------------------------------------------潜空间 开始----------------------------------------------------------------------# -# 潜空间sigma相乘 -class latentNoisy: - @classmethod - def INPUT_TYPES(s): - return {"required": { - "sampler_name": (comfy.samplers.KSampler.SAMPLERS,), - "scheduler": (comfy.samplers.KSampler.SCHEDULERS,), - "steps": ("INT", {"default": 10000, "min": 0, "max": 10000}), - "start_at_step": ("INT", {"default": 0, "min": 0, "max": 10000}), - "end_at_step": ("INT", {"default": 10000, "min": 1, "max": 10000}), - "source": (["CPU", "GPU"],), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), - }, - "optional": { - "pipe": ("PIPE_LINE",), - "optional_model": ("MODEL",), - "optional_latent": ("LATENT",) - }} - - RETURN_TYPES = ("PIPE_LINE", "LATENT", "FLOAT",) - RETURN_NAMES = ("pipe", "latent", "sigma",) - FUNCTION = "run" - - CATEGORY = "EasyUse/Latent" - - def run(self, sampler_name, scheduler, steps, start_at_step, end_at_step, source, seed, pipe=None, optional_model=None, optional_latent=None): - model = optional_model if optional_model is not None else pipe["model"] - batch_size = pipe["loader_settings"]["batch_size"] - empty_latent_height = pipe["loader_settings"]["empty_latent_height"] - empty_latent_width = pipe["loader_settings"]["empty_latent_width"] - - if optional_latent is not None: - samples = optional_latent - else: - torch.manual_seed(seed) - if source == "CPU": - device = "cpu" - else: - device = comfy.model_management.get_torch_device() - noise = torch.randn((batch_size, 4, empty_latent_height // 8, empty_latent_width // 8), dtype=torch.float32, - device=device).cpu() - - samples = {"samples": noise} - - device = comfy.model_management.get_torch_device() - end_at_step = min(steps, end_at_step) - start_at_step = min(start_at_step, end_at_step) - comfy.model_management.load_model_gpu(model) - model_patcher = comfy.model_patcher.ModelPatcher(model.model, load_device=device, offload_device=comfy.model_management.unet_offload_device()) - sampler = comfy.samplers.KSampler(model_patcher, steps=steps, device=device, sampler=sampler_name, - scheduler=scheduler, denoise=1.0, model_options=model.model_options) - sigmas = sampler.sigmas - sigma = sigmas[start_at_step] - sigmas[end_at_step] - sigma /= model.model.latent_format.scale_factor - sigma = sigma.cpu().numpy() - - samples_out = samples.copy() - - s1 = samples["samples"] - samples_out["samples"] = s1 * sigma - - if pipe is None: - pipe = {} - new_pipe = { - **pipe, - "samples": samples_out - } - del pipe - - return (new_pipe, samples_out, sigma) - -# Latent遮罩复合 -class latentCompositeMaskedWithCond: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "pipe": ("PIPE_LINE",), - "text_combine": ("LIST",), - "source_latent": ("LATENT",), - "source_mask": ("MASK",), - "destination_mask": ("MASK",), - "text_combine_mode": (["add", "replace", "cover"], {"default": "add"}), - "replace_text": ("STRING", {"default": ""}) - }, - "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO", "my_unique_id": "UNIQUE_ID"}, - } - - OUTPUT_IS_LIST = (False, False, True) - RETURN_TYPES = ("PIPE_LINE", "LATENT", "CONDITIONING") - RETURN_NAMES = ("pipe", "latent", "conditioning",) - FUNCTION = "run" - - CATEGORY = "EasyUse/Latent" - - def run(self, pipe, text_combine, source_latent, source_mask, destination_mask, text_combine_mode, replace_text, prompt=None, extra_pnginfo=None, my_unique_id=None): - positive = None - clip = pipe["clip"] - destination_latent = pipe["samples"] - - conds = [] - - for text in text_combine: - if text_combine_mode == 'cover': - positive = text - elif text_combine_mode == 'replace' and replace_text != '': - positive = pipe["loader_settings"]["positive"].replace(replace_text, text) - else: - positive = pipe["loader_settings"]["positive"] + ',' + text - positive_token_normalization = pipe["loader_settings"]["positive_token_normalization"] - positive_weight_interpretation = pipe["loader_settings"]["positive_weight_interpretation"] - a1111_prompt_style = pipe["loader_settings"]["a1111_prompt_style"] - positive_cond = pipe["positive"] - - log_node_warn("Positive encoding...") - steps = pipe["loader_settings"]["steps"] if "steps" in pipe["loader_settings"] else 1 - positive_embeddings_final = advanced_encode(clip, positive, - positive_token_normalization, - positive_weight_interpretation, w_max=1.0, - apply_to_pooled='enable', a1111_prompt_style=a1111_prompt_style, steps=steps) - - # source cond - (cond_1,) = ConditioningSetMask().append(positive_cond, source_mask, "default", 1) - (cond_2,) = ConditioningSetMask().append(positive_embeddings_final, destination_mask, "default", 1) - positive_cond = cond_1 + cond_2 - - conds.append(positive_cond) - # latent composite masked - (samples,) = LatentCompositeMasked().composite(destination_latent, source_latent, 0, 0, False) - - new_pipe = { - **pipe, - "samples": samples, - "loader_settings": { - **pipe["loader_settings"], - "positive": positive, - } - } - - del pipe - - return (new_pipe, samples, conds) - -# 噪声注入到潜空间 -class injectNoiseToLatent: - @classmethod - def INPUT_TYPES(s): - return {"required": { - "strength": ("FLOAT", {"default": 0.1, "min": 0.0, "max": 200.0, "step": 0.0001}), - "normalize": ("BOOLEAN", {"default": False}), - "average": ("BOOLEAN", {"default": False}), - }, - "optional": { - "pipe_to_noise": ("PIPE_LINE",), - "image_to_latent": ("IMAGE",), - "latent": ("LATENT",), - "noise": ("LATENT",), - "mask": ("MASK",), - "mix_randn_amount": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1000.0, "step": 0.001}), - "seed": ("INT", {"default": 123, "min": 0, "max": 0xffffffffffffffff, "step": 1}), - } - } - - RETURN_TYPES = ("LATENT",) - FUNCTION = "inject" - CATEGORY = "EasyUse/Latent" - - def inject(self,strength, normalize, average, pipe_to_noise=None, noise=None, image_to_latent=None, latent=None, mix_randn_amount=0, mask=None, seed=None): - - vae = pipe_to_noise["vae"] if pipe_to_noise is not None else pipe_to_noise["vae"] - batch_size = pipe_to_noise["loader_settings"]["batch_size"] if pipe_to_noise is not None and "batch_size" in pipe_to_noise["loader_settings"] else 1 - if noise is None and pipe_to_noise is not None: - noise = pipe_to_noise["samples"] - elif noise is None: - raise Exception("InjectNoiseToLatent: No noise provided") - - if image_to_latent is not None and vae is not None: - samples = {"samples": vae.encode(image_to_latent[:, :, :, :3])} - latents = RepeatLatentBatch().repeat(samples, batch_size)[0] - elif latent is not None: - latents = latent - else: - latents = {"samples": noise["samples"].clone()} - - samples = latents.copy() - if latents["samples"].shape != noise["samples"].shape: - raise ValueError("InjectNoiseToLatent: Latent and noise must have the same shape") - if average: - noised = (samples["samples"].clone() + noise["samples"].clone()) / 2 - else: - noised = samples["samples"].clone() + noise["samples"].clone() * strength - if normalize: - noised = noised / noised.std() - if mask is not None: - mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), - size=(noised.shape[2], noised.shape[3]), mode="bilinear") - mask = mask.expand((-1, noised.shape[1], -1, -1)) - if mask.shape[0] < noised.shape[0]: - mask = mask.repeat((noised.shape[0] - 1) // mask.shape[0] + 1, 1, 1, 1)[:noised.shape[0]] - noised = mask * noised + (1 - mask) * latents["samples"] - if mix_randn_amount > 0: - if seed is not None: - torch.manual_seed(seed) - rand_noise = torch.randn_like(noised) - noised = ((1 - mix_randn_amount) * noised + mix_randn_amount * - rand_noise) / ((mix_randn_amount ** 2 + (1 - mix_randn_amount) ** 2) ** 0.5) - samples["samples"] = noised - return (samples,) - -# ---------------------------------------------------------------潜空间 结束----------------------------------------------------------------------# # ---------------------------------------------------------------随机种 开始----------------------------------------------------------------------# # 随机种 @@ -1646,169 +1435,6 @@ class svdLoader: return (pipe, model, vae) -#dynamiCrafter加载器 -from .dynamiCrafter import DynamiCrafter -class dynamiCrafterLoader(DynamiCrafter): - - def __init__(self): - super().__init__() - - @classmethod - def INPUT_TYPES(cls): - - return {"required": { - "model_name": (list(DYNAMICRAFTER_MODELS.keys()),), - "clip_skip": ("INT", {"default": -2, "min": -24, "max": 0, "step": 1}), - - "init_image": ("IMAGE",), - "resolution": (resolution_strings, {"default": "512 x 512"}), - "empty_latent_width": ("INT", {"default": 256, "min": 16, "max": MAX_RESOLUTION, "step": 8}), - "empty_latent_height": ("INT", {"default": 256, "min": 16, "max": MAX_RESOLUTION, "step": 8}), - - "positive": ("STRING", {"default": "", "multiline": True}), - "negative": ("STRING", {"default": "", "multiline": True}), - - "use_interpolate": ("BOOLEAN", {"default": False}), - "fps": ("INT", {"default": 15, "min": 1, "max": 30, "step": 1},), - "frames": ("INT", {"default": 16}), - "scale_latents": ("BOOLEAN", {"default": False}) - }, - "optional": { - "optional_vae": ("VAE",), - }, - "hidden": {"prompt": "PROMPT", "my_unique_id": "UNIQUE_ID"} - } - - RETURN_TYPES = ("PIPE_LINE", "MODEL", "VAE") - RETURN_NAMES = ("pipe", "model", "vae") - - FUNCTION = "adv_pipeloader" - CATEGORY = "EasyUse/Loaders" - - def get_clip_file(self, node_name): - clip_list = folder_paths.get_filename_list("clip") - pattern = 'sd2-1-open-clip|model.(safetensors|bin)$' - clip_files = [e for e in clip_list if re.search(pattern, e, re.IGNORECASE)] - - clip_name = clip_files[0] if len(clip_files)>0 else None - clip_file = folder_paths.get_full_path("clip", clip_name) if clip_name else None - if clip_name is not None: - log_node_info(node_name, f"Using {clip_name}") - - return clip_file, clip_name - - def get_clipvision_file(self, node_name): - clipvision_list = folder_paths.get_filename_list("clip_vision") - pattern = '(ViT.H.14.*s32B.b79K|ipadapter.*sd15|sd1.?5.*model|open_clip_pytorch_model.(bin|safetensors))' - clipvision_files = [e for e in clipvision_list if re.search(pattern, e, re.IGNORECASE)] - - clipvision_name = clipvision_files[0] if len(clipvision_files)>0 else None - clipvision_file = folder_paths.get_full_path("clip_vision", clipvision_name) if clipvision_name else None - if clipvision_name is not None: - log_node_info(node_name, f"Using {clipvision_name}") - - return clipvision_file, clipvision_name - - def get_vae_file(self, node_name): - vae_list = folder_paths.get_filename_list("vae") - pattern = 'vae-ft-mse-840000-ema-pruned.(pt|bin|safetensors)$' - vae_files = [e for e in vae_list if re.search(pattern, e, re.IGNORECASE)] - - vae_name = vae_files[0] if len(vae_files)>0 else None - vae_file = folder_paths.get_full_path("vae", vae_name) if vae_name else None - if vae_name is not None: - log_node_info(node_name, f"Using {vae_name}") - - return vae_file, vae_name - - def adv_pipeloader(self, model_name, clip_skip, init_image, resolution, empty_latent_width, empty_latent_height, positive, negative, use_interpolate, fps, frames, scale_latents, optional_vae=None, prompt=None, my_unique_id=None): - positive_embeddings_final, negative_embeddings_final = None, None - # resolution - if resolution != "自定义 x 自定义": - try: - width, height = map(int, resolution.split(' x ')) - empty_latent_width = width - empty_latent_height = height - except ValueError: - raise ValueError("Invalid base_resolution format.") - - # Clean models from loaded_objects - easyCache.update_loaded_objects(prompt) - - models_0 = list(DYNAMICRAFTER_MODELS.keys())[0] - - if optional_vae: - vae = optional_vae - vae_name = None - else: - vae_file, vae_name = self.get_vae_file("easy dynamiCrafterLoader") - if vae_file is None: - vae_name = "vae-ft-mse-840000-ema-pruned.safetensors" - get_local_filepath(DYNAMICRAFTER_MODELS[models_0]['vae_url'], os.path.join(folder_paths.models_dir, "vae"), - vae_name) - vae = easyCache.load_vae(vae_name) - - clip_file, clip_name = self.get_clip_file("easy dynamiCrafterLoader") - if clip_file is None: - clip_name = 'sd2-1-open-clip.safetensors' - get_local_filepath(DYNAMICRAFTER_MODELS[models_0]['clip_url'], os.path.join(folder_paths.models_dir, "clip"), - clip_name) - - clip = easyCache.load_clip(clip_name) - # load clip vision - clip_vision_file, clip_vision_name = self.get_clipvision_file("easy dynamiCrafterLoader") - if clip_vision_file is None: - clip_vision_name = 'CLIP-ViT-H-14-laion2B-s32B-b79K.safetensors' - clip_vision_file = get_local_filepath(DYNAMICRAFTER_MODELS[models_0]['clip_vision_url'], os.path.join(folder_paths.models_dir, "clip_vision"), - clip_vision_name) - clip_vision = load_clip_vision(clip_vision_file) - # load unet model - model_path = get_local_filepath(DYNAMICRAFTER_MODELS[model_name]['model_url'], DYNAMICRAFTER_DIR) - model_patcher, image_proj_model = self.load_dynamicrafter(model_path) - - # apply - model, empty_latent, image_latent = self.process_image_conditioning(model_patcher, clip_vision, vae, image_proj_model, init_image, use_interpolate, fps, frames, scale_latents) - - clipped = clip.clone() - if clip_skip != 0: - clipped.clip_layer(clip_skip) - - if positive is not None and positive != '': - if has_chinese(positive): - positive = zh_to_en([positive])[0] - positive_embeddings_final, = CLIPTextEncode().encode(clipped, positive) - if negative is not None and negative != '': - if has_chinese(negative): - negative = zh_to_en([negative])[0] - negative_embeddings_final, = CLIPTextEncode().encode(clipped, negative) - - image = easySampler.pil2tensor(Image.new('RGB', (1, 1), (0, 0, 0))) - - pipe = {"model": model, - "positive": positive_embeddings_final, - "negative": negative_embeddings_final, - "vae": vae, - "clip": clip, - "clip_vision": clip_vision, - - "samples": empty_latent, - "images": image, - "seed": 0, - - "loader_settings": {"ckpt_name": model_name, - "vae_name": vae_name, - - "positive": positive, - "negative": negative, - "resolution": resolution, - "empty_latent_width": empty_latent_width, - "empty_latent_height": empty_latent_height, - "batch_size": 1, - "seed": 0, - } - } - - return (pipe, model, vae) # kolors Loader from .kolors.text_encode import chatglm3_adv_text_encode @@ -7760,37 +7386,6 @@ class pipeXYPlotAdvanced: #---------------------------------------------------------------节点束 结束---------------------------------------------------------------------- -# 显示推理时间 -class showSpentTime: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "pipe": ("PIPE_LINE",), - "spent_time": ("INFO", {"default": 'Time will be displayed when reasoning is complete', "forceInput": False}), - }, - "hidden": { - "unique_id": "UNIQUE_ID", - "extra_pnginfo": "EXTRA_PNGINFO", - }, - } - - FUNCTION = "notify" - OUTPUT_NODE = True - RETURN_TYPES = () - RETURN_NAMES = () - - CATEGORY = "EasyUse/Util" - - def notify(self, pipe, spent_time=None, unique_id=None, extra_pnginfo=None): - if unique_id and extra_pnginfo and "workflow" in extra_pnginfo: - workflow = extra_pnginfo["workflow"] - node = next((x for x in workflow["nodes"] if str(x["id"]) == unique_id), None) - if node: - spent_time = pipe['loader_settings']['spent_time'] if 'spent_time' in pipe['loader_settings'] else '' - node["widgets_values"] = [spent_time] - - return {"ui": {"text": spent_time}, "result": {}} # 显示加载器参数中的各种名称 class showLoaderSettingsNames: @@ -7964,7 +7559,6 @@ NODE_CLASS_MAPPINGS = { "easy svdLoader": svdLoader, "easy sv3dLoader": sv3DLoader, "easy zero123Loader": zero123Loader, - "easy dynamiCrafterLoader": dynamiCrafterLoader, "easy cascadeLoader": cascadeLoader, "easy kolorsLoader": kolorsLoader, "easy fluxLoader": fluxLoader, @@ -7993,16 +7587,11 @@ NODE_CLASS_MAPPINGS = { "easy pulIDApplyADV": applyPulIDADV, "easy styleAlignedBatchAlign": styleAlignedBatchAlign, "easy icLightApply": icLightApply, - # "easy ominiControlApply": applyOminiControl, # Inpaint 内补 "easy applyFooocusInpaint": applyFooocusInpaint, "easy applyBrushNet": applyBrushNet, "easy applyPowerPaint": applyPowerPaint, "easy applyInpaint": applyInpaint, - # latent 潜空间 - "easy latentNoisy": latentNoisy, - "easy latentCompositeMaskedWithCond": latentCompositeMaskedWithCond, - "easy injectNoiseToLatent": injectNoiseToLatent, # preSampling 预采样处理 "easy preSampling": samplerSettings, "easy preSamplingAdvanced": samplerSettingsAdvanced, @@ -8058,7 +7647,6 @@ NODE_CLASS_MAPPINGS = { "easy XYInputs: NegativeCond": XYplot_Negative_Cond, "easy XYInputs: NegativeCondList": XYplot_Negative_Cond_List, # others 其他 - "easy showSpentTime": showSpentTime, "easy showLoaderSettingsNames": showLoaderSettingsNames, "easy sliderControl": sliderControl, "dynamicThresholdingFull": dynamicThresholdingFull, @@ -8092,7 +7680,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "easy svdLoader": "EasyLoader (SVD)", "easy sv3dLoader": "EasyLoader (SV3D)", "easy zero123Loader": "EasyLoader (Zero123)", - "easy dynamiCrafterLoader": "EasyLoader (DynamiCrafter)", "easy cascadeLoader": "EasyCascadeLoader", "easy kolorsLoader": "EasyLoader (Kolors)", "easy fluxLoader": "EasyLoader (Flux)", @@ -8122,16 +7709,11 @@ NODE_DISPLAY_NAME_MAPPINGS = { "easy pulIDApplyADV": "Easy Apply PuLID (Advanced)", "easy styleAlignedBatchAlign": "Easy Apply StyleAlign", "easy icLightApply": "Easy Apply ICLight", - "easy ominiControlApply": "Easy Apply OminiContol", # Inpaint 内补 "easy applyFooocusInpaint": "Easy Apply Fooocus Inpaint", "easy applyBrushNet": "Easy Apply BrushNet", "easy applyPowerPaint": "Easy Apply PowerPaint", "easy applyInpaint": "Easy Apply Inpaint", - # latent 潜空间 - "easy latentNoisy": "LatentNoisy", - "easy latentCompositeMaskedWithCond": "LatentCompositeMaskedWithCond", - "easy injectNoiseToLatent": "InjectNoiseToLatent", # preSampling 预采样处理 "easy preSampling": "PreSampling", "easy preSamplingAdvanced": "PreSampling (Advanced)", @@ -8187,7 +7769,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "easy XYInputs: NegativeCond": "XY Inputs: NegCond //EasyUse", "easy XYInputs: NegativeCondList": "XY Inputs: NegCondList //EasyUse", # others 其他 - "easy showSpentTime": "Show Spent Time", "easy showLoaderSettingsNames": "Show Loader Settings Names", "easy sliderControl": "Easy Slider Control", "dynamicThresholdingFull": "DynamicThresholdingFull", diff --git a/py/logic.py b/py/logic.py index 1eb433e..f14b714 100755 --- a/py/logic.py +++ b/py/logic.py @@ -1499,88 +1499,6 @@ class clearCacheAll: return (anything,) -# Deprecated -class If: - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "any": (any_type,), - "if": (any_type,), - "else": (any_type,), - }, - } - - RETURN_TYPES = (any_type,) - RETURN_NAMES = ("?",) - FUNCTION = "execute" - CATEGORY = "EasyUse/🚫 Deprecated" - DEPRECATED = True - - def execute(self, *args, **kwargs): - return (kwargs['if'] if kwargs['any'] else kwargs['else'],) - - -class poseEditor: - @classmethod - def INPUT_TYPES(s): - return {"required": { - "image": ("STRING", {"default": ""}) - }} - - FUNCTION = "output_pose" - CATEGORY = "EasyUse/🚫 Deprecated" - DEPRECATED = True - RETURN_TYPES = () - RETURN_NAMES = () - - def output_pose(self, image): - return () - - -class imageToMask: - @classmethod - def INPUT_TYPES(s): - return {"required": { - "image": ("IMAGE",), - "channel": (['red', 'green', 'blue'],), - } - } - - RETURN_TYPES = ("MASK",) - FUNCTION = "convert" - CATEGORY = "EasyUse/🚫 Deprecated" - DEPRECATED = True - - def convert_to_single_channel(self, image, channel='red'): - from PIL import Image - # Convert to RGB mode to access individual channels - image = image.convert('RGB') - - # Extract the desired channel and convert to greyscale - if channel == 'red': - channel_img = image.split()[0].convert('L') - elif channel == 'green': - channel_img = image.split()[1].convert('L') - elif channel == 'blue': - channel_img = image.split()[2].convert('L') - else: - raise ValueError( - "Invalid channel option. Please choose 'red', 'green', or 'blue'.") - - # Convert the greyscale channel back to RGB mode - channel_img = Image.merge( - 'RGB', (channel_img, channel_img, channel_img)) - - return channel_img - - def convert(self, image, channel='red'): - from .libs.image import pil2tensor, tensor2pil - image = self.convert_to_single_channel(tensor2pil(image), channel) - image = pil2tensor(image) - return (image.squeeze().mean(2),) - - class saveText: def __init__(self): @@ -1827,10 +1745,7 @@ NODE_CLASS_MAPPINGS = { "easy cleanGpuUsed": cleanGPUUsed, "easy saveText": saveText, "easy saveTextLazy": saveTextLazy, - "easy sleep": sleep, - "easy if": If, - "easy poseEditor": poseEditor, - "easy imageToMask": imageToMask, + "easy sleep": sleep } NODE_DISPLAY_NAME_MAPPINGS = { "easy string": "String", @@ -1877,7 +1792,4 @@ NODE_DISPLAY_NAME_MAPPINGS = { "easy saveText": "Save Text", "easy saveTextLazy": "Save Text (Lazy)", "easy sleep": "Sleep", - "easy if": "If (🚫Deprecated)", - "easy poseEditor": "PoseEditor (🚫Deprecated)", - "easy imageToMask": "ImageToMask (🚫Deprecated)" }