diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..bb46937 --- /dev/null +++ b/.gitignore @@ -0,0 +1,145 @@ +.idea/ + +training/ +lightning_logs/ +image_log/ + +*.pth +*.ckpt +*.safetensors + +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +pip-wheel-metadata/ +share/python-wheels/ +*.egg-info/ +.installed.cfg +*.egg +MANIFEST + +# PyInstaller +# Usually these files are written by a python script from a template +# before PyInstaller builds the exe, so as to inject date/other infos into it. +*.manifest +*.spec + +# Installer logs +pip-log.txt +pip-delete-this-directory.txt + +# Unit test / coverage reports +htmlcov/ +.tox/ +.nox/ +.coverage +.coverage.* +.cache +nosetests.xml +coverage.xml +*.cover +*.py,cover +.hypothesis/ +.pytest_cache/ + +# Translations +*.mo +*.pot + +# Django stuff: +*.log +local_settings.py +db.sqlite3 +db.sqlite3-journal + +# Flask stuff: +instance/ +.webassets-cache + +# Scrapy stuff: +.scrapy + +# Sphinx documentation +docs/_build/ + +# PyBuilder +target/ + +# Jupyter Notebook +.ipynb_checkpoints + +# IPython +profile_default/ +ipython_config.py + +# pyenv +.python-version + +# pipenv +# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. +# However, in case of collaboration, if having platform-specific dependencies or dependencies +# having no cross-platform support, pipenv may install dependencies that don't work, or not +# install all needed dependencies. +#Pipfile.lock + +# PEP 582; used by e.g. github.com/David-OConnor/pyflow +__pypackages__/ + +# Celery stuff +celerybeat-schedule +celerybeat.pid + +# SageMath parsed files +*.sage.py + +# Environments +.env +.venv +env/ +venv/ +ENV/ +env.bak/ +venv.bak/ + +# Spyder project settings +.spyderproject +.spyproject + +# Rope project settings +.ropeproject + +# mkdocs documentation +/site + +# mypy +.mypy_cache/ +.dmypy.json +dmypy.json + +# Pyre type checker +.pyre/ + + +*.safetensors +*.ckpt + +checkpoints \ No newline at end of file diff --git a/README.md b/README.md new file mode 100644 index 0000000..589c298 --- /dev/null +++ b/README.md @@ -0,0 +1,3 @@ +# ComfyUI nodes to use Lotus depth/normal prediction + +Original repo: https://github.com/EnVision-Research/Lotus \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..2e96bd6 --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/configs/lotus_unet_config.json b/configs/lotus_unet_config.json new file mode 100644 index 0000000..ef720ef --- /dev/null +++ b/configs/lotus_unet_config.json @@ -0,0 +1,73 @@ +{ + "_class_name": "UNet2DConditionModel", + "_diffusers_version": "0.28.0.dev0", + "_name_or_path": "../Lotus-weights/lotus-depth-d-v1-1/unet", + "act_fn": "silu", + "addition_embed_type": null, + "addition_embed_type_num_heads": 64, + "addition_time_embed_dim": null, + "attention_head_dim": [ + 5, + 10, + 20, + 20 + ], + "attention_type": "default", + "block_out_channels": [ + 320, + 640, + 1280, + 1280 + ], + "center_input_sample": false, + "class_embed_type": "projection", + "class_embeddings_concat": false, + "conv_in_kernel": 3, + "conv_out_kernel": 3, + "cross_attention_dim": 1024, + "cross_attention_norm": null, + "down_block_types": [ + "CrossAttnDownBlock2D", + "CrossAttnDownBlock2D", + "CrossAttnDownBlock2D", + "DownBlock2D" + ], + "downsample_padding": 1, + "dropout": 0.0, + "dual_cross_attention": false, + "encoder_hid_dim": null, + "encoder_hid_dim_type": null, + "flip_sin_to_cos": true, + "freq_shift": 0, + "in_channels": 4, + "layers_per_block": 2, + "mid_block_only_cross_attention": null, + "mid_block_scale_factor": 1, + "mid_block_type": "UNetMidBlock2DCrossAttn", + "norm_eps": 1e-05, + "norm_num_groups": 32, + "num_attention_heads": null, + "num_class_embeds": null, + "only_cross_attention": false, + "out_channels": 4, + "projection_class_embeddings_input_dim": 4, + "resnet_out_scale_factor": 1.0, + "resnet_skip_time_act": false, + "resnet_time_scale_shift": "default", + "reverse_transformer_layers_per_block": null, + "sample_size": 64, + "time_cond_proj_dim": null, + "time_embedding_act_fn": null, + "time_embedding_dim": null, + "time_embedding_type": "positional", + "timestep_post_act": null, + "transformer_layers_per_block": 1, + "up_block_types": [ + "UpBlock2D", + "CrossAttnUpBlock2D", + "CrossAttnUpBlock2D", + "CrossAttnUpBlock2D" + ], + "upcast_attention": false, + "use_linear_projection": true + } \ No newline at end of file diff --git a/empty_text_embed.pt b/empty_text_embed.pt new file mode 100755 index 0000000..dcedcde Binary files /dev/null and b/empty_text_embed.pt differ diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..1cf81dd --- /dev/null +++ b/nodes.py @@ -0,0 +1,146 @@ +import os +import torch +import folder_paths +import comfy.model_management as mm +from comfy.utils import load_torch_file, ProgressBar + +import logging +import json +from diffusers.models import UNet2DConditionModel + +logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s') +log = logging.getLogger(__name__) + +script_directory = os.path.dirname(os.path.abspath(__file__)) + +class LoadLotusModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": (folder_paths.get_filename_list("diffusion_models"), ), + }, + "optional": { + "precision": (["fp16", "fp32", "bf16"], + {"default": "fp16"} + ), + } + } + + RETURN_TYPES = ("LOTUSUNET",) + RETURN_NAMES = ("lotus_unet", ) + FUNCTION = "loadmodel" + CATEGORY = "ComfyUI-Lotus" + + def loadmodel(self, model, precision): + + dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision] + mm.soft_empty_cache() + + lotus_model_path = folder_paths.get_full_path_or_raise("diffusion_models", model) + + lotus_sd = load_torch_file(lotus_model_path) + in_channels = lotus_sd['conv_in.weight'].shape[1] + lotus_config = os.path.join(script_directory, "configs", "lotus_unet_config.json") + + with open(lotus_config, 'r') as config_file: + config_data = json.load(config_file) + config_data["in_channels"] = in_channels + + lotus_unet = UNet2DConditionModel.from_config(config_data) + + lotus_unet.load_state_dict(lotus_sd) + lotus_unet.to(dtype) + + lotus_model = { + "model": lotus_unet, + "dtype": dtype, + "in_channels": in_channels, + } + + return (lotus_model,) + +class LotusSampler: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "lotus_unet": ("LOTUSUNET",), + "samples": ("LATENT",), + "seed": ("INT", {"default": 123, "min": 0, "max": 2**32, "step": 1}), + "per_batch": ("INT", {"default": 4, "min": 1, "max": 4096, "step": 1}), + "keep_model_loaded": ("BOOLEAN", {"default": False}), + }, + } + + RETURN_TYPES = ("LATENT",) + RETURN_NAMES = ("samples",) + FUNCTION = "loadmodel" + CATEGORY = "ComfyUI-Lotus" + + def loadmodel(self, lotus_unet, seed, samples, per_batch, keep_model_loaded): + + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + mm.soft_empty_cache() + + model = lotus_unet["model"] + dtype = lotus_unet["dtype"] + in_channels = lotus_unet["in_channels"] + + latents = samples["samples"].to(dtype) + latents = latents * 0.18215 + + torch.manual_seed(seed) + torch.cuda.manual_seed(seed) + + if in_channels == 8: # input for g model is 8 channels + single_noise = torch.randn(latents.shape[1:], device=torch.device("cpu"), dtype=dtype, layout=torch.strided) + repeated_noise = single_noise.unsqueeze(0).repeat(latents.shape[0], 1, 1, 1) + latents = torch.cat([latents, repeated_noise], dim=1) + + timesteps = torch.tensor(999, device=device).long() + + task_emb = torch.tensor([1, 0], device=device, dtype=dtype).unsqueeze(0).repeat(1, 1) + task_emb = torch.cat([torch.sin(task_emb), torch.cos(task_emb)], dim=-1).repeat(1, 1) + + prompt_embeds = torch.load(os.path.join(script_directory, "empty_text_embed.pt"), weights_only=True).to(device).to(dtype) + extended_prompt_embeds = prompt_embeds.repeat(latents.shape[0], 1, 1) + + model.to(device) + pbar = ProgressBar(latents.shape[0]) + + results = [] + for start_idx in range(0, latents.shape[0], per_batch): + + sub_images = model( + latents[start_idx:start_idx+per_batch].to(device), + timesteps, + encoder_hidden_states=extended_prompt_embeds[start_idx:start_idx+per_batch], + cross_attention_kwargs=None, + return_dict=False, + class_labels=task_emb, + )[0] + + results.append(sub_images.cpu()) + batch_count = sub_images.shape[0] + pbar.update(batch_count) + + if not keep_model_loaded: + model.to(offload_device) + mm.soft_empty_cache() + + results = torch.cat(results, dim=0) + results = results / 0.18215 + + return {"samples": results}, + + +NODE_CLASS_MAPPINGS = { + "LoadLotusModel": LoadLotusModel, + "LotusSampler": LotusSampler, +} +NODE_DISPLAY_NAME_MAPPINGS = { + "LoadLotusModel": "Load Lotus Model", + "LotusSampler": "Lotus Sampler", + } diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..81caff4 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +diffusers \ No newline at end of file