initial commit

This commit is contained in:
Jukka Seppänen
2024-10-09 19:01:32 +03:00
parent 60a268620c
commit e90605d0b5
7 changed files with 371 additions and 0 deletions
+145
View File
@@ -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
+3
View File
@@ -0,0 +1,3 @@
# ComfyUI nodes to use Lotus depth/normal prediction
Original repo: https://github.com/EnVision-Research/Lotus
+3
View File
@@ -0,0 +1,3 @@
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+73
View File
@@ -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
}
BIN
View File
Binary file not shown.
+146
View File
@@ -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",
}
+1
View File
@@ -0,0 +1 @@
diffusers