diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..e29ac43 --- /dev/null +++ b/.gitignore @@ -0,0 +1,179 @@ +hf_download/ +outputs/ +repo/ + +# 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/ +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/ +cover/ + +# 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 +.pybuilder/ +target/ + +# Jupyter Notebook +.ipynb_checkpoints + +# IPython +profile_default/ +ipython_config.py + +# pyenv +# For a library or package, you might want to ignore these files since the code is +# intended to run in multiple environments; otherwise, check them in: +# .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 + +# UV +# Similar to Pipfile.lock, it is generally recommended to include uv.lock in version control. +# This is especially recommended for binary packages to ensure reproducibility, and is more +# commonly ignored for libraries. +#uv.lock + +# poetry +# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control. +# This is especially recommended for binary packages to ensure reproducibility, and is more +# commonly ignored for libraries. +# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control +#poetry.lock + +# pdm +# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control. +#pdm.lock +# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it +# in version control. +# https://pdm.fming.dev/latest/usage/project/#working-with-version-control +.pdm.toml +.pdm-python +.pdm-build/ + +# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm +__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/ + +# pytype static type analyzer +.pytype/ + +# Cython debug symbols +cython_debug/ + +# PyCharm +# JetBrains specific template is maintained in a separate JetBrains.gitignore that can +# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore +# and can be added to the global gitignore or merged into this file. For a more nuclear +# option (not recommended) you can uncomment the following to ignore the entire idea folder. +.idea/ + +# Ruff stuff: +.ruff_cache/ + +# PyPI configuration file +.pypirc +demo_gradio.py \ 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/example_workflows/lbm_relight_example_01.json b/example_workflows/lbm_relight_example_01.json new file mode 100644 index 0000000..d0abc08 --- /dev/null +++ b/example_workflows/lbm_relight_example_01.json @@ -0,0 +1,1148 @@ +{ + "id": "394ed254-7306-42a2-9ae6-aa880ce4456d", + "revision": 0, + "last_node_id": 1946, + "last_link_id": 5560, + "nodes": [ + { + "id": 1935, + "type": "PreviewImage", + "pos": [ + 3355.5810546875, + 2970.5439453125 + ], + "size": [ + 699.02001953125, + 765.0900268554688 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 5559 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.32", + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 1930, + "type": "LoadLBMModel", + "pos": [ + 2073.0615234375, + 2479.60693359375 + ], + "size": [ + 404.87457275390625, + 130 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [ + { + "name": "compile_args", + "shape": 7, + "type": "FRAMEPACKCOMPILEARGS", + "link": null + } + ], + "outputs": [ + { + "name": "model", + "type": "LBM_MODEL", + "links": [ + 5557 + ] + } + ], + "properties": { + "Node name for S&R": "LoadLBMModel" + }, + "widgets_values": [ + "LBM\\lbm_relight.safetensors", + "bf16", + "main_device" + ] + }, + { + "id": 1940, + "type": "ImageResizeKJv2", + "pos": [ + 1430.7935791015625, + 2517.883056640625 + ], + "size": [ + 270, + 242 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 5546 + }, + { + "name": "width", + "type": "INT", + "widget": { + "name": "width" + }, + "link": 5555 + }, + { + "name": "height", + "type": "INT", + "widget": { + "name": "height" + }, + "link": 5556 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 5547 + ] + }, + { + "name": "width", + "type": "INT", + "links": null + }, + { + "name": "height", + "type": "INT", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "bec42252c690c1b5b2064b5a6732ad11cc452759", + "Node name for S&R": "ImageResizeKJv2" + }, + "widgets_values": [ + 1024, + 1024, + "lanczos", + "crop", + "0, 0, 0", + "center", + 2 + ] + }, + { + "id": 1936, + "type": "ImageCompositeMasked", + "pos": [ + 2054.185791015625, + 2850.79443359375 + ], + "size": [ + 270, + 146 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "destination", + "type": "IMAGE", + "link": 5547 + }, + { + "name": "source", + "type": "IMAGE", + "link": 5544 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": 5543 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 5558, + 5560 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.32", + "Node name for S&R": "ImageCompositeMasked" + }, + "widgets_values": [ + 0, + 0, + false + ] + }, + { + "id": 1932, + "type": "LoadImage", + "pos": [ + 927.2724609375, + 3184.23046875 + ], + "size": [ + 505.674072265625, + 597.9466552734375 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 5536, + 5542 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.32", + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "oldman_upscaled.png", + "image" + ] + }, + { + "id": 1938, + "type": "ImageRemoveBackground+", + "pos": [ + 1759.43017578125, + 3279.618408203125 + ], + "size": [ + 236.54940795898438, + 46 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "rembg_session", + "type": "REMBG_SESSION", + "link": 5541 + }, + { + "name": "image", + "type": "IMAGE", + "link": 5542 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": null + }, + { + "name": "MASK", + "type": "MASK", + "links": [ + 5543 + ] + } + ], + "properties": { + "aux_id": "kijai/ComfyUI_essentials", + "ver": "76e9d1e4399bd025ce8b12c290753d58f9f53e93", + "Node name for S&R": "ImageRemoveBackground+" + }, + "widgets_values": [] + }, + { + "id": 1937, + "type": "TransparentBGSession+", + "pos": [ + 1502.8291015625, + 3405.21533203125 + ], + "size": [ + 299.1265563964844, + 82 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "REMBG_SESSION", + "type": "REMBG_SESSION", + "links": [ + 5541 + ] + } + ], + "properties": { + "aux_id": "kijai/ComfyUI_essentials", + "ver": "76e9d1e4399bd025ce8b12c290753d58f9f53e93", + "Node name for S&R": "TransparentBGSession+" + }, + "widgets_values": [ + "base", + true + ] + }, + { + "id": 1945, + "type": "PreviewImage", + "pos": [ + 2093.712646484375, + 3104.7724609375 + ], + "size": [ + 583.666748046875, + 625.7786254882812 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 5560 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.32", + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 1944, + "type": "LBMSampler", + "pos": [ + 2681.207763671875, + 2832.11962890625 + ], + "size": [ + 270, + 78 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "LBM_MODEL", + "link": 5557 + }, + { + "name": "image", + "type": "IMAGE", + "link": 5558 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [ + 5559 + ] + } + ], + "properties": { + "Node name for S&R": "LBMSampler" + }, + "widgets_values": [ + 20 + ] + }, + { + "id": 1939, + "type": "LoadImage", + "pos": [ + 902.4271240234375, + 2507.15478515625 + ], + "size": [ + 387.65875244140625, + 611.4181518554688 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 5546 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "title": "Load Image: Background", + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.32", + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "pasted/image (902).png", + "image" + ] + }, + { + "id": 1933, + "type": "ImageResizeKJv2", + "pos": [ + 1487.82470703125, + 2934.407470703125 + ], + "size": [ + 270, + 242 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 5536 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 5544 + ] + }, + { + "name": "width", + "type": "INT", + "links": [ + 5555 + ] + }, + { + "name": "height", + "type": "INT", + "links": [ + 5556 + ] + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "bec42252c690c1b5b2064b5a6732ad11cc452759", + "Node name for S&R": "ImageResizeKJv2" + }, + "widgets_values": [ + 1024, + 1024, + "lanczos", + "crop", + "0, 0, 0", + "center", + 2 + ] + } + ], + "links": [ + [ + 5536, + 1932, + 0, + 1933, + 0, + "IMAGE" + ], + [ + 5541, + 1937, + 0, + 1938, + 0, + "REMBG_SESSION" + ], + [ + 5542, + 1932, + 0, + 1938, + 1, + "IMAGE" + ], + [ + 5543, + 1938, + 1, + 1936, + 2, + "MASK" + ], + [ + 5544, + 1933, + 0, + 1936, + 1, + "IMAGE" + ], + [ + 5546, + 1939, + 0, + 1940, + 0, + "IMAGE" + ], + [ + 5547, + 1940, + 0, + 1936, + 0, + "IMAGE" + ], + [ + 5555, + 1933, + 1, + 1940, + 1, + "INT" + ], + [ + 5556, + 1933, + 2, + 1940, + 2, + "INT" + ], + [ + 5557, + 1930, + 0, + 1944, + 0, + "LBM_MODEL" + ], + [ + 5558, + 1936, + 0, + 1944, + 1, + "IMAGE" + ], + [ + 5559, + 1944, + 0, + 1935, + 0, + "IMAGE" + ], + [ + 5560, + 1936, + 0, + 1945, + 0, + "IMAGE" + ] + ], + "groups": [], + "config": {}, + "extra": { + "ds": { + "scale": 0.7513148009015783, + "offset": [ + -444.1598280476131, + -2381.1529304382993 + ] + }, + "frontendVersion": "1.19.4", + "VHS_latentpreview": true, + "VHS_latentpreviewrate": 0, + "VHS_MetadataImage": true, + "VHS_KeepIntermediate": true, + "prompt": { + "6": { + "inputs": { + "text": "", + "clip": [ + "38", + 0 + ] + }, + "class_type": "CLIPTextEncode", + "_meta": { + "title": "CLIP Text Encode (Positive Prompt)" + } + }, + "7": { + "inputs": { + "text": "low quality, worst quality, deformed, distorted, disfigured, motion smear, motion artifacts, fused fingers, bad anatomy, weird hand, ugly", + "clip": [ + "38", + 0 + ] + }, + "class_type": "CLIPTextEncode", + "_meta": { + "title": "CLIP Text Encode (Negative Prompt)" + } + }, + "38": { + "inputs": { + "clip_name": "t5xxl_fp16.safetensors", + "type": "ltxv", + "device": "default" + }, + "class_type": "CLIPLoader", + "_meta": { + "title": "Load CLIP" + } + }, + "44": { + "inputs": { + "ckpt_name": "ltx-video-13b-distilled-step-13000.safetensors" + }, + "class_type": "CheckpointLoaderSimple", + "_meta": { + "title": "Load Checkpoint" + } + }, + "73": { + "inputs": { + "sampler_name": "euler_ancestral" + }, + "class_type": "KSamplerSelect", + "_meta": { + "title": "KSamplerSelect" + } + }, + "1206": { + "inputs": { + "image": "5aa.png" + }, + "class_type": "LoadImage", + "_meta": { + "title": "Load Image" + } + }, + "1241": { + "inputs": { + "frame_rate": 24.000000000000004, + "positive": [ + "6", + 0 + ], + "negative": [ + "7", + 0 + ] + }, + "class_type": "LTXVConditioning", + "_meta": { + "title": "LTXVConditioning" + } + }, + "1335": { + "inputs": { + "samples": [ + "1338", + 0 + ], + "vae": [ + "1870", + 0 + ] + }, + "class_type": "VAEDecode", + "_meta": { + "title": "VAE Decode" + } + }, + "1336": { + "inputs": { + "frame_rate": 24, + "loop_count": 0, + "filename_prefix": "ltxv-base", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "pingpong": false, + "save_output": false, + "images": [ + "1335", + 0 + ] + }, + "class_type": "VHS_VideoCombine", + "_meta": { + "title": "Video Combine 🎥🅥🅗🅢" + } + }, + "1338": { + "inputs": { + "width": 768, + "height": 512, + "num_frames": 97, + "optional_cond_indices": "0, 40, 90", + "strength": 0.8, + "crop": "center", + "crf": 30, + "blur": 1, + "model": [ + "44", + 0 + ], + "vae": [ + "44", + 2 + ], + "guider": [ + "1807", + 0 + ], + "sampler": [ + "73", + 0 + ], + "sigmas": [ + "1872", + 0 + ], + "noise": [ + "1507", + 0 + ], + "optional_cond_images": [ + "1876", + 0 + ] + }, + "class_type": "LTXVBaseSampler", + "_meta": { + "title": "🅛🅣🅧 LTXV Base Sampler" + } + }, + "1507": { + "inputs": { + "noise_seed": 108 + }, + "class_type": "RandomNoise", + "_meta": { + "title": "RandomNoise" + } + }, + "1593": { + "inputs": { + "factor": 0.25, + "latents": [ + "1691", + 0 + ], + "reference": [ + "1338", + 0 + ] + }, + "class_type": "LTXVAdainLatent", + "_meta": { + "title": "🅛🅣🅧 LTXV Adain Latent" + } + }, + "1598": { + "inputs": { + "noise_seed": 414 + }, + "class_type": "RandomNoise", + "_meta": { + "title": "RandomNoise" + } + }, + "1599": { + "inputs": { + "frame_rate": 24, + "loop_count": 0, + "filename_prefix": "ltxv-hd", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 18, + "save_metadata": false, + "pingpong": false, + "save_output": false, + "images": [ + "1699", + 0 + ] + }, + "class_type": "VHS_VideoCombine", + "_meta": { + "title": "Video Combine 🎥🅥🅗🅢" + } + }, + "1601": { + "inputs": { + "tile_size": 1280, + "overlap": 128, + "temporal_size": 128, + "temporal_overlap": 32, + "samples": [ + "1873", + 0 + ], + "vae": [ + "1870", + 0 + ] + }, + "class_type": "VAEDecodeTiled", + "_meta": { + "title": "VAE Decode (Tiled)" + } + }, + "1661": { + "inputs": { + "width": 1280, + "height": 1280, + "upscale_method": "bicubic", + "keep_proportion": true, + "divisible_by": 2, + "crop": "center", + "image": [ + "1601", + 0 + ] + }, + "class_type": "ImageResizeKJ", + "_meta": { + "title": "Resize Image" + } + }, + "1691": { + "inputs": { + "samples": [ + "1338", + 0 + ], + "upscale_model": [ + "1828", + 0 + ], + "vae": [ + "44", + 2 + ] + }, + "class_type": "LTXVLatentUpsampler", + "_meta": { + "title": "🅛🅣🅧 LTXV Latent Upsampler" + } + }, + "1699": { + "inputs": { + "grain_intensity": 0.010000000000000002, + "saturation": 0.5, + "images": [ + "1661", + 0 + ] + }, + "class_type": "LTXVFilmGrain", + "_meta": { + "title": "🅛🅣🅧 LTXV Film Grain" + } + }, + "1807": { + "inputs": { + "skip_steps_sigma_threshold": 0.9970000000000002, + "cfg_star_rescale": true, + "sigmas": "1.0, 0.9933, 0.9850, 0.9767, 0.9008, 0.6180", + "cfg_values": "1,1,1,1,1,1", + "stg_scale_values": "0,0,0,0,0,0", + "stg_rescale_values": "1, 1, 1, 1, 1, 1", + "stg_layers_indices": "[35], [35], [35], [42], [42], [42]", + "model": [ + "44", + 0 + ], + "positive": [ + "1241", + 0 + ], + "negative": [ + "1241", + 1 + ] + }, + "class_type": "STGGuiderAdvanced", + "_meta": { + "title": "🅛🅣🅧 STG Guider Advanced" + } + }, + "1813": { + "inputs": { + "skip_steps_sigma_threshold": 0.9970000000000002, + "cfg_star_rescale": true, + "sigmas": "1", + "cfg_values": "1", + "stg_scale_values": "0", + "stg_rescale_values": "1", + "stg_layers_indices": "[42]", + "model": [ + "44", + 0 + ], + "positive": [ + "1241", + 0 + ], + "negative": [ + "1241", + 1 + ] + }, + "class_type": "STGGuiderAdvanced", + "_meta": { + "title": "🅛🅣🅧 STG Guider Advanced" + } + }, + "1828": { + "inputs": { + "upscale_model": "ltxv-spatial-upscaler-0.9.7.safetensors", + "spatial_upsample": true, + "temporal_upsample": false + }, + "class_type": "LTXVLatentUpsamplerModelLoader", + "_meta": { + "title": "🅛🅣🅧 LTXV Latent Upsampler Model Loader" + } + }, + "1865": { + "inputs": { + "image": "5B.png" + }, + "class_type": "LoadImage", + "_meta": { + "title": "Load Image" + } + }, + "1866": { + "inputs": { + "image": "5C.png" + }, + "class_type": "LoadImage", + "_meta": { + "title": "Load Image" + } + }, + "1867": { + "inputs": { + "image1": [ + "1206", + 0 + ], + "image2": [ + "1865", + 0 + ] + }, + "class_type": "ImageBatch", + "_meta": { + "title": "Batch Images" + } + }, + "1868": { + "inputs": { + "image1": [ + "1867", + 0 + ], + "image2": [ + "1866", + 0 + ] + }, + "class_type": "ImageBatch", + "_meta": { + "title": "Batch Images" + } + }, + "1870": { + "inputs": { + "timestep": 0.05, + "scale": 0.025, + "seed": 42, + "vae": [ + "44", + 2 + ] + }, + "class_type": "Set VAE Decoder Noise", + "_meta": { + "title": "🅛🅣🅧 Set VAE Decoder Noise" + } + }, + "1871": { + "inputs": { + "string": "1.0000, 0.9937, 0.9875, 0.9812, 0.9750, 0.9094, 0.7250, 0.4219, 0.0" + }, + "class_type": "StringToFloatList", + "_meta": { + "title": "String to Float List" + } + }, + "1872": { + "inputs": { + "float_list": [ + "1871", + 0 + ] + }, + "class_type": "FloatToSigmas", + "_meta": { + "title": "Float To Sigmas" + } + }, + "1873": { + "inputs": { + "horizontal_tiles": 1, + "vertical_tiles": 1, + "overlap": 1, + "latents_cond_strength": 0.15, + "boost_latent_similarity": false, + "crop": "disabled", + "optional_cond_indices": "0, 40, 90", + "images_cond_strengths": "0.9", + "model": [ + "44", + 0 + ], + "vae": [ + "44", + 2 + ], + "noise": [ + "1598", + 0 + ], + "sampler": [ + "73", + 0 + ], + "sigmas": [ + "1875", + 0 + ], + "guider": [ + "1813", + 0 + ], + "latents": [ + "1593", + 0 + ], + "optional_cond_images": [ + "1876", + 0 + ] + }, + "class_type": "LTXVTiledSampler", + "_meta": { + "title": "🅛🅣🅧 LTXV Tiled Sampler" + } + }, + "1874": { + "inputs": { + "string": "0.85, 0.7250, 0.6, 0.4219, 0.0" + }, + "class_type": "StringToFloatList", + "_meta": { + "title": "String to Float List" + } + }, + "1875": { + "inputs": { + "float_list": [ + "1874", + 0 + ] + }, + "class_type": "FloatToSigmas", + "_meta": { + "title": "Float To Sigmas" + } + }, + "1876": { + "inputs": { + "radius_x": 1, + "radius_y": 1, + "images": [ + "1868", + 0 + ] + }, + "class_type": "BlurImageFast", + "_meta": { + "title": "Blur Image (Fast)" + } + } + }, + "comfy_fork_version": "develop@580b3007", + "workspace_info": { + "id": "elBQFQknIoLYTEwIloQuw" + }, + "node_versions": { + "comfy-core": "0.3.20" + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/lbm/config.py b/lbm/config.py new file mode 100644 index 0000000..4f69788 --- /dev/null +++ b/lbm/config.py @@ -0,0 +1,141 @@ +import json +import os +import warnings +from dataclasses import asdict, field +from typing import Any, Dict, Union + +import yaml +from pydantic import ValidationError +from pydantic.dataclasses import dataclass +from yaml import safe_load + + +@dataclass +class BaseConfig: + """This is the BaseConfig class which defines all the useful loading and saving methods + of the configs""" + + name: str = field(init=False) + + def __post_init__(self): + self.name = self.__class__.__name__ + + @classmethod + def from_dict(cls, config_dict: Dict[str, Any]) -> "BaseConfig": + """Creates a BaseConfig instance from a dictionnary + + Args: + config_dict (dict): The Python dictionnary containing all the parameters + + Returns: + :class:`BaseConfig`: The created instance + """ + try: + config = cls(**config_dict) + except (ValidationError, TypeError) as e: + raise e + return config + + @classmethod + def _dict_from_json(cls, json_path: Union[str, os.PathLike]) -> Dict[str, Any]: + try: + with open(json_path) as f: + try: + config_dict = json.load(f) + return config_dict + + except (TypeError, json.JSONDecodeError) as e: + raise TypeError( + f"File {json_path} not loadable. Maybe not json ? \n" + f"Catch Exception {type(e)} with message: " + str(e) + ) from e + + except FileNotFoundError: + raise FileNotFoundError( + f"Config file not found. Please check path '{json_path}'" + ) + + @classmethod + def from_json(cls, json_path: str) -> "BaseConfig": + """Creates a BaseConfig instance from a JSON config file + + Args: + json_path (str): The path to the json file containing all the parameters + + Returns: + :class:`BaseConfig`: The created instance + """ + config_dict = cls._dict_from_json(json_path) + + config_name = config_dict.pop("name") + + if cls.__name__ != config_name: + warnings.warn( + f"You are trying to load a " + f"`{ cls.__name__}` while a " + f"`{config_name}` is given." + ) + + return cls.from_dict(config_dict) + + def to_dict(self) -> dict: + """Transforms object into a Python dictionnary + + Returns: + (dict): The dictionnary containing all the parameters""" + return asdict(self) + + def to_json_string(self): + """Transforms object into a JSON string + + Returns: + (str): The JSON str containing all the parameters""" + return json.dumps(self.to_dict()) + + def save_json(self, file_path: str): + """Saves a ``.json`` file from the dataclass + + Args: + file_path (str): path to the file + """ + with open(os.path.join(file_path), "w", encoding="utf-8") as fp: + fp.write(self.to_json_string()) + + def save_yaml(self, file_path: str): + """Saves a ``.yaml`` file from the dataclass + + Args: + file_path (str): path to the file + """ + with open(os.path.join(file_path), "w", encoding="utf-8") as fp: + yaml.dump(self.to_dict(), fp) + + @classmethod + def from_yaml(cls, yaml_path: str) -> "BaseConfig": + """Creates a BaseConfig instance from a YAML config file + + Args: + yaml_path (str): The path to the yaml file containing all the parameters + + Returns: + :class:`BaseConfig`: The created instance + """ + with open(yaml_path, "r") as f: + try: + config_dict = safe_load(f) + except yaml.YAMLError as e: + raise yaml.YAMLError( + f"File {yaml_path} not loadable. Maybe not yaml ? \n" + f"Catch Exception {type(e)} with message: " + str(e) + ) from e + + config_name = config_dict.pop("name") + + if cls.__name__ != config_name: + warnings.warn( + f"You are trying to load a " + f"`{ cls.__name__}` while a " + f"`{config_name}` is given." + ) + + return cls.from_dict(config_dict) diff --git a/lbm/data/__init__.py b/lbm/data/__init__.py new file mode 100644 index 0000000..e38f173 --- /dev/null +++ b/lbm/data/__init__.py @@ -0,0 +1,62 @@ +""" +This module contains a collection of data related classes and functions to train the :mod:`cr.models`. +In a training loop a batch of data is struvtued as a dictionnary on which the modules :mod:`cr.data.datasets` +and :mod:`cr.data.filters` allow to perform several operations. + + +Examples +######## + +Create a DataModule to train a model + +.. code-block::python + + from cr.data import DataModule, DataModuleConfig + from cr.data.filters import KeyFilter, KeyFilterConfig + from cr.data.mappers import KeyRenameMapper, KeyRenameMapperConfig + + # Create the filters and mappers + filters_mappers = [ + KeyFilter(KeyFilterConfig(keys=["image", "txt"])), + KeyRenameMapper( + KeyRenameMapperConfig(key_map={"jpg": "image", "txt": "text"}) + ) + ] + + # Create the DataModule + data_module = DataModule( + train_config=DataModuleConfig( + shards_path_or_urls="your urls or paths", + decoder="pil", + shuffle_buffer_size=100, + per_worker_batch_size=32, + num_workers=4, + ), + train_filters_mappers=filters_mappers, + eval_config=DataModuleConfig( + shards_path_or_urls="your urls or paths", + decoder="pil", + shuffle_buffer_size=100, + per_worker_batch_size=32, + num_workers=4, + ), + eval_filters_mappers=filters_mappers, + ) + + # This can then be passed to a :mod:`pytorch_lightning.Trainer` to train a model + + + + + +The :mod:`cr.data` includes the following submodules: + +- :mod:`cr.data.datasets`: a collection of :mod:`pytorch_lightning.LightningDataModule` used to train the models. In particular, + they can used to create the dataloaders and setup the data pipelines. +- :mod:`cr.data.filters`: a collection of filters used apply filters on a training batch of data/ + +""" + +from .datasets import DataModule + +__all__ = ["DataModule"] diff --git a/lbm/data/datasets/__init__.py b/lbm/data/datasets/__init__.py new file mode 100644 index 0000000..5715c94 --- /dev/null +++ b/lbm/data/datasets/__init__.py @@ -0,0 +1,9 @@ +""" +A collection of :mod:`pytorch_lightning.LightningDataModule` used to train the models. In particular, +they can be used to create the dataloaders and setup the data pipelines. +""" + +from .dataset import DataModule +from .datasets_config import DataModuleConfig + +__all__ = ["DataModule", "DataModuleConfig"] diff --git a/lbm/data/datasets/collation_fn.py b/lbm/data/datasets/collation_fn.py new file mode 100644 index 0000000..046309f --- /dev/null +++ b/lbm/data/datasets/collation_fn.py @@ -0,0 +1,41 @@ +from typing import Dict, List, Union + +import numpy as np +import torch + + +def custom_collation_fn( + samples: List[Dict[str, Union[int, float, np.ndarray, torch.Tensor]]], + combine_tensors: bool = True, + combine_scalars: bool = True, +) -> dict: + """ + Collate function for PyTorch DataLoader. + + Args: + samples(List[Dict[str, Union[int, float, np.ndarray, torch.Tensor]]]): List of samples. + combine_tensors (bool): Whether to turn lists of tensors into a single tensor. + combine_scalars (bool): Whether to turn lists of scalars into a single ndarray. + """ + keys = set.intersection(*[set(sample.keys()) for sample in samples]) + batched = {key: [] for key in keys} + for s in samples: + [batched[key].append(s[key]) for key in batched] + + result = {} + for key in batched: + if isinstance(batched[key][0], (int, float)): + if combine_scalars: + result[key] = np.array(list(batched[key])) + elif isinstance(batched[key][0], torch.Tensor): + if combine_tensors: + result[key] = torch.stack(list(batched[key])) + elif isinstance(batched[key][0], np.ndarray): + if combine_tensors: + result[key] = np.array(list(batched[key])) + else: + result[key] = list(batched[key]) + + del samples + del batched + return result diff --git a/lbm/data/datasets/dataset.py b/lbm/data/datasets/dataset.py new file mode 100644 index 0000000..c895406 --- /dev/null +++ b/lbm/data/datasets/dataset.py @@ -0,0 +1,243 @@ +from typing import Callable, List, Union + +import pytorch_lightning as pl +import webdataset as wds +from webdataset import DataPipeline + +from ..filters import BaseFilter, FilterWrapper +from ..mappers import BaseMapper, MapperWrapper +from .collation_fn import custom_collation_fn +from .datasets_config import DataModuleConfig + + +class DataPipeline: + """ + DataPipeline class for creating a dataloader from a single configuration + + Args: + + config (DataModuleConfig): + Configuration for the dataset + + filters_mappers (Union[List[Union[BaseMapper, BaseFilter, FilterWrapper, MapperWrapper]]): + List of filters and mappers for the dataset. These will be sequentially applied. + + batched_filters_mappers (List[Union[BaseMapper, BaseFilter, FilterWrapper, MapperWrapper]]): + List of batched transforms for the dataset. These will be sequentially applied. + """ + + def __init__( + self, + config: DataModuleConfig, + filters_mappers: List[ + Union[BaseMapper, BaseFilter, FilterWrapper, MapperWrapper] + ], + batched_filters_mappers: List[ + Union[BaseMapper, BaseFilter, FilterWrapper, MapperWrapper] + ] = None, + ): + self.config = config + self.shards_path_or_urls = config.shards_path_or_urls + self.filters_mappers = filters_mappers + self.batched_filters_mappers = batched_filters_mappers or [] + + if filters_mappers is None: + filters_mappers = [] + + # set processing pipeline + self.processing_pipeline = [wds.decode(config.decoder, handler=config.handler)] + self.processing_pipeline.extend( + self._add_filters_mappers( + filters_mappers=filters_mappers, + handler=config.handler, + ) + ) + + def _add_filters_mappers( + self, + filters_mappers: List[ + Union[ + FilterWrapper, + MapperWrapper, + ] + ], + handler: Callable = wds.warn_and_continue, + ) -> List[Union[FilterWrapper, MapperWrapper]]: + tmp_pipeline = [] + for filter_mapper in filters_mappers: + if isinstance(filter_mapper, FilterWrapper) or isinstance( + filter_mapper, BaseFilter + ): + tmp_pipeline.append(wds.select(filter_mapper)) + elif isinstance(filter_mapper, MapperWrapper) or isinstance( + filter_mapper, BaseMapper + ): + tmp_pipeline.append(wds.map(filter_mapper, handler=handler)) + elif isinstance(filter_mapper) or isinstance(filter_mapper): + tmp_pipeline.append(wds.map(filter_mapper, handler=handler)) + else: + raise ValueError("Unknown type of filter/mapper") + return tmp_pipeline + + def setup(self): + pipeline = [wds.SimpleShardList(self.shards_path_or_urls)] + + # shuffle before split by node + if self.config.shuffle_before_split_by_node_buffer_size is not None: + pipeline.append( + wds.shuffle( + self.config.shuffle_before_split_by_node_buffer_size, + handler=self.config.handler, + ) + ) + # split by node + pipeline.append(wds.split_by_node) + + # shuffle before split by workers + if self.config.shuffle_before_split_by_workers_buffer_size is not None: + pipeline.append( + wds.shuffle( + self.config.shuffle_before_split_by_workers_buffer_size, + handler=self.config.handler, + ) + ) + # split by worker + pipeline.extend( + [ + wds.split_by_worker, + wds.tarfile_to_samples( + handler=self.config.handler, + rename_files=self.config.rename_files_fn, + ), + ] + ) + + # shuffle before filter mappers + if self.config.shuffle_before_filter_mappers_buffer_size is not None: + pipeline.append( + wds.shuffle( + self.config.shuffle_before_filter_mappers_buffer_size, + handler=self.config.handler, + ) + ) + + # apply filters and mappers + pipeline.extend(self.processing_pipeline) + + # shuffle after filter mappers + if self.config.shuffle_after_filter_mappers_buffer_size is not None: + pipeline.append( + wds.shuffle( + self.config.shuffle_after_filter_mappers_buffer_size, + handler=self.config.handler, + ), + ) + + # batching + pipeline.append( + wds.batched( + self.config.per_worker_batch_size, + collation_fn=custom_collation_fn, + ) + ) + + # apply batched transforms + pipeline.extend( + self._add_filters_mappers( + filters_mappers=self.batched_filters_mappers, + handler=self.config.handler, + ) + ) + + # create the data pipeline + pipeline = wds.DataPipeline(*pipeline, handler=self.config.handler) + + # set the pipeline + self.pipeline = pipeline + + def dataloader(self): + # return the loader + return wds.WebLoader( + self.pipeline, + batch_size=None, + num_workers=self.config.num_workers, + ) + + +class DataModule(pl.LightningDataModule): + """ + Main DataModule class for creating data loaders and training/evaluating models + + Args: + + train_config (DataModuleConfig): + Configuration for the training dataset + + train_filters_mappers (Union[List[Union[BaseMapper, BaseFilter, FilterWrapper, MapperWrapper]]): + List of filters and mappers for the training dataset. These will be sequentially applied. + + train_batched_filters_mappers (List[Union[BaseMapper, BaseFilter, FilterWrapper, MapperWrapper]]): + List of batched transforms for the training dataset. These will be sequentially applied. + + eval_config (DataModuleConfig): + Configuration for the evaluation dataset + + eval_filters_mappers (List[Union[FilterWrapper, MapperWrapper]]): + List of filters and mappers for the evaluation dataset.These will be sequentially applied. + + eval_batched_filters_mappers (List[Union[BaseMapper, BaseFilter, FilterWrapper, MapperWrapper]]): + List of batched transforms for the evaluation dataset. These will be sequentially applied. + """ + + def __init__( + self, + train_config: DataModuleConfig, + train_filters_mappers: List[ + Union[BaseMapper, BaseFilter, FilterWrapper, MapperWrapper] + ] = None, + train_batched_filters_mappers: List[ + Union[BaseMapper, BaseFilter, FilterWrapper, MapperWrapper] + ] = None, + eval_config: DataModuleConfig = None, + eval_filters_mappers: List[Union[FilterWrapper, MapperWrapper]] = None, + eval_batched_filters_mappers: List[ + Union[BaseMapper, BaseFilter, FilterWrapper, MapperWrapper] + ] = None, + ): + super().__init__() + + self.train_config = train_config + self.train_filters_mappers = train_filters_mappers + self.train_batched_filters_mappers = train_batched_filters_mappers + + self.eval_config = eval_config + self.eval_filters_mappers = eval_filters_mappers + self.eval_batched_filters_mappers = eval_batched_filters_mappers + + def setup(self, stage=None): + """ + Setup the data module and create the webdataset processing pipelines + """ + + # train pipeline + self.train_pipeline = DataPipeline( + config=self.train_config, + filters_mappers=self.train_filters_mappers, + batched_filters_mappers=self.train_batched_filters_mappers, + ) + self.train_pipeline.setup() + + # eval pipeline + if self.eval_config is not None: + self.eval_pipeline = DataPipeline( + config=self.eval_config, + filters_mappers=self.eval_filters_mappers, + batched_filters_mappers=self.eval_batched_filters_mappers, + ) + self.eval_pipeline.setup() + + def train_dataloader(self): + return self.train_pipeline.dataloader() + + def val_dataloader(self): + return self.eval_pipeline.dataloader() diff --git a/lbm/data/datasets/datasets_config.py b/lbm/data/datasets/datasets_config.py new file mode 100644 index 0000000..4ecddee --- /dev/null +++ b/lbm/data/datasets/datasets_config.py @@ -0,0 +1,42 @@ +from typing import Callable, List, Optional, Union + +import webdataset as wds +from pydantic.dataclasses import dataclass + +from ...config import BaseConfig + + +@dataclass +class DataModuleConfig(BaseConfig): + """ + Configuration for the DataModule + + Args: + + shards_path_or_urls (Union[str, List[str]]): The path or url to the shards. Defaults to None. + per_worker_batch_size (int): The batch size for the dataset. Defaults to 16. + num_workers (int): The number of workers to use. Defaults to 1. + shuffle_before_split_by_node_buffer_size (Optional[int]): The buffer size for the shuffle before split by node. Defaults to 100. + shuffle_before_split_by_workers_buffer_size (Optional[int]): The buffer size for the shuffle before split by workers. Defaults to 100. + shuffle_before_filter_mappers_buffer_size (Optional[int]): The buffer size for the shuffle before filter mappers. Defaults to 1000. + shuffle_after_filter_mappers_buffer_size (Optional[int]): The buffer size for the shuffle after filter mappers. Defaults to 1000. + decoder (str): The decoder to use. Defaults to "pil". + handler (Callable): A callable to handle the warnings. Defaults to wds.warn_and_continue. + rename_files_fn (Optional[Callable[[str], str]]): A callable to rename the files. Defaults to None. + """ + + shards_path_or_urls: Union[str, List[str]] = None + per_worker_batch_size: int = 16 + num_workers: int = 1 + shuffle_before_split_by_node_buffer_size: Optional[int] = 100 + shuffle_before_split_by_workers_buffer_size: Optional[int] = 100 + shuffle_before_filter_mappers_buffer_size: Optional[int] = 1000 + shuffle_after_filter_mappers_buffer_size: Optional[int] = 1000 + decoder: str = "pil" + handler: Callable = wds.warn_and_continue + rename_files_fn: Optional[Callable[[str], str]] = None + + def __post_init__(self): + super().__post_init__() + if self.rename_files_fn is not None: + assert callable(self.rename_files_fn), "rename_files must be a callable" diff --git a/lbm/data/filters/__init__.py b/lbm/data/filters/__init__.py new file mode 100644 index 0000000..6a97c41 --- /dev/null +++ b/lbm/data/filters/__init__.py @@ -0,0 +1,12 @@ +from .base import BaseFilter +from .filter_wrapper import FilterWrapper +from .filters import KeyFilter +from .filters_config import BaseFilterConfig, KeyFilterConfig + +__all__ = [ + "BaseFilter", + "FilterWrapper", + "KeyFilter", + "BaseFilterConfig", + "KeyFilterConfig", +] diff --git a/lbm/data/filters/base.py b/lbm/data/filters/base.py new file mode 100644 index 0000000..357dc2e --- /dev/null +++ b/lbm/data/filters/base.py @@ -0,0 +1,21 @@ +from typing import Any, Dict + +from .filters_config import BaseFilterConfig + + +class BaseFilter: + """ + Base class for filters. This class should be subclassed to create a new filter. + + Args: + + config (BaseFilterConfig): + Configuration for the filter + """ + + def __init__(self, config: BaseFilterConfig): + self.verbose = config.verbose + + def __call__(self, sample: Dict[str, Any]) -> bool: + """This function should be implemented by the subclass""" + raise NotImplementedError diff --git a/lbm/data/filters/filter_wrapper.py b/lbm/data/filters/filter_wrapper.py new file mode 100644 index 0000000..e2477e1 --- /dev/null +++ b/lbm/data/filters/filter_wrapper.py @@ -0,0 +1,36 @@ +from typing import Any, Dict, List, Union + +from .base import BaseFilter + + +class FilterWrapper: + """ + Wrapper for multiple filters. This class allows to apply multiple filters to a batch of data. + The filters are applied in the order they are passed to the wrapper. + + Args: + + filters (List[BaseFilter]): + List of filters to apply to the batch of data + """ + + def __init__( + self, + filters: Union[List[BaseFilter], None] = None, + ): + self.filters = filters + + def __call__(self, batch: Dict[str, Any]) -> None: + """ + Forward pass through all filters + + Args: + + batch: batch of data + """ + filter_output = True + for filter in self.filters: + filter_output = filter(batch) + if not filter_output: + return False + return True diff --git a/lbm/data/filters/filters.py b/lbm/data/filters/filters.py new file mode 100644 index 0000000..41ab776 --- /dev/null +++ b/lbm/data/filters/filters.py @@ -0,0 +1,33 @@ +import logging + +from .base import BaseFilter +from .filters_config import KeyFilterConfig + +logging.basicConfig(level=logging.INFO) + + +class KeyFilter(BaseFilter): + """ + This filter checks if ALL the given keys are present in the sample + + Args: + + config (KeyFilterConfig): configuration for the filter + """ + + def __init__(self, config: KeyFilterConfig): + super().__init__(config) + keys = config.keys + if isinstance(keys, str): + keys = [keys] + + self.keys = set(keys) + + def __call__(self, batch: dict) -> bool: + try: + res = self.keys.issubset(set(batch.keys())) + return res + except Exception as e: + if self.verbose: + logging.error(f"Error in KeyFilter: {e}") + return False diff --git a/lbm/data/filters/filters_config.py b/lbm/data/filters/filters_config.py new file mode 100644 index 0000000..731097c --- /dev/null +++ b/lbm/data/filters/filters_config.py @@ -0,0 +1,32 @@ +from typing import List, Union + +from pydantic.dataclasses import dataclass + +from ...config import BaseConfig + + +@dataclass +class BaseFilterConfig(BaseConfig): + """ + Base configuration for filters + + Args: + + verbose (bool): + If True, print debug information. Defaults to False""" + + verbose: bool = False + + +@dataclass +class KeyFilterConfig(BaseFilterConfig): + """ + This filter checks if the keys are present in a sample. + + Args: + + keys (Union[str, List[str]]): + Key or list of keys to check. Defaults to "txt" + """ + + keys: Union[str, List[str]] = "txt" diff --git a/lbm/data/mappers/__init__.py b/lbm/data/mappers/__init__.py new file mode 100644 index 0000000..92a893f --- /dev/null +++ b/lbm/data/mappers/__init__.py @@ -0,0 +1,19 @@ +from .base import BaseMapper +from .mappers import KeyRenameMapper, RescaleMapper, TorchvisionMapper +from .mappers_config import ( + KeyRenameMapperConfig, + RescaleMapperConfig, + TorchvisionMapperConfig, +) +from .mappers_wrapper import MapperWrapper + +__all__ = [ + "BaseMapper", + "KeyRenameMapper", + "RescaleMapper", + "TorchvisionMapper", + "KeyRenameMapperConfig", + "RescaleMapperConfig", + "TorchvisionMapperConfig", + "MapperWrapper", +] diff --git a/lbm/data/mappers/base.py b/lbm/data/mappers/base.py new file mode 100644 index 0000000..49414b4 --- /dev/null +++ b/lbm/data/mappers/base.py @@ -0,0 +1,26 @@ +from typing import Any, Dict + +from .mappers_config import BaseMapperConfig + + +class BaseMapper: + """ + Base class for the mappers used to modify the samples in the data pipeline. + + Args: + + config (BaseMapperConfig): + Configuration for the mapper. + """ + + def __init__(self, config: BaseMapperConfig): + self.config = config + self.key = config.key + + if config.output_key is None: + self.output_key = config.key + else: + self.output_key = config.output_key + + def map(self, batch: Dict[str, Any], *args, **kwargs) -> Dict[str, Any]: + raise NotImplementedError diff --git a/lbm/data/mappers/mappers.py b/lbm/data/mappers/mappers.py new file mode 100644 index 0000000..41df641 --- /dev/null +++ b/lbm/data/mappers/mappers.py @@ -0,0 +1,135 @@ +from typing import Any, Dict + +from torchvision import transforms + +from .base import BaseMapper +from .mappers_config import ( + KeyRenameMapperConfig, + RescaleMapperConfig, + TorchvisionMapperConfig, +) + + +class KeyRenameMapper(BaseMapper): + """ + Rename keys in a sample according to a key map + + Args: + + config (KeyRenameMapperConfig): Configuration for the mapper + + Examples + ######## + + 1. Rename keys in a sample according to a key map + + .. code-block:: python + + from cr.data.mappers import KeyRenameMapper, KeyRenameMapperConfig + + config = KeyRenameMapperConfig( + key_map={"old_key": "new_key"} + ) + + mapper = KeyRenameMapper(config) + + sample = {"old_key": 1} + new_sample = mapper(sample) + print(new_sample) # {"new_key": 1} + + 2. Rename keys in a sample according to a key map and a condition key + + .. code-block:: python + + from cr.data.mappers import KeyRenameMapper, KeyRenameMapperConfig + + config = KeyRenameMapperConfig( + key_map={"old_key": "new_key"}, + condition_key="condition", + condition_fn=lambda x: x == 1 + ) + + mapper = KeyRenameMapper(config) + + sample = {"old_key": 1, "condition": 1} + new_sample = mapper(sample) + print(new_sample) # {"new_key": 1} + + sample = {"old_key": 1, "condition": 0} + new_sample = mapper(sample) + print(new_sample) # {"old_key": 1} + + ``` + """ + + def __init__(self, config: KeyRenameMapperConfig): + super().__init__(config) + self.key_map = config.key_map + self.condition_key = config.condition_key + self.condition_fn = config.condition_fn + self.else_key_map = config.else_key_map + + def __call__(self, batch: Dict[str, Any], *args, **kwrags): + if self.condition_key is not None: + condition_key = batch[self.condition_key] + if self.condition_fn(condition_key): + for old_key, new_key in self.key_map.items(): + if old_key in batch: + batch[new_key] = batch.pop(old_key) + + elif self.else_key_map is not None: + for old_key, new_key in self.else_key_map.items(): + if old_key in batch: + batch[new_key] = batch.pop(old_key) + + else: + for old_key, new_key in self.key_map.items(): + if old_key in batch: + batch[new_key] = batch.pop(old_key) + return batch + + +class TorchvisionMapper(BaseMapper): + """ + Apply torchvision transforms to a sample + + Args: + + config (TorchvisionMapperConfig): Configuration for the mapper + """ + + def __init__(self, config: TorchvisionMapperConfig): + super().__init__(config) + chained_transforms = [] + for transform, kwargs in zip(config.transforms, config.transforms_kwargs): + transform = getattr(transforms, transform) + chained_transforms.append(transform(**kwargs)) + self.transforms = transforms.Compose(chained_transforms) + + def __call__(self, batch: Dict[str, Any], *args, **kwrags) -> Dict[str, Any]: + if self.key in batch: + batch[self.output_key] = self.transforms(batch[self.key]) + return batch + + +class RescaleMapper(BaseMapper): + """ + Rescale a sample from [0, 1] to [-1, 1] + + Args: + + config (RescaleMapperConfig): Configuration for the mapper + """ + + def __init__(self, config: RescaleMapperConfig): + super().__init__(config) + + def __call__(self, batch: Dict[str, Any], *args, **kwrags) -> Dict[str, Any]: + if isinstance(batch[self.key], list): + tmp = [] + for i, image in enumerate(batch[self.key]): + tmp.append(2 * image - 1) + batch[self.output_key] = tmp + else: + batch[self.output_key] = 2 * batch[self.key] - 1 + return batch diff --git a/lbm/data/mappers/mappers_config.py b/lbm/data/mappers/mappers_config.py new file mode 100644 index 0000000..6f0b843 --- /dev/null +++ b/lbm/data/mappers/mappers_config.py @@ -0,0 +1,109 @@ +from typing import Any, Callable, Dict, List, Optional + +from pydantic.dataclasses import dataclass + +from ...config import BaseConfig + + +@dataclass +class BaseMapperConfig(BaseConfig): + """ + Base configuration for mappers. + + Args: + + verbose (bool): + If True, print debug information. Defaults to False + + key (Optional[str]): + Key to apply the mapper to. Defaults to None + + output_key (Optional[str]): + Key to store the output of the mapper. Defaults to None + """ + + verbose: bool = False + key: Optional[str] = None + output_key: Optional[str] = None + + +@dataclass +class KeyRenameMapperConfig(BaseMapperConfig): + """ + Rename keys in a sample according to a key map + + Args: + + key_map (Dict[str, str]): Dictionary with the old keys as keys and the new keys as values + condition_key (Optional[str]): Key to use for the condition. Defaults to None + condition_fn (Optional[Callable[[Any], bool]]): Function to use for the condition to be met so + the key map is applied. Defaults to None. + else_key_map (Optional[Dict[str, str]]): Dictionary with the old keys as keys and the new keys as values + if the condition is not met. Defaults to None *i.e.* the original key will be used. + """ + + key_map: Dict[str, str] = None + condition_key: Optional[str] = None + condition_fn: Optional[Callable[[Any], bool]] = None + else_key_map: Optional[Dict[str, str]] = None + + def __post_init__(self): + super().__post_init__() + assert self.key_map is not None, "key_map should be provided" + assert all( + isinstance(old_key, str) and isinstance(new_key, str) + for old_key, new_key in self.key_map.items() + ), "key_map should be a dictionary with string keys and values" + if self.condition_key is not None: + assert self.condition_fn is not None, "condition_fn should be provided" + assert callable(self.condition_fn), "condition_fn should be callable" + if self.condition_fn is not None: + assert self.condition_key is not None, "condition_key should be provided" + assert isinstance( + self.condition_key, str + ), "condition_key should be a string" + if self.else_key_map is not None: + assert all( + isinstance(old_key, str) and isinstance(new_key, str) + for old_key, new_key in self.else_key_map.items() + ), "else_key_map should be a dictionary with string keys and values" + + +@dataclass +class TorchvisionMapperConfig(BaseMapperConfig): + """ + Apply torchvision transforms to a sample + + Args: + + key (str): Key to apply the transforms to + transforms (torchvision.transforms): List of torchvision transforms to apply + transforms_kwargs (Dict[str, Any]): List of kwargs for the transforms + """ + + key: str = "image" + transforms: List[str] = None + transforms_kwargs: List[Dict[str, Any]] = None + + def __post_init__(self): + super().__post_init__() + if self.transforms is None: + self.transforms = [] + if self.transforms_kwargs is None: + self.transforms_kwargs = [] + assert len(self.transforms) == len( + self.transforms_kwargs + ), "Number of transforms and kwargs should be same" + + +@dataclass +class RescaleMapperConfig(BaseMapperConfig): + """ + Rescale a sample from [0, 1] to [-1, 1] + + Args: + + key (str): Key to rescale + """ + + key: str = "image" diff --git a/lbm/data/mappers/mappers_wrapper.py b/lbm/data/mappers/mappers_wrapper.py new file mode 100644 index 0000000..7c2373c --- /dev/null +++ b/lbm/data/mappers/mappers_wrapper.py @@ -0,0 +1,31 @@ +from typing import Any, Dict, List, Union + +from .base import BaseMapper + + +class MapperWrapper: + """ + Wrapper for the mappers to allow iterating over several mappers in one go. + + Args: + + mappers (Union[List[BaseMapper], None]): List of mappers to apply to the batch + """ + + def __init__( + self, + mappers: Union[List[BaseMapper], None] = None, + ): + self.mappers = mappers + + def __call__(self, batch: Dict[str, Any]) -> Dict[str, Any]: + """ + Forward pass through all mappers + + Args: + + batch (Dict[str, Any]): batch of data + """ + for mapper in self.mappers: + batch = mapper(batch) + return batch diff --git a/lbm/inference/__init__.py b/lbm/inference/__init__.py new file mode 100644 index 0000000..925781d --- /dev/null +++ b/lbm/inference/__init__.py @@ -0,0 +1,4 @@ +from .inference import evaluate +from .utils import get_model + +__all__ = ["evaluate", "get_model"] diff --git a/lbm/inference/inference.py b/lbm/inference/inference.py new file mode 100644 index 0000000..a63f47e --- /dev/null +++ b/lbm/inference/inference.py @@ -0,0 +1,70 @@ +import logging + +import PIL +import torch +from torchvision.transforms import ToPILImage, ToTensor + +from lbm.models.lbm import LBMModel + +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger(__name__) + +ASPECT_RATIOS = { + str(512 / 2048): (512, 2048), + str(1024 / 1024): (1024, 1024), + str(2048 / 512): (2048, 512), + str(896 / 1152): (896, 1152), + str(1152 / 896): (1152, 896), + str(512 / 1920): (512, 1920), + str(640 / 1536): (640, 1536), + str(768 / 1280): (768, 1280), + str(1280 / 768): (1280, 768), + str(1536 / 640): (1536, 640), + str(1920 / 512): (1920, 512), +} + + +@torch.no_grad() +def evaluate( + model: LBMModel, + source_image: PIL.Image.Image, + num_sampling_steps: int = 1, +): + """ + Evaluate the model on an image coming from the source distribution and generate a new image from the target distribution. + + Args: + model (LBMModel): The model to evaluate. + source_image (PIL.Image.Image): The source image to evaluate the model on. + num_sampling_steps (int): The number of sampling steps to use for the model. + + Returns: + PIL.Image.Image: The generated image. + """ + + ori_h_bg, ori_w_bg = source_image.size + ar_bg = ori_h_bg / ori_w_bg + closest_ar_bg = min(ASPECT_RATIOS, key=lambda x: abs(float(x) - ar_bg)) + source_dimensions = ASPECT_RATIOS[closest_ar_bg] + + source_image = source_image.resize(source_dimensions) + + img_pasted_tensor = ToTensor()(source_image).unsqueeze(0) * 2 - 1 + batch = { + "source_image": img_pasted_tensor.cuda().to(torch.bfloat16), + } + + z_source = model.vae.encode(batch[model.source_key]) + + output_image = model.sample( + z=z_source, + num_steps=num_sampling_steps, + conditioner_inputs=batch, + max_samples=1, + ).clamp(-1, 1) + + output_image = (output_image[0].float().cpu() + 1) / 2 + output_image = ToPILImage()(output_image) + output_image.resize((ori_h_bg, ori_w_bg)) + + return output_image diff --git a/lbm/inference/utils.py b/lbm/inference/utils.py new file mode 100644 index 0000000..8681213 --- /dev/null +++ b/lbm/inference/utils.py @@ -0,0 +1,222 @@ +import logging +import os +from typing import List, Optional + +import torch +import yaml +from diffusers import FlowMatchEulerDiscreteScheduler +#from huggingface_hub import snapshot_download +from safetensors.torch import load_file + +from lbm.models.embedders import ( + ConditionerWrapper, + LatentsConcatEmbedder, + LatentsConcatEmbedderConfig, +) +from lbm.models.lbm import LBMConfig, LBMModel +from lbm.models.unets import DiffusersUNet2DCondWrapper +from lbm.models.vae import AutoencoderKLDiffusers, AutoencoderKLDiffusersConfig + + +# def get_model( +# model_dir: str, +# save_dir: Optional[str] = None, +# torch_dtype: torch.dtype = torch.bfloat16, +# device: str = "cuda", +# ) -> LBMModel: +# """Download the model from the model directory using either a local path or a path to HuggingFace Hub + +# Args: +# model_dir (str): The path to the model directory containing the model weights and config, can be a local path or a path to HuggingFace Hub +# save_dir (Optional[str]): The local path to save the model if downloading from HuggingFace Hub. Defaults to None. +# torch_dtype (torch.dtype): The torch dtype to use for the model. Defaults to torch.bfloat16. +# device (str): The device to use for the model. Defaults to "cuda". + +# Returns: +# LBMModel: The loaded model +# """ +# if not os.path.exists(model_dir): +# local_dir = snapshot_download( +# model_dir, +# local_dir=save_dir, +# ) +# model_dir = local_dir + +# model_files = os.listdir(model_dir) + +# # check yaml config file is present +# yaml_file = [f for f in model_files if f.endswith(".yaml")] +# if len(yaml_file) == 0: +# raise ValueError("No yaml file found in the model directory.") + +# # check safetensors weights file is present +# safetensors_files = sorted([f for f in model_files if f.endswith(".safetensors")]) +# ckpt_files = sorted([f for f in model_files if f.endswith(".ckpt")]) +# if len(safetensors_files) == 0 and len(ckpt_files) == 0: +# raise ValueError("No safetensors or ckpt file found in the model directory") + +# if len(model_files) == 0: +# raise ValueError("No model files found in the model directory") + +# with open(os.path.join(model_dir, yaml_file[0]), "r") as f: +# config = yaml.safe_load(f) + +# model = _get_model_from_config(**config, torch_dtype=torch_dtype) + +# if len(safetensors_files) > 0: +# logging.info(f"Loading safetensors file: {safetensors_files[-1]}") +# sd = load_file(os.path.join(model_dir, safetensors_files[-1])) +# model.load_state_dict(sd, strict=True) +# elif len(ckpt_files) > 0: +# logging.info(f"Loading ckpt file: {ckpt_files[-1]}") +# sd = torch.load( +# os.path.join(model_dir, ckpt_files[-1]), +# map_location="cpu", +# )["state_dict"] +# sd = {k[6:]: v for k, v in sd.items() if k.startswith("model.")} +# model.load_state_dict( +# sd, +# strict=True, +# ) +# model.to(device).to(torch_dtype) + +# model.eval() + +# return model + + +def _get_model_from_config( + backbone_signature: str = "stabilityai/stable-diffusion-xl-base-1.0", + vae_num_channels: int = 4, + unet_input_channels: int = 4, + timestep_sampling: str = "log_normal", + selected_timesteps: Optional[List[float]] = None, + prob: Optional[List[float]] = None, + conditioning_images_keys: Optional[List[str]] = [], + conditioning_masks_keys: Optional[List[str]] = [], + source_key: str = "source_image", + target_key: str = "source_image_paste", + bridge_noise_sigma: float = 0.0, + logit_mean: float = 0.0, + logit_std: float = 1.0, + pixel_loss_type: str = "lpips", + latent_loss_type: str = "l2", + latent_loss_weight: float = 1.0, + pixel_loss_weight: float = 0.0, + torch_dtype: torch.dtype = torch.bfloat16, + **kwargs, +): + + conditioners = [] + + denoiser = DiffusersUNet2DCondWrapper( + in_channels=unet_input_channels, # Add downsampled_image + out_channels=vae_num_channels, + center_input_sample=False, + flip_sin_to_cos=True, + freq_shift=0, + down_block_types=[ + "DownBlock2D", + "CrossAttnDownBlock2D", + "CrossAttnDownBlock2D", + ], + mid_block_type="UNetMidBlock2DCrossAttn", + up_block_types=["CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "UpBlock2D"], + only_cross_attention=False, + block_out_channels=[320, 640, 1280], + layers_per_block=2, + downsample_padding=1, + mid_block_scale_factor=1, + dropout=0.0, + act_fn="silu", + norm_num_groups=32, + norm_eps=1e-05, + cross_attention_dim=[320, 640, 1280], + transformer_layers_per_block=[1, 2, 10], + reverse_transformer_layers_per_block=None, + encoder_hid_dim=None, + encoder_hid_dim_type=None, + attention_head_dim=[5, 10, 20], + num_attention_heads=None, + dual_cross_attention=False, + use_linear_projection=True, + class_embed_type=None, + addition_embed_type=None, + addition_time_embed_dim=None, + num_class_embeds=None, + upcast_attention=None, + resnet_time_scale_shift="default", + resnet_skip_time_act=False, + resnet_out_scale_factor=1.0, + time_embedding_type="positional", + time_embedding_dim=None, + time_embedding_act_fn=None, + timestep_post_act=None, + time_cond_proj_dim=None, + conv_in_kernel=3, + conv_out_kernel=3, + projection_class_embeddings_input_dim=None, + attention_type="default", + class_embeddings_concat=False, + mid_block_only_cross_attention=None, + cross_attention_norm=None, + addition_embed_type_num_heads=64, + ).to(torch_dtype) + + if conditioning_images_keys != [] or conditioning_masks_keys != []: + + latents_concat_embedder_config = LatentsConcatEmbedderConfig( + image_keys=conditioning_images_keys, + mask_keys=conditioning_masks_keys, + ) + latent_concat_embedder = LatentsConcatEmbedder(latents_concat_embedder_config) + latent_concat_embedder.freeze() + conditioners.append(latent_concat_embedder) + + # Wrap conditioners and set to device + conditioner = ConditionerWrapper( + conditioners=conditioners, + ) + + ## VAE ## + # Get VAE model + vae_config = AutoencoderKLDiffusersConfig( + version=backbone_signature, + subfolder="vae", + tiling_size=(128, 128), + ) + vae = AutoencoderKLDiffusers(vae_config).to(torch_dtype) + vae.freeze() + vae.to(torch_dtype) + + ## Diffusion Model ## + # Get diffusion model + config = LBMConfig( + source_key=source_key, + target_key=target_key, + latent_loss_weight=latent_loss_weight, + latent_loss_type=latent_loss_type, + pixel_loss_type=pixel_loss_type, + pixel_loss_weight=pixel_loss_weight, + timestep_sampling=timestep_sampling, + logit_mean=logit_mean, + logit_std=logit_std, + selected_timesteps=selected_timesteps, + prob=prob, + bridge_noise_sigma=bridge_noise_sigma, + ) + + sampling_noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + backbone_signature, + subfolder="scheduler", + ) + + model = LBMModel( + config, + denoiser=denoiser, + sampling_noise_scheduler=sampling_noise_scheduler, + vae=vae, + conditioner=conditioner, + ).to(torch_dtype) + + return model diff --git a/lbm/models/__init__.py b/lbm/models/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/lbm/models/base/__init__.py b/lbm/models/base/__init__.py new file mode 100644 index 0000000..3850d04 --- /dev/null +++ b/lbm/models/base/__init__.py @@ -0,0 +1,4 @@ +from .base_model import BaseModel +from .model_config import ModelConfig + +__all__ = ["BaseModel", "ModelConfig"] diff --git a/lbm/models/base/base_model.py b/lbm/models/base/base_model.py new file mode 100644 index 0000000..3302a65 --- /dev/null +++ b/lbm/models/base/base_model.py @@ -0,0 +1,66 @@ +from typing import Any, Dict + +import torch +import torch.nn as nn + +from .model_config import ModelConfig + + +class BaseModel(nn.Module): + def __init__(self, config: ModelConfig): + nn.Module.__init__(self) + self.config = config + self.input_key = config.input_key + self.device = torch.device("cpu") + self.dtype = torch.float32 + + def on_fit_start(self, device: torch.device | None = None, *args, **kwargs): + """Called when the training starts + + Args: + device (Optional[torch.device], optional): The device to use. Usefull to set + relevant parameters on the model and embedder to the right device only + once at the start of the training. Defaults to None. + """ + if device is not None: + self.device = device + self.to(self.device) + + def forward(self, batch: Dict[str, Any], *args, **kwargs): + raise NotImplementedError("forward method is not implemented") + + def freeze(self): + """Freeze the model""" + self.eval() + for param in self.parameters(): + param.requires_grad = False + + def to(self, *args, **kwargs): + device, dtype, non_blocking, _ = torch._C._nn._parse_to(*args, **kwargs) + self = super().to( + device=device, + dtype=dtype, + non_blocking=non_blocking, + ) + + if device is not None: + self.device = device + if dtype is not None: + self.dtype = dtype + return self + + def compute_metrics(self, batch: Dict[str, Any], *args, **kwargs): + """Compute the metrics""" + return {} + + def sample(self, batch: Dict[str, Any], *args, **kwargs): + """Sample from the model""" + return {} + + def log_samples(self, batch: Dict[str, Any], *args, **kwargs): + """Log the samples""" + return None + + def on_train_batch_end(self, batch: Dict[str, Any], *args, **kwargs): + """Update the model an optimization is perforned on a batch.""" + pass diff --git a/lbm/models/base/model_config.py b/lbm/models/base/model_config.py new file mode 100644 index 0000000..0fd0d1b --- /dev/null +++ b/lbm/models/base/model_config.py @@ -0,0 +1,8 @@ +from pydantic.dataclasses import dataclass + +from ...config import BaseConfig + + +@dataclass +class ModelConfig(BaseConfig): + input_key: str = "image" diff --git a/lbm/models/embedders/__init__.py b/lbm/models/embedders/__init__.py new file mode 100644 index 0000000..86813b2 --- /dev/null +++ b/lbm/models/embedders/__init__.py @@ -0,0 +1,4 @@ +from .conditioners_wrapper import ConditionerWrapper +from .latents_concat import LatentsConcatEmbedder, LatentsConcatEmbedderConfig + +__all__ = ["LatentsConcatEmbedder", "LatentsConcatEmbedderConfig", "ConditionerWrapper"] diff --git a/lbm/models/embedders/base/__init__.py b/lbm/models/embedders/base/__init__.py new file mode 100644 index 0000000..6149ca4 --- /dev/null +++ b/lbm/models/embedders/base/__init__.py @@ -0,0 +1,4 @@ +from .base_conditioner import BaseConditioner +from .base_conditioner_config import BaseConditionerConfig + +__all__ = ["BaseConditioner", "BaseConditionerConfig"] diff --git a/lbm/models/embedders/base/base_conditioner.py b/lbm/models/embedders/base/base_conditioner.py new file mode 100644 index 0000000..053a33b --- /dev/null +++ b/lbm/models/embedders/base/base_conditioner.py @@ -0,0 +1,60 @@ +from typing import Any, Dict, List, Optional, Union + +import torch + +from ...base.base_model import BaseModel +from .base_conditioner_config import BaseConditionerConfig + +DIM2CONDITIONING = { + 2: "vector", + 3: "crossattn", + 4: "concat", +} + + +class BaseConditioner(BaseModel): + """This is the base class for all the conditioners. This absctacts the conditioning process + + Args: + + config (BaseConditionerConfig): The configuration of the conditioner + + Examples + ######## + + To use the conditioner, you can import the class and use it as follows: + + .. code-block:: python + + from cr.models.embedders import BaseConditioner, BaseConditionerConfig + + # Create the conditioner config + config = BaseConditionerConfig( + input_key="text", # The key for the input + unconditional_conditioning_rate=0.3, # Drops the conditioning with 30% probability during training + ) + + # Create the conditioner + conditioner = BaseConditioner(config) + """ + + def __init__(self, config: BaseConditionerConfig): + BaseModel.__init__(self, config) + self.config = config + self.input_key = config.input_key + self.dim2outputkey = DIM2CONDITIONING + self.ucg_rate = config.unconditional_conditioning_rate + + def forward( + self, batch: Dict[str, Any], force_zero_embedding: bool = False, *args, **kwargs + ): + """ + Forward pass of the embedder. + + Args: + + batch (Dict[str, Any]): A dictionary containing the input data. + force_zero_embedding (bool): Whether to force zero embedding. + This will return an embedding with all entries set to 0. Defaults to False. + """ + raise NotImplementedError("Forward pass must be implemented in child class") diff --git a/lbm/models/embedders/base/base_conditioner_config.py b/lbm/models/embedders/base/base_conditioner_config.py new file mode 100644 index 0000000..5f6eec2 --- /dev/null +++ b/lbm/models/embedders/base/base_conditioner_config.py @@ -0,0 +1,27 @@ +from typing import Literal + +from pydantic.dataclasses import dataclass + +from ....config import BaseConfig + + +@dataclass +class BaseConditionerConfig(BaseConfig): + """This is the ClipEmbedderConfig class which defines all the useful parameters to instantiate the model + + Args: + + input_key (str): The key for the input. Defaults to "text". + unconditional_conditioning_rate (float): Drops the conditioning with this probability during training. Defaults to 0.0. + """ + + input_key: str = "text" + unconditional_conditioning_rate: float = 0.0 + + def __post_init__(self): + super().__post_init__() + + assert ( + self.unconditional_conditioning_rate >= 0.0 + and self.unconditional_conditioning_rate <= 1.0 + ), "Unconditional conditioning rate should be between 0 and 1" diff --git a/lbm/models/embedders/conditioners_wrapper.py b/lbm/models/embedders/conditioners_wrapper.py new file mode 100644 index 0000000..63d720d --- /dev/null +++ b/lbm/models/embedders/conditioners_wrapper.py @@ -0,0 +1,114 @@ +import logging +from typing import Any, Dict, List, Union + +import torch +import torch.nn as nn + +from .base import BaseConditioner + +KEY2CATDIM = { + "vector": 1, + "crossattn": 2, + "concat": 1, +} + + +class ConditionerWrapper(nn.Module): + """ + Wrapper for conditioners. This class allows to apply multiple conditioners in a single forward pass. + + Args: + + conditioners (List[BaseConditioner]): List of conditioners to apply in the forward pass. + """ + + def __init__( + self, + conditioners: Union[List[BaseConditioner], None] = None, + ): + nn.Module.__init__(self) + self.conditioners = nn.ModuleList(conditioners) + self.device = torch.device("cpu") + self.dtype = torch.float32 + + def conditioner_sanity_check(self): + cond_input_keys = [] + for conditioner in self.conditioners: + cond_input_keys.append(conditioner.input_key) + + assert all([key in set(cond_input_keys) for key in self.ucg_keys]) + + def on_fit_start(self, device: torch.device | None = None, *args, **kwargs): + """Called when the training starts""" + for conditioner in self.conditioners: + conditioner.on_fit_start(device=device, *args, **kwargs) + + def forward( + self, + batch: Dict[str, Any], + ucg_keys: List[str] = None, + set_ucg_rate_zero=False, + *args, + **kwargs, + ): + """ + Forward pass through all conditioners + + Args: + + batch: batch of data + ucg_keys: keys to use for ucg. This will force zero conditioning in all the + conditioners that have input_keys in ucg_keys + set_ucg_rate_zero: set the ucg rate to zero for all the conditioners except the ones in ucg_keys + + Returns: + + Dict[str, Any]: The output of the conditioner. The output of the conditioner is a dictionary with the main key "cond" and value + is a dictionary with the keys as the type of conditioning and the value as the conditioning tensor. + """ + if ucg_keys is None: + ucg_keys = [] + wrapper_outputs = dict(cond={}) + for conditioner in self.conditioners: + if conditioner.input_key in ucg_keys: + force_zero_embedding = True + elif conditioner.ucg_rate > 0 and not set_ucg_rate_zero: + force_zero_embedding = bool(torch.rand(1) < conditioner.ucg_rate) + else: + force_zero_embedding = False + + conditioner_output = conditioner.forward( + batch, force_zero_embedding=force_zero_embedding, *args, **kwargs + ) + logging.debug( + f"conditioner:{conditioner.__class__.__name__}, input_key:{conditioner.input_key}, force_ucg_zero_embedding:{force_zero_embedding}" + ) + for key in conditioner_output: + logging.debug( + f"conditioner_output:{key}:{conditioner_output[key].shape}" + ) + if key in wrapper_outputs["cond"]: + wrapper_outputs["cond"][key] = torch.cat( + [wrapper_outputs["cond"][key], conditioner_output[key]], + KEY2CATDIM[key], + ) + else: + wrapper_outputs["cond"][key] = conditioner_output[key] + + return wrapper_outputs + + def to(self, *args, **kwargs): + """ + Move all conditioners to device and dtype + """ + device, dtype, non_blocking, _ = torch._C._nn._parse_to(*args, **kwargs) + self = super().to(device=device, dtype=dtype, non_blocking=non_blocking) + for conditioner in self.conditioners: + conditioner.to(device=device, dtype=dtype, non_blocking=non_blocking) + + if device is not None: + self.device = device + if dtype is not None: + self.dtype = dtype + + return self diff --git a/lbm/models/embedders/latents_concat/__init__.py b/lbm/models/embedders/latents_concat/__init__.py new file mode 100644 index 0000000..edc2206 --- /dev/null +++ b/lbm/models/embedders/latents_concat/__init__.py @@ -0,0 +1,4 @@ +from .latents_concat_embedder_config import LatentsConcatEmbedderConfig +from .latents_concat_embedder_model import LatentsConcatEmbedder + +__all__ = ["LatentsConcatEmbedder", "LatentsConcatEmbedderConfig"] diff --git a/lbm/models/embedders/latents_concat/latents_concat_embedder_config.py b/lbm/models/embedders/latents_concat/latents_concat_embedder_config.py new file mode 100644 index 0000000..7f8f003 --- /dev/null +++ b/lbm/models/embedders/latents_concat/latents_concat_embedder_config.py @@ -0,0 +1,31 @@ +from dataclasses import field +from typing import List, Union + +from pydantic.dataclasses import dataclass + +from ..base import BaseConditionerConfig + + +@dataclass +class LatentsConcatEmbedderConfig(BaseConditionerConfig): + """ + Configs for the LatentsConcatEmbedder embedder + + Args: + image_keys (Union[List[str], None]): Keys of the images to compute the VAE embeddings + mask_keys (Union[List[str], None]): Keys of the masks to resize + """ + + image_keys: Union[List[str], None] = field(default_factory=lambda: ["image"]) + mask_keys: Union[List[str], None] = field(default_factory=lambda: ["mask"]) + + def __post_init__(self): + super().__post_init__() + + # Make sure that at least one of the image_keys or mask_keys is provided + assert (self.image_keys is not None) or ( + self.mask_keys is not None + ), "At least one of the image_keys or mask_keys must be provided." + + self.image_keys = self.image_keys if self.image_keys is not None else [] + self.mask_keys = self.mask_keys if self.mask_keys is not None else [] diff --git a/lbm/models/embedders/latents_concat/latents_concat_embedder_model.py b/lbm/models/embedders/latents_concat/latents_concat_embedder_model.py new file mode 100644 index 0000000..387f9c6 --- /dev/null +++ b/lbm/models/embedders/latents_concat/latents_concat_embedder_model.py @@ -0,0 +1,80 @@ +from typing import Any, Dict + +import torch +import torchvision.transforms.functional as F + +from .....lbm.models.vae import AutoencoderKLDiffusers + +from ..base import BaseConditioner +from .latents_concat_embedder_config import LatentsConcatEmbedderConfig + + +class LatentsConcatEmbedder(BaseConditioner): + """ + Class computing VAE embeddings from given images and resizing the masks. + Then outputs are then concatenated to the noise in the latent space. + + Args: + config (LatentsConcatEmbedderConfig): Configs to create the embedder + """ + + def __init__(self, config: LatentsConcatEmbedderConfig): + BaseConditioner.__init__(self, config) + + def forward( + self, batch: Dict[str, Any], vae: AutoencoderKLDiffusers, *args, **kwargs + ) -> dict: + """ + Args: + batch (dict): A batch of images to be processed by this embedder. In the batch, + the images must range between [-1, 1] and the masks range between [0, 1]. + vae (AutoencoderKLDiffusers): VAE + + Returns: + output (dict): outputs + """ + + # Check if image are of the same size + dims_list = [] + for image_key in self.config.image_keys: + dims_list.append(batch[image_key].shape[-2:]) + for mask_key in self.config.mask_keys: + dims_list.append(batch[mask_key].shape[-2:]) + assert all( + dims == dims_list[0] for dims in dims_list + ), "All images and masks must have the same dimensions." + + # Find the latent dimensions + if len(self.config.image_keys) > 0: + latent_dims = ( + batch[self.config.image_keys[0]].shape[-2] // vae.downsampling_factor, + batch[self.config.image_keys[0]].shape[-1] // vae.downsampling_factor, + ) + else: + latent_dims = ( + batch[self.config.mask_keys[0]].shape[-2] // vae.downsampling_factor, + batch[self.config.mask_keys[0]].shape[-1] // vae.downsampling_factor, + ) + + outputs = [] + + # Resize the masks and concat them + for mask_key in self.config.mask_keys: + curr_latents = F.resize( + batch[mask_key], + size=latent_dims, + interpolation=F.InterpolationMode.BILINEAR, + ) + outputs.append(curr_latents) + + # Compute VAE embeddings from the images + for image_key in self.config.image_keys: + vae_embs = vae.encode(batch[image_key]) + outputs.append(vae_embs) + + # Concat all the outputs + outputs = torch.concat(outputs, dim=1) + + outputs = {self.dim2outputkey[outputs.dim()]: outputs} + + return outputs diff --git a/lbm/models/lbm/__init__.py b/lbm/models/lbm/__init__.py new file mode 100644 index 0000000..4b8e4d0 --- /dev/null +++ b/lbm/models/lbm/__init__.py @@ -0,0 +1,4 @@ +from .lbm_config import LBMConfig +from .lbm_model import LBMModel + +__all__ = ["LBMModel", "LBMConfig"] diff --git a/lbm/models/lbm/lbm_config.py b/lbm/models/lbm/lbm_config.py new file mode 100644 index 0000000..5ff1886 --- /dev/null +++ b/lbm/models/lbm/lbm_config.py @@ -0,0 +1,101 @@ +from typing import List, Literal, Optional, Tuple + +from pydantic.dataclasses import dataclass + +from ..base import ModelConfig + + +@dataclass +class LBMConfig(ModelConfig): + """This is the Config for LBM Model class which defines all the useful parameters to be used in the model. + + Args: + + source_key (str): + Key for the source image. Defaults to "source_image" + + target_key (str): + Key for the target image. Defaults to "target_image" + + mask_key (Optional[str]): + Key for the mask showing the valid pixels. Defaults to None + + latent_loss_type (str): + Loss type to use. Defaults to "l2". Choices are "l2", "l1" + + pixel_loss_type (str): + Pixel loss type to use. Defaults to "l2". Choices are "l2", "l1", "lpips" + + pixel_loss_max_size (int): + Maximum size of the image for pixel loss. + The image will be cropped to this size to reduce decoding computation cost. Defaults to 512 + + pixel_loss_weight (float): + Weight of the pixel loss. Defaults to 0.0 + + timestep_sampling (str): + Timestep sampling to use. Defaults to "uniform". Choices are "uniform" + + input_key (str): + Key for the input. Defaults to "image" + + controlnet_input_key (str): + Key for the controlnet conditioning. Defaults to "controlnet_conditioning" + + adapter_input_key (str): + Key for the adapter conditioning. Defaults to "adapter_conditioning" + + ucg_keys (Optional[List[str]]): + List of keys for which we enforce zero_conditioning during Classifier-free guidance. Defaults to None + + prediction_type (str): + Type of prediction to use. Defaults to "epsilon". Choices are "epsilon", "v_prediction", "flow + + logit_mean (Optional[float]): + Mean of the logit for the log normal distribution. Defaults to 0.0 + + logit_std (Optional[float]): + Standard deviation of the logit for the log normal distribution. Defaults to 1.0 + + guidance_scale (Optional[float]): + The guidance scale. Useful for finetunning guidance distilled diffusion models. Defaults to None + + selected_timesteps (Optional[List[float]]): + List of selected timesteps to be sampled from if using `custom_timesteps` timestep sampling. Defaults to None + + prob (Optional[List[float]]): + List of probabilities for the selected timesteps if using `custom_timesteps` timestep sampling. Defaults to None + """ + + source_key: str = "source_image" + target_key: str = "target_image" + mask_key: Optional[str] = None + latent_loss_weight: float = 1.0 + latent_loss_type: Literal["l2", "l1"] = "l2" + pixel_loss_type: Literal["l2", "l1", "lpips"] = "l2" + pixel_loss_max_size: int = 512 + pixel_loss_weight: float = 0.0 + timestep_sampling: Literal["uniform", "log_normal", "custom_timesteps"] = "uniform" + logit_mean: Optional[float] = 0.0 + logit_std: Optional[float] = 1.0 + selected_timesteps: Optional[List[float]] = None + prob: Optional[List[float]] = None + bridge_noise_sigma: float = 0.001 + + def __post_init__(self): + super().__post_init__() + if self.timestep_sampling == "log_normal": + assert isinstance(self.logit_mean, float) and isinstance( + self.logit_std, float + ), "logit_mean and logit_std should be float for log_normal timestep sampling" + + if self.timestep_sampling == "custom_timesteps": + assert isinstance(self.selected_timesteps, list) and isinstance( + self.prob, list + ), "timesteps and prob should be list for custom_timesteps timestep sampling" + assert len(self.selected_timesteps) == len( + self.prob + ), "timesteps and prob should be of same length for custom_timesteps timestep sampling" + assert ( + sum(self.prob) == 1 + ), "prob should sum to 1 for custom_timesteps timestep sampling" diff --git a/lbm/models/lbm/lbm_model.py b/lbm/models/lbm/lbm_model.py new file mode 100644 index 0000000..8bbe93d --- /dev/null +++ b/lbm/models/lbm/lbm_model.py @@ -0,0 +1,511 @@ +from typing import Any, Dict, List, Optional, Tuple, Union + +import lpips +import numpy as np +import torch +import torch.nn as nn +from diffusers.schedulers import FlowMatchEulerDiscreteScheduler +from tqdm import tqdm + +from ..base.base_model import BaseModel +from ..embedders import ConditionerWrapper +from ..unets import DiffusersUNet2DCondWrapper, DiffusersUNet2DWrapper +from ..vae import AutoencoderKLDiffusers +from .lbm_config import LBMConfig + +from comfy.utils import ProgressBar + +class LBMModel(BaseModel): + """This is the LBM class which defines the model. + + Args: + + config (LBMConfig): + Configuration for the model + + denoiser (Union[DiffusersUNet2DWrapper, DiffusersTransformer2DWrapper]): + Denoiser to use for the diffusion model. Defaults to None + + training_noise_scheduler (EulerDiscreteScheduler): + Noise scheduler to use for training. Defaults to None + + sampling_noise_scheduler (EulerDiscreteScheduler): + Noise scheduler to use for sampling. Defaults to None + + vae (AutoencoderKLDiffusers): + VAE to use for the diffusion model. Defaults to None + + conditioner (ConditionerWrapper): + Conditioner to use for the diffusion model. Defaults to None + """ + + @classmethod + def load_from_config(cls, config: LBMConfig): + return cls(config=config) + + def __init__( + self, + config: LBMConfig, + denoiser: Union[ + DiffusersUNet2DWrapper, + DiffusersUNet2DCondWrapper, + ] = None, + training_noise_scheduler: FlowMatchEulerDiscreteScheduler = None, + sampling_noise_scheduler: FlowMatchEulerDiscreteScheduler = None, + vae: AutoencoderKLDiffusers = None, + conditioner: ConditionerWrapper = None, + ): + BaseModel.__init__(self, config) + + self.vae = vae + self.denoiser = denoiser + self.conditioner = conditioner + self.sampling_noise_scheduler = sampling_noise_scheduler + self.training_noise_scheduler = training_noise_scheduler + self.timestep_sampling = config.timestep_sampling + self.latent_loss_type = config.latent_loss_type + self.latent_loss_weight = config.latent_loss_weight + self.pixel_loss_type = config.pixel_loss_type + self.pixel_loss_max_size = config.pixel_loss_max_size + self.pixel_loss_weight = config.pixel_loss_weight + self.logit_mean = config.logit_mean + self.logit_std = config.logit_std + self.prob = config.prob + self.selected_timesteps = config.selected_timesteps + self.source_key = config.source_key + self.target_key = config.target_key + self.mask_key = config.mask_key + self.bridge_noise_sigma = config.bridge_noise_sigma + + self.num_iterations = nn.Parameter( + torch.tensor(0, dtype=torch.float32), requires_grad=False + ) + if self.pixel_loss_type == "lpips" and self.pixel_loss_weight > 0: + self.lpips_loss = lpips.LPIPS(net="vgg") + + else: + self.lpips_loss = None + + def on_fit_start(self, device: torch.device | None = None, *args, **kwargs): + """Called when the training starts""" + super().on_fit_start(device=device, *args, **kwargs) + if self.vae is not None: + self.vae.on_fit_start(device=device, *args, **kwargs) + if self.conditioner is not None: + self.conditioner.on_fit_start(device=device, *args, **kwargs) + + def forward(self, batch: Dict[str, Any], step=0, batch_idx=0, *args, **kwargs): + + self.num_iterations += 1 + + # Get inputs/latents + if self.vae is not None: + vae_inputs = batch[self.target_key] + z = self.vae.encode(vae_inputs) + downsampling_factor = self.vae.downsampling_factor + else: + z = batch[self.target_key] + downsampling_factor = 1 + + if self.mask_key in batch: + valid_mask = batch[self.mask_key].bool()[:, 0, :, :].unsqueeze(1) + invalid_mask = ~valid_mask + valid_mask_for_latent = ~torch.max_pool2d( + invalid_mask.float(), + downsampling_factor, + downsampling_factor, + ).bool() + valid_mask_for_latent = valid_mask_for_latent.repeat((1, z.shape[1], 1, 1)) + + else: + valid_mask = torch.ones_like(batch[self.target_key]).bool() + valid_mask_for_latent = torch.ones_like(z).bool() + + source_image = batch[self.source_key] + source_image = torch.nn.functional.interpolate( + source_image, + size=batch[self.target_key].shape[-2:], + mode="bilinear", + align_corners=False, + ).to(z.dtype) + if self.vae is not None: + z_source = self.vae.encode(source_image) + + else: + z_source = source_image + + # Get conditionings + conditioning = self._get_conditioning(batch, *args, **kwargs) + + # Sample a timestep + timestep = self._timestep_sampling(n_samples=z.shape[0], device=z.device) + sigmas = None + + # Create interpolant + sigmas = self._get_sigmas( + self.training_noise_scheduler, timestep, n_dim=4, device=z.device + ) + noisy_sample = ( + sigmas * z_source + + (1.0 - sigmas) * z + + self.bridge_noise_sigma + * (sigmas * (1.0 - sigmas)) ** 0.5 + * torch.randn_like(z) + ) + + for i, t in enumerate(timestep): + if t.item() == self.training_noise_scheduler.timesteps[0]: + noisy_sample[i] = z_source[i] + + # Predict noise level using denoiser + prediction = self.denoiser( + sample=noisy_sample, + timestep=timestep, + conditioning=conditioning, + *args, + **kwargs, + ) + + target = z_source - z + denoised_sample = noisy_sample - prediction * sigmas + target_pixels = batch[self.target_key] + + # Compute loss + if self.latent_loss_weight > 0: + loss = self.latent_loss(prediction, target.detach(), valid_mask_for_latent) + latent_recon_loss = loss.mean() + + else: + loss = torch.zeros(z.shape[0], device=z.device) + latent_recon_loss = torch.zeros_like(loss) + + if self.pixel_loss_weight > 0: + denoised_sample = self._predicted_x_0( + model_output=prediction, + sample=noisy_sample, + sigmas=sigmas, + ) + pixel_loss = self.pixel_loss( + denoised_sample, target_pixels.detach(), valid_mask + ) + loss += self.pixel_loss_weight * pixel_loss + + else: + pixel_loss = torch.zeros_like(latent_recon_loss) + + return { + "loss": loss.mean(), + "latent_recon_loss": latent_recon_loss, + "pixel_recon_loss": pixel_loss.mean(), + "predicted_hr": denoised_sample, + "noisy_sample": noisy_sample, + } + + def latent_loss(self, prediction, model_input, valid_latent_mask): + if self.latent_loss_type == "l2": + return torch.mean( + ( + (prediction * valid_latent_mask - model_input * valid_latent_mask) + ** 2 + ).reshape(model_input.shape[0], -1), + 1, + ) + elif self.latent_loss_type == "l1": + return torch.mean( + torch.abs( + prediction * valid_latent_mask - model_input * valid_latent_mask + ).reshape(model_input.shape[0], -1), + 1, + ) + else: + raise NotImplementedError( + f"Loss type {self.latent_loss_type} not implemented" + ) + + def pixel_loss(self, prediction, model_input, valid_mask): + + latent_crop = self.pixel_loss_max_size // self.vae.downsampling_factor + input_crop = self.pixel_loss_max_size + + crop_h = max((prediction.shape[2] - latent_crop), 0) + crop_w = max((prediction.shape[3] - latent_crop), 0) + + input_crop_h = max((model_input.shape[2] - self.pixel_loss_max_size), 0) + input_crop_w = max((model_input.shape[3] - self.pixel_loss_max_size), 0) + + # image random cropping + if crop_h == 0: + offset_h = 0 + else: + offset_h = torch.randint(0, crop_h, (1,)).item() + + if crop_w == 0: + offset_w = 0 + else: + offset_w = torch.randint(0, crop_w, (1,)).item() + input_offset_h = offset_h * self.vae.downsampling_factor + input_offset_w = offset_w * self.vae.downsampling_factor + + prediction = prediction[ + :, + :, + crop_h + - offset_h : min(crop_h - offset_h + latent_crop, prediction.shape[2]), + crop_w + - offset_w : min(crop_w - offset_w + latent_crop, prediction.shape[3]), + ] + + model_input = model_input[ + :, + :, + input_crop_h + - input_offset_h : min( + input_crop_h - input_offset_h + input_crop, model_input.shape[2] + ), + input_crop_w + - input_offset_w : min( + input_crop_w - input_offset_w + input_crop, model_input.shape[3] + ), + ] + + valid_mask = valid_mask[ + :, + :, + input_crop_h + - input_offset_h : min( + input_crop_h - input_offset_h + input_crop, valid_mask.shape[2] + ), + input_crop_w + - input_offset_w : min( + input_crop_w - input_offset_w + input_crop, valid_mask.shape[3] + ), + ] + + decoded_prediction = self.vae.decode(prediction).clamp(-1, 1) + + if self.pixel_loss_type == "l2": + return torch.mean( + ( + (decoded_prediction * valid_mask - model_input * valid_mask) ** 2 + ).reshape(model_input.shape[0], -1), + 1, + ) + + elif self.pixel_loss_type == "l1": + return torch.mean( + torch.abs( + decoded_prediction * valid_mask - model_input * valid_mask + ).reshape(model_input.shape[0], -1), + 1, + ) + + elif self.pixel_loss_type == "lpips": + return self.lpips_loss( + decoded_prediction * valid_mask, model_input * valid_mask + ).mean() + + def _get_conditioning( + self, + batch: Dict[str, Any], + ucg_keys: List[str] = None, + set_ucg_rate_zero=False, + *args, + **kwargs, + ): + """ + Get the conditionings + """ + if self.conditioner is not None: + return self.conditioner( + batch, + ucg_keys=ucg_keys, + set_ucg_rate_zero=set_ucg_rate_zero, + vae=self.vae, + *args, + **kwargs, + ) + else: + return None + + def _timestep_sampling(self, n_samples=1, device="cpu"): + if self.timestep_sampling == "uniform": + idx = torch.randint( + 0, + self.training_noise_scheduler.config.num_train_timesteps, + (n_samples,), + device="cpu", + ) + return self.training_noise_scheduler.timesteps[idx].to(device=device) + + elif self.timestep_sampling == "log_normal": + u = torch.normal( + mean=self.logit_mean, + std=self.logit_std, + size=(n_samples,), + device="cpu", + ) + u = torch.nn.functional.sigmoid(u) + indices = ( + u * self.training_noise_scheduler.config.num_train_timesteps + ).long() + return self.training_noise_scheduler.timesteps[indices].to(device=device) + + elif self.timestep_sampling == "custom_timesteps": + idx = np.random.choice(len(self.selected_timesteps), n_samples, p=self.prob) + + return torch.tensor( + self.selected_timesteps, device=device, dtype=torch.long + )[idx] + + def _predicted_x_0( + self, + model_output, + sample, + sigmas=None, + ): + """ + Predict x_0, the orinal denoised sample, using the model output and the timesteps depending on the prediction type. + """ + pred_x_0 = sample - model_output * sigmas + return pred_x_0 + + def _get_sigmas( + self, scheduler, timesteps, n_dim=4, dtype=torch.float32, device="cpu" + ): + sigmas = scheduler.sigmas.to(device=device, dtype=dtype) + schedule_timesteps = scheduler.timesteps.to(device) + timesteps = timesteps.to(device) + step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + @torch.no_grad() + def sample( + self, + z: torch.Tensor, + num_steps: int = 20, + conditioner_inputs: Optional[Dict[str, Any]] = None, + max_samples: Optional[int] = None, + verbose: bool = False, + ): + self.sampling_noise_scheduler.set_timesteps( + sigmas=np.linspace(1, 1 / num_steps, num_steps) + ) + + sample = z + + # Get conditioning + conditioning = self._get_conditioning( + conditioner_inputs, set_ucg_rate_zero=True, device=z.device + ) + + # If max_samples parameter is provided, limit the number of samples + if max_samples is not None: + sample = sample[:max_samples] + + if conditioning: + conditioning["cond"] = { + k: v[:max_samples] for k, v in conditioning["cond"].items() + } + comfy_pbar = ProgressBar(num_steps) + for i, t in tqdm(enumerate(self.sampling_noise_scheduler.timesteps), total=num_steps): + if hasattr(self.sampling_noise_scheduler, "scale_model_input"): + denoiser_input = self.sampling_noise_scheduler.scale_model_input( + sample, t + ) + + else: + denoiser_input = sample + + # Predict noise level using denoiser using conditionings + pred = self.denoiser( + sample=denoiser_input, + timestep=t.to(z.device).repeat(denoiser_input.shape[0]), + conditioning=conditioning, + ) + + # Make one step on the reverse diffusion process + sample = self.sampling_noise_scheduler.step( + pred, t, sample, return_dict=False + )[0] + if i < len(self.sampling_noise_scheduler.timesteps) - 1: + timestep = ( + self.sampling_noise_scheduler.timesteps[i + 1] + .to(z.device) + .repeat(sample.shape[0]) + ) + sigmas = self._get_sigmas( + self.sampling_noise_scheduler, timestep, n_dim=4, device=z.device + ) + sample = sample + self.bridge_noise_sigma * ( + sigmas * (1.0 - sigmas) + ) ** 0.5 * torch.randn_like(sample) + sample = sample.to(z.dtype) + comfy_pbar.update(1) + + if self.vae is not None: + decoded_sample = self.vae.decode(sample) + + else: + decoded_sample = sample + + return decoded_sample + + def log_samples( + self, + batch: Dict[str, Any], + input_shape: Optional[Tuple[int, int, int]] = None, + max_samples: Optional[int] = None, + num_steps: Union[int, List[int]] = 20, + ): + if isinstance(num_steps, int): + num_steps = [num_steps] + + logs = {} + + N = max_samples if max_samples is not None else len(batch[self.source_key]) + + batch = {k: v[:N] for k, v in batch.items()} + + # infer input shape based on VAE configuration if not passed + if input_shape is None: + if self.vae is not None: + # get input pixel size of the vae + input_shape = batch[self.target_key].shape[2:] + # rescale to latent size + input_shape = ( + self.vae.latent_channels, + input_shape[0] // self.vae.downsampling_factor, + input_shape[1] // self.vae.downsampling_factor, + ) + else: + raise ValueError( + "input_shape must be passed when no VAE is used in the model" + ) + + for num_step in num_steps: + source_image = batch[self.source_key] + source_image = torch.nn.functional.interpolate( + source_image, + size=batch[self.target_key].shape[2:], + mode="bilinear", + align_corners=False, + ).to(dtype=self.dtype) + if self.vae is not None: + z = self.vae.encode(source_image) + + else: + z = source_image + + with torch.autocast(dtype=self.dtype, device_type="cuda"): + logs[f"samples_{num_step}_steps"] = self.sample( + z, + num_steps=num_step, + conditioner_inputs=batch, + max_samples=N, + ) + + return logs diff --git a/lbm/models/unets/__init__.py b/lbm/models/unets/__init__.py new file mode 100644 index 0000000..380af60 --- /dev/null +++ b/lbm/models/unets/__init__.py @@ -0,0 +1,14 @@ +""" +This module contains a collection of U-Net models. +The :mod:`cr.models.unets` module includes the following classes: + +- :class:`DiffusersUNet2DWrapper`: A 2D U-Net model for diffusers. +- :class:`DiffusersUNet2DCondWrapper`: A 2D U-Net model for diffusers with conditional input. +""" + +from .unet import DiffusersUNet2DCondWrapper, DiffusersUNet2DWrapper + +__all__ = [ + "DiffusersUNet2DWrapper", + "DiffusersUNet2DCondWrapper", +] diff --git a/lbm/models/unets/unet.py b/lbm/models/unets/unet.py new file mode 100644 index 0000000..1b4b7c2 --- /dev/null +++ b/lbm/models/unets/unet.py @@ -0,0 +1,148 @@ +from typing import Dict, List, Optional, Union + +import torch +from diffusers.models import UNet2DConditionModel, UNet2DModel + + +class DiffusersUNet2DWrapper(UNet2DModel): + """ + Wrapper for the UNet2DModel from diffusers + + See diffusers' UNet2DModel for more details + """ + + def __init__(self, *args, **kwargs): + UNet2DModel.__init__(self, *args, **kwargs) + + def forward( + self, + sample: torch.Tensor, + timestep: Union[torch.Tensor, float, int], + conditioning: Dict[str, torch.Tensor] = None, + *args, + **kwargs, + ): + """ + The forward pass of the model + + Args: + + sample (torch.Tensor): The input sample + timesteps (Union[torch.Tensor, float, int]): The number of timesteps + """ + if conditioning is not None: + class_labels = conditioning["cond"].get("vector", None) + concat = conditioning["cond"].get("concat", None) + + else: + class_labels = None + concat = None + + if concat is not None: + sample = torch.cat([sample, concat], dim=1) + + return super().forward(sample, timestep, class_labels).sample + + def freeze(self): + """ + Freeze the model + """ + self.eval() + for param in self.parameters(): + param.requires_grad = False + + +class DiffusersUNet2DCondWrapper(UNet2DConditionModel): + """ + Wrapper for the UNet2DConditionModel from diffusers + + See diffusers' Unet2DConditionModel for more details + """ + + def __init__(self, *args, **kwargs): + UNet2DConditionModel.__init__(self, *args, **kwargs) + # BaseModel.__init__(self, config=ModelConfig()) + + def forward( + self, + sample: torch.Tensor, + timestep: Union[torch.Tensor, float, int], + conditioning: Dict[str, torch.Tensor], + ip_adapter_cond_embedding: Optional[List[torch.Tensor]] = None, + down_block_additional_residuals: torch.Tensor = None, + mid_block_additional_residual: torch.Tensor = None, + down_intrablock_additional_residuals: torch.Tensor = None, + *args, + **kwargs, + ): + """ + The forward pass of the model + + Args: + + sample (torch.Tensor): The input sample + timesteps (Union[torch.Tensor, float, int]): The number of timesteps + conditioning (Dict[str, torch.Tensor]): The conditioning data + down_block_additional_residuals (List[torch.Tensor]): Residuals for the down blocks. + These residuals typically are used for the controlnet. + mid_block_additional_residual (List[torch.Tensor]): Residuals for the mid blocks. + These residuals typically are used for the controlnet. + down_intrablock_additional_residuals (List[torch.Tensor]): Residuals for the down intrablocks. + These residuals typically are used for the T2I adapters.middle block outputs. Defaults to False + """ + + assert isinstance(conditioning, dict), "conditionings must be a dictionary" + # assert "crossattn" in conditioning["cond"], "crossattn must be in conditionings" + + class_labels = conditioning["cond"].get("vector", None) + crossattn = conditioning["cond"].get("crossattn", None) + concat = conditioning["cond"].get("concat", None) + + # concat conditioning + if concat is not None: + sample = torch.cat([sample, concat], dim=1) + + # down_intrablock_additional_residuals needs to be cloned, since unet will modify it + if down_intrablock_additional_residuals is not None: + down_intrablock_additional_residuals_clone = [ + curr_residuals.clone() + for curr_residuals in down_intrablock_additional_residuals + ] + else: + down_intrablock_additional_residuals_clone = None + + # Check diffusers.models.embeddings.py > MultiIPAdapterImageProjectionLayer > forward() for implementation + # Exepected format : List[torch.Tensor] of shape (batch_size, num_image_embeds, embed_dim) + # with length = number of ip_adapters loaded in the ip_adapter_wrapper + if ip_adapter_cond_embedding is not None: + added_cond_kwargs = { + "image_embeds": [ + ip_adapter_embedding.unsqueeze(1) + for ip_adapter_embedding in ip_adapter_cond_embedding + ] + } + else: + added_cond_kwargs = None + + return ( + super() + .forward( + sample=sample, + timestep=timestep, + encoder_hidden_states=crossattn, + class_labels=class_labels, + added_cond_kwargs=added_cond_kwargs, + down_block_additional_residuals=down_block_additional_residuals, + mid_block_additional_residual=mid_block_additional_residual, + down_intrablock_additional_residuals=down_intrablock_additional_residuals_clone, + ) + .sample + ) + + def freeze(self): + """ + Freeze the model + """ + self.eval() + for param in self.parameters(): + param.requires_grad = False diff --git a/lbm/models/utils.py b/lbm/models/utils.py new file mode 100644 index 0000000..adb59fb --- /dev/null +++ b/lbm/models/utils.py @@ -0,0 +1,377 @@ +import logging +import math +from copy import deepcopy +from typing import List, Tuple + +import torch +import torch.nn.functional as F + +TILING_METHODS = ["average", "gaussian", "linear"] + + +class Tiler: + def get_tiles( + self, + input: torch.Tensor, + tile_size: tuple, + overlap_size: tuple, + scale: int = 1, + out_channels: int = 3, + ) -> List[List[torch.tensor]]: + """Get tiles + Args: + input (torch.Tensor): input array of shape (batch_size, channels, height, width) + tile_size (tuple): tile size + overlap_size (tuple): overlap size + scale (int): scaling factor of the output wrt input + out_channels (int): number of output channels + Returns: + List[List[torch.Tensor]]: List of tiles + """ + # assert isinstance(scale, int) + assert ( + overlap_size[0] <= tile_size[0] + ), f"Overlap size {overlap_size} must be smaller than tile size {tile_size}" + assert ( + overlap_size[1] <= tile_size[1] + ), f"Overlap size {overlap_size} must be smaller than tile size {tile_size}" + + B, C, H, W = input.shape + tile_size_H, tile_size_W = tile_size + + # sets overlap to 0 if the input is smaller than the tile size (i.e. no overlap) + overlap_H, overlap_W = ( + overlap_size[0] if H > tile_size_H else 0, + overlap_size[1] if W > tile_size_W else 0, + ) + + self.output_overlap_size = ( + int(overlap_H * scale), + int(overlap_W * scale), + ) + self.tile_size = tile_size + self.output_tile_size = ( + int(tile_size_H * scale), + int(tile_size_W * scale), + ) + self.output_shape = ( + B, + out_channels, + int(H * scale), + int(W * scale), + ) + tiles = [] + logging.debug(f"(Tiler) Input shape: {(B, C, H, W)}") + logging.debug(f"(Tiler) Output shape: {self.output_shape}") + logging.debug(f"(Tiler) Tile size: {(tile_size_H, tile_size_W)}") + logging.debug(f"(Tiler) Overlap size: {(overlap_H, overlap_W)}") + # loop over all tiles in the image with overlap + for i in range(0, H, tile_size_H - overlap_H): + row = [] + for j in range(0, W, tile_size_W - overlap_W): + tile = deepcopy( + input[ + :, + :, + i : i + tile_size_H, + j : j + tile_size_W, + ] + ) + row.append(tile) + tiles.append(row) + return tiles + + def merge_tiles( + self, tiles: List[List[torch.tensor]], tiling_method: str = "gaussian" + ) -> torch.tensor: + """Merge tiles by averaging the overlaping regions + Args: + tiles (Dict[str, Tile]): dictionary of processed tiles + tiling_method (str): tiling method. Can be "average", "gaussian" or "linear" + Returns: + torch.tensor: output image + """ + if tiling_method == "average": + return self._average_merge_tiles(tiles) + elif tiling_method == "gaussian": + return self._gaussian_merge_tiles(tiles) + elif tiling_method == "linear": + return self._linear_merge_tiles(tiles) + else: + raise ValueError( + f"Unknown tiling method {tiling_method}. Available methods are {TILING_METHODS}" + ) + + def _average_merge_tiles(self, tiles: List[List[torch.tensor]]) -> torch.tensor: + """Merge tiles by averaging the overlaping regions + Args: + tiles (Dict[str, Tile]): dictionary of processed tiles + Returns: + torch.tensor: output image + """ + + output = torch.zeros(self.output_shape) + + # weights to store multiplicity + weights = torch.zeros(self.output_shape) + + _, _, output_H, output_W = self.output_shape + output_overlap_size_H, output_overlap_size_W = self.output_overlap_size + output_tile_size_H, output_tile_size_W = self.output_tile_size + + for id_i, i in enumerate( + range( + 0, + output_H, + output_tile_size_H - output_overlap_size_H, + ) + ): + for id_j, j in enumerate( + range( + 0, + output_W, + output_tile_size_W - output_overlap_size_W, + ) + ): + output[ + :, + :, + i : i + output_tile_size_H, + j : j + output_tile_size_W, + ] += ( + tiles[id_i][id_j] * 1 + ) + weights[ + :, + :, + i : i + output_tile_size_H, + j : j + output_tile_size_W, + ] += 1 + + # outputs is summed up with this multiplicity + # so we need to divide by the weights wich is either 1, 2 or 4 depending on the region + output = output / weights + return output + + def _gaussian_weights( + self, tile_width: int, tile_height: int, nbatches: int, channels: int + ): + """Generates a gaussian mask of weights for tile contributions. + + Args: + tile_width (int): width of the tile + tile_height (int): height of the tile + nbatches (int): number of batches + channels (int): number of channels + Returns: + torch.tensor: weights + """ + import numpy as np + from numpy import exp, pi, sqrt + + latent_width = tile_width + latent_height = tile_height + + var = 0.01 + midpoint = ( + latent_width - 1 + ) / 2 # -1 because index goes from 0 to latent_width - 1 + x_probs = [ + exp( + -(x - midpoint) + * (x - midpoint) + / (latent_width * latent_width) + / (2 * var) + ) + / sqrt(2 * pi * var) + for x in range(latent_width) + ] + midpoint = latent_height / 2 + y_probs = [ + exp( + -(y - midpoint) + * (y - midpoint) + / (latent_height * latent_height) + / (2 * var) + ) + / sqrt(2 * pi * var) + for y in range(latent_height) + ] + + weights = np.outer(y_probs, x_probs) + return torch.tile( + torch.tensor(weights, device="cpu"), (nbatches, channels, 1, 1) + ) + + def _gaussian_merge_tiles(self, tiles: List[List[torch.tensor]]) -> torch.tensor: + """Merge tiles by averaging the overlaping regions + Args: + List[List[torch.tensor]]: List of processed tiles + Returns: + torch.tensor: output image + """ + B, output_C, output_H, output_W = self.output_shape + output_overlap_size_H, output_overlap_size_W = self.output_overlap_size + output_tile_size_H, output_tile_size_W = self.output_tile_size + + output = torch.zeros(self.output_shape) + # weights to store multiplicity + weights = torch.zeros(self.output_shape) + + for id_i, i in enumerate( + range( + 0, + output_H, + output_tile_size_H - output_overlap_size_H, + ) + ): + for id_j, j in enumerate( + range( + 0, + output_W, + output_tile_size_W - output_overlap_size_W, + ) + ): + w = self._gaussian_weights( + tiles[id_i][id_j].shape[3], + tiles[id_i][id_j].shape[2], + B, + output_C, + ) + output[ + :, + :, + i : i + output_tile_size_H, + j : j + output_tile_size_W, + ] += ( + tiles[id_i][id_j] * w + ) + weights[ + :, + :, + i : i + output_tile_size_H, + j : j + output_tile_size_W, + ] += w + + # outputs is summed up with this multiplicity + output = output / weights + return output + + def _blend_v( + self, a: torch.Tensor, b: torch.Tensor, blend_extent: int + ) -> torch.Tensor: + blend_extent = min(a.shape[2], b.shape[2], blend_extent) + for y in range(blend_extent): + b[:, :, y, :] = a[:, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[ + :, :, y, : + ] * (y / blend_extent) + return b + + def _blend_h( + self, a: torch.Tensor, b: torch.Tensor, blend_extent: int + ) -> torch.Tensor: + blend_extent = min(a.shape[3], b.shape[3], blend_extent) + for x in range(blend_extent): + b[:, :, :, x] = a[:, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[ + :, :, :, x + ] * (x / blend_extent) + return b + + def _linear_merge_tiles(self, tiles: List[List[torch.tensor]]) -> torch.Tensor: + """Merge tiles by blending the overlaping regions + Args: + tiles (List[List[torch.tensor]]): List of processed tiles + Returns: + torch.Tensor: output image + """ + output_overlap_size_H, output_overlap_size_W = self.output_overlap_size + output_tile_size_H, output_tile_size_W = self.output_tile_size + + res_rows = [] + tiles_copy = deepcopy(tiles) + + # Cut the right and bottom overlap region + limit_i = output_tile_size_H - output_overlap_size_H + limit_j = output_tile_size_W - output_overlap_size_W + for i, tile_row in enumerate(tiles_copy): + res_row = [] + for j, tile in enumerate(tile_row): + tile_val = tile + if j > 0: + tile_val = self._blend_h( + tile_row[j - 1], tile, output_overlap_size_W + ) + tiles_copy[i][j] = tile_val + if i > 0: + tile_val = self._blend_v( + tiles_copy[i - 1][j], tile_val, output_overlap_size_H + ) + tiles_copy[i][j] = tile_val + res_row.append(tile_val[:, :, :limit_i, :limit_j]) + res_rows.append(torch.cat(res_row, dim=3)) + output = torch.cat(res_rows, dim=2) + return output + + +def extract_into_tensor( + a: torch.Tensor, t: torch.Tensor, x_shape: Tuple[int, ...] +) -> torch.Tensor: + """ + Extracts values from a tensor into a new tensor using indices from another tensor. + + :param a: the tensor to extract values from. + :param t: the tensor containing the indices. + :param x_shape: the shape of the tensor to extract values into. + :return: a new tensor containing the extracted values. + """ + + b, *_ = t.shape + out = a.gather(-1, t) + return out.reshape(b, *((1,) * (len(x_shape) - 1))) + + +def pad(x: torch.Tensor, base_h: int, base_w: int) -> torch.Tensor: + """ + Pads a tensor to the nearest multiple of base_h and base_w. + + :param x: the tensor to pad. + :param base_h: the base height. + :param base_w: the base width. + :return: the padded tensor. + """ + h, w = x.shape[-2:] + h_ = math.ceil(h / base_h) * base_h + w_ = math.ceil(w / base_w) * base_w + if w_ != w: + x = F.pad(x, (0, abs(w_ - w), 0, 0)) + if h_ != h: + x = F.pad(x, (0, 0, 0, abs(h_ - h))) + return x + + +def append_dims(x: torch.Tensor, target_dims: int) -> torch.Tensor: + """Appends dimensions to the end of a tensor until it has target_dims dimensions.""" + dims_to_append = target_dims - x.ndim + if dims_to_append < 0: + raise ValueError( + f"input has {x.ndim} dims but target_dims is {target_dims}, which is less" + ) + return x[(...,) + (None,) * dims_to_append] + + +@torch.no_grad() +def update_ema( + target_params: List[torch.Tensor], + source_params: List[torch.Tensor], + rate: float = 0.99, +): + """ + Update target parameters to be closer to those of source parameters using + an exponential moving average. + + :param target_params: the target parameter sequence. + :param source_params: the source parameter sequence. + :param rate: the EMA rate (closer to 1 means slower). + """ + for targ, src in zip(target_params, source_params): + targ.detach().mul_(rate).add_(src, alpha=1 - rate) diff --git a/lbm/models/vae/__init__.py b/lbm/models/vae/__init__.py new file mode 100644 index 0000000..9ad0c4b --- /dev/null +++ b/lbm/models/vae/__init__.py @@ -0,0 +1,4 @@ +from .autoencoderKL import AutoencoderKLDiffusers +from .autoencoderKL_config import AutoencoderKLDiffusersConfig + +__all__ = ["AutoencoderKLDiffusers", "AutoencoderKLDiffusersConfig"] diff --git a/lbm/models/vae/autoencoderKL.py b/lbm/models/vae/autoencoderKL.py new file mode 100644 index 0000000..ed0239f --- /dev/null +++ b/lbm/models/vae/autoencoderKL.py @@ -0,0 +1,136 @@ +import torch +from diffusers.models import AutoencoderKL + +from ..base.base_model import BaseModel +from ..utils import Tiler, pad +from .autoencoderKL_config import AutoencoderKLDiffusersConfig + + +class AutoencoderKLDiffusers(BaseModel): + """This is the VAE class used to work with latent models + + Args: + + config (AutoencoderKLDiffusersConfig): The config class which defines all the required parameters. + """ + + def __init__(self, config: AutoencoderKLDiffusersConfig): + BaseModel.__init__(self, config) + self.config = config + self.vae_model = AutoencoderKL.from_pretrained( + config.version, + subfolder=config.subfolder, + revision=config.revision, + ) + self.tiling_size = config.tiling_size + self.tiling_overlap = config.tiling_overlap + + # get downsampling factor + self._get_properties() + + @torch.no_grad() + def _get_properties(self): + self.has_shift_factor = ( + hasattr(self.vae_model.config, "shift_factor") + and self.vae_model.config.shift_factor is not None + ) + self.shift_factor = ( + self.vae_model.config.shift_factor if self.has_shift_factor else 0 + ) + + # set latent channels + self.latent_channels = self.vae_model.config.latent_channels + self.has_latents_mean = ( + hasattr(self.vae_model.config, "latents_mean") + and self.vae_model.config.latents_mean is not None + ) + self.has_latents_std = ( + hasattr(self.vae_model.config, "latents_std") + and self.vae_model.config.latents_std is not None + ) + self.latents_mean = self.vae_model.config.latents_mean + self.latents_std = self.vae_model.config.latents_std + + x = torch.randn(1, self.vae_model.config.in_channels, 32, 32) + z = self.encode(x) + + # set downsampling factor + self.downsampling_factor = int(x.shape[2] / z.shape[2]) + + def encode(self, x: torch.tensor, batch_size: int = 8): + latents = [] + for i in range(0, x.shape[0], batch_size): + latents.append( + self.vae_model.encode(x[i : i + batch_size]).latent_dist.sample() + ) + latents = torch.cat(latents, dim=0) + latents = (latents - self.shift_factor) * self.vae_model.config.scaling_factor + + return latents + + def decode(self, z: torch.tensor): + + if self.has_latents_mean and self.has_latents_std: + latents_mean = ( + torch.tensor(self.latents_mean) + .view(1, self.latent_channels, 1, 1) + .to(z.device, z.dtype) + ) + latents_std = ( + torch.tensor(self.latents_std) + .view(1, self.latent_channels, 1, 1) + .to(z.device, z.dtype) + ) + z = z * latents_std / self.vae_model.config.scaling_factor + latents_mean + else: + z = z / self.vae_model.config.scaling_factor + self.shift_factor + + use_tiling = ( + z.shape[2] > self.tiling_size[0] or z.shape[3] > self.tiling_size[1] + ) + + if use_tiling: + samples = [] + for i in range(z.shape[0]): + + z_i = z[i].unsqueeze(0) + + tiler = Tiler() + tiles = tiler.get_tiles( + input=z_i, + tile_size=self.tiling_size, + overlap_size=self.tiling_overlap, + scale=self.downsampling_factor, + out_channels=3, + ) + + for i, tile_row in enumerate(tiles): + for j, tile in enumerate(tile_row): + tile_shape = tile.shape + # pad tile to inference size if tile is smaller than inference size + tile = pad( + tile, + base_h=self.tiling_size[0], + base_w=self.tiling_size[1], + ) + tile_decoded = self.vae_model.decode(tile).sample + tiles[i][j] = ( + tile_decoded[ + 0, + :, + : int(tile_shape[2] * self.downsampling_factor), + : int(tile_shape[3] * self.downsampling_factor), + ] + .cpu() + .unsqueeze(0) + ) + + # merge tiles + samples.append(tiler.merge_tiles(tiles=tiles)) + + samples = torch.cat(samples, dim=0) + + else: + samples = self.vae_model.decode(z).sample + + return samples diff --git a/lbm/models/vae/autoencoderKL_config.py b/lbm/models/vae/autoencoderKL_config.py new file mode 100644 index 0000000..f5c31f1 --- /dev/null +++ b/lbm/models/vae/autoencoderKL_config.py @@ -0,0 +1,27 @@ +from typing import Tuple + +from pydantic.dataclasses import dataclass + +from ..base import ModelConfig + + +@dataclass +class AutoencoderKLDiffusersConfig(ModelConfig): + """This is the VAEConfig class which defines all the useful parameters to instantiate the model. + + Args: + + version (str): The version of the model. Defaults to "stabilityai/sdxl-vae". + subfolder (str): The subfolder of the model if loaded from another model. Defaults to "". + revision (str): The revision of the model. Defaults to "main". + input_key (str): The key of the input data in the batch. Defaults to "image". + tiling_size (Tuple[int, int]): The size of the tiling. Defaults to (64, 64). + tiling_overlap (Tuple[int, int]): The overlap of the tiling. Defaults to (16, 16). + """ + + version: str = "stabilityai/sdxl-vae" + subfolder: str = "" + revision: str = "main" + input_key: str = "image" + tiling_size: Tuple[int, int] = (64, 64) + tiling_overlap: Tuple[int, int] = (16, 16) diff --git a/lbm/trainer/__init__.py b/lbm/trainer/__init__.py new file mode 100644 index 0000000..506c801 --- /dev/null +++ b/lbm/trainer/__init__.py @@ -0,0 +1,85 @@ +""" +This module contains the training pipeline and the training configuration along with all relevant parts +of the training pipeline such as loggers and callbacks. + +The :mod:`cr.trainer` includes the following submodules: + +- :mod:`cr.trainer.trainer`: the main training pipeline class for ClipDrop. +- :mod:`cr.trainer.training_config`: the configuration for the training pipeline. +- :mod:`cr.trainer.loggers`: the loggers for logging samples to wandb. + + +Examples +######## + +Train a model using the training pipeline + +.. code-block:: python + + from cr.trainer import TrainingPipeline, TrainingConfig + from cr.data import DataPipeline, DataConfig + from pytorch_lightning import Trainer + from cr.data.datasets import DataModule, DataModuleConfig + + # Create a model to train + model = DummyModel() + + # Create a training configuration + config = TrainingConfig( + experiment_id="test", + optimizers_name=["AdamW"], + optimizers_kwargs=[{}], + learning_rates=[1e-3], + lr_schedulers_name=[None], + lr_schedulers_kwargs=[{}], + trainable_params=[["./*"]], + log_keys="txt", + log_samples_model_kwargs={ + "max_samples": 8, + "num_steps": 20, + "input_shape": (4, 32, 32), + "guidance_scale": 7.5, + } + ) + + # Create a training pipeline + pipeline = TrainingPipeline(model=model, pipeline_config=config) + + # Create a DataModule + data_module = DataModule( + train_config=DataModuleConfig( + shards_path_or_urls="your urls or paths", + decoder="pil", + shuffle_buffer_size=100, + per_worker_batch_size=32, + num_workers=4, + ), + train_filters_mappers=your_mappers_and_filters, + eval_config=DataModuleConfig( + shards_path_or_urls="your urls or paths", + decoder="pil", + shuffle_buffer_size=100, + per_worker_batch_size=32, + num_workers=4, + ), + eval_filters_mappers=your_mappers_and_filters, + ) + + # Create a trainer + trainer = Trainer( + accelerator="cuda", + max_epochs=1, + devices=1, + log_every_n_steps=1, + default_root_dir="your dir", + max_steps=2, + ) + + # Train the model + trainer.fit(pipeline, data_module) +""" + +from .trainer import TrainingPipeline +from .training_config import TrainingConfig + +__all__ = ["TrainingPipeline", "TrainingConfig"] diff --git a/lbm/trainer/loggers.py b/lbm/trainer/loggers.py new file mode 100644 index 0000000..0447239 --- /dev/null +++ b/lbm/trainer/loggers.py @@ -0,0 +1,324 @@ +import logging +import math +from typing import Any, Dict, List, Tuple + +import numpy as np +import torch +import wandb +from PIL import Image, ImageDraw, ImageFont +from pytorch_lightning import Trainer +from pytorch_lightning.callbacks import Callback +from pytorch_lightning.utilities import rank_zero_only +from torchvision.utils import make_grid + +from ..trainer import TrainingPipeline + +logging.basicConfig(level=logging.INFO) + + +def create_grid_texts( + texts: List[str], + n_cols: int = 4, + image_size: Tuple[int] = (512, 512), + font_size: int = 40, + margin: int = 5, + offset: int = 5, +) -> Image.Image: + """ + Create a grid of white images containing the given texts. + + Args: + texts (List[str]): List of strings to be drawn on images. + n_cols (int): Number of columns in the grid. + image_size (tuple): Size of the generated images (width, height). + font_size (int): Font size of the text. + margin (int): Margin around the text. + offset (int): Offset between lines. + + Returns: + PIL.Image: List of generated images as a grid + """ + + images = [] + font = ImageFont.load_default(size=font_size) + + for text in texts: + img = Image.new("RGB", image_size, color="white") + draw = ImageDraw.Draw(img) + margin_ = margin + offset_ = offset + for line in wrap_text( + text=text, draw=draw, max_width=image_size[0] - 2 * margin_, font=font + ): + draw.text((margin_, offset_), line, font=font, fill="black") + offset_ += font_size + images.append(img) + + # create a pil grid + n_rows = math.ceil(len(images) / n_cols) + grid = Image.new( + "RGB", (n_cols * image_size[0], n_rows * image_size[1]), color="white" + ) + for i, img in enumerate(images): + grid.paste(img, (i % n_cols * image_size[0], i // n_cols * image_size[1])) + + return grid + + +def wrap_text( + text: str, draw: ImageDraw.Draw, max_width: int, font: ImageFont +) -> List[str]: + """ + Wrap text to fit within a specified width when drawn. + It will return to the new line when the text is larger than the max_width. + + Args: + text (str): The text to be wrapped. + draw (ImageDraw.Draw): The draw object to calculate text size. + max_width (int): The maximum width for the wrapped text. + font (ImageFont): The font used for the text. + + Returns: + List[str]: List of wrapped lines. + """ + lines = [] + current_line = "" + for letter in text: + if draw.textbbox((0, 0), current_line + letter, font=font)[2] <= max_width: + current_line += letter + else: + lines.append(current_line) + current_line = letter + lines.append(current_line) + return lines + + +class WandbSampleLogger(Callback): + """ + Logger for logging samples to wandb. This logger is used to log images, text, and metrics to wandb. + + Args: + log_batch_freq (int): The frequency of logging samples to wandb. Default is 100. + """ + + def __init__(self, log_batch_freq: int = 100): + super().__init__() + self.log_batch_freq = log_batch_freq + + def on_train_batch_end( + self, + trainer: Trainer, + pl_module: TrainingPipeline, + outputs: Dict[str, Any], + batch: Any, + batch_idx: int, + ) -> None: + self.log_samples(trainer, pl_module, outputs, batch, batch_idx, split="train") + self._process_logs(trainer, outputs, split="train") + + def on_validation_batch_end( + self, + trainer: Trainer, + pl_module: TrainingPipeline, + outputs: Dict[str, Any], + batch: Any, + batch_idx: int, + ) -> None: + self.log_samples(trainer, pl_module, outputs, batch, batch_idx, split="val") + self._process_logs(trainer, outputs, split="val") + + @rank_zero_only + @torch.no_grad() + def log_samples( + self, + trainer: Trainer, + pl_module: TrainingPipeline, + outputs: Dict[str, Any], + batch: Dict[str, Any], + batch_idx: int, + split: str = "train", + ) -> None: + if hasattr(pl_module, "log_samples"): + if batch_idx % self.log_batch_freq == 0: + is_training = pl_module.training + if is_training: + pl_module.eval() + + logs = pl_module.log_samples(batch) + logs = self._process_logs(trainer, logs, split=split) + + if is_training: + pl_module.train() + else: + logging.warning( + "log_img method not found in LightningModule. Skipping image logging." + ) + + @rank_zero_only + def _process_logs( + self, trainer, logs: Dict[str, Any], rescale=True, split="train" + ) -> Dict[str, Any]: + for key, value in logs.items(): + if isinstance(value, torch.Tensor): + value = value.detach().cpu() + if value.dim() == 4: + images = value + if rescale: + images = (images + 1.0) / 2.0 + grid = make_grid(images, nrow=4) + grid = grid.permute(1, 2, 0) + grid = grid.mul(255).clamp(0, 255).to(torch.uint8) + logs[key] = grid.numpy() + trainer.logger.experiment.log( + {f"{key}/{split}": [wandb.Image(Image.fromarray(logs[key]))]}, + step=trainer.global_step, + ) + + # Scalar tensor + if value.dim() == 1 or value.dim() == 0: + value = value.float().numpy() + trainer.logger.experiment.log( + {f"{key}/{split}": value}, step=trainer.global_step + ) + + # list of string (e.g. text) + if isinstance(value, list): + if isinstance(value[0], str): + pil_image_texts = create_grid_texts(value) + wandb_image = wandb.Image(pil_image_texts) + trainer.logger.experiment.log( + {f"{key}/{split}": [wandb_image]}, + step=trainer.global_step, + ) + + # dict of tensors (e.g. metrics) + if isinstance(value, dict): + for k, v in value.items(): + if isinstance(v, torch.Tensor): + value[k] = v.detach().cpu().numpy() + trainer.logger.experiment.log( + {f"{key}/{split}": value}, step=trainer.global_step + ) + + if isinstance(value, int) or isinstance(value, float): + trainer.logger.experiment.log( + {f"{key}/{split}": value}, step=trainer.global_step + ) + + return logs + + +class TensorBoardSampleLogger(Callback): + """ + Logger for logging samples to tensorboard. This logger is used to log images, text, and metrics to tensorboard. + + Args: + log_batch_freq (int): The frequency of logging samples to tensorboard. Default is 100. + """ + + def __init__(self, log_batch_freq: int = 100): + super().__init__() + self.log_batch_freq = log_batch_freq + + def on_train_batch_end( + self, + trainer: Trainer, + pl_module: TrainingPipeline, + outputs: Dict[str, Any], + batch: Any, + batch_idx: int, + ) -> None: + self.log_samples(trainer, pl_module, outputs, batch, batch_idx, split="train") + self._process_logs(trainer, outputs, split="train") + + def on_validation_batch_end( + self, + trainer: Trainer, + pl_module: TrainingPipeline, + outputs: Dict[str, Any], + batch: Any, + batch_idx: int, + ) -> None: + self.log_samples(trainer, pl_module, outputs, batch, batch_idx, split="val") + self._process_logs(trainer, outputs, split="val") + + @rank_zero_only + @torch.no_grad() + def log_samples( + self, + trainer: Trainer, + pl_module: TrainingPipeline, + outputs: Dict[str, Any], + batch: Dict[str, Any], + batch_idx: int, + split: str = "train", + ) -> None: + if hasattr(pl_module, "log_samples"): + if batch_idx % self.log_batch_freq == 0: + is_training = pl_module.training + if is_training: + pl_module.eval() + + logs = pl_module.log_samples(batch) + logs = self._process_logs(trainer, logs, split=split) + + if is_training: + pl_module.train() + else: + logging.warning( + "log_img method not found in LightningModule. Skipping image logging." + ) + + @rank_zero_only + def _process_logs( + self, trainer, logs: Dict[str, Any], rescale=True, split="train" + ) -> Dict[str, Any]: + for key, value in logs.items(): + if isinstance(value, torch.Tensor): + value = value.detach().cpu() + if value.dim() == 4: + images = value + if rescale: + images = (images + 1.0) / 2.0 + grid = make_grid(images, nrow=4) + # grid = grid.permute(1, 2, 0) + grid = grid.mul(255).clamp(0, 255).to(torch.uint8) + logs[key] = grid.numpy() + trainer.logger.experiment.add_image( + f"{key}/{split}", + logs[key], + trainer.global_step, + ) + + # Scalar tensor + if value.dim() == 1 or value.dim() == 0: + value = value.float().numpy() + trainer.logger.experiment.add_scalar( + f"{key}/{split}", value, trainer.global_step + ) + + # list of string (e.g. text) + if isinstance(value, list): + if isinstance(value[0], str): + pil_image_texts = create_grid_texts(value) + trainer.logger.experiment.add_image( + f"{key}/{split}", + np.transpose(np.array(pil_image_texts), (2, 0, 1)), + trainer.global_step, + ) + + # dict of tensors (e.g. metrics) + if isinstance(value, dict): + for k, v in value.items(): + if isinstance(v, torch.Tensor): + value[k] = v.detach().cpu().numpy() + trainer.logger.experiment.add_scalar( + f"{key}/{split}", value, trainer.global_step + ) + + if isinstance(value, int) or isinstance(value, float): + trainer.logger.experiment.add_scalar( + f"{key}/{split}", value, trainer.global_step + ) + + return logs diff --git a/lbm/trainer/trainer.py b/lbm/trainer/trainer.py new file mode 100644 index 0000000..035fe3a --- /dev/null +++ b/lbm/trainer/trainer.py @@ -0,0 +1,199 @@ +import importlib +import logging +import re +import time +from typing import Any, Dict + +import pytorch_lightning as pl +import torch + +from ..models.base.base_model import BaseModel +from .training_config import TrainingConfig + +logging.basicConfig(level=logging.INFO) + + +class TrainingPipeline(pl.LightningModule): + """ + Main Training Pipeline class + + Args: + + model (BaseModel): The model to train + pipeline_config (TrainingConfig): The configuration for the training pipeline + verbose (bool): Whether to print logs in the console. Default is False. + """ + + def __init__( + self, + model: BaseModel, + pipeline_config: TrainingConfig, + verbose: bool = False, + **kwargs, + ): + super().__init__() + + self.model = model + self.pipeline_config = pipeline_config + self.log_samples_model_kwargs = pipeline_config.log_samples_model_kwargs + + # save hyperparameters. + self.save_hyperparameters(ignore="model") + self.save_hyperparameters({"model_config": model.config.to_dict()}) + + # logger. + self.verbose = verbose + + # setup logging. + log_keys = pipeline_config.log_keys + + if isinstance(log_keys, str): + log_keys = [log_keys] + + if log_keys is None: + log_keys = [] + + self.log_keys = log_keys + + def on_fit_start(self) -> None: + self.model.on_fit_start(device=self.device) + if self.global_rank == 0: + self.timer = time.perf_counter() + + def on_train_batch_end( + self, outputs: Dict[str, Any], batch: Any, batch_idx: int + ) -> None: + if self.global_rank == 0: + logging.debug("on_train_batch_end") + self.model.on_train_batch_end(batch) + + average_time_frequency = 10 + if self.global_rank == 0 and batch_idx % average_time_frequency == 0: + delta = time.perf_counter() - self.timer + logging.info( + f"Average time per batch {batch_idx} took {delta / (batch_idx + 1)} seconds" + ) + + def configure_optimizers(self) -> torch.optim.Optimizer: + """ + Setup optimizers and learning rate schedulers. + """ + optimizers = [] + lr = self.pipeline_config.learning_rate + param_list = [] + n_params = 0 + param_list_ = {"params": []} + for name, param in self.model.named_parameters(): + for regex in self.pipeline_config.trainable_params: + pattern = re.compile(regex) + if re.match(pattern, name): + if param.requires_grad: + param_list_["params"].append(param) + n_params += param.numel() + + param_list.append(param_list_) + + logging.info(f"Number of trainable parameters: {n_params}") + + optimizer_cls = getattr( + importlib.import_module("torch.optim"), + self.pipeline_config.optimizer_name, + ) + optimizer = optimizer_cls( + param_list, lr=lr, **self.pipeline_config.optimizer_kwargs + ) + optimizers.append(optimizer) + + self.optims = optimizers + schedulers_config = self.configure_lr_schedulers() + + for name, param in self.model.named_parameters(): + set_grad_false = True + for regex in self.pipeline_config.trainable_params: + pattern = re.compile(regex) + if re.match(pattern, name): + if param.requires_grad: + set_grad_false = False + if set_grad_false: + param.requires_grad = False + + num_trainable_params = sum( + p.numel() for p in self.model.parameters() if p.requires_grad + ) + + logging.info(f"Number of trainable parameters: {num_trainable_params}") + + schedulers_config = self.configure_lr_schedulers() + + if schedulers_config is None: + return optimizers + + return optimizers, [ + schedulers_config_ for schedulers_config_ in schedulers_config + ] + + def configure_lr_schedulers(self): + schedulers_config = [] + if self.pipeline_config.lr_scheduler_name is None: + scheduler = None + schedulers_config.append(scheduler) + else: + scheduler_cls = getattr( + importlib.import_module("torch.optim.lr_scheduler"), + self.pipeline_config.lr_scheduler_name, + ) + scheduler = scheduler_cls( + self.optims[0], + **self.pipeline_config.lr_scheduler_kwargs, + ) + lr_scheduler_config = { + "scheduler": scheduler, + "interval": self.pipeline_config.lr_scheduler_interval, + "monitor": "val_loss", + "frequency": self.pipeline_config.lr_scheduler_frequency, + } + schedulers_config.append(lr_scheduler_config) + + if all([scheduler is None for scheduler in schedulers_config]): + return None + + return schedulers_config + + def training_step(self, train_batch: Dict[str, Any], batch_idx: int) -> dict: + model_output = self.model(train_batch) + loss = model_output["loss"] + logging.info(f"loss: {loss}") + return { + "loss": loss, + "batch_idx": batch_idx, + } + + def validation_step(self, val_batch: Dict[str, Any], val_idx: int) -> dict: + loss = self.model(val_batch, device=self.device)["loss"] + + metrics = self.model.compute_metrics(val_batch) + + return {"loss": loss, "metrics": metrics} + + def log_samples(self, batch: Dict[str, Any]): + logging.debug("log_samples") + logs = self.model.log_samples( + batch, + **self.log_samples_model_kwargs, + ) + + if logs is not None: + N = min([logs[keys].shape[0] for keys in logs]) + else: + N = 0 + + # Log inputs + if self.log_keys is not None: + for key in self.log_keys: + if key in batch: + if N > 0: + logs[key] = batch[key][:N] + else: + logs[key] = batch[key] + + return logs diff --git a/lbm/trainer/training_config.py b/lbm/trainer/training_config.py new file mode 100644 index 0000000..1c86633 --- /dev/null +++ b/lbm/trainer/training_config.py @@ -0,0 +1,82 @@ +from dataclasses import field +from typing import List, Literal, Optional, Union + +from pydantic.dataclasses import dataclass + +from ..config import BaseConfig + + +@dataclass +class TrainingConfig(BaseConfig): + """ + Configuration for the training pipeline + + Args: + + experiment_id (str): + The experiment id for the training run. If not provided, a random id will be generated. + optimizer_name (str): + The optimizer to use. Default is "AdamW". Choices are "Adam", "AdamW", "Adadelta", "Adagrad", "RMSprop", "SGD" + optimizer_kwargs (Dict[str, Any]) + The optimizer kwargs. Default is [{}] + learning_rate (float): + The learning rate to use. Default is 1e-3 + lr_scheduler_name (str): + The learning rate scheduler to use. Default is None. Choices are "StepLR", "CosineAnnealingLR", + "CosineAnnealingWarmRestarts", "ReduceLROnPlateau", "ExponentialLR" + lr_scheduler_kwargs (Dict[str, Any]) + The learning rate scheduler kwargs. Default is [{}] + lr_scheduler_interval (str): + The learning rate scheduler interval. Default is ["step"]. Choices are "step", "epoch" + lr_scheduler_frequency (int): + The learning rate scheduler frequency. Default is 1 + metrics (List[str]) + The metrics to use. Default is None + tracking_metrics: Optional[List[str]] + The metrics to track. Default is None + backup_every (int): + The frequency to backup the model. Default is 50. + trainable_params (Union[str, List[str]]): + Regexes indicateing the parameters to train. + Default is [["./*"]] (i.e. all parameters are trainable) + log_keys: Union[str, List[str]]: + The keys to log when sampling from the model. Default is "txt" + log_samples_model_kwargs (Dict[str, Any]): + The kwargs for logging samples from the model. Default is { + "max_samples": 4, + "num_steps": 20, + "input_shape": None, + } + """ + + experiment_id: Optional[str] = None + optimizer_name: Literal[ + "Adam", "AdamW", "Adadelta", "Adagrad", "RMSprop", "SGD" + ] = field(default_factory=lambda: "AdamW") + optimizer_kwargs: Optional[dict] = field(default_factory=lambda: {}) + learning_rate: float = field(default_factory=lambda: 1e-3) + lr_scheduler_name: Optional[ + Literal[ + "StepLR", + "CosineAnnealingLR", + "CosineAnnealingWarmRestarts", + "ReduceLROnPlateau", + "ExponentialLR", + None, + ] + ] = None + lr_scheduler_kwargs: Optional[dict] = field(default_factory=lambda: {}) + lr_scheduler_interval: Optional[Literal["step", "epoch", None]] = "step" + lr_scheduler_frequency: Optional[int] = 1 + metrics: Optional[List[str]] = None + tracking_metrics: Optional[List[str]] = None + backup_every: int = 50 + trainable_params: List[str] = field(default_factory=lambda: ["./*"]) + log_keys: Optional[Union[str, List[str]]] = "txt" + log_samples_model_kwargs: Optional[dict] = field( + default_factory=lambda: { + "max_samples": 4, + "num_steps": 20, + "input_shape": None, + } + ) diff --git a/lbm/trainer/utils.py b/lbm/trainer/utils.py new file mode 100644 index 0000000..1e827ed --- /dev/null +++ b/lbm/trainer/utils.py @@ -0,0 +1,193 @@ +import logging +import os +import re +import time +from typing import Dict, List, Literal, Optional, Tuple + +import torch + + +class StateDictAdapter: + """ + StateDictAdapter for adapting the state dict of a model to a checkpoint state dict. + + This class will iterate over all keys in the checkpoint state dict and filter them by a list of regex keys. + For each matching key, the class will adapt the checkpoint state dict to the model state dict. + Depending on the target size, the class will add missing blocks or cut the block. + When adding missing blocks, the class will use a strategy to fill the missing blocks: either adding zeros or normal random values. + + Example: + + ``` + adapter = StateDictAdapter() + new_state_dict = adapter( + model_state_dict=model.state_dict(), + checkpoint_state_dict=state_dict, + regex_keys=[ + r"class_embedding.linear_1.weight", + r"conv_in.weight", + r"(down_blocks|up_blocks)\.\d+\.attentions\.\d+\.transformer_blocks\.\d+\.attn\d+\.(to_k|to_v)\.weight", + r"mid_block\.attentions\.\d+\.transformer_blocks\.\d+\.attn\d+\.(to_k|to_v)\.weight" + ] + ) + ``` + + Args: + model_state_dict (Dict[str, torch.Tensor]): The model state dict. + checkpoint_state_dict (Dict[str, torch.Tensor]): The checkpoint state dict. + regex_keys (Optional[List[str]]): A list of regex keys to adapt the checkpoint state dict. Defaults to None. + Passing a list of regex will drastically reduce the latency. + If None, all keys in the checkpoint state dict will be adapted. + strategy (Literal["zeros", "normal"], optional): The strategy to fill the missing blocks. Defaults to "normal". + + """ + + def _create_block( + self, + shape: List[int], + strategy: Literal["zeros", "normal"], + input: torch.Tensor = None, + ): + if strategy == "zeros": + return torch.zeros(shape) + elif strategy == "normal": + if input is not None: + mean = input.mean().item() + std = input.std().item() + return torch.randn(shape) * std + mean + else: + return torch.randn(shape) + else: + raise ValueError(f"Unknown strategy {strategy}") + + def __call__( + self, + model_state_dict: Dict[str, torch.Tensor], + checkpoint_state_dict: Dict[str, torch.Tensor], + regex_keys: Optional[List[str]] = None, + strategy: Literal["zeros", "normal"] = "normal", + ): + start = time.perf_counter() + # if no regex keys are provided, we use all keys in the model state dict + if regex_keys is None: + regex_keys = list(model_state_dict.keys()) + + # iterate over all keys in the checkpoint state dict + for checkpoint_key in list(checkpoint_state_dict.keys()): + # iterate over all regex keys + for regex_key in regex_keys: + if re.match(regex_key, checkpoint_key): + dst_shape = model_state_dict[checkpoint_key].shape + src_shape = checkpoint_state_dict[checkpoint_key].shape + + ## Sizes adapter + # if length of shapes are different, we need to unsqueeze or squeeze the tensor + if len(dst_shape) != len(src_shape): + # in the case [a] vs [a, b] -> unsqueeze [a, 1] + if len(src_shape) == 1: + checkpoint_state_dict[checkpoint_key] = ( + checkpoint_state_dict[checkpoint_key].unsqueeze(1) + ) + logging.info( + f"Unsqueeze {checkpoint_key}: {src_shape} -> {checkpoint_state_dict[checkpoint_key].shape}" + ) + # in the case [a, b] vs [a] -> squeeze [a] + elif len(dst_shape) == 1: + checkpoint_state_dict[checkpoint_key] = ( + checkpoint_state_dict[checkpoint_key][:, 0] + ) + logging.info( + f"Squeeze {checkpoint_key}: {src_shape} -> {checkpoint_state_dict[checkpoint_key].shape}" + ) + # in the other cases, raise an error + else: + raise ValueError( + f"Shapes of {checkpoint_key} are different: {dst_shape} != {src_shape}" + ) + + # update the shapes + dst_shape = model_state_dict[checkpoint_key].shape + src_shape = checkpoint_state_dict[checkpoint_key].shape + assert len(dst_shape) == len( + src_shape + ), f"Shapes of {checkpoint_key} are different: {dst_shape} != {src_shape}" + + ## Shapes adapter + # modify the checkpoint state dict only if the shapes are different + if dst_shape != src_shape: + # create a copy of the tensor + tmp = torch.clone(checkpoint_state_dict[checkpoint_key]) + + # iterate over all dimensions + for i in range(len(dst_shape)): + if dst_shape[i] != src_shape[i]: + diff = dst_shape[i] - src_shape[i] + + # if the difference is greater than 0, we need to add missing blocks + if diff > 0: + missing_shape = list(tmp.shape) + missing_shape[i] = diff + missing = self._create_block( + shape=missing_shape, + strategy=strategy, + input=tmp, + ) + tmp = torch.cat((tmp, missing), dim=i) + logging.info( + f"Adapting {checkpoint_key} with strategy:{strategy} from shape {src_shape} to {dst_shape}" + ) + # if the difference is less than 0, we need to cut the block + else: + tmp = tmp.narrow(i, 0, dst_shape[i]) + logging.info( + f"Adapting {checkpoint_key} by narrowing from shape {src_shape} to {dst_shape}" + ) + + checkpoint_state_dict[checkpoint_key] = tmp + end = time.perf_counter() + logging.info(f"StateDictAdapter took {end-start:.2f} seconds") + return checkpoint_state_dict + + +class StateDictRenamer: + """ + StateDictRenamer for renaming keys in a checkpoint state dict. + This class will iterate over all keys in the checkpoint state dict and rename them according to a rename dict. + + Example: + + ``` + renamer = StateDictRenamer() + new_state_dict = renamer( + checkpoint_state_dict=state_dict, + rename_dict={ + "add_embedding.linear_1.weight": "class_embedding.linear_1.weight", + "add_embedding.linear_1.bias": "class_embedding.linear_1.bias", + "add_embedding.linear_2.weight": "class_embedding.linear_2.weight", + "add_embedding.linear_2.bias": "class_embedding.linear_2.bias", + } + ) + ``` + + Args: + + checkpoint_state_dict (Dict[str, torch.Tensor]): The checkpoint state dict. + rename_dict (Dict[str, str]): The dictionary mapping the old keys to new keys + """ + + def __call__( + self, + checkpoint_state_dict: Dict[str, torch.Tensor], + rename_dict: Dict[str, str], + ) -> Dict[str, torch.Tensor]: + for old_key, new_key in rename_dict.items(): + if old_key not in checkpoint_state_dict: + logging.warning(f"Key {old_key} not found in checkpoint state dict") + continue + else: + assert ( + new_key not in checkpoint_state_dict + ), f"Key {new_key} already exists in checkpoint state dict" + checkpoint_state_dict[new_key] = checkpoint_state_dict.pop(old_key) + logging.info(f"Renaming {old_key} to {new_key}") + return checkpoint_state_dict diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..0182082 --- /dev/null +++ b/nodes.py @@ -0,0 +1,135 @@ +import os +import torch +from tqdm import tqdm + +from accelerate import init_empty_weights +from accelerate.utils import set_module_tensor_to_device + +import folder_paths +import comfy.model_management as mm +from comfy.utils import load_torch_file + +from .utils import get_model_from_config + +script_directory = os.path.dirname(os.path.abspath(__file__)) + + +class LoadLBMModel: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}), + + "base_precision": (["fp32", "bf16", "fp16"], {"default": "bf16"}), + "load_device": (["main_device", "offload_device"], {"default": "cuda", "tooltip": "Initialize the model on the main device or offload device"}), + }, + } + + + RETURN_TYPES = ("LBM_MODEL",) + RETURN_NAMES = ("model", ) + FUNCTION = "loadmodel" + CATEGORY = "LBMWrapper" + + def loadmodel(self, model, base_precision, load_device="main_device"): + + base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision] + + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + if load_device == "main_device": + transformer_load_device = device + else: + transformer_load_device = offload_device + + model_path = folder_paths.get_full_path_or_raise("diffusion_models", model) + + config = { + "vae_num_channels": 4, + "unet_input_channels": 4, + "timestep_sampling": "custom_timesteps", + "selected_timesteps": [250, 500, 750, 1000], + "prob": [0.25, 0.25, 0.25, 0.25], + "conditioning_images_keys": [], + "conditioning_masks_keys": [], + "source_key": "source_image", + "target_key": "source_image_paste", + "bridge_noise_sigma": 0.005, + } + sd = load_torch_file(model_path, device=offload_device, safe_load=True) + + with init_empty_weights(): + unet = get_model_from_config(**config) + + print("Using accelerate to load and assign model weights to device...") + param_count = sum(1 for _ in unet.named_parameters()) + for name, param in tqdm(unet.named_parameters(), + desc=f"Loading transformer parameters to {transformer_load_device}", + total=param_count, + leave=True): + set_module_tensor_to_device(unet, name, device=transformer_load_device, dtype=base_dtype, value=sd[name]) + + return(unet, ) + + +class LBMSampler: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("LBM_MODEL",), + "image": ("IMAGE", ), + "steps": ("INT", {"default": 30, "min": 1}), + }, + } + + RETURN_TYPES = ("IMAGE", ) + RETURN_NAMES = ("image",) + FUNCTION = "process" + CATEGORY = "LBMWrapper" + + def process(self, model, image, steps): + + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + + mm.unload_all_models() + mm.cleanup_models() + mm.soft_empty_cache() + + + input_image = image.clone().permute(0, 3, 1, 2).to(device, model.dtype) * 2 - 1 + + batch = { + "source_image": input_image, + } + + z_source = model.vae.encode(batch[model.source_key]) + + model.to(device) + + result = model.sample( + z=z_source, + num_steps=steps, + conditioner_inputs=batch, + max_samples=1, + ).clamp(-1, 1) + + out = result.permute(0, 2, 3, 1).cpu().float() + out = (out + 1) / 2 + + model.to(offload_device) + mm.soft_empty_cache() + + return out, + +NODE_CLASS_MAPPINGS = { + "LoadLBMModel": LoadLBMModel, + "LBMSampler": LBMSampler, + } +NODE_DISPLAY_NAME_MAPPINGS = { + "LoadLBMModel": "Load LBM Model", + "LBMSampler": "LBMSampler", + } + diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..858f9d6 --- /dev/null +++ b/requirements.txt @@ -0,0 +1,2 @@ +diffusers +accelerate \ No newline at end of file diff --git a/utils.py b/utils.py new file mode 100644 index 0000000..3edd0b6 --- /dev/null +++ b/utils.py @@ -0,0 +1,177 @@ +import os +from typing import List + +import torch +from diffusers import FlowMatchEulerDiscreteScheduler +from PIL import Image +from torchvision import transforms + +from .lbm.models.embedders import ( + ConditionerWrapper, + LatentsConcatEmbedder, + LatentsConcatEmbedderConfig, +) +from .lbm.models.lbm import LBMConfig, LBMModel +from .lbm.models.unets import DiffusersUNet2DCondWrapper +from .lbm.models.vae import AutoencoderKLDiffusers, AutoencoderKLDiffusersConfig + + +def get_model_from_config( + backbone_signature: str = "stabilityai/stable-diffusion-xl-base-1.0", + vae_num_channels: int = 4, + unet_input_channels: int = 4, + timestep_sampling: str = "log_normal", + selected_timesteps: List[float] = None, + prob: List[float] = None, + conditioning_images_keys: List[str] = [], + conditioning_masks_keys: List[str] = ["mask"], + source_key: str = "source_image", + target_key: str = "source_image_paste", + bridge_noise_sigma: float = 0.0, +): + + conditioners = [] + + denoiser = DiffusersUNet2DCondWrapper( + in_channels=unet_input_channels, # Add downsampled_image + out_channels=vae_num_channels, + center_input_sample=False, + flip_sin_to_cos=True, + freq_shift=0, + down_block_types=[ + "DownBlock2D", + "CrossAttnDownBlock2D", + "CrossAttnDownBlock2D", + ], + mid_block_type="UNetMidBlock2DCrossAttn", + up_block_types=["CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "UpBlock2D"], + only_cross_attention=False, + block_out_channels=[320, 640, 1280], + layers_per_block=2, + downsample_padding=1, + mid_block_scale_factor=1, + dropout=0.0, + act_fn="silu", + norm_num_groups=32, + norm_eps=1e-05, + cross_attention_dim=[320, 640, 1280], + transformer_layers_per_block=[1, 2, 10], + reverse_transformer_layers_per_block=None, + encoder_hid_dim=None, + encoder_hid_dim_type=None, + attention_head_dim=[5, 10, 20], + num_attention_heads=None, + dual_cross_attention=False, + use_linear_projection=True, + class_embed_type=None, + addition_embed_type=None, + addition_time_embed_dim=None, + num_class_embeds=None, + upcast_attention=None, + resnet_time_scale_shift="default", + resnet_skip_time_act=False, + resnet_out_scale_factor=1.0, + time_embedding_type="positional", + time_embedding_dim=None, + time_embedding_act_fn=None, + timestep_post_act=None, + time_cond_proj_dim=None, + conv_in_kernel=3, + conv_out_kernel=3, + projection_class_embeddings_input_dim=None, + attention_type="default", + class_embeddings_concat=False, + mid_block_only_cross_attention=None, + cross_attention_norm=None, + addition_embed_type_num_heads=64, + ).to(torch.bfloat16) + + if conditioning_images_keys != [] or conditioning_masks_keys != []: + + latents_concat_embedder_config = LatentsConcatEmbedderConfig( + image_keys=conditioning_images_keys, + mask_keys=conditioning_masks_keys, + ) + latent_concat_embedder = LatentsConcatEmbedder(latents_concat_embedder_config) + latent_concat_embedder.freeze() + conditioners.append(latent_concat_embedder) + + # Wrap conditioners and set to device + conditioner = ConditionerWrapper( + conditioners=conditioners, + ) + + ## VAE ## + # Get VAE model + vae_config = AutoencoderKLDiffusersConfig( + version=backbone_signature, + subfolder="vae", + tiling_size=(128, 128), + ) + vae = AutoencoderKLDiffusers(vae_config).to(torch.bfloat16) + vae.freeze() + vae.to(torch.bfloat16) + + ## Diffusion Model ## + # Get diffusion model + config = LBMConfig( + source_key=source_key, + target_key=target_key, + timestep_sampling=timestep_sampling, + selected_timesteps=selected_timesteps, + prob=prob, + bridge_noise_sigma=bridge_noise_sigma, + ) + + sampling_noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + backbone_signature, + subfolder="scheduler", + ) + + model = LBMModel( + config, + denoiser=denoiser, + sampling_noise_scheduler=sampling_noise_scheduler, + vae=vae, + conditioner=conditioner, + ).to(torch.bfloat16) + + return model + + +def extract_object(birefnet, img): + # Data settings + image_size = (1024, 1024) + transform_image = transforms.Compose( + [ + transforms.Resize(image_size), + transforms.ToTensor(), + transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), + ] + ) + + image = img + input_images = transform_image(image).unsqueeze(0).cuda() + + # Prediction + with torch.no_grad(): + preds = birefnet(input_images)[-1].sigmoid().cpu() + pred = preds[0].squeeze() + pred_pil = transforms.ToPILImage()(pred) + mask = pred_pil.resize(image.size) + image = Image.composite(image, Image.new("RGB", image.size, (127, 127, 127)), mask) + return image, mask + + +def resize_and_center_crop(image, target_width, target_height): + original_width, original_height = image.size + scale_factor = max(target_width / original_width, target_height / original_height) + resized_width = int(round(original_width * scale_factor)) + resized_height = int(round(original_height * scale_factor)) + resized_image = image.resize((resized_width, resized_height), Image.LANCZOS) + left = (resized_width - target_width) / 2 + top = (resized_height - target_height) / 2 + right = (resized_width + target_width) / 2 + bottom = (resized_height + target_height) / 2 + cropped_image = resized_image.crop((left, top, right, bottom)) + return cropped_image \ No newline at end of file