diff --git a/.eslintrc.json b/.eslintrc.json index 22a9edee5..53691dbd9 100644 --- a/.eslintrc.json +++ b/.eslintrc.json @@ -90,6 +90,8 @@ "getENActiveTab": "readonly", "quickApplyStyle": "readonly", "quickSaveStyle": "readonly", + "setupExtraNetworks": "readonly", + "showNetworks": "readonly", // from python "localization": "readonly", // progressbar.js @@ -112,6 +114,8 @@ "idbPut": "readonly", "idbDel": "readonly", "idbAdd": "readonly", + // changelog.js + "initChangelog": "readonly", // notification.js "sendNotification": "readonly" }, diff --git a/.gitignore b/.gitignore index 2bc819860..dca4e17ad 100644 --- a/.gitignore +++ b/.gitignore @@ -44,6 +44,14 @@ tunableop_results*.csv !webui.sh !package.json +# pyinstaller +*.spec +build/ +dist/ + +# dynamically generated +/repositories/ip-instruct/ + # all dynamic stuff /extensions/**/* /outputs/**/* @@ -59,7 +67,6 @@ tunableop_results*.csv .vscode/ .idea/ /localizations - .*/ # force included @@ -67,3 +74,4 @@ tunableop_results*.csv !/models/VAE-approx/model.pt !/models/Reference !/models/Reference/**/* + diff --git a/.pylintrc b/.pylintrc index 45869a8c3..2d8e4869b 100644 --- a/.pylintrc +++ b/.pylintrc @@ -31,6 +31,9 @@ ignore-paths=/usr/lib/.*$, modules/xadapter, modules/meissonic, modules/omnigen, + modules/instantir, + modules/consistory, + modules/pulid/eva_clip, repositories, extensions-builtin/sd-webui-agent-scheduler, extensions-builtin/sd-extension-chainner/nodes, @@ -130,7 +133,8 @@ confidence=HIGH, INFERENCE_FAILURE, UNDEFINED # disable=C,R,W -disable=bad-inline-option, +disable=abstract-method, + bad-inline-option, bare-except, broad-exception-caught, chained-comparison, @@ -174,6 +178,7 @@ disable=bad-inline-option, unnecessary-dict-index-lookup, unnecessary-dunder-call, unnecessary-lambda, + unnecessary-lambda-assigment, use-dict-literal, use-symbolic-message-instead, unknown-option-value, diff --git a/.ruff.toml b/.ruff.toml index fe8ac4f87..2a29fb089 100644 --- a/.ruff.toml +++ b/.ruff.toml @@ -26,6 +26,9 @@ exclude = [ "modules/xadapter", "modules/meissonic", "modules/omnigen", + "modules/instantir", + "modules/consistory", + "modules/pulid/eva_clip", "repositories", "extensions-builtin/sd-extension-chainner/nodes", "extensions-builtin/sd-webui-agent-scheduler", diff --git a/CHANGELOG.md b/CHANGELOG.md index 3c2f1e3bb..2726a1864 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,9 +1,118 @@ # Change Log for SD.Next -## Update for 2024-10-25 +## Update for 2024-11-06 -Improvements: -- Model selector: +Smaller release just few days after the last one, but with some important fixes and improvements. +This release can be considered an LTS release before we kick off the next round of major updates. + +- Docs: + - add built-in [changelog](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) search + since changelog is the best up-to-date source of info + go to system -> changelog and search/highligh/navigate directly in UI! + +- Integrations: + - [PuLID](https://github.com/ToTheBeginning/PuLID): Pure and Lightning ID Customization via Contrastive Alignment + - advanced method of face transfer with better quality as well as control over identity and appearance + try it out, likely the best quality available for sdxl models + - select in *scripts -> pulid* + - compatible with *sdxl* + - can be used in xyz grid + - [InstantIR](https://github.com/instantX-research/InstantIR): Blind Image Restoration with Instant Generative Reference + - alternative to traditional `img2img` with more control over restoration process + - select in *image -> scripts -> instantir* + - compatible with *sdxl* + - *note*: after used once it cannot be unloaded without reloading base model + - [ConsiStory](https://github.com/NVlabs/consistory): Consistent Image Generation + - create consistent anchor image and then generate images that are consistent with anchor + - select in *scripts -> consistory* + - compatible with *sdxl* + - *note*: very resource intensive and not compatible with model offloading + - *note*: changing default parameters can lead to unexpected results and/or failures + - *note*: after used once it cannot be unloaded without reloading base model + - [MiaoshouAI PromptGen v2.0](https://huggingface.co/MiaoshouAI/Florence-2-base-PromptGen-v2.0) base and large: + - *in process -> visual query* + - caption modes: + `` generate tags + ``, ``, `` caption image + `` image composition + ``, `` detailed caption and tags with optional analyze + +- Model improvements: + - SD3: ControlNets: + - *InstantX Canny, Pose, Depth, Tile* + - *Alimama Inpainting, SoftEdge* + - *note*: that just like with FLUX.1 or any large model, ControlNet are also large and can push your system over the limit + e.g. SD3 controlnets vary from 1GB to over 4GB in size + - SD3: all-in-one safetensors + - *examples*: [large](https://civitai.com/models/882666/sd35-large-google-flan?modelVersionId=1003031), [medium](https://civitai.com/models/900327) + - *note*: enable *bnb* on-the-fly quantization for even bigger gains + +- Workflow improvements: + - XYZ grid: + - optional time benchmark info to individual images + - optional add params to individual images + - create video from generated grid images + supports all standard video types and interpolation + - UI: + - add additional [hotkeys](https://github.com/vladmandic/automatic/wiki/Hotkeys) + - add show networks on startup setting + - better mapping of networks previews + - optimize networks display load +- Other: + - Installer: + - Log `venv` and package search paths + - Auto-remove invalid packages from `venv/site-packages` + e.g. packages starting with `~` which are left-over due to windows access violation + - Requirements: update + - Model loader: + - Report modules included in safetensors when attempting to load a model + - CLI: + - refactor command line params + run `webui.sh`/`webui.bat` with `--help` to see all options + - added `cli/model-metadata.py` to display metadata in any safetensors file + - added `cli/model-keys.py` to quicky display content of any safetensors file + - Internal: + - Repo: move screenshots to GH pages + +- Fixes: + - custom watermark add alphablending + - detailer min/max size as fractions of image size + - ipadapter load on-demand + - ipadapter face use correct yolo model + - list diffusers remove duplicates + - fix legacy extensions access to shared objects + - fix diffusers load from folder + - fix lora enum logging on windows + - fix xyz grid with batch count + - fix xyz grid include images + - fix xyz skip on interrupted + - fix vqa models ignoring hfcache folder setting + - fix network height in standard vs modern ui + - fix k-diff enum on startup + - fix text2video scripts + - dont uninstall flash-attn + - move downloads of some auxillary models to hfcache instead of models folder + +## Update for 2024-10-29 + +### Highlights for 2024-10-29 + +- Support for **all SD3.x variants** + *SD3.0-Medium, SD3.5-Medium, SD3.5-Large, SD3.0-Large-Turbo* +- Allow quantization using `bitsandbytes` on-the-fly during models load + Load any variant of SD3.x or FLUX.1 and apply quantization during load without the need for pre-quantized models +- Allow for custom model URL in standard model selector + Can be used to specify any model from *HuggingFace* or *CivitAI* +- Full support for `torch==2.5.1` +- New wiki articles: [Gated Access](https://github.com/vladmandic/automatic/wiki/Gated), [Quantization](https://github.com/vladmandic/automatic/wiki/Quantization), [Offloading](https://github.com/vladmandic/automatic/wiki/Offload) + +Plus tons of smaller improvements and cumulative fixes reported since last release + +[README](https://github.com/vladmandic/automatic/blob/master/README.md) | [CHANGELOG](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) | [WiKi](https://github.com/vladmandic/automatic/wiki) | [Discord](https://discord.com/invite/sd-next-federal-batch-inspectors-1101998836328697867) + +### Details for 2024-10-29 + +- model selector: - change-in-behavior - when typing, it will auto-load model as soon as exactly one match is found - allows entering model that are not on the list which triggers huggingface search @@ -14,16 +123,40 @@ Improvements: e.g. `https://civitai.com/api/download/models/72396?type=Model&format=SafeTensor&size=full&fp=fp16` - auto-search-and-download can be disabled in settings -> models -> auto-download this also disables reference models as they are auto-downloaded on first use as well -- SD3 loader enhancements +- sd3 enhancements: + - allow on-the-fly bnb quantization during load - report when loading incomplete model - - handle missing model components + - handle missing model components during load - handle component preloading - native lora handler + - support for all sd35 variants: *medium/large/large-turbo* - gguf transformer loader (prototype) -- OpenVINO: add accuracy option -- ZLUDA: guess GPU arch +- flux.1 enhancements: + - allow on-the-fly bnb quantization during load +- samplers: + - support for original k-diffusion samplers + select via *scripts -> k-diffusion -> sampler* +- ipadapter: + - list available adapters based on loaded model type + - add adapter `ostris consistency` for sd15/sdxl +- detailer: + - add `[prompt]` to refine/defailer prompts as placeholder referencing original prompt +- torch + - use `torch==2.5.1` by default on supported platforms + - CUDA set device memory limit + in *settings -> compute settings -> torch memory limit* + default=0 meaning no limit, if set torch will limit memory usage to specified fraction + *note*: this is not a hard limit, torch will try to stay under this value +- compute backends: + - OpenVINO: add accuracy option + - ZLUDA: guess GPU arch +- major model load refactor +- wiki: new articles + - [Gated Access Wiki](https://github.com/vladmandic/automatic/wiki/Gated) + - [Quantization Wiki](https://github.com/vladmandic/automatic/wiki/Quantization) + - [Offloading Wiki](https://github.com/vladmandic/automatic/wiki/Offload) -Fixes: +fixes: - fix send-to-control - fix k-diffusion - fix sd3 img2img and hires diff --git a/CITATION.cff b/CITATION.cff index f7fd4bba3..c39efe474 100644 --- a/CITATION.cff +++ b/CITATION.cff @@ -24,5 +24,5 @@ abstract: >- generation keywords: - stablediffusion diffusers sdnext -license: AGPL-3.0 +license: Apache-2.0 date-released: 2022-12-24 diff --git a/LICENSE.txt b/LICENSE.txt index 211d32e75..0ad25db4b 100644 --- a/LICENSE.txt +++ b/LICENSE.txt @@ -1,8 +1,6 @@ GNU AFFERO GENERAL PUBLIC LICENSE Version 3, 19 November 2007 - Copyright (c) 2023 AUTOMATIC1111 - Copyright (C) 2007 Free Software Foundation, Inc. Everyone is permitted to copy and distribute verbatim copies of this license document, but changing it is not allowed. @@ -635,8 +633,8 @@ the "copyright" line and a pointer to where the full notice is found. Copyright (C) This program is free software: you can redistribute it and/or modify - it under the terms of the GNU Affero General Public License as published by - the Free Software Foundation, either version 3 of the License, or + it under the terms of the GNU Affero General Public License as published + by the Free Software Foundation, either version 3 of the License, or (at your option) any later version. This program is distributed in the hope that it will be useful, diff --git a/README.md b/README.md index ac96b2e32..a2caa5cb7 100644 --- a/README.md +++ b/README.md @@ -50,12 +50,13 @@ All individual features are not listed here, instead check [ChangeLog](CHANGELOG
*Main interface using **StandardUI***: -![Screenshot-Dark](html/screenshot-text2image.jpg) +![screenshot-text2image](https://github.com/user-attachments/assets/87ac2813-65c2-45f4-80b8-67b26ccf5cd6) *Main interface using **ModernUI***: -![Screenshot-Dark](html/screenshot-modernui-f1.jpg) -![Screenshot-Dark](html/screenshot-modernui.jpg) -![Screenshot-Dark](html/screenshot-modernui-sd3.jpg) + +![screenshot-modernui-f1](https://github.com/user-attachments/assets/b509a280-8d3b-48b5-8525-363bad8c1ed2) +![screenshot-modernui](https://github.com/user-attachments/assets/fef33127-f733-4e78-b66e-17729539512f) +![screenshot-modernui-sd3](https://github.com/user-attachments/assets/1ed02ecc-23e4-4fda-8ae5-2d7393dc530c) For screenshots and informations on other available themes, see [Themes Wiki](https://github.com/vladmandic/automatic/wiki/Themes) @@ -63,12 +64,13 @@ For screenshots and informations on other available themes, see [Themes Wiki](ht ## Model support -Additional models will be added as they become available and there is public interest in them +Additional models will be added as they become available and there is public interest in them +See [models overview](https://github.com/vladmandic/automatic/wiki/Models) for details on each model, including their architecture, complexity and other info - [RunwayML Stable Diffusion](https://github.com/Stability-AI/stablediffusion/) 1.x and 2.x *(all variants)* - [StabilityAI Stable Diffusion XL](https://github.com/Stability-AI/generative-models) -- [StabilityAI Stable Diffusion 3 Medium](https://stability.ai/news/stable-diffusion-3-medium) -- [Stable Diffusion 3.5 Large](https://huggingface.co/stabilityai/stable-diffusion-3.5-large) +- [StabilityAI Stable Diffusion](https://stability.ai/news/stable-diffusion-3-medium) +- [Stable Diffusion 3.x](https://huggingface.co/stabilityai/stable-diffusion-3.5-large) 3.0 Medium, 3.5 Medium, 3.5 Large, 3.5 Large Turbo - [StabilityAI Stable Video Diffusion](https://huggingface.co/stabilityai/stable-video-diffusion-img2vid) Base, XT 1.0, XT 1.1 - [StabilityAI Stable Cascade](https://github.com/Stability-AI/StableCascade) *Full* and *Lite* - [Black Forest Labs FLUX.1](https://blackforestlabs.ai/announcing-black-forest-labs/) Dev, Schnell @@ -136,13 +138,13 @@ Also supported are modifiers such as: ## Examples *IP Adapters*: -![Screenshot-IPAdapter](html/screenshot-ipadapter.jpg) +![screenshot-ipadapter](https://github.com/user-attachments/assets/92830894-845c-49ec-92d9-18c8a577d04f) *Color grading*: -![Screenshot-Color](html/screenshot-color.jpg) +![screenshot-control](https://github.com/user-attachments/assets/cdad2722-ae7c-4c9c-94d6-5ea35a4b1356) *InstantID*: -![Screenshot-InstantID](html/screenshot-instantid.jpg) +![screenshot-instantid](https://github.com/user-attachments/assets/f38a5660-32b3-4235-9da1-c79eccf5372f) > [!IMPORTANT] > - Loading any model other than standard SD 1.x / SD 2.x requires use of backend **Diffusers** @@ -173,92 +175,34 @@ Also supported are modifiers such as: Once SD.Next is installed, simply run `webui.ps1` or `webui.bat` (*Windows*) or `webui.sh` (*Linux or MacOS*) -List of available parameters, run `webui --help` for the full & up-to-date list: +For list of available command line options, run `webui --help` for the full & up-to-date list - Server options: - --config CONFIG Use specific server configuration file, default: config.json - --ui-config UI_CONFIG Use specific UI configuration file, default: ui-config.json - --medvram Split model stages and keep only active part in VRAM, default: False - --lowvram Split model components and keep only active part in VRAM, default: False - --ckpt CKPT Path to model checkpoint to load immediately, default: None - --vae VAE Path to VAE checkpoint to load immediately, default: None - --data-dir DATA_DIR Base path where all user data is stored, default: - --models-dir MODELS_DIR Base path where all models are stored, default: models - --allow-code Allow custom script execution, default: False - --share Enable UI accessible through Gradio site, default: False - --insecure Enable extensions tab regardless of other options, default: False - --use-cpu USE_CPU [USE_CPU ...] Force use CPU for specified modules, default: [] - --listen Launch web server using public IP address, default: False - --port PORT Launch web server with given server port, default: 7860 - --freeze Disable editing settings - --auth AUTH Set access authentication like "user:pwd,user:pwd"" - --auth-file AUTH_FILE Set access authentication using file, default: None - --autolaunch Open the UI URL in the system's default browser upon launch - --docs Mount API docs, default: False - --api-only Run in API only mode without starting UI - --api-log Enable logging of all API requests, default: False - --device-id DEVICE_ID Select the default CUDA device to use, default: None - --cors-origins CORS_ORIGINS Allowed CORS origins as comma-separated list, default: None - --cors-regex CORS_REGEX Allowed CORS origins as regular expression, default: None - --tls-keyfile TLS_KEYFILE Enable TLS and specify key file, default: None - --tls-certfile TLS_CERTFILE Enable TLS and specify cert file, default: None - --tls-selfsign Enable TLS with self-signed certificates, default: False - --server-name SERVER_NAME Sets hostname of server, default: None - --no-hashing Disable hashing of checkpoints, default: False - --no-metadata Disable reading of metadata from models, default: False - --disable-queue Disable queues, default: False - --subpath SUBPATH Customize the URL subpath for usage with reverse proxy - --backend {original,diffusers} force model pipeline type - --allowed-paths ALLOWED_PATHS [ALLOWED_PATHS ...] add additional paths to paths allowed for web access - - Setup options: - --reset Reset main repository to latest version, default: False - --upgrade Upgrade main repository to latest version, default: False - --requirements Force re-check of requirements, default: False - --quick Bypass version checks, default: False - --use-directml Use DirectML if no compatible GPU is detected, default: False - --use-openvino Use Intel OpenVINO backend, default: False - --use-ipex Force use Intel OneAPI XPU backend, default: False - --use-cuda Force use nVidia CUDA backend, default: False - --use-rocm Force use AMD ROCm backend, default: False - --use-zluda Force use ZLUDA, AMD GPUs only, default: False - --use-xformers Force use xFormers cross-optimization, default: False - --skip-requirements Skips checking and installing requirements, default: False - --skip-extensions Skips running individual extension installers, default: False - --skip-git Skips running all GIT operations, default: False - --skip-torch Skips running Torch checks, default: False - --skip-all Skips running all checks, default: False - --skip-env Skips setting of env variables during startup, default: False - --experimental Allow unsupported versions of libraries, default: False - --reinstall Force reinstallation of all requirements, default: False - --test Run test only and exit - --version Print version information - --ignore Ignore any errors and attempt to continue - --safe Run in safe mode with no user extensions - --uv Use uv as installer, default: False - - Logging options: - --log LOG Set log file, default: None - --debug Run installer with debug logging, default: False - --profile Run profiler, default: False +> [!TIP] +> All command line options can also be set via env variable +> For example `--debug` is same as `set SD_DEBUG=true` ## Notes > [!TIP] > If you don't want to use built-in `venv` support and prefer to run SD.Next in your own environment such as *Docker* container, *Conda* environment or any other virtual environment, you can skip `venv` create/activate and launch SD.Next directly using `python launch.py` (command line flags noted above still apply). +### Quantization + +**SD.Next** comes with broad quantization support, including support for BitsAndBytes, Optimum.Quanto, TorchAO, NNCF and GGUF +See [Quantization Wiki](https://github.com/vladmandic/automatic/wiki/Quantization) + ### Control **SD.Next** comes with built-in control for all types of text2image, image2image, video2video and batch processing *Control interface*: -![Screenshot-Control](html/screenshot-control.jpg) +![screenshot-control](https://github.com/user-attachments/assets/cdad2722-ae7c-4c9c-94d6-5ea35a4b1356) *Control processors*: -![Screenshot-Process](html/screenshot-processors.jpg) +![screenshot-processors](https://github.com/user-attachments/assets/7bccb82b-366e-4bdb-ae57-cc53fac95d3c) *Masking*: -![Screenshot-Mask](html/screenshot-mask.jpg) +![screenshot-mask](https://github.com/user-attachments/assets/4b057e65-64f0-44ea-93b4-c3b69bc55532) ### Extensions diff --git a/TODO.md b/TODO.md index 5726e67da..2f3f90852 100644 --- a/TODO.md +++ b/TODO.md @@ -6,7 +6,7 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma - async lowvram: - fp8: -- ipadapter-negative: https://github.com/huggingface/diffusers/discussions/7167 +- ipadapter-negative: - include reference styles ### Missing diff --git a/cli/README.md b/cli/README.md index 70de255b1..838db7a50 100644 --- a/cli/README.md +++ b/cli/README.md @@ -1,16 +1,43 @@ # Stable-Diffusion Productivity Scripts -Note: All scripts have built-in `--help` parameter that can be used to get more information +## API Examples -
+### Run Generate -## Main Scripts +- `cli/api-txt2img.py` +- `cli/api-img2img.py` +- `cli/api-control.py` -### Generate +### Monitor + +- `cli/api-progress.py` + +### Generic + +- `cli/api-json.py` + +### Process + +- `cli/api-info.py` +- `cli/api-upscale.py` +- `cli/api-vqa.py` +- `cli/api-preprocess.py` + +### Other + +- `cli/api-faceid.py` +- `cli/api-faces.py` +- `cli/api-mask.py` + +### JavaScript + +- `cli/api-txt2img.js` + +## Generate Text-to-image with all of the possible parameters Supports upsampling, face restoration and grid creation -> python generate.py +> python cli/generate.py By default uses parameters from `generate.json` @@ -20,25 +47,6 @@ Parameters that are not specified will be randomized: - Sampler/Scheduler will be randomly picked from available ones - CFG Scale set to 5-10 -### Train - -Combined pipeline for **embeddings**, **lora**, **lycoris**, **dreambooth** and **hypernetwork** -Optionally runs several image processing steps before training: - -- keep original image -- detect and extract face -- detect and extract body -- detect blur -- detect dynamic range -- attempt to upscale low resolution images -- attempt to restore quality of low quality images -- automatically generate captions using interrogate -- resize image -- square image -- run image segmentation to remove background - -> python train.py -
## Auxiliary Scripts diff --git a/cli/image-watermark.py b/cli/image-watermark.py index 4b3396c17..763b6d714 100755 --- a/cli/image-watermark.py +++ b/cli/image-watermark.py @@ -82,6 +82,7 @@ def watermark(params, file): exif = get_exif(image) + wm = None if params.command == 'read': fn = params.input wm = get_watermark(image, params) diff --git a/cli/load_unet.py b/cli/load-unet.py similarity index 100% rename from cli/load_unet.py rename to cli/load-unet.py diff --git a/cli/model-keys.py b/cli/model-keys.py new file mode 100755 index 000000000..bd4a91551 --- /dev/null +++ b/cli/model-keys.py @@ -0,0 +1,89 @@ +#!/usr/bin/env python +import os +import sys +from rich import print as pprint + + +def has(obj, attr, *args): + import functools + if not isinstance(obj, dict): + return False + def _getattr(obj, attr): + return obj.get(attr, args) if isinstance(obj, dict) else False + return functools.reduce(_getattr, [obj] + attr.split('.')) + + +def remove_entries_after_depth(d, depth, current_depth=0): + try: + if current_depth >= depth: + return None + if isinstance(d, dict): + return {k: remove_entries_after_depth(v, depth, current_depth + 1) for k, v in d.items() if remove_entries_after_depth(v, depth, current_depth + 1) is not None} + except Exception: + pass + return d + + +def list_to_dict(flat_list): + result_dict = {} + try: + for item in flat_list: + keys = item.split('.') + d = result_dict + for key in keys[:-1]: + d = d.setdefault(key, {}) + d[keys[-1]] = None + except Exception: + pass + return result_dict + + +def guess_dct(dct: dict): + # if has(dct, 'model.diffusion_model.input_blocks') and has(dct, 'model.diffusion_model.label_emb'): + # return 'sdxl' + if has(dct, 'model.diffusion_model.input_blocks') and len(list(has(dct, 'model.diffusion_model.input_blocks'))) == 12: + return 'sd15' + if has(dct, 'model.diffusion_model.input_blocks') and len(list(has(dct, 'model.diffusion_model.input_blocks'))) == 9: + return 'sdxl' + if has(dct, 'model.diffusion_model.joint_blocks') and len(list(has(dct, 'model.diffusion_model.joint_blocks'))) == 24: + return 'sd35-medium' + if has(dct, 'model.diffusion_model.joint_blocks') and len(list(has(dct, 'model.diffusion_model.joint_blocks'))) == 38: + return 'sd35-large' + if has(dct, 'model.diffusion_model.double_blocks') and len(list(has(dct, 'model.diffusion_model.double_blocks'))) == 19: + return 'flux-dev' + return None + + +def read_keys(fn): + if not fn.lower().endswith(".safetensors"): + return + from safetensors.torch import safe_open + keys = [] + try: + with safe_open(fn, framework="pt", device="cpu") as f: + keys = f.keys() + except Exception as e: + pprint(e) + dct = list_to_dict(keys) + pprint(f'file: {fn}') + pprint(remove_entries_after_depth(dct, 3)) + pprint(remove_entries_after_depth(dct, 6)) + guess = guess_dct(dct) + pprint(f'guess: {guess}') + return keys + + +def main(): + if len(sys.argv) == 0: + print('metadata:', 'no files specified') + for fn in sys.argv: + if os.path.isfile(fn): + read_keys(fn) + elif os.path.isdir(fn): + for root, _dirs, files in os.walk(fn): + for file in files: + read_keys(os.path.join(root, file)) + +if __name__ == '__main__': + sys.argv.pop(0) + main() diff --git a/extensions-builtin/Lora/networks.py b/extensions-builtin/Lora/networks.py index c0a8555e1..160487e88 100644 --- a/extensions-builtin/Lora/networks.py +++ b/extensions-builtin/Lora/networks.py @@ -127,6 +127,8 @@ def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_ def load_network(name, network_on_disk) -> network.Network: + if not shared.sd_loaded: + return None t0 = time.time() cached = lora_cache.get(name, None) if debug: @@ -531,18 +533,14 @@ def network_MultiheadAttention_load_state_dict(self, *args, **kwargs): def list_available_networks(): + t0 = time.time() available_networks.clear() available_network_aliases.clear() forbidden_network_aliases.clear() available_network_hash_lookup.clear() forbidden_network_aliases.update({"none": 1, "Addams": 1}) - directories = [] - if os.path.exists(shared.cmd_opts.lora_dir): - directories.append(shared.cmd_opts.lora_dir) - else: + if not os.path.exists(shared.cmd_opts.lora_dir): shared.log.warning(f'LoRA directory not found: path="{shared.cmd_opts.lora_dir}"') - if os.path.exists(shared.cmd_opts.lyco_dir) and shared.cmd_opts.lyco_dir != shared.cmd_opts.lora_dir: - directories.append(shared.cmd_opts.lyco_dir) def add_network(filename): if not os.path.isfile(filename): @@ -563,11 +561,12 @@ def list_available_networks(): except OSError as e: # should catch FileNotFoundError and PermissionError etc. shared.log.error(f'LoRA: filename="{filename}" {e}') - candidates = list(files_cache.list_files(*directories, ext_filter=[".pt", ".ckpt", ".safetensors"])) + candidates = list(files_cache.list_files(shared.cmd_opts.lora_dir, ext_filter=[".pt", ".ckpt", ".safetensors"])) with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor: for fn in candidates: executor.submit(add_network, fn) - shared.log.info(f'Available LoRAs: items={len(available_networks)} folders={len(forbidden_network_aliases)}') + t1 = time.time() + shared.log.info(f'Available LoRAs: path="{shared.cmd_opts.lora_dir}" items={len(available_networks)} folders={len(forbidden_network_aliases)} time={t1 - t0:.2f}') def infotext_pasted(infotext, params): # pylint: disable=W0613 diff --git a/extensions-builtin/Lora/ui_extra_networks_lora.py b/extensions-builtin/Lora/ui_extra_networks_lora.py index 170c3c7d3..4220b8e02 100644 --- a/extensions-builtin/Lora/ui_extra_networks_lora.py +++ b/extensions-builtin/Lora/ui_extra_networks_lora.py @@ -5,7 +5,7 @@ import networks from modules import shared, ui_extra_networks -debug = os.environ.get('SD_LOAD_DEBUG', None) is not None +debug = os.environ.get('SD_LORA_DEBUG', None) is not None class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): @@ -16,34 +16,18 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): def refresh(self): networks.list_available_networks() - def create_item(self, name): - l = networks.available_networks.get(name) - if l is None: - shared.log.warning(f'Networks: type=lora registered={len(list(networks.available_networks))} file="{name}" not registered') - return None + def get_tags(self, l, info): + tags = {} try: - # path, _ext = os.path.splitext(l.filename) - name = os.path.splitext(os.path.relpath(l.filename, shared.cmd_opts.lora_dir))[0] - item = { - "type": 'Lora', - "name": name, - "filename": l.filename, - "hash": l.shorthash, - "prompt": json.dumps(f" "), - "metadata": json.dumps(l.metadata, indent=4) if l.metadata else None, - "mtime": os.path.getmtime(l.filename), - "size": os.path.getsize(l.filename), - "version": l.sd_version, - } - info = self.find_info(l.filename) - - tags = {} if l.metadata is not None: modelspec_tags = l.metadata.get('modelspec.tags', {}) possible_tags = l.metadata.get('ss_tag_frequency', {}) # tags from model metedata - possible_tags.update(modelspec_tags) if isinstance(possible_tags, str): possible_tags = {} + if isinstance(modelspec_tags, str): + modelspec_tags = {} + if len(list(modelspec_tags)) > 0: + possible_tags.update(modelspec_tags) for k, v in possible_tags.items(): words = k.split('_', 1) if '_' in k else [v, k] words = [str(w).replace('.json', '') for w in words] @@ -80,20 +64,41 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): tag = tag.strip().lower() if tag not in tags: tags[tag] = 0 + except Exception: + pass + bad_chars = [';', ':', '<', ">", "*", '?', '\'', '\"', '(', ')', '[', ']', '{', '}', '\\', '/'] + clean_tags = {} + for k, v in tags.items(): + tag = ''.join(i for i in k if i not in bad_chars).strip() + clean_tags[tag] = v - bad_chars = [';', ':', '<', ">", "*", '?', '\'', '\"', '(', ')', '[', ']', '{', '}', '\\', '/'] - clean_tags = {} - for k, v in tags.items(): - tag = ''.join(i for i in k if i not in bad_chars).strip() - clean_tags[tag] = v - - clean_tags.pop('img', None) - clean_tags.pop('dataset', None) + clean_tags.pop('img', None) + clean_tags.pop('dataset', None) + return clean_tags + def create_item(self, name): + l = networks.available_networks.get(name) + if l is None: + shared.log.warning(f'Networks: type=lora registered={len(list(networks.available_networks))} file="{name}" not registered') + return None + try: + # path, _ext = os.path.splitext(l.filename) + name = os.path.splitext(os.path.relpath(l.filename, shared.cmd_opts.lora_dir))[0] + item = { + "type": 'Lora', + "name": name, + "filename": l.filename, + "hash": l.shorthash, + "prompt": json.dumps(f" "), + "metadata": json.dumps(l.metadata, indent=4) if l.metadata else None, + "mtime": os.path.getmtime(l.filename), + "size": os.path.getsize(l.filename), + "version": l.sd_version, + } + info = self.find_info(l.filename) item["info"] = info item["description"] = self.find_description(l.filename, info) # use existing info instead of double-read - item["tags"] = clean_tags - + item["tags"] = self.get_tags(l, info) return item except Exception as e: shared.log.error(f'Networks: type=lora file="{name}" {e}') diff --git a/extensions-builtin/sdnext-modernui b/extensions-builtin/sdnext-modernui index 906bd2a98..71bdbbd9c 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit 906bd2a98ba0736c235925f2beaea787a050aeed +Subproject commit 71bdbbd9c0a55ccea38cbf6fb01483323ac93676 diff --git a/html/card-no-preview.png b/html/card-no-preview.png index 352b0a9df..6b89aad1a 100644 Binary files a/html/card-no-preview.png and b/html/card-no-preview.png differ diff --git a/html/logo-bg-1.jpg b/html/logo-bg-1.jpg index 4acd348d1..fa3d77cfc 100644 Binary files a/html/logo-bg-1.jpg and b/html/logo-bg-1.jpg differ diff --git a/html/logo-bg-10.jpg b/html/logo-bg-10.jpg deleted file mode 100644 index d2289762c..000000000 Binary files a/html/logo-bg-10.jpg and /dev/null differ diff --git a/html/logo-bg-2.jpg b/html/logo-bg-2.jpg index 313600477..ac3aca0e7 100644 Binary files a/html/logo-bg-2.jpg and b/html/logo-bg-2.jpg differ diff --git a/html/logo-bg-3.jpg b/html/logo-bg-3.jpg index ff412818e..1c0323dee 100644 Binary files a/html/logo-bg-3.jpg and b/html/logo-bg-3.jpg differ diff --git a/html/logo-bg-4.jpg b/html/logo-bg-4.jpg index 095a56ad0..7896c35d8 100644 Binary files a/html/logo-bg-4.jpg and b/html/logo-bg-4.jpg differ diff --git a/html/logo-bg-5.jpg b/html/logo-bg-5.jpg index 3dc774b01..8868ae93e 100644 Binary files a/html/logo-bg-5.jpg and b/html/logo-bg-5.jpg differ diff --git a/html/logo-bg-6.jpg b/html/logo-bg-6.jpg index 5e001eb2b..102642be3 100644 Binary files a/html/logo-bg-6.jpg and b/html/logo-bg-6.jpg differ diff --git a/html/logo-bg-7.jpg b/html/logo-bg-7.jpg index cc8e0cc65..2bccf3230 100644 Binary files a/html/logo-bg-7.jpg and b/html/logo-bg-7.jpg differ diff --git a/html/logo-bg-8.jpg b/html/logo-bg-8.jpg index 70681c2fa..596dca98f 100644 Binary files a/html/logo-bg-8.jpg and b/html/logo-bg-8.jpg differ diff --git a/html/logo-bg-9.jpg b/html/logo-bg-9.jpg index 203f78b9f..3a2c80f2b 100644 Binary files a/html/logo-bg-9.jpg and b/html/logo-bg-9.jpg differ diff --git a/html/previews.json b/html/previews.json new file mode 100644 index 000000000..3e4bc3d41 --- /dev/null +++ b/html/previews.json @@ -0,0 +1,11 @@ +{ + "stabilityai--stable-diffusion-3-medium-diffusers": "models/Reference/stabilityai--stable-diffusion-3.jpg", + "stabilityai--stable-diffusion-3.5-medium": "models/Reference/stabilityai--stable-diffusion-3_5.jpg", + "stabilityai--stable-diffusion-3.5-large": "models/Reference/stabilityai--stable-diffusion-3_5.jpg", + "Disty0--FLUX.1-dev-qint8": "models/Reference/black-forest-labs--FLUX.1-dev.jpg", + "Disty0--FLUX.1-dev-qint4": "models/Reference/black-forest-labs--FLUX.1-dev.jpg", + "sayakpaul--flux.1-dev-nf4": "models/Reference/black-forest-labs--FLUX.1-dev.jpg", + "THUDM--CogVideoX-2b": "models/Reference/THUDM--CogView3-Plus-3B.jpg", + "THUDM--CogVideoX-5b": "models/Reference/THUDM--CogView3-Plus-3B.jpg", + "THUDM--CogVideoX-5b-I2V": "models/Reference/THUDM--CogView3-Plus-3B.jpg" +} diff --git a/html/reference.json b/html/reference.json index 8d26433e7..4a549586f 100644 --- a/html/reference.json +++ b/html/reference.json @@ -119,11 +119,19 @@ "preview": "stabilityai--stable-diffusion-3.jpg", "extras": "sampler: Default, cfg_scale: 7.0" }, + "StabilityAI Stable Diffusion 3.5 Medium": { + "path": "stabilityai/stable-diffusion-3.5-medium", + "skip": true, + "variant": "fp16", + "desc": "Stable Diffusion 3.5 Medium is a Multimodal Diffusion Transformer with improvements (MMDiT-X) text-to-image model that features improved performance in image quality, typography, complex prompt understanding, and resource-efficiency.", + "preview": "stabilityai--stable-diffusion-3_5.jpg", + "extras": "sampler: Default, cfg_scale: 7.0" + }, "StabilityAI Stable Diffusion 3.5 Large": { "path": "stabilityai/stable-diffusion-3.5-large", "skip": true, "variant": "fp16", - "desc": "Stable Diffusion 3 Medium is a Multimodal Diffusion Transformer (MMDiT) text-to-image model that features greatly improved performance in image quality, typography, complex prompt understanding, and resource-efficiency", + "desc": "Stable Diffusion 3.5 Large is a Multimodal Diffusion Transformer (MMDiT) text-to-image model that features improved performance in image quality, typography, complex prompt understanding, and resource-efficiency.", "preview": "stabilityai--stable-diffusion-3_5.jpg", "extras": "sampler: Default, cfg_scale: 7.0" }, @@ -131,7 +139,7 @@ "path": "stabilityai/stable-diffusion-3.5-large-turbo", "skip": true, "variant": "fp16", - "desc": "Stable Diffusion 3 Medium is a Multimodal Diffusion Transformer (MMDiT) text-to-image model that features greatly improved performance in image quality, typography, complex prompt understanding, and resource-efficiency", + "desc": "Stable Diffusion 3.5 Large Turbo is a Multimodal Diffusion Transformer (MMDiT) text-to-image model with Adversarial Diffusion Distillation (ADD) that features improved performance in image quality, typography, complex prompt understanding, and resource-efficiency, with a focus on fewer inference steps.", "preview": "stabilityai--stable-diffusion-3_5.jpg", "extras": "sampler: Default, cfg_scale: 7.0" }, diff --git a/html/screenshot-color.jpg b/html/screenshot-color.jpg deleted file mode 100644 index 2cd070a3a..000000000 Binary files a/html/screenshot-color.jpg and /dev/null differ diff --git a/html/screenshot-control.jpg b/html/screenshot-control.jpg deleted file mode 100644 index ffc71cdbe..000000000 Binary files a/html/screenshot-control.jpg and /dev/null differ diff --git a/html/screenshot-corrections.jpg b/html/screenshot-corrections.jpg deleted file mode 100644 index e202fa746..000000000 Binary files a/html/screenshot-corrections.jpg and /dev/null differ diff --git a/html/screenshot-differential.jpg b/html/screenshot-differential.jpg deleted file mode 100644 index c0bb1e78d..000000000 Binary files a/html/screenshot-differential.jpg and /dev/null differ diff --git a/html/screenshot-instantid.jpg b/html/screenshot-instantid.jpg deleted file mode 100644 index 5f780a3c0..000000000 Binary files a/html/screenshot-instantid.jpg and /dev/null differ diff --git a/html/screenshot-ipadapter-mask.jpg b/html/screenshot-ipadapter-mask.jpg deleted file mode 100644 index 3dd36a2af..000000000 Binary files a/html/screenshot-ipadapter-mask.jpg and /dev/null differ diff --git a/html/screenshot-ipadapter.jpg b/html/screenshot-ipadapter.jpg deleted file mode 100644 index b1d5c80b6..000000000 Binary files a/html/screenshot-ipadapter.jpg and /dev/null differ diff --git a/html/screenshot-ipadapter02.jpg b/html/screenshot-ipadapter02.jpg deleted file mode 100644 index d0e59a2ec..000000000 Binary files a/html/screenshot-ipadapter02.jpg and /dev/null differ diff --git a/html/screenshot-ledit.jpg b/html/screenshot-ledit.jpg deleted file mode 100644 index f43189d81..000000000 Binary files a/html/screenshot-ledit.jpg and /dev/null differ diff --git a/html/screenshot-mask.jpg b/html/screenshot-mask.jpg deleted file mode 100644 index c286803f2..000000000 Binary files a/html/screenshot-mask.jpg and /dev/null differ diff --git a/html/screenshot-modernui-control.jpg b/html/screenshot-modernui-control.jpg deleted file mode 100644 index 20a22a88a..000000000 Binary files a/html/screenshot-modernui-control.jpg and /dev/null differ diff --git a/html/screenshot-modernui-f1.jpg b/html/screenshot-modernui-f1.jpg deleted file mode 100644 index 3e7b27ce1..000000000 Binary files a/html/screenshot-modernui-f1.jpg and /dev/null differ diff --git a/html/screenshot-modernui-img2img.jpg b/html/screenshot-modernui-img2img.jpg deleted file mode 100644 index 22afd166d..000000000 Binary files a/html/screenshot-modernui-img2img.jpg and /dev/null differ diff --git a/html/screenshot-modernui-sd3.jpg b/html/screenshot-modernui-sd3.jpg deleted file mode 100644 index 81d0eb459..000000000 Binary files a/html/screenshot-modernui-sd3.jpg and /dev/null differ diff --git a/html/screenshot-modernui.jpg b/html/screenshot-modernui.jpg deleted file mode 100644 index dafadc28f..000000000 Binary files a/html/screenshot-modernui.jpg and /dev/null differ diff --git a/html/screenshot-outpaint.jpg b/html/screenshot-outpaint.jpg deleted file mode 100644 index 485c67c52..000000000 Binary files a/html/screenshot-outpaint.jpg and /dev/null differ diff --git a/html/screenshot-processors.jpg b/html/screenshot-processors.jpg deleted file mode 100644 index 8d3207fa2..000000000 Binary files a/html/screenshot-processors.jpg and /dev/null differ diff --git a/html/screenshot-regional.jpg b/html/screenshot-regional.jpg deleted file mode 100644 index 335297b37..000000000 Binary files a/html/screenshot-regional.jpg and /dev/null differ diff --git a/html/screenshot-text2image.jpg b/html/screenshot-text2image.jpg deleted file mode 100644 index f5b48dc3a..000000000 Binary files a/html/screenshot-text2image.jpg and /dev/null differ diff --git a/installer.py b/installer.py index ef6ce44ba..b19eae87e 100644 --- a/installer.py +++ b/installer.py @@ -227,9 +227,9 @@ def installed(package, friendly: str = None, reload = False, quiet = False): exact = pkg_version == p[1] if not exact and not quiet: if args.experimental: - log.warning(f"Package: {p[0]} {pkg_version} required {p[1]} allowing experimental") + log.warning(f"Package: {p[0]} installed={pkg_version} required={p[1]} allowing experimental") else: - log.warning(f"Package: {p[0]} {pkg_version} required {p[1]} version mismatch") + log.warning(f"Package: {p[0]} installed={pkg_version} required={p[1]} version mismatch") ok = ok and (exact or args.experimental) else: if not quiet: @@ -254,11 +254,12 @@ def uninstall(package, quiet = False): @lru_cache() def pip(arg: str, ignore: bool = False, quiet: bool = False, uv = True): originalArg = arg - uv = uv and args.uv - pipCmd = "uv pip" if uv else "pip" arg = arg.replace('>=', '==') + package = arg.replace("install", "").replace("--upgrade", "").replace("--no-deps", "").replace("--force", "").replace(" ", " ").strip() + uv = uv and args.uv and not package.startswith('git+') + pipCmd = "uv pip" if uv else "pip" if not quiet and '-r ' not in arg: - log.info(f'Install: package="{arg.replace("install", "").replace("--upgrade", "").replace("--no-deps", "").replace("--force", "").replace(" ", " ").strip()}" mode={"uv" if uv else "pip"}') + log.info(f'Install: package="{package}" mode={"uv" if uv else "pip"}') env_args = os.environ.get("PIP_EXTRA_ARGS", "") all_args = f'{pip_log}{arg} {env_args}'.strip() if not quiet: @@ -376,10 +377,11 @@ def update(folder, keep_branch = False, rebase = True): else: res = git(f'pull origin {b} {arg}', folder) debug(f'Install update: folder={folder} branch={b} args={arg} {res}') - commit = extensions_commit.get(os.path.basename(folder), None) - if commit is not None: - res = git(f'checkout {commit}', folder) - debug(f'Install update: folder={folder} branch={b} args={arg} commit={commit} {res}') + if not args.experimental: + commit = extensions_commit.get(os.path.basename(folder), None) + if commit is not None: + res = git(f'checkout {commit}', folder) + debug(f'Install update: folder={folder} branch={b} args={arg} commit={commit} {res}') return res @@ -454,7 +456,7 @@ def check_python(supported_minors=[9, 10, 11, 12], reason=None): # check diffusers version def check_diffusers(): - sha = '435f6b7e47c031f98b8374b1689e1abeb17bfdb6' + sha = '0d1d267b12e47b40b0e8f265339c76e0f45f8c49' pkg = pkg_resources.working_set.by_key.get('diffusers', None) minor = int(pkg.version.split('.')[1] if pkg is not None else 0) cur = opts.get('diffusers_version', '') if minor > 0 else '' @@ -489,7 +491,7 @@ def install_cuda(): log.info('CUDA: nVidia toolkit detected') install('onnxruntime-gpu', 'onnxruntime-gpu', ignore=True, quiet=True) # return os.environ.get('TORCH_COMMAND', 'torch torchvision --index-url https://download.pytorch.org/whl/cu124') - return os.environ.get('TORCH_COMMAND', 'torch==2.4.1+cu124 torchvision==0.19.1+cu124 --index-url https://download.pytorch.org/whl/cu124') + return os.environ.get('TORCH_COMMAND', 'torch==2.5.1+cu124 torchvision==0.20.1+cu124 --index-url https://download.pytorch.org/whl/cu124') def install_rocm_zluda(): @@ -546,11 +548,11 @@ def install_rocm_zluda(): os.environ['HIP_VISIBLE_DEVICES'] = args.device_id del args.device_id - log.warning("ZLUDA support: experimental") error = None from modules import zluda_installer + zluda_installer.set_default_agent(device) try: - if args.reinstall_zluda: + if args.reinstall: zluda_installer.uninstall() zluda_path = zluda_installer.get_path() zluda_installer.install(zluda_path) @@ -570,8 +572,10 @@ def install_rocm_zluda(): log.info('Using CPU-only torch') torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision') else: - if rocm.version is None or float(rocm.version) >= 6.1: # assume the latest if version check fails - #torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision --index-url https://download.pytorch.org/whl/rocm6.1') + if rocm.version is None or float(rocm.version) > 6.1: # assume the latest if version check fails + # torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.5.1+rocm6.2 torchvision==0.20.1+rocm6.2 --index-url https://download.pytorch.org/whl/rocm6.2') + torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.4.1+rocm6.1 torchvision==0.19.1+rocm6.1 --index-url https://download.pytorch.org/whl/rocm6.1') + elif rocm.version == "6.1": # lock to 2.4.1, older rocm (5.7) uses torch 2.3 torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.4.1+rocm6.1 torchvision==0.19.1+rocm6.1 --index-url https://download.pytorch.org/whl/rocm6.1') elif rocm.version == "6.0": # lock to 2.4.1, older rocm (5.7) uses torch 2.3 torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.4.1+rocm6.0 torchvision==0.19.1+rocm6.0 --index-url https://download.pytorch.org/whl/rocm6.0') @@ -591,11 +595,11 @@ def install_rocm_zluda(): install(ort_package, 'onnxruntime-training') if installed("torch") and device is not None: - if 'Flash attention' in opts.get('sdp_options'): + if 'Flash attention' in opts.get('sdp_options', ''): if not installed('flash-attn'): install(rocm.get_flash_attention_command(device), reinstall=True) - elif not args.experimental: - uninstall('flash-attn') + #elif not args.experimental: + # uninstall('flash-attn') if device is not None and rocm.version != "6.2" and rocm.version == rocm.version_torch and rocm.get_blaslt_enabled(): log.debug(f'ROCm hipBLASLt: arch={device.name} available={device.blaslt_supported}') @@ -730,7 +734,7 @@ def check_torch(): else: if args.use_zluda: log.warning("ZLUDA failed to initialize: no HIP SDK found") - log.info('Using CPU-only Torch') + log.warning('Torch: CPU-only version installed') torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision') if 'torch' in torch_command and not args.version: install(torch_command, 'torch torchvision', quiet=True) @@ -817,6 +821,7 @@ def install_packages(): log.info('Verifying packages') clip_package = os.environ.get('CLIP_PACKAGE', "git+https://github.com/openai/CLIP.git") install(clip_package, 'clip', quiet=True) + install('open-clip-torch', no_deps=True, quiet=True) # tensorflow_package = os.environ.get('TENSORFLOW_PACKAGE', 'tensorflow==2.13.0') # tensorflow_package = os.environ.get('TENSORFLOW_PACKAGE', None) # if tensorflow_package is not None: @@ -1134,6 +1139,28 @@ def check_ui(ver): os.chdir(cwd) +def check_venv(): + import site + pkg_path = [os.path.relpath(p) for p in site.getsitepackages() if os.path.exists(p)] + log.debug(f'Packages: venv={os.path.relpath(sys.prefix)} site={pkg_path}') + for p in pkg_path: + invalid = [] + for f in os.listdir(p): + if f.startswith('~'): + invalid.append(f) + if len(invalid) > 0: + log.warning(f'Packages: site="{p}" invalid={invalid} removing') + for f in invalid: + fn = os.path.join(p, f) + try: + if os.path.isdir(fn): + shutil.rmtree(fn) + elif os.path.isfile(fn): + os.unlink(fn) + except Exception as e: + log.error(f'Packages: site={p} invalid={f} error={e}') + + # check version of the main repo and optionally upgrade it def check_version(offline=False, reset=True): # pylint: disable=unused-argument if args.skip_all: @@ -1229,39 +1256,44 @@ def check_timestamp(): def add_args(parser): - group = parser.add_argument_group('Setup options') - group.add_argument('--reset', default = os.environ.get("SD_RESET",False), action='store_true', help = "Reset main repository to latest version, default: %(default)s") - group.add_argument('--upgrade', '--update', default = os.environ.get("SD_UPGRADE",False), action='store_true', help = "Upgrade main repository to latest version, default: %(default)s") - group.add_argument('--requirements', default = os.environ.get("SD_REQUIREMENTS",False), action='store_true', help = "Force re-check of requirements, default: %(default)s") - group.add_argument('--quick', default = os.environ.get("SD_QUICK",False), action='store_true', help = "Bypass version checks, default: %(default)s") - group.add_argument('--use-directml', default = os.environ.get("SD_USEDIRECTML",False), action='store_true', help = "Use DirectML if no compatible GPU is detected, default: %(default)s") - group.add_argument("--use-openvino", default = os.environ.get("SD_USEOPENVINO",False), action='store_true', help="Use Intel OpenVINO backend, default: %(default)s") - group.add_argument("--use-ipex", default = os.environ.get("SD_USEIPEX",False), action='store_true', help="Force use Intel OneAPI XPU backend, default: %(default)s") - group.add_argument("--use-cuda", default = os.environ.get("SD_USECUDA",False), action='store_true', help="Force use nVidia CUDA backend, default: %(default)s") - group.add_argument("--use-rocm", default = os.environ.get("SD_USEROCM",False), action='store_true', help="Force use AMD ROCm backend, default: %(default)s") - group.add_argument('--use-zluda', default=os.environ.get("SD_USEZLUDA", False), action='store_true', help = "Force use ZLUDA, AMD GPUs only, default: %(default)s") - group.add_argument("--use-xformers", default = os.environ.get("SD_USEXFORMERS",False), action='store_true', help="Force use xFormers cross-optimization, default: %(default)s") - group.add_argument('--skip-requirements', default = os.environ.get("SD_SKIPREQUIREMENTS",False), action='store_true', help = "Skips checking and installing requirements, default: %(default)s") - group.add_argument('--skip-extensions', default = os.environ.get("SD_SKIPEXTENSION",False), action='store_true', help = "Skips running individual extension installers, default: %(default)s") - group.add_argument('--skip-git', default = os.environ.get("SD_SKIPGIT",False), action='store_true', help = "Skips running all GIT operations, default: %(default)s") - group.add_argument('--skip-torch', default = os.environ.get("SD_SKIPTORCH",False), action='store_true', help = "Skips running Torch checks, default: %(default)s") - group.add_argument('--skip-all', default = os.environ.get("SD_SKIPALL",False), action='store_true', help = "Skips running all checks, default: %(default)s") - group.add_argument('--skip-env', default = os.environ.get("SD_SKIPENV",False), action='store_true', help = "Skips setting of env variables during startup, default: %(default)s") - group.add_argument('--experimental', default = os.environ.get("SD_EXPERIMENTAL",False), action='store_true', help = "Allow unsupported versions of libraries, default: %(default)s") - group.add_argument('--reinstall', default = os.environ.get("SD_REINSTALL",False), action='store_true', help = "Force reinstallation of all requirements, default: %(default)s") - group.add_argument('--reinstall-zluda', default = os.environ.get("SD_REINSTALL_ZLUDA",False), action='store_true', help = "Force reinstallation of ZLUDA, default: %(default)s") - group.add_argument('--test', default = os.environ.get("SD_TEST",False), action='store_true', help = "Run test only and exit") - group.add_argument('--version', default = False, action='store_true', help = "Print version information") - group.add_argument('--ignore', default = os.environ.get("SD_IGNORE",False), action='store_true', help = "Ignore any errors and attempt to continue") - group.add_argument('--safe', default = os.environ.get("SD_SAFE",False), action='store_true', help = "Run in safe mode with no user extensions") - group.add_argument('--uv', default = os.environ.get("SD_UV",False), action='store_true', help = "Use uv instead of pip to install the packages") + group_setup = parser.add_argument_group('Setup') + group_setup.add_argument('--reset', default = os.environ.get("SD_RESET",False), action='store_true', help = "Reset main repository to latest version, default: %(default)s") + group_setup.add_argument('--upgrade', '--update', default = os.environ.get("SD_UPGRADE",False), action='store_true', help = "Upgrade main repository to latest version, default: %(default)s") + group_setup.add_argument('--requirements', default = os.environ.get("SD_REQUIREMENTS",False), action='store_true', help = "Force re-check of requirements, default: %(default)s") + group_setup.add_argument('--reinstall', default = os.environ.get("SD_REINSTALL",False), action='store_true', help = "Force reinstallation of all requirements, default: %(default)s") + group_setup.add_argument('--uv', default = os.environ.get("SD_UV",False), action='store_true', help = "Use uv instead of pip to install the packages") - group = parser.add_argument_group('Logging options') - group.add_argument("--log", type=str, default=os.environ.get("SD_LOG", None), help="Set log file, default: %(default)s") - group.add_argument('--debug', default = os.environ.get("SD_DEBUG",False), action='store_true', help = "Run installer with debug logging, default: %(default)s") - group.add_argument("--profile", default=os.environ.get("SD_PROFILE", False), action='store_true', help="Run profiler, default: %(default)s") - group.add_argument('--docs', default=os.environ.get("SD_DOCS", False), action='store_true', help = "Mount API docs, default: %(default)s") - group.add_argument("--api-log", default=os.environ.get("SD_APILOG", False), action='store_true', help="Enable logging of all API requests, default: %(default)s") + group_startup = parser.add_argument_group('Startup') + group_startup.add_argument('--quick', default = os.environ.get("SD_QUICK",False), action='store_true', help = "Bypass version checks, default: %(default)s") + group_startup.add_argument('--skip-requirements', default = os.environ.get("SD_SKIPREQUIREMENTS",False), action='store_true', help = "Skips checking and installing requirements, default: %(default)s") + group_startup.add_argument('--skip-extensions', default = os.environ.get("SD_SKIPEXTENSION",False), action='store_true', help = "Skips running individual extension installers, default: %(default)s") + group_startup.add_argument('--skip-git', default = os.environ.get("SD_SKIPGIT",False), action='store_true', help = "Skips running all GIT operations, default: %(default)s") + group_startup.add_argument('--skip-torch', default = os.environ.get("SD_SKIPTORCH",False), action='store_true', help = "Skips running Torch checks, default: %(default)s") + group_startup.add_argument('--skip-all', default = os.environ.get("SD_SKIPALL",False), action='store_true', help = "Skips running all checks, default: %(default)s") + group_startup.add_argument('--skip-env', default = os.environ.get("SD_SKIPENV",False), action='store_true', help = "Skips setting of env variables during startup, default: %(default)s") + + group_compute = parser.add_argument_group('Compute Engine') + group_compute.add_argument('--use-directml', default = os.environ.get("SD_USEDIRECTML",False), action='store_true', help = "Use DirectML if no compatible GPU is detected, default: %(default)s") + group_compute.add_argument("--use-openvino", default = os.environ.get("SD_USEOPENVINO",False), action='store_true', help="Use Intel OpenVINO backend, default: %(default)s") + group_compute.add_argument("--use-ipex", default = os.environ.get("SD_USEIPEX",False), action='store_true', help="Force use Intel OneAPI XPU backend, default: %(default)s") + group_compute.add_argument("--use-cuda", default = os.environ.get("SD_USECUDA",False), action='store_true', help="Force use nVidia CUDA backend, default: %(default)s") + group_compute.add_argument("--use-rocm", default = os.environ.get("SD_USEROCM",False), action='store_true', help="Force use AMD ROCm backend, default: %(default)s") + group_compute.add_argument('--use-zluda', default=os.environ.get("SD_USEZLUDA", False), action='store_true', help = "Force use ZLUDA, AMD GPUs only, default: %(default)s") + group_compute.add_argument("--use-xformers", default = os.environ.get("SD_USEXFORMERS",False), action='store_true', help="Force use xFormers cross-optimization, default: %(default)s") + + group_diag = parser.add_argument_group('Diagnostics') + group_diag.add_argument('--safe', default = os.environ.get("SD_SAFE",False), action='store_true', help = "Run in safe mode with no user extensions") + group_diag.add_argument('--experimental', default = os.environ.get("SD_EXPERIMENTAL",False), action='store_true', help = "Allow unsupported versions of libraries, default: %(default)s") + group_diag.add_argument('--test', default = os.environ.get("SD_TEST",False), action='store_true', help = "Run test only and exit") + group_diag.add_argument('--version', default = False, action='store_true', help = "Print version information") + group_diag.add_argument('--ignore', default = os.environ.get("SD_IGNORE",False), action='store_true', help = "Ignore any errors and attempt to continue") + + group_log = parser.add_argument_group('Logging') + group_log.add_argument("--log", type=str, default=os.environ.get("SD_LOG", None), help="Set log file, default: %(default)s") + group_log.add_argument('--debug', default = os.environ.get("SD_DEBUG",False), action='store_true', help = "Run installer with debug logging, default: %(default)s") + group_log.add_argument("--profile", default=os.environ.get("SD_PROFILE", False), action='store_true', help="Run profiler, default: %(default)s") + group_log.add_argument('--docs', default=os.environ.get("SD_DOCS", False), action='store_true', help = "Mount API docs, default: %(default)s") + group_log.add_argument("--api-log", default=os.environ.get("SD_APILOG", False), action='store_true', help="Enable logging of all API requests, default: %(default)s") def parse_args(parser): diff --git a/javascript/changelog.js b/javascript/changelog.js new file mode 100644 index 000000000..80c9956b8 --- /dev/null +++ b/javascript/changelog.js @@ -0,0 +1,77 @@ +let changelogElements = []; + +const getAllChildren = (el) => { + const elements = []; + for (let i = 0; i < el.children.length; i++) { + elements.push(el.children[i]); + if (el.children[i].children.length) elements.push(...getAllChildren(el.children[i])); + } + return elements; +}; + +function getText(el) { + let text = ''; + el.childNodes.forEach((node) => { + if (node.nodeType === Node.TEXT_NODE) text += node.nodeValue; + }); + return text.trim(); +} + +let currentElement = -1; + +function changelogNavigate(found) { + const result = gradioApp().getElementById('changelog_result'); + result.innerHTML = ''; + const text = document.createElement('p'); + + const onPrev = () => { + if (currentElement > 0) { + currentElement--; + found[currentElement].scrollIntoView(); + text.innerHTML = `   search item ${currentElement + 1} of ${found.length}`; + } + }; + const onNext = () => { + if (currentElement < found.length - 1) { + currentElement++; + found[currentElement].scrollIntoView(); + text.innerHTML = `   search item ${currentElement + 1} of ${found.length}`; + } + }; + + const prev = document.createElement('p'); + prev.innerHTML = ' ⇦ '; + prev.className = 'changelog_arrow'; + prev.onclick = onPrev; + prev.title = 'Search previous'; + result.appendChild(prev); + + const next = document.createElement('p'); + next.innerHTML = ' ⇨ '; + next.className = 'changelog_arrow'; + next.title = 'Search next'; + next.onclick = onNext; + result.appendChild(next); + + text.innerHTML = `   found ${found.length} items`; + result.appendChild(text); +} + +async function initChangelog() { + const search = gradioApp().querySelector('#changelog_search > label> textarea'); + const md = gradioApp().getElementById('changelog_markdown'); + const searchChangelog = async (e) => { + if (changelogElements.length < 100) changelogElements = getAllChildren(md); + const found = []; + for (const el of changelogElements) { + if (search.value.length > 1 && getText(el).toLowerCase().includes(search.value.toLowerCase())) { + el.classList.add('changelog_highlight'); + found.push(el); + } else { + el.classList.remove('changelog_highlight'); + } + } + changelogNavigate(found); + }; + search.addEventListener('keyup', searchChangelog); +} diff --git a/javascript/extraNetworks.js b/javascript/extraNetworks.js index 032ee4999..9a33baa86 100644 --- a/javascript/extraNetworks.js +++ b/javascript/extraNetworks.js @@ -375,6 +375,7 @@ function setupExtraNetworksForTab(tabname) { if (!tabs) return; // buttons + const btnShow = gradioApp().getElementById(`${tabname}_extra_networks_btn`); const btnRefresh = gradioApp().getElementById(`${tabname}_extra_refresh`); const btnScan = gradioApp().getElementById(`${tabname}_extra_scan`); const btnSave = gradioApp().getElementById(`${tabname}_extra_save`); @@ -448,50 +449,69 @@ function setupExtraNetworksForTab(tabname) { // en style if (!en) return; + let lastView; + let heightInitialized = false; const intersectionObserver = new IntersectionObserver((entries) => { - for (const el of Array.from(gradioApp().querySelectorAll('.extra-networks-page'))) { - el.style.height = `${window.opts.extra_networks_height}vh`; - el.parentElement.style.width = '-webkit-fill-available'; + if (!heightInitialized) { + heightInitialized = true; + let h = 0; + const target = window.opts.extra_networks_card_cover === 'sidebar' ? 0 : window.opts.extra_networks_height; + if (window.opts.theme_type === 'Standard') h = target > 0 ? target : 55; + else h = target > 0 ? target : 87; + for (const el of Array.from(gradioApp().getElementById(`${tabname}_extra_tabs`).querySelectorAll('.extra-networks-page'))) { + if (h > 0) el.style.height = `${h}vh`; + el.parentElement.style.width = '-webkit-fill-available'; + } } - if (entries[0].intersectionRatio > 0) { - refreshENpage(); - // sortExtraNetworks('fixed'); - if (window.opts.extra_networks_card_cover === 'cover') { - en.style.transition = ''; - en.style.zIndex = 100; - en.style.top = '13em'; - en.style.position = 'absolute'; - en.style.right = 'unset'; - en.style.width = 'unset'; - en.style.height = 'unset'; - gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = 'unset'; - } else if (window.opts.extra_networks_card_cover === 'sidebar') { - en.style.zIndex = 100; - en.style.position = 'absolute'; - en.style.right = '0'; - en.style.top = '13em'; - en.style.height = 'auto'; - en.style.transition = 'width 0.3s ease'; - en.style.width = `${window.opts.extra_networks_sidebar_width}vw`; - gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = `${100 - 2 - window.opts.extra_networks_sidebar_width}vw`; + if (lastView !== entries[0].intersectionRatio > 0) { + lastView = entries[0].intersectionRatio > 0; + if (lastView) { + refreshENpage(); + // sortExtraNetworks('fixed'); + if (window.opts.extra_networks_card_cover === 'cover') { + en.style.position = 'absolute'; + en.style.height = 'unset'; + en.style.width = 'unset'; + en.style.right = 'unset'; + en.style.top = '13em'; + en.style.transition = ''; + en.style.zIndex = 100; + gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = 'unset'; + } else if (window.opts.extra_networks_card_cover === 'sidebar') { + en.style.position = 'absolute'; + en.style.height = 'auto'; + en.style.width = `${window.opts.extra_networks_sidebar_width}vw`; + en.style.right = '0'; + en.style.top = '13em'; + en.style.transition = 'width 0.3s ease'; + en.style.zIndex = 100; + gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = `${100 - 2 - window.opts.extra_networks_sidebar_width}vw`; + } else { + en.style.position = 'relative'; + en.style.height = 'unset'; + en.style.width = 'unset'; + en.style.right = 'unset'; + en.style.top = 0; + en.style.transition = ''; + en.style.zIndex = 0; + gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = 'unset'; + } } else { - en.style.transition = ''; - en.style.zIndex = 0; - en.style.top = 0; - en.style.position = 'relative'; - en.style.right = 'unset'; - en.style.width = 'unset'; - en.style.height = 'unset'; + if (window.opts.extra_networks_card_cover === 'sidebar') en.style.width = 0; gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = 'unset'; } - } else { - if (window.opts.extra_networks_card_cover === 'sidebar') en.style.width = 0; - gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = 'unset'; } }); intersectionObserver.observe(en); // monitor visibility } +async function showNetworks() { + for (const tabname of ['txt2img', 'img2img', 'control']) { + if (window.opts.extra_networks_show) gradioApp().getElementById(`${tabname}_extra_networks_btn`).click(); + } + log('showNetworks'); +} + async function setupExtraNetworks() { setupExtraNetworksForTab('txt2img'); setupExtraNetworksForTab('img2img'); diff --git a/javascript/loader.js b/javascript/loader.js index 4fb85462f..f3c7fe60f 100644 --- a/javascript/loader.js +++ b/javascript/loader.js @@ -3,7 +3,7 @@ const appStartTime = performance.now(); async function preloadImages() { const dark = window.matchMedia && window.matchMedia('(prefers-color-scheme: dark)').matches; const imagePromises = []; - const num = Math.floor(10 * Math.random()); + const num = Math.floor(9.99 * Math.random()); const imageUrls = [ `file=html/logo-bg-${dark ? 'dark' : 'light'}.jpg`, `file=html/logo-bg-${num}.jpg`, @@ -27,7 +27,7 @@ async function preloadImages() { async function createSplash() { const dark = window.matchMedia && window.matchMedia('(prefers-color-scheme: dark)').matches; log('createSplash', { theme: dark ? 'dark' : 'light' }); - const num = Math.floor(11 * Math.random()); + const num = Math.floor(9.99 * Math.random()); const splash = `
diff --git a/javascript/script.js b/javascript/script.js index 6f52f4a3d..104567dd7 100644 --- a/javascript/script.js +++ b/javascript/script.js @@ -133,19 +133,23 @@ document.addEventListener('DOMContentLoaded', () => { }); /** - * Add a ctrl+enter as a shortcut to start a generation + * Add a listener to the document for keydown events */ document.addEventListener('keydown', (e) => { - let handled = false; - if (e.key !== undefined) { - if ((e.key === 'Enter' && (e.metaKey || e.ctrlKey || e.altKey))) handled = true; - } else if (e.keyCode !== undefined) { - if ((e.keyCode === 13 && (e.metaKey || e.ctrlKey || e.altKey))) handled = true; - } - if (handled) { - const button = getUICurrentTabContent().querySelector('button[id$=_generate]'); - if (button) button.click(); + let elem; + if (e.key === 'Escape') elem = getUICurrentTabContent().querySelector('button[id$=_interrupt]'); + if (e.key === 'Enter' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id$=_generate]'); + if (e.key === 'Backspace' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id$=_reprocess]'); + if (e.key === ' ' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id$=_extra_networks_btn]'); + if (e.key === 's' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id^=save_]'); + if (e.key === 'Insert' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id^=save_]'); + if (e.key === 'Delete' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id^=delete_]'); + // if (e.key === 'm' && e.ctrlKey) elem = gradioApp().getElementById('setting_sd_model_checkpoint'); + if (elem) { e.preventDefault(); + log('hotkey', { key: e.key, meta: e.metaKey, ctrl: e.ctrlKey, alt: e.altKey }, elem?.id, elem.nodeName); + if (elem.nodeName === 'BUTTON') elem.click(); + else elem.focus(); } }); diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 3c81e7d8e..d412d966f 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -38,7 +38,7 @@ td > div > span { overflow-y: auto; max-height: 3em; overflow-x: hidden; } .gradio-button.secondary-down, .gradio-button.secondary-down:hover { box-shadow: 1px 1px 1px rgba(0,0,0,0.25) inset, 0px 0px 3px rgba(0,0,0,0.15) inset; } .gradio-button.secondary-down:hover { background: var(--button-secondary-background-fill-hover); color: var(--button-secondary-text-color-hover); } .gradio-button.tool { max-width: min-content; min-width: min-content !important; font-size: 20px !important; color: var(--body-text-color) !important; align-self: end; margin-bottom: 4px; } -.gradio-checkbox { margin: 0.75em 1.5em 0 0; align-self: center; } +.gradio-checkbox { margin-right: 1em !important; align-self: center; } .gradio-column { min-width: min(160px, 100%) !important; } .gradio-container { max-width: unset !important; padding: var(--block-label-padding) !important; } .gradio-container .prose a, .gradio-container .prose a:visited{ color: unset; text-decoration: none; } @@ -203,7 +203,7 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt #extensions .date { opacity: 0.85; font-size: var(--text-sm); } /* extra networks */ -#txt2img_extra_networks, #img2img_extra_networks, #control_extra_networks { width: 0; } +.extra_networks_root { width: 0; position: absolute; height: auto; right: 0; top: 13em; z-index: 100; } /* default is sidebar view */ .extra-networks { background: var(--background-color); padding: var(--block-label-padding); } .extra-networks > div { margin: 0; border-bottom: none !important; gap: 0.3em 0; } .extra-networks .second-line { display: flex; width: -moz-available; width: -webkit-fill-available; gap: 0.3em; box-shadow: var(--input-shadow); } @@ -319,6 +319,13 @@ div:has(>#tab-gallery-folders) { flex-grow: 0 !important; background-color: var( .gallery-sort { background: var(--input-background-fill) !important; margin: 0 !important; padding: 6px !important; } .gallery-sort:hover { background: var(--button-primary-background-fill-hover) !important; } +/* changelog */ +#changelog_markdown { max-height: 55vh; margin-top: 1em; } +#changelog_result { display: flex; margin-left: 1em; align-items: center; } +.changelog_arrow { font-size: 2em; padding: 0.1em; cursor: pointer; height: 1em; background-color: var(--button-secondary-background-fill); } +.changelog_arrow:hover { background-color: var(--button-primary-border-color-hover); } +.changelog_highlight { background-color: var(--color-warning); } + /* loader */ .splash { position: fixed; top: 0; left: 0; width: 100vw; height: 100vh; z-index: 1000; display: block; text-align: center; } .motd { margin-top: 2em; color: var(--body-text-color-subdued); font-family: monospace; font-variant: all-petite-caps; } diff --git a/javascript/startup.js b/javascript/startup.js index 81bfa1a61..51e1800fc 100644 --- a/javascript/startup.js +++ b/javascript/startup.js @@ -16,6 +16,7 @@ async function initStartup() { initImageViewer(); initGallery(); initiGenerationParams(); + initChangelog(); setupControlUI(); // reconnect server session @@ -31,6 +32,7 @@ async function initStartup() { removeSplash(); // post startup tasks that may take longer but are not critical + showNetworks(); setHints(); applyStyles(); initIndexDB(); diff --git a/launch.py b/launch.py index 222c24fe2..903234490 100755 --- a/launch.py +++ b/launch.py @@ -164,7 +164,6 @@ def start_server(immediate=True, server=None): module_spec = importlib.util.spec_from_file_location('webui', 'webui.py') server = importlib.util.module_from_spec(module_spec) installer.log.debug(f'Starting module: {server}') - get_custom_args() module_spec.loader.exec_module(server) uvicorn = None if args.test: @@ -209,6 +208,8 @@ def main(): installer.log.info('Skipping GIT operations') installer.check_version() installer.log.info(f'Platform: {installer.print_dict(installer.get_platform())}') + installer.check_venv() + installer.log.info(f'Args: {sys.argv[1:]}') if not args.skip_env: installer.set_environment() if args.uv: @@ -246,6 +247,7 @@ def main(): installer.log.warning(f'See log file for more details: {installer.log_file}') installer.extensions_preload(parser) # adds additional args from extensions args = installer.parse_args(parser) + get_custom_args() uv, instance = start_server(immediate=True, server=None) while True: diff --git a/modules/api/process.py b/modules/api/process.py index dcb3b6487..f50d58381 100644 --- a/modules/api/process.py +++ b/modules/api/process.py @@ -113,10 +113,9 @@ class APIProcess(): images = [] scores = [] with self.queue_lock: - yolo.load() - faces = yolo.predict(image) + faces = yolo.predict('face-yolo8n', image) for face in faces: - images.append(encode_pil_to_base64(face.face)) + images.append(encode_pil_to_base64(face.item)) scores.append(face.score) shared.state.end(api=False) return ResFace(images=images, scores=scores) diff --git a/modules/cmd_args.py b/modules/cmd_args.py index afbe239fe..752ad02c0 100644 --- a/modules/cmd_args.py +++ b/modules/cmd_args.py @@ -2,104 +2,134 @@ import os import argparse from modules.paths import data_path + parser = argparse.ArgumentParser(description="SD.Next", conflict_handler='resolve', epilog='For other options see UI Settings page', prog='', add_help=True, formatter_class=lambda prog: argparse.HelpFormatter(prog, max_help_position=55, indent_increment=2, width=200)) parser._optionals = parser.add_argument_group('Other options') # pylint: disable=protected-access -group = parser.add_argument_group('Server options') - -# main server args -group.add_argument("--config", type=str, default=os.environ.get("SD_CONFIG", os.path.join(data_path, 'config.json')), help="Use specific server configuration file, default: %(default)s") -group.add_argument("--ui-config", type=str, default=os.environ.get("SD_UICONFIG", os.path.join(data_path, 'ui-config.json')), help="Use specific UI configuration file, default: %(default)s") -group.add_argument("--medvram", default=os.environ.get("SD_MEDVRAM", False), action='store_true', help="Split model stages and keep only active part in VRAM, default: %(default)s") -group.add_argument("--lowvram", default=os.environ.get("SD_LOWVRAM", False), action='store_true', help="Split model components and keep only active part in VRAM, default: %(default)s") -group.add_argument("--ckpt", type=str, default=os.environ.get("SD_MODEL", None), help="Path to model checkpoint to load immediately, default: %(default)s") -group.add_argument('--vae', type=str, default=os.environ.get("SD_VAE", None), help='Path to VAE checkpoint to load immediately, default: %(default)s') -group.add_argument("--data-dir", type=str, default=os.environ.get("SD_DATADIR", ''), help="Base path where all user data is stored, default: %(default)s") -group.add_argument("--models-dir", type=str, default=os.environ.get("SD_MODELSDIR", 'models'), help="Base path where all models are stored, default: %(default)s",) -group.add_argument("--allow-code", default=os.environ.get("SD_ALLOWCODE", False), action='store_true', help="Allow custom script execution, default: %(default)s") -group.add_argument("--share", default=os.environ.get("SD_SHARE", False), action='store_true', help="Enable UI accessible through Gradio site, default: %(default)s") -group.add_argument("--insecure", default=os.environ.get("SD_INSECURE", False), action='store_true', help="Enable extensions tab regardless of other options, default: %(default)s") -group.add_argument("--use-cpu", nargs='+', default=[], type=str.lower, help="Force use CPU for specified modules, default: %(default)s") -group.add_argument("--listen", default=os.environ.get("SD_LISTEN", False), action='store_true', help="Launch web server using public IP address, default: %(default)s") -group.add_argument("--port", type=int, default=os.environ.get("SD_PORT", 7860), help="Launch web server with given server port, default: %(default)s") -group.add_argument("--freeze", default=os.environ.get("SD_FREEZE", False), action='store_true', help="Disable editing settings") -group.add_argument("--auth", type=str, default=os.environ.get("SD_AUTH", None), help='Set access authentication like "user:pwd,user:pwd""') -group.add_argument("--auth-file", type=str, default=os.environ.get("SD_AUTHFILE", None), help='Set access authentication using file, default: %(default)s') -group.add_argument("--autolaunch", default=os.environ.get("SD_AUTOLAUNCH", False), action='store_true', help="Open the UI URL in the system's default browser upon launch") -group.add_argument('--docs', default=os.environ.get("SD_DOCS", False), action='store_true', help = "Mount API docs, default: %(default)s") -group.add_argument('--api-only', default=os.environ.get("SD_APIONLY", False), action='store_true', help = "Run in API only mode without starting UI") -group.add_argument("--device-id", type=str, default=os.environ.get("SD_DEVICEID", None), help="Select the default CUDA device to use, default: %(default)s") -group.add_argument("--cors-origins", type=str, default=os.environ.get("SD_CORSORIGINS", None), help="Allowed CORS origins as comma-separated list, default: %(default)s") -group.add_argument("--cors-regex", type=str, default=os.environ.get("SD_CORSREGEX", None), help="Allowed CORS origins as regular expression, default: %(default)s") -group.add_argument("--tls-keyfile", type=str, default=os.environ.get("SD_TLSKEYFILE", None), help="Enable TLS and specify key file, default: %(default)s") -group.add_argument("--tls-certfile", type=str, default=os.environ.get("SD_TLSCERTFILE", None), help="Enable TLS and specify cert file, default: %(default)s") -group.add_argument("--tls-selfsign", action="store_true", default=os.environ.get("SD_TLSSELFSIGN", False), help="Enable TLS with self-signed certificates, default: %(default)s") -group.add_argument("--server-name", type=str, default=os.environ.get("SD_SERVERNAME", None), help="Sets hostname of server, default: %(default)s") -group.add_argument("--no-hashing", default=os.environ.get("SD_NOHASHING", False), action='store_true', help="Disable hashing of checkpoints, default: %(default)s") -group.add_argument("--no-metadata", default=os.environ.get("SD_NOMETADATA", False), action='store_true', help="Disable reading of metadata from models, default: %(default)s") -group.add_argument("--profile", default=os.environ.get("SD_PROFILE", False), action='store_true', help="Run profiler, default: %(default)s") -group.add_argument("--disable-queue", default=os.environ.get("SD_DISABLEQUEUE", False), action='store_true', help="Disable queues, default: %(default)s") -group.add_argument('--debug', default=os.environ.get("SD_DEBUG", False), action='store_true', help = "Run installer with debug logging, default: %(default)s") -group.add_argument('--use-directml', default=os.environ.get("SD_USEDIRECTML", False), action='store_true', help = "Use DirectML if no compatible GPU is detected, default: %(default)s") -group.add_argument('--use-zluda', default=os.environ.get("SD_USEZLUDA", False), action='store_true', help = "Force use ZLUDA, AMD GPUs only, default: %(default)s") -group.add_argument("--use-openvino", default=os.environ.get("SD_USEOPENVINO", False), action='store_true', help="Use Intel OpenVINO backend, default: %(default)s") -group.add_argument("--use-ipex", default=os.environ.get("SD_USEIPX", False), action='store_true', help="Force use Intel OneAPI XPU backend, default: %(default)s") -group.add_argument("--use-cuda", default=os.environ.get("SD_USECUDA", False), action='store_true', help="Force use nVidia CUDA backend, default: %(default)s") -group.add_argument("--use-rocm", default=os.environ.get("SD_USEROCM", False), action='store_true', help="Force use AMD ROCm backend, default: %(default)s") -group.add_argument('--subpath', type=str, default=os.environ.get("SD_SUBPATH", None), help='Customize the URL subpath for usage with reverse proxy') -group.add_argument('--backend', type=str, default=os.environ.get("SD_BACKEND", None), choices=['original', 'diffusers'], required=False, help='force model pipeline type') -group.add_argument('--theme', type=str, default=os.environ.get("SD_THEME", None), help='Override UI theme') -# removed args are added here as hidden in fixed format for compatbility reasons -group.add_argument("-f", action='store_true', help=argparse.SUPPRESS) # allows running as root; implemented outside of webui -group.add_argument("--ui-settings-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'config.json')) -group.add_argument("--ui-config-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'ui-config.json')) -group.add_argument("--hide-ui-dir-config", action='store_true', help=argparse.SUPPRESS, default=False) -group.add_argument("--theme", type=str, help=argparse.SUPPRESS, default=None) -group.add_argument("--disable-console-progressbars", action='store_true', help=argparse.SUPPRESS, default=True) -group.add_argument("--disable-safe-unpickle", action='store_true', help=argparse.SUPPRESS, default=True) -group.add_argument("--lowram", action='store_true', help=argparse.SUPPRESS) -group.add_argument("--disable-extension-access", default=False, action='store_true', help=argparse.SUPPRESS) -group.add_argument("--allowed-paths", nargs='+', default=[], type=str, required=False, help="add additional paths to paths allowed for web access") -group.add_argument("--api", help=argparse.SUPPRESS, default=True) -group.add_argument("--api-auth", type=str, help=argparse.SUPPRESS, default=None) +def main_args(): + # main server args + group_config = parser.add_argument_group('Configuration') + group_config.add_argument('--backend', type=str, default=os.environ.get("SD_BACKEND", None), choices=['original', 'diffusers'], required=False, help='force model pipeline type') + group_config.add_argument("--config", type=str, default=os.environ.get("SD_CONFIG", os.path.join(data_path, 'config.json')), help="Use specific server configuration file, default: %(default)s") + group_config.add_argument("--ui-config", type=str, default=os.environ.get("SD_UICONFIG", os.path.join(data_path, 'ui-config.json')), help="Use specific UI configuration file, default: %(default)s") + group_config.add_argument("--medvram", default=os.environ.get("SD_MEDVRAM", False), action='store_true', help="Split model stages and keep only active part in VRAM, default: %(default)s") + group_config.add_argument("--lowvram", default=os.environ.get("SD_LOWVRAM", False), action='store_true', help="Split model components and keep only active part in VRAM, default: %(default)s") + group_config.add_argument("--freeze", default=os.environ.get("SD_FREEZE", False), action='store_true', help="Disable editing settings") + + group_paths = parser.add_argument_group('Paths') + group_paths.add_argument("--ckpt", type=str, default=os.environ.get("SD_MODEL", None), help="Path to model checkpoint to load immediately, default: %(default)s") + group_paths.add_argument("--data-dir", type=str, default=os.environ.get("SD_DATADIR", ''), help="Base path where all user data is stored, default: %(default)s") + group_paths.add_argument("--models-dir", type=str, default=os.environ.get("SD_MODELSDIR", 'models'), help="Base path where all models are stored, default: %(default)s",) + + group_diag = parser.add_argument_group('Diagnostics') + group_diag.add_argument("--no-hashing", default=os.environ.get("SD_NOHASHING", False), action='store_true', help="Disable hashing of checkpoints, default: %(default)s") + group_diag.add_argument("--no-metadata", default=os.environ.get("SD_NOMETADATA", False), action='store_true', help="Disable reading of metadata from models, default: %(default)s") + group_diag.add_argument("--profile", default=os.environ.get("SD_PROFILE", False), action='store_true', help="Run profiler, default: %(default)s") + group_diag.add_argument("--disable-queue", default=os.environ.get("SD_DISABLEQUEUE", False), action='store_true', help="Disable queues, default: %(default)s") + group_diag.add_argument('--debug', default=os.environ.get("SD_DEBUG", False), action='store_true', help = "Run installer with debug logging, default: %(default)s") + + group_compute = parser.add_argument_group('Compute Engine') + group_compute.add_argument('--use-directml', default=os.environ.get("SD_USEDIRECTML", False), action='store_true', help = "Use DirectML if no compatible GPU is detected, default: %(default)s") + group_compute.add_argument('--use-zluda', default=os.environ.get("SD_USEZLUDA", False), action='store_true', help = "Force use ZLUDA, AMD GPUs only, default: %(default)s") + group_compute.add_argument("--use-openvino", default=os.environ.get("SD_USEOPENVINO", False), action='store_true', help="Use Intel OpenVINO backend, default: %(default)s") + group_compute.add_argument("--use-ipex", default=os.environ.get("SD_USEIPX", False), action='store_true', help="Force use Intel OneAPI XPU backend, default: %(default)s") + group_compute.add_argument("--use-cuda", default=os.environ.get("SD_USECUDA", False), action='store_true', help="Force use nVidia CUDA backend, default: %(default)s") + group_compute.add_argument("--use-rocm", default=os.environ.get("SD_USEROCM", False), action='store_true', help="Force use AMD ROCm backend, default: %(default)s") + group_diag.add_argument("--device-id", type=str, default=os.environ.get("SD_DEVICEID", None), help="Select the default CUDA device to use, default: %(default)s") + + group_http = parser.add_argument_group('HTTP') + group_http.add_argument('--theme', type=str, default=os.environ.get("SD_THEME", None), help='Override UI theme') + group_http.add_argument("--server-name", type=str, default=os.environ.get("SD_SERVERNAME", None), help="Sets hostname of server, default: %(default)s") + group_http.add_argument("--tls-keyfile", type=str, default=os.environ.get("SD_TLSKEYFILE", None), help="Enable TLS and specify key file, default: %(default)s") + group_http.add_argument("--tls-certfile", type=str, default=os.environ.get("SD_TLSCERTFILE", None), help="Enable TLS and specify cert file, default: %(default)s") + group_http.add_argument("--tls-selfsign", action="store_true", default=os.environ.get("SD_TLSSELFSIGN", False), help="Enable TLS with self-signed certificates, default: %(default)s") + group_http.add_argument("--cors-origins", type=str, default=os.environ.get("SD_CORSORIGINS", None), help="Allowed CORS origins as comma-separated list, default: %(default)s") + group_http.add_argument("--cors-regex", type=str, default=os.environ.get("SD_CORSREGEX", None), help="Allowed CORS origins as regular expression, default: %(default)s") + group_http.add_argument('--subpath', type=str, default=os.environ.get("SD_SUBPATH", None), help='Customize the URL subpath for usage with reverse proxy') + group_http.add_argument("--autolaunch", default=os.environ.get("SD_AUTOLAUNCH", False), action='store_true', help="Open the UI URL in the system's default browser upon launch") + group_http.add_argument('--docs', default=os.environ.get("SD_DOCS", False), action='store_true', help = "Mount API docs, default: %(default)s") + group_http.add_argument("--auth", type=str, default=os.environ.get("SD_AUTH", None), help='Set access authentication like "user:pwd,user:pwd""') + group_http.add_argument("--auth-file", type=str, default=os.environ.get("SD_AUTHFILE", None), help='Set access authentication using file, default: %(default)s') + group_http.add_argument('--api-only', default=os.environ.get("SD_APIONLY", False), action='store_true', help = "Run in API only mode without starting UI") + group_http.add_argument("--allowed-paths", nargs='+', default=[], type=str, required=False, help="add additional paths to paths allowed for web access") + group_http.add_argument("--share", default=os.environ.get("SD_SHARE", False), action='store_true', help="Enable UI accessible through Gradio site, default: %(default)s") + group_http.add_argument("--insecure", default=os.environ.get("SD_INSECURE", False), action='store_true', help="Enable extensions tab regardless of other options, default: %(default)s") + group_http.add_argument("--listen", default=os.environ.get("SD_LISTEN", False), action='store_true', help="Launch web server using public IP address, default: %(default)s") + group_http.add_argument("--port", type=int, default=os.environ.get("SD_PORT", 7860), help="Launch web server with given server port, default: %(default)s") -def compatibility_args(opts, args): +def compatibility_args(): + group_compat = parser.add_argument_group('Compatibility options') + # removed args are added here as hidden in fixed format for compatbility reasons + group_compat.add_argument("--allow-code", default=os.environ.get("SD_ALLOWCODE", False), action='store_true', help=argparse.SUPPRESS) + group_compat.add_argument("--use-cpu", nargs='+', default=[], type=str.lower, help=argparse.SUPPRESS) + group_compat.add_argument("-f", action='store_true', help=argparse.SUPPRESS) # allows running as root; implemented outside of webui + group_compat.add_argument('--vae', type=str, default=os.environ.get("SD_VAE", None), help=argparse.SUPPRESS) + group_compat.add_argument("--ui-settings-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'config.json')) + group_compat.add_argument("--ui-config-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'ui-config.json')) + group_compat.add_argument("--hide-ui-dir-config", action='store_true', help=argparse.SUPPRESS, default=False) + group_compat.add_argument("--theme", type=str, help=argparse.SUPPRESS, default=None) + group_compat.add_argument("--disable-console-progressbars", action='store_true', help=argparse.SUPPRESS, default=True) + group_compat.add_argument("--disable-safe-unpickle", action='store_true', help=argparse.SUPPRESS, default=True) + group_compat.add_argument("--lowram", action='store_true', help=argparse.SUPPRESS) + group_compat.add_argument("--disable-extension-access", default=False, action='store_true', help=argparse.SUPPRESS) + group_compat.add_argument("--api", help=argparse.SUPPRESS, default=True) + group_compat.add_argument("--api-auth", type=str, help=argparse.SUPPRESS, default=None) + + +def settings_args(opts, args): + group_compat = parser.add_argument_group('Compatibility options') + # removed args are added here as hidden in fixed format for compatbility reasons + group_compat.add_argument("--allow-code", default=os.environ.get("SD_ALLOWCODE", False), action='store_true', help=argparse.SUPPRESS) + group_compat.add_argument("--use-cpu", nargs='+', default=[], type=str.lower, help=argparse.SUPPRESS) + group_compat.add_argument("-f", action='store_true', help=argparse.SUPPRESS) # allows running as root; implemented outside of webui + group_compat.add_argument('--vae', type=str, default=os.environ.get("SD_VAE", None), help=argparse.SUPPRESS) + group_compat.add_argument("--ui-settings-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'config.json')) + group_compat.add_argument("--ui-config-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'ui-config.json')) + group_compat.add_argument("--hide-ui-dir-config", action='store_true', help=argparse.SUPPRESS, default=False) + group_compat.add_argument("--theme", type=str, help=argparse.SUPPRESS, default=None) + group_compat.add_argument("--disable-console-progressbars", action='store_true', help=argparse.SUPPRESS, default=True) + group_compat.add_argument("--disable-safe-unpickle", action='store_true', help=argparse.SUPPRESS, default=True) + group_compat.add_argument("--lowram", action='store_true', help=argparse.SUPPRESS) + group_compat.add_argument("--disable-extension-access", default=False, action='store_true', help=argparse.SUPPRESS) + group_compat.add_argument("--allowed-paths", nargs='+', default=[], type=str, required=False, help="add additional paths to paths allowed for web access") + group_compat.add_argument("--api", help=argparse.SUPPRESS, default=True) + group_compat.add_argument("--api-auth", type=str, help=argparse.SUPPRESS, default=None) # removed args that have been moved to opts are added here as hidden with default values as defined in opts - group.add_argument("--ckpt-dir", type=str, help=argparse.SUPPRESS, default=opts.ckpt_dir) - group.add_argument("--vae-dir", type=str, help=argparse.SUPPRESS, default=opts.vae_dir) - group.add_argument("--embeddings-dir", type=str, help=argparse.SUPPRESS, default=opts.embeddings_dir) - group.add_argument("--embeddings-templates-dir", type=str, help=argparse.SUPPRESS, default=opts.embeddings_templates_dir) - group.add_argument("--hypernetwork-dir", type=str, help=argparse.SUPPRESS, default=opts.hypernetwork_dir) - group.add_argument("--codeformer-models-path", type=str, help=argparse.SUPPRESS, default=opts.codeformer_models_path) - group.add_argument("--gfpgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.gfpgan_models_path) - group.add_argument("--esrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.esrgan_models_path) - group.add_argument("--bsrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.bsrgan_models_path) - group.add_argument("--realesrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.realesrgan_models_path) - group.add_argument("--scunet-models-path", help=argparse.SUPPRESS, default=opts.scunet_models_path) - group.add_argument("--swinir-models-path", help=argparse.SUPPRESS, default=opts.swinir_models_path) - group.add_argument("--ldsr-models-path", help=argparse.SUPPRESS, default=opts.ldsr_models_path) - group.add_argument("--clip-models-path", type=str, help=argparse.SUPPRESS, default=opts.clip_models_path) - group.add_argument("--opt-channelslast", help=argparse.SUPPRESS, action='store_true', default=opts.opt_channelslast) - group.add_argument("--xformers", default=(opts.cross_attention_optimization == "xFormers"), action='store_true', help=argparse.SUPPRESS) - group.add_argument("--disable-nan-check", help=argparse.SUPPRESS, action='store_true', default=opts.disable_nan_check) - group.add_argument("--rollback-vae", help=argparse.SUPPRESS, default=opts.rollback_vae) - group.add_argument("--no-half", help=argparse.SUPPRESS, action='store_true', default=opts.no_half) - group.add_argument("--no-half-vae", help=argparse.SUPPRESS, action='store_true', default=opts.no_half_vae) - group.add_argument("--precision", help=argparse.SUPPRESS, default=opts.precision) - group.add_argument("--sub-quad-q-chunk-size", help=argparse.SUPPRESS, default=opts.sub_quad_q_chunk_size) - group.add_argument("--sub-quad-kv-chunk-size", help=argparse.SUPPRESS, default=opts.sub_quad_kv_chunk_size) - group.add_argument("--sub-quad-chunk-threshold", help=argparse.SUPPRESS, default=opts.sub_quad_chunk_threshold) - group.add_argument("--lora-dir", help=argparse.SUPPRESS, default=opts.lora_dir) - group.add_argument("--lyco-dir", help=argparse.SUPPRESS, default=opts.lyco_dir) - group.add_argument("--embeddings-dir", help=argparse.SUPPRESS, default=opts.embeddings_dir) - group.add_argument("--hypernetwork-dir", help=argparse.SUPPRESS, default=opts.hypernetwork_dir) - group.add_argument("--lyco-patch-lora", help=argparse.SUPPRESS, action='store_true', default=False) - group.add_argument("--lyco-debug", help=argparse.SUPPRESS, action='store_true', default=False) - group.add_argument("--enable-console-prompts", help=argparse.SUPPRESS, action='store_true', default=False) - group.add_argument("--safe", help=argparse.SUPPRESS, action='store_true', default=False) - group.add_argument("--use-xformers", help=argparse.SUPPRESS, action='store_true', default=False) + group_compat.add_argument("--ckpt-dir", type=str, help=argparse.SUPPRESS, default=opts.ckpt_dir) + group_compat.add_argument("--vae-dir", type=str, help=argparse.SUPPRESS, default=opts.vae_dir) + group_compat.add_argument("--embeddings-dir", type=str, help=argparse.SUPPRESS, default=opts.embeddings_dir) + group_compat.add_argument("--embeddings-templates-dir", type=str, help=argparse.SUPPRESS, default=opts.embeddings_templates_dir) + group_compat.add_argument("--hypernetwork-dir", type=str, help=argparse.SUPPRESS, default=opts.hypernetwork_dir) + group_compat.add_argument("--codeformer-models-path", type=str, help=argparse.SUPPRESS, default=opts.codeformer_models_path) + group_compat.add_argument("--gfpgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.gfpgan_models_path) + group_compat.add_argument("--esrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.esrgan_models_path) + group_compat.add_argument("--bsrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.bsrgan_models_path) + group_compat.add_argument("--realesrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.realesrgan_models_path) + group_compat.add_argument("--scunet-models-path", help=argparse.SUPPRESS, default=opts.scunet_models_path) + group_compat.add_argument("--swinir-models-path", help=argparse.SUPPRESS, default=opts.swinir_models_path) + group_compat.add_argument("--ldsr-models-path", help=argparse.SUPPRESS, default=opts.ldsr_models_path) + group_compat.add_argument("--clip-models-path", type=str, help=argparse.SUPPRESS, default=opts.clip_models_path) + group_compat.add_argument("--opt-channelslast", help=argparse.SUPPRESS, action='store_true', default=opts.opt_channelslast) + group_compat.add_argument("--xformers", default=(opts.cross_attention_optimization == "xFormers"), action='store_true', help=argparse.SUPPRESS) + group_compat.add_argument("--disable-nan-check", help=argparse.SUPPRESS, action='store_true', default=opts.disable_nan_check) + group_compat.add_argument("--rollback-vae", help=argparse.SUPPRESS, default=opts.rollback_vae) + group_compat.add_argument("--no-half", help=argparse.SUPPRESS, action='store_true', default=opts.no_half) + group_compat.add_argument("--no-half-vae", help=argparse.SUPPRESS, action='store_true', default=opts.no_half_vae) + group_compat.add_argument("--precision", help=argparse.SUPPRESS, default=opts.precision) + group_compat.add_argument("--sub-quad-q-chunk-size", help=argparse.SUPPRESS, default=opts.sub_quad_q_chunk_size) + group_compat.add_argument("--sub-quad-kv-chunk-size", help=argparse.SUPPRESS, default=opts.sub_quad_kv_chunk_size) + group_compat.add_argument("--sub-quad-chunk-threshold", help=argparse.SUPPRESS, default=opts.sub_quad_chunk_threshold) + group_compat.add_argument("--lora-dir", help=argparse.SUPPRESS, default=opts.lora_dir) + group_compat.add_argument("--lyco-dir", help=argparse.SUPPRESS, default=opts.lyco_dir) + group_compat.add_argument("--embeddings-dir", help=argparse.SUPPRESS, default=opts.embeddings_dir) + group_compat.add_argument("--hypernetwork-dir", help=argparse.SUPPRESS, default=opts.hypernetwork_dir) + group_compat.add_argument("--lyco-patch-lora", help=argparse.SUPPRESS, action='store_true', default=False) + group_compat.add_argument("--lyco-debug", help=argparse.SUPPRESS, action='store_true', default=False) + group_compat.add_argument("--enable-console-prompts", help=argparse.SUPPRESS, action='store_true', default=False) + group_compat.add_argument("--safe", help=argparse.SUPPRESS, action='store_true', default=False) + group_compat.add_argument("--use-xformers", help=argparse.SUPPRESS, action='store_true', default=False) # removed opts are added here with fixed values for compatibility reasons opts.use_old_emphasis_implementation = False @@ -126,3 +156,7 @@ def compatibility_args(opts, args): args = parser.parse_args() return args + + +main_args() +compatibility_args() diff --git a/modules/consistory/__init__.py b/modules/consistory/__init__.py new file mode 100644 index 000000000..2a06b53a6 --- /dev/null +++ b/modules/consistory/__init__.py @@ -0,0 +1,6 @@ +""" +original code from +""" +from .consistory_pipeline import ConsistoryExtendAttnSDXLPipeline +from .consistory_unet_sdxl import ConsistorySDXLUNet2DConditionModel +from .consistory_run import run_anchor_generation, run_extra_generation diff --git a/modules/consistory/attention_processor.py b/modules/consistory/attention_processor.py new file mode 100644 index 000000000..04985ee4f --- /dev/null +++ b/modules/consistory/attention_processor.py @@ -0,0 +1,287 @@ +# Copyright 2023 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# Not a contribution +# Changes made by NVIDIA CORPORATION & AFFILIATES enabling ConsiStory or otherwise documented as NVIDIA-proprietary +# are not a contribution and subject to the license under the LICENSE file located at the root directory. + + +from typing import Callable, Optional +import torch +import torch.nn.functional as F +from diffusers.utils import USE_PEFT_BACKEND +from diffusers.models.attention_processor import Attention +from .consistory_utils import AnchorCache, FeatureInjector, QueryStore + + +class ConsistoryAttnStoreProcessor: + def __init__(self, attnstore, place_in_unet): + super().__init__() + self.attnstore = attnstore + self.place_in_unet = place_in_unet + + def __call__(self, attn: Attention, hidden_states, encoder_hidden_states=None, attention_mask=None, record_attention=True, **kwargs): + batch_size, sequence_length, _ = hidden_states.shape + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + + query = attn.to_q(hidden_states) + + is_cross = encoder_hidden_states is not None + encoder_hidden_states = encoder_hidden_states if encoder_hidden_states is not None else hidden_states + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + query = attn.head_to_batch_dim(query) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + + attention_probs = attn.get_attention_scores(query, key, attention_mask) + + # only need to store attention maps during the Attend and Excite process + # if attention_probs.requires_grad: + if record_attention: + self.attnstore(attention_probs, is_cross, self.place_in_unet, attn.heads) + + hidden_states = torch.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + return hidden_states + + +class ConsistoryExtendedAttnXFormersAttnProcessor: + r""" + Processor for implementing memory efficient attention using xFormers. + + Args: + attention_op (`Callable`, *optional*, defaults to `None`): + The base + [operator](https://facebookresearch.github.io/xformers/components/ops.html#xformers.ops.AttentionOpBase) to + use as the attention operator. It is recommended to set to `None`, and allow xFormers to choose the best + operator. + """ + + def __init__(self, place_in_unet, attnstore, extended_attn_kwargs, attention_op: Optional[Callable] = None): + self.attention_op = attention_op + self.t_range = extended_attn_kwargs.get('t_range', []) + self.extend_kv_unet_parts = extended_attn_kwargs.get('extend_kv_unet_parts', ['down', 'mid', 'up']) + + self.place_in_unet = place_in_unet + self.curr_unet_part = self.place_in_unet.split('_')[0] + self.attnstore = attnstore + + def __call__( + self, + attn: Attention, + hidden_states: torch.FloatTensor, + encoder_hidden_states: Optional[torch.FloatTensor] = None, + attention_mask: Optional[torch.FloatTensor] = None, + temb: Optional[torch.FloatTensor] = None, + scale: float = 1.0, + perform_extend_attn: bool = False, + query_store: Optional[QueryStore] = None, + feature_injector: Optional[FeatureInjector] = None, + anchors_cache: Optional[AnchorCache] = None, + **kwargs + ) -> torch.FloatTensor: + residual = hidden_states + + args = () if USE_PEFT_BACKEND else (scale,) + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + else: + batch_size, wh, channel = hidden_states.shape + height = width = int(wh ** 0.5) + + is_cross = encoder_hidden_states is not None + perform_extend_attn = perform_extend_attn and (not is_cross) and \ + any([self.attnstore.curr_iter >= x[0] and self.attnstore.curr_iter <= x[1] for x in self.t_range]) and \ + self.curr_unet_part in self.extend_kv_unet_parts + + batch_size, key_tokens, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + + attention_mask = attn.prepare_attention_mask(attention_mask, key_tokens, batch_size) + if attention_mask is not None: + # expand our mask's singleton query_tokens dimension: + # [batch*heads, 1, key_tokens] -> + # [batch*heads, query_tokens, key_tokens] + # so that it can be added as a bias onto the attention scores that xformers computes: + # [batch*heads, query_tokens, key_tokens] + # we do this explicitly because xformers doesn't broadcast the singleton dimension for us. + _, query_tokens, _ = hidden_states.shape + attention_mask = attention_mask.expand(-1, query_tokens, -1) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states, *args) + + if (self.curr_unet_part in self.extend_kv_unet_parts) and query_store and query_store.mode == 'cache': + query_store.cache_query(query, self.place_in_unet) + elif perform_extend_attn and query_store and query_store.mode == 'inject': + query = query_store.inject_query(query, self.place_in_unet, self.attnstore.curr_iter) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states, *args) + value = attn.to_v(encoder_hidden_states, *args) + + query = attn.head_to_batch_dim(query).contiguous() + + if perform_extend_attn: + # Anchor Caching + if anchors_cache and anchors_cache.is_cache_mode(): + if self.place_in_unet not in anchors_cache.input_h_cache: + anchors_cache.input_h_cache[self.place_in_unet] = {} + + # Hidden states inside the mask, for uncond (index 0) and cond (index 1) prompts + subjects_hidden_states = torch.stack([x[self.attnstore.last_mask_dropout[width]] for x in hidden_states.chunk(2)]) + anchors_cache.input_h_cache[self.place_in_unet][self.attnstore.curr_iter] = subjects_hidden_states + + if anchors_cache and anchors_cache.is_inject_mode(): + # We make extended key and value by concatenating the original key and value with the query. + anchors_hidden_states = anchors_cache.input_h_cache[self.place_in_unet][self.attnstore.curr_iter] + + anchors_keys = attn.to_k(anchors_hidden_states, *args) + anchors_values = attn.to_v(anchors_hidden_states, *args) + + extended_key = torch.cat([torch.cat([key.chunk(2, dim=0)[x], anchors_keys[x].unsqueeze(0)], dim=1) for x in range(2)]) + extended_value = torch.cat([torch.cat([value.chunk(2, dim=0)[x], anchors_values[x].unsqueeze(0)], dim=1) for x in range(2)]) + + extended_key = attn.head_to_batch_dim(extended_key).contiguous() + extended_value = attn.head_to_batch_dim(extended_value).contiguous() + + # attn_masks needs to be of shape [batch_size, query_tokens, key_tokens] + # hidden_states = xformers.ops.memory_efficient_attention(query, extended_key, extended_value, op=self.attention_op, scale=attn.scale) + hidden_states = F.scaled_dot_product_attention(query, extended_key, extended_value, scale=attn.scale) + else: + # # We make extended key and value by concatenating the original key and value with the query. + # attention_mask_bias = self.attnstore.get_attn_mask_bias(tgt_size = width, bsz = batch_size) + + # if attention_mask_bias is not None: + # attention_mask_bias = torch.cat([x.unsqueeze(0).expand(attn.heads, -1, -1) for x in attention_mask_bias]) + + # Pre-allocate the output tensor + ex_out = torch.empty_like(query) + + for i in range(batch_size): + start_idx = i * attn.heads + end_idx = start_idx + attn.heads + + attention_mask = self.attnstore.get_extended_attn_mask_instance(width, i%(batch_size//2)) + + curr_q = query[start_idx:end_idx] + + if i < batch_size//2: + curr_k = key[:batch_size//2] + curr_v = value[:batch_size//2] + else: + curr_k = key[batch_size//2:] + curr_v = value[batch_size//2:] + + curr_k = curr_k.flatten(0,1)[attention_mask].unsqueeze(0) + curr_v = curr_v.flatten(0,1)[attention_mask].unsqueeze(0) + + curr_k = attn.head_to_batch_dim(curr_k).contiguous() + curr_v = attn.head_to_batch_dim(curr_v).contiguous() + + # hidden_states = xformers.ops.memory_efficient_attention(curr_q, curr_k, curr_v, op=self.attention_op, scale=attn.scale) + hidden_states = F.scaled_dot_product_attention(curr_q, curr_k, curr_v, scale=attn.scale) + + ex_out[start_idx:end_idx] = hidden_states + + hidden_states = ex_out + else: + key = attn.head_to_batch_dim(key).contiguous() + value = attn.head_to_batch_dim(value).contiguous() + + # attn_masks needs to be of shape [batch_size, query_tokens, key_tokens] + # hidden_states = xformers.ops.memory_efficient_attention(query, key, value, op=self.attention_op, scale=attn.scale) + hidden_states = F.scaled_dot_product_attention(query, key, value, scale=attn.scale) + + hidden_states = hidden_states.to(query.dtype) + hidden_states = attn.batch_to_head_dim(hidden_states) + + # linear proj + hidden_states = attn.to_out[0](hidden_states, *args) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if feature_injector is not None: + output_res = int(hidden_states.shape[1] ** 0.5) + + if anchors_cache and anchors_cache.is_inject_mode(): + hidden_states[batch_size//2:] = feature_injector.inject_anchors(hidden_states[batch_size//2:], self.attnstore.curr_iter, output_res, self.attnstore.extended_mapping, self.place_in_unet, anchors_cache) + else: + hidden_states[batch_size//2:] = feature_injector.inject_outputs(hidden_states[batch_size//2:], self.attnstore.curr_iter, output_res, self.attnstore.extended_mapping, self.place_in_unet, anchors_cache) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +def register_extended_self_attn(unet, attnstore, extended_attn_kwargs): + DICT_PLACE_TO_RES = {'down_0': 64, 'down_1': 64, 'down_2': 64, 'down_3': 64, 'down_4': 64, 'down_5': 64, 'down_6': 64, 'down_7': 64, + 'down_8': 32, 'down_9': 32, 'down_10': 32, 'down_11': 32, 'down_12': 32, 'down_13': 32, 'down_14': 32, 'down_15': 32, + 'down_16': 32, 'down_17': 32, 'down_18': 32, 'down_19': 32, 'down_20': 32, 'down_21': 32, 'down_22': 32, 'down_23': 32, + 'down_24': 32, 'down_25': 32, 'down_26': 32, 'down_27': 32, 'down_28': 32, 'down_29': 32, 'down_30': 32, 'down_31': 32, + 'down_32': 32, 'down_33': 32, 'down_34': 32, 'down_35': 32, 'down_36': 32, 'down_37': 32, 'down_38': 32, 'down_39': 32, + 'down_40': 32, 'down_41': 32, 'down_42': 32, 'down_43': 32, 'down_44': 32, 'down_45': 32, 'down_46': 32, 'down_47': 32, + 'mid_120': 32, 'mid_121': 32, 'mid_122': 32, 'mid_123': 32, 'mid_124': 32, 'mid_125': 32, 'mid_126': 32, 'mid_127': 32, + 'mid_128': 32, 'mid_129': 32, 'mid_130': 32, 'mid_131': 32, 'mid_132': 32, 'mid_133': 32, 'mid_134': 32, 'mid_135': 32, + 'mid_136': 32, 'mid_137': 32, 'mid_138': 32, 'mid_139': 32, 'up_49': 32, 'up_51': 32, 'up_53': 32, 'up_55': 32, 'up_57': 32, + 'up_59': 32, 'up_61': 32, 'up_63': 32, 'up_65': 32, 'up_67': 32, 'up_69': 32, 'up_71': 32, 'up_73': 32, 'up_75': 32, + 'up_77': 32, 'up_79': 32, 'up_81': 32, 'up_83': 32, 'up_85': 32, 'up_87': 32, 'up_89': 32, 'up_91': 32, 'up_93': 32, + 'up_95': 32, 'up_97': 32, 'up_99': 32, 'up_101': 32, 'up_103': 32, 'up_105': 32, 'up_107': 32, 'up_109': 64, 'up_111': 64, + 'up_113': 64, 'up_115': 64, 'up_117': 64, 'up_119': 64} + attn_procs = {} + for i, name in enumerate(unet.attn_processors.keys()): + is_self_attn = i % 2 == 0 + if name.startswith("mid_block"): + place_in_unet = f"mid_{i}" + elif name.startswith("up_blocks"): + place_in_unet = f"up_{i}" + elif name.startswith("down_blocks"): + place_in_unet = f"down_{i}" + else: + continue + + if is_self_attn: + attn_procs[name] = ConsistoryExtendedAttnXFormersAttnProcessor(place_in_unet, attnstore, extended_attn_kwargs) + else: + attn_procs[name] = ConsistoryAttnStoreProcessor(attnstore, place_in_unet) + + unet.set_attn_processor(attn_procs) diff --git a/modules/consistory/consistory_pipeline.py b/modules/consistory/consistory_pipeline.py new file mode 100644 index 000000000..b17fb4143 --- /dev/null +++ b/modules/consistory/consistory_pipeline.py @@ -0,0 +1,519 @@ +# Copyright 2023 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# Not a contribution +# Changes made by NVIDIA CORPORATION & AFFILIATES enabling ConsiStory or otherwise documented as NVIDIA-proprietary +# are not a contribution and subject to the license under the LICENSE file located at the root directory. + +import torch +from diffusers.pipelines.stable_diffusion_xl.pipeline_output import StableDiffusionXLPipelineOutput +from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl import StableDiffusionXLPipeline, \ + rescale_noise_cfg, EXAMPLE_DOC_STRING +from diffusers.utils import ( + deprecate, + is_torch_xla_available, + logging, + replace_example_docstring, +) +from typing import Any, Callable, Dict, List, Optional, Tuple, Union + +from .attention_processor import register_extended_self_attn +from .consistory_utils import FeatureInjector, AnchorCache, QueryStore +from .utils.ptp_utils import AttentionStore + +if is_torch_xla_available(): + # import torch_xla.core.xla_model as xm + + XLA_AVAILABLE = True +else: + XLA_AVAILABLE = False + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + +T = torch.Tensor + +class ConsistoryExtendAttnSDXLPipeline( + StableDiffusionXLPipeline +): + + + @torch.no_grad() + @replace_example_docstring(EXAMPLE_DOC_STRING) + def __call__( + self, + prompt: Union[str, List[str]] = None, + prompt_2: Optional[Union[str, List[str]]] = None, + height: Optional[int] = None, + width: Optional[int] = None, + num_inference_steps: int = 50, + denoising_end: Optional[float] = None, + guidance_scale: float = 5.0, + negative_prompt: Optional[Union[str, List[str]]] = None, + negative_prompt_2: Optional[Union[str, List[str]]] = None, + num_images_per_prompt: Optional[int] = 1, + eta: float = 0.0, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.FloatTensor] = None, + prompt_embeds: Optional[torch.FloatTensor] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None, + pooled_prompt_embeds: Optional[torch.FloatTensor] = None, + negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None, + output_type: Optional[str] = "pil", + return_dict: bool = True, + cross_attention_kwargs: Optional[Dict[str, Any]] = None, + guidance_rescale: float = 0.0, + original_size: Optional[Tuple[int, int]] = None, + crops_coords_top_left: Tuple[int, int] = (0, 0), + target_size: Optional[Tuple[int, int]] = None, + negative_original_size: Optional[Tuple[int, int]] = None, + negative_crops_coords_top_left: Tuple[int, int] = (0, 0), + negative_target_size: Optional[Tuple[int, int]] = None, + clip_skip: Optional[int] = None, + callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None, + callback_on_step_end_tensor_inputs: List[str] = ["latents"], + + attention_store_kwargs: Optional[Dict] = None, + extended_attn_kwargs: Optional[Dict] = None, + share_queries: bool = False, + query_store_kwargs: Optional[Dict] = {}, + feature_injector: Optional[FeatureInjector] = None, + anchors_cache: Optional[AnchorCache] = None, + + instance_latents: Optional[torch.FloatTensor] = None, + **kwargs, + ): + r""" + Function invoked when calling the pipeline for generation. + + Args: + prompt (`str` or `List[str]`, *optional*): + The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`. + instead. + prompt_2 (`str` or `List[str]`, *optional*): + The prompt or prompts to be sent to the `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is + used in both text-encoders + height (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): + The height in pixels of the generated image. This is set to 1024 by default for the best results. + Anything below 512 pixels won't work well for + [stabilityai/stable-diffusion-xl-base-1.0](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) + and checkpoints that are not specifically fine-tuned on low resolutions. + width (`int`, *optional*, defaults to self.unet.config.sample_size * self.vae_scale_factor): + The width in pixels of the generated image. This is set to 1024 by default for the best results. + Anything below 512 pixels won't work well for + [stabilityai/stable-diffusion-xl-base-1.0](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) + and checkpoints that are not specifically fine-tuned on low resolutions. + num_inference_steps (`int`, *optional*, defaults to 50): + The number of denoising steps. More denoising steps usually lead to a higher quality image at the + expense of slower inference. + denoising_end (`float`, *optional*): + When specified, determines the fraction (between 0.0 and 1.0) of the total denoising process to be + completed before it is intentionally prematurely terminated. As a result, the returned sample will + still retain a substantial amount of noise as determined by the discrete timesteps selected by the + scheduler. The denoising_end parameter should ideally be utilized when this pipeline forms a part of a + "Mixture of Denoisers" multi-pipeline setup, as elaborated in [**Refining the Image + Output**](https://huggingface.co/docs/diffusers/api/pipelines/stable_diffusion/stable_diffusion_xl#refining-the-image-output) + guidance_scale (`float`, *optional*, defaults to 5.0): + Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598). + `guidance_scale` is defined as `w` of equation 2. of [Imagen + Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale > + 1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`, + usually at the expense of lower image quality. + negative_prompt (`str` or `List[str]`, *optional*): + The prompt or prompts not to guide the image generation. If not defined, one has to pass + `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is + less than `1`). + negative_prompt_2 (`str` or `List[str]`, *optional*): + The prompt or prompts not to guide the image generation to be sent to `tokenizer_2` and + `text_encoder_2`. If not defined, `negative_prompt` is used in both text-encoders + num_images_per_prompt (`int`, *optional*, defaults to 1): + The number of images to generate per prompt. + eta (`float`, *optional*, defaults to 0.0): + Corresponds to parameter eta (η) in the DDIM paper: https://arxiv.org/abs/2010.02502. Only applies to + [`schedulers.DDIMScheduler`], will be ignored for others. + generator (`torch.Generator` or `List[torch.Generator]`, *optional*): + One or a list of [torch generator(s)](https://pytorch.org/docs/stable/generated/torch.Generator.html) + to make generation deterministic. + latents (`torch.FloatTensor`, *optional*): + Pre-generated noisy latents, sampled from a Gaussian distribution, to be used as inputs for image + generation. Can be used to tweak the same generation with different prompts. If not provided, a latents + tensor will ge generated by sampling using the supplied random `generator`. + prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + negative_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input + argument. + pooled_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. + If not provided, pooled text embeddings will be generated from `prompt` input argument. + negative_pooled_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated negative pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, pooled negative_prompt_embeds will be generated from `negative_prompt` + input argument. + output_type (`str`, *optional*, defaults to `"pil"`): + The output format of the generate image. Choose between + [PIL](https://pillow.readthedocs.io/en/stable/): `PIL.Image.Image` or `np.array`. + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] instead + of a plain tuple. + cross_attention_kwargs (`dict`, *optional*): + A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under + `self.processor` in + [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). + guidance_rescale (`float`, *optional*, defaults to 0.0): + Guidance rescale factor proposed by [Common Diffusion Noise Schedules and Sample Steps are + Flawed](https://arxiv.org/pdf/2305.08891.pdf) `guidance_scale` is defined as `φ` in equation 16. of + [Common Diffusion Noise Schedules and Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). + Guidance rescale factor should fix overexposure when using zero terminal SNR. + original_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)): + If `original_size` is not the same as `target_size` the image will appear to be down- or upsampled. + `original_size` defaults to `(height, width)` if not specified. Part of SDXL's micro-conditioning as + explained in section 2.2 of + [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). + crops_coords_top_left (`Tuple[int]`, *optional*, defaults to (0, 0)): + `crops_coords_top_left` can be used to generate an image that appears to be "cropped" from the position + `crops_coords_top_left` downwards. Favorable, well-centered images are usually achieved by setting + `crops_coords_top_left` to (0, 0). Part of SDXL's micro-conditioning as explained in section 2.2 of + [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). + target_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)): + For most cases, `target_size` should be set to the desired height and width of the generated image. If + not specified it will default to `(height, width)`. Part of SDXL's micro-conditioning as explained in + section 2.2 of [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). + negative_original_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)): + To negatively condition the generation process based on a specific image resolution. Part of SDXL's + micro-conditioning as explained in section 2.2 of + [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more + information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208. + negative_crops_coords_top_left (`Tuple[int]`, *optional*, defaults to (0, 0)): + To negatively condition the generation process based on a specific crop coordinates. Part of SDXL's + micro-conditioning as explained in section 2.2 of + [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more + information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208. + negative_target_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)): + To negatively condition the generation process based on a target image resolution. It should be as same + as the `target_size` for most cases. Part of SDXL's micro-conditioning as explained in section 2.2 of + [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more + information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208. + callback_on_step_end (`Callable`, *optional*): + A function that calls at the end of each denoising steps during the inference. The function is called + with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int, + callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by + `callback_on_step_end_tensor_inputs`. + callback_on_step_end_tensor_inputs (`List`, *optional*): + The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list + will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the + `._callback_tensor_inputs` attribute of your pipeine class. + + Examples: + + Returns: + [`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] or `tuple`: + [`~pipelines.stable_diffusion_xl.StableDiffusionXLPipelineOutput`] if `return_dict` is True, otherwise a + `tuple`. When returning a tuple, the first element is a list with the generated images. + """ + callback = kwargs.pop("callback", None) + callback_steps = kwargs.pop("callback_steps", None) + + if callback is not None: + deprecate( + "callback", + "1.0.0", + "Passing `callback` as an input argument to `__call__` is deprecated, consider use `callback_on_step_end`", + ) + if callback_steps is not None: + deprecate( + "callback_steps", + "1.0.0", + "Passing `callback_steps` as an input argument to `__call__` is deprecated, consider use `callback_on_step_end`", + ) + + # 0. Default height and width to unet + height = height or self.default_sample_size * self.vae_scale_factor + width = width or self.default_sample_size * self.vae_scale_factor + + original_size = original_size or (height, width) + target_size = target_size or (height, width) + + # 1. Check inputs. Raise error if not correct + self.check_inputs( + prompt, + prompt_2, + height, + width, + callback_steps, + negative_prompt, + negative_prompt_2, + prompt_embeds, + negative_prompt_embeds, + pooled_prompt_embeds, + negative_pooled_prompt_embeds, + callback_on_step_end_tensor_inputs, + ) + + self._guidance_scale = guidance_scale + self._guidance_rescale = guidance_rescale + self._clip_skip = clip_skip + self._cross_attention_kwargs = cross_attention_kwargs + self._denoising_end = denoising_end + + # 2. Define call parameters + if prompt is not None and isinstance(prompt, str): + batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + device = self._execution_device + + # 3. Encode input prompt + lora_scale = ( + self.cross_attention_kwargs.get("scale", None) if self.cross_attention_kwargs is not None else None + ) + + ( + prompt_embeds, + negative_prompt_embeds, + pooled_prompt_embeds, + negative_pooled_prompt_embeds, + ) = self.encode_prompt( + prompt=prompt, + prompt_2=prompt_2, + device=device, + num_images_per_prompt=num_images_per_prompt, + do_classifier_free_guidance=self.do_classifier_free_guidance, + negative_prompt=negative_prompt, + negative_prompt_2=negative_prompt_2, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + pooled_prompt_embeds=pooled_prompt_embeds, + negative_pooled_prompt_embeds=negative_pooled_prompt_embeds, + lora_scale=lora_scale, + clip_skip=self.clip_skip, + ) + + # 4. Prepare timesteps + self.scheduler.set_timesteps(num_inference_steps, device=device) + + timesteps = self.scheduler.timesteps + + # 5. Prepare latent variables + num_channels_latents = self.unet.config.in_channels + latents = self.prepare_latents( + batch_size * num_images_per_prompt, + num_channels_latents, + height, + width, + prompt_embeds.dtype, + device, + generator, + latents, + ) + + # 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) + + if share_queries: + query_store = QueryStore(**query_store_kwargs) + else: + query_store = None + + self.attention_store = AttentionStore(attention_store_kwargs) + register_extended_self_attn(self.unet, self.attention_store, extended_attn_kwargs) + + # 7. Prepare added time ids & embeddings + add_text_embeds = pooled_prompt_embeds + if self.text_encoder_2 is None: + text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1]) + else: + text_encoder_projection_dim = self.text_encoder_2.config.projection_dim + + add_time_ids = self._get_add_time_ids( + original_size, + crops_coords_top_left, + target_size, + dtype=prompt_embeds.dtype, + text_encoder_projection_dim=text_encoder_projection_dim, + ) + if negative_original_size is not None and negative_target_size is not None: + negative_add_time_ids = self._get_add_time_ids( + negative_original_size, + negative_crops_coords_top_left, + negative_target_size, + dtype=prompt_embeds.dtype, + text_encoder_projection_dim=text_encoder_projection_dim, + ) + else: + negative_add_time_ids = add_time_ids + + if self.do_classifier_free_guidance: + prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0) + add_text_embeds = torch.cat([negative_pooled_prompt_embeds, add_text_embeds], dim=0) + add_time_ids = torch.cat([negative_add_time_ids, add_time_ids], dim=0) + + prompt_embeds = prompt_embeds.to(device) + add_text_embeds = add_text_embeds.to(device) + add_time_ids = add_time_ids.to(device).repeat(batch_size * num_images_per_prompt, 1) + + # 8. Denoising loop + num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) + + # 8.1 Apply denoising_end + if ( + self.denoising_end is not None + and isinstance(self.denoising_end, float) + and self.denoising_end > 0 + and self.denoising_end < 1 + ): + discrete_timestep_cutoff = int( + round( + self.scheduler.config.num_train_timesteps + - (self.denoising_end * self.scheduler.config.num_train_timesteps) + ) + ) + num_inference_steps = len(list(filter(lambda ts: ts >= discrete_timestep_cutoff, timesteps))) + timesteps = timesteps[:num_inference_steps] + + # 9. Optionally get Guidance Scale Embedding + timestep_cond = None + if self.unet.config.time_cond_proj_dim is not None: + guidance_scale_tensor = torch.tensor(self.guidance_scale - 1).repeat(batch_size * num_images_per_prompt) + timestep_cond = self.get_guidance_scale_embedding( + guidance_scale_tensor, embedding_dim=self.unet.config.time_cond_proj_dim + ).to(device=device, dtype=latents.dtype) + + self._num_timesteps = len(timesteps) + + if instance_latents is not None: + n_instances = instance_latents.shape[0] + instance_noise = latents[:n_instances].clone() + + with self.progress_bar(total=num_inference_steps) as progress_bar: + for i, t in enumerate(timesteps): + self.attention_store.curr_iter = i + + if instance_latents is not None: + noised_instances = self.scheduler.add_noise(instance_latents, instance_noise, t.repeat(n_instances).long()) + latents[:n_instances] = noised_instances + + # expand the latents if we are doing classifier free guidance + latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents + latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) + + # predict the noise residual + added_cond_kwargs = {"text_embeds": add_text_embeds, "time_ids": add_time_ids} + + if share_queries and (i >= query_store.t_range[0] and i <= query_store.t_range[1]): + query_store.set_mode('cache') + noise_pred_vanilla = self.unet( + latent_model_input, + t, + encoder_hidden_states=prompt_embeds, + timestep_cond=timestep_cond, + cross_attention_kwargs={'query_store': query_store, + 'perform_extend_attn': False, + 'record_attention': False}, + added_cond_kwargs=added_cond_kwargs, + return_dict=False, + )[0] + + query_store.set_mode('inject') + + noise_pred = self.unet( + latent_model_input, + t, + encoder_hidden_states=prompt_embeds, + timestep_cond=timestep_cond, + cross_attention_kwargs={'query_store': query_store, + 'perform_extend_attn': True, + 'record_attention': True, + 'feature_injector': feature_injector, + 'anchors_cache': anchors_cache}, + added_cond_kwargs=added_cond_kwargs, + return_dict=False, + )[0] + + # perform guidance + if self.do_classifier_free_guidance: + noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) + noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond) + + if self.do_classifier_free_guidance and self.guidance_rescale > 0.0: + # Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf + noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=self.guidance_rescale) + + # compute the previous noisy sample x_t -> x_t-1 + latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0] + + if callback_on_step_end is not None: + callback_kwargs = {} + for k in callback_on_step_end_tensor_inputs: + callback_kwargs[k] = locals()[k] + callback_outputs = callback_on_step_end(self, i, t, callback_kwargs) + + latents = callback_outputs.pop("latents", latents) + prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds) + negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds) + add_text_embeds = callback_outputs.pop("add_text_embeds", add_text_embeds) + negative_pooled_prompt_embeds = callback_outputs.pop( + "negative_pooled_prompt_embeds", negative_pooled_prompt_embeds + ) + add_time_ids = callback_outputs.pop("add_time_ids", add_time_ids) + negative_add_time_ids = callback_outputs.pop("negative_add_time_ids", negative_add_time_ids) + + # call the callback, if provided + if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): + progress_bar.update() + if callback is not None and i % callback_steps == 0: + step_idx = i // getattr(self.scheduler, "order", 1) + callback(step_idx, t, latents) + + if XLA_AVAILABLE: + # xm.mark_step() + pass + + # Update attention store mask + self.attention_store.aggregate_last_steps_attention() + + if not output_type == "latent": + # make sure the VAE is in float32 mode, as it overflows in float16 + needs_upcasting = self.vae.dtype == torch.float16 and self.vae.config.force_upcast + + if needs_upcasting: + self.upcast_vae() + latents = latents.to(next(iter(self.vae.post_quant_conv.parameters())).dtype) + + image = self.vae.decode(latents / self.vae.config.scaling_factor, return_dict=False)[0] + + # cast back to fp16 if needed + if needs_upcasting: + self.vae.to(dtype=torch.float16) + else: + image = latents + + if not output_type == "latent": + # apply watermark if available + if self.watermark is not None: + image = self.watermark.apply_watermark(image) + + image = self.image_processor.postprocess(image, output_type=output_type) + + # Offload all models + self.maybe_free_model_hooks() + + if not return_dict: + return (image,) + + return StableDiffusionXLPipelineOutput(images=image) \ No newline at end of file diff --git a/modules/consistory/consistory_run.py b/modules/consistory/consistory_run.py new file mode 100644 index 000000000..c51bf08b9 --- /dev/null +++ b/modules/consistory/consistory_run.py @@ -0,0 +1,260 @@ +# Copyright (C) 2024 NVIDIA Corporation. All rights reserved. +# +# This work is licensed under the LICENSE file +# located at the root directory. + +import torch +from diffusers import DDIMScheduler +from diffusers.utils.torch_utils import randn_tensor +from .consistory_unet_sdxl import ConsistorySDXLUNet2DConditionModel +from .consistory_pipeline import ConsistoryExtendAttnSDXLPipeline +from .consistory_utils import FeatureInjector, AnchorCache +# from .utils.general_utils import * +from .utils.general_utils import gaussian_smooth, cyclic_nn_map, anchor_nn_map + + +LATENT_RESOLUTIONS = [32, 64] + + +def load_pipeline(gpu_id=0): + float_type = torch.float16 + sd_id = "stabilityai/stable-diffusion-xl-base-1.0" + device = torch.device(f'cuda:{gpu_id}') if torch.cuda.is_available() else torch.device('cpu') + unet = ConsistorySDXLUNet2DConditionModel.from_pretrained(sd_id, subfolder="unet", torch_dtype=float_type) + scheduler = DDIMScheduler.from_pretrained(sd_id, subfolder="scheduler") + story_pipeline = ConsistoryExtendAttnSDXLPipeline.from_pretrained(sd_id, unet=unet, torch_dtype=float_type, variant="fp16", use_safetensors=True, scheduler=scheduler).to(device) + story_pipeline.enable_freeu(s1=0.6, s2=0.4, b1=1.1, b2=1.2) + return story_pipeline + + +def create_anchor_mapping(bsz, anchor_indices=[0]): + anchor_mapping = torch.eye(bsz, dtype=torch.bool) + for anchor_idx in anchor_indices: + anchor_mapping[:, anchor_idx] = True + return anchor_mapping + + +def create_token_indices(prompts, batch_size, concept_token, tokenizer): + if isinstance(concept_token, str): + concept_token = [concept_token] + concept_token_id = [tokenizer.encode(x, add_special_tokens=False)[0] for x in concept_token] + tokens = tokenizer.batch_encode_plus(prompts, padding=True, return_tensors='pt')['input_ids'] + token_indices = torch.full((len(concept_token), batch_size), -1, dtype=torch.int64) + for i, token_id in enumerate(concept_token_id): + batch_loc, token_loc = torch.where(tokens == token_id) + token_indices[i, batch_loc] = token_loc + return token_indices + + +def create_latents(story_pipeline, seed, batch_size, same_latent, device, float_type): + # if seed is int + if isinstance(seed, int): + g = torch.Generator('cuda').manual_seed(seed) + shape = (batch_size, story_pipeline.unet.config.in_channels, 128, 128) + latents = randn_tensor(shape, generator=g, device=device, dtype=float_type) + elif isinstance(seed, list): + shape = (batch_size, story_pipeline.unet.config.in_channels, 128, 128) + latents = torch.empty(shape, device=device, dtype=float_type) + for i, seed_i in enumerate(seed): + g = torch.Generator('cuda').manual_seed(seed_i) + curr_latent = randn_tensor(shape, generator=g, device=device, dtype=float_type) + latents[i] = curr_latent[i] + if same_latent: + latents = latents[:1].repeat(batch_size, 1, 1, 1) + return latents, g + + +# Batch inference +def run_batch_generation(story_pipeline, prompts, concept_token, + seed=40, n_steps=50, mask_dropout=0.5, + same_latent=False, share_queries=True, + perform_sdsa=True, perform_injection=True, + inject_range_alpha=(10,20,0.8), + n_achors=2): + device = story_pipeline.device + tokenizer = story_pipeline.tokenizer + float_type = story_pipeline.dtype + unet = story_pipeline.unet + batch_size = len(prompts) + token_indices = create_token_indices(prompts, batch_size, concept_token, tokenizer) + anchor_mappings = create_anchor_mapping(batch_size, anchor_indices=list(range(n_achors))) + default_attention_store_kwargs = { + 'token_indices': token_indices, + 'mask_dropout': mask_dropout, + 'extended_mapping': anchor_mappings + } + default_extended_attn_kwargs = {'extend_kv_unet_parts': ['up']} + query_store_kwargs= {'t_range': [0,n_steps//10], 'strength_start': 0.9, 'strength_end': 0.81836735} + latents, g = create_latents(story_pipeline, seed, batch_size, same_latent, device, float_type) + + # ------------------ # + # Extended attention First Run # + if perform_sdsa: + extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': [(1, n_steps)]} + else: + extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': []} + out = story_pipeline(prompt=prompts, generator=g, latents=latents, + attention_store_kwargs=default_attention_store_kwargs, + extended_attn_kwargs=extended_attn_kwargs, + share_queries=share_queries, + query_store_kwargs=query_store_kwargs, + num_inference_steps=n_steps) + last_masks = story_pipeline.attention_store.last_mask + dift_features = unet.latent_store.dift_features['261_0'][batch_size:] + dift_features = torch.stack([gaussian_smooth(x, kernel_size=3, sigma=1) for x in dift_features], dim=0) + nn_map, nn_distances = cyclic_nn_map(dift_features, last_masks, LATENT_RESOLUTIONS, device) + + # ------------------ # + # Extended attention with nn_map # + if perform_injection: + feature_injector = FeatureInjector( + nn_map, + nn_distances, + last_masks, + inject_range_alpha=[inject_range_alpha], + swap_strategy='min', inject_unet_parts=['up', 'down'], dist_thr='dynamic') + out = story_pipeline(prompt=prompts, generator=g, latents=latents, + attention_store_kwargs=default_attention_store_kwargs, + extended_attn_kwargs=extended_attn_kwargs, + share_queries=share_queries, + query_store_kwargs=query_store_kwargs, + feature_injector=feature_injector, + num_inference_steps=n_steps) + # display_attn_maps(story_pipeline.attention_store.last_mask, out.images) + return out.images + + +# Anchors +def run_anchor_generation(story_pipeline, prompts, concept_token, + seed=40, n_steps=50, mask_dropout=0.5, + inject_range_alpha=(10,20,0.8), + same_latent=False, share_queries=True, + perform_sdsa=True, perform_injection=True): + device = story_pipeline.device + tokenizer = story_pipeline.tokenizer + float_type = story_pipeline.dtype + unet = story_pipeline.unet + batch_size = len(prompts) + token_indices = create_token_indices(prompts, batch_size, concept_token, tokenizer) + default_attention_store_kwargs = { + 'token_indices': token_indices, + 'mask_dropout': mask_dropout + } + default_extended_attn_kwargs = {'extend_kv_unet_parts': ['up']} + query_store_kwargs={'t_range': [0,n_steps//10], 'strength_start': 0.9, 'strength_end': 0.81836735} + latents, g = create_latents(story_pipeline, seed, batch_size, same_latent, device, float_type) + anchor_cache_first_stage = AnchorCache() + anchor_cache_second_stage = AnchorCache() + + # ------------------ # + # Extended attention First Run # + if perform_sdsa: + extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': [(1, n_steps)]} + else: + extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': []} + out = story_pipeline(prompt=prompts, generator=g, latents=latents, + attention_store_kwargs=default_attention_store_kwargs, + extended_attn_kwargs=extended_attn_kwargs, + share_queries=share_queries, + query_store_kwargs=query_store_kwargs, + anchors_cache=anchor_cache_first_stage, + num_inference_steps=n_steps) + last_masks = story_pipeline.attention_store.last_mask + dift_features = unet.latent_store.dift_features['261_0'][batch_size:] + dift_features = torch.stack([gaussian_smooth(x, kernel_size=3, sigma=1) for x in dift_features], dim=0) + anchor_cache_first_stage.dift_cache = dift_features + anchor_cache_first_stage.anchors_last_mask = last_masks + nn_map, nn_distances = cyclic_nn_map(dift_features, last_masks, LATENT_RESOLUTIONS, device) + + # ------------------ # + # Extended attention with nn_map # + if perform_injection: + feature_injector = FeatureInjector( + nn_map, + nn_distances, + last_masks, + inject_range_alpha=[inject_range_alpha], + swap_strategy='min', + inject_unet_parts=['up', 'down'], + dist_thr='dynamic') + out = story_pipeline(prompt=prompts, generator=g, latents=latents, + attention_store_kwargs=default_attention_store_kwargs, + extended_attn_kwargs=extended_attn_kwargs, + share_queries=share_queries, + query_store_kwargs=query_store_kwargs, + feature_injector=feature_injector, + anchors_cache=anchor_cache_second_stage, + num_inference_steps=n_steps) + # display_attn_maps(story_pipeline.attention_store.last_mask, out.images) + anchor_cache_second_stage.dift_cache = dift_features + anchor_cache_second_stage.anchors_last_mask = last_masks + return out.images, anchor_cache_first_stage, anchor_cache_second_stage + + +def run_extra_generation(story_pipeline, prompts, concept_token, + anchor_cache_first_stage, anchor_cache_second_stage, + seed=40, n_steps=50, mask_dropout=0.5, + inject_range_alpha=(10,20,0.8), + same_latent=False, share_queries=True, + perform_sdsa=True, perform_injection=True): + device = story_pipeline.device + tokenizer = story_pipeline.tokenizer + float_type = story_pipeline.dtype + unet = story_pipeline.unet + batch_size = len(prompts) + token_indices = create_token_indices(prompts, batch_size, concept_token, tokenizer) + default_attention_store_kwargs = { + 'token_indices': token_indices, + 'mask_dropout': mask_dropout + } + default_extended_attn_kwargs = {'extend_kv_unet_parts': ['up']} + query_store_kwargs={'t_range': [0,n_steps//10], 'strength_start': 0.9, 'strength_end': 0.81836735} + extra_batch_size = batch_size + 2 + if isinstance(seed, list): + seed = [seed[0], seed[0], *seed] + latents, g = create_latents(story_pipeline, seed, extra_batch_size, same_latent, device, float_type) + latents = latents[2:] + anchor_cache_first_stage.set_mode_inject() + anchor_cache_second_stage.set_mode_inject() + + # ------------------ # + # Extended attention First Run # + if perform_sdsa: + extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': [(1, n_steps)]} + else: + extended_attn_kwargs = {**default_extended_attn_kwargs, 't_range': []} + out = story_pipeline(prompt=prompts, generator=g, latents=latents, + attention_store_kwargs=default_attention_store_kwargs, + extended_attn_kwargs=extended_attn_kwargs, + share_queries=share_queries, + query_store_kwargs=query_store_kwargs, + anchors_cache=anchor_cache_first_stage, + num_inference_steps=n_steps) + last_masks = story_pipeline.attention_store.last_mask + dift_features = unet.latent_store.dift_features['261_0'][batch_size:] + dift_features = torch.stack([gaussian_smooth(x, kernel_size=3, sigma=1) for x in dift_features], dim=0) + anchor_dift_features = anchor_cache_first_stage.dift_cache + anchor_last_masks = anchor_cache_first_stage.anchors_last_mask + nn_map, nn_distances = anchor_nn_map(dift_features, anchor_dift_features, last_masks, anchor_last_masks, LATENT_RESOLUTIONS, device) + + # ------------------ # + # Extended attention with nn_map # + if perform_injection: + feature_injector = FeatureInjector( + nn_map, + nn_distances, + last_masks, + inject_range_alpha=[inject_range_alpha], + swap_strategy='min', + inject_unet_parts=['up', 'down'], + dist_thr='dynamic') + out = story_pipeline(prompt=prompts, generator=g, latents=latents, + attention_store_kwargs=default_attention_store_kwargs, + extended_attn_kwargs=extended_attn_kwargs, + share_queries=share_queries, + query_store_kwargs=query_store_kwargs, + feature_injector=feature_injector, + anchors_cache=anchor_cache_second_stage, + num_inference_steps=n_steps) + # display_attn_maps(story_pipeline.attention_store.last_mask, out.images) + return out.images diff --git a/modules/consistory/consistory_unet_sdxl.py b/modules/consistory/consistory_unet_sdxl.py new file mode 100644 index 000000000..4dd9b42d2 --- /dev/null +++ b/modules/consistory/consistory_unet_sdxl.py @@ -0,0 +1,1173 @@ +# Copyright 2023 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# Not a contribution +# Changes made by NVIDIA CORPORATION & AFFILIATES enabling ConsiStory or otherwise documented as NVIDIA-proprietary +# are not a contribution and subject to the license under the LICENSE file located at the root directory. + + +from dataclasses import dataclass +from typing import Any, Dict, List, Optional, Tuple, Union + +import torch +import torch.nn as nn +import torch.utils.checkpoint + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.loaders import UNet2DConditionLoadersMixin +from diffusers.utils import USE_PEFT_BACKEND, BaseOutput, deprecate, logging, scale_lora_layers, unscale_lora_layers +from diffusers.models.activations import get_activation +from diffusers.models.attention_processor import ( + ADDED_KV_ATTENTION_PROCESSORS, + CROSS_ATTENTION_PROCESSORS, + AttentionProcessor, + AttnAddedKVProcessor, + AttnProcessor, +) +from diffusers.models.embeddings import ( + GaussianFourierProjection, + ImageHintTimeEmbedding, + ImageProjection, + ImageTimeEmbedding, + PositionNet, + TextImageProjection, + TextImageTimeEmbedding, + TextTimeEmbedding, + TimestepEmbedding, + Timesteps, +) +from diffusers.models.modeling_utils import ModelMixin +from diffusers.models.unets.unet_2d_blocks import ( + UNetMidBlock2D, + UNetMidBlock2DCrossAttn, + UNetMidBlock2DSimpleCrossAttn, + get_down_block, + get_up_block, +) + +from .consistory_utils import DIFTLatentStore + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +@dataclass +class UNet2DConditionOutput(BaseOutput): + """ + The output of [`UNet2DConditionModel`]. + + Args: + sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)`): + The hidden states output conditioned on `encoder_hidden_states` input. Output of last layer of model. + """ + + sample: torch.FloatTensor = None + + +class ConsistorySDXLUNet2DConditionModel(ModelMixin, ConfigMixin, UNet2DConditionLoadersMixin): + r""" + A conditional 2D UNet model that takes a noisy sample, conditional state, and a timestep and returns a sample + shaped output. + + This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented + for all models (such as downloading or saving). + + Parameters: + sample_size (`int` or `Tuple[int, int]`, *optional*, defaults to `None`): + Height and width of input/output sample. + in_channels (`int`, *optional*, defaults to 4): Number of channels in the input sample. + out_channels (`int`, *optional*, defaults to 4): Number of channels in the output. + center_input_sample (`bool`, *optional*, defaults to `False`): Whether to center the input sample. + flip_sin_to_cos (`bool`, *optional*, defaults to `False`): + Whether to flip the sin to cos in the time embedding. + freq_shift (`int`, *optional*, defaults to 0): The frequency shift to apply to the time embedding. + down_block_types (`Tuple[str]`, *optional*, defaults to `("CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "DownBlock2D")`): + The tuple of downsample blocks to use. + mid_block_type (`str`, *optional*, defaults to `"UNetMidBlock2DCrossAttn"`): + Block type for middle of UNet, it can be one of `UNetMidBlock2DCrossAttn`, `UNetMidBlock2D`, or + `UNetMidBlock2DSimpleCrossAttn`. If `None`, the mid block layer is skipped. + up_block_types (`Tuple[str]`, *optional*, defaults to `("UpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D")`): + The tuple of upsample blocks to use. + only_cross_attention(`bool` or `Tuple[bool]`, *optional*, default to `False`): + Whether to include self-attention in the basic transformer blocks, see + [`~models.attention.BasicTransformerBlock`]. + block_out_channels (`Tuple[int]`, *optional*, defaults to `(320, 640, 1280, 1280)`): + The tuple of output channels for each block. + layers_per_block (`int`, *optional*, defaults to 2): The number of layers per block. + downsample_padding (`int`, *optional*, defaults to 1): The padding to use for the downsampling convolution. + mid_block_scale_factor (`float`, *optional*, defaults to 1.0): The scale factor to use for the mid block. + dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use. + act_fn (`str`, *optional*, defaults to `"silu"`): The activation function to use. + norm_num_groups (`int`, *optional*, defaults to 32): The number of groups to use for the normalization. + If `None`, normalization and activation layers is skipped in post-processing. + norm_eps (`float`, *optional*, defaults to 1e-5): The epsilon to use for the normalization. + cross_attention_dim (`int` or `Tuple[int]`, *optional*, defaults to 1280): + The dimension of the cross attention features. + transformer_layers_per_block (`int`, `Tuple[int]`, or `Tuple[Tuple]` , *optional*, defaults to 1): + The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`]. Only relevant for + [`~models.unet_2d_blocks.CrossAttnDownBlock2D`], [`~models.unet_2d_blocks.CrossAttnUpBlock2D`], + [`~models.unet_2d_blocks.UNetMidBlock2DCrossAttn`]. + reverse_transformer_layers_per_block : (`Tuple[Tuple]`, *optional*, defaults to None): + The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`], in the upsampling + blocks of the U-Net. Only relevant if `transformer_layers_per_block` is of type `Tuple[Tuple]` and for + [`~models.unet_2d_blocks.CrossAttnDownBlock2D`], [`~models.unet_2d_blocks.CrossAttnUpBlock2D`], + [`~models.unet_2d_blocks.UNetMidBlock2DCrossAttn`]. + encoder_hid_dim (`int`, *optional*, defaults to None): + If `encoder_hid_dim_type` is defined, `encoder_hidden_states` will be projected from `encoder_hid_dim` + dimension to `cross_attention_dim`. + encoder_hid_dim_type (`str`, *optional*, defaults to `None`): + If given, the `encoder_hidden_states` and potentially other embeddings are down-projected to text + embeddings of dimension `cross_attention` according to `encoder_hid_dim_type`. + attention_head_dim (`int`, *optional*, defaults to 8): The dimension of the attention heads. + num_attention_heads (`int`, *optional*): + The number of attention heads. If not defined, defaults to `attention_head_dim` + resnet_time_scale_shift (`str`, *optional*, defaults to `"default"`): Time scale shift config + for ResNet blocks (see [`~models.resnet.ResnetBlock2D`]). Choose from `default` or `scale_shift`. + class_embed_type (`str`, *optional*, defaults to `None`): + The type of class embedding to use which is ultimately summed with the time embeddings. Choose from `None`, + `"timestep"`, `"identity"`, `"projection"`, or `"simple_projection"`. + addition_embed_type (`str`, *optional*, defaults to `None`): + Configures an optional embedding which will be summed with the time embeddings. Choose from `None` or + "text". "text" will use the `TextTimeEmbedding` layer. + addition_time_embed_dim: (`int`, *optional*, defaults to `None`): + Dimension for the timestep embeddings. + num_class_embeds (`int`, *optional*, defaults to `None`): + Input dimension of the learnable embedding matrix to be projected to `time_embed_dim`, when performing + class conditioning with `class_embed_type` equal to `None`. + time_embedding_type (`str`, *optional*, defaults to `positional`): + The type of position embedding to use for timesteps. Choose from `positional` or `fourier`. + time_embedding_dim (`int`, *optional*, defaults to `None`): + An optional override for the dimension of the projected time embedding. + time_embedding_act_fn (`str`, *optional*, defaults to `None`): + Optional activation function to use only once on the time embeddings before they are passed to the rest of + the UNet. Choose from `silu`, `mish`, `gelu`, and `swish`. + timestep_post_act (`str`, *optional*, defaults to `None`): + The second activation function to use in timestep embedding. Choose from `silu`, `mish` and `gelu`. + time_cond_proj_dim (`int`, *optional*, defaults to `None`): + The dimension of `cond_proj` layer in the timestep embedding. + conv_in_kernel (`int`, *optional*, default to `3`): The kernel size of `conv_in` layer. conv_out_kernel (`int`, + *optional*, default to `3`): The kernel size of `conv_out` layer. projection_class_embeddings_input_dim (`int`, + *optional*): The dimension of the `class_labels` input when + `class_embed_type="projection"`. Required when `class_embed_type="projection"`. + class_embeddings_concat (`bool`, *optional*, defaults to `False`): Whether to concatenate the time + embeddings with the class embeddings. + mid_block_only_cross_attention (`bool`, *optional*, defaults to `None`): + Whether to use cross attention with the mid block when using the `UNetMidBlock2DSimpleCrossAttn`. If + `only_cross_attention` is given as a single boolean and `mid_block_only_cross_attention` is `None`, the + `only_cross_attention` value is used as the value for `mid_block_only_cross_attention`. Default to `False` + otherwise. + """ + + _supports_gradient_checkpointing = True + + @register_to_config + def __init__( + self, + sample_size: Optional[int] = None, + in_channels: int = 4, + out_channels: int = 4, + center_input_sample: bool = False, + flip_sin_to_cos: bool = True, + freq_shift: int = 0, + down_block_types: Tuple[str] = ( + "CrossAttnDownBlock2D", + "CrossAttnDownBlock2D", + "CrossAttnDownBlock2D", + "DownBlock2D", + ), + mid_block_type: Optional[str] = "UNetMidBlock2DCrossAttn", + up_block_types: Tuple[str] = ("UpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D", "CrossAttnUpBlock2D"), + only_cross_attention: Union[bool, Tuple[bool]] = False, + block_out_channels: Tuple[int] = (320, 640, 1280, 1280), + layers_per_block: Union[int, Tuple[int]] = 2, + downsample_padding: int = 1, + mid_block_scale_factor: float = 1, + dropout: float = 0.0, + act_fn: str = "silu", + norm_num_groups: Optional[int] = 32, + norm_eps: float = 1e-5, + cross_attention_dim: Union[int, Tuple[int]] = 1280, + transformer_layers_per_block: Union[int, Tuple[int], Tuple[Tuple]] = 1, + reverse_transformer_layers_per_block: Optional[Tuple[Tuple[int]]] = None, + encoder_hid_dim: Optional[int] = None, + encoder_hid_dim_type: Optional[str] = None, + attention_head_dim: Union[int, Tuple[int]] = 8, + num_attention_heads: Optional[Union[int, Tuple[int]]] = None, + dual_cross_attention: bool = False, + use_linear_projection: bool = False, + class_embed_type: Optional[str] = None, + addition_embed_type: Optional[str] = None, + addition_time_embed_dim: Optional[int] = None, + num_class_embeds: Optional[int] = None, + upcast_attention: bool = False, + resnet_time_scale_shift: str = "default", + resnet_skip_time_act: bool = False, + resnet_out_scale_factor: int = 1.0, + time_embedding_type: str = "positional", + time_embedding_dim: Optional[int] = None, + time_embedding_act_fn: Optional[str] = None, + timestep_post_act: Optional[str] = None, + time_cond_proj_dim: Optional[int] = None, + conv_in_kernel: int = 3, + conv_out_kernel: int = 3, + projection_class_embeddings_input_dim: Optional[int] = None, + attention_type: str = "default", + class_embeddings_concat: bool = False, + mid_block_only_cross_attention: Optional[bool] = None, + cross_attention_norm: Optional[str] = None, + addition_embed_type_num_heads=64, + ): + super().__init__() + + self.latent_store = DIFTLatentStore(steps=[261], up_ft_indices=[0]) + self.sample_size = sample_size + + if num_attention_heads is not None: + raise ValueError( + "At the moment it is not possible to define the number of attention heads via `num_attention_heads` because of a naming issue as described in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131. Passing `num_attention_heads` will only be supported in diffusers v0.19." + ) + + # If `num_attention_heads` is not defined (which is the case for most models) + # it will default to `attention_head_dim`. This looks weird upon first reading it and it is. + # The reason for this behavior is to correct for incorrectly named variables that were introduced + # when this library was created. The incorrect naming was only discovered much later in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131 + # Changing `attention_head_dim` to `num_attention_heads` for 40,000+ configurations is too backwards breaking + # which is why we correct for the naming here. + num_attention_heads = num_attention_heads or attention_head_dim + + # Check inputs + if len(down_block_types) != len(up_block_types): + raise ValueError( + f"Must provide the same number of `down_block_types` as `up_block_types`. `down_block_types`: {down_block_types}. `up_block_types`: {up_block_types}." + ) + + if len(block_out_channels) != len(down_block_types): + raise ValueError( + f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." + ) + + if not isinstance(only_cross_attention, bool) and len(only_cross_attention) != len(down_block_types): + raise ValueError( + f"Must provide the same number of `only_cross_attention` as `down_block_types`. `only_cross_attention`: {only_cross_attention}. `down_block_types`: {down_block_types}." + ) + + if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types): + raise ValueError( + f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}." + ) + + if not isinstance(attention_head_dim, int) and len(attention_head_dim) != len(down_block_types): + raise ValueError( + f"Must provide the same number of `attention_head_dim` as `down_block_types`. `attention_head_dim`: {attention_head_dim}. `down_block_types`: {down_block_types}." + ) + + if isinstance(cross_attention_dim, list) and len(cross_attention_dim) != len(down_block_types): + raise ValueError( + f"Must provide the same number of `cross_attention_dim` as `down_block_types`. `cross_attention_dim`: {cross_attention_dim}. `down_block_types`: {down_block_types}." + ) + + if not isinstance(layers_per_block, int) and len(layers_per_block) != len(down_block_types): + raise ValueError( + f"Must provide the same number of `layers_per_block` as `down_block_types`. `layers_per_block`: {layers_per_block}. `down_block_types`: {down_block_types}." + ) + if isinstance(transformer_layers_per_block, list) and reverse_transformer_layers_per_block is None: + for layer_number_per_block in transformer_layers_per_block: + if isinstance(layer_number_per_block, list): + raise ValueError("Must provide 'reverse_transformer_layers_per_block` if using asymmetrical UNet.") + + # input + conv_in_padding = (conv_in_kernel - 1) // 2 + self.conv_in = nn.Conv2d( + in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding + ) + + # time + if time_embedding_type == "fourier": + time_embed_dim = time_embedding_dim or block_out_channels[0] * 2 + if time_embed_dim % 2 != 0: + raise ValueError(f"`time_embed_dim` should be divisible by 2, but is {time_embed_dim}.") + self.time_proj = GaussianFourierProjection( + time_embed_dim // 2, set_W_to_weight=False, log=False, flip_sin_to_cos=flip_sin_to_cos + ) + timestep_input_dim = time_embed_dim + elif time_embedding_type == "positional": + time_embed_dim = time_embedding_dim or block_out_channels[0] * 4 + + self.time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos, freq_shift) + timestep_input_dim = block_out_channels[0] + else: + raise ValueError( + f"{time_embedding_type} does not exist. Please make sure to use one of `fourier` or `positional`." + ) + + self.time_embedding = TimestepEmbedding( + timestep_input_dim, + time_embed_dim, + act_fn=act_fn, + post_act_fn=timestep_post_act, + cond_proj_dim=time_cond_proj_dim, + ) + + if encoder_hid_dim_type is None and encoder_hid_dim is not None: + encoder_hid_dim_type = "text_proj" + self.register_to_config(encoder_hid_dim_type=encoder_hid_dim_type) + logger.info("encoder_hid_dim_type defaults to 'text_proj' as `encoder_hid_dim` is defined.") + + if encoder_hid_dim is None and encoder_hid_dim_type is not None: + raise ValueError( + f"`encoder_hid_dim` has to be defined when `encoder_hid_dim_type` is set to {encoder_hid_dim_type}." + ) + + if encoder_hid_dim_type == "text_proj": + self.encoder_hid_proj = nn.Linear(encoder_hid_dim, cross_attention_dim) + elif encoder_hid_dim_type == "text_image_proj": + # image_embed_dim DOESN'T have to be `cross_attention_dim`. To not clutter the __init__ too much + # they are set to `cross_attention_dim` here as this is exactly the required dimension for the currently only use + # case when `addition_embed_type == "text_image_proj"` (Kadinsky 2.1)` + self.encoder_hid_proj = TextImageProjection( + text_embed_dim=encoder_hid_dim, + image_embed_dim=cross_attention_dim, + cross_attention_dim=cross_attention_dim, + ) + elif encoder_hid_dim_type == "image_proj": + # Kandinsky 2.2 + self.encoder_hid_proj = ImageProjection( + image_embed_dim=encoder_hid_dim, + cross_attention_dim=cross_attention_dim, + ) + elif encoder_hid_dim_type is not None: + raise ValueError( + f"encoder_hid_dim_type: {encoder_hid_dim_type} must be None, 'text_proj' or 'text_image_proj'." + ) + else: + self.encoder_hid_proj = None + + # class embedding + if class_embed_type is None and num_class_embeds is not None: + self.class_embedding = nn.Embedding(num_class_embeds, time_embed_dim) + elif class_embed_type == "timestep": + self.class_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim, act_fn=act_fn) + elif class_embed_type == "identity": + self.class_embedding = nn.Identity(time_embed_dim, time_embed_dim) + elif class_embed_type == "projection": + if projection_class_embeddings_input_dim is None: + raise ValueError( + "`class_embed_type`: 'projection' requires `projection_class_embeddings_input_dim` be set" + ) + # The projection `class_embed_type` is the same as the timestep `class_embed_type` except + # 1. the `class_labels` inputs are not first converted to sinusoidal embeddings + # 2. it projects from an arbitrary input dimension. + # + # Note that `TimestepEmbedding` is quite general, being mainly linear layers and activations. + # When used for embedding actual timesteps, the timesteps are first converted to sinusoidal embeddings. + # As a result, `TimestepEmbedding` can be passed arbitrary vectors. + self.class_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) + elif class_embed_type == "simple_projection": + if projection_class_embeddings_input_dim is None: + raise ValueError( + "`class_embed_type`: 'simple_projection' requires `projection_class_embeddings_input_dim` be set" + ) + self.class_embedding = nn.Linear(projection_class_embeddings_input_dim, time_embed_dim) + else: + self.class_embedding = None + + if addition_embed_type == "text": + if encoder_hid_dim is not None: + text_time_embedding_from_dim = encoder_hid_dim + else: + text_time_embedding_from_dim = cross_attention_dim + + self.add_embedding = TextTimeEmbedding( + text_time_embedding_from_dim, time_embed_dim, num_heads=addition_embed_type_num_heads + ) + elif addition_embed_type == "text_image": + # text_embed_dim and image_embed_dim DON'T have to be `cross_attention_dim`. To not clutter the __init__ too much + # they are set to `cross_attention_dim` here as this is exactly the required dimension for the currently only use + # case when `addition_embed_type == "text_image"` (Kadinsky 2.1)` + self.add_embedding = TextImageTimeEmbedding( + text_embed_dim=cross_attention_dim, image_embed_dim=cross_attention_dim, time_embed_dim=time_embed_dim + ) + elif addition_embed_type == "text_time": + self.add_time_proj = Timesteps(addition_time_embed_dim, flip_sin_to_cos, freq_shift) + self.add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) + elif addition_embed_type == "image": + # Kandinsky 2.2 + self.add_embedding = ImageTimeEmbedding(image_embed_dim=encoder_hid_dim, time_embed_dim=time_embed_dim) + elif addition_embed_type == "image_hint": + # Kandinsky 2.2 ControlNet + self.add_embedding = ImageHintTimeEmbedding(image_embed_dim=encoder_hid_dim, time_embed_dim=time_embed_dim) + elif addition_embed_type is not None: + raise ValueError(f"addition_embed_type: {addition_embed_type} must be None, 'text' or 'text_image'.") + + if time_embedding_act_fn is None: + self.time_embed_act = None + else: + self.time_embed_act = get_activation(time_embedding_act_fn) + + self.down_blocks = nn.ModuleList([]) + self.up_blocks = nn.ModuleList([]) + + if isinstance(only_cross_attention, bool): + if mid_block_only_cross_attention is None: + mid_block_only_cross_attention = only_cross_attention + + only_cross_attention = [only_cross_attention] * len(down_block_types) + + if mid_block_only_cross_attention is None: + mid_block_only_cross_attention = False + + if isinstance(num_attention_heads, int): + num_attention_heads = (num_attention_heads,) * len(down_block_types) + + if isinstance(attention_head_dim, int): + attention_head_dim = (attention_head_dim,) * len(down_block_types) + + if isinstance(cross_attention_dim, int): + cross_attention_dim = (cross_attention_dim,) * len(down_block_types) + + if isinstance(layers_per_block, int): + layers_per_block = [layers_per_block] * len(down_block_types) + + if isinstance(transformer_layers_per_block, int): + transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types) + + if class_embeddings_concat: + # The time embeddings are concatenated with the class embeddings. The dimension of the + # time embeddings passed to the down, middle, and up blocks is twice the dimension of the + # regular time embeddings + blocks_time_embed_dim = time_embed_dim * 2 + else: + blocks_time_embed_dim = time_embed_dim + + # down + output_channel = block_out_channels[0] + for i, down_block_type in enumerate(down_block_types): + input_channel = output_channel + output_channel = block_out_channels[i] + is_final_block = i == len(block_out_channels) - 1 + + down_block = get_down_block( + down_block_type, + num_layers=layers_per_block[i], + transformer_layers_per_block=transformer_layers_per_block[i], + in_channels=input_channel, + out_channels=output_channel, + temb_channels=blocks_time_embed_dim, + add_downsample=not is_final_block, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + resnet_groups=norm_num_groups, + cross_attention_dim=cross_attention_dim[i], + num_attention_heads=num_attention_heads[i], + downsample_padding=downsample_padding, + dual_cross_attention=dual_cross_attention, + use_linear_projection=use_linear_projection, + only_cross_attention=only_cross_attention[i], + upcast_attention=upcast_attention, + resnet_time_scale_shift=resnet_time_scale_shift, + attention_type=attention_type, + resnet_skip_time_act=resnet_skip_time_act, + resnet_out_scale_factor=resnet_out_scale_factor, + cross_attention_norm=cross_attention_norm, + attention_head_dim=attention_head_dim[i] if attention_head_dim[i] is not None else output_channel, + dropout=dropout, + ) + self.down_blocks.append(down_block) + + # mid + if mid_block_type == "UNetMidBlock2DCrossAttn": + self.mid_block = UNetMidBlock2DCrossAttn( + transformer_layers_per_block=transformer_layers_per_block[-1], + in_channels=block_out_channels[-1], + temb_channels=blocks_time_embed_dim, + dropout=dropout, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + output_scale_factor=mid_block_scale_factor, + resnet_time_scale_shift=resnet_time_scale_shift, + cross_attention_dim=cross_attention_dim[-1], + num_attention_heads=num_attention_heads[-1], + resnet_groups=norm_num_groups, + dual_cross_attention=dual_cross_attention, + use_linear_projection=use_linear_projection, + upcast_attention=upcast_attention, + attention_type=attention_type, + ) + elif mid_block_type == "UNetMidBlock2DSimpleCrossAttn": + self.mid_block = UNetMidBlock2DSimpleCrossAttn( + in_channels=block_out_channels[-1], + temb_channels=blocks_time_embed_dim, + dropout=dropout, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + output_scale_factor=mid_block_scale_factor, + cross_attention_dim=cross_attention_dim[-1], + attention_head_dim=attention_head_dim[-1], + resnet_groups=norm_num_groups, + resnet_time_scale_shift=resnet_time_scale_shift, + skip_time_act=resnet_skip_time_act, + only_cross_attention=mid_block_only_cross_attention, + cross_attention_norm=cross_attention_norm, + ) + elif mid_block_type == "UNetMidBlock2D": + self.mid_block = UNetMidBlock2D( + in_channels=block_out_channels[-1], + temb_channels=blocks_time_embed_dim, + dropout=dropout, + num_layers=0, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + output_scale_factor=mid_block_scale_factor, + resnet_groups=norm_num_groups, + resnet_time_scale_shift=resnet_time_scale_shift, + add_attention=False, + ) + elif mid_block_type is None: + self.mid_block = None + else: + raise ValueError(f"unknown mid_block_type : {mid_block_type}") + + # count how many layers upsample the images + self.num_upsamplers = 0 + + # up + reversed_block_out_channels = list(reversed(block_out_channels)) + reversed_num_attention_heads = list(reversed(num_attention_heads)) + reversed_layers_per_block = list(reversed(layers_per_block)) + reversed_cross_attention_dim = list(reversed(cross_attention_dim)) + reversed_transformer_layers_per_block = ( + list(reversed(transformer_layers_per_block)) + if reverse_transformer_layers_per_block is None + else reverse_transformer_layers_per_block + ) + only_cross_attention = list(reversed(only_cross_attention)) + + output_channel = reversed_block_out_channels[0] + for i, up_block_type in enumerate(up_block_types): + is_final_block = i == len(block_out_channels) - 1 + + prev_output_channel = output_channel + output_channel = reversed_block_out_channels[i] + input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)] + + # add upsample block for all BUT final layer + if not is_final_block: + add_upsample = True + self.num_upsamplers += 1 + else: + add_upsample = False + + up_block = get_up_block( + up_block_type, + num_layers=reversed_layers_per_block[i] + 1, + transformer_layers_per_block=reversed_transformer_layers_per_block[i], + in_channels=input_channel, + out_channels=output_channel, + prev_output_channel=prev_output_channel, + temb_channels=blocks_time_embed_dim, + add_upsample=add_upsample, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + resolution_idx=i, + resnet_groups=norm_num_groups, + cross_attention_dim=reversed_cross_attention_dim[i], + num_attention_heads=reversed_num_attention_heads[i], + dual_cross_attention=dual_cross_attention, + use_linear_projection=use_linear_projection, + only_cross_attention=only_cross_attention[i], + upcast_attention=upcast_attention, + resnet_time_scale_shift=resnet_time_scale_shift, + attention_type=attention_type, + resnet_skip_time_act=resnet_skip_time_act, + resnet_out_scale_factor=resnet_out_scale_factor, + cross_attention_norm=cross_attention_norm, + attention_head_dim=attention_head_dim[i] if attention_head_dim[i] is not None else output_channel, + dropout=dropout, + ) + self.up_blocks.append(up_block) + prev_output_channel = output_channel + + # out + if norm_num_groups is not None: + self.conv_norm_out = nn.GroupNorm( + num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=norm_eps + ) + + self.conv_act = get_activation(act_fn) + + else: + self.conv_norm_out = None + self.conv_act = None + + conv_out_padding = (conv_out_kernel - 1) // 2 + self.conv_out = nn.Conv2d( + block_out_channels[0], out_channels, kernel_size=conv_out_kernel, padding=conv_out_padding + ) + + if attention_type in ["gated", "gated-text-image"]: + positive_len = 768 + if isinstance(cross_attention_dim, int): + positive_len = cross_attention_dim + elif isinstance(cross_attention_dim, tuple) or isinstance(cross_attention_dim, list): + positive_len = cross_attention_dim[0] + + feature_type = "text-only" if attention_type == "gated" else "text-image" + self.position_net = PositionNet( + positive_len=positive_len, out_dim=cross_attention_dim, feature_type=feature_type + ) + + @property + def attn_processors(self) -> Dict[str, AttentionProcessor]: + r""" + Returns: + `dict` of attention processors: A dictionary containing all attention processors used in the model with + indexed by its weight name. + """ + # set recursively + processors = {} + + def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: Dict[str, AttentionProcessor]): + if hasattr(module, "get_processor"): + processors[f"{name}.processor"] = module.get_processor(return_deprecated_lora=True) + + for sub_name, child in module.named_children(): + fn_recursive_add_processors(f"{name}.{sub_name}", child, processors) + + return processors + + for name, module in self.named_children(): + fn_recursive_add_processors(name, module, processors) + + return processors + + def set_attn_processor( + self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]] + ): + r""" + Sets the attention processor to use to compute attention. + + Parameters: + processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`): + The instantiated processor class or a dictionary of processor classes that will be set as the processor + for **all** `Attention` layers. + + If `processor` is a dict, the key needs to define the path to the corresponding cross attention + processor. This is strongly recommended when setting trainable attention processors. + + """ + count = len(self.attn_processors.keys()) + + if isinstance(processor, dict) and len(processor) != count: + raise ValueError( + f"A dict of processors was passed, but the number of processors {len(processor)} does not match the" + f" number of attention layers: {count}. Please make sure to pass {count} processor classes." + ) + + def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor): + if hasattr(module, "set_processor"): + if not isinstance(processor, dict): + module.set_processor(processor) + else: + module.set_processor(processor.pop(f"{name}.processor")) + + for sub_name, child in module.named_children(): + fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor) + + for name, module in self.named_children(): + fn_recursive_attn_processor(name, module, processor) + + def set_default_attn_processor(self): + """ + Disables custom attention processors and sets the default attention implementation. + """ + if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): + processor = AttnAddedKVProcessor() + elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): + processor = AttnProcessor() + else: + raise ValueError( + f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" + ) + + self.set_attn_processor(processor) + + def set_attention_slice(self, slice_size): + r""" + Enable sliced attention computation. + + When this option is enabled, the attention module splits the input tensor in slices to compute attention in + several steps. This is useful for saving some memory in exchange for a small decrease in speed. + + Args: + slice_size (`str` or `int` or `list(int)`, *optional*, defaults to `"auto"`): + When `"auto"`, input to the attention heads is halved, so attention is computed in two steps. If + `"max"`, maximum amount of memory is saved by running only one slice at a time. If a number is + provided, uses as many slices as `attention_head_dim // slice_size`. In this case, `attention_head_dim` + must be a multiple of `slice_size`. + """ + sliceable_head_dims = [] + + def fn_recursive_retrieve_sliceable_dims(module: torch.nn.Module): + if hasattr(module, "set_attention_slice"): + sliceable_head_dims.append(module.sliceable_head_dim) + + for child in module.children(): + fn_recursive_retrieve_sliceable_dims(child) + + # retrieve number of attention layers + for module in self.children(): + fn_recursive_retrieve_sliceable_dims(module) + + num_sliceable_layers = len(sliceable_head_dims) + + if slice_size == "auto": + # half the attention head size is usually a good trade-off between + # speed and memory + slice_size = [dim // 2 for dim in sliceable_head_dims] + elif slice_size == "max": + # make smallest slice possible + slice_size = num_sliceable_layers * [1] + + slice_size = num_sliceable_layers * [slice_size] if not isinstance(slice_size, list) else slice_size + + if len(slice_size) != len(sliceable_head_dims): + raise ValueError( + f"You have provided {len(slice_size)}, but {self.config} has {len(sliceable_head_dims)} different" + f" attention layers. Make sure to match `len(slice_size)` to be {len(sliceable_head_dims)}." + ) + + for i in range(len(slice_size)): + size = slice_size[i] + dim = sliceable_head_dims[i] + if size is not None and size > dim: + raise ValueError(f"size {size} has to be smaller or equal to {dim}.") + + # Recursively walk through all the children. + # Any children which exposes the set_attention_slice method + # gets the message + def fn_recursive_set_attention_slice(module: torch.nn.Module, slice_size: List[int]): + if hasattr(module, "set_attention_slice"): + module.set_attention_slice(slice_size.pop()) + + for child in module.children(): + fn_recursive_set_attention_slice(child, slice_size) + + reversed_slice_size = list(reversed(slice_size)) + for module in self.children(): + fn_recursive_set_attention_slice(module, reversed_slice_size) + + def _set_gradient_checkpointing(self, module, value=False): + if hasattr(module, "gradient_checkpointing"): + module.gradient_checkpointing = value + + def enable_freeu(self, s1, s2, b1, b2): + r"""Enables the FreeU mechanism from https://arxiv.org/abs/2309.11497. + + The suffixes after the scaling factors represent the stage blocks where they are being applied. + + Please refer to the [official repository](https://github.com/ChenyangSi/FreeU) for combinations of values that + are known to work well for different pipelines such as Stable Diffusion v1, v2, and Stable Diffusion XL. + + Args: + s1 (`float`): + Scaling factor for stage 1 to attenuate the contributions of the skip features. This is done to + mitigate the "oversmoothing effect" in the enhanced denoising process. + s2 (`float`): + Scaling factor for stage 2 to attenuate the contributions of the skip features. This is done to + mitigate the "oversmoothing effect" in the enhanced denoising process. + b1 (`float`): Scaling factor for stage 1 to amplify the contributions of backbone features. + b2 (`float`): Scaling factor for stage 2 to amplify the contributions of backbone features. + """ + for i, upsample_block in enumerate(self.up_blocks): + setattr(upsample_block, "s1", s1) + setattr(upsample_block, "s2", s2) + setattr(upsample_block, "b1", b1) + setattr(upsample_block, "b2", b2) + + def disable_freeu(self): + """Disables the FreeU mechanism.""" + freeu_keys = {"s1", "s2", "b1", "b2"} + for i, upsample_block in enumerate(self.up_blocks): + for k in freeu_keys: + if hasattr(upsample_block, k) or getattr(upsample_block, k, None) is not None: + setattr(upsample_block, k, None) + + def forward( + self, + sample: torch.FloatTensor, + timestep: Union[torch.Tensor, float, int], + encoder_hidden_states: torch.Tensor, + class_labels: Optional[torch.Tensor] = None, + timestep_cond: Optional[torch.Tensor] = None, + attention_mask: Optional[torch.Tensor] = None, + cross_attention_kwargs: Optional[Dict[str, Any]] = None, + added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None, + down_block_additional_residuals: Optional[Tuple[torch.Tensor]] = None, + mid_block_additional_residual: Optional[torch.Tensor] = None, + down_intrablock_additional_residuals: Optional[Tuple[torch.Tensor]] = None, + encoder_attention_mask: Optional[torch.Tensor] = None, + return_dict: bool = True, + ) -> Union[UNet2DConditionOutput, Tuple]: + r""" + The [`UNet2DConditionModel`] forward method. + + Args: + sample (`torch.FloatTensor`): + The noisy input tensor with the following shape `(batch, channel, height, width)`. + timestep (`torch.FloatTensor` or `float` or `int`): The number of timesteps to denoise an input. + encoder_hidden_states (`torch.FloatTensor`): + The encoder hidden states with shape `(batch, sequence_length, feature_dim)`. + class_labels (`torch.Tensor`, *optional*, defaults to `None`): + Optional class labels for conditioning. Their embeddings will be summed with the timestep embeddings. + timestep_cond: (`torch.Tensor`, *optional*, defaults to `None`): + Conditional embeddings for timestep. If provided, the embeddings will be summed with the samples passed + through the `self.time_embedding` layer to obtain the timestep embeddings. + attention_mask (`torch.Tensor`, *optional*, defaults to `None`): + An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask + is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large + negative values to the attention scores corresponding to "discard" tokens. + cross_attention_kwargs (`dict`, *optional*): + A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under + `self.processor` in + [diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). + added_cond_kwargs: (`dict`, *optional*): + A kwargs dictionary containing additional embeddings that if specified are added to the embeddings that + are passed along to the UNet blocks. + down_block_additional_residuals: (`tuple` of `torch.Tensor`, *optional*): + A tuple of tensors that if specified are added to the residuals of down unet blocks. + mid_block_additional_residual: (`torch.Tensor`, *optional*): + A tensor that if specified is added to the residual of the middle unet block. + encoder_attention_mask (`torch.Tensor`): + A cross-attention mask of shape `(batch, sequence_length)` is applied to `encoder_hidden_states`. If + `True` the mask is kept, otherwise if `False` it is discarded. Mask will be converted into a bias, + which adds large negative values to the attention scores corresponding to "discard" tokens. + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~models.unet_2d_condition.UNet2DConditionOutput`] instead of a plain + tuple. + cross_attention_kwargs (`dict`, *optional*): + A kwargs dictionary that if specified is passed along to the [`AttnProcessor`]. + added_cond_kwargs: (`dict`, *optional*): + A kwargs dictionary containin additional embeddings that if specified are added to the embeddings that + are passed along to the UNet blocks. + down_block_additional_residuals (`tuple` of `torch.Tensor`, *optional*): + additional residuals to be added to UNet long skip connections from down blocks to up blocks for + example from ControlNet side model(s) + mid_block_additional_residual (`torch.Tensor`, *optional*): + additional residual to be added to UNet mid block output, for example from ControlNet side model + down_intrablock_additional_residuals (`tuple` of `torch.Tensor`, *optional*): + additional residuals to be added within UNet down blocks, for example from T2I-Adapter side model(s) + + Returns: + [`~models.unet_2d_condition.UNet2DConditionOutput`] or `tuple`: + If `return_dict` is True, an [`~models.unet_2d_condition.UNet2DConditionOutput`] is returned, otherwise + a `tuple` is returned where the first element is the sample tensor. + """ + # By default samples have to be AT least a multiple of the overall upsampling factor. + # The overall upsampling factor is equal to 2 ** (# num of upsampling layers). + # However, the upsampling interpolation output size can be forced to fit any upsampling size + # on the fly if necessary. + default_overall_up_factor = 2**self.num_upsamplers + + # upsample size should be forwarded when sample is not a multiple of `default_overall_up_factor` + forward_upsample_size = False + upsample_size = None + + for dim in sample.shape[-2:]: + if dim % default_overall_up_factor != 0: + # Forward upsample size to force interpolation output size. + forward_upsample_size = True + break + + # ensure attention_mask is a bias, and give it a singleton query_tokens dimension + # expects mask of shape: + # [batch, key_tokens] + # adds singleton query_tokens dimension: + # [batch, 1, key_tokens] + # this helps to broadcast it as a bias over attention scores, which will be in one of the following shapes: + # [batch, heads, query_tokens, key_tokens] (e.g. torch sdp attn) + # [batch * heads, query_tokens, key_tokens] (e.g. xformers or classic attn) + if attention_mask is not None: + # assume that mask is expressed as: + # (1 = keep, 0 = discard) + # convert mask into a bias that can be added to attention scores: + # (keep = +0, discard = -10000.0) + attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0 + attention_mask = attention_mask.unsqueeze(1) + + # convert encoder_attention_mask to a bias the same way we do for attention_mask + if encoder_attention_mask is not None: + encoder_attention_mask = (1 - encoder_attention_mask.to(sample.dtype)) * -10000.0 + encoder_attention_mask = encoder_attention_mask.unsqueeze(1) + + # 0. center input if necessary + if self.config.center_input_sample: + sample = 2 * sample - 1.0 + + # 1. time + timesteps = timestep + if not torch.is_tensor(timesteps): + # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can + # This would be a good case for the `match` statement (Python 3.10+) + is_mps = sample.device.type == "mps" + if isinstance(timestep, float): + dtype = torch.float32 if is_mps else torch.float64 + else: + dtype = torch.int32 if is_mps else torch.int64 + timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) + elif len(timesteps.shape) == 0: + timesteps = timesteps[None].to(sample.device) + + # broadcast to batch dimension in a way that's compatible with ONNX/Core ML + timesteps = timesteps.expand(sample.shape[0]) + + t_emb = self.time_proj(timesteps) + + # `Timesteps` does not contain any weights and will always return f32 tensors + # but time_embedding might actually be running in fp16. so we need to cast here. + # there might be better ways to encapsulate this. + t_emb = t_emb.to(dtype=sample.dtype) + + emb = self.time_embedding(t_emb, timestep_cond) + aug_emb = None + + if self.class_embedding is not None: + if class_labels is None: + raise ValueError("class_labels should be provided when num_class_embeds > 0") + + if self.config.class_embed_type == "timestep": + class_labels = self.time_proj(class_labels) + + # `Timesteps` does not contain any weights and will always return f32 tensors + # there might be better ways to encapsulate this. + class_labels = class_labels.to(dtype=sample.dtype) + + class_emb = self.class_embedding(class_labels).to(dtype=sample.dtype) + + if self.config.class_embeddings_concat: + emb = torch.cat([emb, class_emb], dim=-1) + else: + emb = emb + class_emb + + if self.config.addition_embed_type == "text": + aug_emb = self.add_embedding(encoder_hidden_states) + elif self.config.addition_embed_type == "text_image": + # Kandinsky 2.1 - style + if "image_embeds" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `addition_embed_type` set to 'text_image' which requires the keyword argument `image_embeds` to be passed in `added_cond_kwargs`" + ) + + image_embs = added_cond_kwargs.get("image_embeds") + text_embs = added_cond_kwargs.get("text_embeds", encoder_hidden_states) + aug_emb = self.add_embedding(text_embs, image_embs) + elif self.config.addition_embed_type == "text_time": + # SDXL - style + if "text_embeds" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `text_embeds` to be passed in `added_cond_kwargs`" + ) + text_embeds = added_cond_kwargs.get("text_embeds") + if "time_ids" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `time_ids` to be passed in `added_cond_kwargs`" + ) + time_ids = added_cond_kwargs.get("time_ids") + time_embeds = self.add_time_proj(time_ids.flatten()) + time_embeds = time_embeds.reshape((text_embeds.shape[0], -1)) + add_embeds = torch.concat([text_embeds, time_embeds], dim=-1) + add_embeds = add_embeds.to(emb.dtype) + aug_emb = self.add_embedding(add_embeds) + elif self.config.addition_embed_type == "image": + # Kandinsky 2.2 - style + if "image_embeds" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `addition_embed_type` set to 'image' which requires the keyword argument `image_embeds` to be passed in `added_cond_kwargs`" + ) + image_embs = added_cond_kwargs.get("image_embeds") + aug_emb = self.add_embedding(image_embs) + elif self.config.addition_embed_type == "image_hint": + # Kandinsky 2.2 - style + if "image_embeds" not in added_cond_kwargs or "hint" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `addition_embed_type` set to 'image_hint' which requires the keyword arguments `image_embeds` and `hint` to be passed in `added_cond_kwargs`" + ) + image_embs = added_cond_kwargs.get("image_embeds") + hint = added_cond_kwargs.get("hint") + aug_emb, hint = self.add_embedding(image_embs, hint) + sample = torch.cat([sample, hint], dim=1) + + emb = emb + aug_emb if aug_emb is not None else emb + + if self.time_embed_act is not None: + emb = self.time_embed_act(emb) + + if self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "text_proj": + encoder_hidden_states = self.encoder_hid_proj(encoder_hidden_states) + elif self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "text_image_proj": + # Kadinsky 2.1 - style + if "image_embeds" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'text_image_proj' which requires the keyword argument `image_embeds` to be passed in `added_conditions`" + ) + + image_embeds = added_cond_kwargs.get("image_embeds") + encoder_hidden_states = self.encoder_hid_proj(encoder_hidden_states, image_embeds) + elif self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "image_proj": + # Kandinsky 2.2 - style + if "image_embeds" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'image_proj' which requires the keyword argument `image_embeds` to be passed in `added_conditions`" + ) + image_embeds = added_cond_kwargs.get("image_embeds") + encoder_hidden_states = self.encoder_hid_proj(image_embeds) + # 2. pre-process + sample = self.conv_in(sample) + + # 2.5 GLIGEN position net + if cross_attention_kwargs is not None and cross_attention_kwargs.get("gligen", None) is not None: + cross_attention_kwargs = cross_attention_kwargs.copy() + gligen_args = cross_attention_kwargs.pop("gligen") + cross_attention_kwargs["gligen"] = {"objs": self.position_net(**gligen_args)} + + # 3. down + lora_scale = cross_attention_kwargs.get("scale", 1.0) if cross_attention_kwargs is not None else 1.0 + if USE_PEFT_BACKEND: + # weight the lora layers by setting `lora_scale` for each PEFT layer + scale_lora_layers(self, lora_scale) + + is_controlnet = mid_block_additional_residual is not None and down_block_additional_residuals is not None + # using new arg down_intrablock_additional_residuals for T2I-Adapters, to distinguish from controlnets + is_adapter = down_intrablock_additional_residuals is not None + # maintain backward compatibility for legacy usage, where + # T2I-Adapter and ControlNet both use down_block_additional_residuals arg + # but can only use one or the other + if not is_adapter and mid_block_additional_residual is None and down_block_additional_residuals is not None: + deprecate( + "T2I should not use down_block_additional_residuals", + "1.3.0", + "Passing intrablock residual connections with `down_block_additional_residuals` is deprecated \ + and will be removed in diffusers 1.3.0. `down_block_additional_residuals` should only be used \ + for ControlNet. Please make sure use `down_intrablock_additional_residuals` instead. ", + standard_warn=False, + ) + down_intrablock_additional_residuals = down_block_additional_residuals + is_adapter = True + + down_block_res_samples = (sample,) + for downsample_block in self.down_blocks: + if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention: + # For t2i-adapter CrossAttnDownBlock2D + additional_residuals = {} + if is_adapter and len(down_intrablock_additional_residuals) > 0: + additional_residuals["additional_residuals"] = down_intrablock_additional_residuals.pop(0) + + sample, res_samples = downsample_block( + hidden_states=sample, + temb=emb, + encoder_hidden_states=encoder_hidden_states, + attention_mask=attention_mask, + cross_attention_kwargs=cross_attention_kwargs, + encoder_attention_mask=encoder_attention_mask, + **additional_residuals, + ) + else: + sample, res_samples = downsample_block(hidden_states=sample, temb=emb, scale=lora_scale) + if is_adapter and len(down_intrablock_additional_residuals) > 0: + sample += down_intrablock_additional_residuals.pop(0) + + down_block_res_samples += res_samples + + if is_controlnet: + new_down_block_res_samples = () + + for down_block_res_sample, down_block_additional_residual in zip( + down_block_res_samples, down_block_additional_residuals + ): + down_block_res_sample = down_block_res_sample + down_block_additional_residual + new_down_block_res_samples = new_down_block_res_samples + (down_block_res_sample,) + + down_block_res_samples = new_down_block_res_samples + + # 4. mid + if self.mid_block is not None: + if hasattr(self.mid_block, "has_cross_attention") and self.mid_block.has_cross_attention: + sample = self.mid_block( + sample, + emb, + encoder_hidden_states=encoder_hidden_states, + attention_mask=attention_mask, + cross_attention_kwargs=cross_attention_kwargs, + encoder_attention_mask=encoder_attention_mask, + ) + else: + sample = self.mid_block(sample, emb) + + # To support T2I-Adapter-XL + if ( + is_adapter + and len(down_intrablock_additional_residuals) > 0 + and sample.shape == down_intrablock_additional_residuals[0].shape + ): + sample += down_intrablock_additional_residuals.pop(0) + + if is_controlnet: + sample = sample + mid_block_additional_residual + + # 5. up + for i, upsample_block in enumerate(self.up_blocks): + is_final_block = i == len(self.up_blocks) - 1 + + res_samples = down_block_res_samples[-len(upsample_block.resnets) :] + down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)] + + # if we have not reached the final block and need to forward the + # upsample size, we do it here + if not is_final_block and forward_upsample_size: + upsample_size = down_block_res_samples[-1].shape[2:] + + if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention: + sample = upsample_block( + hidden_states=sample, + temb=emb, + res_hidden_states_tuple=res_samples, + encoder_hidden_states=encoder_hidden_states, + cross_attention_kwargs=cross_attention_kwargs, + upsample_size=upsample_size, + attention_mask=attention_mask, + encoder_attention_mask=encoder_attention_mask, + ) + else: + sample = upsample_block( + hidden_states=sample, + temb=emb, + res_hidden_states_tuple=res_samples, + upsample_size=upsample_size, + scale=lora_scale, + ) + + self.latent_store(sample.detach(), t=timestep, layer_index=i) + + # 6. post-process + if self.conv_norm_out: + sample = self.conv_norm_out(sample) + sample = self.conv_act(sample) + sample = self.conv_out(sample) + + if USE_PEFT_BACKEND: + # remove `lora_scale` from each PEFT layer + unscale_lora_layers(self, lora_scale) + + if not return_dict: + return (sample,) + + return UNet2DConditionOutput(sample=sample) diff --git a/modules/consistory/consistory_utils.py b/modules/consistory/consistory_utils.py new file mode 100644 index 000000000..4255ac3ed --- /dev/null +++ b/modules/consistory/consistory_utils.py @@ -0,0 +1,192 @@ +# Copyright (C) 2024 NVIDIA Corporation. All rights reserved. +# +# This work is licensed under the LICENSE file +# located at the root directory. + +from typing import List +from collections import defaultdict +import numpy as np +import torch +from .utils.general_utils import get_dynamic_threshold + + +class FeatureInjector: + def __init__(self, nn_map, nn_distances, attn_masks, inject_range_alpha=[(10,20,0.8)], swap_strategy='min', dist_thr='dynamic', inject_unet_parts=['up']): + self.nn_map = nn_map + self.nn_distances = nn_distances + self.attn_masks = attn_masks + self.inject_range_alpha = inject_range_alpha if isinstance(inject_range_alpha, list) else [inject_range_alpha] + self.swap_strategy = swap_strategy # 'min / 'mean' / 'first' + self.dist_thr = dist_thr + self.inject_unet_parts = inject_unet_parts + self.inject_res = [64] + + def inject_outputs(self, output, curr_iter, output_res, extended_mapping, place_in_unet, anchors_cache=None): + curr_unet_part = place_in_unet.split('_')[0] + + # Inject only in the specified unet parts (up, mid, down) + if (curr_unet_part not in self.inject_unet_parts) or output_res not in self.inject_res: + return output + + bsz = output.shape[0] + nn_map = self.nn_map[output_res] + nn_distances = self.nn_distances[output_res] + attn_masks = self.attn_masks[output_res] + vector_dim = output_res**2 + + alpha = next((alpha for min_range, max_range, alpha in self.inject_range_alpha if min_range <= curr_iter <= max_range), None) + if alpha: + old_output = output#.clone() + for i in range(bsz): + other_outputs = [] + + if self.swap_strategy == 'min': + curr_mapping = extended_mapping[i] + + # If the current image is not mapped to any other image, skip + if not torch.any(torch.cat([curr_mapping[:i], curr_mapping[i+1:]])): + continue + + min_dists = nn_distances[i][curr_mapping].argmin(dim=0) + curr_nn_map = nn_map[i][curr_mapping][min_dists, torch.arange(vector_dim)] + + curr_nn_distances = nn_distances[i][curr_mapping][min_dists, torch.arange(vector_dim)] + dist_thr = get_dynamic_threshold(curr_nn_distances) if self.dist_thr == 'dynamic' else self.dist_thr + dist_mask = curr_nn_distances < dist_thr + final_mask_tgt = attn_masks[i] & dist_mask + + other_outputs = old_output[curr_mapping][min_dists, curr_nn_map][final_mask_tgt] + + output[i][final_mask_tgt] = alpha * other_outputs + (1 - alpha)*old_output[i][final_mask_tgt] + + if anchors_cache and anchors_cache.is_cache_mode(): + if place_in_unet not in anchors_cache.h_out_cache: + anchors_cache.h_out_cache[place_in_unet] = {} + + anchors_cache.h_out_cache[place_in_unet][curr_iter] = output + + return output + + def inject_anchors(self, output, curr_iter, output_res, extended_mapping, place_in_unet, anchors_cache): + curr_unet_part = place_in_unet.split('_')[0] + + # Inject only in the specified unet parts (up, mid, down) + if (curr_unet_part not in self.inject_unet_parts) or output_res not in self.inject_res: + return output + + bsz = output.shape[0] + nn_map = self.nn_map[output_res] + nn_distances = self.nn_distances[output_res] + attn_masks = self.attn_masks[output_res] + vector_dim = output_res**2 + + alpha = next((alpha for min_range, max_range, alpha in self.inject_range_alpha if min_range <= curr_iter <= max_range), None) + if alpha: + + anchor_outputs = anchors_cache.h_out_cache[place_in_unet][curr_iter] + + old_output = output#.clone() + for i in range(bsz): + other_outputs = [] + + if self.swap_strategy == 'min': + min_dists = nn_distances[i].argmin(dim=0) + curr_nn_map = nn_map[i][min_dists, torch.arange(vector_dim)] + + curr_nn_distances = nn_distances[i][min_dists, torch.arange(vector_dim)] + dist_thr = get_dynamic_threshold(curr_nn_distances) if self.dist_thr == 'dynamic' else self.dist_thr + dist_mask = curr_nn_distances < dist_thr + final_mask_tgt = attn_masks[i] & dist_mask + + other_outputs = anchor_outputs[min_dists, curr_nn_map][final_mask_tgt] + + output[i][final_mask_tgt] = alpha * other_outputs + (1 - alpha)*old_output[i][final_mask_tgt] + + return output + + +class AnchorCache: + def __init__(self): + self.input_h_cache = {} # place_in_unet, iter, h_in + self.h_out_cache = {} # place_in_unet, iter, h_out + self.anchors_last_mask = None + self.dift_cache = None + + self.mode = 'cache' # mode can be 'cache' or 'inject' + + def set_mode(self, mode): + self.mode = mode + + def set_mode_inject(self): + self.mode = 'inject' + + def set_mode_cache(self): + self.mode = 'cache' + + def is_inject_mode(self): + return self.mode == 'inject' + + def is_cache_mode(self): + return self.mode == 'cache' + + + def to_device(self, device): + for key, value in self.input_h_cache.items(): + self.input_h_cache[key] = {k: v.to(device) for k, v in value.items()} + + for key, value in self.h_out_cache.items(): + self.h_out_cache[key] = {k: v.to(device) for k, v in value.items()} + + if self.anchors_last_mask: + self.anchors_last_mask = {k: v.to(device) for k, v in self.anchors_last_mask.items()} + + if self.dift_cache is not None: + self.dift_cache = self.dift_cache.to(device) + + +class QueryStore: + def __init__(self, mode='store', t_range=[0, 1000], strength_start=1, strength_end=1): + """ + Initialize an empty ActivationsStore + """ + self.query_store = defaultdict(list) + self.mode = mode + self.t_range = t_range + self.strengthes = np.linspace(strength_start, strength_end, (t_range[1] - t_range[0])+1) + + def set_mode(self, mode): # mode can be 'cache' or 'inject' + self.mode = mode + + def cache_query(self, query, place_in_unet: str): + self.query_store[place_in_unet] = query + + def inject_query(self, query, place_in_unet, t): + if t >= self.t_range[0] and t <= self.t_range[1]: + relative_t = t - self.t_range[0] + strength = self.strengthes[relative_t] + new_query = strength * self.query_store[place_in_unet] + (1 - strength) * query + else: + new_query = query + + return new_query + +class DIFTLatentStore: + def __init__(self, steps: List[int], up_ft_indices: List[int]): + self.steps = steps + self.up_ft_indices = up_ft_indices + self.dift_features = {} + + def __call__(self, features: torch.Tensor, t: int, layer_index: int): + if t in self.steps and layer_index in self.up_ft_indices: + self.dift_features[f'{int(t)}_{layer_index}'] = features + + def copy(self): + copy_dift = DIFTLatentStore(self.steps, self.up_ft_indices) + + for key, value in self.dift_features.items(): + copy_dift.dift_features[key] = value.clone() + + return copy_dift + + def reset(self): + self.dift_features = {} diff --git a/modules/consistory/utils/general_utils.py b/modules/consistory/utils/general_utils.py new file mode 100644 index 000000000..4493fa96e --- /dev/null +++ b/modules/consistory/utils/general_utils.py @@ -0,0 +1,117 @@ +# Copyright (C) 2024 NVIDIA Corporation. All rights reserved. +# +# This work is licensed under the LICENSE file +# located at the root directory. + +import torch +import torch.nn.functional as F +import numpy as np +from skimage import filters + + +## Attention Utils +def get_dynamic_threshold(tensor): + return filters.threshold_otsu(tensor.float().cpu().numpy()) + + +def attn_map_to_binary(attention_map, scaler=1.): + attention_map_np = attention_map.float().cpu().numpy() + threshold_value = filters.threshold_otsu(attention_map_np) * scaler + binary_mask = (attention_map_np > threshold_value).astype(np.uint8) + + return binary_mask + + +## Features + +def gaussian_smooth(input_tensor, kernel_size=3, sigma=1): + """ + Function to apply Gaussian smoothing on each 2D slice of a 3D tensor. + """ + kernel = np.fromfunction( + lambda x, y: (1/ (2 * np.pi * sigma ** 2)) * + np.exp(-((x - (kernel_size - 1) / 2) ** 2 + (y - (kernel_size - 1) / 2) ** 2) / (2 * sigma ** 2)), + (kernel_size, kernel_size) + ) + kernel = torch.Tensor(kernel / kernel.sum()).to(input_tensor.dtype).to(input_tensor.device) + # Add batch and channel dimensions to the kernel + kernel = kernel.unsqueeze(0).unsqueeze(0) + # Iterate over each 2D slice and apply convolution + smoothed_slices = [] + for i in range(input_tensor.size(0)): + slice_tensor = input_tensor[i, :, :] + slice_tensor = F.conv2d(slice_tensor.unsqueeze(0).unsqueeze(0), kernel, padding=kernel_size // 2)[0, 0] + smoothed_slices.append(slice_tensor) + # Stack the smoothed slices to get the final tensor + smoothed_tensor = torch.stack(smoothed_slices, dim=0) + return smoothed_tensor + + +## Dense correspondence utils + +def cos_dist(a, b): + a_norm = F.normalize(a, dim=-1) + b_norm = F.normalize(b, dim=-1) + res = a_norm @ b_norm.T + return 1 - res + + +def gen_nn_map(src_features, src_mask, tgt_features, tgt_mask, device, batch_size=100, tgt_size=768): + resized_src_features = F.interpolate(src_features.unsqueeze(0), size=tgt_size, mode='bilinear', align_corners=False).squeeze(0) + resized_src_features = resized_src_features.permute(1,2,0).view(tgt_size**2, -1) + resized_tgt_features = F.interpolate(tgt_features.unsqueeze(0), size=tgt_size, mode='bilinear', align_corners=False).squeeze(0) + resized_tgt_features = resized_tgt_features.permute(1,2,0).view(tgt_size**2, -1) + nearest_neighbor_indices = torch.zeros(tgt_size**2, dtype=torch.long, device=device) + nearest_neighbor_distances = torch.zeros(tgt_size**2, dtype=src_features.dtype, device=device) + if not batch_size: + batch_size = tgt_size**2 + for i in range(0, tgt_size**2, batch_size): + distances = cos_dist(resized_src_features, resized_tgt_features[i:i+batch_size]) + distances[~src_mask] = 2. + min_distances, min_indices = torch.min(distances, dim=0) + nearest_neighbor_indices[i:i+batch_size] = min_indices + nearest_neighbor_distances[i:i+batch_size] = min_distances + return nearest_neighbor_indices, nearest_neighbor_distances + + +def cyclic_nn_map(features, masks, latent_resolutions, device): + bsz = features.shape[0] + nn_map_dict = {} + nn_distances_dict = {} + + for tgt_size in latent_resolutions: + nn_map = torch.empty(bsz, bsz, tgt_size**2, dtype=torch.long, device=device) + nn_distances = torch.full((bsz, bsz, tgt_size**2), float('inf'), dtype=features.dtype, device=device) + + for i in range(bsz): + for j in range(bsz): + if i != j: + nearest_neighbor_indices, nearest_neighbor_distances = gen_nn_map(features[j], masks[tgt_size][j], features[i], masks[tgt_size][i], device, batch_size=None, tgt_size=tgt_size) + nn_map[i,j] = nearest_neighbor_indices + nn_distances[i,j] = nearest_neighbor_distances + + nn_map_dict[tgt_size] = nn_map + nn_distances_dict[tgt_size] = nn_distances + + return nn_map_dict, nn_distances_dict + + +def anchor_nn_map(features, anchor_features, masks, anchor_masks, latent_resolutions, device): + bsz = features.shape[0] + anchor_bsz = anchor_features.shape[0] + nn_map_dict = {} + nn_distances_dict = {} + + for tgt_size in latent_resolutions: + nn_map = torch.empty(bsz, anchor_bsz, tgt_size**2, dtype=torch.long, device=device) + nn_distances = torch.full((bsz, anchor_bsz, tgt_size**2), float('inf'), dtype=features.dtype, device=device) + + for i in range(bsz): + for j in range(anchor_bsz): + nearest_neighbor_indices, nearest_neighbor_distances = gen_nn_map(anchor_features[j], anchor_masks[tgt_size][j], features[i], masks[tgt_size][i], device, batch_size=None, tgt_size=tgt_size) + nn_map[i,j] = nearest_neighbor_indices + nn_distances[i,j] = nearest_neighbor_distances + nn_map_dict[tgt_size] = nn_map + nn_distances_dict[tgt_size] = nn_distances + + return nn_map_dict, nn_distances_dict diff --git a/modules/consistory/utils/ptp_utils.py b/modules/consistory/utils/ptp_utils.py new file mode 100644 index 000000000..0ab2d8d93 --- /dev/null +++ b/modules/consistory/utils/ptp_utils.py @@ -0,0 +1,194 @@ +# Copyright 2022 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# MIT License +# +# Copyright (c) 2023 AttendAndExcite +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in all +# copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. + +# Copyright 2022 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# Not a contribution +# Changes made by NVIDIA CORPORATION & AFFILIATES enabling ConsiStory or otherwise documented as NVIDIA-proprietary +# are not a contribution and subject to the license under the LICENSE file located at the root directory. + +import torch +from collections import defaultdict +import numpy as np +from typing import Union, List +from PIL import Image + +from modules.consistory.utils.general_utils import attn_map_to_binary +import torch.nn.functional as F + + +class AttentionStore: + def __init__(self, attention_store_kwargs): + """ + Initialize an empty AttentionStore :param step_index: used to visualize only a specific step in the diffusion + process + """ + self.attn_res = attention_store_kwargs.get('attn_res', (32,32)) + self.token_indices = attention_store_kwargs['token_indices'] + bsz = self.token_indices.size(1) + self.mask_background_query = attention_store_kwargs.get('mask_background_query', False) + self.original_attn_masks = attention_store_kwargs.get('original_attn_masks', None) + self.extended_mapping = attention_store_kwargs.get('extended_mapping', torch.ones(bsz, bsz).bool()) + self.mask_dropout = attention_store_kwargs.get('mask_dropout', 0.0) + torch.manual_seed(0) # For dropout mask reproducibility + + self.curr_iter = 0 + self.ALL_RES = [32, 64] + self.step_store = defaultdict(list) + self.attn_masks = {res: None for res in self.ALL_RES} + self.last_mask = {res: None for res in self.ALL_RES} + self.last_mask_dropout = {res: None for res in self.ALL_RES} + + def __call__(self, attn, is_cross: bool, place_in_unet: str, attn_heads: int): + if is_cross and attn.shape[1] == np.prod(self.attn_res): + guidance_attention = attn[attn.size(0)//2:] + batched_guidance_attention = guidance_attention.reshape([guidance_attention.shape[0]//attn_heads, attn_heads, *guidance_attention.shape[1:]]) + batched_guidance_attention = batched_guidance_attention.mean(dim=1) + self.step_store[place_in_unet].append(batched_guidance_attention) + + def reset(self): + self.step_store = defaultdict(list) + self.attn_masks = {res: None for res in self.ALL_RES} + self.last_mask = {res: None for res in self.ALL_RES} + self.last_mask_dropout = {res: None for res in self.ALL_RES} + + torch.cuda.empty_cache() + + def aggregate_last_steps_attention(self) -> torch.Tensor: + """Aggregates the attention across the different layers and heads at the specified resolution.""" + attention_maps = torch.cat([torch.stack(x[-20:]) for x in self.step_store.values()]).mean(dim=0) + bsz, wh, _ = attention_maps.shape + + # Create attention maps for each concept token, for each batch item + agg_attn_maps = [] + for i in range(bsz): + curr_prompt_indices = [] + + for concept_token_indices in self.token_indices: + if concept_token_indices[i] != -1: + curr_prompt_indices.append(attention_maps[i, :, concept_token_indices[i]].view(*self.attn_res)) + + agg_attn_maps.append(torch.stack(curr_prompt_indices)) + + # Upsample the attention maps to the target resolution + # and create the attention masks, unifying masks across the different concepts + for tgt_size in self.ALL_RES: + pixels = tgt_size ** 2 + tgt_agg_attn_maps = [F.interpolate(x.unsqueeze(1), size=tgt_size, mode='bilinear').squeeze(1) for x in agg_attn_maps] + + attn_masks = [] + for batch_item_map in tgt_agg_attn_maps: + concept_attn_masks = [] + + for concept_maps in batch_item_map: + concept_attn_masks.append(torch.from_numpy(attn_map_to_binary(concept_maps, 1.)).to(attention_maps.device).bool().view(-1)) + + concept_attn_masks = torch.stack(concept_attn_masks, dim=0).max(dim=0).values + attn_masks.append(concept_attn_masks) + + attn_masks = torch.stack(attn_masks) + self.last_mask[tgt_size] = attn_masks.clone() + + # Add mask dropout + if self.curr_iter < 1000: + rand_mask = (torch.rand_like(attn_masks.float()) < self.mask_dropout) + attn_masks[rand_mask] = False + + self.last_mask_dropout[tgt_size] = attn_masks.clone() + + # # Create subject driven extended self attention masks + # output_attn_mask = torch.zeros((bsz, tgt_size**2, attn_masks.view(-1).size(0)), device=attn_masks.device).bool() + + # for i in range(bsz): + # for j in range(bsz): + # if i==j: + # output_attn_mask[i, :, j*pixels:(j+1)*pixels] = 1 + # else: + # if self.extended_mapping[i,j]: + # if not self.mask_background_query: + # output_attn_mask[i, :, j*pixels:(j+1)*pixels] = attn_masks[j].unsqueeze(0).expand(pixels, -1) + # else: + # output_attn_mask[i, attn_masks[i], j*pixels:(j+1)*pixels] = attn_masks[j].unsqueeze(0).expand(attn_masks[i].sum(), -1) + + # self.attn_masks[tgt_size] = output_attn_mask + + def get_attn_mask_bias(self, tgt_size, bsz=None): + attn_mask = self.attn_masks[tgt_size] if self.original_attn_masks is None else self.original_attn_masks[tgt_size] + + if attn_mask is None: + return None + + attn_bias = torch.zeros_like(attn_mask, dtype=torch.float16) + attn_bias[~attn_mask] = float('-inf') + + if bsz and bsz != attn_bias.shape[0]: + attn_bias = attn_bias.repeat(bsz // attn_bias.shape[0], 1, 1) + + return attn_bias + + def get_extended_attn_mask_instance(self, width, i): + attn_mask = self.last_mask_dropout[width] + if attn_mask is None: + return None + + n_patches = width**2 + + + output_attn_mask = torch.zeros((attn_mask.shape[0] * attn_mask.shape[1],), device=attn_mask.device, dtype=torch.bool) + for j in range(attn_mask.shape[0]): + if i==j: + output_attn_mask[j*n_patches:(j+1)*n_patches] = 1 + else: + if self.extended_mapping[i,j]: + if not self.mask_background_query: + output_attn_mask[j*n_patches:(j+1)*n_patches] = attn_mask[j].unsqueeze(0) #.expand(n_patches, -1) + else: + raise NotImplementedError('mask_background_query is not supported anymore') + output_attn_mask[0, attn_mask[i], k*n_patches:(k+1)*n_patches] = attn_mask[j].unsqueeze(0).expand(attn_mask[i].sum(), -1) + + return output_attn_mask \ No newline at end of file diff --git a/modules/control/run.py b/modules/control/run.py index e3c7bbf76..74bab35c4 100644 --- a/modules/control/run.py +++ b/modules/control/run.py @@ -564,8 +564,8 @@ def control_run(state: str = '', return [], '', '', 'Reference mode without image' elif unit_type == 'controlnet' and has_models: if input_type == 0: # Control only - if shared.sd_model_type == 'f1' and 'control_image' not in p.task_args: - p.task_args['control_image'] = p.init_images # flux controlnet mandates this + if shared.sd_model_type in ['f1', 'sd3'] and 'control_image' not in p.task_args: + p.task_args['control_image'] = p.init_images # some controlnets mandate this p.task_args['strength'] = p.denoising_strength elif input_type == 1: # Init image same as control p.task_args['control_image'] = p.init_images # switch image and control_image diff --git a/modules/control/units/controlnet.py b/modules/control/units/controlnet.py index 236892eb5..48a3d440d 100644 --- a/modules/control/units/controlnet.py +++ b/modules/control/units/controlnet.py @@ -1,7 +1,7 @@ import os import time from typing import Union -from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline, FluxPipeline, ControlNetModel +from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline, FluxPipeline, StableDiffusion3Pipeline, ControlNetModel from modules.control.units import detect from modules.shared import log, opts, listdir from modules import errors, sd_models, devices, model_quant @@ -50,7 +50,6 @@ predefined_sdxl = { 'Depth Zoe XL': 'diffusers/controlnet-zoe-depth-sdxl-1.0', 'Depth Mid XL': 'diffusers/controlnet-depth-sdxl-1.0-mid', 'OpenPose XL': 'thibaud/controlnet-openpose-sdxl-1.0/bin', - # 'OpenPose XL': 'thibaud/controlnet-openpose-sdxl-1.0/OpenPoseXL2.safetensors', 'Xinsir Union XL': 'xinsir/controlnet-union-sdxl-1.0', 'Xinsir OpenPose XL': 'xinsir/controlnet-openpose-sdxl-1.0', 'Xinsir Canny XL': 'xinsir/controlnet-canny-sdxl-1.0', @@ -79,11 +78,20 @@ predefined_f1 = { "XLabs-AI Depth": 'XLabs-AI/flux-controlnet-depth-diffusers', "XLabs-AI HED": 'XLabs-AI/flux-controlnet-hed-diffusers' } +predefined_sd3 = { + "InstantX Canny": 'InstantX/SD3-Controlnet-Canny', + "InstantX Pose": 'InstantX/SD3-Controlnet-Pose', + "InstantX Depth": 'InstantX/SD3-Controlnet-Depth', + "InstantX Tile": 'InstantX/SD3-Controlnet-Tile', + "Alimama Inpainting": 'alimama-creative/SD3-Controlnet-Inpainting', + "Alimama SoftEdge": 'alimama-creative/SD3-Controlnet-Softedge', +} models = {} all_models = {} all_models.update(predefined_sd15) all_models.update(predefined_sdxl) all_models.update(predefined_f1) +all_models.update(predefined_sd3) cache_dir = 'models/control/controlnet' @@ -118,9 +126,11 @@ def list_models(refresh=False): models = ['None'] + list(predefined_sd15) + sorted(find_models()) elif modules.shared.sd_model_type == 'f1': models = ['None'] + list(predefined_f1) + sorted(find_models()) + elif modules.shared.sd_model_type == 'sd3': + models = ['None'] + list(predefined_sd3) + sorted(find_models()) else: log.warning(f'Control {what} model list failed: unknown model type') - models = ['None'] + sorted(predefined_sd15) + sorted(predefined_sdxl) + sorted(find_models()) + models = ['None'] + sorted(predefined_sd15) + sorted(predefined_sdxl) + sorted(predefined_f1) + sorted(predefined_sd3) + sorted(find_models()) debug(f'Control list {what}: path={cache_dir} models={models}') return models @@ -151,6 +161,8 @@ class ControlNet(): from diffusers import ControlNetModel as model_class # pylint: disable=reimported # sdxl shares same model class elif modules.shared.sd_model_type == 'f1': from diffusers import FluxControlNetModel as model_class + elif modules.shared.sd_model_type == 'sd3': + from diffusers import SD3ControlNetModel as model_class else: log.error(f'Control {what}: type={modules.shared.sd_model_type} unsupported model') return None @@ -247,7 +259,11 @@ class ControlNet(): class ControlNetPipeline(): - def __init__(self, controlnet: Union[ControlNetModel, list[ControlNetModel]], pipeline: Union[StableDiffusionXLPipeline, StableDiffusionPipeline, FluxPipeline], dtype = None): + def __init__(self, + controlnet: Union[ControlNetModel, list[ControlNetModel]], + pipeline: Union[StableDiffusionXLPipeline, StableDiffusionPipeline, FluxPipeline, StableDiffusion3Pipeline], + dtype = None, + ): t0 = time.time() self.orig_pipeline = pipeline self.pipeline = None @@ -293,6 +309,20 @@ class ControlNetPipeline(): scheduler=pipeline.scheduler, controlnet=controlnet, # can be a list ) + elif detect.is_sd3(pipeline): + from diffusers import StableDiffusion3ControlNetPipeline + self.pipeline = StableDiffusion3ControlNetPipeline( + vae=pipeline.vae, + text_encoder=pipeline.text_encoder, + text_encoder_2=pipeline.text_encoder_2, + text_encoder_3=pipeline.text_encoder_3, + tokenizer=pipeline.tokenizer, + tokenizer_2=pipeline.tokenizer_2, + tokenizer_3=pipeline.tokenizer_3, + transformer=pipeline.transformer, + scheduler=pipeline.scheduler, + controlnet=controlnet, # can be a list + ) else: log.error(f'Control {what} pipeline: class={pipeline.__class__.__name__} unsupported model type') return diff --git a/modules/control/units/detect.py b/modules/control/units/detect.py index 70a1c8000..7cc8d144d 100644 --- a/modules/control/units/detect.py +++ b/modules/control/units/detect.py @@ -20,5 +20,12 @@ def is_f1(model): if model is None: return False if hasattr(model, '__name__'): - return model.__name__ == p.FluxPipeline.__name__ - return isinstance(model, p.FluxPipeline) + return model.__name__ == p.FluxPipeline.__name__ or model.__name__ == p.FluxImg2ImgPipeline.__name__ or model.__name__ == p.FluxInpaintPipeline.__name__ + return isinstance(model, p.FluxPipeline) or isinstance(model, p.FluxImg2ImgPipeline) or isinstance(model, p.FluxInpaintPipeline) + +def is_sd3(model): + if model is None: + return False + if hasattr(model, '__name__'): + return model.__name__ == p.StableDiffusion3Pipeline.__name__ or model.__name__ == p.StableDiffusion3Img2ImgPipeline.__name__ or model.__name__ == p.StableDiffusion3InpaintPipeline.__name__ + return isinstance(model, p.StableDiffusion3Pipeline) or isinstance(model, p.StableDiffusion3Img2ImgPipeline) or isinstance(model, p.StableDiffusion3InpaintPipeline) diff --git a/modules/devices.py b/modules/devices.py index 490d2a54d..56ac50091 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -4,6 +4,7 @@ import time import contextlib from functools import wraps import torch +from modules import rocm from modules.errors import log, display, install as install_traceback from installer import install @@ -50,8 +51,8 @@ def has_zluda() -> bool: if not cuda_ok: return False try: - device = torch.device("cuda") - return torch.cuda.get_device_name(device).endswith("[ZLUDA]") + dev = torch.device("cuda") + return torch.cuda.get_device_name(dev).endswith("[ZLUDA]") except Exception: return False @@ -206,7 +207,7 @@ def torch_gc(force=False, fast=False): force = True if oom > previous_oom: previous_oom = oom - log.warning(f'GPU out-of-memory error: {mem}') + log.warning(f'Torch GPU out-of-memory error: {mem}') force = True if force: # actual gc @@ -246,13 +247,26 @@ def set_cuda_sync_mode(mode): return try: import ctypes - log.info(f'Set cuda sync: mode={mode}') + log.info(f'Torch CUDA sync: mode={mode}') torch.cuda.set_device(torch.device(get_optimal_device_name())) ctypes.CDLL('libcudart.so').cudaSetDeviceFlags({'auto': 0, 'spin': 1, 'yield': 2, 'block': 4}[mode]) except Exception: pass +def set_cuda_memory_limit(): + if not cuda_ok or opts.cuda_mem_fraction == 0: + return + from modules.shared import cmd_opts + try: + torch_gc(force=True) + mem = torch.cuda.get_device_properties(device).total_memory + torch.cuda.set_per_process_memory_fraction(float(opts.cuda_mem_fraction), cmd_opts.device_id if cmd_opts.device_id is not None else 0) + log.info(f'Torch CUDA memory limit: fraction={opts.cuda_mem_fraction:.2f} limit={round(opts.cuda_mem_fraction * mem / 1024 / 1024)} total={round(mem / 1024 / 1024)}') + except Exception as e: + log.warning(f'Torch CUDA memory limit: fraction={opts.cuda_mem_fraction:.2f} {e}') + + def test_fp16(): global fp16_ok # pylint: disable=global-statement if fp16_ok is not None: @@ -283,16 +297,14 @@ def test_bf16(): if sys.platform == "darwin" or backend == 'openvino' or backend == 'directml': # override bf16_ok = False return bf16_ok - elif backend == 'zluda': - device_name = torch.cuda.get_device_name(device) - if device_name.startswith("AMD Radeon RX "): # only force AMD - device_name = device_name.replace("AMD Radeon RX ", "").split(" ", maxsplit=1)[0] - if len(device_name) == 4 and device_name[0] in {"5", "6"}: # RDNA 1 and 2 - bf16_ok = False - return bf16_ok - elif backend == 'rocm': - gcn_arch = getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000")[3:7] - if len(gcn_arch) == 4 and gcn_arch[0:2] == "10": # RDNA 1 and 2 + elif backend == 'rocm' or backend == 'zluda': + agent = None + if backend == 'rocm': + agent = rocm.Agent(getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000")) + else: + from modules.zluda_installer import default_agent + agent = default_agent + if agent is not None and agent.gfx_version < 0x1100 and agent.arch != rocm.MicroArchitecture.CDNA: # all cards before RDNA 3 except for CDNA cards bf16_ok = False return bf16_ok try: @@ -450,6 +462,7 @@ def set_dtype(): def set_cuda_params(): override_ipex_math() + set_cuda_memory_limit() set_cudnn_params() set_sdpa_params() set_dtype() diff --git a/modules/extras.py b/modules/extras.py index e22360f8a..162491580 100644 --- a/modules/extras.py +++ b/modules/extras.py @@ -188,7 +188,7 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument _, extension = os.path.splitext(output_modelname) if os.path.exists(output_modelname) and not kwargs.get("overwrite", False): - return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], f"Model alredy exists: {output_modelname}"] + return [*[gr.Dropdown.update(choices=sd_models.checkpoint_titles()) for _ in range(4)], f"Model alredy exists: {output_modelname}"] if extension.lower() == ".safetensors": safetensors.torch.save_file(theta_0, output_modelname, metadata=metadata) else: @@ -202,7 +202,7 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument created_model.calculate_shorthash() devices.torch_gc(force=True) shared.state.end() - return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], f"Model saved to {output_modelname}"] + return [*[gr.Dropdown.update(choices=sd_models.checkpoint_titles()) for _ in range(4)], f"Model saved to {output_modelname}"] def run_modelconvert(model, checkpoint_formats, precision, conv_type, custom_name, unet_conv, text_encoder_conv, diff --git a/modules/face/faceid.py b/modules/face/faceid.py index 754ce59a3..b74e15dc5 100644 --- a/modules/face/faceid.py +++ b/modules/face/faceid.py @@ -6,9 +6,10 @@ import numpy as np import diffusers import huggingface_hub as hf from PIL import Image -from modules import processing, shared, devices, extra_networks, sd_models, sd_hijack_freeu, script_callbacks, ipadapter +from modules import processing, shared, devices, extra_networks, sd_hijack_freeu, script_callbacks, ipadapter, token_merge from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet + FACEID_MODELS = { "FaceID Base": "h94/IP-Adapter-FaceID/ip-adapter-faceid_sd15.bin", "FaceID Plus v1": "h94/IP-Adapter-FaceID/ip-adapter-faceid-plus_sd15.bin", @@ -69,7 +70,7 @@ def face_id( shared.prompt_styles.apply_styles_to_extra(p) if shared.opts.cuda_compile_backend == 'none': - sd_models.apply_token_merging(p.sd_model) + token_merge.apply_token_merging(p.sd_model) sd_hijack_freeu.apply_freeu(p, not shared.native) script_callbacks.before_process_callback(p) @@ -78,7 +79,7 @@ def face_id( ip_ckpt = FACEID_MODELS[model] folder, filename = os.path.split(ip_ckpt) basename, _ext = os.path.splitext(filename) - model_path = hf.hf_hub_download(repo_id=folder, filename=filename, cache_dir=shared.opts.diffusers_dir) + model_path = hf.hf_hub_download(repo_id=folder, filename=filename, cache_dir=shared.opts.hfcache_dir) if model_path is None: shared.log.error(f'FaceID download failed: model={model} file="{ip_ckpt}"') return None @@ -246,7 +247,7 @@ def face_id( if faceid_model is not None and original_load_ip_adapter is not None: faceid_model.__class__.load_ip_adapter = original_load_ip_adapter if shared.opts.cuda_compile_backend == 'none': - sd_models.remove_token_merging(p.sd_model) + token_merge.remove_token_merging(p.sd_model) script_callbacks.after_process_callback(p) return processed_images diff --git a/modules/face/insightface.py b/modules/face/insightface.py index 3eb7171bf..a59e48b41 100644 --- a/modules/face/insightface.py +++ b/modules/face/insightface.py @@ -32,7 +32,7 @@ def get_app(mp_name): repo_id='vladmandic/insightface-faceanalysis', filename=f'{mp_name}.zip', local_dir_use_symlinks=False, - cache_dir=opts.diffusers_dir, + cache_dir=opts.hfcache_dir, local_dir=local_dir ) if not os.path.exists(extract_dir): diff --git a/modules/generation_parameters_copypaste.py b/modules/generation_parameters_copypaste.py index cf4278e80..36b5ea382 100644 --- a/modules/generation_parameters_copypaste.py +++ b/modules/generation_parameters_copypaste.py @@ -189,7 +189,7 @@ def create_override_settings_dict(text_pairs): def connect_paste(button, local_paste_fields, input_comp, override_settings_component, tabname): def paste_func(prompt): - if prompt is None or len(prompt.strip()) == 0 and not shared.cmd_opts.hide_ui_dir_config: + if prompt is None or len(prompt.strip()) == 0: filename = os.path.join(data_path, "params.txt") if os.path.exists(filename): with open(filename, "r", encoding="utf8") as file: diff --git a/modules/images.py b/modules/images.py index add7c5c7d..fb9cc9652 100644 --- a/modules/images.py +++ b/modules/images.py @@ -10,7 +10,7 @@ import threading import numpy as np import piexif import piexif.helper -from PIL import Image, PngImagePlugin, ExifTags +from PIL import Image, PngImagePlugin, ExifTags, ImageDraw from modules import sd_samplers, shared, script_callbacks, errors, paths from modules.images_grid import image_grid, get_grid_size, split_grid, combine_grid, check_grid_size, get_font, draw_grid_annotations, draw_prompt_matrix, GridAnnotation, Grid # pylint: disable=unused-import from modules.images_resize import resize_image # pylint: disable=unused-import @@ -361,11 +361,21 @@ def flatten(img, bgcolor): return img.convert('RGB') +def draw_overlay(im, text: str = '', y_offset: int = 0): + d = ImageDraw.Draw(im) + fontsize = (im.width + im.height) // 50 + font = get_font(fontsize) + d.text((fontsize//2, fontsize//2 + y_offset), text, font=font, fill=shared.opts.font_color) + return im + + def set_watermark(image, watermark): if shared.opts.image_watermark_position != 'none': # visible watermark wm_image = None try: wm_image = Image.open(shared.opts.image_watermark_image) + if wm_image.mode != 'RGBA': + wm_image = wm_image.convert('RGBA') except Exception as e: shared.log.warning(f'Set image watermark: fn="{shared.opts.image_watermark_image}" {e}') if wm_image is not None: @@ -384,8 +394,14 @@ def set_watermark(image, watermark): try: for x in range(wm_image.width): for y in range(wm_image.height): - r, g, b, _a = wm_image.getpixel((x, y)) - if not (r == 0 and g == 0 and b == 0): + rgba = wm_image.getpixel((x, y)) + orig = image.getpixel((x+position[0], y+position[1])) + # alpha blend + a = rgba[3] / 255 + r = int(rgba[0] * a + orig[0] * (1 - a)) + g = int(rgba[1] * a + orig[1] * (1 - a)) + b = int(rgba[2] * a + orig[2] * (1 - a)) + if not a == 0: image.putpixel((x+position[0], y+position[1]), (r, g, b)) shared.log.debug(f'Set image watermark: fn="{shared.opts.image_watermark_image}" image={wm_image} position={position}') except Exception as e: diff --git a/modules/images_grid.py b/modules/images_grid.py index 194b0d1cd..51371a0fe 100644 --- a/modules/images_grid.py +++ b/modules/images_grid.py @@ -113,7 +113,7 @@ def get_font(fontsize): return ImageFont.truetype("javascript/notosans-nerdfont-regular.ttf", fontsize) -def draw_grid_annotations(im, width, height, hor_texts, ver_texts, margin=0, title=None): +def draw_grid_annotations(im, width, height, x_texts, y_texts, margin=0, title=None): def wrap(drawing, text, font, line_length): lines = [''] for word in text.split(): @@ -140,15 +140,15 @@ def draw_grid_annotations(im, width, height, hor_texts, ver_texts, margin=0, tit line_spacing = fontsize // 2 font = get_font(fontsize) color_inactive = (127, 127, 127) - pad_left = 0 if sum([sum([len(line.text) for line in lines]) for lines in ver_texts]) == 0 else width * 3 // 4 - cols = len(hor_texts) - rows = len(ver_texts) + pad_left = 0 if sum([sum([len(line.text) for line in lines]) for lines in y_texts]) == 0 else width * 3 // 4 + cols = len(x_texts) + rows = len(y_texts) # assert cols == len(hor_texts), f'bad number of horizontal texts: {len(hor_texts)}; must be {cols}' # assert rows == len(hor_texts), f'bad number of vertical texts: {len(ver_texts)}; must be {rows}' calc_img = Image.new("RGB", (1, 1), shared.opts.grid_background) calc_d = ImageDraw.Draw(calc_img) title_texts = [title] if title else [[GridAnnotation()]] - for texts, allowed_width in zip(hor_texts + ver_texts + title_texts, [width] * len(hor_texts) + [pad_left] * len(ver_texts) + [(width+margin)*cols]): + for texts, allowed_width in zip(x_texts + y_texts + title_texts, [width] * len(x_texts) + [pad_left] * len(y_texts) + [(width+margin)*cols]): items = [] + texts texts.clear() for line in items: @@ -158,8 +158,8 @@ def draw_grid_annotations(im, width, height, hor_texts, ver_texts, margin=0, tit bbox = calc_d.multiline_textbbox((0, 0), line.text, font=font) line.size = (bbox[2] - bbox[0], bbox[3] - bbox[1]) line.allowed_width = allowed_width - hor_text_heights = [sum([line.size[1] + line_spacing for line in lines]) - line_spacing for lines in hor_texts] - ver_text_heights = [sum([line.size[1] + line_spacing for line in lines]) - line_spacing * len(lines) for lines in ver_texts] + hor_text_heights = [sum([line.size[1] + line_spacing for line in lines]) - line_spacing for lines in x_texts] + ver_text_heights = [sum([line.size[1] + line_spacing for line in lines]) - line_spacing * len(lines) for lines in y_texts] pad_top = 0 if sum(hor_text_heights) == 0 else max(hor_text_heights) + line_spacing * 2 title_pad = 0 if title: @@ -178,11 +178,11 @@ def draw_grid_annotations(im, width, height, hor_texts, ver_texts, margin=0, tit for col in range(cols): x = pad_left + (width + margin) * col + width / 2 y = (pad_top / 2 - hor_text_heights[col] / 2) + title_pad - draw_texts(d, x, y, hor_texts[col], font, fontsize) + draw_texts(d, x, y, x_texts[col], font, fontsize) for row in range(rows): x = pad_left / 2 y = (pad_top + (height + margin) * row + height / 2 - ver_text_heights[row] / 2) + title_pad - draw_texts(d, x, y, ver_texts[row], font, fontsize) + draw_texts(d, x, y, y_texts[row], font, fontsize) return result diff --git a/modules/img2img.py b/modules/img2img.py index faf65161a..f3bec5b02 100644 --- a/modules/img2img.py +++ b/modules/img2img.py @@ -267,6 +267,7 @@ def img2img(id_task: str, state: str, mode: int, processed = modules.scripts.scripts_img2img.run(p, *args) if processed is None: processed = processing.process_images(p) + processed = modules.scripts.scripts_img2img.after(p, processed, *args) p.close() generation_info_js = processed.js() if processed is not None else '' if processed is None: diff --git a/modules/instantir/__init__.py b/modules/instantir/__init__.py new file mode 100644 index 000000000..fdd6cf8a0 --- /dev/null +++ b/modules/instantir/__init__.py @@ -0,0 +1,3 @@ +from .sdxl_instantir import InstantIRPipeline +from .lcm_single_step_scheduler import LCMSingleStepScheduler +from .ip_adapter.utils import init_adapter_in_unet, load_adapter_to_pipe diff --git a/modules/instantir/aggregator.py b/modules/instantir/aggregator.py new file mode 100644 index 000000000..fd6151003 --- /dev/null +++ b/modules/instantir/aggregator.py @@ -0,0 +1,983 @@ +from dataclasses import dataclass +from typing import Any, Dict, List, Optional, Tuple, Union + +import torch +from torch import nn +from torch.nn import functional as F + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.loaders.single_file_model import FromOriginalModelMixin +from diffusers.utils import BaseOutput, logging +from diffusers.models.attention_processor import ( + ADDED_KV_ATTENTION_PROCESSORS, + CROSS_ATTENTION_PROCESSORS, + AttentionProcessor, + AttnAddedKVProcessor, + AttnProcessor, +) +from diffusers.models.embeddings import TextImageProjection, TextImageTimeEmbedding, TextTimeEmbedding, TimestepEmbedding, Timesteps +from diffusers.models.modeling_utils import ModelMixin +from diffusers.models.unets.unet_2d_blocks import ( + CrossAttnDownBlock2D, + DownBlock2D, + UNetMidBlock2D, + UNetMidBlock2DCrossAttn, + get_down_block, +) +from diffusers.models.unets.unet_2d_condition import UNet2DConditionModel + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +class ZeroConv(nn.Module): + def __init__(self, label_nc, norm_nc, mask=False): + super().__init__() + self.zero_conv = zero_module(nn.Conv2d(label_nc+norm_nc, norm_nc, 1, 1, 0)) + self.mask = mask + + def forward(self, hidden_states, h_ori=None): + # with torch.cuda.amp.autocast(enabled=False, dtype=torch.float32): + c, h = hidden_states + if not self.mask: + h = self.zero_conv(torch.cat([c, h], dim=1)) + else: + h = self.zero_conv(torch.cat([c, h], dim=1)) * torch.zeros_like(h) + if h_ori is not None: + h = torch.cat([h_ori, h], dim=1) + return h + + +class SFT(nn.Module): + def __init__(self, label_nc, norm_nc, mask=False): + super().__init__() + + # param_free_norm_type = str(parsed.group(1)) + ks = 3 + pw = ks // 2 + + self.mask = mask + + nhidden = 128 + + self.mlp_shared = nn.Sequential( + nn.Conv2d(label_nc, nhidden, kernel_size=ks, padding=pw), + nn.SiLU() + ) + self.mul = nn.Conv2d(nhidden, norm_nc, kernel_size=ks, padding=pw) + self.add = nn.Conv2d(nhidden, norm_nc, kernel_size=ks, padding=pw) + + def forward(self, hidden_states, mask=False): + + c, h = hidden_states + mask = mask or self.mask + assert mask is False + + actv = self.mlp_shared(c) + gamma = self.mul(actv) + beta = self.add(actv) + + if self.mask: + gamma = gamma * torch.zeros_like(gamma) + beta = beta * torch.zeros_like(beta) + # gamma_ori, gamma_res = torch.split(gamma, [h_ori_c, h_c], dim=1) + # beta_ori, beta_res = torch.split(beta, [h_ori_c, h_c], dim=1) + # print(gamma_ori.mean(), gamma_res.mean(), beta_ori.mean(), beta_res.mean()) + h = h * (gamma + 1) + beta + # sample_ori, sample_res = torch.split(h, [h_ori_c, h_c], dim=1) + # print(sample_ori.mean(), sample_res.mean()) + + return h + + +@dataclass +class AggregatorOutput(BaseOutput): + """ + The output of [`Aggregator`]. + + Args: + down_block_res_samples (`tuple[torch.Tensor]`): + A tuple of downsample activations at different resolutions for each downsampling block. Each tensor should + be of shape `(batch_size, channel * resolution, height //resolution, width // resolution)`. Output can be + used to condition the original UNet's downsampling activations. + mid_down_block_re_sample (`torch.Tensor`): + The activation of the midde block (the lowest sample resolution). Each tensor should be of shape + `(batch_size, channel * lowest_resolution, height // lowest_resolution, width // lowest_resolution)`. + Output can be used to condition the original UNet's middle block activation. + """ + + down_block_res_samples: Tuple[torch.Tensor] + mid_block_res_sample: torch.Tensor + + +class ConditioningEmbedding(nn.Module): + """ + Quoting from https://arxiv.org/abs/2302.05543: "Stable Diffusion uses a pre-processing method similar to VQ-GAN + [11] to convert the entire dataset of 512 × 512 images into smaller 64 × 64 “latent images” for stabilized + training. This requires ControlNets to convert image-based conditions to 64 × 64 feature space to match the + convolution size. We use a tiny network E(·) of four convolution layers with 4 × 4 kernels and 2 × 2 strides + (activated by ReLU, channels are 16, 32, 64, 128, initialized with Gaussian weights, trained jointly with the full + model) to encode image-space conditions ... into feature maps ..." + """ + + def __init__( + self, + conditioning_embedding_channels: int, + conditioning_channels: int = 3, + block_out_channels: Tuple[int, ...] = (16, 32, 96, 256), + ): + super().__init__() + + self.conv_in = nn.Conv2d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1) + + self.blocks = nn.ModuleList([]) + + for i in range(len(block_out_channels) - 1): + channel_in = block_out_channels[i] + channel_out = block_out_channels[i + 1] + self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1)) + self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2)) + + self.conv_out = zero_module( + nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1) + ) + + def forward(self, conditioning): + embedding = self.conv_in(conditioning) + embedding = F.silu(embedding) + + for block in self.blocks: + embedding = block(embedding) + embedding = F.silu(embedding) + + embedding = self.conv_out(embedding) + + return embedding + + +class Aggregator(ModelMixin, ConfigMixin, FromOriginalModelMixin): + """ + Aggregator model. + + Args: + in_channels (`int`, defaults to 4): + The number of channels in the input sample. + flip_sin_to_cos (`bool`, defaults to `True`): + Whether to flip the sin to cos in the time embedding. + freq_shift (`int`, defaults to 0): + The frequency shift to apply to the time embedding. + down_block_types (`tuple[str]`, defaults to `("CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "CrossAttnDownBlock2D", "DownBlock2D")`): + The tuple of downsample blocks to use. + only_cross_attention (`Union[bool, Tuple[bool]]`, defaults to `False`): + block_out_channels (`tuple[int]`, defaults to `(320, 640, 1280, 1280)`): + The tuple of output channels for each block. + layers_per_block (`int`, defaults to 2): + The number of layers per block. + downsample_padding (`int`, defaults to 1): + The padding to use for the downsampling convolution. + mid_block_scale_factor (`float`, defaults to 1): + The scale factor to use for the mid block. + act_fn (`str`, defaults to "silu"): + The activation function to use. + norm_num_groups (`int`, *optional*, defaults to 32): + The number of groups to use for the normalization. If None, normalization and activation layers is skipped + in post-processing. + norm_eps (`float`, defaults to 1e-5): + The epsilon to use for the normalization. + cross_attention_dim (`int`, defaults to 1280): + The dimension of the cross attention features. + transformer_layers_per_block (`int` or `Tuple[int]`, *optional*, defaults to 1): + The number of transformer blocks of type [`~models.attention.BasicTransformerBlock`]. Only relevant for + [`~models.unet_2d_blocks.CrossAttnDownBlock2D`], [`~models.unet_2d_blocks.CrossAttnUpBlock2D`], + [`~models.unet_2d_blocks.UNetMidBlock2DCrossAttn`]. + encoder_hid_dim (`int`, *optional*, defaults to None): + If `encoder_hid_dim_type` is defined, `encoder_hidden_states` will be projected from `encoder_hid_dim` + dimension to `cross_attention_dim`. + encoder_hid_dim_type (`str`, *optional*, defaults to `None`): + If given, the `encoder_hidden_states` and potentially other embeddings are down-projected to text + embeddings of dimension `cross_attention` according to `encoder_hid_dim_type`. + attention_head_dim (`Union[int, Tuple[int]]`, defaults to 8): + The dimension of the attention heads. + use_linear_projection (`bool`, defaults to `False`): + class_embed_type (`str`, *optional*, defaults to `None`): + The type of class embedding to use which is ultimately summed with the time embeddings. Choose from None, + `"timestep"`, `"identity"`, `"projection"`, or `"simple_projection"`. + addition_embed_type (`str`, *optional*, defaults to `None`): + Configures an optional embedding which will be summed with the time embeddings. Choose from `None` or + "text". "text" will use the `TextTimeEmbedding` layer. + num_class_embeds (`int`, *optional*, defaults to 0): + Input dimension of the learnable embedding matrix to be projected to `time_embed_dim`, when performing + class conditioning with `class_embed_type` equal to `None`. + upcast_attention (`bool`, defaults to `False`): + resnet_time_scale_shift (`str`, defaults to `"default"`): + Time scale shift config for ResNet blocks (see `ResnetBlock2D`). Choose from `default` or `scale_shift`. + projection_class_embeddings_input_dim (`int`, *optional*, defaults to `None`): + The dimension of the `class_labels` input when `class_embed_type="projection"`. Required when + `class_embed_type="projection"`. + controlnet_conditioning_channel_order (`str`, defaults to `"rgb"`): + The channel order of conditional image. Will convert to `rgb` if it's `bgr`. + conditioning_embedding_out_channels (`tuple[int]`, *optional*, defaults to `(16, 32, 96, 256)`): + The tuple of output channel for each block in the `conditioning_embedding` layer. + global_pool_conditions (`bool`, defaults to `False`): + TODO(Patrick) - unused parameter. + addition_embed_type_num_heads (`int`, defaults to 64): + The number of heads to use for the `TextTimeEmbedding` layer. + """ + + _supports_gradient_checkpointing = True + + @register_to_config + def __init__( + self, + in_channels: int = 4, + conditioning_channels: int = 3, + flip_sin_to_cos: bool = True, + freq_shift: int = 0, + down_block_types: Tuple[str, ...] = ( + "CrossAttnDownBlock2D", + "CrossAttnDownBlock2D", + "CrossAttnDownBlock2D", + "DownBlock2D", + ), + mid_block_type: Optional[str] = "UNetMidBlock2DCrossAttn", + only_cross_attention: Union[bool, Tuple[bool]] = False, + block_out_channels: Tuple[int, ...] = (320, 640, 1280, 1280), + layers_per_block: int = 2, + downsample_padding: int = 1, + mid_block_scale_factor: float = 1, + act_fn: str = "silu", + norm_num_groups: Optional[int] = 32, + norm_eps: float = 1e-5, + cross_attention_dim: int = 1280, + transformer_layers_per_block: Union[int, Tuple[int, ...]] = 1, + encoder_hid_dim: Optional[int] = None, + encoder_hid_dim_type: Optional[str] = None, + attention_head_dim: Union[int, Tuple[int, ...]] = 8, + num_attention_heads: Optional[Union[int, Tuple[int, ...]]] = None, + use_linear_projection: bool = False, + class_embed_type: Optional[str] = None, + addition_embed_type: Optional[str] = None, + addition_time_embed_dim: Optional[int] = None, + num_class_embeds: Optional[int] = None, + upcast_attention: bool = False, + resnet_time_scale_shift: str = "default", + projection_class_embeddings_input_dim: Optional[int] = None, + controlnet_conditioning_channel_order: str = "rgb", + conditioning_embedding_out_channels: Optional[Tuple[int, ...]] = (16, 32, 96, 256), + global_pool_conditions: bool = False, + addition_embed_type_num_heads: int = 64, + pad_concat: bool = False, + ): + super().__init__() + + # If `num_attention_heads` is not defined (which is the case for most models) + # it will default to `attention_head_dim`. This looks weird upon first reading it and it is. + # The reason for this behavior is to correct for incorrectly named variables that were introduced + # when this library was created. The incorrect naming was only discovered much later in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131 + # Changing `attention_head_dim` to `num_attention_heads` for 40,000+ configurations is too backwards breaking + # which is why we correct for the naming here. + num_attention_heads = num_attention_heads or attention_head_dim + self.pad_concat = pad_concat + + # Check inputs + if len(block_out_channels) != len(down_block_types): + raise ValueError( + f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}." + ) + + if not isinstance(only_cross_attention, bool) and len(only_cross_attention) != len(down_block_types): + raise ValueError( + f"Must provide the same number of `only_cross_attention` as `down_block_types`. `only_cross_attention`: {only_cross_attention}. `down_block_types`: {down_block_types}." + ) + + if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types): + raise ValueError( + f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}." + ) + + if isinstance(transformer_layers_per_block, int): + transformer_layers_per_block = [transformer_layers_per_block] * len(down_block_types) + + # input + conv_in_kernel = 3 + conv_in_padding = (conv_in_kernel - 1) // 2 + self.conv_in = nn.Conv2d( + in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding + ) + + # time + time_embed_dim = block_out_channels[0] * 4 + self.time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos, freq_shift) + timestep_input_dim = block_out_channels[0] + self.time_embedding = TimestepEmbedding( + timestep_input_dim, + time_embed_dim, + act_fn=act_fn, + ) + + if encoder_hid_dim_type is None and encoder_hid_dim is not None: + encoder_hid_dim_type = "text_proj" + self.register_to_config(encoder_hid_dim_type=encoder_hid_dim_type) + logger.info("encoder_hid_dim_type defaults to 'text_proj' as `encoder_hid_dim` is defined.") + + if encoder_hid_dim is None and encoder_hid_dim_type is not None: + raise ValueError( + f"`encoder_hid_dim` has to be defined when `encoder_hid_dim_type` is set to {encoder_hid_dim_type}." + ) + + if encoder_hid_dim_type == "text_proj": + self.encoder_hid_proj = nn.Linear(encoder_hid_dim, cross_attention_dim) + elif encoder_hid_dim_type == "text_image_proj": + # image_embed_dim DOESN'T have to be `cross_attention_dim`. To not clutter the __init__ too much + # they are set to `cross_attention_dim` here as this is exactly the required dimension for the currently only use + # case when `addition_embed_type == "text_image_proj"` (Kandinsky 2.1)` + self.encoder_hid_proj = TextImageProjection( + text_embed_dim=encoder_hid_dim, + image_embed_dim=cross_attention_dim, + cross_attention_dim=cross_attention_dim, + ) + + elif encoder_hid_dim_type is not None: + raise ValueError( + f"encoder_hid_dim_type: {encoder_hid_dim_type} must be None, 'text_proj' or 'text_image_proj'." + ) + else: + self.encoder_hid_proj = None + + # class embedding + if class_embed_type is None and num_class_embeds is not None: + self.class_embedding = nn.Embedding(num_class_embeds, time_embed_dim) + elif class_embed_type == "timestep": + self.class_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim) + elif class_embed_type == "identity": + self.class_embedding = nn.Identity(time_embed_dim, time_embed_dim) + elif class_embed_type == "projection": + if projection_class_embeddings_input_dim is None: + raise ValueError( + "`class_embed_type`: 'projection' requires `projection_class_embeddings_input_dim` be set" + ) + # The projection `class_embed_type` is the same as the timestep `class_embed_type` except + # 1. the `class_labels` inputs are not first converted to sinusoidal embeddings + # 2. it projects from an arbitrary input dimension. + # + # Note that `TimestepEmbedding` is quite general, being mainly linear layers and activations. + # When used for embedding actual timesteps, the timesteps are first converted to sinusoidal embeddings. + # As a result, `TimestepEmbedding` can be passed arbitrary vectors. + self.class_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) + else: + self.class_embedding = None + + if addition_embed_type == "text": + if encoder_hid_dim is not None: + text_time_embedding_from_dim = encoder_hid_dim + else: + text_time_embedding_from_dim = cross_attention_dim + + self.add_embedding = TextTimeEmbedding( + text_time_embedding_from_dim, time_embed_dim, num_heads=addition_embed_type_num_heads + ) + elif addition_embed_type == "text_image": + # text_embed_dim and image_embed_dim DON'T have to be `cross_attention_dim`. To not clutter the __init__ too much + # they are set to `cross_attention_dim` here as this is exactly the required dimension for the currently only use + # case when `addition_embed_type == "text_image"` (Kandinsky 2.1)` + self.add_embedding = TextImageTimeEmbedding( + text_embed_dim=cross_attention_dim, image_embed_dim=cross_attention_dim, time_embed_dim=time_embed_dim + ) + elif addition_embed_type == "text_time": + self.add_time_proj = Timesteps(addition_time_embed_dim, flip_sin_to_cos, freq_shift) + self.add_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim) + + elif addition_embed_type is not None: + raise ValueError(f"addition_embed_type: {addition_embed_type} must be None, 'text' or 'text_image'.") + + # control net conditioning embedding + self.ref_conv_in = nn.Conv2d( + in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding + ) + + self.down_blocks = nn.ModuleList([]) + self.controlnet_down_blocks = nn.ModuleList([]) + + if isinstance(only_cross_attention, bool): + only_cross_attention = [only_cross_attention] * len(down_block_types) + + if isinstance(attention_head_dim, int): + attention_head_dim = (attention_head_dim,) * len(down_block_types) + + if isinstance(num_attention_heads, int): + num_attention_heads = (num_attention_heads,) * len(down_block_types) + + # down + output_channel = block_out_channels[0] + + # controlnet_block = ZeroConv(output_channel, output_channel) + controlnet_block = nn.Sequential( + SFT(output_channel, output_channel), + zero_module(nn.Conv2d(output_channel, output_channel, kernel_size=1)) + ) + self.controlnet_down_blocks.append(controlnet_block) + + for i, down_block_type in enumerate(down_block_types): + input_channel = output_channel + output_channel = block_out_channels[i] + is_final_block = i == len(block_out_channels) - 1 + + down_block = get_down_block( + down_block_type, + num_layers=layers_per_block, + transformer_layers_per_block=transformer_layers_per_block[i], + in_channels=input_channel, + out_channels=output_channel, + temb_channels=time_embed_dim, + add_downsample=not is_final_block, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + resnet_groups=norm_num_groups, + cross_attention_dim=cross_attention_dim, + num_attention_heads=num_attention_heads[i], + attention_head_dim=attention_head_dim[i] if attention_head_dim[i] is not None else output_channel, + downsample_padding=downsample_padding, + use_linear_projection=use_linear_projection, + only_cross_attention=only_cross_attention[i], + upcast_attention=upcast_attention, + resnet_time_scale_shift=resnet_time_scale_shift, + ) + self.down_blocks.append(down_block) + + for _ in range(layers_per_block): + # controlnet_block = ZeroConv(output_channel, output_channel) + controlnet_block = nn.Sequential( + SFT(output_channel, output_channel), + zero_module(nn.Conv2d(output_channel, output_channel, kernel_size=1)) + ) + self.controlnet_down_blocks.append(controlnet_block) + + if not is_final_block: + # controlnet_block = ZeroConv(output_channel, output_channel) + controlnet_block = nn.Sequential( + SFT(output_channel, output_channel), + zero_module(nn.Conv2d(output_channel, output_channel, kernel_size=1)) + ) + self.controlnet_down_blocks.append(controlnet_block) + + # mid + mid_block_channel = block_out_channels[-1] + + # controlnet_block = ZeroConv(mid_block_channel, mid_block_channel) + controlnet_block = nn.Sequential( + SFT(mid_block_channel, mid_block_channel), + zero_module(nn.Conv2d(mid_block_channel, mid_block_channel, kernel_size=1)) + ) + self.controlnet_mid_block = controlnet_block + + if mid_block_type == "UNetMidBlock2DCrossAttn": + self.mid_block = UNetMidBlock2DCrossAttn( + transformer_layers_per_block=transformer_layers_per_block[-1], + in_channels=mid_block_channel, + temb_channels=time_embed_dim, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + output_scale_factor=mid_block_scale_factor, + resnet_time_scale_shift=resnet_time_scale_shift, + cross_attention_dim=cross_attention_dim, + num_attention_heads=num_attention_heads[-1], + resnet_groups=norm_num_groups, + use_linear_projection=use_linear_projection, + upcast_attention=upcast_attention, + ) + elif mid_block_type == "UNetMidBlock2D": + self.mid_block = UNetMidBlock2D( + in_channels=block_out_channels[-1], + temb_channels=time_embed_dim, + num_layers=0, + resnet_eps=norm_eps, + resnet_act_fn=act_fn, + output_scale_factor=mid_block_scale_factor, + resnet_groups=norm_num_groups, + resnet_time_scale_shift=resnet_time_scale_shift, + add_attention=False, + ) + else: + raise ValueError(f"unknown mid_block_type : {mid_block_type}") + + @classmethod + def from_unet( + cls, + unet: UNet2DConditionModel, + controlnet_conditioning_channel_order: str = "rgb", + conditioning_embedding_out_channels: Optional[Tuple[int, ...]] = (16, 32, 96, 256), + load_weights_from_unet: bool = True, + conditioning_channels: int = 3, + ): + r""" + Instantiate a [`ControlNetModel`] from [`UNet2DConditionModel`]. + + Parameters: + unet (`UNet2DConditionModel`): + The UNet model weights to copy to the [`ControlNetModel`]. All configuration options are also copied + where applicable. + """ + transformer_layers_per_block = ( + unet.config.transformer_layers_per_block if "transformer_layers_per_block" in unet.config else 1 + ) + encoder_hid_dim = unet.config.encoder_hid_dim if "encoder_hid_dim" in unet.config else None + encoder_hid_dim_type = unet.config.encoder_hid_dim_type if "encoder_hid_dim_type" in unet.config else None + addition_embed_type = unet.config.addition_embed_type if "addition_embed_type" in unet.config else None + addition_time_embed_dim = ( + unet.config.addition_time_embed_dim if "addition_time_embed_dim" in unet.config else None + ) + + controlnet = cls( + encoder_hid_dim=encoder_hid_dim, + encoder_hid_dim_type=encoder_hid_dim_type, + addition_embed_type=addition_embed_type, + addition_time_embed_dim=addition_time_embed_dim, + transformer_layers_per_block=transformer_layers_per_block, + in_channels=unet.config.in_channels, + flip_sin_to_cos=unet.config.flip_sin_to_cos, + freq_shift=unet.config.freq_shift, + down_block_types=unet.config.down_block_types, + only_cross_attention=unet.config.only_cross_attention, + block_out_channels=unet.config.block_out_channels, + layers_per_block=unet.config.layers_per_block, + downsample_padding=unet.config.downsample_padding, + mid_block_scale_factor=unet.config.mid_block_scale_factor, + act_fn=unet.config.act_fn, + norm_num_groups=unet.config.norm_num_groups, + norm_eps=unet.config.norm_eps, + cross_attention_dim=unet.config.cross_attention_dim, + attention_head_dim=unet.config.attention_head_dim, + num_attention_heads=unet.config.num_attention_heads, + use_linear_projection=unet.config.use_linear_projection, + class_embed_type=unet.config.class_embed_type, + num_class_embeds=unet.config.num_class_embeds, + upcast_attention=unet.config.upcast_attention, + resnet_time_scale_shift=unet.config.resnet_time_scale_shift, + projection_class_embeddings_input_dim=unet.config.projection_class_embeddings_input_dim, + mid_block_type=unet.config.mid_block_type, + controlnet_conditioning_channel_order=controlnet_conditioning_channel_order, + conditioning_embedding_out_channels=conditioning_embedding_out_channels, + conditioning_channels=conditioning_channels, + ) + + if load_weights_from_unet: + controlnet.conv_in.load_state_dict(unet.conv_in.state_dict()) + controlnet.ref_conv_in.load_state_dict(unet.conv_in.state_dict()) + controlnet.time_proj.load_state_dict(unet.time_proj.state_dict()) + controlnet.time_embedding.load_state_dict(unet.time_embedding.state_dict()) + + if controlnet.class_embedding: + controlnet.class_embedding.load_state_dict(unet.class_embedding.state_dict()) + + if hasattr(controlnet, "add_embedding"): + controlnet.add_embedding.load_state_dict(unet.add_embedding.state_dict()) + + controlnet.down_blocks.load_state_dict(unet.down_blocks.state_dict()) + controlnet.mid_block.load_state_dict(unet.mid_block.state_dict()) + + return controlnet + + @property + # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors + def attn_processors(self) -> Dict[str, AttentionProcessor]: + r""" + Returns: + `dict` of attention processors: A dictionary containing all attention processors used in the model with + indexed by its weight name. + """ + # set recursively + processors = {} + + def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: Dict[str, AttentionProcessor]): + if hasattr(module, "get_processor"): + processors[f"{name}.processor"] = module.get_processor(return_deprecated_lora=True) + + for sub_name, child in module.named_children(): + fn_recursive_add_processors(f"{name}.{sub_name}", child, processors) + + return processors + + for name, module in self.named_children(): + fn_recursive_add_processors(name, module, processors) + + return processors + + # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor + def set_attn_processor(self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]): + r""" + Sets the attention processor to use to compute attention. + + Parameters: + processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`): + The instantiated processor class or a dictionary of processor classes that will be set as the processor + for **all** `Attention` layers. + + If `processor` is a dict, the key needs to define the path to the corresponding cross attention + processor. This is strongly recommended when setting trainable attention processors. + + """ + count = len(self.attn_processors.keys()) + + if isinstance(processor, dict) and len(processor) != count: + raise ValueError( + f"A dict of processors was passed, but the number of processors {len(processor)} does not match the" + f" number of attention layers: {count}. Please make sure to pass {count} processor classes." + ) + + def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor): + if hasattr(module, "set_processor"): + if not isinstance(processor, dict): + module.set_processor(processor) + else: + module.set_processor(processor.pop(f"{name}.processor")) + + for sub_name, child in module.named_children(): + fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor) + + for name, module in self.named_children(): + fn_recursive_attn_processor(name, module, processor) + + # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_default_attn_processor + def set_default_attn_processor(self): + """ + Disables custom attention processors and sets the default attention implementation. + """ + if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): + processor = AttnAddedKVProcessor() + elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()): + processor = AttnProcessor() + else: + raise ValueError( + f"Cannot call `set_default_attn_processor` when attention processors are of type {next(iter(self.attn_processors.values()))}" + ) + + self.set_attn_processor(processor) + + # Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attention_slice + def set_attention_slice(self, slice_size: Union[str, int, List[int]]) -> None: + r""" + Enable sliced attention computation. + + When this option is enabled, the attention module splits the input tensor in slices to compute attention in + several steps. This is useful for saving some memory in exchange for a small decrease in speed. + + Args: + slice_size (`str` or `int` or `list(int)`, *optional*, defaults to `"auto"`): + When `"auto"`, input to the attention heads is halved, so attention is computed in two steps. If + `"max"`, maximum amount of memory is saved by running only one slice at a time. If a number is + provided, uses as many slices as `attention_head_dim // slice_size`. In this case, `attention_head_dim` + must be a multiple of `slice_size`. + """ + sliceable_head_dims = [] + + def fn_recursive_retrieve_sliceable_dims(module: torch.nn.Module): + if hasattr(module, "set_attention_slice"): + sliceable_head_dims.append(module.sliceable_head_dim) + + for child in module.children(): + fn_recursive_retrieve_sliceable_dims(child) + + # retrieve number of attention layers + for module in self.children(): + fn_recursive_retrieve_sliceable_dims(module) + + num_sliceable_layers = len(sliceable_head_dims) + + if slice_size == "auto": + # half the attention head size is usually a good trade-off between + # speed and memory + slice_size = [dim // 2 for dim in sliceable_head_dims] + elif slice_size == "max": + # make smallest slice possible + slice_size = num_sliceable_layers * [1] + + slice_size = num_sliceable_layers * [slice_size] if not isinstance(slice_size, list) else slice_size + + if len(slice_size) != len(sliceable_head_dims): + raise ValueError( + f"You have provided {len(slice_size)}, but {self.config} has {len(sliceable_head_dims)} different" + f" attention layers. Make sure to match `len(slice_size)` to be {len(sliceable_head_dims)}." + ) + + for i in range(len(slice_size)): + size = slice_size[i] + dim = sliceable_head_dims[i] + if size is not None and size > dim: + raise ValueError(f"size {size} has to be smaller or equal to {dim}.") + + # Recursively walk through all the children. + # Any children which exposes the set_attention_slice method + # gets the message + def fn_recursive_set_attention_slice(module: torch.nn.Module, slice_size: List[int]): + if hasattr(module, "set_attention_slice"): + module.set_attention_slice(slice_size.pop()) + + for child in module.children(): + fn_recursive_set_attention_slice(child, slice_size) + + reversed_slice_size = list(reversed(slice_size)) + for module in self.children(): + fn_recursive_set_attention_slice(module, reversed_slice_size) + + def process_encoder_hidden_states( + self, encoder_hidden_states: torch.Tensor, added_cond_kwargs: Dict[str, Any] + ) -> torch.Tensor: + if self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "text_proj": + encoder_hidden_states = self.encoder_hid_proj(encoder_hidden_states) + elif self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "text_image_proj": + # Kandinsky 2.1 - style + if "image_embeds" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'text_image_proj' which requires the keyword argument `image_embeds` to be passed in `added_conditions`" + ) + + image_embeds = added_cond_kwargs.get("image_embeds") + encoder_hidden_states = self.encoder_hid_proj(encoder_hidden_states, image_embeds) + elif self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "image_proj": + # Kandinsky 2.2 - style + if "image_embeds" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'image_proj' which requires the keyword argument `image_embeds` to be passed in `added_conditions`" + ) + image_embeds = added_cond_kwargs.get("image_embeds") + encoder_hidden_states = self.encoder_hid_proj(image_embeds) + elif self.encoder_hid_proj is not None and self.config.encoder_hid_dim_type == "ip_image_proj": + if "image_embeds" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `encoder_hid_dim_type` set to 'ip_image_proj' which requires the keyword argument `image_embeds` to be passed in `added_conditions`" + ) + image_embeds = added_cond_kwargs.get("image_embeds") + image_embeds = self.encoder_hid_proj(image_embeds) + encoder_hidden_states = (encoder_hidden_states, image_embeds) + return encoder_hidden_states + + def _set_gradient_checkpointing(self, module, value: bool = False) -> None: + if isinstance(module, (CrossAttnDownBlock2D, DownBlock2D)): + module.gradient_checkpointing = value + + def forward( + self, + sample: torch.FloatTensor, + timestep: Union[torch.Tensor, float, int], + encoder_hidden_states: torch.Tensor, + controlnet_cond: torch.FloatTensor, + cat_dim: int = -2, + conditioning_scale: float = 1.0, + class_labels: Optional[torch.Tensor] = None, + timestep_cond: Optional[torch.Tensor] = None, + attention_mask: Optional[torch.Tensor] = None, + added_cond_kwargs: Optional[Dict[str, torch.Tensor]] = None, + cross_attention_kwargs: Optional[Dict[str, Any]] = None, + return_dict: bool = True, + ) -> Union[AggregatorOutput, Tuple[Tuple[torch.FloatTensor, ...], torch.FloatTensor]]: + """ + The [`Aggregator`] forward method. + + Args: + sample (`torch.FloatTensor`): + The noisy input tensor. + timestep (`Union[torch.Tensor, float, int]`): + The number of timesteps to denoise an input. + encoder_hidden_states (`torch.Tensor`): + The encoder hidden states. + controlnet_cond (`torch.FloatTensor`): + The conditional input tensor of shape `(batch_size, sequence_length, hidden_size)`. + conditioning_scale (`float`, defaults to `1.0`): + The scale factor for ControlNet outputs. + class_labels (`torch.Tensor`, *optional*, defaults to `None`): + Optional class labels for conditioning. Their embeddings will be summed with the timestep embeddings. + timestep_cond (`torch.Tensor`, *optional*, defaults to `None`): + Additional conditional embeddings for timestep. If provided, the embeddings will be summed with the + timestep_embedding passed through the `self.time_embedding` layer to obtain the final timestep + embeddings. + attention_mask (`torch.Tensor`, *optional*, defaults to `None`): + An attention mask of shape `(batch, key_tokens)` is applied to `encoder_hidden_states`. If `1` the mask + is kept, otherwise if `0` it is discarded. Mask will be converted into a bias, which adds large + negative values to the attention scores corresponding to "discard" tokens. + added_cond_kwargs (`dict`): + Additional conditions for the Stable Diffusion XL UNet. + cross_attention_kwargs (`dict[str]`, *optional*, defaults to `None`): + A kwargs dictionary that if specified is passed along to the `AttnProcessor`. + return_dict (`bool`, defaults to `True`): + Whether or not to return a [`~models.controlnet.ControlNetOutput`] instead of a plain tuple. + + Returns: + [`~models.controlnet.ControlNetOutput`] **or** `tuple`: + If `return_dict` is `True`, a [`~models.controlnet.ControlNetOutput`] is returned, otherwise a tuple is + returned where the first element is the sample tensor. + """ + # check channel order + channel_order = self.config.controlnet_conditioning_channel_order + + if channel_order == "rgb": + # in rgb order by default + ... + else: + raise ValueError(f"unknown `controlnet_conditioning_channel_order`: {channel_order}") + + # prepare attention_mask + if attention_mask is not None: + attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0 + attention_mask = attention_mask.unsqueeze(1) + + # 1. time + timesteps = timestep + if not torch.is_tensor(timesteps): + # TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can + # This would be a good case for the `match` statement (Python 3.10+) + is_mps = sample.device.type == "mps" + if isinstance(timestep, float): + dtype = torch.float32 if is_mps else torch.float64 + else: + dtype = torch.int32 if is_mps else torch.int64 + timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device) + elif len(timesteps.shape) == 0: + timesteps = timesteps[None].to(sample.device) + + # broadcast to batch dimension in a way that's compatible with ONNX/Core ML + timesteps = timesteps.expand(sample.shape[0]) + + t_emb = self.time_proj(timesteps) + + # timesteps does not contain any weights and will always return f32 tensors + # but time_embedding might actually be running in fp16. so we need to cast here. + # there might be better ways to encapsulate this. + t_emb = t_emb.to(dtype=sample.dtype) + + emb = self.time_embedding(t_emb, timestep_cond) + aug_emb = None + + if self.class_embedding is not None: + if class_labels is None: + raise ValueError("class_labels should be provided when num_class_embeds > 0") + + if self.config.class_embed_type == "timestep": + class_labels = self.time_proj(class_labels) + + class_emb = self.class_embedding(class_labels).to(dtype=self.dtype) + emb = emb + class_emb + + if self.config.addition_embed_type is not None: + if self.config.addition_embed_type == "text": + aug_emb = self.add_embedding(encoder_hidden_states) + + elif self.config.addition_embed_type == "text_time": + if "text_embeds" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `text_embeds` to be passed in `added_cond_kwargs`" + ) + text_embeds = added_cond_kwargs.get("text_embeds") + if "time_ids" not in added_cond_kwargs: + raise ValueError( + f"{self.__class__} has the config param `addition_embed_type` set to 'text_time' which requires the keyword argument `time_ids` to be passed in `added_cond_kwargs`" + ) + time_ids = added_cond_kwargs.get("time_ids") + time_embeds = self.add_time_proj(time_ids.flatten()) + time_embeds = time_embeds.reshape((text_embeds.shape[0], -1)) + + add_embeds = torch.concat([text_embeds, time_embeds], dim=-1) + add_embeds = add_embeds.to(emb.dtype) + aug_emb = self.add_embedding(add_embeds) + + emb = emb + aug_emb if aug_emb is not None else emb + + encoder_hidden_states = self.process_encoder_hidden_states( + encoder_hidden_states=encoder_hidden_states, added_cond_kwargs=added_cond_kwargs + ) + + # 2. prepare input + cond_latent = self.conv_in(sample) + ref_latent = self.ref_conv_in(controlnet_cond) + batch_size, channel, height, width = cond_latent.shape + if self.pad_concat: + if cat_dim == -2 or cat_dim == 2: + concat_pad = torch.zeros(batch_size, channel, 1, width) + elif cat_dim == -1 or cat_dim == 3: + concat_pad = torch.zeros(batch_size, channel, height, 1) + else: + raise ValueError(f"Aggregator shall concat along spatial dimension, but is asked to concat dim: {cat_dim}.") + concat_pad = concat_pad.to(cond_latent.device, dtype=cond_latent.dtype) + sample = torch.cat([cond_latent, concat_pad, ref_latent], dim=cat_dim) + else: + sample = torch.cat([cond_latent, ref_latent], dim=cat_dim) + + # 3. down + down_block_res_samples = (sample,) + for downsample_block in self.down_blocks: + sample, res_samples = downsample_block( + hidden_states=sample, + temb=emb, + cross_attention_kwargs=cross_attention_kwargs, + ) + + # rebuild sample: split and concat + if self.pad_concat: + batch_size, channel, height, width = sample.shape + if cat_dim == -2 or cat_dim == 2: + cond_latent = sample[:, :, :height//2, :] + ref_latent = sample[:, :, -(height//2):, :] + concat_pad = torch.zeros(batch_size, channel, 1, width) + elif cat_dim == -1 or cat_dim == 3: + cond_latent = sample[:, :, :, :width//2] + ref_latent = sample[:, :, :, -(width//2):] + concat_pad = torch.zeros(batch_size, channel, height, 1) + concat_pad = concat_pad.to(cond_latent.device, dtype=cond_latent.dtype) + sample = torch.cat([cond_latent, concat_pad, ref_latent], dim=cat_dim) + res_samples = res_samples[:-1] + (sample,) + + down_block_res_samples += res_samples + + # 4. mid + if self.mid_block is not None: + sample = self.mid_block( + sample, + emb, + cross_attention_kwargs=cross_attention_kwargs, + ) + + # 5. split samples and SFT. + controlnet_down_block_res_samples = () + for down_block_res_sample, controlnet_block in zip(down_block_res_samples, self.controlnet_down_blocks): + batch_size, channel, height, width = down_block_res_sample.shape + if cat_dim == -2 or cat_dim == 2: + cond_latent = down_block_res_sample[:, :, :height//2, :] + ref_latent = down_block_res_sample[:, :, -(height//2):, :] + elif cat_dim == -1 or cat_dim == 3: + cond_latent = down_block_res_sample[:, :, :, :width//2] + ref_latent = down_block_res_sample[:, :, :, -(width//2):] + down_block_res_sample = controlnet_block((cond_latent, ref_latent), ) + controlnet_down_block_res_samples = controlnet_down_block_res_samples + (down_block_res_sample,) + + down_block_res_samples = controlnet_down_block_res_samples + + batch_size, channel, height, width = sample.shape + if cat_dim == -2 or cat_dim == 2: + cond_latent = sample[:, :, :height//2, :] + ref_latent = sample[:, :, -(height//2):, :] + elif cat_dim == -1 or cat_dim == 3: + cond_latent = sample[:, :, :, :width//2] + ref_latent = sample[:, :, :, -(width//2):] + mid_block_res_sample = self.controlnet_mid_block((cond_latent, ref_latent), ) + + # 6. scaling + down_block_res_samples = [sample*conditioning_scale for sample in down_block_res_samples] + mid_block_res_sample = mid_block_res_sample*conditioning_scale + + if self.config.global_pool_conditions: + down_block_res_samples = [ + torch.mean(sample, dim=(2, 3), keepdim=True) for sample in down_block_res_samples + ] + mid_block_res_sample = torch.mean(mid_block_res_sample, dim=(2, 3), keepdim=True) + + if not return_dict: + return (down_block_res_samples, mid_block_res_sample) + + return AggregatorOutput( + down_block_res_samples=down_block_res_samples, mid_block_res_sample=mid_block_res_sample + ) + + +def zero_module(module): + for p in module.parameters(): + nn.init.zeros_(p) + return module diff --git a/modules/instantir/ip_adapter/__init__.py b/modules/instantir/ip_adapter/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/modules/instantir/ip_adapter/attention_processor.py b/modules/instantir/ip_adapter/attention_processor.py new file mode 100644 index 000000000..ed6cf755f --- /dev/null +++ b/modules/instantir/ip_adapter/attention_processor.py @@ -0,0 +1,1467 @@ +# modified from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py +import torch +import torch.nn as nn +import torch.nn.functional as F + +class AdaLayerNorm(nn.Module): + def __init__(self, embedding_dim: int, time_embedding_dim: int = None): + super().__init__() + + if time_embedding_dim is None: + time_embedding_dim = embedding_dim + + self.silu = nn.SiLU() + self.linear = nn.Linear(time_embedding_dim, 2 * embedding_dim, bias=True) + nn.init.zeros_(self.linear.weight) + nn.init.zeros_(self.linear.bias) + + self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6) + + def forward( + self, x: torch.Tensor, timestep_embedding: torch.Tensor + ): + emb = self.linear(self.silu(timestep_embedding)) + shift, scale = emb.view(len(x), 1, -1).chunk(2, dim=-1) + x = self.norm(x) * (1 + scale) + shift + return x + + +class AttnProcessor(nn.Module): + r""" + Default processor for performing attention-related computations. + """ + + def __init__( + self, + hidden_size=None, + cross_attention_dim=None, + ): + super().__init__() + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + temb=None, + ): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + query = attn.head_to_batch_dim(query) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + + attention_probs = attn.get_attention_scores(query, key, attention_mask) + hidden_states = torch.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class IPAttnProcessor(nn.Module): + r""" + Attention processor for IP-Adapater. + Args: + hidden_size (`int`): + The hidden size of the attention layer. + cross_attention_dim (`int`): + The number of channels in the `encoder_hidden_states`. + scale (`float`, defaults to 1.0): + the weight scale of image prompt. + num_tokens (`int`, defaults to 4 when do ip_adapter_plus it should be 16): + The context length of the image features. + """ + + def __init__(self, hidden_size, cross_attention_dim=None, scale=1.0, num_tokens=4): + super().__init__() + + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + self.scale = scale + self.num_tokens = num_tokens + + self.to_k_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + self.to_v_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + temb=None, + ): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + else: + # get encoder_hidden_states, ip_hidden_states + end_pos = encoder_hidden_states.shape[1] - self.num_tokens + encoder_hidden_states, ip_hidden_states = ( + encoder_hidden_states[:, :end_pos, :], + encoder_hidden_states[:, end_pos:, :], + ) + if attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + query = attn.head_to_batch_dim(query) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + + attention_probs = attn.get_attention_scores(query, key, attention_mask) + hidden_states = torch.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + + # for ip-adapter + ip_key = self.to_k_ip(ip_hidden_states) + ip_value = self.to_v_ip(ip_hidden_states) + + ip_key = attn.head_to_batch_dim(ip_key) + ip_value = attn.head_to_batch_dim(ip_value) + + ip_attention_probs = attn.get_attention_scores(query, ip_key, None) + ip_hidden_states = torch.bmm(ip_attention_probs, ip_value) + ip_hidden_states = attn.batch_to_head_dim(ip_hidden_states) + + hidden_states = hidden_states + self.scale * ip_hidden_states + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class TA_IPAttnProcessor(nn.Module): + r""" + Attention processor for IP-Adapater. + Args: + hidden_size (`int`): + The hidden size of the attention layer. + cross_attention_dim (`int`): + The number of channels in the `encoder_hidden_states`. + scale (`float`, defaults to 1.0): + the weight scale of image prompt. + num_tokens (`int`, defaults to 4 when do ip_adapter_plus it should be 16): + The context length of the image features. + """ + + def __init__(self, hidden_size, cross_attention_dim=None, time_embedding_dim: int = None, scale=1.0, num_tokens=4): + super().__init__() + + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + self.scale = scale + self.num_tokens = num_tokens + + self.to_k_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + self.to_v_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + + self.ln_k_ip = AdaLayerNorm(hidden_size, time_embedding_dim) + self.ln_v_ip = AdaLayerNorm(hidden_size, time_embedding_dim) + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + temb=None, + ): + assert temb is not None, "Timestep embedding is needed for a time-aware attention processor." + + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + else: + # get encoder_hidden_states, ip_hidden_states + end_pos = encoder_hidden_states.shape[1] - self.num_tokens + encoder_hidden_states, ip_hidden_states = ( + encoder_hidden_states[:, :end_pos, :], + encoder_hidden_states[:, end_pos:, :], + ) + if attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + query = attn.head_to_batch_dim(query) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + + attention_probs = attn.get_attention_scores(query, key, attention_mask) + hidden_states = torch.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + + # for ip-adapter + ip_key = self.to_k_ip(ip_hidden_states) + ip_value = self.to_v_ip(ip_hidden_states) + + # time-dependent adaLN + ip_key = self.ln_k_ip(ip_key, temb) + ip_value = self.ln_v_ip(ip_value, temb) + + ip_key = attn.head_to_batch_dim(ip_key) + ip_value = attn.head_to_batch_dim(ip_value) + + ip_attention_probs = attn.get_attention_scores(query, ip_key, None) + ip_hidden_states = torch.bmm(ip_attention_probs, ip_value) + ip_hidden_states = attn.batch_to_head_dim(ip_hidden_states) + + hidden_states = hidden_states + self.scale * ip_hidden_states + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class AttnProcessor2_0(torch.nn.Module): + r""" + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). + """ + + def __init__( + self, + hidden_size=None, + cross_attention_dim=None, + ): + super().__init__() + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + external_kv=None, + temb=None, + ): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + if external_kv: + key = torch.cat([key, external_kv.k], axis=1) + value = torch.cat([value, external_kv.v], axis=1) + + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + # TODO: add support for attn.scale when we move to Torch 2.1 + hidden_states = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states.to(query.dtype) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class split_AttnProcessor2_0(torch.nn.Module): + r""" + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). + """ + + def __init__( + self, + hidden_size=None, + cross_attention_dim=None, + time_embedding_dim=None, + ): + super().__init__() + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + external_kv=None, + temb=None, + cat_dim=-2, + original_shape=None, + ): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + # 2d to sequence. + height, width = hidden_states.shape[-2:] + if cat_dim==-2 or cat_dim==2: + hidden_states_0 = hidden_states[:, :, :height//2, :] + hidden_states_1 = hidden_states[:, :, -(height//2):, :] + elif cat_dim==-1 or cat_dim==3: + hidden_states_0 = hidden_states[:, :, :, :width//2] + hidden_states_1 = hidden_states[:, :, :, -(width//2):] + batch_size, channel, height, width = hidden_states_0.shape + hidden_states_0 = hidden_states_0.view(batch_size, channel, height * width).transpose(1, 2) + hidden_states_1 = hidden_states_1.view(batch_size, channel, height * width).transpose(1, 2) + else: + # directly split sqeuence according to concat dim. + single_dim = original_shape[2] if cat_dim==-2 or cat_dim==2 else original_shape[1] + hidden_states_0 = hidden_states[:, :single_dim*single_dim,:] + hidden_states_1 = hidden_states[:, single_dim*(single_dim+1):,:] + + hidden_states = torch.cat([hidden_states_0, hidden_states_1], dim=1) + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + key = attn.to_k(hidden_states) + value = attn.to_v(hidden_states) + + if external_kv: + key = torch.cat([key, external_kv.k], dim=1) + value = torch.cat([value, external_kv.v], dim=1) + + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + # TODO: add support for attn.scale when we move to Torch 2.1 + hidden_states = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states.to(query.dtype) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + # spatially split. + hidden_states_0, hidden_states_1 = hidden_states.chunk(2, dim=1) + + if input_ndim == 4: + hidden_states_0 = hidden_states_0.transpose(-1, -2).reshape(batch_size, channel, height, width) + hidden_states_1 = hidden_states_1.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if cat_dim==-2 or cat_dim==2: + hidden_states_pad = torch.zeros(batch_size, channel, 1, width) + elif cat_dim==-1 or cat_dim==3: + hidden_states_pad = torch.zeros(batch_size, channel, height, 1) + hidden_states_pad = hidden_states_pad.to(hidden_states_0.device, dtype=hidden_states_0.dtype) + hidden_states = torch.cat([hidden_states_0, hidden_states_pad, hidden_states_1], dim=cat_dim) + assert hidden_states.shape == residual.shape, f"{hidden_states.shape} != {residual.shape}" + else: + batch_size, sequence_length, inner_dim = hidden_states.shape + hidden_states_pad = torch.zeros(batch_size, single_dim, inner_dim) + hidden_states_pad = hidden_states_pad.to(hidden_states_0.device, dtype=hidden_states_0.dtype) + hidden_states = torch.cat([hidden_states_0, hidden_states_pad, hidden_states_1], dim=1) + assert hidden_states.shape == residual.shape, f"{hidden_states.shape} != {residual.shape}" + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class sep_split_AttnProcessor2_0(torch.nn.Module): + r""" + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). + """ + + def __init__( + self, + hidden_size=None, + cross_attention_dim=None, + time_embedding_dim=None, + ): + super().__init__() + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") + self.ln_k_ref = AdaLayerNorm(hidden_size, time_embedding_dim) + self.ln_v_ref = AdaLayerNorm(hidden_size, time_embedding_dim) + # self.hidden_size = hidden_size + # self.cross_attention_dim = cross_attention_dim + # self.scale = scale + # self.num_tokens = num_tokens + + # self.to_q_ref = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + # self.to_k_ref = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + # self.to_v_ref = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + external_kv=None, + temb=None, + cat_dim=-2, + original_shape=None, + ref_scale=1.0, + ): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + # 2d to sequence. + height, width = hidden_states.shape[-2:] + if cat_dim==-2 or cat_dim==2: + hidden_states_0 = hidden_states[:, :, :height//2, :] + hidden_states_1 = hidden_states[:, :, -(height//2):, :] + elif cat_dim==-1 or cat_dim==3: + hidden_states_0 = hidden_states[:, :, :, :width//2] + hidden_states_1 = hidden_states[:, :, :, -(width//2):] + batch_size, channel, height, width = hidden_states_0.shape + hidden_states_0 = hidden_states_0.view(batch_size, channel, height * width).transpose(1, 2) + hidden_states_1 = hidden_states_1.view(batch_size, channel, height * width).transpose(1, 2) + else: + # directly split sqeuence according to concat dim. + single_dim = original_shape[2] if cat_dim==-2 or cat_dim==2 else original_shape[1] + hidden_states_0 = hidden_states[:, :single_dim*single_dim,:] + hidden_states_1 = hidden_states[:, single_dim*(single_dim+1):,:] + + batch_size, sequence_length, _ = ( + hidden_states_0.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + + if attn.group_norm is not None: + hidden_states_0 = attn.group_norm(hidden_states_0.transpose(1, 2)).transpose(1, 2) + hidden_states_1 = attn.group_norm(hidden_states_1.transpose(1, 2)).transpose(1, 2) + + query_0 = attn.to_q(hidden_states_0) + query_1 = attn.to_q(hidden_states_1) + key_0 = attn.to_k(hidden_states_0) + key_1 = attn.to_k(hidden_states_1) + value_0 = attn.to_v(hidden_states_0) + value_1 = attn.to_v(hidden_states_1) + + # time-dependent adaLN + key_1 = self.ln_k_ref(key_1, temb) + value_1 = self.ln_v_ref(value_1, temb) + + if external_kv: + key_1 = torch.cat([key_1, external_kv.k], dim=1) + value_1 = torch.cat([value_1, external_kv.v], dim=1) + + inner_dim = key_0.shape[-1] + head_dim = inner_dim // attn.heads + + query_0 = query_0.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + query_1 = query_1.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + key_0 = key_0.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + key_1 = key_1.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value_0 = value_0.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value_1 = value_1.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + # TODO: add support for attn.scale when we move to Torch 2.1 + hidden_states_0 = F.scaled_dot_product_attention( + query_0, key_0, value_0, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + hidden_states_1 = F.scaled_dot_product_attention( + query_1, key_1, value_1, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + # cross-attn + _hidden_states_0 = F.scaled_dot_product_attention( + query_0, key_1, value_1, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + hidden_states_0 = hidden_states_0 + ref_scale * _hidden_states_0 * 10 + + # TODO: drop this cross-attn + _hidden_states_1 = F.scaled_dot_product_attention( + query_1, key_0, value_0, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + hidden_states_1 = hidden_states_1 + ref_scale * _hidden_states_1 + + hidden_states_0 = hidden_states_0.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states_1 = hidden_states_1.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states_0 = hidden_states_0.to(query_0.dtype) + hidden_states_1 = hidden_states_1.to(query_1.dtype) + + + # linear proj + hidden_states_0 = attn.to_out[0](hidden_states_0) + hidden_states_1 = attn.to_out[0](hidden_states_1) + # dropout + hidden_states_0 = attn.to_out[1](hidden_states_0) + hidden_states_1 = attn.to_out[1](hidden_states_1) + + + if input_ndim == 4: + hidden_states_0 = hidden_states_0.transpose(-1, -2).reshape(batch_size, channel, height, width) + hidden_states_1 = hidden_states_1.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if cat_dim==-2 or cat_dim==2: + hidden_states_pad = torch.zeros(batch_size, channel, 1, width) + elif cat_dim==-1 or cat_dim==3: + hidden_states_pad = torch.zeros(batch_size, channel, height, 1) + hidden_states_pad = hidden_states_pad.to(hidden_states_0.device, dtype=hidden_states_0.dtype) + hidden_states = torch.cat([hidden_states_0, hidden_states_pad, hidden_states_1], dim=cat_dim) + assert hidden_states.shape == residual.shape, f"{hidden_states.shape} != {residual.shape}" + else: + batch_size, sequence_length, inner_dim = hidden_states.shape + hidden_states_pad = torch.zeros(batch_size, single_dim, inner_dim) + hidden_states_pad = hidden_states_pad.to(hidden_states_0.device, dtype=hidden_states_0.dtype) + hidden_states = torch.cat([hidden_states_0, hidden_states_pad, hidden_states_1], dim=1) + assert hidden_states.shape == residual.shape, f"{hidden_states.shape} != {residual.shape}" + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class AdditiveKV_AttnProcessor2_0(torch.nn.Module): + r""" + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). + """ + + def __init__( + self, + hidden_size: int = None, + cross_attention_dim: int = None, + time_embedding_dim: int = None, + additive_scale: float = 1.0, + ): + super().__init__() + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") + self.additive_scale = additive_scale + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + external_kv=None, + attention_mask=None, + temb=None, + ): + assert temb is not None, "Timestep embedding is needed for a time-aware attention processor." + + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + # TODO: add support for attn.scale when we move to Torch 2.1 + hidden_states = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + + if external_kv: + key = external_kv.k + value = external_kv.v + + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + external_attn_output = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + external_attn_output = external_attn_output.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states + self.additive_scale * external_attn_output + + hidden_states = hidden_states.to(query.dtype) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class TA_AdditiveKV_AttnProcessor2_0(torch.nn.Module): + r""" + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). + """ + + def __init__( + self, + hidden_size: int = None, + cross_attention_dim: int = None, + time_embedding_dim: int = None, + additive_scale: float = 1.0, + ): + super().__init__() + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") + self.ln_k = AdaLayerNorm(hidden_size, time_embedding_dim) + self.ln_v = AdaLayerNorm(hidden_size, time_embedding_dim) + self.additive_scale = additive_scale + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + external_kv=None, + attention_mask=None, + temb=None, + ): + assert temb is not None, "Timestep embedding is needed for a time-aware attention processor." + + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + # TODO: add support for attn.scale when we move to Torch 2.1 + hidden_states = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + + if external_kv: + key = external_kv.k + value = external_kv.v + + # time-dependent adaLN + key = self.ln_k(key, temb) + value = self.ln_v(value, temb) + + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + external_attn_output = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + external_attn_output = external_attn_output.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states + self.additive_scale * external_attn_output + + hidden_states = hidden_states.to(query.dtype) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class IPAttnProcessor2_0(torch.nn.Module): + r""" + Attention processor for IP-Adapater for PyTorch 2.0. + Args: + hidden_size (`int`): + The hidden size of the attention layer. + cross_attention_dim (`int`): + The number of channels in the `encoder_hidden_states`. + scale (`float`, defaults to 1.0): + the weight scale of image prompt. + num_tokens (`int`, defaults to 4 when do ip_adapter_plus it should be 16): + The context length of the image features. + """ + + def __init__(self, hidden_size, cross_attention_dim=None, scale=1.0, num_tokens=4): + super().__init__() + + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") + + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + self.scale = scale + self.num_tokens = num_tokens + + self.to_k_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + self.to_v_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + temb=None, + ): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + if isinstance(encoder_hidden_states, tuple): + # FIXME: now hard coded to single image prompt. + batch_size, _, hid_dim = encoder_hidden_states[0].shape + ip_tokens = encoder_hidden_states[1][0] + encoder_hidden_states = torch.cat([encoder_hidden_states[0], ip_tokens], dim=1) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + else: + # get encoder_hidden_states, ip_hidden_states + end_pos = encoder_hidden_states.shape[1] - self.num_tokens + encoder_hidden_states, ip_hidden_states = ( + encoder_hidden_states[:, :end_pos, :], + encoder_hidden_states[:, end_pos:, :], + ) + if attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + # TODO: add support for attn.scale when we move to Torch 2.1 + hidden_states = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states.to(query.dtype) + + # for ip-adapter + ip_key = self.to_k_ip(ip_hidden_states) + ip_value = self.to_v_ip(ip_hidden_states) + + ip_key = ip_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + ip_value = ip_value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + # TODO: add support for attn.scale when we move to Torch 2.1 + ip_hidden_states = F.scaled_dot_product_attention( + query, ip_key, ip_value, attn_mask=None, dropout_p=0.0, is_causal=False + ) + + ip_hidden_states = ip_hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + ip_hidden_states = ip_hidden_states.to(query.dtype) + + hidden_states = hidden_states + self.scale * ip_hidden_states + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class TA_IPAttnProcessor2_0(torch.nn.Module): + r""" + Attention processor for IP-Adapater for PyTorch 2.0. + Args: + hidden_size (`int`): + The hidden size of the attention layer. + cross_attention_dim (`int`): + The number of channels in the `encoder_hidden_states`. + scale (`float`, defaults to 1.0): + the weight scale of image prompt. + num_tokens (`int`, defaults to 4 when do ip_adapter_plus it should be 16): + The context length of the image features. + """ + + def __init__(self, hidden_size, cross_attention_dim=None, time_embedding_dim: int = None, scale=1.0, num_tokens=4): + super().__init__() + + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") + + self.hidden_size = hidden_size + self.cross_attention_dim = cross_attention_dim + self.scale = scale + self.num_tokens = num_tokens + + self.to_k_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + self.to_v_ip = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + self.ln_k_ip = AdaLayerNorm(hidden_size, time_embedding_dim) + self.ln_v_ip = AdaLayerNorm(hidden_size, time_embedding_dim) + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + external_kv=None, + temb=None, + ): + assert temb is not None, "Timestep embedding is needed for a time-aware attention processor." + + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + if not isinstance(encoder_hidden_states, tuple): + # get encoder_hidden_states, ip_hidden_states + end_pos = encoder_hidden_states.shape[1] - self.num_tokens + encoder_hidden_states, ip_hidden_states = ( + encoder_hidden_states[:, :end_pos, :], + encoder_hidden_states[:, end_pos:, :], + ) + else: + # FIXME: now hard coded to single image prompt. + batch_size, _, hid_dim = encoder_hidden_states[0].shape + ip_hidden_states = encoder_hidden_states[1][0] + encoder_hidden_states = encoder_hidden_states[0] + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + else: + if attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + if external_kv: + key = torch.cat([key, external_kv.k], axis=1) + value = torch.cat([value, external_kv.v], axis=1) + + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + # TODO: add support for attn.scale when we move to Torch 2.1 + hidden_states = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states.to(query.dtype) + + # for ip-adapter + ip_key = self.to_k_ip(ip_hidden_states) + ip_value = self.to_v_ip(ip_hidden_states) + + # time-dependent adaLN + ip_key = self.ln_k_ip(ip_key, temb) + ip_value = self.ln_v_ip(ip_value, temb) + + ip_key = ip_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + ip_value = ip_value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + # TODO: add support for attn.scale when we move to Torch 2.1 + ip_hidden_states = F.scaled_dot_product_attention( + query, ip_key, ip_value, attn_mask=None, dropout_p=0.0, is_causal=False + ) + + ip_hidden_states = ip_hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + ip_hidden_states = ip_hidden_states.to(query.dtype) + + hidden_states = hidden_states + self.scale * ip_hidden_states + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +## for controlnet +class CNAttnProcessor: + r""" + Default processor for performing attention-related computations. + """ + + def __init__(self, num_tokens=4): + self.num_tokens = num_tokens + + def __call__(self, attn, hidden_states, encoder_hidden_states=None, attention_mask=None, temb=None): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + else: + end_pos = encoder_hidden_states.shape[1] - self.num_tokens + encoder_hidden_states = encoder_hidden_states[:, :end_pos] # only use text + if attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + query = attn.head_to_batch_dim(query) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + + attention_probs = attn.get_attention_scores(query, key, attention_mask) + hidden_states = torch.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class CNAttnProcessor2_0: + r""" + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). + """ + + def __init__(self, num_tokens=4): + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") + self.num_tokens = num_tokens + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + temb=None, + ): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + else: + end_pos = encoder_hidden_states.shape[1] - self.num_tokens + encoder_hidden_states = encoder_hidden_states[:, :end_pos] # only use text + if attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + # TODO: add support for attn.scale when we move to Torch 2.1 + hidden_states = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states.to(query.dtype) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +def init_attn_proc(unet, ip_adapter_tokens=16, use_lcm=False, use_adaln=True, use_external_kv=False): + attn_procs = {} + unet_sd = unet.state_dict() + for name in unet.attn_processors.keys(): + cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim + if name.startswith("mid_block"): + hidden_size = unet.config.block_out_channels[-1] + elif name.startswith("up_blocks"): + block_id = int(name[len("up_blocks.")]) + hidden_size = list(reversed(unet.config.block_out_channels))[block_id] + elif name.startswith("down_blocks"): + block_id = int(name[len("down_blocks.")]) + hidden_size = unet.config.block_out_channels[block_id] + if cross_attention_dim is None: + if use_external_kv: + attn_procs[name] = AdditiveKV_AttnProcessor2_0( + hidden_size=hidden_size, + cross_attention_dim=cross_attention_dim, + time_embedding_dim=1280, + ) if hasattr(F, "scaled_dot_product_attention") else AdditiveKV_AttnProcessor() + else: + attn_procs[name] = AttnProcessor2_0() if hasattr(F, "scaled_dot_product_attention") else AttnProcessor() + else: + if use_adaln: + layer_name = name.split(".processor")[0] + if use_lcm: + weights = { + "to_k_ip.weight": unet_sd[layer_name + ".to_k.base_layer.weight"], + "to_v_ip.weight": unet_sd[layer_name + ".to_v.base_layer.weight"], + } + else: + weights = { + "to_k_ip.weight": unet_sd[layer_name + ".to_k.weight"], + "to_v_ip.weight": unet_sd[layer_name + ".to_v.weight"], + } + attn_procs[name] = TA_IPAttnProcessor2_0( + hidden_size=hidden_size, + cross_attention_dim=cross_attention_dim, + num_tokens=ip_adapter_tokens, + time_embedding_dim=1280, + ) if hasattr(F, "scaled_dot_product_attention") else \ + TA_IPAttnProcessor( + hidden_size=hidden_size, + cross_attention_dim=cross_attention_dim, + num_tokens=ip_adapter_tokens, + time_embedding_dim=1280, + ) + attn_procs[name].load_state_dict(weights, strict=False) + else: + attn_procs[name] = AttnProcessor2_0() if hasattr(F, "scaled_dot_product_attention") else AttnProcessor() + + return attn_procs + + +def init_aggregator_attn_proc(unet, use_adaln=False, split_attn=False): + attn_procs = {} + unet_sd = unet.state_dict() + for name in unet.attn_processors.keys(): + # get layer name and hidden dim + cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim + if name.startswith("mid_block"): + hidden_size = unet.config.block_out_channels[-1] + elif name.startswith("up_blocks"): + block_id = int(name[len("up_blocks.")]) + hidden_size = list(reversed(unet.config.block_out_channels))[block_id] + elif name.startswith("down_blocks"): + block_id = int(name[len("down_blocks.")]) + hidden_size = unet.config.block_out_channels[block_id] + # init attn proc + if split_attn: + # layer_name = name.split(".processor")[0] + # weights = { + # "to_q_ref.weight": unet_sd[layer_name + ".to_q.weight"], + # "to_k_ref.weight": unet_sd[layer_name + ".to_k.weight"], + # "to_v_ref.weight": unet_sd[layer_name + ".to_v.weight"], + # } + attn_procs[name] = ( + sep_split_AttnProcessor2_0( + hidden_size=hidden_size, + cross_attention_dim=hidden_size, + time_embedding_dim=1280, + ) + if use_adaln + else split_AttnProcessor2_0( + hidden_size=hidden_size, + cross_attention_dim=cross_attention_dim, + time_embedding_dim=1280, + ) + ) + # attn_procs[name].load_state_dict(weights, strict=False) + else: + attn_procs[name] = ( + AttnProcessor2_0( + hidden_size=hidden_size, + cross_attention_dim=hidden_size, + ) + if hasattr(F, "scaled_dot_product_attention") + else AttnProcessor( + hidden_size=hidden_size, + cross_attention_dim=hidden_size, + ) + ) + + return attn_procs diff --git a/modules/instantir/ip_adapter/ip_adapter.py b/modules/instantir/ip_adapter/ip_adapter.py new file mode 100644 index 000000000..10f01d4f3 --- /dev/null +++ b/modules/instantir/ip_adapter/ip_adapter.py @@ -0,0 +1,236 @@ +import os +import torch +from typing import List +from collections import namedtuple, OrderedDict + +def is_torch2_available(): + return hasattr(torch.nn.functional, "scaled_dot_product_attention") + +if is_torch2_available(): + from .attention_processor import ( + AttnProcessor2_0 as AttnProcessor, + ) + from .attention_processor import ( + CNAttnProcessor2_0 as CNAttnProcessor, + ) + from .attention_processor import ( + IPAttnProcessor2_0 as IPAttnProcessor, + ) + from .attention_processor import ( + TA_IPAttnProcessor2_0 as TA_IPAttnProcessor, + ) +else: + from .attention_processor import AttnProcessor, CNAttnProcessor, IPAttnProcessor, TA_IPAttnProcessor + + +class ImageProjModel(torch.nn.Module): + """Projection Model""" + + def __init__(self, cross_attention_dim=2048, clip_embeddings_dim=1280, clip_extra_context_tokens=4): + super().__init__() + + self.cross_attention_dim = cross_attention_dim + self.clip_extra_context_tokens = clip_extra_context_tokens + self.proj = torch.nn.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim) + self.norm = torch.nn.LayerNorm(cross_attention_dim) + + def forward(self, image_embeds): + embeds = image_embeds + clip_extra_context_tokens = self.proj(embeds).reshape( + -1, self.clip_extra_context_tokens, self.cross_attention_dim + ) + clip_extra_context_tokens = self.norm(clip_extra_context_tokens) + return clip_extra_context_tokens + + +class MLPProjModel(torch.nn.Module): + """SD model with image prompt""" + def __init__(self, cross_attention_dim=2048, clip_embeddings_dim=1280): + super().__init__() + + self.proj = torch.nn.Sequential( + torch.nn.Linear(clip_embeddings_dim, clip_embeddings_dim), + torch.nn.GELU(), + torch.nn.Linear(clip_embeddings_dim, cross_attention_dim), + torch.nn.LayerNorm(cross_attention_dim) + ) + + def forward(self, image_embeds): + clip_extra_context_tokens = self.proj(image_embeds) + return clip_extra_context_tokens + + +class MultiIPAdapterImageProjection(torch.nn.Module): + def __init__(self, IPAdapterImageProjectionLayers): + super().__init__() + self.image_projection_layers = torch.nn.ModuleList(IPAdapterImageProjectionLayers) + + def forward(self, image_embeds: List[torch.FloatTensor]): + projected_image_embeds = [] + + # currently, we accept `image_embeds` as + # 1. a tensor (deprecated) with shape [batch_size, embed_dim] or [batch_size, sequence_length, embed_dim] + # 2. list of `n` tensors where `n` is number of ip-adapters, each tensor can hae shape [batch_size, num_images, embed_dim] or [batch_size, num_images, sequence_length, embed_dim] + if not isinstance(image_embeds, list): + image_embeds = [image_embeds.unsqueeze(1)] + + if len(image_embeds) != len(self.image_projection_layers): + raise ValueError( + f"image_embeds must have the same length as image_projection_layers, got {len(image_embeds)} and {len(self.image_projection_layers)}" + ) + + for image_embed, image_projection_layer in zip(image_embeds, self.image_projection_layers): + batch_size, num_images = image_embed.shape[0], image_embed.shape[1] + image_embed = image_embed.reshape((batch_size * num_images,) + image_embed.shape[2:]) + image_embed = image_projection_layer(image_embed) + # image_embed = image_embed.reshape((batch_size, num_images) + image_embed.shape[1:]) + + projected_image_embeds.append(image_embed) + + return projected_image_embeds + + +class IPAdapter(torch.nn.Module): + """IP-Adapter""" + def __init__(self, unet, image_proj_model, adapter_modules, ckpt_path=None): + super().__init__() + self.unet = unet + self.image_proj = image_proj_model + self.ip_adapter = adapter_modules + + if ckpt_path is not None: + self.load_from_checkpoint(ckpt_path) + + def forward(self, noisy_latents, timesteps, encoder_hidden_states, image_embeds): + ip_tokens = self.image_proj(image_embeds) + encoder_hidden_states = torch.cat([encoder_hidden_states, ip_tokens], dim=1) + # Predict the noise residual + noise_pred = self.unet(noisy_latents, timesteps, encoder_hidden_states).sample + return noise_pred + + def load_from_checkpoint(self, ckpt_path: str): + # Calculate original checksums + orig_ip_proj_sum = torch.sum(torch.stack([torch.sum(p) for p in self.image_proj.parameters()])) + orig_adapter_sum = torch.sum(torch.stack([torch.sum(p) for p in self.ip_adapter.parameters()])) + + state_dict = torch.load(ckpt_path, map_location="cpu") + keys = list(state_dict.keys()) + if keys != ["image_proj", "ip_adapter"]: + state_dict = revise_state_dict(state_dict) + + # Load state dict for image_proj_model and adapter_modules + self.image_proj.load_state_dict(state_dict["image_proj"], strict=True) + self.ip_adapter.load_state_dict(state_dict["ip_adapter"], strict=True) + + # Calculate new checksums + new_ip_proj_sum = torch.sum(torch.stack([torch.sum(p) for p in self.image_proj.parameters()])) + new_adapter_sum = torch.sum(torch.stack([torch.sum(p) for p in self.ip_adapter.parameters()])) + + # Verify if the weights have changed + assert orig_ip_proj_sum != new_ip_proj_sum, "Weights of image_proj_model did not change!" + assert orig_adapter_sum != new_adapter_sum, "Weights of adapter_modules did not change!" + + +class IPAdapterPlus(torch.nn.Module): + """IP-Adapter""" + def __init__(self, unet, image_proj_model, adapter_modules, ckpt_path=None): + super().__init__() + self.unet = unet + self.image_proj = image_proj_model + self.ip_adapter = adapter_modules + + if ckpt_path is not None: + self.load_from_checkpoint(ckpt_path) + + def forward(self, noisy_latents, timesteps, encoder_hidden_states, image_embeds): + ip_tokens = self.image_proj(image_embeds) + encoder_hidden_states = torch.cat([encoder_hidden_states, ip_tokens], dim=1) + # Predict the noise residual + noise_pred = self.unet(noisy_latents, timesteps, encoder_hidden_states).sample + return noise_pred + + def load_from_checkpoint(self, ckpt_path: str): + # Calculate original checksums + orig_ip_proj_sum = torch.sum(torch.stack([torch.sum(p) for p in self.image_proj.parameters()])) + orig_adapter_sum = torch.sum(torch.stack([torch.sum(p) for p in self.ip_adapter.parameters()])) + org_unet_sum = [] + for attn_name, attn_proc in self.unet.attn_processors.items(): + if isinstance(attn_proc, (TA_IPAttnProcessor, IPAttnProcessor)): + org_unet_sum.append(torch.sum(torch.stack([torch.sum(p) for p in attn_proc.parameters()]))) + org_unet_sum = torch.sum(torch.stack(org_unet_sum)) + + state_dict = torch.load(ckpt_path, map_location="cpu") + keys = list(state_dict.keys()) + if keys != ["image_proj", "ip_adapter"]: + state_dict = revise_state_dict(state_dict) + + # Check if 'latents' exists in both the saved state_dict and the current model's state_dict + strict_load_image_proj_model = True + if "latents" in state_dict["image_proj"] and "latents" in self.image_proj.state_dict(): + # Check if the shapes are mismatched + if state_dict["image_proj"]["latents"].shape != self.image_proj.state_dict()["latents"].shape: + print(f"Shapes of 'image_proj.latents' in checkpoint {ckpt_path} and current model do not match.") + print("Removing 'latents' from checkpoint and loading the rest of the weights.") + del state_dict["image_proj"]["latents"] + strict_load_image_proj_model = False + + # Load state dict for image_proj_model and adapter_modules + self.image_proj.load_state_dict(state_dict["image_proj"], strict=strict_load_image_proj_model) + missing_key, unexpected_key = self.ip_adapter.load_state_dict(state_dict["ip_adapter"], strict=False) + if len(missing_key) > 0: + for ms in missing_key: + if "ln" not in ms: + raise ValueError(f"Missing key in adapter_modules: {len(missing_key)}") + if len(unexpected_key) > 0: + raise ValueError(f"Unexpected key in adapter_modules: {len(unexpected_key)}") + + # Calculate new checksums + new_ip_proj_sum = torch.sum(torch.stack([torch.sum(p) for p in self.image_proj.parameters()])) + new_adapter_sum = torch.sum(torch.stack([torch.sum(p) for p in self.ip_adapter.parameters()])) + + # Verify if the weights loaded to unet + unet_sum = [] + for attn_name, attn_proc in self.unet.attn_processors.items(): + if isinstance(attn_proc, (TA_IPAttnProcessor, IPAttnProcessor)): + unet_sum.append(torch.sum(torch.stack([torch.sum(p) for p in attn_proc.parameters()]))) + unet_sum = torch.sum(torch.stack(unet_sum)) + + assert org_unet_sum != unet_sum, "Weights of adapter_modules in unet did not change!" + assert (unet_sum - new_adapter_sum < 1e-4), "Weights of adapter_modules did not load to unet!" + + # Verify if the weights have changed + assert orig_ip_proj_sum != new_ip_proj_sum, "Weights of image_proj_model did not change!" + assert orig_adapter_sum != new_adapter_sum, "Weights of adapter_mod`ules did not change!" + + +class IPAdapterXL(IPAdapter): + """SDXL""" + + def forward(self, noisy_latents, timesteps, encoder_hidden_states, unet_added_cond_kwargs, image_embeds): + ip_tokens = self.image_proj(image_embeds) + encoder_hidden_states = torch.cat([encoder_hidden_states, ip_tokens], dim=1) + # Predict the noise residual + noise_pred = self.unet(noisy_latents, timesteps, encoder_hidden_states, added_cond_kwargs=unet_added_cond_kwargs).sample + return noise_pred + + +class IPAdapterPlusXL(IPAdapterPlus): + """IP-Adapter with fine-grained features""" + + def forward(self, noisy_latents, timesteps, encoder_hidden_states, unet_added_cond_kwargs, image_embeds): + ip_tokens = self.image_proj(image_embeds) + encoder_hidden_states = torch.cat([encoder_hidden_states, ip_tokens], dim=1) + # Predict the noise residual + noise_pred = self.unet(noisy_latents, timesteps, encoder_hidden_states, added_cond_kwargs=unet_added_cond_kwargs).sample + return noise_pred + + +class IPAdapterFull(IPAdapterPlus): + """IP-Adapter with full features""" + + def init_proj(self): + image_proj_model = MLPProjModel( + cross_attention_dim=self.pipe.unet.config.cross_attention_dim, + clip_embeddings_dim=self.image_encoder.config.hidden_size, + ).to(self.device, dtype=torch.float16) + return image_proj_model diff --git a/modules/instantir/ip_adapter/resampler.py b/modules/instantir/ip_adapter/resampler.py new file mode 100644 index 000000000..72295f90b --- /dev/null +++ b/modules/instantir/ip_adapter/resampler.py @@ -0,0 +1,158 @@ +# modified from https://github.com/mlfoundations/open_flamingo/blob/main/open_flamingo/src/helpers.py +# and https://github.com/lucidrains/imagen-pytorch/blob/main/imagen_pytorch/imagen_pytorch.py + +import math + +import torch +import torch.nn as nn +from einops import rearrange +from einops.layers.torch import Rearrange + + +# FFN +def FeedForward(dim, mult=4): + inner_dim = int(dim * mult) + return nn.Sequential( + nn.LayerNorm(dim), + nn.Linear(dim, inner_dim, bias=False), + nn.GELU(), + nn.Linear(inner_dim, dim, bias=False), + ) + + +def reshape_tensor(x, heads): + bs, length, width = x.shape + # (bs, length, width) --> (bs, length, n_heads, dim_per_head) + x = x.view(bs, length, heads, -1) + # (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head) + x = x.transpose(1, 2) + # (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head) + x = x.reshape(bs, heads, length, -1) + return x + + +class PerceiverAttention(nn.Module): + def __init__(self, *, dim, dim_head=64, heads=8): + super().__init__() + self.scale = dim_head**-0.5 + self.dim_head = dim_head + self.heads = heads + inner_dim = dim_head * heads + + self.norm1 = nn.LayerNorm(dim) + self.norm2 = nn.LayerNorm(dim) + + self.to_q = nn.Linear(dim, inner_dim, bias=False) + self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False) + self.to_out = nn.Linear(inner_dim, dim, bias=False) + + def forward(self, x, latents): + """ + Args: + x (torch.Tensor): image features + shape (b, n1, D) + latent (torch.Tensor): latent features + shape (b, n2, D) + """ + x = self.norm1(x) + latents = self.norm2(latents) + + b, l, _ = latents.shape + + q = self.to_q(latents) + kv_input = torch.cat((x, latents), dim=-2) + k, v = self.to_kv(kv_input).chunk(2, dim=-1) + + q = reshape_tensor(q, self.heads) + k = reshape_tensor(k, self.heads) + v = reshape_tensor(v, self.heads) + + # attention + scale = 1 / math.sqrt(math.sqrt(self.dim_head)) + weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards + weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) + out = weight @ v + + out = out.permute(0, 2, 1, 3).reshape(b, l, -1) + + return self.to_out(out) + + +class Resampler(nn.Module): + def __init__( + self, + dim=1280, + depth=4, + dim_head=64, + heads=20, + num_queries=64, + embedding_dim=768, + output_dim=1024, + ff_mult=4, + max_seq_len: int = 257, # CLIP tokens + CLS token + apply_pos_emb: bool = False, + num_latents_mean_pooled: int = 0, # number of latents derived from mean pooled representation of the sequence + ): + super().__init__() + self.pos_emb = nn.Embedding(max_seq_len, embedding_dim) if apply_pos_emb else None + + self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5) + + self.proj_in = nn.Linear(embedding_dim, dim) + + self.proj_out = nn.Linear(dim, output_dim) + self.norm_out = nn.LayerNorm(output_dim) + + self.to_latents_from_mean_pooled_seq = ( + nn.Sequential( + nn.LayerNorm(dim), + nn.Linear(dim, dim * num_latents_mean_pooled), + Rearrange("b (n d) -> b n d", n=num_latents_mean_pooled), + ) + if num_latents_mean_pooled > 0 + else None + ) + + self.layers = nn.ModuleList([]) + for _ in range(depth): + self.layers.append( + nn.ModuleList( + [ + PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads), + FeedForward(dim=dim, mult=ff_mult), + ] + ) + ) + + def forward(self, x): + if self.pos_emb is not None: + n, device = x.shape[1], x.device + pos_emb = self.pos_emb(torch.arange(n, device=device)) + x = x + pos_emb + + latents = self.latents.repeat(x.size(0), 1, 1) + + x = self.proj_in(x) + + if self.to_latents_from_mean_pooled_seq: + meanpooled_seq = masked_mean(x, dim=1, mask=torch.ones(x.shape[:2], device=x.device, dtype=torch.bool)) + meanpooled_latents = self.to_latents_from_mean_pooled_seq(meanpooled_seq) + latents = torch.cat((meanpooled_latents, latents), dim=-2) + + for attn, ff in self.layers: + latents = attn(x, latents) + latents + latents = ff(latents) + latents + + latents = self.proj_out(latents) + return self.norm_out(latents) + + +def masked_mean(t, *, dim, mask=None): + if mask is None: + return t.mean(dim=dim) + + denom = mask.sum(dim=dim, keepdim=True) + mask = rearrange(mask, "b n -> b n 1") + masked_t = t.masked_fill(~mask, 0.0) + + return masked_t.sum(dim=dim) / denom.clamp(min=1e-5) diff --git a/modules/instantir/ip_adapter/utils.py b/modules/instantir/ip_adapter/utils.py new file mode 100644 index 000000000..64c45cd85 --- /dev/null +++ b/modules/instantir/ip_adapter/utils.py @@ -0,0 +1,248 @@ +import torch +from collections import namedtuple, OrderedDict +from safetensors import safe_open +from .attention_processor import init_attn_proc +from .ip_adapter import MultiIPAdapterImageProjection +from .resampler import Resampler +from transformers import ( + AutoModel, AutoImageProcessor, + CLIPVisionModelWithProjection, CLIPImageProcessor) + + +def init_adapter_in_unet( + unet, + image_proj_model=None, + pretrained_model_path_or_dict=None, + adapter_tokens=64, + embedding_dim=None, + use_lcm=False, + use_adaln=True, + ): + device = unet.device + dtype = unet.dtype + if image_proj_model is None: + assert embedding_dim is not None, "embedding_dim must be provided if image_proj_model is None." + image_proj_model = Resampler( + embedding_dim=embedding_dim, + output_dim=unet.config.cross_attention_dim, + num_queries=adapter_tokens, + ) + if pretrained_model_path_or_dict is not None: + if not isinstance(pretrained_model_path_or_dict, dict): + if pretrained_model_path_or_dict.endswith(".safetensors"): + state_dict = {"image_proj": {}, "ip_adapter": {}} + with safe_open(pretrained_model_path_or_dict, framework="pt", device=unet.device) as f: + for key in f.keys(): + if key.startswith("image_proj."): + state_dict["image_proj"][key.replace("image_proj.", "")] = f.get_tensor(key) + elif key.startswith("ip_adapter."): + state_dict["ip_adapter"][key.replace("ip_adapter.", "")] = f.get_tensor(key) + else: + state_dict = torch.load(pretrained_model_path_or_dict, map_location=unet.device) + else: + state_dict = pretrained_model_path_or_dict + keys = list(state_dict.keys()) + if "image_proj" not in keys and "ip_adapter" not in keys: + state_dict = revise_state_dict(state_dict) + + # Creat IP cross-attention in unet. + attn_procs = init_attn_proc(unet, adapter_tokens, use_lcm, use_adaln) + unet.set_attn_processor(attn_procs) + + # Load pretrinaed model if needed. + if pretrained_model_path_or_dict is not None: + if "ip_adapter" in state_dict.keys(): + adapter_modules = torch.nn.ModuleList(unet.attn_processors.values()) + missing, unexpected = adapter_modules.load_state_dict(state_dict["ip_adapter"], strict=False) + for mk in missing: + if "ln" not in mk: + raise ValueError(f"Missing keys in adapter_modules: {missing}") + if "image_proj" in state_dict.keys(): + image_proj_model.load_state_dict(state_dict["image_proj"]) + + # Load image projectors into iterable ModuleList. + image_projection_layers = [] + image_projection_layers.append(image_proj_model) + unet.encoder_hid_proj = MultiIPAdapterImageProjection(image_projection_layers) + + # Adjust unet config to handle addtional ip hidden states. + unet.config.encoder_hid_dim_type = "ip_image_proj" + unet.to(dtype=dtype, device=device) + + +def load_adapter_to_pipe( + pipe, + pretrained_model_path_or_dict, + image_encoder_or_path=None, + feature_extractor_or_path=None, + use_clip_encoder=False, + adapter_tokens=64, + use_lcm=False, + use_adaln=True, + ): + + if not isinstance(pretrained_model_path_or_dict, dict): + if pretrained_model_path_or_dict.endswith(".safetensors"): + state_dict = {"image_proj": {}, "ip_adapter": {}} + with safe_open(pretrained_model_path_or_dict, framework="pt", device=pipe.device) as f: + for key in f.keys(): + if key.startswith("image_proj."): + state_dict["image_proj"][key.replace("image_proj.", "")] = f.get_tensor(key) + elif key.startswith("ip_adapter."): + state_dict["ip_adapter"][key.replace("ip_adapter.", "")] = f.get_tensor(key) + else: + state_dict = torch.load(pretrained_model_path_or_dict, map_location=pipe.device) + else: + state_dict = pretrained_model_path_or_dict + keys = list(state_dict.keys()) + if "image_proj" not in keys and "ip_adapter" not in keys: + state_dict = revise_state_dict(state_dict) + + # load CLIP image encoder here if it has not been registered to the pipeline yet + if image_encoder_or_path is not None: + if isinstance(image_encoder_or_path, str): + feature_extractor_or_path = image_encoder_or_path if feature_extractor_or_path is None else feature_extractor_or_path + + image_encoder_or_path = ( + CLIPVisionModelWithProjection.from_pretrained( + image_encoder_or_path + ) if use_clip_encoder else + AutoModel.from_pretrained(image_encoder_or_path) + ) + + if feature_extractor_or_path is not None: + if isinstance(feature_extractor_or_path, str): + feature_extractor_or_path = ( + CLIPImageProcessor() if use_clip_encoder else + AutoImageProcessor.from_pretrained(feature_extractor_or_path) + ) + + # create image encoder if it has not been registered to the pipeline yet + if hasattr(pipe, "image_encoder") and getattr(pipe, "image_encoder", None) is None: + image_encoder = image_encoder_or_path.to(pipe.device, dtype=pipe.dtype) + pipe.register_modules(image_encoder=image_encoder) + else: + image_encoder = pipe.image_encoder + + # create feature extractor if it has not been registered to the pipeline yet + if hasattr(pipe, "feature_extractor") and getattr(pipe, "feature_extractor", None) is None: + feature_extractor = feature_extractor_or_path + pipe.register_modules(feature_extractor=feature_extractor) + else: + feature_extractor = pipe.feature_extractor + + # load adapter into unet + unet = getattr(pipe, pipe.unet_name) if not hasattr(pipe, "unet") else pipe.unet + attn_procs = init_attn_proc(unet, adapter_tokens, use_lcm, use_adaln) + unet.set_attn_processor(attn_procs) + image_proj_model = Resampler( + embedding_dim=image_encoder.config.hidden_size, + output_dim=unet.config.cross_attention_dim, + num_queries=adapter_tokens, + ) + + # Load pretrinaed model if needed. + if "ip_adapter" in state_dict.keys(): + adapter_modules = torch.nn.ModuleList(unet.attn_processors.values()) + missing, unexpected = adapter_modules.load_state_dict(state_dict["ip_adapter"], strict=False) + for mk in missing: + if "ln" not in mk: + raise ValueError(f"Missing keys in adapter_modules: {missing}") + if "image_proj" in state_dict.keys(): + image_proj_model.load_state_dict(state_dict["image_proj"]) + + # convert IP-Adapter Image Projection layers to diffusers + image_projection_layers = [] + image_projection_layers.append(image_proj_model) + unet.encoder_hid_proj = MultiIPAdapterImageProjection(image_projection_layers) + + # Adjust unet config to handle addtional ip hidden states. + unet.config.encoder_hid_dim_type = "ip_image_proj" + unet.to(dtype=pipe.dtype, device=pipe.device) + + +def revise_state_dict(old_state_dict_or_path, map_location="cpu"): + new_state_dict = OrderedDict() + new_state_dict["image_proj"] = OrderedDict() + new_state_dict["ip_adapter"] = OrderedDict() + if isinstance(old_state_dict_or_path, str): + old_state_dict = torch.load(old_state_dict_or_path, map_location=map_location) + else: + old_state_dict = old_state_dict_or_path + for name, weight in old_state_dict.items(): + if name.startswith("image_proj_model."): + new_state_dict["image_proj"][name[len("image_proj_model."):]] = weight + elif name.startswith("adapter_modules."): + new_state_dict["ip_adapter"][name[len("adapter_modules."):]] = weight + return new_state_dict + + +# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.encode_image +def encode_image(image_encoder, feature_extractor, image, device, num_images_per_prompt, output_hidden_states=None): + dtype = next(image_encoder.parameters()).dtype + + if not isinstance(image, torch.Tensor): + image = feature_extractor(image, return_tensors="pt").pixel_values + + image = image.to(device=device, dtype=dtype) + if output_hidden_states: + image_enc_hidden_states = image_encoder(image, output_hidden_states=True).hidden_states[-2] + image_enc_hidden_states = image_enc_hidden_states.repeat_interleave(num_images_per_prompt, dim=0) + return image_enc_hidden_states + else: + if isinstance(image_encoder, CLIPVisionModelWithProjection): + # CLIP image encoder. + image_embeds = image_encoder(image).image_embeds + else: + # DINO image encoder. + image_embeds = image_encoder(image).last_hidden_state + image_embeds = image_embeds.repeat_interleave(num_images_per_prompt, dim=0) + return image_embeds + + +def prepare_training_image_embeds( + image_encoder, feature_extractor, + ip_adapter_image, ip_adapter_image_embeds, + device, drop_rate, output_hidden_state, idx_to_replace=None +): + if ip_adapter_image_embeds is None: + if not isinstance(ip_adapter_image, list): + ip_adapter_image = [ip_adapter_image] + + # if len(ip_adapter_image) != len(unet.encoder_hid_proj.image_projection_layers): + # raise ValueError( + # f"`ip_adapter_image` must have same length as the number of IP Adapters. Got {len(ip_adapter_image)} images and {len(unet.encoder_hid_proj.image_projection_layers)} IP Adapters." + # ) + + image_embeds = [] + for single_ip_adapter_image in ip_adapter_image: + if idx_to_replace is None: + idx_to_replace = torch.rand(len(single_ip_adapter_image)) < drop_rate + zero_ip_adapter_image = torch.zeros_like(single_ip_adapter_image) + single_ip_adapter_image[idx_to_replace] = zero_ip_adapter_image[idx_to_replace] + single_image_embeds = encode_image( + image_encoder, feature_extractor, single_ip_adapter_image, device, 1, output_hidden_state + ) + single_image_embeds = torch.stack([single_image_embeds], dim=1) # FIXME + + image_embeds.append(single_image_embeds) + else: + repeat_dims = [1] + image_embeds = [] + for single_image_embeds in ip_adapter_image_embeds: + if do_classifier_free_guidance: + single_negative_image_embeds, single_image_embeds = single_image_embeds.chunk(2) + single_image_embeds = single_image_embeds.repeat( + num_images_per_prompt, *(repeat_dims * len(single_image_embeds.shape[1:])) + ) + single_negative_image_embeds = single_negative_image_embeds.repeat( + num_images_per_prompt, *(repeat_dims * len(single_negative_image_embeds.shape[1:])) + ) + single_image_embeds = torch.cat([single_negative_image_embeds, single_image_embeds]) + else: + single_image_embeds = single_image_embeds.repeat( + num_images_per_prompt, *(repeat_dims * len(single_image_embeds.shape[1:])) + ) + image_embeds.append(single_image_embeds) + + return image_embeds \ No newline at end of file diff --git a/modules/instantir/lcm_single_step_scheduler.py b/modules/instantir/lcm_single_step_scheduler.py new file mode 100644 index 000000000..a32affdc2 --- /dev/null +++ b/modules/instantir/lcm_single_step_scheduler.py @@ -0,0 +1,537 @@ +# Copyright 2023 Stanford University Team and The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# DISCLAIMER: This code is strongly influenced by https://github.com/pesser/pytorch_diffusion +# and https://github.com/hojonathanho/diffusion + +import math +from dataclasses import dataclass +from typing import List, Optional, Tuple, Union + +import numpy as np +import torch + +from diffusers.configuration_utils import ConfigMixin, register_to_config +from diffusers.utils import BaseOutput, logging +from diffusers.utils.torch_utils import randn_tensor +from diffusers.schedulers.scheduling_utils import SchedulerMixin + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +@dataclass +class LCMSingleStepSchedulerOutput(BaseOutput): + """ + Output class for the scheduler's `step` function output. + + Args: + pred_original_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images): + The predicted denoised sample `(x_{0})` based on the model output from the current timestep. + `pred_original_sample` can be used to preview progress or for guidance. + """ + + denoised: Optional[torch.FloatTensor] = None + + +# Copied from diffusers.schedulers.scheduling_ddpm.betas_for_alpha_bar +def betas_for_alpha_bar( + num_diffusion_timesteps, + max_beta=0.999, + alpha_transform_type="cosine", +): + """ + Create a beta schedule that discretizes the given alpha_t_bar function, which defines the cumulative product of + (1-beta) over time from t = [0,1]. + + Contains a function alpha_bar that takes an argument t and transforms it to the cumulative product of (1-beta) up + to that part of the diffusion process. + + + Args: + num_diffusion_timesteps (`int`): the number of betas to produce. + max_beta (`float`): the maximum beta to use; use values lower than 1 to + prevent singularities. + alpha_transform_type (`str`, *optional*, default to `cosine`): the type of noise schedule for alpha_bar. + Choose from `cosine` or `exp` + + Returns: + betas (`np.ndarray`): the betas used by the scheduler to step the model outputs + """ + if alpha_transform_type == "cosine": + + def alpha_bar_fn(t): + return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2 + + elif alpha_transform_type == "exp": + + def alpha_bar_fn(t): + return math.exp(t * -12.0) + + else: + raise ValueError(f"Unsupported alpha_tranform_type: {alpha_transform_type}") + + betas = [] + for i in range(num_diffusion_timesteps): + t1 = i / num_diffusion_timesteps + t2 = (i + 1) / num_diffusion_timesteps + betas.append(min(1 - alpha_bar_fn(t2) / alpha_bar_fn(t1), max_beta)) + return torch.tensor(betas, dtype=torch.float32) + + +# Copied from diffusers.schedulers.scheduling_ddim.rescale_zero_terminal_snr +def rescale_zero_terminal_snr(betas: torch.FloatTensor) -> torch.FloatTensor: + """ + Rescales betas to have zero terminal SNR Based on https://arxiv.org/pdf/2305.08891.pdf (Algorithm 1) + + + Args: + betas (`torch.FloatTensor`): + the betas that the scheduler is being initialized with. + + Returns: + `torch.FloatTensor`: rescaled betas with zero terminal SNR + """ + # Convert betas to alphas_bar_sqrt + alphas = 1.0 - betas + alphas_cumprod = torch.cumprod(alphas, dim=0) + alphas_bar_sqrt = alphas_cumprod.sqrt() + + # Store old values. + alphas_bar_sqrt_0 = alphas_bar_sqrt[0].clone() + alphas_bar_sqrt_T = alphas_bar_sqrt[-1].clone() + + # Shift so the last timestep is zero. + alphas_bar_sqrt -= alphas_bar_sqrt_T + + # Scale so the first timestep is back to the old value. + alphas_bar_sqrt *= alphas_bar_sqrt_0 / (alphas_bar_sqrt_0 - alphas_bar_sqrt_T) + + # Convert alphas_bar_sqrt to betas + alphas_bar = alphas_bar_sqrt**2 # Revert sqrt + alphas = alphas_bar[1:] / alphas_bar[:-1] # Revert cumprod + alphas = torch.cat([alphas_bar[0:1], alphas]) + betas = 1 - alphas + + return betas + + +class LCMSingleStepScheduler(SchedulerMixin, ConfigMixin): + """ + `LCMSingleStepScheduler` extends the denoising procedure introduced in denoising diffusion probabilistic models (DDPMs) with + non-Markovian guidance. + + This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. [`~ConfigMixin`] takes care of storing all config + attributes that are passed in the scheduler's `__init__` function, such as `num_train_timesteps`. They can be + accessed via `scheduler.config.num_train_timesteps`. [`SchedulerMixin`] provides general loading and saving + functionality via the [`SchedulerMixin.save_pretrained`] and [`~SchedulerMixin.from_pretrained`] functions. + + Args: + num_train_timesteps (`int`, defaults to 1000): + The number of diffusion steps to train the model. + beta_start (`float`, defaults to 0.0001): + The starting `beta` value of inference. + beta_end (`float`, defaults to 0.02): + The final `beta` value. + beta_schedule (`str`, defaults to `"linear"`): + The beta schedule, a mapping from a beta range to a sequence of betas for stepping the model. Choose from + `linear`, `scaled_linear`, or `squaredcos_cap_v2`. + trained_betas (`np.ndarray`, *optional*): + Pass an array of betas directly to the constructor to bypass `beta_start` and `beta_end`. + original_inference_steps (`int`, *optional*, defaults to 50): + The default number of inference steps used to generate a linearly-spaced timestep schedule, from which we + will ultimately take `num_inference_steps` evenly spaced timesteps to form the final timestep schedule. + clip_sample (`bool`, defaults to `True`): + Clip the predicted sample for numerical stability. + clip_sample_range (`float`, defaults to 1.0): + The maximum magnitude for sample clipping. Valid only when `clip_sample=True`. + set_alpha_to_one (`bool`, defaults to `True`): + Each diffusion step uses the alphas product value at that step and at the previous one. For the final step + there is no previous alpha. When this option is `True` the previous alpha product is fixed to `1`, + otherwise it uses the alpha value at step 0. + steps_offset (`int`, defaults to 0): + An offset added to the inference steps. You can use a combination of `offset=1` and + `set_alpha_to_one=False` to make the last step use step 0 for the previous alpha product like in Stable + Diffusion. + prediction_type (`str`, defaults to `epsilon`, *optional*): + Prediction type of the scheduler function; can be `epsilon` (predicts the noise of the diffusion process), + `sample` (directly predicts the noisy sample`) or `v_prediction` (see section 2.4 of [Imagen + Video](https://imagen.research.google/video/paper.pdf) paper). + thresholding (`bool`, defaults to `False`): + Whether to use the "dynamic thresholding" method. This is unsuitable for latent-space diffusion models such + as Stable Diffusion. + dynamic_thresholding_ratio (`float`, defaults to 0.995): + The ratio for the dynamic thresholding method. Valid only when `thresholding=True`. + sample_max_value (`float`, defaults to 1.0): + The threshold value for dynamic thresholding. Valid only when `thresholding=True`. + timestep_spacing (`str`, defaults to `"leading"`): + The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and + Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information. + timestep_scaling (`float`, defaults to 10.0): + The factor the timesteps will be multiplied by when calculating the consistency model boundary conditions + `c_skip` and `c_out`. Increasing this will decrease the approximation error (although the approximation + error at the default of `10.0` is already pretty small). + rescale_betas_zero_snr (`bool`, defaults to `False`): + Whether to rescale the betas to have zero terminal SNR. This enables the model to generate very bright and + dark samples instead of limiting it to samples with medium brightness. Loosely related to + [`--offset_noise`](https://github.com/huggingface/diffusers/blob/74fd735eb073eb1d774b1ab4154a0876eb82f055/examples/dreambooth/train_dreambooth.py#L506). + """ + + order = 1 + + @register_to_config + def __init__( + self, + num_train_timesteps: int = 1000, + beta_start: float = 0.00085, + beta_end: float = 0.012, + beta_schedule: str = "scaled_linear", + trained_betas: Optional[Union[np.ndarray, List[float]]] = None, + original_inference_steps: int = 50, + clip_sample: bool = False, + clip_sample_range: float = 1.0, + set_alpha_to_one: bool = True, + steps_offset: int = 0, + prediction_type: str = "epsilon", + thresholding: bool = False, + dynamic_thresholding_ratio: float = 0.995, + sample_max_value: float = 1.0, + timestep_spacing: str = "leading", + timestep_scaling: float = 10.0, + rescale_betas_zero_snr: bool = False, + ): + if trained_betas is not None: + self.betas = torch.tensor(trained_betas, dtype=torch.float32) + elif beta_schedule == "linear": + self.betas = torch.linspace(beta_start, beta_end, num_train_timesteps, dtype=torch.float32) + elif beta_schedule == "scaled_linear": + # this schedule is very specific to the latent diffusion model. + self.betas = ( + torch.linspace(beta_start**0.5, beta_end**0.5, num_train_timesteps, dtype=torch.float32) ** 2 + ) + elif beta_schedule == "squaredcos_cap_v2": + # Glide cosine schedule + self.betas = betas_for_alpha_bar(num_train_timesteps) + else: + raise NotImplementedError(f"{beta_schedule} does is not implemented for {self.__class__}") + + # Rescale for zero SNR + if rescale_betas_zero_snr: + self.betas = rescale_zero_terminal_snr(self.betas) + + self.alphas = 1.0 - self.betas + self.alphas_cumprod = torch.cumprod(self.alphas, dim=0) + + # At every step in ddim, we are looking into the previous alphas_cumprod + # For the final step, there is no previous alphas_cumprod because we are already at 0 + # `set_alpha_to_one` decides whether we set this parameter simply to one or + # whether we use the final alpha of the "non-previous" one. + self.final_alpha_cumprod = torch.tensor(1.0) if set_alpha_to_one else self.alphas_cumprod[0] + + # standard deviation of the initial noise distribution + self.init_noise_sigma = 1.0 + + # setable values + self.num_inference_steps = None + self.timesteps = torch.from_numpy(np.arange(0, num_train_timesteps)[::-1].copy().astype(np.int64)) + + self._step_index = None + + # Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._init_step_index + def _init_step_index(self, timestep): + if isinstance(timestep, torch.Tensor): + timestep = timestep.to(self.timesteps.device) + + index_candidates = (self.timesteps == timestep).nonzero() + + # The sigma index that is taken for the **very** first `step` + # is always the second index (or the last index if there is only 1) + # This way we can ensure we don't accidentally skip a sigma in + # case we start in the middle of the denoising schedule (e.g. for image-to-image) + if len(index_candidates) > 1: + step_index = index_candidates[1] + else: + step_index = index_candidates[0] + + self._step_index = step_index.item() + + @property + def step_index(self): + return self._step_index + + def scale_model_input(self, sample: torch.FloatTensor, timestep: Optional[int] = None) -> torch.FloatTensor: + """ + Ensures interchangeability with schedulers that need to scale the denoising model input depending on the + current timestep. + + Args: + sample (`torch.FloatTensor`): + The input sample. + timestep (`int`, *optional*): + The current timestep in the diffusion chain. + Returns: + `torch.FloatTensor`: + A scaled input sample. + """ + return sample + + # Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler._threshold_sample + def _threshold_sample(self, sample: torch.FloatTensor) -> torch.FloatTensor: + """ + "Dynamic thresholding: At each sampling step we set s to a certain percentile absolute pixel value in xt0 (the + prediction of x_0 at timestep t), and if s > 1, then we threshold xt0 to the range [-s, s] and then divide by + s. Dynamic thresholding pushes saturated pixels (those near -1 and 1) inwards, thereby actively preventing + pixels from saturation at each step. We find that dynamic thresholding results in significantly better + photorealism as well as better image-text alignment, especially when using very large guidance weights." + + https://arxiv.org/abs/2205.11487 + """ + dtype = sample.dtype + batch_size, channels, *remaining_dims = sample.shape + + if dtype not in (torch.float32, torch.float64): + sample = sample.float() # upcast for quantile calculation, and clamp not implemented for cpu half + + # Flatten sample for doing quantile calculation along each image + sample = sample.reshape(batch_size, channels * np.prod(remaining_dims)) + + abs_sample = sample.abs() # "a certain percentile absolute pixel value" + + s = torch.quantile(abs_sample, self.config.dynamic_thresholding_ratio, dim=1) + s = torch.clamp( + s, min=1, max=self.config.sample_max_value + ) # When clamped to min=1, equivalent to standard clipping to [-1, 1] + s = s.unsqueeze(1) # (batch_size, 1) because clamp will broadcast along dim=0 + sample = torch.clamp(sample, -s, s) / s # "we threshold xt0 to the range [-s, s] and then divide by s" + + sample = sample.reshape(batch_size, channels, *remaining_dims) + sample = sample.to(dtype) + + return sample + + def set_timesteps( + self, + num_inference_steps: int = None, + device: Union[str, torch.device] = None, + original_inference_steps: Optional[int] = None, + strength: int = 1.0, + timesteps: Optional[list] = None, + ): + """ + Sets the discrete timesteps used for the diffusion chain (to be run before inference). + + Args: + num_inference_steps (`int`): + The number of diffusion steps used when generating samples with a pre-trained model. + device (`str` or `torch.device`, *optional*): + The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. + original_inference_steps (`int`, *optional*): + The original number of inference steps, which will be used to generate a linearly-spaced timestep + schedule (which is different from the standard `diffusers` implementation). We will then take + `num_inference_steps` timesteps from this schedule, evenly spaced in terms of indices, and use that as + our final timestep schedule. If not set, this will default to the `original_inference_steps` attribute. + """ + + if num_inference_steps is not None and timesteps is not None: + raise ValueError("Can only pass one of `num_inference_steps` or `custom_timesteps`.") + + if timesteps is not None: + for i in range(1, len(timesteps)): + if timesteps[i] >= timesteps[i - 1]: + raise ValueError("`custom_timesteps` must be in descending order.") + + if timesteps[0] >= self.config.num_train_timesteps: + raise ValueError( + f"`timesteps` must start before `self.config.train_timesteps`:" + f" {self.config.num_train_timesteps}." + ) + + timesteps = np.array(timesteps, dtype=np.int64) + else: + if num_inference_steps > self.config.num_train_timesteps: + raise ValueError( + f"`num_inference_steps`: {num_inference_steps} cannot be larger than `self.config.train_timesteps`:" + f" {self.config.num_train_timesteps} as the unet model trained with this scheduler can only handle" + f" maximal {self.config.num_train_timesteps} timesteps." + ) + + self.num_inference_steps = num_inference_steps + original_steps = ( + original_inference_steps if original_inference_steps is not None else self.config.original_inference_steps + ) + + if original_steps > self.config.num_train_timesteps: + raise ValueError( + f"`original_steps`: {original_steps} cannot be larger than `self.config.train_timesteps`:" + f" {self.config.num_train_timesteps} as the unet model trained with this scheduler can only handle" + f" maximal {self.config.num_train_timesteps} timesteps." + ) + + if num_inference_steps > original_steps: + raise ValueError( + f"`num_inference_steps`: {num_inference_steps} cannot be larger than `original_inference_steps`:" + f" {original_steps} because the final timestep schedule will be a subset of the" + f" `original_inference_steps`-sized initial timestep schedule." + ) + + # LCM Timesteps Setting + # Currently, only linear spacing is supported. + c = self.config.num_train_timesteps // original_steps + # LCM Training Steps Schedule + lcm_origin_timesteps = np.asarray(list(range(1, int(original_steps * strength) + 1))) * c - 1 + skipping_step = len(lcm_origin_timesteps) // num_inference_steps + # LCM Inference Steps Schedule + timesteps = lcm_origin_timesteps[::-skipping_step][:num_inference_steps] + + self.timesteps = torch.from_numpy(timesteps.copy()).to(device=device, dtype=torch.long) + + self._step_index = None + + def get_scalings_for_boundary_condition_discrete(self, timestep): + self.sigma_data = 0.5 # Default: 0.5 + scaled_timestep = timestep * self.config.timestep_scaling + + c_skip = self.sigma_data**2 / (scaled_timestep**2 + self.sigma_data**2) + c_out = scaled_timestep / (scaled_timestep**2 + self.sigma_data**2) ** 0.5 + return c_skip, c_out + + def append_dims(self, x, target_dims): + """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] + + def extract_into_tensor(self, a, t, x_shape): + b, *_ = t.shape + out = a.gather(-1, t) + return out.reshape(b, *((1,) * (len(x_shape) - 1))) + + def step( + self, + model_output: torch.FloatTensor, + timestep: torch.Tensor, + sample: torch.FloatTensor, + generator: Optional[torch.Generator] = None, + return_dict: bool = True, + ) -> Union[LCMSingleStepSchedulerOutput, Tuple]: + """ + Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion + process from the learned model outputs (most often the predicted noise). + + Args: + model_output (`torch.FloatTensor`): + The direct output from learned diffusion model. + timestep (`float`): + The current discrete timestep in the diffusion chain. + sample (`torch.FloatTensor`): + A current instance of a sample created by the diffusion process. + generator (`torch.Generator`, *optional*): + A random number generator. + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~schedulers.scheduling_lcm.LCMSchedulerOutput`] or `tuple`. + Returns: + [`~schedulers.scheduling_utils.LCMSchedulerOutput`] or `tuple`: + If return_dict is `True`, [`~schedulers.scheduling_lcm.LCMSchedulerOutput`] is returned, otherwise a + tuple is returned where the first element is the sample tensor. + """ + # 0. make sure everything is on the same device + alphas_cumprod = self.alphas_cumprod.to(sample.device) + + # 1. compute alphas, betas + if timestep.ndim == 0: + timestep = timestep.unsqueeze(0) + alpha_prod_t = self.extract_into_tensor(alphas_cumprod, timestep, sample.shape) + beta_prod_t = 1 - alpha_prod_t + + # 2. Get scalings for boundary conditions + c_skip, c_out = self.get_scalings_for_boundary_condition_discrete(timestep) + c_skip, c_out = [self.append_dims(x, sample.ndim) for x in [c_skip, c_out]] + + # 3. Compute the predicted original sample x_0 based on the model parameterization + if self.config.prediction_type == "epsilon": # noise-prediction + predicted_original_sample = (sample - torch.sqrt(beta_prod_t) * model_output) / torch.sqrt(alpha_prod_t) + elif self.config.prediction_type == "sample": # x-prediction + predicted_original_sample = model_output + elif self.config.prediction_type == "v_prediction": # v-prediction + predicted_original_sample = torch.sqrt(alpha_prod_t) * sample - torch.sqrt(beta_prod_t) * model_output + else: + raise ValueError( + f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample` or" + " `v_prediction` for `LCMScheduler`." + ) + + # 4. Clip or threshold "predicted x_0" + if self.config.thresholding: + predicted_original_sample = self._threshold_sample(predicted_original_sample) + elif self.config.clip_sample: + predicted_original_sample = predicted_original_sample.clamp( + -self.config.clip_sample_range, self.config.clip_sample_range + ) + + # 5. Denoise model output using boundary conditions + denoised = c_out * predicted_original_sample + c_skip * sample + + if not return_dict: + return (denoised, ) + + return LCMSingleStepSchedulerOutput(denoised=denoised) + + # Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler.add_noise + def add_noise( + self, + original_samples: torch.FloatTensor, + noise: torch.FloatTensor, + timesteps: torch.IntTensor, + ) -> torch.FloatTensor: + # Make sure alphas_cumprod and timestep have same device and dtype as original_samples + alphas_cumprod = self.alphas_cumprod.to(device=original_samples.device, dtype=original_samples.dtype) + timesteps = timesteps.to(original_samples.device) + + sqrt_alpha_prod = alphas_cumprod[timesteps] ** 0.5 + sqrt_alpha_prod = sqrt_alpha_prod.flatten() + while len(sqrt_alpha_prod.shape) < len(original_samples.shape): + sqrt_alpha_prod = sqrt_alpha_prod.unsqueeze(-1) + + sqrt_one_minus_alpha_prod = (1 - alphas_cumprod[timesteps]) ** 0.5 + sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.flatten() + while len(sqrt_one_minus_alpha_prod.shape) < len(original_samples.shape): + sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.unsqueeze(-1) + + noisy_samples = sqrt_alpha_prod * original_samples + sqrt_one_minus_alpha_prod * noise + return noisy_samples + + # Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler.get_velocity + def get_velocity( + self, sample: torch.FloatTensor, noise: torch.FloatTensor, timesteps: torch.IntTensor + ) -> torch.FloatTensor: + # Make sure alphas_cumprod and timestep have same device and dtype as sample + alphas_cumprod = self.alphas_cumprod.to(device=sample.device, dtype=sample.dtype) + timesteps = timesteps.to(sample.device) + + sqrt_alpha_prod = alphas_cumprod[timesteps] ** 0.5 + sqrt_alpha_prod = sqrt_alpha_prod.flatten() + while len(sqrt_alpha_prod.shape) < len(sample.shape): + sqrt_alpha_prod = sqrt_alpha_prod.unsqueeze(-1) + + sqrt_one_minus_alpha_prod = (1 - alphas_cumprod[timesteps]) ** 0.5 + sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.flatten() + while len(sqrt_one_minus_alpha_prod.shape) < len(sample.shape): + sqrt_one_minus_alpha_prod = sqrt_one_minus_alpha_prod.unsqueeze(-1) + + velocity = sqrt_alpha_prod * noise - sqrt_one_minus_alpha_prod * sample + return velocity + + def __len__(self): + return self.config.num_train_timesteps diff --git a/modules/instantir/sdxl_instantir.py b/modules/instantir/sdxl_instantir.py new file mode 100644 index 000000000..b279d8445 --- /dev/null +++ b/modules/instantir/sdxl_instantir.py @@ -0,0 +1,1738 @@ +# Copyright 2024 The InstantX Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + + +import inspect +from typing import Any, Callable, Dict, List, Optional, Tuple, Union + +import numpy as np +import PIL.Image +import torch +import torch.nn.functional as F +from transformers import ( + CLIPImageProcessor, + CLIPTextModel, + CLIPTextModelWithProjection, + CLIPTokenizer, + CLIPVisionModelWithProjection, +) + +from diffusers.utils.import_utils import is_invisible_watermark_available + +from diffusers.image_processor import PipelineImageInput, VaeImageProcessor +from diffusers.loaders import ( + FromSingleFileMixin, + IPAdapterMixin, + StableDiffusionXLLoraLoaderMixin, + TextualInversionLoaderMixin, +) +from diffusers.models import AutoencoderKL, ImageProjection, UNet2DConditionModel +from diffusers.models.attention_processor import ( + AttnProcessor2_0, + LoRAAttnProcessor2_0, + LoRAXFormersAttnProcessor, + XFormersAttnProcessor, +) +from diffusers.models.lora import adjust_lora_scale_text_encoder +from diffusers.schedulers import KarrasDiffusionSchedulers, LCMScheduler +from diffusers.utils import ( + USE_PEFT_BACKEND, + deprecate, + logging, + replace_example_docstring, + scale_lora_layers, + unscale_lora_layers, + convert_unet_state_dict_to_peft +) +from diffusers.utils.torch_utils import is_compiled_module, is_torch_version, randn_tensor +from diffusers.pipelines.pipeline_utils import DiffusionPipeline, StableDiffusionMixin +from diffusers.pipelines.stable_diffusion_xl.pipeline_output import StableDiffusionXLPipelineOutput + + +if is_invisible_watermark_available(): + from diffusers.pipelines.stable_diffusion_xl.watermark import StableDiffusionXLWatermarker + +from peft import LoraConfig, set_peft_model_state_dict +from .aggregator import Aggregator + + +logger = logging.get_logger(__name__) # pylint: disable=invalid-name + + +EXAMPLE_DOC_STRING = """ + Examples: + ```py + >>> # !pip install diffusers pillow transformers accelerate + >>> import torch + >>> from PIL import Image + + >>> from diffusers import DDPMScheduler + >>> from schedulers.lcm_single_step_scheduler import LCMSingleStepScheduler + + >>> from module.ip_adapter.utils import load_adapter_to_pipe + >>> from pipelines.sdxl_instantir import InstantIRPipeline + + >>> # download models under ./models + >>> dcp_adapter = f'./models/adapter.pt' + >>> previewer_lora_path = f'./models' + >>> instantir_path = f'./models/aggregator.pt' + + >>> # load pretrained models + >>> pipe = InstantIRPipeline.from_pretrained( + ... "stabilityai/stable-diffusion-xl-base-1.0", controlnet=controlnet, vae=vae, torch_dtype=torch.float16 + ... ) + >>> # load adapter + >>> load_adapter_to_pipe( + ... pipe, + ... dcp_adapter, + ... image_encoder_or_path = 'facebook/dinov2-large', + ... ) + >>> # load previewer lora + >>> pipe.prepare_previewers(previewer_lora_path) + >>> pipe.scheduler = DDPMScheduler.from_pretrained('stabilityai/stable-diffusion-xl-base-1.0', subfolder="scheduler") + >>> lcm_scheduler = LCMSingleStepScheduler.from_config(pipe.scheduler.config) + + >>> # load aggregator weights + >>> pretrained_state_dict = torch.load(instantir_path) + >>> pipe.aggregator.load_state_dict(pretrained_state_dict) + + >>> # send to GPU and fp16 + >>> pipe.to(device="cuda", dtype=torch.float16) + >>> pipe.aggregator.to(device="cuda", dtype=torch.float16) + >>> pipe.enable_model_cpu_offload() + + >>> # load a broken image + >>> low_quality_image = Image.open('path/to/your-image').convert("RGB") + + >>> # restoration + >>> image = pipe( + ... image=low_quality_image, + ... previewer_scheduler=lcm_scheduler, + ... ).images[0] + ``` +""" + +LCM_LORA_MODULES = [ + "to_q", + "to_k", + "to_v", + "to_out.0", + "proj_in", + "proj_out", + "ff.net.0.proj", + "ff.net.2", + "conv1", + "conv2", + "conv_shortcut", + "downsamplers.0.conv", + "upsamplers.0.conv", + "time_emb_proj", +] +PREVIEWER_LORA_MODULES = [ + "to_q", + "to_kv", + "0.to_out", + "attn1.to_k", + "attn1.to_v", + "to_k_ip", + "to_v_ip", + "ln_k_ip.linear", + "ln_v_ip.linear", + "to_out.0", + "proj_in", + "proj_out", + "ff.net.0.proj", + "ff.net.2", + "conv1", + "conv2", + "conv_shortcut", + "downsamplers.0.conv", + "upsamplers.0.conv", + "time_emb_proj", +] + + +def remove_attn2(model): + def recursive_find_module(name, module): + if not "up_blocks" in name and not "down_blocks" in name and not "mid_block" in name: return + elif "resnets" in name: return + if hasattr(module, "attn2"): + setattr(module, "attn2", None) + setattr(module, "norm2", None) + return + for sub_name, sub_module in module.named_children(): + recursive_find_module(f"{name}.{sub_name}", sub_module) + + for name, module in model.named_children(): + recursive_find_module(name, module) + + +# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.rescale_noise_cfg +def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0): + """ + Rescale `noise_cfg` according to `guidance_rescale`. Based on findings of [Common Diffusion Noise Schedules and + Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). See Section 3.4 + """ + std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True) + std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True) + # rescale the results from guidance (fixes overexposure) + noise_pred_rescaled = noise_cfg * (std_text / std_cfg) + # mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images + noise_cfg = guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg + return noise_cfg + + +# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps +def retrieve_timesteps( + scheduler, + num_inference_steps: Optional[int] = None, + device: Optional[Union[str, torch.device]] = None, + timesteps: Optional[List[int]] = None, + **kwargs, +): + """ + Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles + custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`. + + Args: + scheduler (`SchedulerMixin`): + The scheduler to get timesteps from. + num_inference_steps (`int`): + The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps` + must be `None`. + device (`str` or `torch.device`, *optional*): + The device to which the timesteps should be moved to. If `None`, the timesteps are not moved. + timesteps (`List[int]`, *optional*): + Custom timesteps used to support arbitrary spacing between timesteps. If `None`, then the default + timestep spacing strategy of the scheduler is used. If `timesteps` is passed, `num_inference_steps` + must be `None`. + + Returns: + `Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the + second element is the number of inference steps. + """ + if timesteps is not None: + accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys()) + if not accepts_timesteps: + raise ValueError( + f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom" + f" timestep schedules. Please check whether you are using the correct scheduler." + ) + scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs) + timesteps = scheduler.timesteps + num_inference_steps = len(timesteps) + else: + scheduler.set_timesteps(num_inference_steps, device=device, **kwargs) + timesteps = scheduler.timesteps + return timesteps, num_inference_steps + + +class InstantIRPipeline( + DiffusionPipeline, + StableDiffusionMixin, + TextualInversionLoaderMixin, + StableDiffusionXLLoraLoaderMixin, + IPAdapterMixin, + FromSingleFileMixin, +): + r""" + Pipeline for text-to-image generation using Stable Diffusion XL with ControlNet guidance. + + This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods + implemented for all pipelines (downloading, saving, running on a particular device, etc.). + + The pipeline also inherits the following loading methods: + - [`~loaders.TextualInversionLoaderMixin.load_textual_inversion`] for loading textual inversion embeddings + - [`~loaders.StableDiffusionXLLoraLoaderMixin.load_lora_weights`] for loading LoRA weights + - [`~loaders.StableDiffusionXLLoraLoaderMixin.save_lora_weights`] for saving LoRA weights + - [`~loaders.FromSingleFileMixin.from_single_file`] for loading `.ckpt` files + - [`~loaders.IPAdapterMixin.load_ip_adapter`] for loading IP Adapters + + Args: + vae ([`AutoencoderKL`]): + Variational Auto-Encoder (VAE) model to encode and decode images to and from latent representations. + text_encoder ([`~transformers.CLIPTextModel`]): + Frozen text-encoder ([clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14)). + text_encoder_2 ([`~transformers.CLIPTextModelWithProjection`]): + Second frozen text-encoder + ([laion/CLIP-ViT-bigG-14-laion2B-39B-b160k](https://huggingface.co/laion/CLIP-ViT-bigG-14-laion2B-39B-b160k)). + tokenizer ([`~transformers.CLIPTokenizer`]): + A `CLIPTokenizer` to tokenize text. + tokenizer_2 ([`~transformers.CLIPTokenizer`]): + A `CLIPTokenizer` to tokenize text. + unet ([`UNet2DConditionModel`]): + A `UNet2DConditionModel` to denoise the encoded image latents. + controlnet ([`ControlNetModel`] or `List[ControlNetModel]`): + Provides additional conditioning to the `unet` during the denoising process. If you set multiple + ControlNets as a list, the outputs from each ControlNet are added together to create one combined + additional conditioning. + scheduler ([`SchedulerMixin`]): + A scheduler to be used in combination with `unet` to denoise the encoded image latents. Can be one of + [`DDIMScheduler`], [`LMSDiscreteScheduler`], or [`PNDMScheduler`]. + force_zeros_for_empty_prompt (`bool`, *optional*, defaults to `"True"`): + Whether the negative prompt embeddings should always be set to 0. Also see the config of + `stabilityai/stable-diffusion-xl-base-1-0`. + add_watermarker (`bool`, *optional*): + Whether to use the [invisible_watermark](https://github.com/ShieldMnt/invisible-watermark/) library to + watermark output images. If not defined, it defaults to `True` if the package is installed; otherwise no + watermarker is used. + """ + + # leave controlnet out on purpose because it iterates with unet + model_cpu_offload_seq = "text_encoder->text_encoder_2->image_encoder->unet->vae" + _optional_components = [ + "tokenizer", + "tokenizer_2", + "text_encoder", + "text_encoder_2", + "feature_extractor", + "image_encoder", + ] + _callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"] + + def __init__( + self, + vae: AutoencoderKL, + text_encoder: CLIPTextModel, + text_encoder_2: CLIPTextModelWithProjection, + tokenizer: CLIPTokenizer, + tokenizer_2: CLIPTokenizer, + unet: UNet2DConditionModel, + scheduler: KarrasDiffusionSchedulers, + aggregator: Aggregator = None, + force_zeros_for_empty_prompt: bool = True, + add_watermarker: Optional[bool] = None, + feature_extractor: CLIPImageProcessor = None, + image_encoder: CLIPVisionModelWithProjection = None, + ): + super().__init__() + + if aggregator is None: + aggregator = Aggregator.from_unet(unet) + remove_attn2(aggregator) + + self.register_modules( + vae=vae, + text_encoder=text_encoder, + text_encoder_2=text_encoder_2, + tokenizer=tokenizer, + tokenizer_2=tokenizer_2, + unet=unet, + aggregator=aggregator, + scheduler=scheduler, + feature_extractor=feature_extractor, + image_encoder=image_encoder, + ) + self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1) + self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor, do_convert_rgb=True) + self.control_image_processor = VaeImageProcessor( + vae_scale_factor=self.vae_scale_factor, do_convert_rgb=True, do_normalize=True + ) + add_watermarker = add_watermarker if add_watermarker is not None else is_invisible_watermark_available() + + if add_watermarker: + self.watermark = StableDiffusionXLWatermarker() + else: + self.watermark = None + + self.register_to_config(force_zeros_for_empty_prompt=force_zeros_for_empty_prompt) + + def prepare_previewers(self, previewer_lora_path: str, use_lcm=False): + if use_lcm: + lora_state_dict, alpha_dict = self.lora_state_dict( + previewer_lora_path, + ) + else: + lora_state_dict, alpha_dict = self.lora_state_dict( + previewer_lora_path, + weight_name="previewer_lora_weights.bin" + ) + unet_state_dict = { + f'{k.replace("unet.", "")}': v for k, v in lora_state_dict.items() if k.startswith("unet.") + } + unet_state_dict = convert_unet_state_dict_to_peft(unet_state_dict) + lora_state_dict = dict() + for k, v in unet_state_dict.items(): + if "ip" in k: + k = k.replace("attn2", "attn2.processor") + lora_state_dict[k] = v + else: + lora_state_dict[k] = v + if alpha_dict: + lora_alpha = next(iter(alpha_dict.values())) + else: + lora_alpha = 1 + logger.info(f"use lora alpha {lora_alpha}") + lora_config = LoraConfig( + r=64, + target_modules=LCM_LORA_MODULES if use_lcm else PREVIEWER_LORA_MODULES, + lora_alpha=lora_alpha, + lora_dropout=0.0, + ) + + adapter_name = "lcm" if use_lcm else "previewer" + self.unet.add_adapter(lora_config, adapter_name) + incompatible_keys = set_peft_model_state_dict(self.unet, lora_state_dict, adapter_name=adapter_name) + if incompatible_keys is not None: + # check only for unexpected keys + unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None) + missing_keys = getattr(incompatible_keys, "missing_keys", None) + if unexpected_keys: + raise ValueError( + f"Loading adapter weights from state_dict led to unexpected keys not found in the model: " + f" {unexpected_keys}. " + ) + self.unet.disable_adapters() + + return lora_alpha + + # Copied from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl.StableDiffusionXLPipeline.encode_prompt + def encode_prompt( + self, + prompt: str, + prompt_2: Optional[str] = None, + device: Optional[torch.device] = None, + num_images_per_prompt: int = 1, + do_classifier_free_guidance: bool = True, + negative_prompt: Optional[str] = None, + negative_prompt_2: Optional[str] = None, + prompt_embeds: Optional[torch.FloatTensor] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None, + pooled_prompt_embeds: Optional[torch.FloatTensor] = None, + negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None, + lora_scale: Optional[float] = None, + clip_skip: Optional[int] = None, + ): + r""" + Encodes the prompt into text encoder hidden states. + + Args: + prompt (`str` or `List[str]`, *optional*): + prompt to be encoded + prompt_2 (`str` or `List[str]`, *optional*): + The prompt or prompts to be sent to the `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is + used in both text-encoders + device: (`torch.device`): + torch device + num_images_per_prompt (`int`): + number of images that should be generated per prompt + do_classifier_free_guidance (`bool`): + whether to use classifier free guidance or not + negative_prompt (`str` or `List[str]`, *optional*): + The prompt or prompts not to guide the image generation. If not defined, one has to pass + `negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is + less than `1`). + negative_prompt_2 (`str` or `List[str]`, *optional*): + The prompt or prompts not to guide the image generation to be sent to `tokenizer_2` and + `text_encoder_2`. If not defined, `negative_prompt` is used in both text-encoders + prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not + provided, text embeddings will be generated from `prompt` input argument. + negative_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input + argument. + pooled_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. + If not provided, pooled text embeddings will be generated from `prompt` input argument. + negative_pooled_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated negative pooled text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt + weighting. If not provided, pooled negative_prompt_embeds will be generated from `negative_prompt` + input argument. + lora_scale (`float`, *optional*): + A lora scale that will be applied to all LoRA layers of the text encoder if LoRA layers are loaded. + clip_skip (`int`, *optional*): + Number of layers to be skipped from CLIP while computing the prompt embeddings. A value of 1 means that + the output of the pre-final layer will be used for computing the prompt embeddings. + """ + device = device or self._execution_device + + # set lora scale so that monkey patched LoRA + # function of text encoder can correctly access it + if lora_scale is not None and isinstance(self, StableDiffusionXLLoraLoaderMixin): + self._lora_scale = lora_scale + + # dynamically adjust the LoRA scale + if self.text_encoder is not None: + if not USE_PEFT_BACKEND: + adjust_lora_scale_text_encoder(self.text_encoder, lora_scale) + else: + scale_lora_layers(self.text_encoder, lora_scale) + + if self.text_encoder_2 is not None: + if not USE_PEFT_BACKEND: + adjust_lora_scale_text_encoder(self.text_encoder_2, lora_scale) + else: + scale_lora_layers(self.text_encoder_2, lora_scale) + + prompt = [prompt] if isinstance(prompt, str) else prompt + + if prompt is not None: + batch_size = len(prompt) + else: + batch_size = prompt_embeds.shape[0] + + # Define tokenizers and text encoders + tokenizers = [self.tokenizer, self.tokenizer_2] if self.tokenizer is not None else [self.tokenizer_2] + text_encoders = ( + [self.text_encoder, self.text_encoder_2] if self.text_encoder is not None else [self.text_encoder_2] + ) + + if prompt_embeds is None: + prompt_2 = prompt_2 or prompt + prompt_2 = [prompt_2] if isinstance(prompt_2, str) else prompt_2 + + # textual inversion: process multi-vector tokens if necessary + prompt_embeds_list = [] + prompts = [prompt, prompt_2] + for prompt, tokenizer, text_encoder in zip(prompts, tokenizers, text_encoders): + if isinstance(self, TextualInversionLoaderMixin): + prompt = self.maybe_convert_prompt(prompt, tokenizer) + + text_inputs = tokenizer( + prompt, + padding="max_length", + max_length=tokenizer.model_max_length, + truncation=True, + return_tensors="pt", + ) + + text_input_ids = text_inputs.input_ids + untruncated_ids = tokenizer(prompt, padding="longest", return_tensors="pt").input_ids + + if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal( + text_input_ids, untruncated_ids + ): + removed_text = tokenizer.batch_decode(untruncated_ids[:, tokenizer.model_max_length - 1 : -1]) + logger.warning( + "The following part of your input was truncated because CLIP can only handle sequences up to" + f" {tokenizer.model_max_length} tokens: {removed_text}" + ) + + prompt_embeds = text_encoder(text_input_ids.to(device), output_hidden_states=True) + + # We are only ALWAYS interested in the pooled output of the final text encoder + pooled_prompt_embeds = prompt_embeds[0] + if clip_skip is None: + prompt_embeds = prompt_embeds.hidden_states[-2] + else: + # "2" because SDXL always indexes from the penultimate layer. + prompt_embeds = prompt_embeds.hidden_states[-(clip_skip + 2)] + + prompt_embeds_list.append(prompt_embeds) + + prompt_embeds = torch.concat(prompt_embeds_list, dim=-1) + + # get unconditional embeddings for classifier free guidance + zero_out_negative_prompt = negative_prompt is None and self.config.force_zeros_for_empty_prompt + if do_classifier_free_guidance and negative_prompt_embeds is None and zero_out_negative_prompt: + negative_prompt_embeds = torch.zeros_like(prompt_embeds) + negative_pooled_prompt_embeds = torch.zeros_like(pooled_prompt_embeds) + elif do_classifier_free_guidance and negative_prompt_embeds is None: + negative_prompt = negative_prompt or "" + negative_prompt_2 = negative_prompt_2 or negative_prompt + + # normalize str to list + negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt + negative_prompt_2 = ( + batch_size * [negative_prompt_2] if isinstance(negative_prompt_2, str) else negative_prompt_2 + ) + + uncond_tokens: List[str] + if prompt is not None and type(prompt) is not type(negative_prompt): + raise TypeError( + f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !=" + f" {type(prompt)}." + ) + elif batch_size != len(negative_prompt): + raise ValueError( + f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:" + f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches" + " the batch size of `prompt`." + ) + else: + uncond_tokens = [negative_prompt, negative_prompt_2] + + negative_prompt_embeds_list = [] + for negative_prompt, tokenizer, text_encoder in zip(uncond_tokens, tokenizers, text_encoders): + if isinstance(self, TextualInversionLoaderMixin): + negative_prompt = self.maybe_convert_prompt(negative_prompt, tokenizer) + + max_length = prompt_embeds.shape[1] + uncond_input = tokenizer( + negative_prompt, + padding="max_length", + max_length=max_length, + truncation=True, + return_tensors="pt", + ) + + negative_prompt_embeds = text_encoder( + uncond_input.input_ids.to(device), + output_hidden_states=True, + ) + # We are only ALWAYS interested in the pooled output of the final text encoder + negative_pooled_prompt_embeds = negative_prompt_embeds[0] + negative_prompt_embeds = negative_prompt_embeds.hidden_states[-2] + + negative_prompt_embeds_list.append(negative_prompt_embeds) + + negative_prompt_embeds = torch.concat(negative_prompt_embeds_list, dim=-1) + + if self.text_encoder_2 is not None: + prompt_embeds = prompt_embeds.to(dtype=self.text_encoder_2.dtype, device=device) + else: + prompt_embeds = prompt_embeds.to(dtype=self.unet.dtype, device=device) + + bs_embed, seq_len, _ = prompt_embeds.shape + # duplicate text embeddings for each generation per prompt, using mps friendly method + prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1) + prompt_embeds = prompt_embeds.view(bs_embed * num_images_per_prompt, seq_len, -1) + + if do_classifier_free_guidance: + # duplicate unconditional embeddings for each generation per prompt, using mps friendly method + seq_len = negative_prompt_embeds.shape[1] + + if self.text_encoder_2 is not None: + negative_prompt_embeds = negative_prompt_embeds.to(dtype=self.text_encoder_2.dtype, device=device) + else: + negative_prompt_embeds = negative_prompt_embeds.to(dtype=self.unet.dtype, device=device) + + negative_prompt_embeds = negative_prompt_embeds.repeat(1, num_images_per_prompt, 1) + negative_prompt_embeds = negative_prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1) + + pooled_prompt_embeds = pooled_prompt_embeds.repeat(1, num_images_per_prompt).view( + bs_embed * num_images_per_prompt, -1 + ) + if do_classifier_free_guidance: + negative_pooled_prompt_embeds = negative_pooled_prompt_embeds.repeat(1, num_images_per_prompt).view( + bs_embed * num_images_per_prompt, -1 + ) + + if self.text_encoder is not None: + if isinstance(self, StableDiffusionXLLoraLoaderMixin) and USE_PEFT_BACKEND: + # Retrieve the original scale by scaling back the LoRA layers + unscale_lora_layers(self.text_encoder, lora_scale) + + if self.text_encoder_2 is not None: + if isinstance(self, StableDiffusionXLLoraLoaderMixin) and USE_PEFT_BACKEND: + # Retrieve the original scale by scaling back the LoRA layers + unscale_lora_layers(self.text_encoder_2, lora_scale) + + return prompt_embeds, negative_prompt_embeds, pooled_prompt_embeds, negative_pooled_prompt_embeds + + # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.encode_image + def encode_image(self, image, device, num_images_per_prompt, output_hidden_states=None): + dtype = next(self.image_encoder.parameters()).dtype + + if not isinstance(image, torch.Tensor): + image = self.feature_extractor(image, return_tensors="pt").pixel_values + + image = image.to(device=device, dtype=dtype) + if output_hidden_states: + image_enc_hidden_states = self.image_encoder(image, output_hidden_states=True).hidden_states[-2] + image_enc_hidden_states = image_enc_hidden_states.repeat_interleave(num_images_per_prompt, dim=0) + uncond_image_enc_hidden_states = self.image_encoder( + torch.zeros_like(image), output_hidden_states=True + ).hidden_states[-2] + uncond_image_enc_hidden_states = uncond_image_enc_hidden_states.repeat_interleave( + num_images_per_prompt, dim=0 + ) + return image_enc_hidden_states, uncond_image_enc_hidden_states + else: + if isinstance(self.image_encoder, CLIPVisionModelWithProjection): + # CLIP image encoder. + image_embeds = self.image_encoder(image).image_embeds + image_embeds = image_embeds.repeat_interleave(num_images_per_prompt, dim=0) + uncond_image_embeds = torch.zeros_like(image_embeds) + else: + # DINO image encoder. + image_embeds = self.image_encoder(image).last_hidden_state + image_embeds = image_embeds.repeat_interleave(num_images_per_prompt, dim=0) + uncond_image_embeds = self.image_encoder( + torch.zeros_like(image) + ).last_hidden_state + uncond_image_embeds = uncond_image_embeds.repeat_interleave( + num_images_per_prompt, dim=0 + ) + + return image_embeds, uncond_image_embeds + + # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_ip_adapter_image_embeds + def prepare_ip_adapter_image_embeds( + self, ip_adapter_image, ip_adapter_image_embeds, device, num_images_per_prompt, do_classifier_free_guidance + ): + if ip_adapter_image_embeds is None: + if not isinstance(ip_adapter_image, list): + ip_adapter_image = [ip_adapter_image] + + if len(ip_adapter_image) != len(self.unet.encoder_hid_proj.image_projection_layers): + if isinstance(ip_adapter_image[0], list): + raise ValueError( + f"`ip_adapter_image` must have same length as the number of IP Adapters. Got {len(ip_adapter_image)} images and {len(self.unet.encoder_hid_proj.image_projection_layers)} IP Adapters." + ) + else: + logger.warning( + f"Got {len(ip_adapter_image)} images for {len(self.unet.encoder_hid_proj.image_projection_layers)} IP Adapters." + " By default, these images will be sent to each IP-Adapter. If this is not your use-case, please specify `ip_adapter_image` as a list of image-list, with" + f" length equals to the number of IP-Adapters." + ) + ip_adapter_image = [ip_adapter_image] * len(self.unet.encoder_hid_proj.image_projection_layers) + + image_embeds = [] + for single_ip_adapter_image, image_proj_layer in zip( + ip_adapter_image, self.unet.encoder_hid_proj.image_projection_layers + ): + output_hidden_state = isinstance(self.image_encoder, CLIPVisionModelWithProjection) and not isinstance(image_proj_layer, ImageProjection) + single_image_embeds, single_negative_image_embeds = self.encode_image( + single_ip_adapter_image, device, 1, output_hidden_state + ) + single_image_embeds = torch.stack([single_image_embeds] * (num_images_per_prompt//single_image_embeds.shape[0]), dim=0) + single_negative_image_embeds = torch.stack( + [single_negative_image_embeds] * (num_images_per_prompt//single_negative_image_embeds.shape[0]), dim=0 + ) + + if do_classifier_free_guidance: + single_image_embeds = torch.cat([single_negative_image_embeds, single_image_embeds]) + single_image_embeds = single_image_embeds.to(device) + + image_embeds.append(single_image_embeds) + else: + repeat_dims = [1] + image_embeds = [] + for single_image_embeds in ip_adapter_image_embeds: + if do_classifier_free_guidance: + single_negative_image_embeds, single_image_embeds = single_image_embeds.chunk(2) + single_image_embeds = single_image_embeds.repeat( + num_images_per_prompt, *(repeat_dims * len(single_image_embeds.shape[1:])) + ) + single_negative_image_embeds = single_negative_image_embeds.repeat( + num_images_per_prompt, *(repeat_dims * len(single_negative_image_embeds.shape[1:])) + ) + single_image_embeds = torch.cat([single_negative_image_embeds, single_image_embeds]) + else: + single_image_embeds = single_image_embeds.repeat( + num_images_per_prompt, *(repeat_dims * len(single_image_embeds.shape[1:])) + ) + image_embeds.append(single_image_embeds) + + return image_embeds + + # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_extra_step_kwargs + def prepare_extra_step_kwargs(self, generator, eta): + # prepare extra kwargs for the scheduler step, since not all schedulers have the same signature + # eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers. + # eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502 + # and should be between [0, 1] + + accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys()) + extra_step_kwargs = {} + if accepts_eta: + extra_step_kwargs["eta"] = eta + + # check if the scheduler accepts generator + accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys()) + if accepts_generator: + extra_step_kwargs["generator"] = generator + return extra_step_kwargs + + def check_inputs( + self, + prompt, + prompt_2, + image, + callback_steps, + negative_prompt=None, + negative_prompt_2=None, + prompt_embeds=None, + negative_prompt_embeds=None, + pooled_prompt_embeds=None, + ip_adapter_image=None, + ip_adapter_image_embeds=None, + negative_pooled_prompt_embeds=None, + controlnet_conditioning_scale=1.0, + control_guidance_start=0.0, + control_guidance_end=1.0, + callback_on_step_end_tensor_inputs=None, + ): + if callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0): + raise ValueError( + f"`callback_steps` has to be a positive integer but is {callback_steps} of type" + f" {type(callback_steps)}." + ) + + if callback_on_step_end_tensor_inputs is not None and not all( + k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs + ): + raise ValueError( + f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}" + ) + + if prompt is not None and prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to" + " only forward one of the two." + ) + elif prompt_2 is not None and prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `prompt_2`: {prompt_2} and `prompt_embeds`: {prompt_embeds}. Please make sure to" + " only forward one of the two." + ) + elif prompt is None and prompt_embeds is None: + raise ValueError( + "Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined." + ) + elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)): + raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}") + elif prompt_2 is not None and (not isinstance(prompt_2, str) and not isinstance(prompt_2, list)): + raise ValueError(f"`prompt_2` has to be of type `str` or `list` but is {type(prompt_2)}") + + if negative_prompt is not None and negative_prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:" + f" {negative_prompt_embeds}. Please make sure to only forward one of the two." + ) + elif negative_prompt_2 is not None and negative_prompt_embeds is not None: + raise ValueError( + f"Cannot forward both `negative_prompt_2`: {negative_prompt_2} and `negative_prompt_embeds`:" + f" {negative_prompt_embeds}. Please make sure to only forward one of the two." + ) + + if prompt_embeds is not None and negative_prompt_embeds is not None: + if prompt_embeds.shape != negative_prompt_embeds.shape: + raise ValueError( + "`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but" + f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`" + f" {negative_prompt_embeds.shape}." + ) + + if prompt_embeds is not None and pooled_prompt_embeds is None: + raise ValueError( + "If `prompt_embeds` are provided, `pooled_prompt_embeds` also have to be passed. Make sure to generate `pooled_prompt_embeds` from the same text encoder that was used to generate `prompt_embeds`." + ) + + if negative_prompt_embeds is not None and negative_pooled_prompt_embeds is None: + raise ValueError( + "If `negative_prompt_embeds` are provided, `negative_pooled_prompt_embeds` also have to be passed. Make sure to generate `negative_pooled_prompt_embeds` from the same text encoder that was used to generate `negative_prompt_embeds`." + ) + + # Check `image` + is_compiled = hasattr(F, "scaled_dot_product_attention") and isinstance( + self.aggregator, torch._dynamo.eval_frame.OptimizedModule + ) + if ( + isinstance(self.aggregator, Aggregator) + or is_compiled + and isinstance(self.aggregator._orig_mod, Aggregator) + ): + self.check_image(image, prompt, prompt_embeds) + else: + assert False + + if control_guidance_start >= control_guidance_end: + raise ValueError( + f"control guidance start: {control_guidance_start} cannot be larger or equal to control guidance end: {control_guidance_end}." + ) + if control_guidance_start < 0.0: + raise ValueError(f"control guidance start: {control_guidance_start} can't be smaller than 0.") + if control_guidance_end > 1.0: + raise ValueError(f"control guidance end: {control_guidance_end} can't be larger than 1.0.") + + if ip_adapter_image is not None and ip_adapter_image_embeds is not None: + raise ValueError( + "Provide either `ip_adapter_image` or `ip_adapter_image_embeds`. Cannot leave both `ip_adapter_image` and `ip_adapter_image_embeds` defined." + ) + + if ip_adapter_image_embeds is not None: + if not isinstance(ip_adapter_image_embeds, list): + raise ValueError( + f"`ip_adapter_image_embeds` has to be of type `list` but is {type(ip_adapter_image_embeds)}" + ) + elif ip_adapter_image_embeds[0].ndim not in [3, 4]: + raise ValueError( + f"`ip_adapter_image_embeds` has to be a list of 3D or 4D tensors but is {ip_adapter_image_embeds[0].ndim}D" + ) + + # Copied from diffusers.pipelines.controlnet.pipeline_controlnet.StableDiffusionControlNetPipeline.check_image + def check_image(self, image, prompt, prompt_embeds): + image_is_pil = isinstance(image, PIL.Image.Image) + image_is_tensor = isinstance(image, torch.Tensor) + image_is_np = isinstance(image, np.ndarray) + image_is_pil_list = isinstance(image, list) and isinstance(image[0], PIL.Image.Image) + image_is_tensor_list = isinstance(image, list) and isinstance(image[0], torch.Tensor) + image_is_np_list = isinstance(image, list) and isinstance(image[0], np.ndarray) + + if ( + not image_is_pil + and not image_is_tensor + and not image_is_np + and not image_is_pil_list + and not image_is_tensor_list + and not image_is_np_list + ): + raise TypeError( + f"image must be passed and be one of PIL image, numpy array, torch tensor, list of PIL images, list of numpy arrays or list of torch tensors, but is {type(image)}" + ) + + if image_is_pil: + image_batch_size = 1 + else: + image_batch_size = len(image) + + if prompt is not None and isinstance(prompt, str): + prompt_batch_size = 1 + elif prompt is not None and isinstance(prompt, list): + prompt_batch_size = len(prompt) + elif prompt_embeds is not None: + prompt_batch_size = prompt_embeds.shape[0] + + if image_batch_size != 1 and image_batch_size != prompt_batch_size: + raise ValueError( + f"If image batch size is not 1, image batch size must be same as prompt batch size. image batch size: {image_batch_size}, prompt batch size: {prompt_batch_size}" + ) + + # Copied from diffusers.pipelines.controlnet.pipeline_controlnet.StableDiffusionControlNetPipeline.prepare_image + def prepare_image( + self, + image, + width, + height, + batch_size, + num_images_per_prompt, + device, + dtype, + do_classifier_free_guidance=False, + ): + image = self.control_image_processor.preprocess(image, height=height, width=width).to(dtype=torch.float32) + image_batch_size = image.shape[0] + + if image_batch_size == 1: + repeat_by = batch_size + else: + # image batch size is the same as prompt batch size + repeat_by = num_images_per_prompt + + image = image.repeat_interleave(repeat_by, dim=0) + + image = image.to(device=device, dtype=dtype) + + return image + + @torch.no_grad() + def init_latents(self, latents, generator, timestep): + noise = torch.randn(latents.shape, generator=generator[0] if isinstance(generator, list) else generator, device=self.vae.device, dtype=self.vae.dtype, layout=torch.strided) + bsz = latents.shape[0] + timestep = torch.tensor([timestep]*bsz, device=self.vae.device) + # Note that the latents will be scaled aleady by scheduler.add_noise + latents = self.scheduler.add_noise(latents, noise, timestep) + return latents + + # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.StableDiffusionPipeline.prepare_latents + def prepare_latents(self, batch_size, num_channels_latents, height, width, dtype, device, generator, latents=None): + shape = ( + batch_size, + num_channels_latents, + int(height) // self.vae_scale_factor, + int(width) // self.vae_scale_factor, + ) + if isinstance(generator, list) and len(generator) != batch_size: + raise ValueError( + f"You have passed a list of generators of length {len(generator)}, but requested an effective batch" + f" size of {batch_size}. Make sure the batch size matches the length of the generators." + ) + + if latents is None: + latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype) + else: + latents = latents.to(device) + + # scale the initial noise by the standard deviation required by the scheduler + latents = latents * self.scheduler.init_noise_sigma + return latents + + # Copied from diffusers.pipelines.stable_diffusion_xl.pipeline_stable_diffusion_xl.StableDiffusionXLPipeline._get_add_time_ids + def _get_add_time_ids( + self, original_size, crops_coords_top_left, target_size, dtype, text_encoder_projection_dim=None + ): + add_time_ids = list(original_size + crops_coords_top_left + target_size) + + passed_add_embed_dim = ( + self.unet.config.addition_time_embed_dim * len(add_time_ids) + text_encoder_projection_dim + ) + expected_add_embed_dim = self.unet.add_embedding.linear_1.in_features + + if expected_add_embed_dim != passed_add_embed_dim: + raise ValueError( + f"Model expects an added time embedding vector of length {expected_add_embed_dim}, but a vector of {passed_add_embed_dim} was created. The model has an incorrect config. Please check `unet.config.time_embedding_type` and `text_encoder_2.config.projection_dim`." + ) + + add_time_ids = torch.tensor([add_time_ids], dtype=dtype) + return add_time_ids + + # Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion_upscale.StableDiffusionUpscalePipeline.upcast_vae + def upcast_vae(self): + dtype = self.vae.dtype + self.vae.to(dtype=torch.float32) + use_torch_2_0_or_xformers = isinstance( + self.vae.decoder.mid_block.attentions[0].processor, + ( + AttnProcessor2_0, + XFormersAttnProcessor, + LoRAXFormersAttnProcessor, + LoRAAttnProcessor2_0, + ), + ) + # if xformers or torch_2_0 is used attention block does not need + # to be in float32 which can save lots of memory + if use_torch_2_0_or_xformers: + self.vae.post_quant_conv.to(dtype) + self.vae.decoder.conv_in.to(dtype) + self.vae.decoder.mid_block.to(dtype) + + # Copied from diffusers.pipelines.latent_consistency_models.pipeline_latent_consistency_text2img.LatentConsistencyModelPipeline.get_guidance_scale_embedding + def get_guidance_scale_embedding( + self, w: torch.Tensor, embedding_dim: int = 512, dtype: torch.dtype = torch.float32 + ) -> torch.FloatTensor: + """ + See https://github.com/google-research/vdm/blob/dc27b98a554f65cdc654b800da5aa1846545d41b/model_vdm.py#L298 + + Args: + w (`torch.Tensor`): + Generate embedding vectors with a specified guidance scale to subsequently enrich timestep embeddings. + embedding_dim (`int`, *optional*, defaults to 512): + Dimension of the embeddings to generate. + dtype (`torch.dtype`, *optional*, defaults to `torch.float32`): + Data type of the generated embeddings. + + Returns: + `torch.FloatTensor`: Embedding vectors with shape `(len(w), embedding_dim)`. + """ + assert len(w.shape) == 1 + w = w * 1000.0 + + half_dim = embedding_dim // 2 + emb = torch.log(torch.tensor(10000.0)) / (half_dim - 1) + emb = torch.exp(torch.arange(half_dim, dtype=dtype) * -emb) + emb = w.to(dtype)[:, None] * emb[None, :] + emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=1) + if embedding_dim % 2 == 1: # zero pad + emb = torch.nn.functional.pad(emb, (0, 1)) + assert emb.shape == (w.shape[0], embedding_dim) + return emb + + @property + def guidance_scale(self): + return self._guidance_scale + + @property + def guidance_rescale(self): + return self._guidance_rescale + + @property + def clip_skip(self): + return self._clip_skip + + # here `guidance_scale` is defined analog to the guidance weight `w` of equation (2) + # of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1` + # corresponds to doing no classifier free guidance. + @property + def do_classifier_free_guidance(self): + return self._guidance_scale > 1 and self.unet.config.time_cond_proj_dim is None + + @property + def cross_attention_kwargs(self): + return self._cross_attention_kwargs + + @property + def denoising_end(self): + return self._denoising_end + + @property + def num_timesteps(self): + return self._num_timesteps + + @torch.no_grad() + @replace_example_docstring(EXAMPLE_DOC_STRING) + def __call__( + self, + prompt: Union[str, List[str]] = None, + prompt_2: Optional[Union[str, List[str]]] = None, + image: PipelineImageInput = None, + height: Optional[int] = None, + width: Optional[int] = None, + num_inference_steps: int = 30, + timesteps: List[int] = None, + denoising_end: Optional[float] = None, + guidance_scale: float = 7.0, + negative_prompt: Optional[Union[str, List[str]]] = None, + negative_prompt_2: Optional[Union[str, List[str]]] = None, + num_images_per_prompt: Optional[int] = 1, + eta: float = 0.0, + generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, + latents: Optional[torch.FloatTensor] = None, + prompt_embeds: Optional[torch.FloatTensor] = None, + negative_prompt_embeds: Optional[torch.FloatTensor] = None, + pooled_prompt_embeds: Optional[torch.FloatTensor] = None, + negative_pooled_prompt_embeds: Optional[torch.FloatTensor] = None, + ip_adapter_image: Optional[PipelineImageInput] = None, + ip_adapter_image_embeds: Optional[List[torch.FloatTensor]] = None, + output_type: Optional[str] = "pil", + return_dict: bool = True, + save_preview_row: bool = False, + init_latents_with_lq: bool = True, + multistep_restore: bool = False, + adastep_restore: bool = False, + cross_attention_kwargs: Optional[Dict[str, Any]] = None, + guidance_rescale: float = 0.0, + controlnet_conditioning_scale: float = 1.0, + control_guidance_start: float = 0.0, + control_guidance_end: float = 1.0, + preview_start: float = 0.0, + preview_end: float = 1.0, + original_size: Tuple[int, int] = None, + crops_coords_top_left: Tuple[int, int] = (0, 0), + target_size: Tuple[int, int] = None, + negative_original_size: Optional[Tuple[int, int]] = None, + negative_crops_coords_top_left: Tuple[int, int] = (0, 0), + negative_target_size: Optional[Tuple[int, int]] = None, + clip_skip: Optional[int] = None, + callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None, + callback_on_step_end_tensor_inputs: List[str] = ["latents"], + previewer_scheduler: KarrasDiffusionSchedulers = None, + reference_latents: Optional[torch.FloatTensor] = None, + **kwargs, + ): + r""" + The call function to the pipeline for generation. + + Args: + prompt (`str` or `List[str]`, *optional*): + The prompt or prompts to guide image generation. If not defined, you need to pass `prompt_embeds`. + prompt_2 (`str` or `List[str]`, *optional*): + The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is + used in both text-encoders. + image (`torch.FloatTensor`, `PIL.Image.Image`, `np.ndarray`, `List[torch.FloatTensor]`, `List[PIL.Image.Image]`, `List[np.ndarray]`,: + `List[List[torch.FloatTensor]]`, `List[List[np.ndarray]]` or `List[List[PIL.Image.Image]]`): + The ControlNet input condition to provide guidance to the `unet` for generation. If the type is + specified as `torch.FloatTensor`, it is passed to ControlNet as is. `PIL.Image.Image` can also be + accepted as an image. The dimensions of the output image defaults to `image`'s dimensions. If height + and/or width are passed, `image` is resized accordingly. If multiple ControlNets are specified in + `init`, images must be passed as a list such that each element of the list can be correctly batched for + input to a single ControlNet. + height (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`): + The height in pixels of the generated image. Anything below 512 pixels won't work well for + [stabilityai/stable-diffusion-xl-base-1.0](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) + and checkpoints that are not specifically fine-tuned on low resolutions. + width (`int`, *optional*, defaults to `self.unet.config.sample_size * self.vae_scale_factor`): + The width in pixels of the generated image. Anything below 512 pixels won't work well for + [stabilityai/stable-diffusion-xl-base-1.0](https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0) + and checkpoints that are not specifically fine-tuned on low resolutions. + num_inference_steps (`int`, *optional*, defaults to 50): + The number of denoising steps. More denoising steps usually lead to a higher quality image at the + expense of slower inference. + timesteps (`List[int]`, *optional*): + Custom timesteps to use for the denoising process with schedulers which support a `timesteps` argument + in their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is + passed will be used. Must be in descending order. + denoising_end (`float`, *optional*): + When specified, determines the fraction (between 0.0 and 1.0) of the total denoising process to be + completed before it is intentionally prematurely terminated. As a result, the returned sample will + still retain a substantial amount of noise as determined by the discrete timesteps selected by the + scheduler. The denoising_end parameter should ideally be utilized when this pipeline forms a part of a + "Mixture of Denoisers" multi-pipeline setup, as elaborated in [**Refining the Image + Output**](https://huggingface.co/docs/diffusers/api/pipelines/stable_diffusion/stable_diffusion_xl#refining-the-image-output) + guidance_scale (`float`, *optional*, defaults to 5.0): + A higher guidance scale value encourages the model to generate images closely linked to the text + `prompt` at the expense of lower image quality. Guidance scale is enabled when `guidance_scale > 1`. + negative_prompt (`str` or `List[str]`, *optional*): + The prompt or prompts to guide what to not include in image generation. If not defined, you need to + pass `negative_prompt_embeds` instead. Ignored when not using guidance (`guidance_scale < 1`). + negative_prompt_2 (`str` or `List[str]`, *optional*): + The prompt or prompts to guide what to not include in image generation. This is sent to `tokenizer_2` + and `text_encoder_2`. If not defined, `negative_prompt` is used in both text-encoders. + num_images_per_prompt (`int`, *optional*, defaults to 1): + The number of images to generate per prompt. + eta (`float`, *optional*, defaults to 0.0): + Corresponds to parameter eta (η) from the [DDIM](https://arxiv.org/abs/2010.02502) paper. Only applies + to the [`~schedulers.DDIMScheduler`], and is ignored in other schedulers. + generator (`torch.Generator` or `List[torch.Generator]`, *optional*): + A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make + generation deterministic. + latents (`torch.FloatTensor`, *optional*): + Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image + generation. Can be used to tweak the same generation with different prompts. If not provided, a latents + tensor is generated by sampling using the supplied random `generator`. + prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not + provided, text embeddings are generated from the `prompt` input argument. + negative_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated negative text embeddings. Can be used to easily tweak text inputs (prompt weighting). If + not provided, `negative_prompt_embeds` are generated from the `negative_prompt` input argument. + pooled_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated pooled text embeddings. Can be used to easily tweak text inputs (prompt weighting). If + not provided, pooled text embeddings are generated from `prompt` input argument. + negative_pooled_prompt_embeds (`torch.FloatTensor`, *optional*): + Pre-generated negative pooled text embeddings. Can be used to easily tweak text inputs (prompt + weighting). If not provided, pooled `negative_prompt_embeds` are generated from `negative_prompt` input + argument. + ip_adapter_image: (`PipelineImageInput`, *optional*): Optional image input to work with IP Adapters. + ip_adapter_image_embeds (`List[torch.FloatTensor]`, *optional*): + Pre-generated image embeddings for IP-Adapter. It should be a list of length same as number of + IP-adapters. Each element should be a tensor of shape `(batch_size, num_images, emb_dim)`. It should + contain the negative image embedding if `do_classifier_free_guidance` is set to `True`. If not + provided, embeddings are computed from the `ip_adapter_image` input argument. + output_type (`str`, *optional*, defaults to `"pil"`): + The output format of the generated image. Choose between `PIL.Image` or `np.array`. + return_dict (`bool`, *optional*, defaults to `True`): + Whether or not to return a [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] instead of a + plain tuple. + cross_attention_kwargs (`dict`, *optional*): + A kwargs dictionary that if specified is passed along to the [`AttentionProcessor`] as defined in + [`self.processor`](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py). + controlnet_conditioning_scale (`float` or `List[float]`, *optional*, defaults to 1.0): + The outputs of the ControlNet are multiplied by `controlnet_conditioning_scale` before they are added + to the residual in the original `unet`. If multiple ControlNets are specified in `init`, you can set + the corresponding scale as a list. + control_guidance_start (`float` or `List[float]`, *optional*, defaults to 0.0): + The percentage of total steps at which the ControlNet starts applying. + control_guidance_end (`float` or `List[float]`, *optional*, defaults to 1.0): + The percentage of total steps at which the ControlNet stops applying. + original_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)): + If `original_size` is not the same as `target_size` the image will appear to be down- or upsampled. + `original_size` defaults to `(height, width)` if not specified. Part of SDXL's micro-conditioning as + explained in section 2.2 of + [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). + crops_coords_top_left (`Tuple[int]`, *optional*, defaults to (0, 0)): + `crops_coords_top_left` can be used to generate an image that appears to be "cropped" from the position + `crops_coords_top_left` downwards. Favorable, well-centered images are usually achieved by setting + `crops_coords_top_left` to (0, 0). Part of SDXL's micro-conditioning as explained in section 2.2 of + [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). + target_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)): + For most cases, `target_size` should be set to the desired height and width of the generated image. If + not specified it will default to `(height, width)`. Part of SDXL's micro-conditioning as explained in + section 2.2 of [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). + negative_original_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)): + To negatively condition the generation process based on a specific image resolution. Part of SDXL's + micro-conditioning as explained in section 2.2 of + [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more + information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208. + negative_crops_coords_top_left (`Tuple[int]`, *optional*, defaults to (0, 0)): + To negatively condition the generation process based on a specific crop coordinates. Part of SDXL's + micro-conditioning as explained in section 2.2 of + [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more + information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208. + negative_target_size (`Tuple[int]`, *optional*, defaults to (1024, 1024)): + To negatively condition the generation process based on a target image resolution. It should be as same + as the `target_size` for most cases. Part of SDXL's micro-conditioning as explained in section 2.2 of + [https://huggingface.co/papers/2307.01952](https://huggingface.co/papers/2307.01952). For more + information, refer to this issue thread: https://github.com/huggingface/diffusers/issues/4208. + clip_skip (`int`, *optional*): + Number of layers to be skipped from CLIP while computing the prompt embeddings. A value of 1 means that + the output of the pre-final layer will be used for computing the prompt embeddings. + callback_on_step_end (`Callable`, *optional*): + A function that calls at the end of each denoising steps during the inference. The function is called + with the following arguments: `callback_on_step_end(self: DiffusionPipeline, step: int, timestep: int, + callback_kwargs: Dict)`. `callback_kwargs` will include a list of all tensors as specified by + `callback_on_step_end_tensor_inputs`. + callback_on_step_end_tensor_inputs (`List`, *optional*): + The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list + will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the + `._callback_tensor_inputs` attribute of your pipeline class. + + Examples: + + Returns: + [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] or `tuple`: + If `return_dict` is `True`, [`~pipelines.stable_diffusion.StableDiffusionPipelineOutput`] is returned, + otherwise a `tuple` is returned containing the output images. + """ + + callback = kwargs.pop("callback", None) + callback_steps = kwargs.pop("callback_steps", None) + + if callback is not None: + deprecate( + "callback", + "1.0.0", + "Passing `callback` as an input argument to `__call__` is deprecated, consider using `callback_on_step_end`", + ) + if callback_steps is not None: + deprecate( + "callback_steps", + "1.0.0", + "Passing `callback_steps` as an input argument to `__call__` is deprecated, consider using `callback_on_step_end`", + ) + + aggregator = self.aggregator._orig_mod if is_compiled_module(self.aggregator) else self.aggregator + if not isinstance(ip_adapter_image, list): + ip_adapter_image = [ip_adapter_image] if ip_adapter_image is not None else [image] + + # 1. Check inputs. Raise error if not correct + self.check_inputs( + prompt, + prompt_2, + image, + callback_steps, + negative_prompt, + negative_prompt_2, + prompt_embeds, + negative_prompt_embeds, + pooled_prompt_embeds, + ip_adapter_image, + ip_adapter_image_embeds, + negative_pooled_prompt_embeds, + controlnet_conditioning_scale, + control_guidance_start, + control_guidance_end, + callback_on_step_end_tensor_inputs, + ) + + self._guidance_scale = guidance_scale + self._guidance_rescale = guidance_rescale + self._clip_skip = clip_skip + self._cross_attention_kwargs = cross_attention_kwargs + self._denoising_end = denoising_end + + # 2. Define call parameters + if prompt is not None and isinstance(prompt, str): + if not isinstance(image, PIL.Image.Image): + batch_size = len(image) + else: + batch_size = 1 + prompt = [prompt] * batch_size + elif prompt is not None and isinstance(prompt, list): + batch_size = len(prompt) + assert batch_size == len(image) or (isinstance(image, PIL.Image.Image) or len(image) == 1) + else: + batch_size = prompt_embeds.shape[0] + assert batch_size == len(image) or (isinstance(image, PIL.Image.Image) or len(image) == 1) + + device = self._execution_device + + # 3.1 Encode input prompt + text_encoder_lora_scale = ( + self.cross_attention_kwargs.get("scale", None) if self.cross_attention_kwargs is not None else None + ) + ( + prompt_embeds, + negative_prompt_embeds, + pooled_prompt_embeds, + negative_pooled_prompt_embeds, + ) = self.encode_prompt( + prompt=prompt, + prompt_2=prompt_2, + device=device, + num_images_per_prompt=num_images_per_prompt, + do_classifier_free_guidance=self.do_classifier_free_guidance, + negative_prompt=negative_prompt, + negative_prompt_2=negative_prompt_2, + prompt_embeds=prompt_embeds, + negative_prompt_embeds=negative_prompt_embeds, + pooled_prompt_embeds=pooled_prompt_embeds, + negative_pooled_prompt_embeds=negative_pooled_prompt_embeds, + lora_scale=text_encoder_lora_scale, + clip_skip=self.clip_skip, + ) + # 3.2 Encode ip_adapter_image + if ip_adapter_image is not None or ip_adapter_image_embeds is not None: + image_embeds = self.prepare_ip_adapter_image_embeds( + ip_adapter_image, + ip_adapter_image_embeds, + device, + batch_size * num_images_per_prompt, + self.do_classifier_free_guidance, + ) + + # 4. Prepare image + image = self.prepare_image( + image=image, + width=width, + height=height, + batch_size=batch_size * num_images_per_prompt, + num_images_per_prompt=num_images_per_prompt, + device=device, + dtype=aggregator.dtype, + do_classifier_free_guidance=self.do_classifier_free_guidance, + ) + height, width = image.shape[-2:] + if image.shape[1] != 4: + needs_upcasting = self.vae.dtype == torch.float16 and self.vae.config.force_upcast + if needs_upcasting: + image = image.float() + self.vae.to(dtype=torch.float32) + image = self.vae.encode(image).latent_dist.sample() + image = image * self.vae.config.scaling_factor + if needs_upcasting: + self.vae.to(dtype=torch.float16) + image = image.to(dtype=torch.float16) + else: + height = int(height * self.vae_scale_factor) + width = int(width * self.vae_scale_factor) + + # 5. Prepare timesteps + timesteps, num_inference_steps = retrieve_timesteps(self.scheduler, num_inference_steps, device, timesteps) + + # 6. Prepare latent variables + if init_latents_with_lq: + latents = self.init_latents(image, generator, timesteps[0]) + else: + num_channels_latents = self.unet.config.in_channels + latents = self.prepare_latents( + batch_size * num_images_per_prompt, + num_channels_latents, + height, + width, + prompt_embeds.dtype, + device, + generator, + latents, + ) + + # 6.5 Optionally get Guidance Scale Embedding + timestep_cond = None + if self.unet.config.time_cond_proj_dim is not None: + guidance_scale_tensor = torch.tensor(self.guidance_scale - 1).repeat(batch_size * num_images_per_prompt) + timestep_cond = self.get_guidance_scale_embedding( + guidance_scale_tensor, embedding_dim=self.unet.config.time_cond_proj_dim + ).to(device=device, dtype=latents.dtype) + + # 7. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline + extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) + + # 7.1 Create tensor stating which controlnets to keep + controlnet_keep = [] + previewing = [] + for i in range(len(timesteps)): + keeps = 1.0 - float(i / len(timesteps) < control_guidance_start or (i + 1) / len(timesteps) > control_guidance_end) + controlnet_keep.append(keeps) + use_preview = 1.0 - float(i / len(timesteps) < preview_start or (i + 1) / len(timesteps) > preview_end) + previewing.append(use_preview) + if isinstance(controlnet_conditioning_scale, list): + assert len(controlnet_conditioning_scale) == len(timesteps), f"{len(controlnet_conditioning_scale)} controlnet scales do not match number of sampling steps {len(timesteps)}" + else: + controlnet_conditioning_scale = [controlnet_conditioning_scale] * len(controlnet_keep) + + # 7.2 Prepare added time ids & embeddings + original_size = original_size or (height, width) + target_size = target_size or (height, width) + + add_text_embeds = pooled_prompt_embeds + if self.text_encoder_2 is None: + text_encoder_projection_dim = int(pooled_prompt_embeds.shape[-1]) + else: + text_encoder_projection_dim = self.text_encoder_2.config.projection_dim + + add_time_ids = self._get_add_time_ids( + original_size, + crops_coords_top_left, + target_size, + dtype=prompt_embeds.dtype, + text_encoder_projection_dim=text_encoder_projection_dim, + ) + + if negative_original_size is not None and negative_target_size is not None: + negative_add_time_ids = self._get_add_time_ids( + negative_original_size, + negative_crops_coords_top_left, + negative_target_size, + dtype=prompt_embeds.dtype, + text_encoder_projection_dim=text_encoder_projection_dim, + ) + else: + negative_add_time_ids = add_time_ids + + if self.do_classifier_free_guidance: + prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0) + add_text_embeds = torch.cat([negative_pooled_prompt_embeds, add_text_embeds], dim=0) + add_time_ids = torch.cat([negative_add_time_ids, add_time_ids], dim=0) + image = torch.cat([image] * 2, dim=0) + + prompt_embeds = prompt_embeds.to(device) + add_text_embeds = add_text_embeds.to(device) + add_time_ids = add_time_ids.to(device).repeat(batch_size * num_images_per_prompt, 1) + + # 8. Denoising loop + num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) + + # 8.1 Apply denoising_end + if ( + self.denoising_end is not None + and isinstance(self.denoising_end, float) + and self.denoising_end > 0 + and self.denoising_end < 1 + ): + discrete_timestep_cutoff = int( + round( + self.scheduler.config.num_train_timesteps + - (self.denoising_end * self.scheduler.config.num_train_timesteps) + ) + ) + num_inference_steps = len(list(filter(lambda ts: ts >= discrete_timestep_cutoff, timesteps))) + timesteps = timesteps[:num_inference_steps] + + is_unet_compiled = is_compiled_module(self.unet) + is_aggregator_compiled = is_compiled_module(self.aggregator) + is_torch_higher_equal_2_1 = is_torch_version(">=", "2.1") + previewer_mean = torch.zeros_like(latents) + unet_mean = torch.zeros_like(latents) + preview_factor = torch.ones( + (latents.shape[0], *((1,) * (len(latents.shape) - 1))), dtype=latents.dtype, device=latents.device + ) + + self._num_timesteps = len(timesteps) + preview_row = [] + with self.progress_bar(total=num_inference_steps) as progress_bar: + for i, t in enumerate(timesteps): + # Relevant thread: + # https://dev-discuss.pytorch.org/t/cudagraphs-in-pytorch-2-0/1428 + if (is_unet_compiled and is_aggregator_compiled) and is_torch_higher_equal_2_1: + torch._inductor.cudagraph_mark_step_begin() + # expand the latents if we are doing classifier free guidance + latent_model_input = torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents + latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) + prev_t = t + unet_model_input = latent_model_input + + added_cond_kwargs = { + "text_embeds": add_text_embeds, + "time_ids": add_time_ids, + "image_embeds": image_embeds + } + aggregator_added_cond_kwargs = {"text_embeds": add_text_embeds, "time_ids": add_time_ids} + + # prepare time_embeds in advance as adapter input + cross_attention_t_emb = self.unet.get_time_embed(sample=latent_model_input, timestep=t) + cross_attention_emb = self.unet.time_embedding(cross_attention_t_emb, timestep_cond) + cross_attention_aug_emb = None + + cross_attention_aug_emb = self.unet.get_aug_embed( + emb=cross_attention_emb, + encoder_hidden_states=prompt_embeds, + added_cond_kwargs=added_cond_kwargs + ) + + cross_attention_emb = cross_attention_emb + cross_attention_aug_emb if cross_attention_aug_emb is not None else cross_attention_emb + + if self.unet.time_embed_act is not None: + cross_attention_emb = self.unet.time_embed_act(cross_attention_emb) + + current_cross_attention_kwargs = {"temb": cross_attention_emb} + if cross_attention_kwargs is not None: + for k,v in cross_attention_kwargs.items(): + current_cross_attention_kwargs[k] = v + self._cross_attention_kwargs = current_cross_attention_kwargs + + # adaptive restoration factors + adaRes_scale = preview_factor.to(latent_model_input.dtype).clamp(0.0, controlnet_conditioning_scale[i]) + cond_scale = adaRes_scale * controlnet_keep[i] + cond_scale = torch.cat([cond_scale] * 2) if self.do_classifier_free_guidance else cond_scale + + if (cond_scale>0.1).sum().item() > 0: + if previewing[i] > 0: + # preview with LCM + self.unet.enable_adapters() + preview_noise = self.unet( + latent_model_input, + t, + encoder_hidden_states=prompt_embeds, + timestep_cond=timestep_cond, + cross_attention_kwargs=self.cross_attention_kwargs, + added_cond_kwargs=added_cond_kwargs, + return_dict=False, + )[0] + preview_latent = previewer_scheduler.step( + preview_noise, + t.to(dtype=torch.int64), + # torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents, + latent_model_input, # scaled latents here for compatibility + return_dict=False + )[0] + self.unet.disable_adapters() + + if self.do_classifier_free_guidance: + preview_row.append(preview_latent.chunk(2)[1].to('cpu')) + else: + preview_row.append(preview_latent.to('cpu')) + # Prepare 2nd order step. + if multistep_restore and i+1 < len(timesteps): + noise_preview = preview_noise.chunk(2)[1] if self.do_classifier_free_guidance else preview_noise + first_step = self.scheduler.step( + noise_preview, t, latents, + **extra_step_kwargs, return_dict=True, step_forward=False + ) + prev_t = timesteps[i + 1] + unet_model_input = torch.cat([first_step.prev_sample] * 2) if self.do_classifier_free_guidance else first_step.prev_sample + unet_model_input = self.scheduler.scale_model_input(unet_model_input, prev_t, heun_step=True) + + elif reference_latents is not None: + preview_latent = torch.cat([reference_latents] * 2) if self.do_classifier_free_guidance else reference_latents + else: + preview_latent = image + + # Add fresh noise + # preview_noise = torch.randn_like(preview_latent) + # preview_latent = self.scheduler.add_noise(preview_latent, preview_noise, t) + + preview_latent=preview_latent.to(dtype=next(aggregator.parameters()).dtype) + + # Aggregator inference + down_block_res_samples, mid_block_res_sample = aggregator( + image, + prev_t, + encoder_hidden_states=prompt_embeds, + controlnet_cond=preview_latent, + # conditioning_scale=cond_scale, + added_cond_kwargs=aggregator_added_cond_kwargs, + return_dict=False, + ) + + # aggregator features scaling + down_block_res_samples = [sample*cond_scale for sample in down_block_res_samples] + mid_block_res_sample = mid_block_res_sample*cond_scale + + # predict the noise residual + noise_pred = self.unet( + unet_model_input, + prev_t, + encoder_hidden_states=prompt_embeds, + timestep_cond=timestep_cond, + cross_attention_kwargs=self.cross_attention_kwargs, + down_block_additional_residuals=down_block_res_samples, + mid_block_additional_residual=mid_block_res_sample, + added_cond_kwargs=added_cond_kwargs, + return_dict=False, + )[0] + + # perform guidance + if self.do_classifier_free_guidance: + noise_pred_uncond, noise_pred_text = noise_pred.chunk(2) + noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond) + + if self.do_classifier_free_guidance and self.guidance_rescale > 0.0: + # Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf + noise_pred = rescale_noise_cfg(noise_pred, noise_pred_text, guidance_rescale=self.guidance_rescale) + + # compute the previous noisy sample x_t -> x_t-1 + latents_dtype = latents.dtype + unet_step = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=True) + latents = unet_step.prev_sample + + # Update adaRes factors + unet_pred_latent = unet_step.pred_original_sample + + # Adaptive restoration. + if adastep_restore: + pred_x0_l2 = ((preview_latent[latents.shape[0]:].float()-unet_pred_latent.float())).pow(2).sum(dim=(1,2,3)) + previewer_l2 = ((preview_latent[latents.shape[0]:].float()-previewer_mean.float())).pow(2).sum(dim=(1,2,3)) + # unet_l2 = ((unet_pred_latent.float()-unet_mean.float())).pow(2).sum(dim=(1,2,3)).sqrt() + # l2_error = (((preview_latent[latents.shape[0]:]-previewer_mean) - (unet_pred_latent-unet_mean))).pow(2).mean(dim=(1,2,3)) + # preview_error = torch.nn.functional.cosine_similarity(preview_latent[latents.shape[0]:].reshape(latents.shape[0], -1), unet_pred_latent.reshape(latents.shape[0],-1)) + previewer_mean = preview_latent[latents.shape[0]:] + unet_mean = unet_pred_latent + preview_factor = (pred_x0_l2 / previewer_l2).reshape(-1, 1, 1, 1) + + if latents.dtype != latents_dtype: + if torch.backends.mps.is_available(): + # some platforms (eg. apple mps) misbehave due to a pytorch bug: https://github.com/pytorch/pytorch/pull/99272 + latents = latents.to(latents_dtype) + + if callback_on_step_end is not None: + callback_kwargs = {} + for k in callback_on_step_end_tensor_inputs: + callback_kwargs[k] = locals()[k] + callback_outputs = callback_on_step_end(self, i, t, callback_kwargs) + + latents = callback_outputs.pop("latents", latents) + prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds) + negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds) + + # call the callback, if provided + if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0): + progress_bar.update() + if callback is not None and i % callback_steps == 0: + step_idx = i // getattr(self.scheduler, "order", 1) + callback(step_idx, t, latents) + + if not output_type == "latent": + # make sure the VAE is in float32 mode, as it overflows in float16 + needs_upcasting = self.vae.dtype == torch.float16 and self.vae.config.force_upcast + + if needs_upcasting: + self.upcast_vae() + latents = latents.to(next(iter(self.vae.post_quant_conv.parameters())).dtype) + + # unscale/denormalize the latents + # denormalize with the mean and std if available and not None + has_latents_mean = hasattr(self.vae.config, "latents_mean") and self.vae.config.latents_mean is not None + has_latents_std = hasattr(self.vae.config, "latents_std") and self.vae.config.latents_std is not None + if has_latents_mean and has_latents_std: + latents_mean = ( + torch.tensor(self.vae.config.latents_mean).view(1, 4, 1, 1).to(latents.device, latents.dtype) + ) + latents_std = ( + torch.tensor(self.vae.config.latents_std).view(1, 4, 1, 1).to(latents.device, latents.dtype) + ) + latents = latents * latents_std / self.vae.config.scaling_factor + latents_mean + else: + latents = latents / self.vae.config.scaling_factor + + image = self.vae.decode(latents, return_dict=False)[0] + + # cast back to fp16 if needed + if needs_upcasting: + self.vae.to(dtype=torch.float16) + else: + image = latents + + if not output_type == "latent": + # apply watermark if available + if self.watermark is not None: + image = self.watermark.apply_watermark(image) + + image = self.image_processor.postprocess(image, output_type=output_type) + + if save_preview_row: + preview_image_row = [] + if needs_upcasting: + self.upcast_vae() + for preview_latents in preview_row: + preview_latents = preview_latents.to(device=self.device, dtype=next(iter(self.vae.post_quant_conv.parameters())).dtype) + if has_latents_mean and has_latents_std: + latents_mean = ( + torch.tensor(self.vae.config.latents_mean).view(1, 4, 1, 1).to(preview_latents.device, preview_latents.dtype) + ) + latents_std = ( + torch.tensor(self.vae.config.latents_std).view(1, 4, 1, 1).to(preview_latents.device, preview_latents.dtype) + ) + preview_latents = preview_latents * latents_std / self.vae.config.scaling_factor + latents_mean + else: + preview_latents = preview_latents / self.vae.config.scaling_factor + + preview_image = self.vae.decode(preview_latents, return_dict=False)[0] + preview_image = self.image_processor.postprocess(preview_image, output_type=output_type) + preview_image_row.append(preview_image) + + # cast back to fp16 if needed + if needs_upcasting: + self.vae.to(dtype=torch.float16) + + # Offload all models + self.maybe_free_model_hooks() + + if not return_dict: + if save_preview_row: + return (image, preview_image_row) + return (image,) + + return StableDiffusionXLPipelineOutput(images=image) diff --git a/modules/ipadapter.py b/modules/ipadapter.py index 0ab27f03c..4b67b9bb4 100644 --- a/modules/ipadapter.py +++ b/modules/ipadapter.py @@ -3,8 +3,6 @@ Lightweight IP-Adapter applied to existing pipeline in Diffusers - Downloads image_encoder or first usage (2.5GB) - Introduced via: https://github.com/huggingface/diffusers/pull/5713 - IP adapters: https://huggingface.co/h94/IP-Adapter -TODO ipadapter items: -- SD/SDXL autodetect """ import os @@ -14,21 +12,42 @@ from PIL import Image from modules import processing, shared, devices, sd_models -base_repo = "h94/IP-Adapter" +clip_repo = "h94/IP-Adapter" clip_loaded = None -ADAPTERS = { - 'None': 'none', - 'Base': 'ip-adapter_sd15.safetensors', - 'Base ViT-G': 'ip-adapter_sd15_vit-G.safetensors', - 'Light': 'ip-adapter_sd15_light.safetensors', - 'Plus': 'ip-adapter-plus_sd15.safetensors', - 'Plus Face': 'ip-adapter-plus-face_sd15.safetensors', - 'Full Face': 'ip-adapter-full-face_sd15.safetensors', - 'Base SDXL': 'ip-adapter_sdxl.safetensors', - 'Base ViT-H SDXL': 'ip-adapter_sdxl_vit-h.safetensors', - 'Plus ViT-H SDXL': 'ip-adapter-plus_sdxl_vit-h.safetensors', - 'Plus Face ViT-H SDXL': 'ip-adapter-plus-face_sdxl_vit-h.safetensors', +ADAPTERS_NONE = { + 'None': { 'name': 'none', 'repo': 'none', 'subfolder': 'none' }, } +ADAPTERS_SD15 = { + 'None': { 'name': 'none', 'repo': 'none', 'subfolder': 'none' }, + 'Base': { 'name': 'ip-adapter_sd15.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' }, + 'Base ViT-G': { 'name': 'ip-adapter_sd15_vit-G.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' }, + 'Light': { 'name': 'ip-adapter_sd15_light.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' }, + 'Plus': { 'name': 'ip-adapter-plus_sd15.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' }, + 'Plus Face': { 'name': 'ip-adapter-plus-face_sd15.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' }, + 'Full Face': { 'name': 'ip-adapter-full-face_sd15.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' }, + 'Ostris Composition ViT-H': { 'name': 'ip_plus_composition_sd15.safetensors', 'repo': 'ostris/ip-composition-adapter', 'subfolder': '' }, +} +ADAPTERS_SDXL = { + 'None': { 'name': 'none', 'repo': 'none', 'subfolder': 'none' }, + 'Base SDXL': { 'name': 'ip-adapter_sdxl.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'sdxl_models' }, + 'Base ViT-H SDXL': { 'name': 'ip-adapter_sdxl_vit-h.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'sdxl_models' }, + 'Plus ViT-H SDXL': { 'name': 'ip-adapter-plus_sdxl_vit-h.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'sdxl_models' }, + 'Plus Face ViT-H SDXL': { 'name': 'ip-adapter-plus-face_sdxl_vit-h.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'sdxl_models' }, + 'Ostris Composition ViT-H SDXL': { 'name': 'ip_plus_composition_sdxl.safetensors', 'repo': 'ostris/ip-composition-adapter', 'subfolder': '' }, +} +ADAPTERS = { **ADAPTERS_SD15, **ADAPTERS_SDXL } +ADAPTERS_ALL = { **ADAPTERS_SD15, **ADAPTERS_SDXL } + + +def get_adapters(): + global ADAPTERS # pylint: disable=global-statement + if shared.sd_model_type == 'sd': + ADAPTERS = ADAPTERS_SD15 + elif shared.sd_model_type == 'sdxl': + ADAPTERS = ADAPTERS_SDXL + else: + ADAPTERS = ADAPTERS_NONE + return list(ADAPTERS) def get_images(input_images): @@ -83,13 +102,12 @@ def crop_images(images, crops): try: for i in range(len(images)): if crops[i]: - from shared import yolo # pylint: disable=no-name-in-module - yolo.load() + from modules.shared import yolo # pylint: disable=no-name-in-module cropped = [] for image in images[i]: - faces = yolo.predict(image) + faces = yolo.predict('face-yolo8n', image) if len(faces) > 0: - cropped.append(faces[0].face) + cropped.append(faces[0].item) if len(cropped) == len(images[i]): images[i] = cropped else: @@ -117,13 +135,13 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt if hasattr(p, 'ip_adapter_names'): if isinstance(p.ip_adapter_names, str): p.ip_adapter_names = [p.ip_adapter_names] - adapters = [ADAPTERS.get(adapter, None) for adapter in p.ip_adapter_names if adapter is not None and adapter.lower() != 'none'] + adapters = [ADAPTERS_ALL.get(adapter_name, None) for adapter_name in p.ip_adapter_names if adapter_name is not None and adapter_name.lower() != 'none'] adapter_names = p.ip_adapter_names else: if isinstance(adapter_names, str): adapter_names = [adapter_names] adapters = [ADAPTERS.get(adapter, None) for adapter in adapter_names] - adapters = [adapter for adapter in adapters if adapter is not None and adapter.lower() != 'none'] + adapters = [adapter for adapter in adapters if adapter is not None and adapter['name'].lower() != 'none'] if len(adapters) == 0: unapply(pipe) if hasattr(p, 'ip_adapter_images'): @@ -189,41 +207,48 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt for adapter_name in adapter_names: # which clip to use - if 'ViT' not in adapter_name: - clip_repo = base_repo - clip_subfolder = 'models/image_encoder' if shared.sd_model_type == 'sd' else 'sdxl_models/image_encoder' # defaults per model + if 'ViT' not in adapter_name: # defaults per model + if shared.sd_model_type == 'sd': + clip_subfolder = 'models/image_encoder' + else: + clip_subfolder = 'sdxl_models/image_encoder' elif 'ViT-H' in adapter_name: - clip_repo = base_repo clip_subfolder = 'models/image_encoder' # this is vit-h elif 'ViT-G' in adapter_name: - clip_repo = base_repo clip_subfolder = 'sdxl_models/image_encoder' # this is vit-g else: shared.log.error(f'IP adapter: unknown model type: {adapter_name}') return False - # load feature extractor used by ip adapter - if pipe.feature_extractor is None: + # load feature extractor used by ip adapter + if pipe.feature_extractor is None: + try: from transformers import CLIPImageProcessor shared.log.debug('IP adapter load: feature extractor') pipe.feature_extractor = CLIPImageProcessor() - # load image encoder used by ip adapter - if pipe.image_encoder is None or clip_loaded != f'{clip_repo}/{clip_subfolder}': - try: - from transformers import CLIPVisionModelWithProjection - shared.log.debug(f'IP adapter load: image encoder="{clip_repo}/{clip_subfolder}"') - pipe.image_encoder = CLIPVisionModelWithProjection.from_pretrained(clip_repo, subfolder=clip_subfolder, torch_dtype=devices.dtype, cache_dir=shared.opts.diffusers_dir, use_safetensors=True) - clip_loaded = f'{clip_repo}/{clip_subfolder}' - except Exception as e: - shared.log.error(f'IP adapter: failed to load image encoder: {e}') - return False - sd_models.move_model(pipe.image_encoder, devices.device) + except Exception as e: + shared.log.error(f'IP adapter load: feature extractor {e}') + return False + + # load image encoder used by ip adapter + if pipe.image_encoder is None or clip_loaded != f'{clip_repo}/{clip_subfolder}': + try: + from transformers import CLIPVisionModelWithProjection + shared.log.debug(f'IP adapter load: image encoder="{clip_repo}/{clip_subfolder}"') + pipe.image_encoder = CLIPVisionModelWithProjection.from_pretrained(clip_repo, subfolder=clip_subfolder, torch_dtype=devices.dtype, cache_dir=shared.opts.diffusers_dir, use_safetensors=True) + clip_loaded = f'{clip_repo}/{clip_subfolder}' + except Exception as e: + shared.log.error(f'IP adapter load: image encoder="{clip_repo}/{clip_subfolder}" {e}') + return False + sd_models.move_model(pipe.image_encoder, devices.device) # main code - t0 = time.time() - ip_subfolder = 'models' if shared.sd_model_type == 'sd' else 'sdxl_models' try: - pipe.load_ip_adapter([base_repo], subfolder=[ip_subfolder], weight_name=adapters) + t0 = time.time() + repos = [adapter['repo'] for adapter in adapters] + subfolders = [adapter['subfolder'] for adapter in adapters] + names = [adapter['name'] for adapter in adapters] + pipe.load_ip_adapter(repos, subfolder=subfolders, weight_name=names) if hasattr(p, 'ip_adapter_layers'): pipe.set_ip_adapter_scale(p.ip_adapter_layers) ip_str = ';'.join(adapter_names) + ':' + json.dumps(p.ip_adapter_layers) @@ -240,5 +265,5 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt t1 = time.time() shared.log.info(f'IP adapter: {ip_str} image={adapter_images} mask={adapter_masks is not None} time={t1-t0:.2f}') except Exception as e: - shared.log.error(f'IP adapter failed to load: repo="{base_repo}" folder="{ip_subfolder}" weights={adapters} names={adapter_names} {e}') + shared.log.error(f'IP adapter load: adapters={adapter_names} repo={repos} folders={subfolders} names={names} {e}') return True diff --git a/modules/loader.py b/modules/loader.py index a2970abfd..0711c2906 100644 --- a/modules/loader.py +++ b/modules/loader.py @@ -44,6 +44,8 @@ if ".dev" in torch.__version__ or "+git" in torch.__version__: timer.startup.record("torch") import transformers # pylint: disable=W0611,C0411 +from transformers import logging as transformers_logging # pylint: disable=W0611,C0411 +transformers_logging.set_verbosity_error() timer.startup.record("transformers") import accelerate # pylint: disable=W0611,C0411 @@ -61,6 +63,9 @@ errors.install([gradio]) import pydantic # pylint: disable=W0611,C0411 timer.startup.record("pydantic") +import diffusers.utils.import_utils # pylint: disable=W0611,C0411 +diffusers.utils.import_utils._k_diffusion_available = True # pylint: disable=protected-access # monkey-patch since we use k-diffusion from git +diffusers.utils.import_utils._k_diffusion_version = '0.0.12' # pylint: disable=protected-access import diffusers # pylint: disable=W0611,C0411 import diffusers.loaders.single_file # pylint: disable=W0611,C0411 import huggingface_hub # pylint: disable=W0611,C0411 diff --git a/modules/model_flux.py b/modules/model_flux.py index 38207f73b..c605702c8 100644 --- a/modules/model_flux.py +++ b/modules/model_flux.py @@ -122,10 +122,12 @@ def quant_flux_bnb(checkpoint_info, transformer, text_encoder_2): bnb_4bit_quant_type=shared.opts.bnb_quantization_type, bnb_4bit_compute_dtype=devices.dtype ) - if 'Model' in shared.opts.bnb_quantization and transformer is None: + if ('Model' in shared.opts.bnb_quantization) and (transformer is None): transformer = diffusers.FluxTransformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype) shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') - if 'Text Encoder' in shared.opts.bnb_quantization and text_encoder_2 is None: + if ('Text Encoder' in shared.opts.bnb_quantization) and (text_encoder_2 is None): + if repo_id == 'sayakpaul/flux.1-dev-nf4': + repo_id = 'black-forest-labs/FLUX.1-dev' # workaround since sayakpaul model is missing model_index.json text_encoder_2 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_2", cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype) shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') except Exception as e: @@ -285,25 +287,26 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch errors.display(e, 'FLUX Quanto:') # initialize pipeline with pre-loaded components - components = {} - transformer, text_encoder_2 = quant_flux_bnb(checkpoint_info, transformer, text_encoder_2) + kwargs = {} + # transformer, text_encoder_2 = quant_flux_bnb(checkpoint_info, transformer, text_encoder_2) if transformer is not None: - components['transformer'] = transformer + kwargs['transformer'] = transformer sd_unet.loaded_unet = shared.opts.sd_unet if text_encoder_1 is not None: - components['text_encoder'] = text_encoder_1 + kwargs['text_encoder'] = text_encoder_1 model_te.loaded_te = shared.opts.sd_text_encoder if text_encoder_2 is not None: - components['text_encoder_2'] = text_encoder_2 + kwargs['text_encoder_2'] = text_encoder_2 model_te.loaded_te = shared.opts.sd_text_encoder if vae is not None: - components['vae'] = vae - shared.log.debug(f'Load model: type=FLUX preloaded={list(components)}') + kwargs['vae'] = vae + shared.log.debug(f'Load model: type=FLUX preloaded={list(kwargs)}') if repo_id == 'sayakpaul/flux.1-dev-nf4': repo_id = 'black-forest-labs/FLUX.1-dev' # workaround since sayakpaul model is missing model_index.json - for c in components: - if components[c].dtype == torch.float32 and devices.dtype != torch.float32: - shared.log.warning(f'Load model: type=FLUX component={c} dtype={components[c].dtype} cast dtype={devices.dtype}') - components[c] = components[c].to(dtype=devices.dtype) - pipe = diffusers.FluxPipeline.from_pretrained(repo_id, cache_dir=shared.opts.diffusers_dir, **components, **diffusers_load_config) + for c in kwargs: + if kwargs[c].dtype == torch.float32 and devices.dtype != torch.float32: + shared.log.warning(f'Load model: type=FLUX component={c} dtype={kwargs[c].dtype} cast dtype={devices.dtype}') + kwargs[c] = kwargs[c].to(dtype=devices.dtype) + kwargs = model_quant.create_bnb_config(kwargs) + pipe = diffusers.FluxPipeline.from_pretrained(repo_id, cache_dir=shared.opts.diffusers_dir, **kwargs, **diffusers_load_config) return pipe diff --git a/modules/model_quant.py b/modules/model_quant.py index d54d6ff6d..547a3d7ae 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -1,4 +1,5 @@ import sys +import diffusers from installer import install, log @@ -6,6 +7,27 @@ bnb = None quanto = None +def create_bnb_config(kwargs = None): + from modules import shared, devices + if len(shared.opts.bnb_quantization) > 0: + if 'Model' in shared.opts.bnb_quantization and 'transformer' not in (kwargs or {}): + load_bnb() + bnb_config = diffusers.BitsAndBytesConfig( + load_in_8bit=shared.opts.bnb_quantization_type in ['fp8'], + load_in_4bit=shared.opts.bnb_quantization_type in ['nf4', 'fp4'], + bnb_4bit_quant_storage=shared.opts.bnb_quantization_storage, + bnb_4bit_quant_type=shared.opts.bnb_quantization_type, + bnb_4bit_compute_dtype=devices.dtype + ) + shared.log.debug(f'Quantization: module=all type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}') + if kwargs is None: + return bnb_config + else: + kwargs['quantization_config'] = bnb_config + return kwargs + return kwargs + + def load_bnb(msg='', silent=False): global bnb # pylint: disable=global-statement if bnb is not None: @@ -16,6 +38,8 @@ def load_bnb(msg='', silent=False): try: import bitsandbytes bnb = bitsandbytes + diffusers.utils.import_utils._bitsandbytes_available = True # pylint: disable=protected-access + diffusers.utils.import_utils._bitsandbytes_version = '0.43.3' # pylint: disable=protected-access return bnb except Exception as e: if len(msg) > 0: @@ -23,6 +47,7 @@ def load_bnb(msg='', silent=False): bnb = None if not silent: raise + return None def load_quanto(msg='', silent=False): @@ -42,6 +67,7 @@ def load_quanto(msg='', silent=False): quanto = None if not silent: raise + return None def get_quant(name): diff --git a/modules/model_sd3.py b/modules/model_sd3.py index 639f6e4eb..78eee7b4d 100644 --- a/modules/model_sd3.py +++ b/modules/model_sd3.py @@ -1,7 +1,7 @@ import os import diffusers import transformers -from modules import shared, devices, sd_models, sd_unet, model_te +from modules import shared, devices, sd_models, sd_unet, model_te, model_quant, model_tools def load_overrides(kwargs, cache_dir): @@ -51,8 +51,7 @@ def load_overrides(kwargs, cache_dir): def load_quants(kwargs, repo_id, cache_dir): if len(shared.opts.bnb_quantization) > 0: - from modules.model_quant import load_bnb - load_bnb('Load model: type=SD3') + model_quant.load_bnb('Load model: type=SD3') bnb_config = diffusers.BitsAndBytesConfig( load_in_8bit=shared.opts.bnb_quantization_type in ['fp8'], load_in_4bit=shared.opts.bnb_quantization_type in ['nf4', 'fp4'], @@ -70,7 +69,7 @@ def load_quants(kwargs, repo_id, cache_dir): def load_missing(kwargs, fn, cache_dir): - keys = sd_models.get_safetensor_keys(fn) + keys = model_tools.get_safetensor_keys(fn) size = os.stat(fn).st_size // 1024 // 1024 if size > 15000: repo_id = 'stabilityai/stable-diffusion-3.5-large' @@ -85,6 +84,9 @@ def load_missing(kwargs, fn, cache_dir): if 'text_encoder_3' not in kwargs and 'text_encoder_3' not in keys: kwargs['text_encoder_3'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, torch_dtype=devices.dtype) shared.log.debug(f'Load model: type=SD3 missing=te3 repo="{repo_id}"') + if 'vae' not in kwargs and 'vae' not in keys: + kwargs['vae'] = diffusers.AutoencoderKL.from_pretrained(repo_id, subfolder='vae', cache_dir=cache_dir, torch_dtype=devices.dtype) + shared.log.debug(f'Load model: type=SD3 missing=vae repo="{repo_id}"') # if 'transformer' not in kwargs and 'transformer' not in keys: # kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_pretrained(default_repo_id, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype) return kwargs @@ -120,13 +122,18 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None): kwargs = {} kwargs = load_overrides(kwargs, cache_dir) - kwargs = load_quants(kwargs, repo_id, cache_dir) + if fn is None or not os.path.exists(fn): + kwargs = load_quants(kwargs, repo_id, cache_dir) loader = diffusers.StableDiffusion3Pipeline.from_pretrained - if fn is not None and os.path.exists(fn): + if fn is not None and os.path.exists(fn) and os.path.isfile(fn): if fn.endswith('.safetensors'): loader = diffusers.StableDiffusion3Pipeline.from_single_file - kwargs = load_missing(kwargs, fn, cache_dir) + # required_modules = model_tools.get_modules(diffusers.StableDiffusion3Pipeline) + # have_modules = model_tools.get_safetensor_keys(fn) + # loaded_modules = model_tools.load_modules('stabilityai/stable-diffusion-3.5-medium', required_modules) + # kwargs = {**kwargs, **loaded_modules} + # kwargs = load_missing(kwargs, fn, cache_dir) repo_id = fn elif fn.endswith('.gguf'): kwargs = load_gguf(kwargs, fn) @@ -135,8 +142,9 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None): else: kwargs['variant'] = 'fp16' - shared.log.debug(f'Load model: type=SD3 preloaded={list(kwargs)}') + shared.log.debug(f'Load model: type=SD3 kwargs={list(kwargs)} repo="{repo_id}"') + kwargs = model_quant.create_bnb_config(kwargs) pipe = loader( repo_id, torch_dtype=devices.dtype, @@ -144,5 +152,5 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None): config=config, **kwargs, ) - devices.torch_gc() + devices.torch_gc(force=True) return pipe diff --git a/modules/model_tools.py b/modules/model_tools.py new file mode 100644 index 000000000..1d016a19e --- /dev/null +++ b/modules/model_tools.py @@ -0,0 +1,71 @@ +import inspect +import diffusers +import transformers +import safetensors.torch +from modules import shared, devices, model_quant + + +def remove_entries_after_depth(d, depth, current_depth=0): + if current_depth >= depth: + return None + if isinstance(d, dict): + return {k: remove_entries_after_depth(v, depth, current_depth + 1) for k, v in d.items() if remove_entries_after_depth(v, depth, current_depth + 1) is not None} + return d + + +def list_to_dict(flat_list): + result_dict = {} + try: + for item in flat_list: + keys = item.split('.') + d = result_dict + for key in keys[:-1]: + d = d.setdefault(key, {}) + d[keys[-1]] = None + except Exception: + pass + return result_dict + + +def get_safetensor_keys(filename): + keys = [] + try: + with safetensors.torch.safe_open(filename, framework="pt", device="cpu") as f: + keys = f.keys() + except Exception: + pass + return keys + + +def get_modules(model: callable): + signature = inspect.signature(model.__init__, follow_wrapped=True) + params = {param.name: param.annotation for param in signature.parameters.values() if param.annotation != inspect._empty and hasattr(param.annotation, 'from_pretrained')} # pylint: disable=protected-access + for name, cls in params.items(): + shared.log.debug(f'Analyze: model={model} module={name} class={cls.__name__} loadable={getattr(cls, "from_pretrained", None)}') + return params + + +def load_modules(repo_id: str, params: dict): + cache_dir = shared.opts.hfcache_dir + modules = {} + for name, cls in params.items(): + subfolder = None + kwargs = {} + if cls == diffusers.AutoencoderKL: + subfolder = 'vae' + if cls == transformers.CLIPTextModel: # clip-vit-l + subfolder = 'text_encoder' + if cls == transformers.CLIPTextModelWithProjection: # clip-vit-g + subfolder = 'text_encoder_2' + if cls == transformers.T5EncoderModel: # t5-xxl + subfolder = 'text_encoder_3' + kwargs['quantization_config'] = model_quant.create_bnb_config() + kwargs['variant'] = 'fp16' + if cls == diffusers.SD3Transformer2DModel: + subfolder = 'transformer' + kwargs['quantization_config'] = model_quant.create_bnb_config() + if subfolder is None: + continue + shared.log.debug(f'Load: module={name} class={cls.__name__} repo={repo_id} location={subfolder}') + modules[name] = cls.from_pretrained(repo_id, subfolder=subfolder, cache_dir=cache_dir, torch_dtype=devices.dtype, **kwargs) + return modules diff --git a/modules/modelloader.py b/modules/modelloader.py index 70ea7cabb..ce36a739b 100644 --- a/modules/modelloader.py +++ b/modules/modelloader.py @@ -273,6 +273,7 @@ def load_diffusers_models(clear=True): place = os.path.join(models_path, 'Diffusers') if clear: diffuser_repos.clear() + already_found = [] try: for folder in os.listdir(place): try: @@ -303,7 +304,11 @@ def load_diffusers_models(clear=True): if (not os.path.exists(index)) and (not os.path.exists(info)) and (not os.path.exists(config)): debug(f'Diffusers skip model no info: {name}') continue + if name in already_found: + debug(f'Diffusers skip model already found: {name}') + continue repo = { 'name': name, 'filename': name, 'friendly': friendly, 'folder': folder, 'path': commit, 'hash': snapshot, 'mtime': mtime, 'model_info': info, 'model_index': index, 'model_config': config } + already_found.append(name) diffuser_repos.append(repo) if os.path.exists(os.path.join(folder, 'hidden')): continue @@ -318,7 +323,7 @@ def load_diffusers_models(clear=True): def find_diffuser(name: str, full=False): repo = [r for r in diffuser_repos if name == r['name'] or name == r['friendly'] or name == r['path']] if len(repo) > 0: - return repo['name'] + return [repo[0]['name']] hf_api = hf.HfApi() models = list(hf_api.list_models(model_name=name, library=['diffusers'], full=True, limit=20, sort="downloads", direction=-1)) shared.log.debug(f'Searching diffusers models: {name} {len(models) > 0}') diff --git a/modules/onnx_impl/__init__.py b/modules/onnx_impl/__init__.py index 013387d8f..42f22ea3f 100644 --- a/modules/onnx_impl/__init__.py +++ b/modules/onnx_impl/__init__.py @@ -184,8 +184,8 @@ def preprocess_pipeline(p): return shared.sd_model -def ORTDiffusionModelPart_to(self, *args, **kwargs): - self.parent_model = self.parent_model.to(*args, **kwargs) +def ORTPipelinePart_to(self, *args, **kwargs): + self.parent_pipeline = self.parent_pipeline.to(*args, **kwargs) return self @@ -241,9 +241,9 @@ def initialize_onnx(): diffusers.ORTStableDiffusionXLPipeline = diffusers.OnnxStableDiffusionXLPipeline # Huggingface model compatibility diffusers.ORTStableDiffusionXLImg2ImgPipeline = diffusers.OnnxStableDiffusionXLImg2ImgPipeline - optimum.onnxruntime.modeling_diffusion._ORTDiffusionModelPart.to = ORTDiffusionModelPart_to # pylint: disable=protected-access - except Exception: - pass + optimum.onnxruntime.modeling_diffusion.ORTPipelinePart.to = ORTPipelinePart_to # pylint: disable=protected-access + except Exception as e: + log.debug(f'ONNX failed to initialize XL pipelines: {e}') initialized = True diff --git a/modules/onnx_impl/ui.py b/modules/onnx_impl/ui.py index f73e477c4..49af8d98b 100644 --- a/modules/onnx_impl/ui.py +++ b/modules/onnx_impl/ui.py @@ -15,7 +15,7 @@ def create_ui(): from modules.ui_common import create_refresh_button from modules.ui_components import DropdownMulti from modules.shared import log, opts, cmd_opts, refresh_checkpoints - from modules.sd_models import checkpoint_tiles, get_closet_checkpoint_match + from modules.sd_models import checkpoint_titles, get_closet_checkpoint_match from modules.paths import sd_configs_path from .execution_providers import ExecutionProvider, install_execution_provider from .utils import check_diffusers_cache @@ -46,7 +46,7 @@ def create_ui(): with gr.TabItem("Manage cache", id="manage_cache"): cache_state_dirname = gr.Textbox(value=None, visible=False) with gr.Row(): - model_dropdown = gr.Dropdown(label="Model", value="Please select model", choices=checkpoint_tiles()) + model_dropdown = gr.Dropdown(label="Model", value="Please select model", choices=checkpoint_titles()) create_refresh_button(model_dropdown, refresh_checkpoints, {}, "onnx_cache_refresh_diffusers_model") with gr.Row(): def remove_cache_onnx_converted(dirname: str): diff --git a/modules/postprocess/yolo.py b/modules/postprocess/yolo.py index b162240e0..1a52e1334 100644 --- a/modules/postprocess/yolo.py +++ b/modules/postprocess/yolo.py @@ -18,14 +18,13 @@ PREDEFINED = [ # class YoloResult: - def __init__(self, cls: int, label: str, score: float, box: list[int], mask: Image.Image = None, item: Image.Image = None, size: float = 0, width = 0, height = 0, args = {}): + def __init__(self, cls: int, label: str, score: float, box: list[int], mask: Image.Image = None, item: Image.Image = None, width = 0, height = 0, args = {}): self.cls = cls self.label = label self.score = score self.box = box self.mask = mask self.item = item - self.size = size self.width = width self.height = height self.args = args @@ -78,6 +77,10 @@ class YoloRestorer(Detailer): ) -> list[YoloResult]: result = [] + if isinstance(model, str): + model = self.models.get(model, None) + if model is None: + _, model = self.load(model) if model is None: return result args = { @@ -123,21 +126,24 @@ class YoloRestorer(Detailer): box = box.tolist() mask_image = None w, h = box[2] - box[0], box[3] - box[1] - size = w * h / (image.width * image.height) - if (min(w, h) > shared.opts.detailer_min_size if shared.opts.detailer_min_size > 0 else True) and (max(w, h) < shared.opts.detailer_max_size if shared.opts.detailer_max_size > 0 else True): + x_size, y_size = w/image.width, h/image.height + min_size = shared.opts.detailer_min_size if shared.opts.detailer_min_size > 0 and shared.opts.detailer_min_size < 1 else 0 + max_size = shared.opts.detailer_max_size if shared.opts.detailer_max_size > 0 and shared.opts.detailer_max_size < 1 else 1 + if x_size >= min_size and y_size >=min_size and x_size <= max_size and y_size <= max_size: if mask: mask_image = image.copy() mask_image = Image.new('L', image.size, 0) draw = ImageDraw.Draw(mask_image) draw.rectangle(box, fill="white", outline=None, width=0) cropped = image.crop(box) - result.append(YoloResult(cls=cls, label=label, score=round(score, 2), box=box, mask=mask_image, item=cropped, size=size, width=w, height=h, args=args)) + result.append(YoloResult(cls=cls, label=label, score=round(score, 2), box=box, mask=mask_image, item=cropped, width=w, height=h, args=args)) if len(result) >= shared.opts.detailer_max: break return result def load(self, model_name: str = None): from modules import modelloader + model = None self.dependencies() if model_name is None: model_name = list(self.list)[0] @@ -150,15 +156,15 @@ class YoloRestorer(Detailer): try: model_file = modelloader.load_file_from_url(url=model_url, model_dir=shared.opts.yolo_dir, file_name=file_name) if model_file is not None: - from ultralytics import YOLO # pylint: disable=import-outside-toplevel - model = YOLO(model_file) + import ultralytics + model = ultralytics.YOLO(model_file) classes = list(model.names.values()) - shared.log.info(f'Load: type=Detailer name="{model_name}" model="{model_file}" classes={classes}') + shared.log.info(f'Load: type=Detailer name="{model_name}" model="{model_file}" ultralytics={ultralytics.__version__} classes={classes}') self.models[model_name] = model return model_name, model except Exception as e: shared.log.error(f'Load: type=Detailer name="{model_name}" error="{e}"') - return None + return None, None def restore(self, np_image, p: processing.StableDiffusionProcessing = None): if hasattr(p, 'recursion'): @@ -191,12 +197,20 @@ class YoloRestorer(Detailer): pp = None shared.opts.data['mask_apply_overlay'] = True resolution = 512 if shared.sd_model_type in ['none', 'sd', 'lcm', 'unknown'] else 1024 + orig_prompt: str = orig_p.get('all_prompts', [''])[0] + orig_negative: str = orig_p.get('all_negative_prompts', [''])[0] prompt: str = orig_p.get('refiner_prompt', '') negative: str = orig_p.get('refiner_negative', '') if len(prompt) == 0: - prompt = orig_p.get('all_prompts', [''])[0] + prompt = orig_prompt + else: + prompt = prompt.replace('[PROMPT]', orig_prompt) + prompt = prompt.replace('[prompt]', orig_prompt) if len(negative) == 0: - negative = orig_p.get('all_negative_prompts', [''])[0] + negative = orig_negative + else: + negative = negative.replace('[PROMPT]', orig_negative) + negative = negative.replace('[prompt]', orig_negative) prompt_lines = prompt.split('\n') negative_lines = negative.split('\n') prompt = prompt_lines[i % len(prompt_lines)] @@ -315,8 +329,10 @@ class YoloRestorer(Detailer): min_confidence = gr.Slider(label="Min confidence", elem_id=f"{tab}_detailer_conf", value=shared.opts.detailer_conf, minimum=0.0, maximum=1.0, step=0.05) iou = gr.Slider(label="Max overlap", elem_id=f"{tab}_detailer_iou", value=shared.opts.detailer_iou, minimum=0, maximum=1.0, step=0.05) with gr.Row(): - min_size = gr.Slider(label="Min size", elem_id=f"{tab}_detailer_min_size", value=shared.opts.detailer_min_size, minimum=0, maximum=1024, step=1) - max_size = gr.Slider(label="Max size", elem_id=f"{tab}_detailer_max_size", value=shared.opts.detailer_max_size, minimum=0, maximum=1024, step=1) + min_size = shared.opts.detailer_min_size if shared.opts.detailer_min_size < 1 else 0.0 + min_size = gr.Slider(label="Min size", elem_id=f"{tab}_detailer_min_size", value=min_size, minimum=0.0, maximum=1.0, step=0.05) + max_size = shared.opts.detailer_max_size if shared.opts.detailer_max_size < 1 and shared.opts.detailer_max_size > 0 else 1.0 + max_size = gr.Slider(label="Max size", elem_id=f"{tab}_detailer_max_size", value=max_size, minimum=0.0, maximum=1.0, step=0.05) detailers.change(fn=ui_settings_change, inputs=[detailers, classes, strength, padding, blur, min_confidence, max_detected, min_size, max_size, iou], outputs=[]) classes.change(fn=ui_settings_change, inputs=[detailers, classes, strength, padding, blur, min_confidence, max_detected, min_size, max_size, iou], outputs=[]) strength.change(fn=ui_settings_change, inputs=[detailers, classes, strength, padding, blur, min_confidence, max_detected, min_size, max_size, iou], outputs=[]) diff --git a/modules/postprocessing.py b/modules/postprocessing.py index 346758851..967bd04cc 100644 --- a/modules/postprocessing.py +++ b/modules/postprocessing.py @@ -37,7 +37,6 @@ def run_postprocessing(extras_mode, image, image_folder: List[tempfile.NamedTemp image_ext.append(ext) shared.log.debug(f'Process: mode=batch inputs={len(image_folder)} images={len(image_data)}') elif extras_mode == 2: - assert not shared.cmd_opts.hide_ui_dir_config, '--hide-ui-dir-config option must be disabled' assert input_dir, 'input directory not selected' image_list = os.listdir(input_dir) for filename in image_list: diff --git a/modules/processing.py b/modules/processing.py index 04350ee39..0d557e64e 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -4,7 +4,7 @@ import time from contextlib import nullcontext import numpy as np from PIL import Image, ImageOps -from modules import shared, devices, errors, images, scripts, memstats, lowvram, script_callbacks, extra_networks, detailer, sd_hijack_freeu, sd_models, sd_vae, processing_helpers, timer, face_restoration +from modules import shared, devices, errors, images, scripts, memstats, lowvram, script_callbacks, extra_networks, detailer, sd_hijack_freeu, sd_models, sd_checkpoint, sd_vae, processing_helpers, timer, face_restoration, token_merge from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet from modules.processing_class import StableDiffusionProcessing, StableDiffusionProcessingTxt2Img, StableDiffusionProcessingImg2Img, StableDiffusionProcessingControl # pylint: disable=unused-import from modules.processing_info import create_infotext @@ -34,19 +34,20 @@ images_tensor_to_samples = processing_helpers.images_tensor_to_samples class Processed: - def __init__(self, p: StableDiffusionProcessing, images_list, seed=-1, info="", subseed=None, all_prompts=None, all_negative_prompts=None, all_seeds=None, all_subseeds=None, index_of_first_image=0, infotexts=None, comments=""): + def __init__(self, p: StableDiffusionProcessing, images_list, seed=-1, info=None, subseed=None, all_prompts=None, all_negative_prompts=None, all_seeds=None, all_subseeds=None, index_of_first_image=0, infotexts=None, comments=""): self.images = images_list self.prompt = p.prompt or '' self.negative_prompt = p.negative_prompt or '' self.seed = seed if seed != -1 else p.seed self.subseed = subseed self.subseed_strength = p.subseed_strength - self.info = info + self.info = info or create_infotext(p) self.comments = comments or '' self.width = p.width if hasattr(p, 'width') else (self.images[0].width if len(self.images) > 0 else 0) self.height = p.height if hasattr(p, 'height') else (self.images[0].height if len(self.images) > 0 else 0) self.sampler_name = p.sampler_name or '' - self.cfg_scale = p.cfg_scale or 0 + self.cfg_scale = p.cfg_scale if p.cfg_scale > 1 else None + self.cfg_end = p.cfg_end if p.cfg_end < 0 else None self.image_cfg_scale = p.image_cfg_scale or 0 self.steps = p.steps or 0 self.batch_size = max(1, p.batch_size) @@ -79,7 +80,7 @@ class Processed: self.all_negative_prompts = all_negative_prompts or p.all_negative_prompts or [self.negative_prompt] self.all_seeds = all_seeds or p.all_seeds or [self.seed] self.all_subseeds = all_subseeds or p.all_subseeds or [self.subseed] - self.infotexts = infotexts or [info] + self.infotexts = infotexts or [self.info] def js(self): obj = { @@ -96,6 +97,7 @@ class Processed: "height": self.height, "sampler_name": self.sampler_name, "cfg_scale": self.cfg_scale, + "cfg_end": self.cfg_end, "steps": self.steps, "batch_size": self.batch_size, "detailer": self.detailer, @@ -136,11 +138,11 @@ def process_images(p: StableDiffusionProcessing) -> Processed: processed = None try: # if no checkpoint override or the override checkpoint can't be found, remove override entry and load opts checkpoint - if p.override_settings.get('sd_model_checkpoint', None) is not None and sd_models.checkpoint_aliases.get(p.override_settings.get('sd_model_checkpoint')) is None: + if p.override_settings.get('sd_model_checkpoint', None) is not None and sd_checkpoint.checkpoint_aliases.get(p.override_settings.get('sd_model_checkpoint')) is None: shared.log.warning(f"Override not found: checkpoint={p.override_settings.get('sd_model_checkpoint', None)}") p.override_settings.pop('sd_model_checkpoint', None) sd_models.reload_model_weights() - if p.override_settings.get('sd_model_refiner', None) is not None and sd_models.checkpoint_aliases.get(p.override_settings.get('sd_model_refiner')) is None: + if p.override_settings.get('sd_model_refiner', None) is not None and sd_checkpoint.checkpoint_aliases.get(p.override_settings.get('sd_model_refiner')) is None: shared.log.warning(f"Override not found: refiner={p.override_settings.get('sd_model_refiner', None)}") p.override_settings.pop('sd_model_refiner', None) sd_models.reload_model_weights() @@ -162,7 +164,7 @@ def process_images(p: StableDiffusionProcessing) -> Processed: shared.prompt_styles.apply_styles_to_extra(p) shared.prompt_styles.extract_comments(p) if shared.opts.cuda_compile_backend == 'none': - sd_models.apply_token_merging(p.sd_model) + token_merge.apply_token_merging(p.sd_model) sd_hijack_freeu.apply_freeu(p, not shared.native) if p.width is not None: @@ -205,7 +207,7 @@ def process_images(p: StableDiffusionProcessing) -> Processed: finally: pag.unapply() if shared.opts.cuda_compile_backend == 'none': - sd_models.remove_token_merging(p.sd_model) + token_merge.remove_token_merging(p.sd_model) script_callbacks.after_process_callback(p) diff --git a/modules/processing_args.py b/modules/processing_args.py index 5cdf290a7..601615238 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -100,10 +100,11 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 if hasattr(model, "set_progress_bar_config"): model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m' + desc, ncols=80, colour='#327fba') args = {} - if hasattr(model, 'pipe'): # recurse + if hasattr(model, 'pipe') and not hasattr(model, 'no_recurse'): # recurse model = model.pipe signature = inspect.signature(type(model).__call__, follow_wrapped=True) possible = list(signature.parameters) + debug(f'Diffusers pipeline possible: {possible}') prompts, negative_prompts, prompts_2, negative_prompts_2 = fix_prompts(prompts, negative_prompts, prompts_2, negative_prompts_2) parser = 'Fixed attention' @@ -128,7 +129,7 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 if 'prompt' in possible: if 'OmniGen' in model.__class__.__name__: prompts = [p.replace('|image|', '<|image_1|>') for p in prompts] - if hasattr(model, 'text_encoder') and 'prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None: + if hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'prompt_embeds' in possible and len(p.prompt_embeds) > 0 and p.prompt_embeds[0] is not None: args['prompt_embeds'] = prompt_parser_diffusers.embedder('prompt_embeds') if 'StableCascade' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0: args['prompt_embeds_pooled'] = prompt_parser_diffusers.embedder('positive_pooleds').unsqueeze(0) @@ -141,7 +142,7 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 else: args['prompt'] = prompts if 'negative_prompt' in possible: - if hasattr(model, 'text_encoder') and 'negative_prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None: + if hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'negative_prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None: args['negative_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_prompt_embeds') if 'StableCascade' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None: args['negative_prompt_embeds_pooled'] = prompt_parser_diffusers.embedder('negative_pooleds').unsqueeze(0) diff --git a/modules/processing_callbacks.py b/modules/processing_callbacks.py index 3ace64ed8..59887c5c4 100644 --- a/modules/processing_callbacks.py +++ b/modules/processing_callbacks.py @@ -32,7 +32,7 @@ def diffusers_callback_legacy(step: int, timestep: int, latents: typing.Union[to time.sleep(0.1) -def diffusers_callback(pipe, step: int, timestep: int, kwargs: dict): +def diffusers_callback(pipe, step: int = 0, timestep: int = 0, kwargs: dict = {}): t0 = time.time() if p is None: return kwargs diff --git a/modules/processing_correction.py b/modules/processing_correction.py index c52f30ab3..e715d8c49 100644 --- a/modules/processing_correction.py +++ b/modules/processing_correction.py @@ -7,7 +7,8 @@ import os import torch from modules import shared, sd_vae_taesd, devices -debug = shared.log.trace if os.environ.get('SD_HDR_DEBUG', None) is not None else lambda *args, **kwargs: None +debug_enabled = os.environ.get('SD_HDR_DEBUG', None) is not None +debug = shared.log.trace if debug_enabled else lambda *args, **kwargs: None debug('Trace: HDR') @@ -119,16 +120,18 @@ def correction_callback(p, timestep, kwargs): if not any([p.hdr_clamp, p.hdr_mode, p.hdr_maximize, p.hdr_sharpen, p.hdr_color, p.hdr_brightness, p.hdr_tint_ratio]): return kwargs latents = kwargs["latents"] - debug('') - debug(f' Timestep: {timestep}') + if debug_enabled: + debug('') + debug(f' Timestep: {timestep}') # debug(f'HDR correction: latents={latents.shape}') if len(latents.shape) == 4: # standard batched latent for i in range(latents.shape[0]): latents[i] = correction(p, timestep, latents[i]) - debug(f"Full Mean: {latents[i].mean().item()}") - debug(f"Channel Means: {latents[i].mean(dim=(-1, -2), keepdim=True).flatten().float().cpu().numpy()}") - debug(f"Channel Mins: {latents[i].min(-1, keepdim=True)[0].min(-2, keepdim=True)[0].flatten().float().cpu().numpy()}") - debug(f"Channel Maxes: {latents[i].max(-1, keepdim=True)[0].min(-2, keepdim=True)[0].flatten().float().cpu().numpy()}") + if debug_enabled: + debug(f"Full Mean: {latents[i].mean().item()}") + debug(f"Channel Means: {latents[i].mean(dim=(-1, -2), keepdim=True).flatten().float().cpu().numpy()}") + debug(f"Channel Mins: {latents[i].min(-1, keepdim=True)[0].min(-2, keepdim=True)[0].flatten().float().cpu().numpy()}") + debug(f"Channel Maxes: {latents[i].max(-1, keepdim=True)[0].min(-2, keepdim=True)[0].flatten().float().cpu().numpy()}") elif len(latents.shape) == 5 and latents.shape[0] == 1: # probably animatediff latents = latents.squeeze(0).permute(1, 0, 2, 3) for i in range(latents.shape[0]): diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 12fa4bc53..7ec0dd08a 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -101,7 +101,8 @@ def process_base(p: processing.StableDiffusionProcessing): output = SimpleNamespace(**output) if isinstance(output, list): output = SimpleNamespace(images=output) - shared.history.add(output.images, info=processing.create_infotext(p), ops=p.ops) + if hasattr(output, 'images'): + shared.history.add(output.images, info=processing.create_infotext(p), ops=p.ops) timer.process.record('pipeline') hidiffusion.unapply() sd_models_compile.openvino_post_compile(op="base") # only executes on compiled vino models diff --git a/modules/processing_info.py b/modules/processing_info.py index 29513167d..e798211b1 100644 --- a/modules/processing_info.py +++ b/modules/processing_info.py @@ -41,11 +41,12 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No # basic "Steps": p.steps, "Seed": all_seeds[index], - "Sampler": p.sampler_name, - "CFG scale": p.cfg_scale, + "Sampler": p.sampler_name if p.sampler_name != 'Default' else None, + "CFG scale": p.cfg_scale if p.cfg_scale > 1.0 else None, + "CFG end": p.cfg_end if p.cfg_end < 1.0 else None, "Size": f"{p.width}x{p.height}" if hasattr(p, 'width') and hasattr(p, 'height') else None, "Batch": f'{p.n_iter}x{p.batch_size}' if p.n_iter > 1 or p.batch_size > 1 else None, - "Parser": shared.opts.prompt_attention, + "Parser": shared.opts.prompt_attention.split()[0], "Model": None if (not shared.opts.add_model_name_to_info) or (not shared.sd_model.sd_checkpoint_info.model_name) else shared.sd_model.sd_checkpoint_info.model_name.replace(',', '').replace(':', ''), "Model hash": getattr(p, 'sd_model_hash', None if (not shared.opts.add_model_hash_to_info) or (not shared.sd_model.sd_model_hash) else shared.sd_model.sd_model_hash), "VAE": (None if not shared.opts.add_model_name_to_info or sd_vae.loaded_vae_file is None else os.path.splitext(os.path.basename(sd_vae.loaded_vae_file))[0]) if p.full_quality else 'TAESD', diff --git a/modules/processing_original.py b/modules/processing_original.py index 852eb9a37..649023aae 100644 --- a/modules/processing_original.py +++ b/modules/processing_original.py @@ -1,7 +1,7 @@ import torch import numpy as np from PIL import Image -from modules import shared, devices, processing, images, sd_models, sd_vae, sd_samplers, processing_helpers, prompt_parser +from modules import shared, devices, processing, images, sd_vae, sd_samplers, processing_helpers, prompt_parser, token_merge from modules.sd_hijack_hypertile import hypertile_set @@ -135,10 +135,10 @@ def sample_txt2img(p: processing.StableDiffusionProcessingTxt2Img, conditioning, p.sampler.initialize(p) samples = samples[:, :, p.truncate_y//2:samples.shape[2]-(p.truncate_y+1)//2, p.truncate_x//2:samples.shape[3]-(p.truncate_x+1)//2] noise = create_random_tensors(samples.shape[1:], seeds=seeds, subseeds=subseeds, subseed_strength=subseed_strength, p=p) - sd_models.apply_token_merging(p.sd_model) + token_merge.apply_token_merging(p.sd_model) hypertile_set(p, hr=True) samples = p.sampler.sample_img2img(p, samples, noise, conditioning, unconditional_conditioning, steps=p.hr_second_pass_steps or p.steps, image_conditioning=image_conditioning) - sd_models.apply_token_merging(p.sd_model) + token_merge.apply_token_merging(p.sd_model) else: p.ops.append('upscale') x = None diff --git a/modules/pulid/__init__.py b/modules/pulid/__init__.py new file mode 100644 index 000000000..000f45293 --- /dev/null +++ b/modules/pulid/__init__.py @@ -0,0 +1,10 @@ +""" +Credit and original implementation: +""" + +import os +import sys +sys.path.append(os.path.dirname(__file__)) +from pulid_sdxl import StableDiffusionXLPuLIDPipeline +from pulid_utils import resize_numpy_image_long as resize +import attention_processor as attention diff --git a/modules/pulid/attention_processor.py b/modules/pulid/attention_processor.py new file mode 100644 index 000000000..9756decc1 --- /dev/null +++ b/modules/pulid/attention_processor.py @@ -0,0 +1,423 @@ +# modified from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py +import torch +import torch.nn as nn +import torch.nn.functional as F + + +NUM_ZERO = 0 +ORTHO = False +ORTHO_v2 = False + + +class AttnProcessor(nn.Module): + def __init__(self): + super().__init__() + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + temb=None, + id_embedding=None, + id_scale=1.0, + ): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + query = attn.head_to_batch_dim(query) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + + attention_probs = attn.get_attention_scores(query, key, attention_mask) + hidden_states = torch.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class IDAttnProcessor(nn.Module): + r""" + Attention processor for ID-Adapater. + Args: + hidden_size (`int`): + The hidden size of the attention layer. + cross_attention_dim (`int`): + The number of channels in the `encoder_hidden_states`. + scale (`float`, defaults to 1.0): + the weight scale of image prompt. + """ + + def __init__(self, hidden_size, cross_attention_dim=None): + super().__init__() + self.id_to_k = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + self.id_to_v = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + temb=None, + id_embedding=None, + id_scale=1.0, + ): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + query = attn.head_to_batch_dim(query) + key = attn.head_to_batch_dim(key) + value = attn.head_to_batch_dim(value) + + attention_probs = attn.get_attention_scores(query, key, attention_mask) + hidden_states = torch.bmm(attention_probs, value) + hidden_states = attn.batch_to_head_dim(hidden_states) + + # for id-adapter + if id_embedding is not None: + if NUM_ZERO == 0: + id_key = self.id_to_k(id_embedding) + id_value = self.id_to_v(id_embedding) + else: + zero_tensor = torch.zeros( + (id_embedding.size(0), NUM_ZERO, id_embedding.size(-1)), + dtype=id_embedding.dtype, + device=id_embedding.device, + ) + id_key = self.id_to_k(torch.cat((id_embedding, zero_tensor), dim=1)) + id_value = self.id_to_v(torch.cat((id_embedding, zero_tensor), dim=1)) + + id_key = attn.head_to_batch_dim(id_key).to(query.dtype) + id_value = attn.head_to_batch_dim(id_value).to(query.dtype) + + id_attention_probs = attn.get_attention_scores(query, id_key, None) + id_hidden_states = torch.bmm(id_attention_probs, id_value) + id_hidden_states = attn.batch_to_head_dim(id_hidden_states) + + if not ORTHO: + hidden_states = hidden_states + id_scale * id_hidden_states + else: + projection = ( + torch.sum((hidden_states * id_hidden_states), dim=-2, keepdim=True) + / torch.sum((hidden_states * hidden_states), dim=-2, keepdim=True) + * hidden_states + ) + orthogonal = id_hidden_states - projection + hidden_states = hidden_states + id_scale * orthogonal + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class AttnProcessor2_0(nn.Module): + r""" + Processor for implementing scaled dot-product attention (enabled by default if you're using PyTorch 2.0). + """ + + def __init__(self): + super().__init__() + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + temb=None, + id_embedding=None, + id_scale=1.0, + ): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + hidden_states = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states.to(query.dtype) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states + + +class IDAttnProcessor2_0(torch.nn.Module): + r""" + Attention processor for ID-Adapater for PyTorch 2.0. + Args: + hidden_size (`int`): + The hidden size of the attention layer. + cross_attention_dim (`int`): + The number of channels in the `encoder_hidden_states`. + """ + + def __init__(self, hidden_size, cross_attention_dim=None): + super().__init__() + if not hasattr(F, "scaled_dot_product_attention"): + raise ImportError("AttnProcessor2_0 requires PyTorch 2.0, to use it, please upgrade PyTorch to 2.0.") + + self.id_to_k = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + self.id_to_v = nn.Linear(cross_attention_dim or hidden_size, hidden_size, bias=False) + + def __call__( + self, + attn, + hidden_states, + encoder_hidden_states=None, + attention_mask=None, + temb=None, + id_embedding=None, + id_scale=1.0, + ): + residual = hidden_states + + if attn.spatial_norm is not None: + hidden_states = attn.spatial_norm(hidden_states, temb) + + input_ndim = hidden_states.ndim + + if input_ndim == 4: + batch_size, channel, height, width = hidden_states.shape + hidden_states = hidden_states.view(batch_size, channel, height * width).transpose(1, 2) + + batch_size, sequence_length, _ = ( + hidden_states.shape if encoder_hidden_states is None else encoder_hidden_states.shape + ) + + if attention_mask is not None: + attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size) + # scaled_dot_product_attention expects attention_mask shape to be + # (batch, heads, source_length, target_length) + attention_mask = attention_mask.view(batch_size, attn.heads, -1, attention_mask.shape[-1]) + + if attn.group_norm is not None: + hidden_states = attn.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2) + + query = attn.to_q(hidden_states) + + if encoder_hidden_states is None: + encoder_hidden_states = hidden_states + elif attn.norm_cross: + encoder_hidden_states = attn.norm_encoder_hidden_states(encoder_hidden_states) + + key = attn.to_k(encoder_hidden_states) + value = attn.to_v(encoder_hidden_states) + + inner_dim = key.shape[-1] + head_dim = inner_dim // attn.heads + + query = query.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + key = key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + value = value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + hidden_states = F.scaled_dot_product_attention( + query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False + ) + + hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + hidden_states = hidden_states.to(query.dtype) + + # for id embedding + if id_embedding is not None: + if NUM_ZERO == 0: + id_key = self.id_to_k(id_embedding).to(query.dtype) + id_value = self.id_to_v(id_embedding).to(query.dtype) + else: + zero_tensor = torch.zeros( + (id_embedding.size(0), NUM_ZERO, id_embedding.size(-1)), + dtype=id_embedding.dtype, + device=id_embedding.device, + ) + id_key = self.id_to_k(torch.cat((id_embedding, zero_tensor), dim=1)).to(query.dtype) + id_value = self.id_to_v(torch.cat((id_embedding, zero_tensor), dim=1)).to(query.dtype) + + id_key = id_key.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + id_value = id_value.view(batch_size, -1, attn.heads, head_dim).transpose(1, 2) + + # the output of sdp = (batch, num_heads, seq_len, head_dim) + id_hidden_states = F.scaled_dot_product_attention( + query, id_key, id_value, attn_mask=None, dropout_p=0.0, is_causal=False + ) + + id_hidden_states = id_hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim) + id_hidden_states = id_hidden_states.to(query.dtype) + + if not ORTHO and not ORTHO_v2: + hidden_states = hidden_states + id_scale * id_hidden_states + elif ORTHO_v2: + orig_dtype = hidden_states.dtype + hidden_states = hidden_states.to(torch.float32) + id_hidden_states = id_hidden_states.to(torch.float32) + attn_map = query @ id_key.transpose(-2, -1) + attn_mean = attn_map.softmax(dim=-1).mean(dim=1) + attn_mean = attn_mean[:, :, :5].sum(dim=-1, keepdim=True) + projection = ( + torch.sum((hidden_states * id_hidden_states), dim=-2, keepdim=True) + / torch.sum((hidden_states * hidden_states), dim=-2, keepdim=True) + * hidden_states + ) + orthogonal = id_hidden_states + (attn_mean - 1) * projection + hidden_states = hidden_states + id_scale * orthogonal + hidden_states = hidden_states.to(orig_dtype) + else: + orig_dtype = hidden_states.dtype + hidden_states = hidden_states.to(torch.float32) + id_hidden_states = id_hidden_states.to(torch.float32) + projection = ( + torch.sum((hidden_states * id_hidden_states), dim=-2, keepdim=True) + / torch.sum((hidden_states * hidden_states), dim=-2, keepdim=True) + * hidden_states + ) + orthogonal = id_hidden_states - projection + hidden_states = hidden_states + id_scale * orthogonal + hidden_states = hidden_states.to(orig_dtype) + + # linear proj + hidden_states = attn.to_out[0](hidden_states) + # dropout + hidden_states = attn.to_out[1](hidden_states) + + if input_ndim == 4: + hidden_states = hidden_states.transpose(-1, -2).reshape(batch_size, channel, height, width) + + if attn.residual_connection: + hidden_states = hidden_states + residual + + hidden_states = hidden_states / attn.rescale_output_factor + + return hidden_states diff --git a/modules/pulid/encoders_transformer.py b/modules/pulid/encoders_transformer.py new file mode 100644 index 000000000..d1ecef2c6 --- /dev/null +++ b/modules/pulid/encoders_transformer.py @@ -0,0 +1,208 @@ +import math +import torch +import torch.nn as nn + + +# FFN +def FeedForward(dim, mult=4): + inner_dim = int(dim * mult) + return nn.Sequential( + nn.LayerNorm(dim), + nn.Linear(dim, inner_dim, bias=False), + nn.GELU(), + nn.Linear(inner_dim, dim, bias=False), + ) + + +def reshape_tensor(x, heads): + bs, length, _width = x.shape + # (bs, length, width) --> (bs, length, n_heads, dim_per_head) + x = x.view(bs, length, heads, -1) + # (bs, length, n_heads, dim_per_head) --> (bs, n_heads, length, dim_per_head) + x = x.transpose(1, 2) + # (bs, n_heads, length, dim_per_head) --> (bs*n_heads, length, dim_per_head) + x = x.reshape(bs, heads, length, -1) + return x + + +class PerceiverAttentionCA(nn.Module): + def __init__(self, *, dim=3072, dim_head=128, heads=16, kv_dim=2048): + super().__init__() + self.scale = dim_head ** -0.5 + self.dim_head = dim_head + self.heads = heads + inner_dim = dim_head * heads + + self.norm1 = nn.LayerNorm(dim if kv_dim is None else kv_dim) + self.norm2 = nn.LayerNorm(dim) + + self.to_q = nn.Linear(dim, inner_dim, bias=False) + self.to_kv = nn.Linear(dim if kv_dim is None else kv_dim, inner_dim * 2, bias=False) + self.to_out = nn.Linear(inner_dim, dim, bias=False) + + def forward(self, x, latents): + """ + Args: + x (torch.Tensor): image features + shape (b, n1, D) + latent (torch.Tensor): latent features + shape (b, n2, D) + """ + x = self.norm1(x) + latents = self.norm2(latents) + + b, seq_len, _ = latents.shape + + q = self.to_q(latents) + k, v = self.to_kv(x).chunk(2, dim=-1) + + q = reshape_tensor(q, self.heads) + k = reshape_tensor(k, self.heads) + v = reshape_tensor(v, self.heads) + + # attention + scale = 1 / math.sqrt(math.sqrt(self.dim_head)) + weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards + weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) + out = weight @ v + + out = out.permute(0, 2, 1, 3).reshape(b, seq_len, -1) + + return self.to_out(out) + + +class PerceiverAttention(nn.Module): + def __init__(self, *, dim, dim_head=64, heads=8, kv_dim=None): + super().__init__() + self.scale = dim_head ** -0.5 + self.dim_head = dim_head + self.heads = heads + inner_dim = dim_head * heads + + self.norm1 = nn.LayerNorm(dim if kv_dim is None else kv_dim) + self.norm2 = nn.LayerNorm(dim) + + self.to_q = nn.Linear(dim, inner_dim, bias=False) + self.to_kv = nn.Linear(dim if kv_dim is None else kv_dim, inner_dim * 2, bias=False) + self.to_out = nn.Linear(inner_dim, dim, bias=False) + + def forward(self, x, latents): + """ + Args: + x (torch.Tensor): image features + shape (b, n1, D) + latent (torch.Tensor): latent features + shape (b, n2, D) + """ + x = self.norm1(x) + latents = self.norm2(latents) + + b, seq_len, _ = latents.shape + + q = self.to_q(latents) + kv_input = torch.cat((x, latents), dim=-2) + k, v = self.to_kv(kv_input).chunk(2, dim=-1) + + q = reshape_tensor(q, self.heads) + k = reshape_tensor(k, self.heads) + v = reshape_tensor(v, self.heads) + + # attention + scale = 1 / math.sqrt(math.sqrt(self.dim_head)) + weight = (q * scale) @ (k * scale).transpose(-2, -1) # More stable with f16 than dividing afterwards + weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) + out = weight @ v + + out = out.permute(0, 2, 1, 3).reshape(b, seq_len, -1) + + return self.to_out(out) + + +class IDFormer(nn.Module): + """ + - perceiver resampler like arch (compared with previous MLP-like arch) + - we concat id embedding (generated by arcface) and query tokens as latents + - latents will attend each other and interact with vit features through cross-attention + - vit features are multi-scaled and inserted into IDFormer in order, currently, each scale corresponds to two + IDFormer layers + """ + def __init__( + self, + dim=1024, + depth=10, + dim_head=64, + heads=16, + num_id_token=5, + num_queries=32, + output_dim=2048, + ff_mult=4, + ): + super().__init__() + + self.num_id_token = num_id_token + self.dim = dim + self.num_queries = num_queries + assert depth % 5 == 0 + self.depth = depth // 5 + scale = dim ** -0.5 + + self.latents = nn.Parameter(torch.randn(1, num_queries, dim) * scale) + self.proj_out = nn.Parameter(scale * torch.randn(dim, output_dim)) + + self.layers = nn.ModuleList([]) + for _ in range(depth): + self.layers.append( + nn.ModuleList( + [ + PerceiverAttention(dim=dim, dim_head=dim_head, heads=heads), + FeedForward(dim=dim, mult=ff_mult), + ] + ) + ) + + for i in range(5): + setattr( + self, + f'mapping_{i}', + nn.Sequential( + nn.Linear(1024, 1024), + nn.LayerNorm(1024), + nn.LeakyReLU(), + nn.Linear(1024, 1024), + nn.LayerNorm(1024), + nn.LeakyReLU(), + nn.Linear(1024, dim), + ), + ) + + self.id_embedding_mapping = nn.Sequential( + nn.Linear(1280, 1024), + nn.LayerNorm(1024), + nn.LeakyReLU(), + nn.Linear(1024, 1024), + nn.LayerNorm(1024), + nn.LeakyReLU(), + nn.Linear(1024, dim * num_id_token), + ) + + def forward(self, x, y): + + latents = self.latents.repeat(x.size(0), 1, 1) + + num_duotu = x.shape[1] if x.ndim == 3 else 1 + + x = self.id_embedding_mapping(x) + x = x.reshape(-1, self.num_id_token * num_duotu, self.dim) + + latents = torch.cat((latents, x), dim=1) + + for i in range(5): + vit_feature = getattr(self, f'mapping_{i}')(y[i]) + ctx_feature = torch.cat((x, vit_feature), dim=1) + for attn, ff in self.layers[i * self.depth: (i + 1) * self.depth]: + latents = attn(ctx_feature, latents) + latents + latents = ff(latents) + latents + + latents = latents[:, :self.num_queries] + latents = latents @ self.proj_out + return latents diff --git a/modules/pulid/eva_clip/__init__.py b/modules/pulid/eva_clip/__init__.py new file mode 100644 index 000000000..fa2d014bb --- /dev/null +++ b/modules/pulid/eva_clip/__init__.py @@ -0,0 +1,11 @@ +from .constants import OPENAI_DATASET_MEAN, OPENAI_DATASET_STD +from .factory import create_model, create_model_and_transforms, create_model_from_pretrained, get_tokenizer, create_transforms +from .factory import list_models, add_model_config, get_model_config, load_checkpoint +from .loss import ClipLoss +from .model import CLIP, CustomCLIP, CLIPTextCfg, CLIPVisionCfg,\ + convert_weights_to_lp, convert_weights_to_fp16, trace_model, get_cast_dtype +from .openai import load_openai_model, list_openai_models +from .pretrained import list_pretrained, list_pretrained_models_by_tag, list_pretrained_tags_by_model,\ + get_pretrained_url, download_pretrained_from_url, is_pretrained_cfg, get_pretrained_cfg, download_pretrained +from .tokenizer import SimpleTokenizer, tokenize +from .transform import image_transform \ No newline at end of file diff --git a/modules/pulid/eva_clip/bpe_simple_vocab_16e6.txt.gz b/modules/pulid/eva_clip/bpe_simple_vocab_16e6.txt.gz new file mode 100644 index 000000000..7b5088a52 Binary files /dev/null and b/modules/pulid/eva_clip/bpe_simple_vocab_16e6.txt.gz differ diff --git a/modules/pulid/eva_clip/constants.py b/modules/pulid/eva_clip/constants.py new file mode 100644 index 000000000..a670bb3fa --- /dev/null +++ b/modules/pulid/eva_clip/constants.py @@ -0,0 +1,2 @@ +OPENAI_DATASET_MEAN = (0.48145466, 0.4578275, 0.40821073) +OPENAI_DATASET_STD = (0.26862954, 0.26130258, 0.27577711) diff --git a/modules/pulid/eva_clip/eva_vit_model.py b/modules/pulid/eva_clip/eva_vit_model.py new file mode 100644 index 000000000..51db88cf0 --- /dev/null +++ b/modules/pulid/eva_clip/eva_vit_model.py @@ -0,0 +1,548 @@ +# -------------------------------------------------------- +# Adapted from https://github.com/microsoft/unilm/tree/master/beit +# -------------------------------------------------------- +import math +import os +from functools import partial +import torch +import torch.nn as nn +import torch.nn.functional as F +try: + from timm.models.layers import drop_path, to_2tuple, trunc_normal_ +except: + from timm.layers import drop_path, to_2tuple, trunc_normal_ + +from .transformer import PatchDropout +from .rope import VisionRotaryEmbedding, VisionRotaryEmbeddingFast + +if os.getenv('ENV_TYPE') == 'deepspeed': + try: + from deepspeed.runtime.activation_checkpointing.checkpointing import checkpoint + except: + from torch.utils.checkpoint import checkpoint +else: + from torch.utils.checkpoint import checkpoint + +try: + import xformers + import xformers.ops as xops + XFORMERS_IS_AVAILBLE = True +except: + XFORMERS_IS_AVAILBLE = False + +class DropPath(nn.Module): + """Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks). + """ + def __init__(self, drop_prob=None): + super(DropPath, self).__init__() + self.drop_prob = drop_prob + + def forward(self, x): + return drop_path(x, self.drop_prob, self.training) + + def extra_repr(self) -> str: + return 'p={}'.format(self.drop_prob) + + +class Mlp(nn.Module): + def __init__( + self, + in_features, + hidden_features=None, + out_features=None, + act_layer=nn.GELU, + norm_layer=nn.LayerNorm, + drop=0., + subln=False, + + ): + super().__init__() + out_features = out_features or in_features + hidden_features = hidden_features or in_features + self.fc1 = nn.Linear(in_features, hidden_features) + self.act = act_layer() + + self.ffn_ln = norm_layer(hidden_features) if subln else nn.Identity() + + self.fc2 = nn.Linear(hidden_features, out_features) + self.drop = nn.Dropout(drop) + + def forward(self, x): + x = self.fc1(x) + x = self.act(x) + # x = self.drop(x) + # commit this for the orignal BERT implement + x = self.ffn_ln(x) + + x = self.fc2(x) + x = self.drop(x) + return x + +class SwiGLU(nn.Module): + def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.SiLU, drop=0., + norm_layer=nn.LayerNorm, subln=False): + super().__init__() + out_features = out_features or in_features + hidden_features = hidden_features or in_features + + self.w1 = nn.Linear(in_features, hidden_features) + self.w2 = nn.Linear(in_features, hidden_features) + + self.act = act_layer() + self.ffn_ln = norm_layer(hidden_features) if subln else nn.Identity() + self.w3 = nn.Linear(hidden_features, out_features) + + self.drop = nn.Dropout(drop) + + def forward(self, x): + x1 = self.w1(x) + x2 = self.w2(x) + hidden = self.act(x1) * x2 + x = self.ffn_ln(hidden) + x = self.w3(x) + x = self.drop(x) + return x + +class Attention(nn.Module): + def __init__( + self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., + proj_drop=0., window_size=None, attn_head_dim=None, xattn=False, rope=None, subln=False, norm_layer=nn.LayerNorm): + super().__init__() + self.num_heads = num_heads + head_dim = dim // num_heads + if attn_head_dim is not None: + head_dim = attn_head_dim + all_head_dim = head_dim * self.num_heads + self.scale = qk_scale or head_dim ** -0.5 + + self.subln = subln + if self.subln: + self.q_proj = nn.Linear(dim, all_head_dim, bias=False) + self.k_proj = nn.Linear(dim, all_head_dim, bias=False) + self.v_proj = nn.Linear(dim, all_head_dim, bias=False) + else: + self.qkv = nn.Linear(dim, all_head_dim * 3, bias=False) + + if qkv_bias: + self.q_bias = nn.Parameter(torch.zeros(all_head_dim)) + self.v_bias = nn.Parameter(torch.zeros(all_head_dim)) + else: + self.q_bias = None + self.v_bias = None + + if window_size: + self.window_size = window_size + self.num_relative_distance = (2 * window_size[0] - 1) * (2 * window_size[1] - 1) + 3 + self.relative_position_bias_table = nn.Parameter( + torch.zeros(self.num_relative_distance, num_heads)) # 2*Wh-1 * 2*Ww-1, nH + # cls to token & token 2 cls & cls to cls + + # get pair-wise relative position index for each token inside the window + coords_h = torch.arange(window_size[0]) + coords_w = torch.arange(window_size[1]) + coords = torch.stack(torch.meshgrid([coords_h, coords_w])) # 2, Wh, Ww + coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww + relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww + relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2 + relative_coords[:, :, 0] += window_size[0] - 1 # shift to start from 0 + relative_coords[:, :, 1] += window_size[1] - 1 + relative_coords[:, :, 0] *= 2 * window_size[1] - 1 + relative_position_index = \ + torch.zeros(size=(window_size[0] * window_size[1] + 1, ) * 2, dtype=relative_coords.dtype) + relative_position_index[1:, 1:] = relative_coords.sum(-1) # Wh*Ww, Wh*Ww + relative_position_index[0, 0:] = self.num_relative_distance - 3 + relative_position_index[0:, 0] = self.num_relative_distance - 2 + relative_position_index[0, 0] = self.num_relative_distance - 1 + + self.register_buffer("relative_position_index", relative_position_index) + else: + self.window_size = None + self.relative_position_bias_table = None + self.relative_position_index = None + + self.attn_drop = nn.Dropout(attn_drop) + self.inner_attn_ln = norm_layer(all_head_dim) if subln else nn.Identity() + # self.proj = nn.Linear(all_head_dim, all_head_dim) + self.proj = nn.Linear(all_head_dim, dim) + self.proj_drop = nn.Dropout(proj_drop) + self.xattn = xattn + self.xattn_drop = attn_drop + + self.rope = rope + + def forward(self, x, rel_pos_bias=None, attn_mask=None): + B, N, C = x.shape + if self.subln: + q = F.linear(input=x, weight=self.q_proj.weight, bias=self.q_bias) + k = F.linear(input=x, weight=self.k_proj.weight, bias=None) + v = F.linear(input=x, weight=self.v_proj.weight, bias=self.v_bias) + + q = q.reshape(B, N, self.num_heads, -1).permute(0, 2, 1, 3) # B, num_heads, N, C + k = k.reshape(B, N, self.num_heads, -1).permute(0, 2, 1, 3) + v = v.reshape(B, N, self.num_heads, -1).permute(0, 2, 1, 3) + else: + + qkv_bias = None + if self.q_bias is not None: + qkv_bias = torch.cat((self.q_bias, torch.zeros_like(self.v_bias, requires_grad=False), self.v_bias)) + + qkv = F.linear(input=x, weight=self.qkv.weight, bias=qkv_bias) + qkv = qkv.reshape(B, N, 3, self.num_heads, -1).permute(2, 0, 3, 1, 4) # 3, B, num_heads, N, C + q, k, v = qkv[0], qkv[1], qkv[2] + + if self.rope: + # slightly fast impl + q_t = q[:, :, 1:, :] + ro_q_t = self.rope(q_t) + q = torch.cat((q[:, :, :1, :], ro_q_t), -2).type_as(v) + + k_t = k[:, :, 1:, :] + ro_k_t = self.rope(k_t) + k = torch.cat((k[:, :, :1, :], ro_k_t), -2).type_as(v) + + if self.xattn: + q = q.permute(0, 2, 1, 3) # B, num_heads, N, C -> B, N, num_heads, C + k = k.permute(0, 2, 1, 3) + v = v.permute(0, 2, 1, 3) + + x = xops.memory_efficient_attention( + q, k, v, + p=self.xattn_drop, + scale=self.scale, + ) + x = x.reshape(B, N, -1) + x = self.inner_attn_ln(x) + x = self.proj(x) + x = self.proj_drop(x) + else: + q = q * self.scale + attn = (q @ k.transpose(-2, -1)) + + if self.relative_position_bias_table is not None: + relative_position_bias = \ + self.relative_position_bias_table[self.relative_position_index.view(-1)].view( + self.window_size[0] * self.window_size[1] + 1, + self.window_size[0] * self.window_size[1] + 1, -1) # Wh*Ww,Wh*Ww,nH + relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() # nH, Wh*Ww, Wh*Ww + attn = attn + relative_position_bias.unsqueeze(0).type_as(attn) + + if rel_pos_bias is not None: + attn = attn + rel_pos_bias.type_as(attn) + + if attn_mask is not None: + attn_mask = attn_mask.bool() + attn = attn.masked_fill(~attn_mask[:, None, None, :], float("-inf")) + + attn = attn.softmax(dim=-1) + attn = self.attn_drop(attn) + + x = (attn @ v).transpose(1, 2).reshape(B, N, -1) + x = self.inner_attn_ln(x) + x = self.proj(x) + x = self.proj_drop(x) + return x + + +class Block(nn.Module): + + def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0., + drop_path=0., init_values=None, act_layer=nn.GELU, norm_layer=nn.LayerNorm, + window_size=None, attn_head_dim=None, xattn=False, rope=None, postnorm=False, + subln=False, naiveswiglu=False): + super().__init__() + self.norm1 = norm_layer(dim) + self.attn = Attention( + dim, num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, + attn_drop=attn_drop, proj_drop=drop, window_size=window_size, attn_head_dim=attn_head_dim, + xattn=xattn, rope=rope, subln=subln, norm_layer=norm_layer) + # NOTE: drop path for stochastic depth, we shall see if this is better than dropout here + self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity() + self.norm2 = norm_layer(dim) + mlp_hidden_dim = int(dim * mlp_ratio) + + if naiveswiglu: + self.mlp = SwiGLU( + in_features=dim, + hidden_features=mlp_hidden_dim, + subln=subln, + norm_layer=norm_layer, + ) + else: + self.mlp = Mlp( + in_features=dim, + hidden_features=mlp_hidden_dim, + act_layer=act_layer, + subln=subln, + drop=drop + ) + + if init_values is not None and init_values > 0: + self.gamma_1 = nn.Parameter(init_values * torch.ones((dim)),requires_grad=True) + self.gamma_2 = nn.Parameter(init_values * torch.ones((dim)),requires_grad=True) + else: + self.gamma_1, self.gamma_2 = None, None + + self.postnorm = postnorm + + def forward(self, x, rel_pos_bias=None, attn_mask=None): + if self.gamma_1 is None: + if self.postnorm: + x = x + self.drop_path(self.norm1(self.attn(x, rel_pos_bias=rel_pos_bias, attn_mask=attn_mask))) + x = x + self.drop_path(self.norm2(self.mlp(x))) + else: + x = x + self.drop_path(self.attn(self.norm1(x), rel_pos_bias=rel_pos_bias, attn_mask=attn_mask)) + x = x + self.drop_path(self.mlp(self.norm2(x))) + else: + if self.postnorm: + x = x + self.drop_path(self.gamma_1 * self.norm1(self.attn(x, rel_pos_bias=rel_pos_bias, attn_mask=attn_mask))) + x = x + self.drop_path(self.gamma_2 * self.norm2(self.mlp(x))) + else: + x = x + self.drop_path(self.gamma_1 * self.attn(self.norm1(x), rel_pos_bias=rel_pos_bias, attn_mask=attn_mask)) + x = x + self.drop_path(self.gamma_2 * self.mlp(self.norm2(x))) + return x + + +class PatchEmbed(nn.Module): + """ Image to Patch Embedding + """ + def __init__(self, img_size=224, patch_size=16, in_chans=3, embed_dim=768): + super().__init__() + img_size = to_2tuple(img_size) + patch_size = to_2tuple(patch_size) + num_patches = (img_size[1] // patch_size[1]) * (img_size[0] // patch_size[0]) + self.patch_shape = (img_size[0] // patch_size[0], img_size[1] // patch_size[1]) + self.img_size = img_size + self.patch_size = patch_size + self.num_patches = num_patches + + self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) + + def forward(self, x, **kwargs): + B, C, H, W = x.shape + # FIXME look at relaxing size constraints + assert H == self.img_size[0] and W == self.img_size[1], \ + f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})." + x = self.proj(x).flatten(2).transpose(1, 2) + return x + + +class RelativePositionBias(nn.Module): + + def __init__(self, window_size, num_heads): + super().__init__() + self.window_size = window_size + self.num_relative_distance = (2 * window_size[0] - 1) * (2 * window_size[1] - 1) + 3 + self.relative_position_bias_table = nn.Parameter( + torch.zeros(self.num_relative_distance, num_heads)) # 2*Wh-1 * 2*Ww-1, nH + # cls to token & token 2 cls & cls to cls + + # get pair-wise relative position index for each token inside the window + coords_h = torch.arange(window_size[0]) + coords_w = torch.arange(window_size[1]) + coords = torch.stack(torch.meshgrid([coords_h, coords_w])) # 2, Wh, Ww + coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww + relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww + relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2 + relative_coords[:, :, 0] += window_size[0] - 1 # shift to start from 0 + relative_coords[:, :, 1] += window_size[1] - 1 + relative_coords[:, :, 0] *= 2 * window_size[1] - 1 + relative_position_index = \ + torch.zeros(size=(window_size[0] * window_size[1] + 1,) * 2, dtype=relative_coords.dtype) + relative_position_index[1:, 1:] = relative_coords.sum(-1) # Wh*Ww, Wh*Ww + relative_position_index[0, 0:] = self.num_relative_distance - 3 + relative_position_index[0:, 0] = self.num_relative_distance - 2 + relative_position_index[0, 0] = self.num_relative_distance - 1 + + self.register_buffer("relative_position_index", relative_position_index) + + def forward(self): + relative_position_bias = \ + self.relative_position_bias_table[self.relative_position_index.view(-1)].view( + self.window_size[0] * self.window_size[1] + 1, + self.window_size[0] * self.window_size[1] + 1, -1) # Wh*Ww,Wh*Ww,nH + return relative_position_bias.permute(2, 0, 1).contiguous() # nH, Wh*Ww, Wh*Ww + + +class EVAVisionTransformer(nn.Module): + """ Vision Transformer with support for patch or hybrid CNN input stage + """ + def __init__(self, img_size=224, patch_size=16, in_chans=3, num_classes=1000, embed_dim=768, depth=12, + num_heads=12, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop_rate=0., attn_drop_rate=0., + drop_path_rate=0., norm_layer=nn.LayerNorm, init_values=None, patch_dropout=0., + use_abs_pos_emb=True, use_rel_pos_bias=False, use_shared_rel_pos_bias=False, rope=False, + use_mean_pooling=True, init_scale=0.001, grad_checkpointing=False, xattn=False, postnorm=False, + pt_hw_seq_len=16, intp_freq=False, naiveswiglu=False, subln=False): + super().__init__() + + if not XFORMERS_IS_AVAILBLE: + xattn = False + + self.image_size = img_size + self.num_classes = num_classes + self.num_features = self.embed_dim = embed_dim # num_features for consistency with other models + + self.patch_embed = PatchEmbed( + img_size=img_size, patch_size=patch_size, in_chans=in_chans, embed_dim=embed_dim) + num_patches = self.patch_embed.num_patches + + self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) + # self.mask_token = nn.Parameter(torch.zeros(1, 1, embed_dim)) + if use_abs_pos_emb: + self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim)) + else: + self.pos_embed = None + self.pos_drop = nn.Dropout(p=drop_rate) + + if use_shared_rel_pos_bias: + self.rel_pos_bias = RelativePositionBias(window_size=self.patch_embed.patch_shape, num_heads=num_heads) + else: + self.rel_pos_bias = None + + if rope: + half_head_dim = embed_dim // num_heads // 2 + hw_seq_len = img_size // patch_size + self.rope = VisionRotaryEmbeddingFast( + dim=half_head_dim, + pt_seq_len=pt_hw_seq_len, + ft_seq_len=hw_seq_len if intp_freq else None, + # patch_dropout=patch_dropout + ) + else: + self.rope = None + + self.naiveswiglu = naiveswiglu + + dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)] # stochastic depth decay rule + self.use_rel_pos_bias = use_rel_pos_bias + self.blocks = nn.ModuleList([ + Block( + dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, qk_scale=qk_scale, + drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[i], norm_layer=norm_layer, + init_values=init_values, window_size=self.patch_embed.patch_shape if use_rel_pos_bias else None, + xattn=xattn, rope=self.rope, postnorm=postnorm, subln=subln, naiveswiglu=naiveswiglu) + for i in range(depth)]) + self.norm = nn.Identity() if use_mean_pooling else norm_layer(embed_dim) + self.fc_norm = norm_layer(embed_dim) if use_mean_pooling else None + self.head = nn.Linear(embed_dim, num_classes) if num_classes > 0 else nn.Identity() + + if self.pos_embed is not None: + trunc_normal_(self.pos_embed, std=.02) + + trunc_normal_(self.cls_token, std=.02) + # trunc_normal_(self.mask_token, std=.02) + + self.apply(self._init_weights) + self.fix_init_weight() + + if isinstance(self.head, nn.Linear): + trunc_normal_(self.head.weight, std=.02) + self.head.weight.data.mul_(init_scale) + self.head.bias.data.mul_(init_scale) + + # setting a patch_dropout of 0. would mean it is disabled and this function would be the identity fn + self.patch_dropout = PatchDropout(patch_dropout) if patch_dropout > 0. else nn.Identity() + + self.grad_checkpointing = grad_checkpointing + + def fix_init_weight(self): + def rescale(param, layer_id): + param.div_(math.sqrt(2.0 * layer_id)) + + for layer_id, layer in enumerate(self.blocks): + rescale(layer.attn.proj.weight.data, layer_id + 1) + if self.naiveswiglu: + rescale(layer.mlp.w3.weight.data, layer_id + 1) + else: + rescale(layer.mlp.fc2.weight.data, layer_id + 1) + + def get_cast_dtype(self) -> torch.dtype: + return self.blocks[0].mlp.fc2.weight.dtype + + def _init_weights(self, m): + if isinstance(m, nn.Linear): + trunc_normal_(m.weight, std=.02) + if m.bias is not None: + nn.init.constant_(m.bias, 0) + elif isinstance(m, nn.LayerNorm): + nn.init.constant_(m.bias, 0) + nn.init.constant_(m.weight, 1.0) + + def get_num_layers(self): + return len(self.blocks) + + def lock(self, unlocked_groups=0, freeze_bn_stats=False): + assert unlocked_groups == 0, 'partial locking not currently supported for this model' + for param in self.parameters(): + param.requires_grad = False + + @torch.jit.ignore + def set_grad_checkpointing(self, enable=True): + self.grad_checkpointing = enable + + @torch.jit.ignore + def no_weight_decay(self): + return {'pos_embed', 'cls_token'} + + def get_classifier(self): + return self.head + + def reset_classifier(self, num_classes, global_pool=''): + self.num_classes = num_classes + self.head = nn.Linear(self.embed_dim, num_classes) if num_classes > 0 else nn.Identity() + + def forward_features(self, x, return_all_features=False, return_hidden=False, shuffle=False): + + x = self.patch_embed(x) + batch_size, seq_len, _ = x.size() + + if shuffle: + idx = torch.randperm(x.shape[1]) + 1 + zero = torch.LongTensor([0, ]) + idx = torch.cat([zero, idx]) + pos_embed = self.pos_embed[:, idx] + + cls_tokens = self.cls_token.expand(batch_size, -1, -1) # stole cls_tokens impl from Phil Wang, thanks + x = torch.cat((cls_tokens, x), dim=1) + if shuffle: + x = x + pos_embed + elif self.pos_embed is not None: + x = x + self.pos_embed + x = self.pos_drop(x) + + # a patch_dropout of 0. would mean it is disabled and this function would do nothing but return what was passed in + if os.getenv('RoPE') == '1': + if self.training and not isinstance(self.patch_dropout, nn.Identity): + x, patch_indices_keep = self.patch_dropout(x) + self.rope.forward = partial(self.rope.forward, patch_indices_keep=patch_indices_keep) + else: + self.rope.forward = partial(self.rope.forward, patch_indices_keep=None) + x = self.patch_dropout(x) + else: + x = self.patch_dropout(x) + + rel_pos_bias = self.rel_pos_bias() if self.rel_pos_bias is not None else None + hidden_states = [] + for idx, blk in enumerate(self.blocks): + if (0 < idx <= 20) and (idx % 4 == 0) and return_hidden: + hidden_states.append(x) + if self.grad_checkpointing: + x = checkpoint(blk, x, (rel_pos_bias,)) + else: + x = blk(x, rel_pos_bias=rel_pos_bias) + + if not return_all_features: + x = self.norm(x) + if self.fc_norm is not None: + return self.fc_norm(x.mean(1)), hidden_states + else: + return x[:, 0], hidden_states + return x + + def forward(self, x, return_all_features=False, return_hidden=False, shuffle=False): + if return_all_features: + return self.forward_features(x, return_all_features, return_hidden, shuffle) + x, hidden_states = self.forward_features(x, return_all_features, return_hidden, shuffle) + x = self.head(x) + if return_hidden: + return x, hidden_states + return x diff --git a/modules/pulid/eva_clip/factory.py b/modules/pulid/eva_clip/factory.py new file mode 100644 index 000000000..ced899999 --- /dev/null +++ b/modules/pulid/eva_clip/factory.py @@ -0,0 +1,517 @@ +import json +import logging +import os +import pathlib +import re +from copy import deepcopy +from pathlib import Path +from typing import Optional, Tuple, Union, Dict, Any +import torch + +from .constants import OPENAI_DATASET_MEAN, OPENAI_DATASET_STD +from .model import CLIP, CustomCLIP, convert_weights_to_lp, convert_to_custom_text_state_dict,\ + get_cast_dtype +from .openai import load_openai_model +from .pretrained import is_pretrained_cfg, get_pretrained_cfg, download_pretrained, list_pretrained_tags_by_model +from .transform import image_transform +from .tokenizer import HFTokenizer, tokenize +from .utils import resize_clip_pos_embed, resize_evaclip_pos_embed, resize_visual_pos_embed, resize_eva_pos_embed + + +_MODEL_CONFIG_PATHS = [Path(__file__).parent / f"model_configs/"] +_MODEL_CONFIGS = {} # directory (model_name: config) of model architecture configs + + +def _natural_key(string_): + return [int(s) if s.isdigit() else s for s in re.split(r'(\d+)', string_.lower())] + + +def _rescan_model_configs(): + global _MODEL_CONFIGS + + config_ext = ('.json',) + config_files = [] + for config_path in _MODEL_CONFIG_PATHS: + if config_path.is_file() and config_path.suffix in config_ext: + config_files.append(config_path) + elif config_path.is_dir(): + for ext in config_ext: + config_files.extend(config_path.glob(f'*{ext}')) + + for cf in config_files: + with open(cf, "r", encoding="utf8") as f: + model_cfg = json.load(f) + if all(a in model_cfg for a in ('embed_dim', 'vision_cfg', 'text_cfg')): + _MODEL_CONFIGS[cf.stem] = model_cfg + + _MODEL_CONFIGS = dict(sorted(_MODEL_CONFIGS.items(), key=lambda x: _natural_key(x[0]))) + + +_rescan_model_configs() # initial populate of model config registry + + +def list_models(): + """ enumerate available model architectures based on config files """ + return list(_MODEL_CONFIGS.keys()) + + +def add_model_config(path): + """ add model config path or file and update registry """ + if not isinstance(path, Path): + path = Path(path) + _MODEL_CONFIG_PATHS.append(path) + _rescan_model_configs() + + +def get_model_config(model_name): + if model_name in _MODEL_CONFIGS: + return deepcopy(_MODEL_CONFIGS[model_name]) + else: + return None + + +def get_tokenizer(model_name): + config = get_model_config(model_name) + tokenizer = HFTokenizer(config['text_cfg']['hf_tokenizer_name']) if 'hf_tokenizer_name' in config['text_cfg'] else tokenize + return tokenizer + + +# loading openai CLIP weights when is_openai=True for training +def load_state_dict(checkpoint_path: str, map_location: str='cpu', model_key: str='model|module|state_dict', is_openai: bool=False, skip_list: list=[]): + if is_openai: + model = torch.jit.load(checkpoint_path, map_location="cpu").eval() + state_dict = model.state_dict() + for key in ["input_resolution", "context_length", "vocab_size"]: + state_dict.pop(key, None) + else: + checkpoint = torch.load(checkpoint_path, map_location=map_location) + for mk in model_key.split('|'): + if isinstance(checkpoint, dict) and mk in checkpoint: + state_dict = checkpoint[mk] + break + else: + state_dict = checkpoint + if next(iter(state_dict.items()))[0].startswith('module'): + state_dict = {k[7:]: v for k, v in state_dict.items()} + + for k in skip_list: + if k in list(state_dict.keys()): + logging.info(f"Removing key {k} from pretrained checkpoint") + del state_dict[k] + + if os.getenv('RoPE') == '1': + for k in list(state_dict.keys()): + if 'freqs_cos' in k or 'freqs_sin' in k: + del state_dict[k] + return state_dict + + + +def load_checkpoint(model, checkpoint_path, model_key="model|module|state_dict", strict=True): + state_dict = load_state_dict(checkpoint_path, model_key=model_key, is_openai=False) + # detect old format and make compatible with new format + if 'positional_embedding' in state_dict and not hasattr(model, 'positional_embedding'): + state_dict = convert_to_custom_text_state_dict(state_dict) + if 'text.logit_scale' in state_dict and hasattr(model, 'logit_scale'): + state_dict['logit_scale'] = state_dict['text.logit_scale'] + del state_dict['text.logit_scale'] + + # resize_clip_pos_embed for CLIP and open CLIP + if 'visual.positional_embedding' in state_dict: + resize_clip_pos_embed(state_dict, model) + # specified to eva_vit_model + elif 'visual.pos_embed' in state_dict: + resize_evaclip_pos_embed(state_dict, model) + + # resize_clip_pos_embed(state_dict, model) + incompatible_keys = model.load_state_dict(state_dict, strict=strict) + logging.info(f"incompatible_keys.missing_keys: {incompatible_keys.missing_keys}") + return incompatible_keys + +def load_clip_visual_state_dict(checkpoint_path: str, map_location: str='cpu', is_openai: bool=False, skip_list:list=[]): + state_dict = load_state_dict(checkpoint_path, map_location=map_location, is_openai=is_openai, skip_list=skip_list) + + for k in list(state_dict.keys()): + if not k.startswith('visual.'): + del state_dict[k] + for k in list(state_dict.keys()): + if k.startswith('visual.'): + new_k = k[7:] + state_dict[new_k] = state_dict[k] + del state_dict[k] + return state_dict + +def load_clip_text_state_dict(checkpoint_path: str, map_location: str='cpu', is_openai: bool=False, skip_list:list=[]): + state_dict = load_state_dict(checkpoint_path, map_location=map_location, is_openai=is_openai, skip_list=skip_list) + + for k in list(state_dict.keys()): + if k.startswith('visual.'): + del state_dict[k] + return state_dict + +def get_pretrained_tag(pretrained_model): + pretrained_model = pretrained_model.lower() + if "laion" in pretrained_model or "open_clip" in pretrained_model: + return "open_clip" + elif "openai" in pretrained_model: + return "clip" + elif "eva" in pretrained_model and "clip" in pretrained_model: + return "eva_clip" + else: + return "other" + +def load_pretrained_checkpoint( + model, + visual_checkpoint_path, + text_checkpoint_path, + strict=True, + visual_model=None, + text_model=None, + model_key="model|module|state_dict", + skip_list=[]): + visual_tag = get_pretrained_tag(visual_model) + text_tag = get_pretrained_tag(text_model) + + logging.info(f"num of model state_dict keys: {len(model.state_dict().keys())}") + visual_incompatible_keys, text_incompatible_keys = None, None + if visual_checkpoint_path: + if visual_tag == "eva_clip" or visual_tag == "open_clip": + visual_state_dict = load_clip_visual_state_dict(visual_checkpoint_path, is_openai=False, skip_list=skip_list) + elif visual_tag == "clip": + visual_state_dict = load_clip_visual_state_dict(visual_checkpoint_path, is_openai=True, skip_list=skip_list) + else: + visual_state_dict = load_state_dict(visual_checkpoint_path, model_key=model_key, is_openai=False, skip_list=skip_list) + + # resize_clip_pos_embed for CLIP and open CLIP + if 'positional_embedding' in visual_state_dict: + resize_visual_pos_embed(visual_state_dict, model) + # specified to EVA model + elif 'pos_embed' in visual_state_dict: + resize_eva_pos_embed(visual_state_dict, model) + + visual_incompatible_keys = model.visual.load_state_dict(visual_state_dict, strict=strict) + logging.info(f"num of loaded visual_state_dict keys: {len(visual_state_dict.keys())}") + logging.info(f"visual_incompatible_keys.missing_keys: {visual_incompatible_keys.missing_keys}") + + if text_checkpoint_path: + if text_tag == "eva_clip" or text_tag == "open_clip": + text_state_dict = load_clip_text_state_dict(text_checkpoint_path, is_openai=False, skip_list=skip_list) + elif text_tag == "clip": + text_state_dict = load_clip_text_state_dict(text_checkpoint_path, is_openai=True, skip_list=skip_list) + else: + text_state_dict = load_state_dict(visual_checkpoint_path, model_key=model_key, is_openai=False, skip_list=skip_list) + + text_incompatible_keys = model.text.load_state_dict(text_state_dict, strict=strict) + + logging.info(f"num of loaded text_state_dict keys: {len(text_state_dict.keys())}") + logging.info(f"text_incompatible_keys.missing_keys: {text_incompatible_keys.missing_keys}") + + return visual_incompatible_keys, text_incompatible_keys + +def create_model( + model_name: str, + pretrained: Optional[str] = None, + precision: str = 'fp32', + device: Union[str, torch.device] = 'cpu', + jit: bool = False, + force_quick_gelu: bool = False, + force_custom_clip: bool = False, + force_patch_dropout: Optional[float] = None, + pretrained_image: str = '', + pretrained_text: str = '', + pretrained_hf: bool = True, + pretrained_visual_model: str = None, + pretrained_text_model: str = None, + cache_dir: Optional[str] = None, + skip_list: list = [], +): + model_name = model_name.replace('/', '-') # for callers using old naming with / in ViT names + if isinstance(device, str): + device = torch.device(device) + + if pretrained and pretrained.lower() == 'openai': + logging.info(f'Loading pretrained {model_name} from OpenAI.') + model = load_openai_model( + model_name, + precision=precision, + device=device, + jit=jit, + cache_dir=cache_dir, + ) + else: + model_cfg = get_model_config(model_name) + if model_cfg is not None: + logging.info(f'Loaded {model_name} model config.') + else: + logging.error(f'Model config for {model_name} not found; available models {list_models()}.') + raise RuntimeError(f'Model config for {model_name} not found.') + + if 'rope' in model_cfg.get('vision_cfg', {}): + if model_cfg['vision_cfg']['rope']: + os.environ['RoPE'] = "1" + else: + os.environ['RoPE'] = "0" + + if force_quick_gelu: + # override for use of QuickGELU on non-OpenAI transformer models + model_cfg["quick_gelu"] = True + + if force_patch_dropout is not None: + # override the default patch dropout value + model_cfg['vision_cfg']["patch_dropout"] = force_patch_dropout + + cast_dtype = get_cast_dtype(precision) + custom_clip = model_cfg.pop('custom_text', False) or force_custom_clip or ('hf_model_name' in model_cfg['text_cfg']) + + + if custom_clip: + if 'hf_model_name' in model_cfg.get('text_cfg', {}): + model_cfg['text_cfg']['hf_model_pretrained'] = pretrained_hf + model = CustomCLIP(**model_cfg, cast_dtype=cast_dtype) + else: + model = CLIP(**model_cfg, cast_dtype=cast_dtype) + + pretrained_cfg = {} + if pretrained: + checkpoint_path = '' + pretrained_cfg = get_pretrained_cfg(model_name, pretrained) + if pretrained_cfg: + checkpoint_path = download_pretrained(pretrained_cfg, cache_dir=cache_dir) + elif os.path.exists(pretrained): + checkpoint_path = pretrained + + if checkpoint_path: + logging.info(f'Loading pretrained {model_name} weights ({pretrained}).') + load_checkpoint(model, + checkpoint_path, + model_key="model|module|state_dict", + strict=False + ) + else: + error_str = ( + f'Pretrained weights ({pretrained}) not found for model {model_name}.' + f'Available pretrained tags ({list_pretrained_tags_by_model(model_name)}.') + logging.warning(error_str) + raise RuntimeError(error_str) + else: + visual_checkpoint_path = '' + text_checkpoint_path = '' + + if pretrained_image: + pretrained_visual_model = pretrained_visual_model.replace('/', '-') # for callers using old naming with / in ViT names + pretrained_image_cfg = get_pretrained_cfg(pretrained_visual_model, pretrained_image) + if 'timm_model_name' in model_cfg.get('vision_cfg', {}): + # pretrained weight loading for timm models set via vision_cfg + model_cfg['vision_cfg']['timm_model_pretrained'] = True + elif pretrained_image_cfg: + visual_checkpoint_path = download_pretrained(pretrained_image_cfg, cache_dir=cache_dir) + elif os.path.exists(pretrained_image): + visual_checkpoint_path = pretrained_image + else: + logging.warning(f'Pretrained weights ({visual_checkpoint_path}) not found for model {model_name}.visual.') + raise RuntimeError(f'Pretrained weights ({visual_checkpoint_path}) not found for model {model_name}.visual.') + + if pretrained_text: + pretrained_text_model = pretrained_text_model.replace('/', '-') # for callers using old naming with / in ViT names + pretrained_text_cfg = get_pretrained_cfg(pretrained_text_model, pretrained_text) + if pretrained_image_cfg: + text_checkpoint_path = download_pretrained(pretrained_text_cfg, cache_dir=cache_dir) + elif os.path.exists(pretrained_text): + text_checkpoint_path = pretrained_text + else: + logging.warning(f'Pretrained weights ({text_checkpoint_path}) not found for model {model_name}.text.') + raise RuntimeError(f'Pretrained weights ({text_checkpoint_path}) not found for model {model_name}.text.') + + if visual_checkpoint_path: + logging.info(f'Loading pretrained {model_name}.visual weights ({visual_checkpoint_path}).') + if text_checkpoint_path: + logging.info(f'Loading pretrained {model_name}.text weights ({text_checkpoint_path}).') + + if visual_checkpoint_path or text_checkpoint_path: + load_pretrained_checkpoint( + model, + visual_checkpoint_path, + text_checkpoint_path, + strict=False, + visual_model=pretrained_visual_model, + text_model=pretrained_text_model, + model_key="model|module|state_dict", + skip_list=skip_list + ) + + if "fp16" in precision or "bf16" in precision: + logging.info(f'convert precision to {precision}') + model = model.to(torch.bfloat16) if 'bf16' in precision else model.to(torch.float16) + + model.to(device=device) + + # set image / mean metadata from pretrained_cfg if available, or use default + model.visual.image_mean = pretrained_cfg.get('mean', None) or OPENAI_DATASET_MEAN + model.visual.image_std = pretrained_cfg.get('std', None) or OPENAI_DATASET_STD + + if jit: + model = torch.jit.script(model) + + return model + + +def create_model_and_transforms( + model_name: str, + pretrained: Optional[str] = None, + precision: str = 'fp32', + device: Union[str, torch.device] = 'cpu', + jit: bool = False, + force_quick_gelu: bool = False, + force_custom_clip: bool = False, + force_patch_dropout: Optional[float] = None, + pretrained_image: str = '', + pretrained_text: str = '', + pretrained_hf: bool = True, + pretrained_visual_model: str = None, + pretrained_text_model: str = None, + image_mean: Optional[Tuple[float, ...]] = None, + image_std: Optional[Tuple[float, ...]] = None, + cache_dir: Optional[str] = None, + skip_list: list = [], +): + model = create_model( + model_name, + pretrained, + precision=precision, + device=device, + jit=jit, + force_quick_gelu=force_quick_gelu, + force_custom_clip=force_custom_clip, + force_patch_dropout=force_patch_dropout, + pretrained_image=pretrained_image, + pretrained_text=pretrained_text, + pretrained_hf=pretrained_hf, + pretrained_visual_model=pretrained_visual_model, + pretrained_text_model=pretrained_text_model, + cache_dir=cache_dir, + skip_list=skip_list, + ) + + image_mean = image_mean or getattr(model.visual, 'image_mean', None) + image_std = image_std or getattr(model.visual, 'image_std', None) + preprocess_train = image_transform( + model.visual.image_size, + is_train=True, + mean=image_mean, + std=image_std + ) + preprocess_val = image_transform( + model.visual.image_size, + is_train=False, + mean=image_mean, + std=image_std + ) + + return model, preprocess_train, preprocess_val + + +def create_transforms( + model_name: str, + pretrained: Optional[str] = None, + precision: str = 'fp32', + device: Union[str, torch.device] = 'cpu', + jit: bool = False, + force_quick_gelu: bool = False, + force_custom_clip: bool = False, + force_patch_dropout: Optional[float] = None, + pretrained_image: str = '', + pretrained_text: str = '', + pretrained_hf: bool = True, + pretrained_visual_model: str = None, + pretrained_text_model: str = None, + image_mean: Optional[Tuple[float, ...]] = None, + image_std: Optional[Tuple[float, ...]] = None, + cache_dir: Optional[str] = None, + skip_list: list = [], +): + model = create_model( + model_name, + pretrained, + precision=precision, + device=device, + jit=jit, + force_quick_gelu=force_quick_gelu, + force_custom_clip=force_custom_clip, + force_patch_dropout=force_patch_dropout, + pretrained_image=pretrained_image, + pretrained_text=pretrained_text, + pretrained_hf=pretrained_hf, + pretrained_visual_model=pretrained_visual_model, + pretrained_text_model=pretrained_text_model, + cache_dir=cache_dir, + skip_list=skip_list, + ) + + + image_mean = image_mean or getattr(model.visual, 'image_mean', None) + image_std = image_std or getattr(model.visual, 'image_std', None) + preprocess_train = image_transform( + model.visual.image_size, + is_train=True, + mean=image_mean, + std=image_std + ) + preprocess_val = image_transform( + model.visual.image_size, + is_train=False, + mean=image_mean, + std=image_std + ) + del model + + return preprocess_train, preprocess_val + +def create_model_from_pretrained( + model_name: str, + pretrained: str, + precision: str = 'fp32', + device: Union[str, torch.device] = 'cpu', + jit: bool = False, + force_quick_gelu: bool = False, + force_custom_clip: bool = False, + force_patch_dropout: Optional[float] = None, + return_transform: bool = True, + image_mean: Optional[Tuple[float, ...]] = None, + image_std: Optional[Tuple[float, ...]] = None, + cache_dir: Optional[str] = None, + is_frozen: bool = False, +): + if not is_pretrained_cfg(model_name, pretrained) and not os.path.exists(pretrained): + raise RuntimeError( + f'{pretrained} is not a valid pretrained cfg or checkpoint for {model_name}.' + f' Use open_clip.list_pretrained() to find one.') + + model = create_model( + model_name, + pretrained, + precision=precision, + device=device, + jit=jit, + force_quick_gelu=force_quick_gelu, + force_custom_clip=force_custom_clip, + force_patch_dropout=force_patch_dropout, + cache_dir=cache_dir, + ) + + if is_frozen: + for param in model.parameters(): + param.requires_grad = False + + if not return_transform: + return model + + image_mean = image_mean or getattr(model.visual, 'image_mean', None) + image_std = image_std or getattr(model.visual, 'image_std', None) + preprocess = image_transform( + model.visual.image_size, + is_train=False, + mean=image_mean, + std=image_std + ) + + return model, preprocess diff --git a/modules/pulid/eva_clip/hf_configs.py b/modules/pulid/eva_clip/hf_configs.py new file mode 100644 index 000000000..a8c9b704d --- /dev/null +++ b/modules/pulid/eva_clip/hf_configs.py @@ -0,0 +1,57 @@ +# HF architecture dict: +arch_dict = { + # https://huggingface.co/docs/transformers/model_doc/roberta#roberta + "roberta": { + "config_names": { + "context_length": "max_position_embeddings", + "vocab_size": "vocab_size", + "width": "hidden_size", + "heads": "num_attention_heads", + "layers": "num_hidden_layers", + "layer_attr": "layer", + "token_embeddings_attr": "embeddings" + }, + "pooler": "mean_pooler", + }, + # https://huggingface.co/docs/transformers/model_doc/xlm-roberta#transformers.XLMRobertaConfig + "xlm-roberta": { + "config_names": { + "context_length": "max_position_embeddings", + "vocab_size": "vocab_size", + "width": "hidden_size", + "heads": "num_attention_heads", + "layers": "num_hidden_layers", + "layer_attr": "layer", + "token_embeddings_attr": "embeddings" + }, + "pooler": "mean_pooler", + }, + # https://huggingface.co/docs/transformers/model_doc/mt5#mt5 + "mt5": { + "config_names": { + # unlimited seqlen + # https://github.com/google-research/text-to-text-transfer-transformer/issues/273 + # https://github.com/huggingface/transformers/blob/v4.24.0/src/transformers/models/t5/modeling_t5.py#L374 + "context_length": "", + "vocab_size": "vocab_size", + "width": "d_model", + "heads": "num_heads", + "layers": "num_layers", + "layer_attr": "block", + "token_embeddings_attr": "embed_tokens" + }, + "pooler": "mean_pooler", + }, + "bert": { + "config_names": { + "context_length": "max_position_embeddings", + "vocab_size": "vocab_size", + "width": "hidden_size", + "heads": "num_attention_heads", + "layers": "num_hidden_layers", + "layer_attr": "layer", + "token_embeddings_attr": "embeddings" + }, + "pooler": "mean_pooler", + } +} diff --git a/modules/pulid/eva_clip/hf_model.py b/modules/pulid/eva_clip/hf_model.py new file mode 100644 index 000000000..c4b9fd85b --- /dev/null +++ b/modules/pulid/eva_clip/hf_model.py @@ -0,0 +1,248 @@ +""" huggingface model adapter + +Wraps HuggingFace transformers (https://github.com/huggingface/transformers) models for use as a text tower in CLIP model. +""" + +import re + +import torch +import torch.nn as nn +from torch.nn import functional as F +from torch import TensorType +try: + import transformers + from transformers import AutoModel, AutoModelForMaskedLM, AutoTokenizer, AutoConfig, PretrainedConfig + from transformers.modeling_outputs import BaseModelOutput, BaseModelOutputWithPooling, \ + BaseModelOutputWithPoolingAndCrossAttentions +except ImportError as e: + transformers = None + + + class BaseModelOutput: + pass + + + class PretrainedConfig: + pass + +from .hf_configs import arch_dict + +# utils +def _camel2snake(s): + return re.sub(r'(? TensorType: + # image_atts = torch.ones(image_embeds.size()[:-1],dtype=torch.long).to(x.device) + # attn_mask = (x != self.config.pad_token_id).long() + # out = self.transformer( + # input_ids=x, + # attention_mask=attn_mask, + # encoder_hidden_states = image_embeds, + # encoder_attention_mask = image_atts, + # ) + # pooled_out = self.pooler(out, attn_mask) + + # return self.itm_proj(pooled_out) + + def mask(self, input_ids, vocab_size, device, targets=None, masked_indices=None, probability_matrix=None): + if masked_indices is None: + masked_indices = torch.bernoulli(probability_matrix).bool() + + masked_indices[input_ids == self.tokenizer.pad_token_id] = False + masked_indices[input_ids == self.tokenizer.cls_token_id] = False + + if targets is not None: + targets[~masked_indices] = -100 # We only compute loss on masked tokens + + # 80% of the time, we replace masked input tokens with tokenizer.mask_token ([MASK]) + indices_replaced = torch.bernoulli(torch.full(input_ids.shape, 0.8)).bool() & masked_indices + input_ids[indices_replaced] = self.tokenizer.mask_token_id + + # 10% of the time, we replace masked input tokens with random word + indices_random = torch.bernoulli(torch.full(input_ids.shape, 0.5)).bool() & masked_indices & ~indices_replaced + random_words = torch.randint(vocab_size, input_ids.shape, dtype=torch.long).to(device) + input_ids[indices_random] = random_words[indices_random] + # The rest of the time (10% of the time) we keep the masked input tokens unchanged + + if targets is not None: + return input_ids, targets + else: + return input_ids + + def forward_mlm(self, input_ids, image_embeds, mlm_probability=0.25): + labels = input_ids.clone() + attn_mask = (input_ids != self.config.pad_token_id).long() + image_atts = torch.ones(image_embeds.size()[:-1],dtype=torch.long).to(input_ids.device) + vocab_size = getattr(self.config, arch_dict[self.config.model_type]["config_names"]["vocab_size"]) + probability_matrix = torch.full(labels.shape, mlm_probability) + input_ids, labels = self.mask(input_ids, vocab_size, input_ids.device, targets=labels, + probability_matrix = probability_matrix) + mlm_output = self.transformer(input_ids, + attention_mask = attn_mask, + encoder_hidden_states = image_embeds, + encoder_attention_mask = image_atts, + return_dict = True, + labels = labels, + ) + return mlm_output.loss + # mlm_output = self.transformer(input_ids, + # attention_mask = attn_mask, + # encoder_hidden_states = image_embeds, + # encoder_attention_mask = image_atts, + # return_dict = True, + # ).last_hidden_state + # logits = self.mlm_proj(mlm_output) + + # # logits = logits[:, :-1, :].contiguous().view(-1, vocab_size) + # logits = logits[:, 1:, :].contiguous().view(-1, vocab_size) + # labels = labels[:, 1:].contiguous().view(-1) + + # mlm_loss = F.cross_entropy( + # logits, + # labels, + # # label_smoothing=0.1, + # ) + # return mlm_loss + + + def forward(self, x:TensorType) -> TensorType: + attn_mask = (x != self.config.pad_token_id).long() + out = self.transformer(input_ids=x, attention_mask=attn_mask) + pooled_out = self.pooler(out, attn_mask) + + return self.proj(pooled_out) + + def lock(self, unlocked_layers:int=0, freeze_layer_norm:bool=True): + if not unlocked_layers: # full freezing + for n, p in self.transformer.named_parameters(): + p.requires_grad = (not freeze_layer_norm) if "LayerNorm" in n.split(".") else False + return + + encoder = self.transformer.encoder if hasattr(self.transformer, 'encoder') else self.transformer + layer_list = getattr(encoder, arch_dict[self.config.model_type]["config_names"]["layer_attr"]) + print(f"Unlocking {unlocked_layers}/{len(layer_list) + 1} layers of hf model") + embeddings = getattr( + self.transformer, arch_dict[self.config.model_type]["config_names"]["token_embeddings_attr"]) + modules = [embeddings, *layer_list][:-unlocked_layers] + # freeze layers + for module in modules: + for n, p in module.named_parameters(): + p.requires_grad = (not freeze_layer_norm) if "LayerNorm" in n.split(".") else False + + + @torch.jit.ignore + def set_grad_checkpointing(self, enable=True): + self.transformer.gradient_checkpointing_enable() + + def get_num_layers(self): + encoder = self.transformer.encoder if hasattr(self.transformer, 'encoder') else self.transformer + layer_list = getattr(encoder, arch_dict[self.config.model_type]["config_names"]["layer_attr"]) + return len(layer_list) + + def init_parameters(self): + pass diff --git a/modules/pulid/eva_clip/loss.py b/modules/pulid/eva_clip/loss.py new file mode 100644 index 000000000..473f60d98 --- /dev/null +++ b/modules/pulid/eva_clip/loss.py @@ -0,0 +1,138 @@ +import math +import torch +import torch.nn as nn +from torch.nn import functional as F + +try: + import torch.distributed.nn + from torch import distributed as dist + has_distributed = True +except ImportError: + has_distributed = False + +try: + import horovod.torch as hvd +except ImportError: + hvd = None + +from timm.loss import LabelSmoothingCrossEntropy + + +def gather_features( + image_features, + text_features, + local_loss=False, + gather_with_grad=False, + rank=0, + world_size=1, + use_horovod=False +): + assert has_distributed, 'torch.distributed did not import correctly, please use a PyTorch version with support.' + if use_horovod: + assert hvd is not None, 'Please install horovod' + if gather_with_grad: + all_image_features = hvd.allgather(image_features) + all_text_features = hvd.allgather(text_features) + else: + with torch.no_grad(): + all_image_features = hvd.allgather(image_features) + all_text_features = hvd.allgather(text_features) + if not local_loss: + # ensure grads for local rank when all_* features don't have a gradient + gathered_image_features = list(all_image_features.chunk(world_size, dim=0)) + gathered_text_features = list(all_text_features.chunk(world_size, dim=0)) + gathered_image_features[rank] = image_features + gathered_text_features[rank] = text_features + all_image_features = torch.cat(gathered_image_features, dim=0) + all_text_features = torch.cat(gathered_text_features, dim=0) + else: + # We gather tensors from all gpus + if gather_with_grad: + all_image_features = torch.cat(torch.distributed.nn.all_gather(image_features), dim=0) + all_text_features = torch.cat(torch.distributed.nn.all_gather(text_features), dim=0) + # all_image_features = torch.cat(torch.distributed.nn.all_gather(image_features, async_op=True), dim=0) + # all_text_features = torch.cat(torch.distributed.nn.all_gather(text_features, async_op=True), dim=0) + else: + gathered_image_features = [torch.zeros_like(image_features) for _ in range(world_size)] + gathered_text_features = [torch.zeros_like(text_features) for _ in range(world_size)] + dist.all_gather(gathered_image_features, image_features) + dist.all_gather(gathered_text_features, text_features) + if not local_loss: + # ensure grads for local rank when all_* features don't have a gradient + gathered_image_features[rank] = image_features + gathered_text_features[rank] = text_features + all_image_features = torch.cat(gathered_image_features, dim=0) + all_text_features = torch.cat(gathered_text_features, dim=0) + + return all_image_features, all_text_features + + +class ClipLoss(nn.Module): + + def __init__( + self, + local_loss=False, + gather_with_grad=False, + cache_labels=False, + rank=0, + world_size=1, + use_horovod=False, + smoothing=0., + ): + super().__init__() + self.local_loss = local_loss + self.gather_with_grad = gather_with_grad + self.cache_labels = cache_labels + self.rank = rank + self.world_size = world_size + self.use_horovod = use_horovod + self.label_smoothing_cross_entropy = LabelSmoothingCrossEntropy(smoothing=smoothing) if smoothing > 0 else None + + # cache state + self.prev_num_logits = 0 + self.labels = {} + + def forward(self, image_features, text_features, logit_scale=1.): + device = image_features.device + if self.world_size > 1: + all_image_features, all_text_features = gather_features( + image_features, text_features, + self.local_loss, self.gather_with_grad, self.rank, self.world_size, self.use_horovod) + + if self.local_loss: + logits_per_image = logit_scale * image_features @ all_text_features.T + logits_per_text = logit_scale * text_features @ all_image_features.T + else: + logits_per_image = logit_scale * all_image_features @ all_text_features.T + logits_per_text = logits_per_image.T + else: + logits_per_image = logit_scale * image_features @ text_features.T + logits_per_text = logit_scale * text_features @ image_features.T + # calculated ground-truth and cache if enabled + num_logits = logits_per_image.shape[0] + if self.prev_num_logits != num_logits or device not in self.labels: + labels = torch.arange(num_logits, device=device, dtype=torch.long) + if self.world_size > 1 and self.local_loss: + labels = labels + num_logits * self.rank + if self.cache_labels: + self.labels[device] = labels + self.prev_num_logits = num_logits + else: + labels = self.labels[device] + + if self.label_smoothing_cross_entropy: + total_loss = ( + self.label_smoothing_cross_entropy(logits_per_image, labels) + + self.label_smoothing_cross_entropy(logits_per_text, labels) + ) / 2 + else: + total_loss = ( + F.cross_entropy(logits_per_image, labels) + + F.cross_entropy(logits_per_text, labels) + ) / 2 + + acc = None + i2t_acc = (logits_per_image.argmax(-1) == labels).sum() / len(logits_per_image) + t2i_acc = (logits_per_text.argmax(-1) == labels).sum() / len(logits_per_text) + acc = {"i2t": i2t_acc, "t2i": t2i_acc} + return total_loss, acc \ No newline at end of file diff --git a/modules/pulid/eva_clip/model.py b/modules/pulid/eva_clip/model.py new file mode 100644 index 000000000..abd8c02db --- /dev/null +++ b/modules/pulid/eva_clip/model.py @@ -0,0 +1,432 @@ +""" CLIP Model + +Adapted from https://github.com/openai/CLIP. Originally MIT License, Copyright (c) 2021 OpenAI. +""" +import os +from dataclasses import dataclass +from typing import Optional, Tuple, Union +from functools import partial + +import numpy as np +import torch +import torch.nn.functional as F +from torch import nn + +try: + from .hf_model import HFTextEncoder +except: + HFTextEncoder = None +from .modified_resnet import ModifiedResNet +from .timm_model import TimmModel +from .eva_vit_model import EVAVisionTransformer +from .transformer import LayerNorm, QuickGELU, Attention, VisionTransformer, TextTransformer + +try: + from apex.normalization import FusedLayerNorm +except: + FusedLayerNorm = LayerNorm + +@dataclass +class CLIPVisionCfg: + layers: Union[Tuple[int, int, int, int], int] = 12 + width: int = 768 + head_width: int = 64 + mlp_ratio: float = 4.0 + patch_size: int = 16 + image_size: Union[Tuple[int, int], int] = 224 + ls_init_value: Optional[float] = None # layer scale initial value + patch_dropout: float = 0. # what fraction of patches to dropout during training (0 would mean disabled and no patches dropped) - 0.5 to 0.75 recommended in the paper for optimal results + global_average_pool: bool = False # whether to global average pool the last embedding layer, instead of using CLS token (https://arxiv.org/abs/2205.01580) + drop_path_rate: Optional[float] = None # drop path rate + timm_model_name: str = None # a valid model name overrides layers, width, patch_size + timm_model_pretrained: bool = False # use (imagenet) pretrained weights for named model + timm_pool: str = 'avg' # feature pooling for timm model ('abs_attn', 'rot_attn', 'avg', '') + timm_proj: str = 'linear' # linear projection for timm model output ('linear', 'mlp', '') + timm_proj_bias: bool = False # enable bias final projection + eva_model_name: str = None # a valid eva model name overrides layers, width, patch_size + qkv_bias: bool = True + fusedLN: bool = False + xattn: bool = False + postnorm: bool = False + rope: bool = False + pt_hw_seq_len: int = 16 # 224/14 + intp_freq: bool = False + naiveswiglu: bool = False + subln: bool = False + + +@dataclass +class CLIPTextCfg: + context_length: int = 77 + vocab_size: int = 49408 + width: int = 512 + heads: int = 8 + layers: int = 12 + ls_init_value: Optional[float] = None # layer scale initial value + hf_model_name: str = None + hf_tokenizer_name: str = None + hf_model_pretrained: bool = True + proj: str = 'mlp' + pooler_type: str = 'mean_pooler' + masked_language_modeling: bool = False + fusedLN: bool = False + xattn: bool = False + attn_mask: bool = True + +def get_cast_dtype(precision: str): + cast_dtype = None + if precision == 'bf16': + cast_dtype = torch.bfloat16 + elif precision == 'fp16': + cast_dtype = torch.float16 + return cast_dtype + + +def _build_vision_tower( + embed_dim: int, + vision_cfg: CLIPVisionCfg, + quick_gelu: bool = False, + cast_dtype: Optional[torch.dtype] = None +): + if isinstance(vision_cfg, dict): + vision_cfg = CLIPVisionCfg(**vision_cfg) + + # OpenAI models are pretrained w/ QuickGELU but native nn.GELU is both faster and more + # memory efficient in recent PyTorch releases (>= 1.10). + # NOTE: timm models always use native GELU regardless of quick_gelu flag. + act_layer = QuickGELU if quick_gelu else nn.GELU + + if vision_cfg.eva_model_name: + vision_heads = vision_cfg.width // vision_cfg.head_width + norm_layer = LayerNorm + + visual = EVAVisionTransformer( + img_size=vision_cfg.image_size, + patch_size=vision_cfg.patch_size, + num_classes=embed_dim, + use_mean_pooling=vision_cfg.global_average_pool, #False + init_values=vision_cfg.ls_init_value, + patch_dropout=vision_cfg.patch_dropout, + embed_dim=vision_cfg.width, + depth=vision_cfg.layers, + num_heads=vision_heads, + mlp_ratio=vision_cfg.mlp_ratio, + qkv_bias=vision_cfg.qkv_bias, + drop_path_rate=vision_cfg.drop_path_rate, + norm_layer= partial(FusedLayerNorm, eps=1e-6) if vision_cfg.fusedLN else partial(norm_layer, eps=1e-6), + xattn=vision_cfg.xattn, + rope=vision_cfg.rope, + postnorm=vision_cfg.postnorm, + pt_hw_seq_len= vision_cfg.pt_hw_seq_len, # 224/14 + intp_freq= vision_cfg.intp_freq, + naiveswiglu= vision_cfg.naiveswiglu, + subln= vision_cfg.subln + ) + elif vision_cfg.timm_model_name: + visual = TimmModel( + vision_cfg.timm_model_name, + pretrained=vision_cfg.timm_model_pretrained, + pool=vision_cfg.timm_pool, + proj=vision_cfg.timm_proj, + proj_bias=vision_cfg.timm_proj_bias, + embed_dim=embed_dim, + image_size=vision_cfg.image_size + ) + act_layer = nn.GELU # so that text transformer doesn't use QuickGELU w/ timm models + elif isinstance(vision_cfg.layers, (tuple, list)): + vision_heads = vision_cfg.width * 32 // vision_cfg.head_width + visual = ModifiedResNet( + layers=vision_cfg.layers, + output_dim=embed_dim, + heads=vision_heads, + image_size=vision_cfg.image_size, + width=vision_cfg.width + ) + else: + vision_heads = vision_cfg.width // vision_cfg.head_width + norm_layer = LayerNormFp32 if cast_dtype in (torch.float16, torch.bfloat16) else LayerNorm + visual = VisionTransformer( + image_size=vision_cfg.image_size, + patch_size=vision_cfg.patch_size, + width=vision_cfg.width, + layers=vision_cfg.layers, + heads=vision_heads, + mlp_ratio=vision_cfg.mlp_ratio, + ls_init_value=vision_cfg.ls_init_value, + patch_dropout=vision_cfg.patch_dropout, + global_average_pool=vision_cfg.global_average_pool, + output_dim=embed_dim, + act_layer=act_layer, + norm_layer=norm_layer, + ) + + return visual + + +def _build_text_tower( + embed_dim: int, + text_cfg: CLIPTextCfg, + quick_gelu: bool = False, + cast_dtype: Optional[torch.dtype] = None, +): + if isinstance(text_cfg, dict): + text_cfg = CLIPTextCfg(**text_cfg) + + if text_cfg.hf_model_name: + text = HFTextEncoder( + text_cfg.hf_model_name, + output_dim=embed_dim, + tokenizer_name=text_cfg.hf_tokenizer_name, + proj=text_cfg.proj, + pooler_type=text_cfg.pooler_type, + masked_language_modeling=text_cfg.masked_language_modeling + ) + else: + act_layer = QuickGELU if quick_gelu else nn.GELU + norm_layer = LayerNorm + + text = TextTransformer( + context_length=text_cfg.context_length, + vocab_size=text_cfg.vocab_size, + width=text_cfg.width, + heads=text_cfg.heads, + layers=text_cfg.layers, + ls_init_value=text_cfg.ls_init_value, + output_dim=embed_dim, + act_layer=act_layer, + norm_layer= FusedLayerNorm if text_cfg.fusedLN else norm_layer, + xattn=text_cfg.xattn, + attn_mask=text_cfg.attn_mask, + ) + return text + +class CLIP(nn.Module): + def __init__( + self, + embed_dim: int, + vision_cfg: CLIPVisionCfg, + text_cfg: CLIPTextCfg, + quick_gelu: bool = False, + cast_dtype: Optional[torch.dtype] = None, + ): + super().__init__() + self.visual = _build_vision_tower(embed_dim, vision_cfg, quick_gelu, cast_dtype) + + text = _build_text_tower(embed_dim, text_cfg, quick_gelu, cast_dtype) + self.transformer = text.transformer + self.vocab_size = text.vocab_size + self.token_embedding = text.token_embedding + self.positional_embedding = text.positional_embedding + self.ln_final = text.ln_final + self.text_projection = text.text_projection + self.register_buffer('attn_mask', text.attn_mask, persistent=False) + + self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07)) + + def lock_image_tower(self, unlocked_groups=0, freeze_bn_stats=False): + # lock image tower as per LiT - https://arxiv.org/abs/2111.07991 + self.visual.lock(unlocked_groups=unlocked_groups, freeze_bn_stats=freeze_bn_stats) + + @torch.jit.ignore + def set_grad_checkpointing(self, enable=True): + self.visual.set_grad_checkpointing(enable) + self.transformer.grad_checkpointing = enable + + @torch.jit.ignore + def no_weight_decay(self): + return {'logit_scale'} + + def encode_image(self, image, normalize: bool = False): + features = self.visual(image) + return F.normalize(features, dim=-1) if normalize else features + + def encode_text(self, text, normalize: bool = False): + cast_dtype = self.transformer.get_cast_dtype() + + x = self.token_embedding(text).to(cast_dtype) # [batch_size, n_ctx, d_model] + + x = x + self.positional_embedding.to(cast_dtype) + x = x.permute(1, 0, 2) # NLD -> LND + x = self.transformer(x, attn_mask=self.attn_mask) + x = x.permute(1, 0, 2) # LND -> NLD + x = self.ln_final(x) # [batch_size, n_ctx, transformer.width] + # take features from the eot embedding (eot_token is the highest number in each sequence) + x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection + return F.normalize(x, dim=-1) if normalize else x + + def forward(self, image, text): + image_features = self.encode_image(image, normalize=True) + text_features = self.encode_text(text, normalize=True) + return image_features, text_features, self.logit_scale.exp() + + +class CustomCLIP(nn.Module): + def __init__( + self, + embed_dim: int, + vision_cfg: CLIPVisionCfg, + text_cfg: CLIPTextCfg, + quick_gelu: bool = False, + cast_dtype: Optional[torch.dtype] = None, + itm_task: bool = False, + ): + super().__init__() + self.visual = _build_vision_tower(embed_dim, vision_cfg, quick_gelu, cast_dtype) + self.text = _build_text_tower(embed_dim, text_cfg, quick_gelu, cast_dtype) + self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07)) + + def lock_image_tower(self, unlocked_groups=0, freeze_bn_stats=False): + # lock image tower as per LiT - https://arxiv.org/abs/2111.07991 + self.visual.lock(unlocked_groups=unlocked_groups, freeze_bn_stats=freeze_bn_stats) + + def lock_text_tower(self, unlocked_layers:int=0, freeze_layer_norm:bool=True): + self.text.lock(unlocked_layers, freeze_layer_norm) + + @torch.jit.ignore + def set_grad_checkpointing(self, enable=True): + self.visual.set_grad_checkpointing(enable) + self.text.set_grad_checkpointing(enable) + + @torch.jit.ignore + def no_weight_decay(self): + return {'logit_scale'} + + def encode_image(self, image, normalize: bool = False): + features = self.visual(image) + return F.normalize(features, dim=-1) if normalize else features + + def encode_text(self, text, normalize: bool = False): + features = self.text(text) + return F.normalize(features, dim=-1) if normalize else features + + def forward(self, image, text): + image_features = self.encode_image(image, normalize=True) + text_features = self.encode_text(text, normalize=True) + return image_features, text_features, self.logit_scale.exp() + + +def convert_weights_to_lp(model: nn.Module, dtype=torch.float16): + """Convert applicable model parameters to low-precision (bf16 or fp16)""" + + def _convert_weights(l): + + if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Linear)): + l.weight.data = l.weight.data.to(dtype) + if l.bias is not None: + l.bias.data = l.bias.data.to(dtype) + + if isinstance(l, (nn.MultiheadAttention, Attention)): + for attr in [*[f"{s}_proj_weight" for s in ["in", "q", "k", "v"]], "in_proj_bias", "bias_k", "bias_v"]: + tensor = getattr(l, attr, None) + if tensor is not None: + tensor.data = tensor.data.to(dtype) + + if isinstance(l, nn.Parameter): + l.data = l.data.to(dtype) + + for name in ["text_projection", "proj"]: + if hasattr(l, name) and isinstance(l, nn.Parameter): + attr = getattr(l, name, None) + if attr is not None: + attr.data = attr.data.to(dtype) + + model.apply(_convert_weights) + + +convert_weights_to_fp16 = convert_weights_to_lp # backwards compat + + +# used to maintain checkpoint compatibility +def convert_to_custom_text_state_dict(state_dict: dict): + if 'text_projection' in state_dict: + # old format state_dict, move text tower -> .text + new_state_dict = {} + for k, v in state_dict.items(): + if any(k.startswith(p) for p in ( + 'text_projection', + 'positional_embedding', + 'token_embedding', + 'transformer', + 'ln_final', + 'logit_scale' + )): + k = 'text.' + k + new_state_dict[k] = v + return new_state_dict + return state_dict + + +def build_model_from_openai_state_dict( + state_dict: dict, + quick_gelu=True, + cast_dtype=torch.float16, +): + vit = "visual.proj" in state_dict + + if vit: + vision_width = state_dict["visual.conv1.weight"].shape[0] + vision_layers = len( + [k for k in state_dict.keys() if k.startswith("visual.") and k.endswith(".attn.in_proj_weight")]) + vision_patch_size = state_dict["visual.conv1.weight"].shape[-1] + grid_size = round((state_dict["visual.positional_embedding"].shape[0] - 1) ** 0.5) + image_size = vision_patch_size * grid_size + else: + counts: list = [ + len(set(k.split(".")[2] for k in state_dict if k.startswith(f"visual.layer{b}"))) for b in [1, 2, 3, 4]] + vision_layers = tuple(counts) + vision_width = state_dict["visual.layer1.0.conv1.weight"].shape[0] + output_width = round((state_dict["visual.attnpool.positional_embedding"].shape[0] - 1) ** 0.5) + vision_patch_size = None + assert output_width ** 2 + 1 == state_dict["visual.attnpool.positional_embedding"].shape[0] + image_size = output_width * 32 + + embed_dim = state_dict["text_projection"].shape[1] + context_length = state_dict["positional_embedding"].shape[0] + vocab_size = state_dict["token_embedding.weight"].shape[0] + transformer_width = state_dict["ln_final.weight"].shape[0] + transformer_heads = transformer_width // 64 + transformer_layers = len(set(k.split(".")[2] for k in state_dict if k.startswith(f"transformer.resblocks"))) + + vision_cfg = CLIPVisionCfg( + layers=vision_layers, + width=vision_width, + patch_size=vision_patch_size, + image_size=image_size, + ) + text_cfg = CLIPTextCfg( + context_length=context_length, + vocab_size=vocab_size, + width=transformer_width, + heads=transformer_heads, + layers=transformer_layers + ) + model = CLIP( + embed_dim, + vision_cfg=vision_cfg, + text_cfg=text_cfg, + quick_gelu=quick_gelu, # OpenAI models were trained with QuickGELU + cast_dtype=cast_dtype, + ) + + for key in ["input_resolution", "context_length", "vocab_size"]: + state_dict.pop(key, None) + + convert_weights_to_fp16(model) # OpenAI state dicts are partially converted to float16 + model.load_state_dict(state_dict) + return model.eval() + + +def trace_model(model, batch_size=256, device=torch.device('cpu')): + model.eval() + image_size = model.visual.image_size + example_images = torch.ones((batch_size, 3, image_size, image_size), device=device) + example_text = torch.zeros((batch_size, model.context_length), dtype=torch.int, device=device) + model = torch.jit.trace_module( + model, + inputs=dict( + forward=(example_images, example_text), + encode_text=(example_text,), + encode_image=(example_images,) + )) + model.visual.image_size = image_size + return model diff --git a/modules/pulid/eva_clip/model_configs/EVA01-CLIP-B-16.json b/modules/pulid/eva_clip/model_configs/EVA01-CLIP-B-16.json new file mode 100644 index 000000000..aad205800 --- /dev/null +++ b/modules/pulid/eva_clip/model_configs/EVA01-CLIP-B-16.json @@ -0,0 +1,19 @@ +{ + "embed_dim": 512, + "vision_cfg": { + "image_size": 224, + "layers": 12, + "width": 768, + "patch_size": 16, + "eva_model_name": "eva-clip-b-16", + "ls_init_value": 0.1, + "drop_path_rate": 0.0 + }, + "text_cfg": { + "context_length": 77, + "vocab_size": 49408, + "width": 512, + "heads": 8, + "layers": 12 + } +} \ No newline at end of file diff --git a/modules/pulid/eva_clip/model_configs/EVA01-CLIP-g-14-plus.json b/modules/pulid/eva_clip/model_configs/EVA01-CLIP-g-14-plus.json new file mode 100644 index 000000000..100279572 --- /dev/null +++ b/modules/pulid/eva_clip/model_configs/EVA01-CLIP-g-14-plus.json @@ -0,0 +1,24 @@ +{ + "embed_dim": 1024, + "vision_cfg": { + "image_size": 224, + "layers": 40, + "width": 1408, + "head_width": 88, + "mlp_ratio": 4.3637, + "patch_size": 14, + "eva_model_name": "eva-clip-g-14-x", + "drop_path_rate": 0, + "xattn": true, + "fusedLN": true + }, + "text_cfg": { + "context_length": 77, + "vocab_size": 49408, + "width": 1024, + "heads": 16, + "layers": 24, + "xattn": false, + "fusedLN": true + } +} \ No newline at end of file diff --git a/modules/pulid/eva_clip/model_configs/EVA01-CLIP-g-14.json b/modules/pulid/eva_clip/model_configs/EVA01-CLIP-g-14.json new file mode 100644 index 000000000..5d338b4e6 --- /dev/null +++ b/modules/pulid/eva_clip/model_configs/EVA01-CLIP-g-14.json @@ -0,0 +1,24 @@ +{ + "embed_dim": 1024, + "vision_cfg": { + "image_size": 224, + "layers": 40, + "width": 1408, + "head_width": 88, + "mlp_ratio": 4.3637, + "patch_size": 14, + "eva_model_name": "eva-clip-g-14-x", + "drop_path_rate": 0.4, + "xattn": true, + "fusedLN": true + }, + "text_cfg": { + "context_length": 77, + "vocab_size": 49408, + "width": 768, + "heads": 12, + "layers": 12, + "xattn": false, + "fusedLN": true + } +} \ No newline at end of file diff --git a/modules/pulid/eva_clip/model_configs/EVA02-CLIP-B-16.json b/modules/pulid/eva_clip/model_configs/EVA02-CLIP-B-16.json new file mode 100644 index 000000000..e4a6e723f --- /dev/null +++ b/modules/pulid/eva_clip/model_configs/EVA02-CLIP-B-16.json @@ -0,0 +1,29 @@ +{ + "embed_dim": 512, + "vision_cfg": { + "image_size": 224, + "layers": 12, + "width": 768, + "head_width": 64, + "patch_size": 16, + "mlp_ratio": 2.6667, + "eva_model_name": "eva-clip-b-16-X", + "drop_path_rate": 0.0, + "xattn": true, + "fusedLN": true, + "rope": true, + "pt_hw_seq_len": 16, + "intp_freq": true, + "naiveswiglu": true, + "subln": true + }, + "text_cfg": { + "context_length": 77, + "vocab_size": 49408, + "width": 512, + "heads": 8, + "layers": 12, + "xattn": true, + "fusedLN": true + } +} \ No newline at end of file diff --git a/modules/pulid/eva_clip/model_configs/EVA02-CLIP-L-14-336.json b/modules/pulid/eva_clip/model_configs/EVA02-CLIP-L-14-336.json new file mode 100644 index 000000000..3e1d124e1 --- /dev/null +++ b/modules/pulid/eva_clip/model_configs/EVA02-CLIP-L-14-336.json @@ -0,0 +1,29 @@ +{ + "embed_dim": 768, + "vision_cfg": { + "image_size": 336, + "layers": 24, + "width": 1024, + "drop_path_rate": 0, + "head_width": 64, + "mlp_ratio": 2.6667, + "patch_size": 14, + "eva_model_name": "eva-clip-l-14-336", + "xattn": true, + "fusedLN": true, + "rope": true, + "pt_hw_seq_len": 16, + "intp_freq": true, + "naiveswiglu": true, + "subln": true + }, + "text_cfg": { + "context_length": 77, + "vocab_size": 49408, + "width": 768, + "heads": 12, + "layers": 12, + "xattn": false, + "fusedLN": true + } +} \ No newline at end of file diff --git a/modules/pulid/eva_clip/model_configs/EVA02-CLIP-L-14.json b/modules/pulid/eva_clip/model_configs/EVA02-CLIP-L-14.json new file mode 100644 index 000000000..03b22ad3c --- /dev/null +++ b/modules/pulid/eva_clip/model_configs/EVA02-CLIP-L-14.json @@ -0,0 +1,29 @@ +{ + "embed_dim": 768, + "vision_cfg": { + "image_size": 224, + "layers": 24, + "width": 1024, + "drop_path_rate": 0, + "head_width": 64, + "mlp_ratio": 2.6667, + "patch_size": 14, + "eva_model_name": "eva-clip-l-14", + "xattn": true, + "fusedLN": true, + "rope": true, + "pt_hw_seq_len": 16, + "intp_freq": true, + "naiveswiglu": true, + "subln": true + }, + "text_cfg": { + "context_length": 77, + "vocab_size": 49408, + "width": 768, + "heads": 12, + "layers": 12, + "xattn": false, + "fusedLN": true + } +} \ No newline at end of file diff --git a/modules/pulid/eva_clip/model_configs/EVA02-CLIP-bigE-14-plus.json b/modules/pulid/eva_clip/model_configs/EVA02-CLIP-bigE-14-plus.json new file mode 100644 index 000000000..aa04e2545 --- /dev/null +++ b/modules/pulid/eva_clip/model_configs/EVA02-CLIP-bigE-14-plus.json @@ -0,0 +1,25 @@ +{ + "embed_dim": 1024, + "vision_cfg": { + "image_size": 224, + "layers": 64, + "width": 1792, + "head_width": 112, + "mlp_ratio": 8.571428571428571, + "patch_size": 14, + "eva_model_name": "eva-clip-4b-14-x", + "drop_path_rate": 0, + "xattn": true, + "postnorm": true, + "fusedLN": true + }, + "text_cfg": { + "context_length": 77, + "vocab_size": 49408, + "width": 1280, + "heads": 20, + "layers": 32, + "xattn": false, + "fusedLN": true + } +} diff --git a/modules/pulid/eva_clip/model_configs/EVA02-CLIP-bigE-14.json b/modules/pulid/eva_clip/model_configs/EVA02-CLIP-bigE-14.json new file mode 100644 index 000000000..747ffccc8 --- /dev/null +++ b/modules/pulid/eva_clip/model_configs/EVA02-CLIP-bigE-14.json @@ -0,0 +1,25 @@ +{ + "embed_dim": 1024, + "vision_cfg": { + "image_size": 224, + "layers": 64, + "width": 1792, + "head_width": 112, + "mlp_ratio": 8.571428571428571, + "patch_size": 14, + "eva_model_name": "eva-clip-4b-14-x", + "drop_path_rate": 0, + "xattn": true, + "postnorm": true, + "fusedLN": true + }, + "text_cfg": { + "context_length": 77, + "vocab_size": 49408, + "width": 1024, + "heads": 16, + "layers": 24, + "xattn": false, + "fusedLN": true + } +} \ No newline at end of file diff --git a/modules/pulid/eva_clip/modified_resnet.py b/modules/pulid/eva_clip/modified_resnet.py new file mode 100644 index 000000000..151bfdd0b --- /dev/null +++ b/modules/pulid/eva_clip/modified_resnet.py @@ -0,0 +1,181 @@ +from collections import OrderedDict + +import torch +from torch import nn +from torch.nn import functional as F + +from eva_clip.utils import freeze_batch_norm_2d + + +class Bottleneck(nn.Module): + expansion = 4 + + def __init__(self, inplanes, planes, stride=1): + super().__init__() + + # all conv layers have stride 1. an avgpool is performed after the second convolution when stride > 1 + self.conv1 = nn.Conv2d(inplanes, planes, 1, bias=False) + self.bn1 = nn.BatchNorm2d(planes) + self.act1 = nn.ReLU(inplace=True) + + self.conv2 = nn.Conv2d(planes, planes, 3, padding=1, bias=False) + self.bn2 = nn.BatchNorm2d(planes) + self.act2 = nn.ReLU(inplace=True) + + self.avgpool = nn.AvgPool2d(stride) if stride > 1 else nn.Identity() + + self.conv3 = nn.Conv2d(planes, planes * self.expansion, 1, bias=False) + self.bn3 = nn.BatchNorm2d(planes * self.expansion) + self.act3 = nn.ReLU(inplace=True) + + self.downsample = None + self.stride = stride + + if stride > 1 or inplanes != planes * Bottleneck.expansion: + # downsampling layer is prepended with an avgpool, and the subsequent convolution has stride 1 + self.downsample = nn.Sequential(OrderedDict([ + ("-1", nn.AvgPool2d(stride)), + ("0", nn.Conv2d(inplanes, planes * self.expansion, 1, stride=1, bias=False)), + ("1", nn.BatchNorm2d(planes * self.expansion)) + ])) + + def forward(self, x: torch.Tensor): + identity = x + + out = self.act1(self.bn1(self.conv1(x))) + out = self.act2(self.bn2(self.conv2(out))) + out = self.avgpool(out) + out = self.bn3(self.conv3(out)) + + if self.downsample is not None: + identity = self.downsample(x) + + out += identity + out = self.act3(out) + return out + + +class AttentionPool2d(nn.Module): + def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: int = None): + super().__init__() + self.positional_embedding = nn.Parameter(torch.randn(spacial_dim ** 2 + 1, embed_dim) / embed_dim ** 0.5) + self.k_proj = nn.Linear(embed_dim, embed_dim) + self.q_proj = nn.Linear(embed_dim, embed_dim) + self.v_proj = nn.Linear(embed_dim, embed_dim) + self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim) + self.num_heads = num_heads + + def forward(self, x): + x = x.reshape(x.shape[0], x.shape[1], x.shape[2] * x.shape[3]).permute(2, 0, 1) # NCHW -> (HW)NC + x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (HW+1)NC + x = x + self.positional_embedding[:, None, :].to(x.dtype) # (HW+1)NC + x, _ = F.multi_head_attention_forward( + query=x, key=x, value=x, + embed_dim_to_check=x.shape[-1], + num_heads=self.num_heads, + q_proj_weight=self.q_proj.weight, + k_proj_weight=self.k_proj.weight, + v_proj_weight=self.v_proj.weight, + in_proj_weight=None, + in_proj_bias=torch.cat([self.q_proj.bias, self.k_proj.bias, self.v_proj.bias]), + bias_k=None, + bias_v=None, + add_zero_attn=False, + dropout_p=0., + out_proj_weight=self.c_proj.weight, + out_proj_bias=self.c_proj.bias, + use_separate_proj_weight=True, + training=self.training, + need_weights=False + ) + + return x[0] + + +class ModifiedResNet(nn.Module): + """ + A ResNet class that is similar to torchvision's but contains the following changes: + - There are now 3 "stem" convolutions as opposed to 1, with an average pool instead of a max pool. + - Performs anti-aliasing strided convolutions, where an avgpool is prepended to convolutions with stride > 1 + - The final pooling layer is a QKV attention instead of an average pool + """ + + def __init__(self, layers, output_dim, heads, image_size=224, width=64): + super().__init__() + self.output_dim = output_dim + self.image_size = image_size + + # the 3-layer stem + self.conv1 = nn.Conv2d(3, width // 2, kernel_size=3, stride=2, padding=1, bias=False) + self.bn1 = nn.BatchNorm2d(width // 2) + self.act1 = nn.ReLU(inplace=True) + self.conv2 = nn.Conv2d(width // 2, width // 2, kernel_size=3, padding=1, bias=False) + self.bn2 = nn.BatchNorm2d(width // 2) + self.act2 = nn.ReLU(inplace=True) + self.conv3 = nn.Conv2d(width // 2, width, kernel_size=3, padding=1, bias=False) + self.bn3 = nn.BatchNorm2d(width) + self.act3 = nn.ReLU(inplace=True) + self.avgpool = nn.AvgPool2d(2) + + # residual layers + self._inplanes = width # this is a *mutable* variable used during construction + self.layer1 = self._make_layer(width, layers[0]) + self.layer2 = self._make_layer(width * 2, layers[1], stride=2) + self.layer3 = self._make_layer(width * 4, layers[2], stride=2) + self.layer4 = self._make_layer(width * 8, layers[3], stride=2) + + embed_dim = width * 32 # the ResNet feature dimension + self.attnpool = AttentionPool2d(image_size // 32, embed_dim, heads, output_dim) + + self.init_parameters() + + def _make_layer(self, planes, blocks, stride=1): + layers = [Bottleneck(self._inplanes, planes, stride)] + + self._inplanes = planes * Bottleneck.expansion + for _ in range(1, blocks): + layers.append(Bottleneck(self._inplanes, planes)) + + return nn.Sequential(*layers) + + def init_parameters(self): + if self.attnpool is not None: + std = self.attnpool.c_proj.in_features ** -0.5 + nn.init.normal_(self.attnpool.q_proj.weight, std=std) + nn.init.normal_(self.attnpool.k_proj.weight, std=std) + nn.init.normal_(self.attnpool.v_proj.weight, std=std) + nn.init.normal_(self.attnpool.c_proj.weight, std=std) + + for resnet_block in [self.layer1, self.layer2, self.layer3, self.layer4]: + for name, param in resnet_block.named_parameters(): + if name.endswith("bn3.weight"): + nn.init.zeros_(param) + + def lock(self, unlocked_groups=0, freeze_bn_stats=False): + assert unlocked_groups == 0, 'partial locking not currently supported for this model' + for param in self.parameters(): + param.requires_grad = False + if freeze_bn_stats: + freeze_batch_norm_2d(self) + + @torch.jit.ignore + def set_grad_checkpointing(self, enable=True): + # FIXME support for non-transformer + pass + + def stem(self, x): + x = self.act1(self.bn1(self.conv1(x))) + x = self.act2(self.bn2(self.conv2(x))) + x = self.act3(self.bn3(self.conv3(x))) + x = self.avgpool(x) + return x + + def forward(self, x): + x = self.stem(x) + x = self.layer1(x) + x = self.layer2(x) + x = self.layer3(x) + x = self.layer4(x) + x = self.attnpool(x) + + return x diff --git a/modules/pulid/eva_clip/openai.py b/modules/pulid/eva_clip/openai.py new file mode 100644 index 000000000..cc4e13e87 --- /dev/null +++ b/modules/pulid/eva_clip/openai.py @@ -0,0 +1,144 @@ +""" OpenAI pretrained model functions + +Adapted from https://github.com/openai/CLIP. Originally MIT License, Copyright (c) 2021 OpenAI. +""" + +import os +import warnings +from typing import List, Optional, Union + +import torch + +from .model import build_model_from_openai_state_dict, convert_weights_to_lp, get_cast_dtype +from .pretrained import get_pretrained_url, list_pretrained_models_by_tag, download_pretrained_from_url + +__all__ = ["list_openai_models", "load_openai_model"] + + +def list_openai_models() -> List[str]: + """Returns the names of available CLIP models""" + return list_pretrained_models_by_tag('openai') + + +def load_openai_model( + name: str, + precision: Optional[str] = None, + device: Optional[Union[str, torch.device]] = None, + jit: bool = True, + cache_dir: Optional[str] = None, +): + """Load a CLIP model + + Parameters + ---------- + name : str + A model name listed by `clip.available_models()`, or the path to a model checkpoint containing the state_dict + precision: str + Model precision, if None defaults to 'fp32' if device == 'cpu' else 'fp16'. + device : Union[str, torch.device] + The device to put the loaded model + jit : bool + Whether to load the optimized JIT model (default) or more hackable non-JIT model. + cache_dir : Optional[str] + The directory to cache the downloaded model weights + + Returns + ------- + model : torch.nn.Module + The CLIP model + preprocess : Callable[[PIL.Image], torch.Tensor] + A torchvision transform that converts a PIL image into a tensor that the returned model can take as its input + """ + if device is None: + device = "cuda" if torch.cuda.is_available() else "cpu" + if precision is None: + precision = 'fp32' if device == 'cpu' else 'fp16' + + if get_pretrained_url(name, 'openai'): + model_path = download_pretrained_from_url(get_pretrained_url(name, 'openai'), cache_dir=cache_dir) + elif os.path.isfile(name): + model_path = name + else: + raise RuntimeError(f"Model {name} not found; available models = {list_openai_models()}") + + try: + # loading JIT archive + model = torch.jit.load(model_path, map_location=device if jit else "cpu").eval() + state_dict = None + except RuntimeError: + # loading saved state dict + if jit: + warnings.warn(f"File {model_path} is not a JIT archive. Loading as a state dict instead") + jit = False + state_dict = torch.load(model_path, map_location="cpu") + + if not jit: + # Build a non-jit model from the OpenAI jitted model state dict + cast_dtype = get_cast_dtype(precision) + try: + model = build_model_from_openai_state_dict(state_dict or model.state_dict(), cast_dtype=cast_dtype) + except KeyError: + sd = {k[7:]: v for k, v in state_dict["state_dict"].items()} + model = build_model_from_openai_state_dict(sd, cast_dtype=cast_dtype) + + # model from OpenAI state dict is in manually cast fp16 mode, must be converted for AMP/fp32/bf16 use + model = model.to(device) + if precision.startswith('amp') or precision == 'fp32': + model.float() + elif precision == 'bf16': + convert_weights_to_lp(model, dtype=torch.bfloat16) + + return model + + # patch the device names + device_holder = torch.jit.trace(lambda: torch.ones([]).to(torch.device(device)), example_inputs=[]) + device_node = [n for n in device_holder.graph.findAllNodes("prim::Constant") if "Device" in repr(n)][-1] + + def patch_device(module): + try: + graphs = [module.graph] if hasattr(module, "graph") else [] + except RuntimeError: + graphs = [] + + if hasattr(module, "forward1"): + graphs.append(module.forward1.graph) + + for graph in graphs: + for node in graph.findAllNodes("prim::Constant"): + if "value" in node.attributeNames() and str(node["value"]).startswith("cuda"): + node.copyAttributes(device_node) + + model.apply(patch_device) + patch_device(model.encode_image) + patch_device(model.encode_text) + + # patch dtype to float32 (typically for CPU) + if precision == 'fp32': + float_holder = torch.jit.trace(lambda: torch.ones([]).float(), example_inputs=[]) + float_input = list(float_holder.graph.findNode("aten::to").inputs())[1] + float_node = float_input.node() + + def patch_float(module): + try: + graphs = [module.graph] if hasattr(module, "graph") else [] + except RuntimeError: + graphs = [] + + if hasattr(module, "forward1"): + graphs.append(module.forward1.graph) + + for graph in graphs: + for node in graph.findAllNodes("aten::to"): + inputs = list(node.inputs()) + for i in [1, 2]: # dtype can be the second or third argument to aten::to() + if inputs[i].node()["value"] == 5: + inputs[i].node().copyAttributes(float_node) + + model.apply(patch_float) + patch_float(model.encode_image) + patch_float(model.encode_text) + model.float() + + # ensure image_size attr available at consistent location for both jit and non-jit + model.visual.image_size = model.input_resolution.item() + return model diff --git a/modules/pulid/eva_clip/pretrained.py b/modules/pulid/eva_clip/pretrained.py new file mode 100644 index 000000000..a1e55dcf3 --- /dev/null +++ b/modules/pulid/eva_clip/pretrained.py @@ -0,0 +1,332 @@ +import hashlib +import os +import urllib +import warnings +from functools import partial +from typing import Dict, Union + +from tqdm import tqdm + +try: + from huggingface_hub import hf_hub_download + _has_hf_hub = True +except ImportError: + hf_hub_download = None + _has_hf_hub = False + + +def _pcfg(url='', hf_hub='', filename='', mean=None, std=None): + return dict( + url=url, + hf_hub=hf_hub, + mean=mean, + std=std, + ) + +_VITB32 = dict( + openai=_pcfg( + "https://openaipublic.azureedge.net/clip/models/40d365715913c9da98579312b702a82c18be219cc2a73407c4526f58eba950af/ViT-B-32.pt"), + laion400m_e31=_pcfg( + "https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_32-quickgelu-laion400m_e31-d867053b.pt"), + laion400m_e32=_pcfg( + "https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_32-quickgelu-laion400m_e32-46683a32.pt"), + laion2b_e16=_pcfg( + "https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_32-laion2b_e16-af8dbd0c.pth"), + laion2b_s34b_b79k=_pcfg(hf_hub='laion/CLIP-ViT-B-32-laion2B-s34B-b79K/') +) + +_VITB32_quickgelu = dict( + openai=_pcfg( + "https://openaipublic.azureedge.net/clip/models/40d365715913c9da98579312b702a82c18be219cc2a73407c4526f58eba950af/ViT-B-32.pt"), + laion400m_e31=_pcfg( + "https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_32-quickgelu-laion400m_e31-d867053b.pt"), + laion400m_e32=_pcfg( + "https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_32-quickgelu-laion400m_e32-46683a32.pt"), +) + +_VITB16 = dict( + openai=_pcfg( + "https://openaipublic.azureedge.net/clip/models/5806e77cd80f8b59890b7e101eabd078d9fb84e6937f9e85e4ecb61988df416f/ViT-B-16.pt"), + laion400m_e31=_pcfg( + "https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_16-laion400m_e31-00efa78f.pt"), + laion400m_e32=_pcfg( + "https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_16-laion400m_e32-55e67d44.pt"), + laion2b_s34b_b88k=_pcfg(hf_hub='laion/CLIP-ViT-B-16-laion2B-s34B-b88K/'), +) + +_EVAB16 = dict( + eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_B_psz14to16.pt'), + eva02=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_B_psz14to16.pt'), + eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_B_psz16_s8B.pt'), + eva02_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_B_psz16_s8B.pt'), +) + +_VITB16_PLUS_240 = dict( + laion400m_e31=_pcfg( + "https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_16_plus_240-laion400m_e31-8fb26589.pt"), + laion400m_e32=_pcfg( + "https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_b_16_plus_240-laion400m_e32-699c4b84.pt"), +) + +_VITL14 = dict( + openai=_pcfg( + "https://openaipublic.azureedge.net/clip/models/b8cca3fd41ae0c99ba7e8951adf17d267cdb84cd88be6f7c2e0eca1737a03836/ViT-L-14.pt"), + laion400m_e31=_pcfg( + "https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_l_14-laion400m_e31-69988bb6.pt"), + laion400m_e32=_pcfg( + "https://github.com/mlfoundations/open_clip/releases/download/v0.2-weights/vit_l_14-laion400m_e32-3d133497.pt"), + laion2b_s32b_b82k=_pcfg( + hf_hub='laion/CLIP-ViT-L-14-laion2B-s32B-b82K/', + mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)), +) + +_EVAL14 = dict( + eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_L_psz14.pt'), + eva02=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_L_psz14.pt'), + eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_psz14_s4B.pt'), + eva02_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_psz14_s4B.pt'), +) + +_VITL14_336 = dict( + openai=_pcfg( + "https://openaipublic.azureedge.net/clip/models/3035c92b350959924f9f00213499208652fc7ea050643e8b385c2dac08641f02/ViT-L-14-336px.pt"), +) + +_EVAL14_336 = dict( + eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_336_psz14_s6B.pt'), + eva02_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_336_psz14_s6B.pt'), + eva_clip_224to336=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_psz14_224to336.pt'), + eva02_clip_224to336=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_L_psz14_224to336.pt'), +) + +_VITH14 = dict( + laion2b_s32b_b79k=_pcfg(hf_hub='laion/CLIP-ViT-H-14-laion2B-s32B-b79K/'), +) + +_VITg14 = dict( + laion2b_s12b_b42k=_pcfg(hf_hub='laion/CLIP-ViT-g-14-laion2B-s12B-b42K/'), + laion2b_s34b_b88k=_pcfg(hf_hub='laion/CLIP-ViT-g-14-laion2B-s34B-b88K/'), +) + +_EVAg14 = dict( + eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/'), + eva01=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_g_psz14.pt'), + eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_CLIP_g_14_psz14_s11B.pt'), + eva01_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_CLIP_g_14_psz14_s11B.pt'), +) + +_EVAg14_PLUS = dict( + eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/'), + eva01=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_g_psz14.pt'), + eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_CLIP_g_14_plus_psz14_s11B.pt'), + eva01_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA01_CLIP_g_14_plus_psz14_s11B.pt'), +) + +_VITbigG14 = dict( + laion2b_s39b_b160k=_pcfg(hf_hub='laion/CLIP-ViT-bigG-14-laion2B-39B-b160k/'), +) + +_EVAbigE14 = dict( + eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_E_psz14.pt'), + eva02=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_E_psz14.pt'), + eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_E_psz14_s4B.pt'), + eva02_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_E_psz14_s4B.pt'), +) + +_EVAbigE14_PLUS = dict( + eva=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_E_psz14.pt'), + eva02=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_E_psz14.pt'), + eva_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_E_psz14_plus_s9B.pt'), + eva02_clip=_pcfg(hf_hub='QuanSun/EVA-CLIP/EVA02_CLIP_E_psz14_plus_s9B.pt'), +) + + +_PRETRAINED = { + # "ViT-B-32": _VITB32, + "OpenaiCLIP-B-32": _VITB32, + "OpenCLIP-B-32": _VITB32, + + # "ViT-B-32-quickgelu": _VITB32_quickgelu, + "OpenaiCLIP-B-32-quickgelu": _VITB32_quickgelu, + "OpenCLIP-B-32-quickgelu": _VITB32_quickgelu, + + # "ViT-B-16": _VITB16, + "OpenaiCLIP-B-16": _VITB16, + "OpenCLIP-B-16": _VITB16, + + "EVA02-B-16": _EVAB16, + "EVA02-CLIP-B-16": _EVAB16, + + # "ViT-B-16-plus-240": _VITB16_PLUS_240, + "OpenCLIP-B-16-plus-240": _VITB16_PLUS_240, + + # "ViT-L-14": _VITL14, + "OpenaiCLIP-L-14": _VITL14, + "OpenCLIP-L-14": _VITL14, + + "EVA02-L-14": _EVAL14, + "EVA02-CLIP-L-14": _EVAL14, + + # "ViT-L-14-336": _VITL14_336, + "OpenaiCLIP-L-14-336": _VITL14_336, + + "EVA02-CLIP-L-14-336": _EVAL14_336, + + # "ViT-H-14": _VITH14, + # "ViT-g-14": _VITg14, + "OpenCLIP-H-14": _VITH14, + "OpenCLIP-g-14": _VITg14, + + "EVA01-CLIP-g-14": _EVAg14, + "EVA01-CLIP-g-14-plus": _EVAg14_PLUS, + + # "ViT-bigG-14": _VITbigG14, + "OpenCLIP-bigG-14": _VITbigG14, + + "EVA02-CLIP-bigE-14": _EVAbigE14, + "EVA02-CLIP-bigE-14-plus": _EVAbigE14_PLUS, +} + + +def _clean_tag(tag: str): + # normalize pretrained tags + return tag.lower().replace('-', '_') + + +def list_pretrained(as_str: bool = False): + """ returns list of pretrained models + Returns a tuple (model_name, pretrain_tag) by default or 'name:tag' if as_str == True + """ + return [':'.join([k, t]) if as_str else (k, t) for k in _PRETRAINED.keys() for t in _PRETRAINED[k].keys()] + + +def list_pretrained_models_by_tag(tag: str): + """ return all models having the specified pretrain tag """ + models = [] + tag = _clean_tag(tag) + for k in _PRETRAINED.keys(): + if tag in _PRETRAINED[k]: + models.append(k) + return models + + +def list_pretrained_tags_by_model(model: str): + """ return all pretrain tags for the specified model architecture """ + tags = [] + if model in _PRETRAINED: + tags.extend(_PRETRAINED[model].keys()) + return tags + + +def is_pretrained_cfg(model: str, tag: str): + if model not in _PRETRAINED: + return False + return _clean_tag(tag) in _PRETRAINED[model] + + +def get_pretrained_cfg(model: str, tag: str): + if model not in _PRETRAINED: + return {} + model_pretrained = _PRETRAINED[model] + return model_pretrained.get(_clean_tag(tag), {}) + + +def get_pretrained_url(model: str, tag: str): + cfg = get_pretrained_cfg(model, _clean_tag(tag)) + return cfg.get('url', '') + + +def download_pretrained_from_url( + url: str, + cache_dir: Union[str, None] = None, +): + if not cache_dir: + cache_dir = os.path.expanduser("~/.cache/clip") + os.makedirs(cache_dir, exist_ok=True) + filename = os.path.basename(url) + + if 'openaipublic' in url: + expected_sha256 = url.split("/")[-2] + elif 'mlfoundations' in url: + expected_sha256 = os.path.splitext(filename)[0].split("-")[-1] + else: + expected_sha256 = '' + + download_target = os.path.join(cache_dir, filename) + + if os.path.exists(download_target) and not os.path.isfile(download_target): + raise RuntimeError(f"{download_target} exists and is not a regular file") + + if os.path.isfile(download_target): + if expected_sha256: + if hashlib.sha256(open(download_target, "rb").read()).hexdigest().startswith(expected_sha256): + return download_target + else: + warnings.warn(f"{download_target} exists, but the SHA256 checksum does not match; re-downloading the file") + else: + return download_target + + with urllib.request.urlopen(url) as source, open(download_target, "wb") as output: + with tqdm(total=int(source.headers.get("Content-Length")), ncols=80, unit='iB', unit_scale=True) as loop: + while True: + buffer = source.read(8192) + if not buffer: + break + + output.write(buffer) + loop.update(len(buffer)) + + if expected_sha256 and not hashlib.sha256(open(download_target, "rb").read()).hexdigest().startswith(expected_sha256): + raise RuntimeError(f"Model has been downloaded but the SHA256 checksum does not not match") + + return download_target + + +def has_hf_hub(necessary=False): + if not _has_hf_hub and necessary: + # if no HF Hub module installed, and it is necessary to continue, raise error + raise RuntimeError( + 'Hugging Face hub model specified but package not installed. Run `pip install huggingface_hub`.') + return _has_hf_hub + + +def download_pretrained_from_hf( + model_id: str, + filename: str = 'open_clip_pytorch_model.bin', + revision=None, + cache_dir: Union[str, None] = None, +): + has_hf_hub(True) + cached_file = hf_hub_download(model_id, filename, revision=revision, cache_dir=cache_dir) + return cached_file + + +def download_pretrained( + cfg: Dict, + force_hf_hub: bool = False, + cache_dir: Union[str, None] = None, +): + target = '' + if not cfg: + return target + + download_url = cfg.get('url', '') + download_hf_hub = cfg.get('hf_hub', '') + if download_hf_hub and force_hf_hub: + # use HF hub even if url exists + download_url = '' + + if download_url: + target = download_pretrained_from_url(download_url, cache_dir=cache_dir) + elif download_hf_hub: + has_hf_hub(True) + # we assume the hf_hub entries in pretrained config combine model_id + filename in + # 'org/model_name/filename.pt' form. To specify just the model id w/o filename and + # use 'open_clip_pytorch_model.bin' default, there must be a trailing slash 'org/model_name/'. + model_id, filename = os.path.split(download_hf_hub) + if filename: + target = download_pretrained_from_hf(model_id, filename=filename, cache_dir=cache_dir) + else: + target = download_pretrained_from_hf(model_id, cache_dir=cache_dir) + + return target diff --git a/modules/pulid/eva_clip/rope.py b/modules/pulid/eva_clip/rope.py new file mode 100644 index 000000000..69030c35e --- /dev/null +++ b/modules/pulid/eva_clip/rope.py @@ -0,0 +1,137 @@ +from math import pi +import torch +from torch import nn +from einops import rearrange, repeat +import logging + +def broadcat(tensors, dim = -1): + num_tensors = len(tensors) + shape_lens = set(list(map(lambda t: len(t.shape), tensors))) + assert len(shape_lens) == 1, 'tensors must all have the same number of dimensions' + shape_len = list(shape_lens)[0] + dim = (dim + shape_len) if dim < 0 else dim + dims = list(zip(*map(lambda t: list(t.shape), tensors))) + expandable_dims = [(i, val) for i, val in enumerate(dims) if i != dim] + assert all([*map(lambda t: len(set(t[1])) <= 2, expandable_dims)]), 'invalid dimensions for broadcastable concatentation' + max_dims = list(map(lambda t: (t[0], max(t[1])), expandable_dims)) + expanded_dims = list(map(lambda t: (t[0], (t[1],) * num_tensors), max_dims)) + expanded_dims.insert(dim, (dim, dims[dim])) + expandable_shapes = list(zip(*map(lambda t: t[1], expanded_dims))) + tensors = list(map(lambda t: t[0].expand(*t[1]), zip(tensors, expandable_shapes))) + return torch.cat(tensors, dim = dim) + +def rotate_half(x): + x = rearrange(x, '... (d r) -> ... d r', r = 2) + x1, x2 = x.unbind(dim = -1) + x = torch.stack((-x2, x1), dim = -1) + return rearrange(x, '... d r -> ... (d r)') + + +class VisionRotaryEmbedding(nn.Module): + def __init__( + self, + dim, + pt_seq_len, + ft_seq_len=None, + custom_freqs = None, + freqs_for = 'lang', + theta = 10000, + max_freq = 10, + num_freqs = 1, + ): + super().__init__() + if custom_freqs: + freqs = custom_freqs + elif freqs_for == 'lang': + freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim)) + elif freqs_for == 'pixel': + freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi + elif freqs_for == 'constant': + freqs = torch.ones(num_freqs).float() + else: + raise ValueError(f'unknown modality {freqs_for}') + + if ft_seq_len is None: ft_seq_len = pt_seq_len + t = torch.arange(ft_seq_len) / ft_seq_len * pt_seq_len + + freqs_h = torch.einsum('..., f -> ... f', t, freqs) + freqs_h = repeat(freqs_h, '... n -> ... (n r)', r = 2) + + freqs_w = torch.einsum('..., f -> ... f', t, freqs) + freqs_w = repeat(freqs_w, '... n -> ... (n r)', r = 2) + + freqs = broadcat((freqs_h[:, None, :], freqs_w[None, :, :]), dim = -1) + + self.register_buffer("freqs_cos", freqs.cos()) + self.register_buffer("freqs_sin", freqs.sin()) + + logging.info(f'Shape of rope freq: {self.freqs_cos.shape}') + + def forward(self, t, start_index = 0): + rot_dim = self.freqs_cos.shape[-1] + end_index = start_index + rot_dim + assert rot_dim <= t.shape[-1], f'feature dimension {t.shape[-1]} is not of sufficient size to rotate in all the positions {rot_dim}' + t_left, t, t_right = t[..., :start_index], t[..., start_index:end_index], t[..., end_index:] + t = (t * self.freqs_cos) + (rotate_half(t) * self.freqs_sin) + + return torch.cat((t_left, t, t_right), dim = -1) + +class VisionRotaryEmbeddingFast(nn.Module): + def __init__( + self, + dim, + pt_seq_len, + ft_seq_len=None, + custom_freqs = None, + freqs_for = 'lang', + theta = 10000, + max_freq = 10, + num_freqs = 1, + patch_dropout = 0. + ): + super().__init__() + if custom_freqs: + freqs = custom_freqs + elif freqs_for == 'lang': + freqs = 1. / (theta ** (torch.arange(0, dim, 2)[:(dim // 2)].float() / dim)) + elif freqs_for == 'pixel': + freqs = torch.linspace(1., max_freq / 2, dim // 2) * pi + elif freqs_for == 'constant': + freqs = torch.ones(num_freqs).float() + else: + raise ValueError(f'unknown modality {freqs_for}') + + if ft_seq_len is None: ft_seq_len = pt_seq_len + t = torch.arange(ft_seq_len) / ft_seq_len * pt_seq_len + + freqs = torch.einsum('..., f -> ... f', t, freqs) + freqs = repeat(freqs, '... n -> ... (n r)', r = 2) + freqs = broadcat((freqs[:, None, :], freqs[None, :, :]), dim = -1) + + freqs_cos = freqs.cos().view(-1, freqs.shape[-1]) + freqs_sin = freqs.sin().view(-1, freqs.shape[-1]) + + self.patch_dropout = patch_dropout + + self.register_buffer("freqs_cos", freqs_cos) + self.register_buffer("freqs_sin", freqs_sin) + + logging.info(f'Shape of rope freq: {self.freqs_cos.shape}') + + def forward(self, t, patch_indices_keep=None): + if patch_indices_keep is not None: + batch = t.size()[0] + batch_indices = torch.arange(batch) + batch_indices = batch_indices[..., None] + + freqs_cos = repeat(self.freqs_cos, 'i j -> n i m j', n=t.shape[0], m=t.shape[1]) + freqs_sin = repeat(self.freqs_sin, 'i j -> n i m j', n=t.shape[0], m=t.shape[1]) + + freqs_cos = freqs_cos[batch_indices, patch_indices_keep] + freqs_cos = rearrange(freqs_cos, 'n i m j -> n m i j') + freqs_sin = freqs_sin[batch_indices, patch_indices_keep] + freqs_sin = rearrange(freqs_sin, 'n i m j -> n m i j') + + return t * freqs_cos + rotate_half(t) * freqs_sin + + return t * self.freqs_cos + rotate_half(t) * self.freqs_sin \ No newline at end of file diff --git a/modules/pulid/eva_clip/timm_model.py b/modules/pulid/eva_clip/timm_model.py new file mode 100644 index 000000000..53bc4d469 --- /dev/null +++ b/modules/pulid/eva_clip/timm_model.py @@ -0,0 +1,119 @@ +""" timm model adapter + +Wraps timm (https://github.com/rwightman/pytorch-image-models) models for use as a vision tower in CLIP model. +""" +import logging +from collections import OrderedDict + +import torch +import torch.nn as nn + +try: + import timm + from timm.models.layers import Mlp, to_2tuple + try: + # old timm imports < 0.8.1 + from timm.models.layers.attention_pool2d import RotAttentionPool2d + from timm.models.layers.attention_pool2d import AttentionPool2d as AbsAttentionPool2d + except ImportError: + # new timm imports >= 0.8.1 + from timm.layers import RotAttentionPool2d + from timm.layers import AttentionPool2d as AbsAttentionPool2d +except ImportError: + timm = None + +from .utils import freeze_batch_norm_2d + + +class TimmModel(nn.Module): + """ timm model adapter + # FIXME this adapter is a work in progress, may change in ways that break weight compat + """ + + def __init__( + self, + model_name, + embed_dim, + image_size=224, + pool='avg', + proj='linear', + proj_bias=False, + drop=0., + pretrained=False): + super().__init__() + + self.image_size = to_2tuple(image_size) + self.trunk = timm.create_model(model_name, pretrained=pretrained) + feat_size = self.trunk.default_cfg.get('pool_size', None) + feature_ndim = 1 if not feat_size else 2 + if pool in ('abs_attn', 'rot_attn'): + assert feature_ndim == 2 + # if attn pooling used, remove both classifier and default pool + self.trunk.reset_classifier(0, global_pool='') + else: + # reset global pool if pool config set, otherwise leave as network default + reset_kwargs = dict(global_pool=pool) if pool else {} + self.trunk.reset_classifier(0, **reset_kwargs) + prev_chs = self.trunk.num_features + + head_layers = OrderedDict() + if pool == 'abs_attn': + head_layers['pool'] = AbsAttentionPool2d(prev_chs, feat_size=feat_size, out_features=embed_dim) + prev_chs = embed_dim + elif pool == 'rot_attn': + head_layers['pool'] = RotAttentionPool2d(prev_chs, out_features=embed_dim) + prev_chs = embed_dim + else: + assert proj, 'projection layer needed if non-attention pooling is used.' + + # NOTE attention pool ends with a projection layer, so proj should usually be set to '' if such pooling is used + if proj == 'linear': + head_layers['drop'] = nn.Dropout(drop) + head_layers['proj'] = nn.Linear(prev_chs, embed_dim, bias=proj_bias) + elif proj == 'mlp': + head_layers['mlp'] = Mlp(prev_chs, 2 * embed_dim, embed_dim, drop=drop, bias=(True, proj_bias)) + + self.head = nn.Sequential(head_layers) + + def lock(self, unlocked_groups=0, freeze_bn_stats=False): + """ lock modules + Args: + unlocked_groups (int): leave last n layer groups unlocked (default: 0) + """ + if not unlocked_groups: + # lock full model + for param in self.trunk.parameters(): + param.requires_grad = False + if freeze_bn_stats: + freeze_batch_norm_2d(self.trunk) + else: + # NOTE: partial freeze requires latest timm (master) branch and is subject to change + try: + # FIXME import here until API stable and in an official release + from timm.models.helpers import group_parameters, group_modules + except ImportError: + raise RuntimeError('Please install latest timm `pip install git+https://github.com/rwightman/pytorch-image-models`') + matcher = self.trunk.group_matcher() + gparams = group_parameters(self.trunk, matcher) + max_layer_id = max(gparams.keys()) + max_layer_id = max_layer_id - unlocked_groups + for group_idx in range(max_layer_id + 1): + group = gparams[group_idx] + for param in group: + self.trunk.get_parameter(param).requires_grad = False + if freeze_bn_stats: + gmodules = group_modules(self.trunk, matcher, reverse=True) + gmodules = {k for k, v in gmodules.items() if v <= max_layer_id} + freeze_batch_norm_2d(self.trunk, gmodules) + + @torch.jit.ignore + def set_grad_checkpointing(self, enable=True): + try: + self.trunk.set_grad_checkpointing(enable) + except Exception as e: + logging.warning('grad checkpointing not supported for this timm image tower, continuing without...') + + def forward(self, x): + x = self.trunk(x) + x = self.head(x) + return x diff --git a/modules/pulid/eva_clip/tokenizer.py b/modules/pulid/eva_clip/tokenizer.py new file mode 100644 index 000000000..41482f82a --- /dev/null +++ b/modules/pulid/eva_clip/tokenizer.py @@ -0,0 +1,201 @@ +""" CLIP tokenizer + +Copied from https://github.com/openai/CLIP. Originally MIT License, Copyright (c) 2021 OpenAI. +""" +import gzip +import html +import os +from functools import lru_cache +from typing import Union, List + +import ftfy +import regex as re +import torch + +# https://stackoverflow.com/q/62691279 +import os +os.environ["TOKENIZERS_PARALLELISM"] = "false" + + +@lru_cache() +def default_bpe(): + return os.path.join(os.path.dirname(os.path.abspath(__file__)), "bpe_simple_vocab_16e6.txt.gz") + + +@lru_cache() +def bytes_to_unicode(): + """ + Returns list of utf-8 byte and a corresponding list of unicode strings. + The reversible bpe codes work on unicode strings. + This means you need a large # of unicode characters in your vocab if you want to avoid UNKs. + When you're at something like a 10B token dataset you end up needing around 5K for decent coverage. + This is a signficant percentage of your normal, say, 32K bpe vocab. + To avoid that, we want lookup tables between utf-8 bytes and unicode strings. + And avoids mapping to whitespace/control characters the bpe code barfs on. + """ + bs = list(range(ord("!"), ord("~")+1))+list(range(ord("¡"), ord("¬")+1))+list(range(ord("®"), ord("ÿ")+1)) + cs = bs[:] + n = 0 + for b in range(2**8): + if b not in bs: + bs.append(b) + cs.append(2**8+n) + n += 1 + cs = [chr(n) for n in cs] + return dict(zip(bs, cs)) + + +def get_pairs(word): + """Return set of symbol pairs in a word. + Word is represented as tuple of symbols (symbols being variable-length strings). + """ + pairs = set() + prev_char = word[0] + for char in word[1:]: + pairs.add((prev_char, char)) + prev_char = char + return pairs + + +def basic_clean(text): + text = ftfy.fix_text(text) + text = html.unescape(html.unescape(text)) + return text.strip() + + +def whitespace_clean(text): + text = re.sub(r'\s+', ' ', text) + text = text.strip() + return text + + +class SimpleTokenizer(object): + def __init__(self, bpe_path: str = default_bpe(), special_tokens=None): + self.byte_encoder = bytes_to_unicode() + self.byte_decoder = {v: k for k, v in self.byte_encoder.items()} + merges = gzip.open(bpe_path).read().decode("utf-8").split('\n') + merges = merges[1:49152-256-2+1] + merges = [tuple(merge.split()) for merge in merges] + vocab = list(bytes_to_unicode().values()) + vocab = vocab + [v+'' for v in vocab] + for merge in merges: + vocab.append(''.join(merge)) + if not special_tokens: + special_tokens = ['', ''] + else: + special_tokens = ['', ''] + special_tokens + vocab.extend(special_tokens) + self.encoder = dict(zip(vocab, range(len(vocab)))) + self.decoder = {v: k for k, v in self.encoder.items()} + self.bpe_ranks = dict(zip(merges, range(len(merges)))) + self.cache = {t:t for t in special_tokens} + special = "|".join(special_tokens) + self.pat = re.compile(special + r"""|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+""", re.IGNORECASE) + + self.vocab_size = len(self.encoder) + self.all_special_ids = [self.encoder[t] for t in special_tokens] + + def bpe(self, token): + if token in self.cache: + return self.cache[token] + word = tuple(token[:-1]) + ( token[-1] + '',) + pairs = get_pairs(word) + + if not pairs: + return token+'' + + while True: + bigram = min(pairs, key = lambda pair: self.bpe_ranks.get(pair, float('inf'))) + if bigram not in self.bpe_ranks: + break + first, second = bigram + new_word = [] + i = 0 + while i < len(word): + try: + j = word.index(first, i) + new_word.extend(word[i:j]) + i = j + except: + new_word.extend(word[i:]) + break + + if word[i] == first and i < len(word)-1 and word[i+1] == second: + new_word.append(first+second) + i += 2 + else: + new_word.append(word[i]) + i += 1 + new_word = tuple(new_word) + word = new_word + if len(word) == 1: + break + else: + pairs = get_pairs(word) + word = ' '.join(word) + self.cache[token] = word + return word + + def encode(self, text): + bpe_tokens = [] + text = whitespace_clean(basic_clean(text)).lower() + for token in re.findall(self.pat, text): + token = ''.join(self.byte_encoder[b] for b in token.encode('utf-8')) + bpe_tokens.extend(self.encoder[bpe_token] for bpe_token in self.bpe(token).split(' ')) + return bpe_tokens + + def decode(self, tokens): + text = ''.join([self.decoder[token] for token in tokens]) + text = bytearray([self.byte_decoder[c] for c in text]).decode('utf-8', errors="replace").replace('', ' ') + return text + + +_tokenizer = SimpleTokenizer() + + +def tokenize(texts: Union[str, List[str]], context_length: int = 77) -> torch.LongTensor: + """ + Returns the tokenized representation of given input string(s) + + Parameters + ---------- + texts : Union[str, List[str]] + An input string or a list of input strings to tokenize + context_length : int + The context length to use; all CLIP models use 77 as the context length + + Returns + ------- + A two-dimensional tensor containing the resulting tokens, shape = [number of input strings, context_length] + """ + if isinstance(texts, str): + texts = [texts] + + sot_token = _tokenizer.encoder[""] + eot_token = _tokenizer.encoder[""] + all_tokens = [[sot_token] + _tokenizer.encode(text) + [eot_token] for text in texts] + result = torch.zeros(len(all_tokens), context_length, dtype=torch.long) + + for i, tokens in enumerate(all_tokens): + if len(tokens) > context_length: + tokens = tokens[:context_length] # Truncate + tokens[-1] = eot_token + result[i, :len(tokens)] = torch.tensor(tokens) + + return result + + +class HFTokenizer: + "HuggingFace tokenizer wrapper" + def __init__(self, tokenizer_name:str): + from transformers import AutoTokenizer + self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_name) + + def __call__(self, texts:Union[str, List[str]], context_length:int=77) -> torch.Tensor: + # same cleaning as for default tokenizer, except lowercasing + # adding lower (for case-sensitive tokenizers) will make it more robust but less sensitive to nuance + if isinstance(texts, str): + texts = [texts] + texts = [whitespace_clean(basic_clean(text)) for text in texts] + input_ids = self.tokenizer(texts, return_tensors='pt', max_length=context_length, padding='max_length', truncation=True).input_ids + return input_ids diff --git a/modules/pulid/eva_clip/transform.py b/modules/pulid/eva_clip/transform.py new file mode 100644 index 000000000..39f3e4cf6 --- /dev/null +++ b/modules/pulid/eva_clip/transform.py @@ -0,0 +1,103 @@ +from typing import Optional, Sequence, Tuple + +import torch +import torch.nn as nn +import torchvision.transforms.functional as F + +from torchvision.transforms import Normalize, Compose, RandomResizedCrop, InterpolationMode, ToTensor, Resize, \ + CenterCrop + +from .constants import OPENAI_DATASET_MEAN, OPENAI_DATASET_STD + + +class ResizeMaxSize(nn.Module): + + def __init__(self, max_size, interpolation=InterpolationMode.BICUBIC, fn='max', fill=0): + super().__init__() + if not isinstance(max_size, int): + raise TypeError(f"Size should be int. Got {type(max_size)}") + self.max_size = max_size + self.interpolation = interpolation + self.fn = min if fn == 'min' else min + self.fill = fill + + def forward(self, img): + if isinstance(img, torch.Tensor): + height, width = img.shape[:2] + else: + width, height = img.size + scale = self.max_size / float(max(height, width)) + if scale != 1.0: + new_size = tuple(round(dim * scale) for dim in (height, width)) + img = F.resize(img, new_size, self.interpolation) + pad_h = self.max_size - new_size[0] + pad_w = self.max_size - new_size[1] + img = F.pad(img, padding=[pad_w//2, pad_h//2, pad_w - pad_w//2, pad_h - pad_h//2], fill=self.fill) + return img + + +def _convert_to_rgb(image): + return image.convert('RGB') + + +# class CatGen(nn.Module): +# def __init__(self, num=4): +# self.num = num +# def mixgen_batch(image, text): +# batch_size = image.shape[0] +# index = np.random.permutation(batch_size) + +# cat_images = [] +# for i in range(batch_size): +# # image mixup +# image[i,:] = lam * image[i,:] + (1 - lam) * image[index[i],:] +# # text concat +# text[i] = tokenizer((str(text[i]) + " " + str(text[index[i]])))[0] +# text = torch.stack(text) +# return image, text + + +def image_transform( + image_size: int, + is_train: bool, + mean: Optional[Tuple[float, ...]] = None, + std: Optional[Tuple[float, ...]] = None, + resize_longest_max: bool = False, + fill_color: int = 0, +): + mean = mean or OPENAI_DATASET_MEAN + if not isinstance(mean, (list, tuple)): + mean = (mean,) * 3 + + std = std or OPENAI_DATASET_STD + if not isinstance(std, (list, tuple)): + std = (std,) * 3 + + if isinstance(image_size, (list, tuple)) and image_size[0] == image_size[1]: + # for square size, pass size as int so that Resize() uses aspect preserving shortest edge + image_size = image_size[0] + + normalize = Normalize(mean=mean, std=std) + if is_train: + return Compose([ + RandomResizedCrop(image_size, scale=(0.9, 1.0), interpolation=InterpolationMode.BICUBIC), + _convert_to_rgb, + ToTensor(), + normalize, + ]) + else: + if resize_longest_max: + transforms = [ + ResizeMaxSize(image_size, fill=fill_color) + ] + else: + transforms = [ + Resize(image_size, interpolation=InterpolationMode.BICUBIC), + CenterCrop(image_size), + ] + transforms.extend([ + _convert_to_rgb, + ToTensor(), + normalize, + ]) + return Compose(transforms) diff --git a/modules/pulid/eva_clip/transformer.py b/modules/pulid/eva_clip/transformer.py new file mode 100644 index 000000000..1e0a52ceb --- /dev/null +++ b/modules/pulid/eva_clip/transformer.py @@ -0,0 +1,721 @@ +import os +import logging +from collections import OrderedDict +import math +from typing import Callable, Optional, Sequence +import numpy as np +import torch +from torch import nn +from torch.nn import functional as F + +try: + from timm.models.layers import trunc_normal_ +except: + from timm.layers import trunc_normal_ + +from .rope import VisionRotaryEmbedding, VisionRotaryEmbeddingFast +from .utils import to_2tuple + + +class LayerNormFp32(nn.LayerNorm): + """Subclass torch's LayerNorm to handle fp16 (by casting to float32 and back).""" + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + def forward(self, x: torch.Tensor): + output = F.layer_norm( + x.float(), + self.normalized_shape, + self.weight.float() if self.weight is not None else None, + self.bias.float() if self.bias is not None else None, + self.eps, + ) + return output.type_as(x) + + +class LayerNorm(nn.LayerNorm): + """Subclass torch's LayerNorm (with cast back to input dtype).""" + + def forward(self, x: torch.Tensor): + orig_type = x.dtype + x = F.layer_norm(x, self.normalized_shape, self.weight, self.bias, self.eps) + return x.to(orig_type) + +class QuickGELU(nn.Module): + # NOTE This is slower than nn.GELU or nn.SiLU and uses more GPU memory + def forward(self, x: torch.Tensor): + return x * torch.sigmoid(1.702 * x) + + +class LayerScale(nn.Module): + def __init__(self, dim, init_values=1e-5, inplace=False): + super().__init__() + self.inplace = inplace + self.gamma = nn.Parameter(init_values * torch.ones(dim)) + + def forward(self, x): + return x.mul_(self.gamma) if self.inplace else x * self.gamma + +class PatchDropout(nn.Module): + """ + https://arxiv.org/abs/2212.00794 + """ + + def __init__(self, prob, exclude_first_token=True): + super().__init__() + assert 0 <= prob < 1. + self.prob = prob + self.exclude_first_token = exclude_first_token # exclude CLS token + logging.info(f"os.getenv('RoPE')={os.getenv('RoPE')}") + + def forward(self, x): + if not self.training or self.prob == 0.: + return x + + if self.exclude_first_token: + cls_tokens, x = x[:, :1], x[:, 1:] + else: + cls_tokens = torch.jit.annotate(torch.Tensor, x[:, :1]) + + batch = x.size()[0] + num_tokens = x.size()[1] + + batch_indices = torch.arange(batch) + batch_indices = batch_indices[..., None] + + keep_prob = 1 - self.prob + num_patches_keep = max(1, int(num_tokens * keep_prob)) + + rand = torch.randn(batch, num_tokens) + patch_indices_keep = rand.topk(num_patches_keep, dim=-1).indices + + x = x[batch_indices, patch_indices_keep] + + if self.exclude_first_token: + x = torch.cat((cls_tokens, x), dim=1) + + if self.training and os.getenv('RoPE') == '1': + return x, patch_indices_keep + + return x + + +def _in_projection_packed( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + w: torch.Tensor, + b: Optional[torch.Tensor] = None, + ): + """ + https://github.com/pytorch/pytorch/blob/db2a237763eb8693a20788be94f8c192e762baa8/torch/nn/functional.py#L4726 + """ + E = q.size(-1) + if k is v: + if q is k: + # self-attention + return F.linear(q, w, b).chunk(3, dim=-1) + else: + # encoder-decoder attention + w_q, w_kv = w.split([E, E * 2]) + if b is None: + b_q = b_kv = None + else: + b_q, b_kv = b.split([E, E * 2]) + return (F.linear(q, w_q, b_q),) + F.linear(k, w_kv, b_kv).chunk(2, dim=-1) + else: + w_q, w_k, w_v = w.chunk(3) + if b is None: + b_q = b_k = b_v = None + else: + b_q, b_k, b_v = b.chunk(3) + return F.linear(q, w_q, b_q), F.linear(k, w_k, b_k), F.linear(v, w_v, b_v) + +class Attention(nn.Module): + def __init__( + self, + dim, + num_heads=8, + qkv_bias=True, + scaled_cosine=False, + scale_heads=False, + logit_scale_max=math.log(1. / 0.01), + attn_drop=0., + proj_drop=0., + xattn=False, + rope=False + ): + super().__init__() + self.scaled_cosine = scaled_cosine + self.scale_heads = scale_heads + assert dim % num_heads == 0, 'dim should be divisible by num_heads' + self.num_heads = num_heads + self.head_dim = dim // num_heads + self.scale = self.head_dim ** -0.5 + self.logit_scale_max = logit_scale_max + + # keeping in_proj in this form (instead of nn.Linear) to match weight scheme of original + self.in_proj_weight = nn.Parameter(torch.randn((dim * 3, dim)) * self.scale) + if qkv_bias: + self.in_proj_bias = nn.Parameter(torch.zeros(dim * 3)) + else: + self.in_proj_bias = None + + if self.scaled_cosine: + self.logit_scale = nn.Parameter(torch.log(10 * torch.ones((num_heads, 1, 1)))) + else: + self.logit_scale = None + self.attn_drop = nn.Dropout(attn_drop) + if self.scale_heads: + self.head_scale = nn.Parameter(torch.ones((num_heads, 1, 1))) + else: + self.head_scale = None + self.out_proj = nn.Linear(dim, dim) + self.out_drop = nn.Dropout(proj_drop) + self.xattn = xattn + self.xattn_drop = attn_drop + self.rope = rope + + def forward(self, x, attn_mask: Optional[torch.Tensor] = None): + L, N, C = x.shape + q, k, v = F.linear(x, self.in_proj_weight, self.in_proj_bias).chunk(3, dim=-1) + if self.xattn: + q = q.contiguous().view(L, N, self.num_heads, -1).transpose(0, 1) + k = k.contiguous().view(L, N, self.num_heads, -1).transpose(0, 1) + v = v.contiguous().view(L, N, self.num_heads, -1).transpose(0, 1) + + x = xops.memory_efficient_attention( + q, k, v, + p=self.xattn_drop, + scale=self.scale if self.logit_scale is None else None, + attn_bias=xops.LowerTriangularMask() if attn_mask is not None else None, + ) + else: + q = q.contiguous().view(L, N * self.num_heads, -1).transpose(0, 1) + k = k.contiguous().view(L, N * self.num_heads, -1).transpose(0, 1) + v = v.contiguous().view(L, N * self.num_heads, -1).transpose(0, 1) + + if self.logit_scale is not None: + attn = torch.bmm(F.normalize(q, dim=-1), F.normalize(k, dim=-1).transpose(-1, -2)) + logit_scale = torch.clamp(self.logit_scale, max=self.logit_scale_max).exp() + attn = attn.view(N, self.num_heads, L, L) * logit_scale + attn = attn.view(-1, L, L) + else: + q = q * self.scale + attn = torch.bmm(q, k.transpose(-1, -2)) + + if attn_mask is not None: + if attn_mask.dtype == torch.bool: + new_attn_mask = torch.zeros_like(attn_mask, dtype=q.dtype) + new_attn_mask.masked_fill_(attn_mask, float("-inf")) + attn_mask = new_attn_mask + attn += attn_mask + + attn = attn.softmax(dim=-1) + attn = self.attn_drop(attn) + + x = torch.bmm(attn, v) + + if self.head_scale is not None: + x = x.view(N, self.num_heads, L, C) * self.head_scale + x = x.view(-1, L, C) + x = x.transpose(0, 1).reshape(L, N, C) + x = self.out_proj(x) + x = self.out_drop(x) + return x + +class CustomAttention(nn.Module): + def __init__( + self, + dim, + num_heads=8, + qkv_bias=True, + scaled_cosine=True, + scale_heads=False, + logit_scale_max=math.log(1. / 0.01), + attn_drop=0., + proj_drop=0., + xattn=False + ): + super().__init__() + self.scaled_cosine = scaled_cosine + self.scale_heads = scale_heads + assert dim % num_heads == 0, 'dim should be divisible by num_heads' + self.num_heads = num_heads + self.head_dim = dim // num_heads + self.scale = self.head_dim ** -0.5 + self.logit_scale_max = logit_scale_max + + # keeping in_proj in this form (instead of nn.Linear) to match weight scheme of original + self.in_proj_weight = nn.Parameter(torch.randn((dim * 3, dim)) * self.scale) + if qkv_bias: + self.in_proj_bias = nn.Parameter(torch.zeros(dim * 3)) + else: + self.in_proj_bias = None + + if self.scaled_cosine: + self.logit_scale = nn.Parameter(torch.log(10 * torch.ones((num_heads, 1, 1)))) + else: + self.logit_scale = None + self.attn_drop = nn.Dropout(attn_drop) + if self.scale_heads: + self.head_scale = nn.Parameter(torch.ones((num_heads, 1, 1))) + else: + self.head_scale = None + self.out_proj = nn.Linear(dim, dim) + self.out_drop = nn.Dropout(proj_drop) + self.xattn = xattn + self.xattn_drop = attn_drop + + def forward(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attn_mask: Optional[torch.Tensor] = None): + q, k, v = _in_projection_packed(query, key, value, self.in_proj_weight, self.in_proj_bias) + N_q, B_q, C_q = q.shape + N_k, B_k, C_k = k.shape + N_v, B_v, C_v = v.shape + if self.xattn: + # B, N, C -> B, N, num_heads, C + q = q.permute(1, 0, 2).reshape(B_q, N_q, self.num_heads, -1) + k = k.permute(1, 0, 2).reshape(B_k, N_k, self.num_heads, -1) + v = v.permute(1, 0, 2).reshape(B_v, N_v, self.num_heads, -1) + + x = xops.memory_efficient_attention( + q, k, v, + p=self.xattn_drop, + scale=self.scale if self.logit_scale is None else None, + attn_bias=xops.LowerTriangularMask() if attn_mask is not None else None + ) + else: + # B*H, L, C + q = q.contiguous().view(N_q, B_q * self.num_heads, -1).transpose(0, 1) + k = k.contiguous().view(N_k, B_k * self.num_heads, -1).transpose(0, 1) + v = v.contiguous().view(N_v, B_v * self.num_heads, -1).transpose(0, 1) + + if self.logit_scale is not None: + # B*H, N_q, N_k + attn = torch.bmm(F.normalize(q, dim=-1), F.normalize(k, dim=-1).transpose(-1, -2)) + logit_scale = torch.clamp(self.logit_scale, max=self.logit_scale_max).exp() + attn = attn.view(B_q, self.num_heads, N_q, N_k) * logit_scale + attn = attn.view(-1, N_q, N_k) + else: + q = q * self.scale + attn = torch.bmm(q, k.transpose(-1, -2)) + + if attn_mask is not None: + if attn_mask.dtype == torch.bool: + new_attn_mask = torch.zeros_like(attn_mask, dtype=q.dtype) + new_attn_mask.masked_fill_(attn_mask, float("-inf")) + attn_mask = new_attn_mask + attn += attn_mask + + attn = attn.softmax(dim=-1) + attn = self.attn_drop(attn) + + x = torch.bmm(attn, v) + + if self.head_scale is not None: + x = x.view(B_q, self.num_heads, N_q, C_q) * self.head_scale + x = x.view(-1, N_q, C_q) + x = x.transpose(0, 1).reshape(N_q, B_q, C_q) + x = self.out_proj(x) + x = self.out_drop(x) + return x + +class CustomResidualAttentionBlock(nn.Module): + def __init__( + self, + d_model: int, + n_head: int, + mlp_ratio: float = 4.0, + ls_init_value: float = None, + act_layer: Callable = nn.GELU, + norm_layer: Callable = LayerNorm, + scale_cosine_attn: bool = False, + scale_heads: bool = False, + scale_attn: bool = False, + scale_fc: bool = False, + cross_attn: bool = False, + xattn: bool = False, + ): + super().__init__() + + self.ln_1 = norm_layer(d_model) + self.ln_1_k = norm_layer(d_model) if cross_attn else self.ln_1 + self.ln_1_v = norm_layer(d_model) if cross_attn else self.ln_1 + self.attn = CustomAttention( + d_model, n_head, + qkv_bias=True, + attn_drop=0., + proj_drop=0., + scaled_cosine=scale_cosine_attn, + scale_heads=scale_heads, + xattn=xattn + ) + + self.ln_attn = norm_layer(d_model) if scale_attn else nn.Identity() + self.ls_1 = LayerScale(d_model, ls_init_value) if ls_init_value is not None else nn.Identity() + + self.ln_2 = norm_layer(d_model) + mlp_width = int(d_model * mlp_ratio) + self.mlp = nn.Sequential(OrderedDict([ + ("c_fc", nn.Linear(d_model, mlp_width)), + ('ln', norm_layer(mlp_width) if scale_fc else nn.Identity()), + ("gelu", act_layer()), + ("c_proj", nn.Linear(mlp_width, d_model)) + ])) + + self.ls_2 = LayerScale(d_model, ls_init_value) if ls_init_value is not None else nn.Identity() + + def forward(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, attn_mask: Optional[torch.Tensor] = None): + q = q + self.ls_1(self.ln_attn(self.attn(self.ln_1(q), self.ln_1_k(k), self.ln_1_v(v), attn_mask=attn_mask))) + q = q + self.ls_2(self.mlp(self.ln_2(q))) + return q + +class CustomTransformer(nn.Module): + def __init__( + self, + width: int, + layers: int, + heads: int, + mlp_ratio: float = 4.0, + ls_init_value: float = None, + act_layer: Callable = nn.GELU, + norm_layer: Callable = LayerNorm, + scale_cosine_attn: bool = True, + scale_heads: bool = False, + scale_attn: bool = False, + scale_fc: bool = False, + cross_attn: bool = False, + xattn: bool = False, + ): + super().__init__() + self.width = width + self.layers = layers + self.grad_checkpointing = False + self.xattn = xattn + + self.resblocks = nn.ModuleList([ + CustomResidualAttentionBlock( + width, + heads, + mlp_ratio, + ls_init_value=ls_init_value, + act_layer=act_layer, + norm_layer=norm_layer, + scale_cosine_attn=scale_cosine_attn, + scale_heads=scale_heads, + scale_attn=scale_attn, + scale_fc=scale_fc, + cross_attn=cross_attn, + xattn=xattn) + for _ in range(layers) + ]) + + def get_cast_dtype(self) -> torch.dtype: + return self.resblocks[0].mlp.c_fc.weight.dtype + + def forward(self, q: torch.Tensor, k: torch.Tensor = None, v: torch.Tensor = None, attn_mask: Optional[torch.Tensor] = None): + if k is None and v is None: + k = v = q + for r in self.resblocks: + if self.grad_checkpointing and not torch.jit.is_scripting(): + q = checkpoint(r, q, k, v, attn_mask) + else: + q = r(q, k, v, attn_mask=attn_mask) + return q + + +class ResidualAttentionBlock(nn.Module): + def __init__( + self, + d_model: int, + n_head: int, + mlp_ratio: float = 4.0, + ls_init_value: float = None, + act_layer: Callable = nn.GELU, + norm_layer: Callable = LayerNorm, + xattn: bool = False, + ): + super().__init__() + + self.ln_1 = norm_layer(d_model) + if xattn: + self.attn = Attention(d_model, n_head, xattn=True) + else: + self.attn = nn.MultiheadAttention(d_model, n_head) + self.ls_1 = LayerScale(d_model, ls_init_value) if ls_init_value is not None else nn.Identity() + + self.ln_2 = norm_layer(d_model) + mlp_width = int(d_model * mlp_ratio) + self.mlp = nn.Sequential(OrderedDict([ + ("c_fc", nn.Linear(d_model, mlp_width)), + ("gelu", act_layer()), + ("c_proj", nn.Linear(mlp_width, d_model)) + ])) + + self.ls_2 = LayerScale(d_model, ls_init_value) if ls_init_value is not None else nn.Identity() + self.xattn = xattn + + def attention(self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None): + attn_mask = attn_mask.to(x.dtype) if attn_mask is not None else None + if self.xattn: + return self.attn(x, attn_mask=attn_mask) + return self.attn(x, x, x, need_weights=False, attn_mask=attn_mask)[0] + + def forward(self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None): + x = x + self.ls_1(self.attention(self.ln_1(x), attn_mask=attn_mask)) + x = x + self.ls_2(self.mlp(self.ln_2(x))) + return x + +class Transformer(nn.Module): + def __init__( + self, + width: int, + layers: int, + heads: int, + mlp_ratio: float = 4.0, + ls_init_value: float = None, + act_layer: Callable = nn.GELU, + norm_layer: Callable = LayerNorm, + xattn: bool = False, + ): + super().__init__() + self.width = width + self.layers = layers + self.grad_checkpointing = False + + self.resblocks = nn.ModuleList([ + ResidualAttentionBlock( + width, heads, mlp_ratio, ls_init_value=ls_init_value, act_layer=act_layer, norm_layer=norm_layer, xattn=xattn) + for _ in range(layers) + ]) + + def get_cast_dtype(self) -> torch.dtype: + return self.resblocks[0].mlp.c_fc.weight.dtype + + def forward(self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None): + for r in self.resblocks: + if self.grad_checkpointing and not torch.jit.is_scripting(): + x = checkpoint(r, x, attn_mask) + else: + x = r(x, attn_mask=attn_mask) + return x + + +class VisionTransformer(nn.Module): + def __init__( + self, + image_size: int, + patch_size: int, + width: int, + layers: int, + heads: int, + mlp_ratio: float, + ls_init_value: float = None, + patch_dropout: float = 0., + global_average_pool: bool = False, + output_dim: int = 512, + act_layer: Callable = nn.GELU, + norm_layer: Callable = LayerNorm, + xattn: bool = False, + ): + super().__init__() + self.image_size = to_2tuple(image_size) + self.patch_size = to_2tuple(patch_size) + self.grid_size = (self.image_size[0] // self.patch_size[0], self.image_size[1] // self.patch_size[1]) + self.output_dim = output_dim + self.conv1 = nn.Conv2d(in_channels=3, out_channels=width, kernel_size=patch_size, stride=patch_size, bias=False) + + scale = width ** -0.5 + self.class_embedding = nn.Parameter(scale * torch.randn(width)) + self.positional_embedding = nn.Parameter(scale * torch.randn(self.grid_size[0] * self.grid_size[1] + 1, width)) + + # setting a patch_dropout of 0. would mean it is disabled and this function would be the identity fn + self.patch_dropout = PatchDropout(patch_dropout) if patch_dropout > 0. else nn.Identity() + self.ln_pre = norm_layer(width) + + self.transformer = Transformer( + width, + layers, + heads, + mlp_ratio, + ls_init_value=ls_init_value, + act_layer=act_layer, + norm_layer=norm_layer, + xattn=xattn + ) + + self.global_average_pool = global_average_pool + self.ln_post = norm_layer(width) + self.proj = nn.Parameter(scale * torch.randn(width, output_dim)) + + def lock(self, unlocked_groups=0, freeze_bn_stats=False): + for param in self.parameters(): + param.requires_grad = False + + if unlocked_groups != 0: + groups = [ + [ + self.conv1, + self.class_embedding, + self.positional_embedding, + self.ln_pre, + ], + *self.transformer.resblocks[:-1], + [ + self.transformer.resblocks[-1], + self.ln_post, + ], + self.proj, + ] + + def _unlock(x): + if isinstance(x, Sequence): + for g in x: + _unlock(g) + else: + if isinstance(x, torch.nn.Parameter): + x.requires_grad = True + else: + for p in x.parameters(): + p.requires_grad = True + + _unlock(groups[-unlocked_groups:]) + + def get_num_layers(self): + return self.transformer.layers + + @torch.jit.ignore + def set_grad_checkpointing(self, enable=True): + self.transformer.grad_checkpointing = enable + + @torch.jit.ignore + def no_weight_decay(self): + return {'positional_embedding', 'class_embedding'} + + def forward(self, x: torch.Tensor, return_all_features: bool=False): + x = self.conv1(x) # shape = [*, width, grid, grid] + x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2] + x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width] + x = torch.cat( + [self.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device), + x], dim=1) # shape = [*, grid ** 2 + 1, width] + x = x + self.positional_embedding.to(x.dtype) + + # a patch_dropout of 0. would mean it is disabled and this function would do nothing but return what was passed in + x = self.patch_dropout(x) + x = self.ln_pre(x) + + x = x.permute(1, 0, 2) # NLD -> LND + x = self.transformer(x) + x = x.permute(1, 0, 2) # LND -> NLD + + if not return_all_features: + if self.global_average_pool: + x = x.mean(dim=1) #x = x[:,1:,:].mean(dim=1) + else: + x = x[:, 0] + + x = self.ln_post(x) + + if self.proj is not None: + x = x @ self.proj + + return x + + +class TextTransformer(nn.Module): + def __init__( + self, + context_length: int = 77, + vocab_size: int = 49408, + width: int = 512, + heads: int = 8, + layers: int = 12, + ls_init_value: float = None, + output_dim: int = 512, + act_layer: Callable = nn.GELU, + norm_layer: Callable = LayerNorm, + xattn: bool= False, + attn_mask: bool = True + ): + super().__init__() + self.context_length = context_length + self.vocab_size = vocab_size + self.width = width + self.output_dim = output_dim + + self.token_embedding = nn.Embedding(vocab_size, width) + self.positional_embedding = nn.Parameter(torch.empty(self.context_length, width)) + self.transformer = Transformer( + width=width, + layers=layers, + heads=heads, + ls_init_value=ls_init_value, + act_layer=act_layer, + norm_layer=norm_layer, + xattn=xattn + ) + + self.xattn = xattn + self.ln_final = norm_layer(width) + self.text_projection = nn.Parameter(torch.empty(width, output_dim)) + + if attn_mask: + self.register_buffer('attn_mask', self.build_attention_mask(), persistent=False) + else: + self.attn_mask = None + + self.init_parameters() + + def init_parameters(self): + nn.init.normal_(self.token_embedding.weight, std=0.02) + nn.init.normal_(self.positional_embedding, std=0.01) + + proj_std = (self.transformer.width ** -0.5) * ((2 * self.transformer.layers) ** -0.5) + attn_std = self.transformer.width ** -0.5 + fc_std = (2 * self.transformer.width) ** -0.5 + for block in self.transformer.resblocks: + nn.init.normal_(block.attn.in_proj_weight, std=attn_std) + nn.init.normal_(block.attn.out_proj.weight, std=proj_std) + nn.init.normal_(block.mlp.c_fc.weight, std=fc_std) + nn.init.normal_(block.mlp.c_proj.weight, std=proj_std) + + if self.text_projection is not None: + nn.init.normal_(self.text_projection, std=self.transformer.width ** -0.5) + + @torch.jit.ignore + def set_grad_checkpointing(self, enable=True): + self.transformer.grad_checkpointing = enable + + @torch.jit.ignore + def no_weight_decay(self): + # return {'positional_embedding', 'token_embedding'} + return {'positional_embedding'} + + def get_num_layers(self): + return self.transformer.layers + + def build_attention_mask(self): + # lazily create causal attention mask, with full attention between the vision tokens + # pytorch uses additive attention mask; fill with -inf + mask = torch.empty(self.context_length, self.context_length) + mask.fill_(float("-inf")) + mask.triu_(1) # zero out the lower diagonal + return mask + + def forward(self, text, return_all_features: bool=False): + cast_dtype = self.transformer.get_cast_dtype() + x = self.token_embedding(text).to(cast_dtype) # [batch_size, n_ctx, d_model] + + x = x + self.positional_embedding.to(cast_dtype) + x = x.permute(1, 0, 2) # NLD -> LND + x = self.transformer(x, attn_mask=self.attn_mask) + # x = self.transformer(x) # no attention mask is applied + x = x.permute(1, 0, 2) # LND -> NLD + x = self.ln_final(x) + + if not return_all_features: + # x.shape = [batch_size, n_ctx, transformer.width] + # take features from the eot embedding (eot_token is the highest number in each sequence) + x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection + return x diff --git a/modules/pulid/eva_clip/utils.py b/modules/pulid/eva_clip/utils.py new file mode 100644 index 000000000..bdc5a7a45 --- /dev/null +++ b/modules/pulid/eva_clip/utils.py @@ -0,0 +1,326 @@ +from itertools import repeat +import collections.abc +import logging +import math +import numpy as np + +import torch +from torch import nn as nn +from torchvision.ops.misc import FrozenBatchNorm2d +import torch.nn.functional as F + +# open CLIP +def resize_clip_pos_embed(state_dict, model, interpolation: str = 'bicubic', seq_dim=1): + # Rescale the grid of position embeddings when loading from state_dict + old_pos_embed = state_dict.get('visual.positional_embedding', None) + if old_pos_embed is None or not hasattr(model.visual, 'grid_size'): + return + grid_size = to_2tuple(model.visual.grid_size) + extra_tokens = 1 # FIXME detect different token configs (ie no class token, or more) + new_seq_len = grid_size[0] * grid_size[1] + extra_tokens + if new_seq_len == old_pos_embed.shape[0]: + return + + if extra_tokens: + pos_emb_tok, pos_emb_img = old_pos_embed[:extra_tokens], old_pos_embed[extra_tokens:] + else: + pos_emb_tok, pos_emb_img = None, old_pos_embed + old_grid_size = to_2tuple(int(math.sqrt(len(pos_emb_img)))) + + logging.info('Resizing position embedding grid-size from %s to %s', old_grid_size, grid_size) + pos_emb_img = pos_emb_img.reshape(1, old_grid_size[0], old_grid_size[1], -1).permute(0, 3, 1, 2) + pos_emb_img = F.interpolate( + pos_emb_img, + size=grid_size, + mode=interpolation, + align_corners=True, + ) + pos_emb_img = pos_emb_img.permute(0, 2, 3, 1).reshape(1, grid_size[0] * grid_size[1], -1)[0] + if pos_emb_tok is not None: + new_pos_embed = torch.cat([pos_emb_tok, pos_emb_img], dim=0) + else: + new_pos_embed = pos_emb_img + state_dict['visual.positional_embedding'] = new_pos_embed + + +def resize_visual_pos_embed(state_dict, model, interpolation: str = 'bicubic', seq_dim=1): + # Rescale the grid of position embeddings when loading from state_dict + old_pos_embed = state_dict.get('positional_embedding', None) + if old_pos_embed is None or not hasattr(model.visual, 'grid_size'): + return + grid_size = to_2tuple(model.visual.grid_size) + extra_tokens = 1 # FIXME detect different token configs (ie no class token, or more) + new_seq_len = grid_size[0] * grid_size[1] + extra_tokens + if new_seq_len == old_pos_embed.shape[0]: + return + + if extra_tokens: + pos_emb_tok, pos_emb_img = old_pos_embed[:extra_tokens], old_pos_embed[extra_tokens:] + else: + pos_emb_tok, pos_emb_img = None, old_pos_embed + old_grid_size = to_2tuple(int(math.sqrt(len(pos_emb_img)))) + + logging.info('Resizing position embedding grid-size from %s to %s', old_grid_size, grid_size) + pos_emb_img = pos_emb_img.reshape(1, old_grid_size[0], old_grid_size[1], -1).permute(0, 3, 1, 2) + pos_emb_img = F.interpolate( + pos_emb_img, + size=grid_size, + mode=interpolation, + align_corners=True, + ) + pos_emb_img = pos_emb_img.permute(0, 2, 3, 1).reshape(1, grid_size[0] * grid_size[1], -1)[0] + if pos_emb_tok is not None: + new_pos_embed = torch.cat([pos_emb_tok, pos_emb_img], dim=0) + else: + new_pos_embed = pos_emb_img + state_dict['positional_embedding'] = new_pos_embed + +def resize_evaclip_pos_embed(state_dict, model, interpolation: str = 'bicubic', seq_dim=1): + all_keys = list(state_dict.keys()) + # interpolate position embedding + if 'visual.pos_embed' in state_dict: + pos_embed_checkpoint = state_dict['visual.pos_embed'] + embedding_size = pos_embed_checkpoint.shape[-1] + num_patches = model.visual.patch_embed.num_patches + num_extra_tokens = model.visual.pos_embed.shape[-2] - num_patches + # height (== width) for the checkpoint position embedding + orig_size = int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5) + # height (== width) for the new position embedding + new_size = int(num_patches ** 0.5) + # class_token and dist_token are kept unchanged + if orig_size != new_size: + print("Position interpolate from %dx%d to %dx%d" % (orig_size, orig_size, new_size, new_size)) + extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens] + # only the position tokens are interpolated + pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:] + pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, embedding_size).permute(0, 3, 1, 2) + pos_tokens = torch.nn.functional.interpolate( + pos_tokens, size=(new_size, new_size), mode='bicubic', align_corners=False) + pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2) + new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1) + state_dict['visual.pos_embed'] = new_pos_embed + + patch_embed_proj = state_dict['visual.patch_embed.proj.weight'] + patch_size = model.visual.patch_embed.patch_size + state_dict['visual.patch_embed.proj.weight'] = torch.nn.functional.interpolate( + patch_embed_proj.float(), size=patch_size, mode='bicubic', align_corners=False) + + +def resize_eva_pos_embed(state_dict, model, interpolation: str = 'bicubic', seq_dim=1): + all_keys = list(state_dict.keys()) + # interpolate position embedding + if 'pos_embed' in state_dict: + pos_embed_checkpoint = state_dict['pos_embed'] + embedding_size = pos_embed_checkpoint.shape[-1] + num_patches = model.visual.patch_embed.num_patches + num_extra_tokens = model.visual.pos_embed.shape[-2] - num_patches + # height (== width) for the checkpoint position embedding + orig_size = int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5) + # height (== width) for the new position embedding + new_size = int(num_patches ** 0.5) + # class_token and dist_token are kept unchanged + if orig_size != new_size: + print("Position interpolate from %dx%d to %dx%d" % (orig_size, orig_size, new_size, new_size)) + extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens] + # only the position tokens are interpolated + pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:] + pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, embedding_size).permute(0, 3, 1, 2) + pos_tokens = torch.nn.functional.interpolate( + pos_tokens, size=(new_size, new_size), mode='bicubic', align_corners=False) + pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2) + new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1) + state_dict['pos_embed'] = new_pos_embed + + patch_embed_proj = state_dict['patch_embed.proj.weight'] + patch_size = model.visual.patch_embed.patch_size + state_dict['patch_embed.proj.weight'] = torch.nn.functional.interpolate( + patch_embed_proj.float(), size=patch_size, mode='bicubic', align_corners=False) + + +def resize_rel_pos_embed(state_dict, model, interpolation: str = 'bicubic', seq_dim=1): + all_keys = list(state_dict.keys()) + for key in all_keys: + if "relative_position_index" in key: + state_dict.pop(key) + + if "relative_position_bias_table" in key: + rel_pos_bias = state_dict[key] + src_num_pos, num_attn_heads = rel_pos_bias.size() + dst_num_pos, _ = model.visual.state_dict()[key].size() + dst_patch_shape = model.visual.patch_embed.patch_shape + if dst_patch_shape[0] != dst_patch_shape[1]: + raise NotImplementedError() + num_extra_tokens = dst_num_pos - (dst_patch_shape[0] * 2 - 1) * (dst_patch_shape[1] * 2 - 1) + src_size = int((src_num_pos - num_extra_tokens) ** 0.5) + dst_size = int((dst_num_pos - num_extra_tokens) ** 0.5) + if src_size != dst_size: + print("Position interpolate for %s from %dx%d to %dx%d" % ( + key, src_size, src_size, dst_size, dst_size)) + extra_tokens = rel_pos_bias[-num_extra_tokens:, :] + rel_pos_bias = rel_pos_bias[:-num_extra_tokens, :] + + def geometric_progression(a, r, n): + return a * (1.0 - r ** n) / (1.0 - r) + + left, right = 1.01, 1.5 + while right - left > 1e-6: + q = (left + right) / 2.0 + gp = geometric_progression(1, q, src_size // 2) + if gp > dst_size // 2: + right = q + else: + left = q + + # if q > 1.090307: + # q = 1.090307 + + dis = [] + cur = 1 + for i in range(src_size // 2): + dis.append(cur) + cur += q ** (i + 1) + + r_ids = [-_ for _ in reversed(dis)] + + x = r_ids + [0] + dis + y = r_ids + [0] + dis + + t = dst_size // 2.0 + dx = np.arange(-t, t + 0.1, 1.0) + dy = np.arange(-t, t + 0.1, 1.0) + + print("Original positions = %s" % str(x)) + print("Target positions = %s" % str(dx)) + + all_rel_pos_bias = [] + + for i in range(num_attn_heads): + z = rel_pos_bias[:, i].view(src_size, src_size).float().numpy() + f = F.interpolate.interp2d(x, y, z, kind='cubic') + all_rel_pos_bias.append( + torch.Tensor(f(dx, dy)).contiguous().view(-1, 1).to(rel_pos_bias.device)) + + rel_pos_bias = torch.cat(all_rel_pos_bias, dim=-1) + + new_rel_pos_bias = torch.cat((rel_pos_bias, extra_tokens), dim=0) + state_dict[key] = new_rel_pos_bias + + # interpolate position embedding + if 'pos_embed' in state_dict: + pos_embed_checkpoint = state_dict['pos_embed'] + embedding_size = pos_embed_checkpoint.shape[-1] + num_patches = model.visual.patch_embed.num_patches + num_extra_tokens = model.visual.pos_embed.shape[-2] - num_patches + # height (== width) for the checkpoint position embedding + orig_size = int((pos_embed_checkpoint.shape[-2] - num_extra_tokens) ** 0.5) + # height (== width) for the new position embedding + new_size = int(num_patches ** 0.5) + # class_token and dist_token are kept unchanged + if orig_size != new_size: + print("Position interpolate from %dx%d to %dx%d" % (orig_size, orig_size, new_size, new_size)) + extra_tokens = pos_embed_checkpoint[:, :num_extra_tokens] + # only the position tokens are interpolated + pos_tokens = pos_embed_checkpoint[:, num_extra_tokens:] + pos_tokens = pos_tokens.reshape(-1, orig_size, orig_size, embedding_size).permute(0, 3, 1, 2) + pos_tokens = torch.nn.functional.interpolate( + pos_tokens, size=(new_size, new_size), mode='bicubic', align_corners=False) + pos_tokens = pos_tokens.permute(0, 2, 3, 1).flatten(1, 2) + new_pos_embed = torch.cat((extra_tokens, pos_tokens), dim=1) + state_dict['pos_embed'] = new_pos_embed + + patch_embed_proj = state_dict['patch_embed.proj.weight'] + patch_size = model.visual.patch_embed.patch_size + state_dict['patch_embed.proj.weight'] = torch.nn.functional.interpolate( + patch_embed_proj.float(), size=patch_size, mode='bicubic', align_corners=False) + + +def freeze_batch_norm_2d(module, module_match={}, name=''): + """ + Converts all `BatchNorm2d` and `SyncBatchNorm` layers of provided module into `FrozenBatchNorm2d`. If `module` is + itself an instance of either `BatchNorm2d` or `SyncBatchNorm`, it is converted into `FrozenBatchNorm2d` and + returned. Otherwise, the module is walked recursively and submodules are converted in place. + + Args: + module (torch.nn.Module): Any PyTorch module. + module_match (dict): Dictionary of full module names to freeze (all if empty) + name (str): Full module name (prefix) + + Returns: + torch.nn.Module: Resulting module + + Inspired by https://github.com/pytorch/pytorch/blob/a5895f85be0f10212791145bfedc0261d364f103/torch/nn/modules/batchnorm.py#L762 + """ + res = module + is_match = True + if module_match: + is_match = name in module_match + if is_match and isinstance(module, (nn.modules.batchnorm.BatchNorm2d, nn.modules.batchnorm.SyncBatchNorm)): + res = FrozenBatchNorm2d(module.num_features) + res.num_features = module.num_features + res.affine = module.affine + if module.affine: + res.weight.data = module.weight.data.clone().detach() + res.bias.data = module.bias.data.clone().detach() + res.running_mean.data = module.running_mean.data + res.running_var.data = module.running_var.data + res.eps = module.eps + else: + for child_name, child in module.named_children(): + full_child_name = '.'.join([name, child_name]) if name else child_name + new_child = freeze_batch_norm_2d(child, module_match, full_child_name) + if new_child is not child: + res.add_module(child_name, new_child) + return res + + +# From PyTorch internals +def _ntuple(n): + def parse(x): + if isinstance(x, collections.abc.Iterable): + return x + return tuple(repeat(x, n)) + return parse + + +to_1tuple = _ntuple(1) +to_2tuple = _ntuple(2) +to_3tuple = _ntuple(3) +to_4tuple = _ntuple(4) +to_ntuple = lambda n, x: _ntuple(n)(x) + + +def is_logging(args): + def is_global_master(args): + return args.rank == 0 + + def is_local_master(args): + return args.local_rank == 0 + + def is_master(args, local=False): + return is_local_master(args) if local else is_global_master(args) + return is_master + + +class AllGather(torch.autograd.Function): + """An autograd function that performs allgather on a tensor. + Performs all_gather operation on the provided tensors. + *** Warning ***: torch.distributed.all_gather has no gradient. + """ + + @staticmethod + def forward(ctx, tensor, rank, world_size): + tensors_gather = [torch.empty_like(tensor) for _ in range(world_size)] + torch.distributed.all_gather(tensors_gather, tensor) + ctx.rank = rank + ctx.batch_size = tensor.shape[0] + return torch.cat(tensors_gather, 0) + + @staticmethod + def backward(ctx, grad_output): + return ( + grad_output[ctx.batch_size * ctx.rank: ctx.batch_size * (ctx.rank + 1)], + None, + None + ) + +allgather = AllGather.apply \ No newline at end of file diff --git a/modules/pulid/pulid_sdxl.py b/modules/pulid/pulid_sdxl.py new file mode 100644 index 000000000..0651efab4 --- /dev/null +++ b/modules/pulid/pulid_sdxl.py @@ -0,0 +1,356 @@ +import os +import cv2 +import insightface +import numpy as np +import torch +import torch.nn as nn +from diffusers import DPMSolverMultistepScheduler, StableDiffusionXLPipeline + +from huggingface_hub import hf_hub_download, snapshot_download +from safetensors.torch import load_file +from torchvision.transforms import InterpolationMode +from torchvision.transforms.functional import normalize, resize + +from basicsr.utils import img2tensor, tensor2img +from facexlib.parsing import init_parsing_model +from facexlib.utils.face_restoration_helper import FaceRestoreHelper +from insightface.app import FaceAnalysis + +from eva_clip import create_model_and_transforms +from eva_clip.constants import OPENAI_DATASET_MEAN, OPENAI_DATASET_STD +from encoders_transformer import IDFormer +from pulid_utils import sample_dpmpp_2m, sample_dpmpp_sde +from attention_processor import AttnProcessor2_0 as AttnProcessor +from attention_processor import IDAttnProcessor2_0 as IDAttnProcessor + + +class StableDiffusionXLPuLIDPipeline: + def __init__(self, pipe: StableDiffusionXLPipeline, device: torch.device, sampler='dpmpp_sde', cache_dir=None): + super().__init__() + self.device = device + self.pipe = pipe + self.cache_dir = cache_dir + self.hack_unet_attn_layers(self.pipe.unet) + self.pipe.scheduler = DPMSolverMultistepScheduler.from_config(self.pipe.scheduler.config) + self.id_adapter = IDFormer().to(self.device) + + # preprocessors + # face align and parsing + self.face_helper = FaceRestoreHelper( + upscale_factor=1, + face_size=512, + crop_ratio=(1, 1), + det_model='retinaface_resnet50', + save_ext='png', + device=self.device, + ) + self.face_helper.face_parse = None + self.face_helper.face_parse = init_parsing_model(model_name='bisenet', device=self.device) + + # clip-vit backbone + model, _, _ = create_model_and_transforms('EVA02-CLIP-L-14-336', 'eva_clip', force_custom_clip=True) + model = model.visual + self.clip_vision_model = model.to(self.device) + eva_transform_mean = getattr(self.clip_vision_model, 'image_mean', OPENAI_DATASET_MEAN) + eva_transform_std = getattr(self.clip_vision_model, 'image_std', OPENAI_DATASET_STD) + if not isinstance(eva_transform_mean, (list, tuple)): + eva_transform_mean = (eva_transform_mean,) * 3 + if not isinstance(eva_transform_std, (list, tuple)): + eva_transform_std = (eva_transform_std,) * 3 + self.eva_transform_mean = eva_transform_mean + self.eva_transform_std = eva_transform_std + + # antelopev2 + # snapshot_download('DIAMONIK7777/antelopev2', local_dir='models/antelopev2') + local_dir = os.path.join(self.cache_dir, 'pulid', 'models', 'antelopev2') + _loc = snapshot_download('DIAMONIK7777/antelopev2', local_dir=local_dir) + self.app = FaceAnalysis( + name='antelopev2', + root=os.path.join(self.cache_dir, 'pulid'), + providers=['CUDAExecutionProvider', 'CPUExecutionProvider'], + ) + self.app.prepare(ctx_id=0, det_size=(640, 640)) + self.handler_ante = insightface.model_zoo.get_model(os.path.join(local_dir, 'glintr100.onnx')) + self.handler_ante.prepare(ctx_id=0) + + self.load_pretrain() + + # other configs + self.debug_img_list = [] + + # karras schedule related code, borrow from lllyasviel/Omost + linear_start = 0.00085 + linear_end = 0.012 + timesteps = 1000 + betas = torch.linspace(linear_start**0.5, linear_end**0.5, timesteps, dtype=torch.float64) ** 2 + alphas = 1.0 - betas + alphas_cumprod = torch.tensor(np.cumprod(alphas, axis=0), dtype=torch.float32) + + self.sigmas = ((1 - alphas_cumprod) / alphas_cumprod) ** 0.5 + self.log_sigmas = self.sigmas.log() + self.sigma_data = 1.0 + + if sampler == 'dpmpp_sde': + self.sampler = sample_dpmpp_sde + elif sampler == 'dpmpp_2m': + self.sampler = sample_dpmpp_2m + else: + raise NotImplementedError(f'sampler {sampler} not implemented') + + @property + def sigma_min(self): + return self.sigmas[0] + + @property + def sigma_max(self): + return self.sigmas[-1] + + def timestep(self, sigma): + log_sigma = sigma.log() + dists = log_sigma.to(self.log_sigmas.device) - self.log_sigmas[:, None] + return dists.abs().argmin(dim=0).view(sigma.shape).to(sigma.device) + + def get_sigmas_karras(self, n, rho=7.0): + ramp = torch.linspace(0, 1, n) + min_inv_rho = self.sigma_min ** (1 / rho) + max_inv_rho = self.sigma_max ** (1 / rho) + sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho + return torch.cat([sigmas, sigmas.new_zeros([1])]) + + def hack_unet_attn_layers(self, unet): + id_adapter_attn_procs = {} + for name, _ in unet.attn_processors.items(): + cross_attention_dim = None if name.endswith("attn1.processor") else unet.config.cross_attention_dim + if name.startswith("mid_block"): + hidden_size = unet.config.block_out_channels[-1] + elif name.startswith("up_blocks"): + block_id = int(name[len("up_blocks.")]) + hidden_size = list(reversed(unet.config.block_out_channels))[block_id] + elif name.startswith("down_blocks"): + block_id = int(name[len("down_blocks.")]) + hidden_size = unet.config.block_out_channels[block_id] + else: + hidden_size = None + if cross_attention_dim is not None: + id_adapter_attn_procs[name] = IDAttnProcessor( + hidden_size=hidden_size, + cross_attention_dim=cross_attention_dim, + ).to(unet.device) + else: + id_adapter_attn_procs[name] = AttnProcessor() + unet.set_attn_processor(id_adapter_attn_procs) + self.id_adapter_attn_layers = nn.ModuleList(unet.attn_processors.values()) + + def load_pretrain(self): + ckpt_path = hf_hub_download('guozinan/PuLID', 'pulid_v1.1.safetensors', local_dir=os.path.join(self.cache_dir, 'pulid')) + state_dict = load_file(ckpt_path) + state_dict_dict = {} + for k, v in state_dict.items(): + module = k.split('.')[0] + state_dict_dict.setdefault(module, {}) + new_k = k[len(module) + 1 :] + state_dict_dict[module][new_k] = v + + for module in state_dict_dict: + print(f'loading from {module}') + getattr(self, module).load_state_dict(state_dict_dict[module], strict=True) + + def to_gray(self, img): + x = 0.299 * img[:, 0:1] + 0.587 * img[:, 1:2] + 0.114 * img[:, 2:3] + x = x.repeat(1, 3, 1, 1) + return x + + def get_id_embedding(self, image_list): + """ + Args: + image in image_list: numpy rgb image, range [0, 255] + """ + id_cond_list = [] + id_vit_hidden_list = [] + for _ii, image in enumerate(image_list): + self.face_helper.clean_all() + image_bgr = cv2.cvtColor(image, cv2.COLOR_RGB2BGR) + # get antelopev2 embedding + face_info = self.app.get(image_bgr) + if len(face_info) > 0: + face_info = sorted( + face_info, key=lambda x: (x['bbox'][2] - x['bbox'][0]) * (x['bbox'][3] - x['bbox'][1]) + )[ + -1 + ] # only use the maximum face + id_ante_embedding = face_info['embedding'] + self.debug_img_list.append( + image[ + int(face_info['bbox'][1]) : int(face_info['bbox'][3]), + int(face_info['bbox'][0]) : int(face_info['bbox'][2]), + ] + ) + else: + id_ante_embedding = None + + # using facexlib to detect and align face + self.face_helper.read_image(image_bgr) + self.face_helper.get_face_landmarks_5(only_center_face=True) + self.face_helper.align_warp_face() + if len(self.face_helper.cropped_faces) == 0: + raise RuntimeError('facexlib align face fail') + align_face = self.face_helper.cropped_faces[0] + # incase insightface didn't detect face + if id_ante_embedding is None: + print('fail to detect face using insightface, extract embedding on align face') + id_ante_embedding = self.handler_ante.get_feat(align_face) + + id_ante_embedding = torch.from_numpy(id_ante_embedding).to(self.device) + if id_ante_embedding.ndim == 1: + id_ante_embedding = id_ante_embedding.unsqueeze(0) + + # parsing + input = img2tensor(align_face, bgr2rgb=True).unsqueeze(0) / 255.0 # pylint: disable=redefined-builtin + input = input.to(self.device) + parsing_out = self.face_helper.face_parse(normalize(input, [0.485, 0.456, 0.406], [0.229, 0.224, 0.225]))[0] + parsing_out = parsing_out.argmax(dim=1, keepdim=True) + bg_label = [0, 16, 18, 7, 8, 9, 14, 15] + bg = sum(parsing_out == i for i in bg_label).bool() + white_image = torch.ones_like(input) + # only keep the face features + face_features_image = torch.where(bg, white_image, self.to_gray(input)) + self.debug_img_list.append(tensor2img(face_features_image, rgb2bgr=False)) + + # transform img before sending to eva-clip-vit + face_features_image = resize( + face_features_image, self.clip_vision_model.image_size, InterpolationMode.BICUBIC + ) + face_features_image = normalize(face_features_image, self.eva_transform_mean, self.eva_transform_std) + id_cond_vit, id_vit_hidden = self.clip_vision_model( + face_features_image, return_all_features=False, return_hidden=True, shuffle=False + ) + id_cond_vit_norm = torch.norm(id_cond_vit, 2, 1, True) + id_cond_vit = torch.div(id_cond_vit, id_cond_vit_norm) + + id_cond = torch.cat([id_ante_embedding, id_cond_vit], dim=-1) + + id_cond_list.append(id_cond) + id_vit_hidden_list.append(id_vit_hidden) + + id_uncond = torch.zeros_like(id_cond_list[0]) + id_vit_hidden_uncond = [] + for layer_idx in range(0, len(id_vit_hidden_list[0])): + id_vit_hidden_uncond.append(torch.zeros_like(id_vit_hidden_list[0][layer_idx])) + + id_cond = torch.stack(id_cond_list, dim=1) + id_vit_hidden = id_vit_hidden_list[0] + for i in range(1, len(image_list)): + for j, x in enumerate(id_vit_hidden_list[i]): + id_vit_hidden[j] = torch.cat([id_vit_hidden[j], x], dim=1) + id_embedding = self.id_adapter(id_cond, id_vit_hidden) + uncond_id_embedding = self.id_adapter(id_uncond, id_vit_hidden_uncond) + + # return id_embedding + return uncond_id_embedding, id_embedding + + def set_progress_bar_config(self, bar_format: str = None, ncols: int = 80, colour: str = None): + import functools + from tqdm.auto import trange as trange_orig + import pulid_utils + pulid_utils.trange = functools.partial(trange_orig, bar_format=bar_format, ncols=ncols, colour=colour) + + def sample(self, x, sigma, **extra_args): + x_ddim_space = x / (sigma[:, None, None, None] ** 2 + self.sigma_data**2) ** 0.5 + t = self.timestep(sigma) + cfg_scale = extra_args['cfg_scale'] + eps_positive = self.pipe.unet(x_ddim_space, t, return_dict=False, **extra_args['positive'])[0] + eps_negative = self.pipe.unet(x_ddim_space, t, return_dict=False, **extra_args['negative'])[0] + noise_pred = eps_negative + cfg_scale * (eps_positive - eps_negative) + latent = x - noise_pred * sigma[:, None, None, None] + if self.callback_on_step_end is not None: + self.step += 1 + self.callback_on_step_end(self.pipe, step=self.step, timestep=t, kwargs={ 'latents': latent }) + return latent + + def init_latent(self, seed, size, image, strength): # pylint: disable=unused-argument + if image is not None and strength > 0: + # TODO pulid img2img + # input can be PIL.Image or np.ndarray so it needs to be converted to rgb tensor + # image must be resized, encoded and noised according to denoising strength + # see below for example from StableDiffusionXLImg2ImgPipeline + latents = None + """ + image = self.image_processor.preprocess(image) + latents = self.prepare_latents( + image, + latent_timestep, + batch_size, + num_images_per_prompt, + prompt_embeds.dtype, + device, + generator, + add_noise, + ) + """ + raise NotImplementedError('pulid: img2img') + else: + # standard txt2img will full noise + latents = torch.randn((size[0], 4, size[1] // 8, size[2] // 8), device="cpu", generator=torch.manual_seed(seed)) + latents = latents.to(dtype=self.pipe.unet.dtype, device=self.device) + return latents + + def __call__( + self, + prompt: str='', + negative_prompt: str='', + width: int=1024, + height: int=1024, + guidance_scale: float=7.0, + num_inference_steps: int=50, + seed: int=-1, + image: np.ndarray=None, + strength: float=0.3, + id_embedding=None, + uncond_id_embedding=None, + id_scale: float=1.0, + callback_on_step_end=None, + ): + self.step = 0 # pylint: disable=attribute-defined-outside-init + self.callback_on_step_end = callback_on_step_end # pylint: disable=attribute-defined-outside-init + size = (1, height, width) + # sigmas + sigmas = self.get_sigmas_karras(num_inference_steps).to(self.device) + + # latents + noise = self.init_latent(seed, size, image, strength) + latents = noise * sigmas[0].to(noise) + + ( + prompt_embeds, + negative_prompt_embeds, + pooled_prompt_embeds, + negative_pooled_prompt_embeds, + ) = self.pipe.encode_prompt( + prompt=prompt, + negative_prompt=negative_prompt, + ) + + add_time_ids = list((size[1], size[2]) + (0, 0) + (size[1], size[2])) + add_time_ids = torch.tensor([add_time_ids], dtype=self.pipe.unet.dtype, device=self.device) + add_neg_time_ids = add_time_ids.clone() + + sampler_kwargs = dict( + cfg_scale=guidance_scale, + positive=dict( + encoder_hidden_states=prompt_embeds, + added_cond_kwargs={"text_embeds": pooled_prompt_embeds, "time_ids": add_time_ids}, + cross_attention_kwargs={'id_embedding': id_embedding, 'id_scale': id_scale}, + ), + negative=dict( + encoder_hidden_states=negative_prompt_embeds, + added_cond_kwargs={"text_embeds": negative_pooled_prompt_embeds, "time_ids": add_neg_time_ids}, + cross_attention_kwargs={'id_embedding': uncond_id_embedding, 'id_scale': id_scale}, + ), + ) + + latents = self.sampler(self.sample, latents, sigmas, extra_args=sampler_kwargs, disable=False) + latents = latents.to(dtype=self.pipe.vae.dtype, device=self.device) / self.pipe.vae.config.scaling_factor + images = self.pipe.vae.decode(latents).sample + images = self.pipe.image_processor.postprocess(images, output_type='pil') + + return images diff --git a/modules/pulid/pulid_utils.py b/modules/pulid/pulid_utils.py new file mode 100644 index 000000000..1a8d3ff06 --- /dev/null +++ b/modules/pulid/pulid_utils.py @@ -0,0 +1,337 @@ +import importlib +import math +import os +import random + +import cv2 +import numpy as np +import torch +import torch.nn.functional as F +import torchsde +from torchvision.utils import make_grid +from tqdm.auto import trange +from transformers import PretrainedConfig + + +def seed_everything(seed): + os.environ["PL_GLOBAL_SEED"] = str(seed) + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + + +def is_torch2_available(): + return hasattr(F, "scaled_dot_product_attention") + + +def instantiate_from_config(config): + if "target" not in config: + if config == '__is_first_stage__' or config == "__is_unconditional__": + return None + raise KeyError("Expected key `target` to instantiate.") + return get_obj_from_str(config["target"])(**config.get("params", {})) + + +def get_obj_from_str(string, reload=False): + module, cls = string.rsplit(".", 1) + if reload: + module_imp = importlib.import_module(module) + importlib.reload(module_imp) + return getattr(importlib.import_module(module, package=None), cls) + + +def drop_seq_token(seq, drop_rate=0.5): + idx = torch.randperm(seq.size(1)) + num_keep_tokens = int(len(idx) * (1 - drop_rate)) + idx = idx[:num_keep_tokens] + seq = seq[:, idx] + return seq + + +def import_model_class_from_model_name_or_path( + pretrained_model_name_or_path: str, revision: str, subfolder: str = "text_encoder" +): + text_encoder_config = PretrainedConfig.from_pretrained( + pretrained_model_name_or_path, subfolder=subfolder, revision=revision + ) + model_class = text_encoder_config.architectures[0] + + if model_class == "CLIPTextModel": + from transformers import CLIPTextModel + + return CLIPTextModel + elif model_class == "CLIPTextModelWithProjection": + from transformers import CLIPTextModelWithProjection + + return CLIPTextModelWithProjection + else: + raise ValueError(f"{model_class} is not supported.") + + +def resize_numpy_image_long(image, resize_long_edge=768): + h, w = image.shape[:2] + if max(h, w) <= resize_long_edge: + return image + k = resize_long_edge / max(h, w) + h = int(h * k) + w = int(w * k) + image = cv2.resize(image, (w, h), interpolation=cv2.INTER_LANCZOS4) + return image + + +# from basicsr +def img2tensor(imgs, bgr2rgb=True, float32=True): + """Numpy array to tensor. + + Args: + imgs (list[ndarray] | ndarray): Input images. + bgr2rgb (bool): Whether to change bgr to rgb. + float32 (bool): Whether to change to float32. + + Returns: + list[tensor] | tensor: Tensor images. If returned results only have + one element, just return tensor. + """ + + def _totensor(img, bgr2rgb, float32): + if img.shape[2] == 3 and bgr2rgb: + if img.dtype == 'float64': + img = img.astype('float32') + img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) + img = torch.from_numpy(img.transpose(2, 0, 1)) + if float32: + img = img.float() + return img + + if isinstance(imgs, list): + return [_totensor(img, bgr2rgb, float32) for img in imgs] + return _totensor(imgs, bgr2rgb, float32) + + +def tensor2img(tensor, rgb2bgr=True, out_type=np.uint8, min_max=(0, 1)): + """Convert torch Tensors into image numpy arrays. + + After clamping to [min, max], values will be normalized to [0, 1]. + + Args: + tensor (Tensor or list[Tensor]): Accept shapes: + 1) 4D mini-batch Tensor of shape (B x 3/1 x H x W); + 2) 3D Tensor of shape (3/1 x H x W); + 3) 2D Tensor of shape (H x W). + Tensor channel should be in RGB order. + rgb2bgr (bool): Whether to change rgb to bgr. + out_type (numpy type): output types. If ``np.uint8``, transform outputs + to uint8 type with range [0, 255]; otherwise, float type with + range [0, 1]. Default: ``np.uint8``. + min_max (tuple[int]): min and max values for clamp. + + Returns: + (Tensor or list): 3D ndarray of shape (H x W x C) OR 2D ndarray of + shape (H x W). The channel order is BGR. + """ + if not (torch.is_tensor(tensor) or (isinstance(tensor, list) and all(torch.is_tensor(t) for t in tensor))): + raise TypeError(f'tensor or list of tensors expected, got {type(tensor)}') + + if torch.is_tensor(tensor): + tensor = [tensor] + result = [] + for _tensor in tensor: + _tensor = _tensor.squeeze(0).float().detach().cpu().clamp_(*min_max) + _tensor = (_tensor - min_max[0]) / (min_max[1] - min_max[0]) + + n_dim = _tensor.dim() + if n_dim == 4: + img_np = make_grid(_tensor, nrow=int(math.sqrt(_tensor.size(0))), normalize=False).numpy() + img_np = img_np.transpose(1, 2, 0) + if rgb2bgr: + img_np = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR) + elif n_dim == 3: + img_np = _tensor.numpy() + img_np = img_np.transpose(1, 2, 0) + if img_np.shape[2] == 1: # gray image + img_np = np.squeeze(img_np, axis=2) + else: + if rgb2bgr: + img_np = cv2.cvtColor(img_np, cv2.COLOR_RGB2BGR) + elif n_dim == 2: + img_np = _tensor.numpy() + else: + raise TypeError(f'Only support 4D, 3D or 2D tensor. But received with dimension: {n_dim}') + if out_type == np.uint8: + # Unlike MATLAB, numpy.unit8() WILL NOT round by default. + img_np = (img_np * 255.0).round() + img_np = img_np.astype(out_type) + result.append(img_np) + if len(result) == 1: + result = result[0] + return result + + +# We didn't find a correct configuration to make the diffusers scheduler align with dpm++2m (karras) in ComfyUI, +# so we copied the ComfyUI code directly. + + +def append_dims(x, target_dims): + """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') + expanded = x[(...,) + (None,) * dims_to_append] + # MPS will get inf values if it tries to index into the new axes, but detaching fixes this. + # https://github.com/pytorch/pytorch/issues/84364 + return expanded.detach().clone() if expanded.device.type == 'mps' else expanded + + +def to_d(x, sigma, denoised): + """Converts a denoiser output to a Karras ODE derivative.""" + return (x - denoised) / append_dims(sigma, x.ndim) + + +def get_ancestral_step(sigma_from, sigma_to, eta=1.0): + """Calculates the noise level (sigma_down) to step down to and the amount + of noise to add (sigma_up) when doing an ancestral sampling step.""" + if not eta: + return sigma_to, 0.0 + sigma_up = min(sigma_to, eta * (sigma_to**2 * (sigma_from**2 - sigma_to**2) / sigma_from**2) ** 0.5) + sigma_down = (sigma_to**2 - sigma_up**2) ** 0.5 + return sigma_down, sigma_up + + +class BatchedBrownianTree: + """A wrapper around torchsde.BrownianTree that enables batches of entropy.""" + + def __init__(self, x, t0, t1, seed=None, **kwargs): + self.cpu_tree = True + if "cpu" in kwargs: + self.cpu_tree = kwargs.pop("cpu") + t0, t1, self.sign = self.sort(t0, t1) + w0 = kwargs.get('w0', torch.zeros_like(x)) + if seed is None: + seed = torch.randint(0, 2**63 - 1, []).item() + self.batched = True + try: + assert len(seed) == x.shape[0] + w0 = w0[0] + except TypeError: + seed = [seed] + self.batched = False + if self.cpu_tree: + self.trees = [torchsde.BrownianTree(t0.cpu(), w0.cpu(), t1.cpu(), entropy=s, **kwargs) for s in seed] + else: + self.trees = [torchsde.BrownianTree(t0, w0, t1, entropy=s, **kwargs) for s in seed] + + @staticmethod + def sort(a, b): + return (a, b, 1) if a < b else (b, a, -1) + + def __call__(self, t0, t1): + t0, t1, sign = self.sort(t0, t1) + if self.cpu_tree: + w = torch.stack( + [tree(t0.cpu().float(), t1.cpu().float()).to(t0.dtype).to(t0.device) for tree in self.trees] + ) * (self.sign * sign) + else: + w = torch.stack([tree(t0, t1) for tree in self.trees]) * (self.sign * sign) + + return w if self.batched else w[0] + + +class BrownianTreeNoiseSampler: + """A noise sampler backed by a torchsde.BrownianTree. + + Args: + x (Tensor): The tensor whose shape, device and dtype to use to generate + random samples. + sigma_min (float): The low end of the valid interval. + sigma_max (float): The high end of the valid interval. + seed (int or List[int]): The random seed. If a list of seeds is + supplied instead of a single integer, then the noise sampler will + use one BrownianTree per batch item, each with its own seed. + transform (callable): A function that maps sigma to the sampler's + internal timestep. + """ + + def __init__(self, x, sigma_min, sigma_max, seed=None, transform=lambda x: x, cpu=False): + self.transform = transform + t0, t1 = self.transform(torch.as_tensor(sigma_min)), self.transform(torch.as_tensor(sigma_max)) + self.tree = BatchedBrownianTree(x, t0, t1, seed, cpu=cpu) + + def __call__(self, sigma, sigma_next): + t0, t1 = self.transform(torch.as_tensor(sigma)), self.transform(torch.as_tensor(sigma_next)) + return self.tree(t0, t1) / (t1 - t0).abs().sqrt() + + +@torch.no_grad() +def sample_dpmpp_2m(model, x, sigmas, extra_args=None, callback=None, disable=None): + """DPM-Solver++(2M).""" + extra_args = {} if extra_args is None else extra_args + s_in = x.new_ones([x.shape[0]]) + sigma_fn = lambda t: t.neg().exp() # pylint: disable=unnecessary-lambda-assignment + t_fn = lambda sigma: sigma.log().neg() # pylint: disable=unnecessary-lambda-assignment + old_denoised = None + + for i in trange(len(sigmas) - 1, disable=disable): + denoised = model(x, sigmas[i] * s_in, **extra_args) + if callback is not None: + callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised}) + t, t_next = t_fn(sigmas[i]), t_fn(sigmas[i + 1]) + h = t_next - t + if old_denoised is None or sigmas[i + 1] == 0: + x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised + else: + h_last = t - t_fn(sigmas[i - 1]) + r = h_last / h + denoised_d = (1 + 1 / (2 * r)) * denoised - (1 / (2 * r)) * old_denoised + x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_d + old_denoised = denoised + return x + + +@torch.no_grad() +def sample_dpmpp_sde( + model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1.0, s_noise=1.0, noise_sampler=None, r=1 / 2 +): + """DPM-Solver++ (stochastic).""" + sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() + seed = extra_args.get("seed", None) + noise_sampler = ( + BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=False) + if noise_sampler is None + else noise_sampler + ) + extra_args = {} if extra_args is None else extra_args + s_in = x.new_ones([x.shape[0]]) + sigma_fn = lambda t: t.neg().exp() # pylint: disable=unnecessary-lambda-assignment + t_fn = lambda sigma: sigma.log().neg() # pylint: disable=unnecessary-lambda-assignment + + for i in trange(len(sigmas) - 1, disable=disable): + denoised = model(x, sigmas[i] * s_in, **extra_args) + if callback is not None: + callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised}) + if sigmas[i + 1] == 0: + # Euler method + d = to_d(x, sigmas[i], denoised) + dt = sigmas[i + 1] - sigmas[i] + x = x + d * dt + else: + # DPM-Solver++ + t, t_next = t_fn(sigmas[i]), t_fn(sigmas[i + 1]) + h = t_next - t + s = t + h * r + fac = 1 / (2 * r) + + # Step 1 + sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(s), eta) + s_ = t_fn(sd) + x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - (t - s_).expm1() * denoised + x_2 = x_2 + noise_sampler(sigma_fn(t), sigma_fn(s)) * s_noise * su + denoised_2 = model(x_2, sigma_fn(s) * s_in, **extra_args) + + # Step 2 + sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(t_next), eta) + t_next_ = t_fn(sd) + denoised_d = (1 - fac) * denoised + fac * denoised_2 + x = (sigma_fn(t_next_) / sigma_fn(t)) * x - (t - t_next_).expm1() * denoised_d + x = x + noise_sampler(sigma_fn(t), sigma_fn(t_next)) * s_noise * su + return x diff --git a/modules/rocm.py b/modules/rocm.py index 831932199..ef76a1cfa 100644 --- a/modules/rocm.py +++ b/modules/rocm.py @@ -52,37 +52,49 @@ class MicroArchitecture(Enum): class Agent: name: str + gfx_version: int arch: MicroArchitecture is_apu: bool if sys.platform != "win32": blaslt_supported: bool + @staticmethod + def parse_gfx_version(name: str) -> int: + result = 0 + for i in range(3, len(name)): + if name[i].isdigit(): + result *= 0x10 + result += ord(name[i]) - 48 + continue + if name[i] in "abcdef": + result *= 0x10 + result += ord(name[i]) - 87 + continue + break + return result + def __init__(self, name: str): self.name = name - gfx = name[3:7] - if len(gfx) == 4: + self.gfx_version = Agent.parse_gfx_version(name) + if self.gfx_version > 0x1000: self.arch = MicroArchitecture.RDNA - elif gfx in ("908", "90a", "942",): + elif self.gfx_version in (0x908, 0x90a, 0x942,): self.arch = MicroArchitecture.CDNA else: self.arch = MicroArchitecture.GCN - self.is_apu = gfx.startswith("115") or gfx in ("801", "902", "90c", "1013", "1033", "1035", "1036", "1103",) + self.is_apu = (self.gfx_version & 0xFFF0 == 0x1150) or self.gfx_version in (0x801, 0x902, 0x90c, 0x1013, 0x1033, 0x1035, 0x1036, 0x1103,) if sys.platform != "win32": self.blaslt_supported = os.path.exists(os.path.join(HIPBLASLT_TENSILE_LIBPATH, f"extop_{name}.co")) def get_gfx_version(self) -> Union[str, None]: - if self.name.startswith("gfx12"): + if self.gfx_version >= 0x1200: return "12.0.0" - elif self.name.startswith("gfx11"): + elif self.gfx_version >= 0x1100: return "11.0.0" - elif self.name.startswith("gfx103"): + elif self.gfx_version >= 0x1000: + # gfx1010 users had to override gfx version to 10.3.0 in Linux + # it is unknown whether overriding is needed in ZLUDA return "10.3.0" - elif self.name.startswith("gfx102"): - return "10.2.0" - elif self.name.startswith("gfx101"): - return "10.1.0" - elif self.name.startswith("gfx100"): - return "10.0.0" return None @@ -198,7 +210,7 @@ else: if os.environ.get("FLASH_ATTENTION_USE_TRITON_ROCM", "FALSE") == "TRUE": return "pytest git+https://github.com/ROCm/flash-attention@micmelesse/upstream_pr" default = "git+https://github.com/ROCm/flash-attention" - if agent.arch == MicroArchitecture.RDNA: + if agent.gfx_version >= 0x1100: default = "git+https://github.com/ROCm/flash-attention@howiejay/navi_support" return os.environ.get("FLASH_ATTENTION_PACKAGE", default) diff --git a/modules/script_loading.py b/modules/script_loading.py index f39394625..37f64b33f 100644 --- a/modules/script_loading.py +++ b/modules/script_loading.py @@ -14,7 +14,7 @@ def load_module(path): module_spec = importlib.util.spec_from_file_location(os.path.basename(path), path) module = importlib.util.module_from_spec(module_spec) try: - if '/sd-extension-' in path or '/Lora' in path: # safe extensions without stdout intercept + if 'sd-extension-' in path or 'Lora' in path: # safe extensions without stdout intercept module_spec.loader.exec_module(module) else: if debug: diff --git a/modules/scripts.py b/modules/scripts.py index 15da9c070..8a67d0a50 100644 --- a/modules/scripts.py +++ b/modules/scripts.py @@ -3,6 +3,7 @@ import re import sys import time from collections import namedtuple +from dataclasses import dataclass import gradio as gr from modules import paths, script_callbacks, extensions, script_loading, scripts_postprocessing, errors, timer @@ -23,6 +24,11 @@ class PostprocessBatchListArgs: self.images = images +@dataclass +class OnComponent: + component: gr.blocks.Block + + class Script: parent = None name = None diff --git a/modules/sd_checkpoint.py b/modules/sd_checkpoint.py new file mode 100644 index 000000000..e1787246f --- /dev/null +++ b/modules/sd_checkpoint.py @@ -0,0 +1,385 @@ +import os +import re +import time +import json +import collections +from modules import shared, paths, modelloader, hashes, sd_hijack_accelerate + + +checkpoints_list = {} +checkpoint_aliases = {} +checkpoints_loaded = collections.OrderedDict() +model_dir = "Stable-diffusion" +model_path = os.path.abspath(os.path.join(paths.models_path, model_dir)) +sd_metadata_file = os.path.join(paths.data_path, "metadata.json") +sd_metadata = None +sd_metadata_pending = 0 +sd_metadata_timer = 0 + + +class CheckpointInfo: + def __init__(self, filename, sha=None): + self.name = None + self.hash = sha + self.filename = filename + self.type = '' + relname = filename + app_path = os.path.abspath(paths.script_path) + + def rel(fn, path): + try: + return os.path.relpath(fn, path) + except Exception: + return fn + + if relname.startswith('..'): + relname = os.path.abspath(relname) + if relname.startswith(shared.opts.ckpt_dir): + relname = rel(filename, shared.opts.ckpt_dir) + elif relname.startswith(shared.opts.diffusers_dir): + relname = rel(filename, shared.opts.diffusers_dir) + elif relname.startswith(model_path): + relname = rel(filename, model_path) + elif relname.startswith(paths.script_path): + relname = rel(filename, paths.script_path) + elif relname.startswith(app_path): + relname = rel(filename, app_path) + else: + relname = os.path.abspath(relname) + relname, ext = os.path.splitext(relname) + ext = ext.lower()[1:] + + if os.path.isfile(filename): # ckpt or safetensor + self.name = relname + self.filename = filename + self.sha256 = hashes.sha256_from_cache(self.filename, f"checkpoint/{relname}") + self.type = ext + if 'nf4' in filename: + self.type = 'transformer' + else: # maybe a diffuser + if self.hash is None: + repo = [r for r in modelloader.diffuser_repos if self.filename == r['name']] + else: + repo = [r for r in modelloader.diffuser_repos if self.hash == r['hash']] + if len(repo) == 0: + self.name = filename + self.filename = filename + self.sha256 = None + self.type = 'unknown' + else: + self.name = os.path.join(os.path.basename(shared.opts.diffusers_dir), repo[0]['name']) + self.filename = repo[0]['path'] + self.sha256 = repo[0]['hash'] + self.type = 'diffusers' + + self.shorthash = self.sha256[0:10] if self.sha256 else None + self.title = self.name if self.shorthash is None else f'{self.name} [{self.shorthash}]' + self.path = self.filename + self.model_name = os.path.basename(self.name) + self.metadata = read_metadata_from_safetensors(filename) + # shared.log.debug(f'Checkpoint: type={self.type} name={self.name} filename={self.filename} hash={self.shorthash} title={self.title}') + + def register(self): + checkpoints_list[self.title] = self + for i in [self.name, self.filename, self.shorthash, self.title]: + if i is not None: + checkpoint_aliases[i] = self + + def calculate_shorthash(self): + self.sha256 = hashes.sha256(self.filename, f"checkpoint/{self.name}") + if self.sha256 is None: + return None + self.shorthash = self.sha256[0:10] + if self.title in checkpoints_list: + checkpoints_list.pop(self.title) + self.title = f'{self.name} [{self.shorthash}]' + self.register() + return self.shorthash + + def __str__(self): + return f'checkpoint: type={self.type} title="{self.title}" path="{self.path}"' + + +def setup_model(): + list_models() + sd_hijack_accelerate.hijack_hfhub() + # sd_hijack_accelerate.hijack_torch_conv() + if not shared.native: + enable_midas_autodownload() + + +def checkpoint_titles(use_short=False): # pylint: disable=unused-argument + def convert(name): + return int(name) if name.isdigit() else name.lower() + def alphanumeric_key(key): + return [convert(c) for c in re.split('([0-9]+)', key)] + return sorted([x.title for x in checkpoints_list.values()], key=alphanumeric_key) + + +def list_models(): + t0 = time.time() + global checkpoints_list # pylint: disable=global-statement + checkpoints_list.clear() + checkpoint_aliases.clear() + ext_filter = [".safetensors"] if shared.opts.sd_disable_ckpt or shared.native else [".ckpt", ".safetensors"] + model_list = list(modelloader.load_models(model_path=model_path, model_url=None, command_path=shared.opts.ckpt_dir, ext_filter=ext_filter, download_name=None, ext_blacklist=[".vae.ckpt", ".vae.safetensors"])) + for filename in sorted(model_list, key=str.lower): + checkpoint_info = CheckpointInfo(filename) + if checkpoint_info.name is not None: + checkpoint_info.register() + if shared.native: + for repo in modelloader.load_diffusers_models(clear=True): + checkpoint_info = CheckpointInfo(repo['name'], sha=repo['hash']) + if checkpoint_info.name is not None: + checkpoint_info.register() + if shared.cmd_opts.ckpt is not None: + if not os.path.exists(shared.cmd_opts.ckpt) and not shared.native: + if shared.cmd_opts.ckpt.lower() != "none": + shared.log.warning(f'Load model: path="{shared.cmd_opts.ckpt}" not found') + else: + checkpoint_info = CheckpointInfo(shared.cmd_opts.ckpt) + if checkpoint_info.name is not None: + checkpoint_info.register() + shared.opts.data['sd_model_checkpoint'] = checkpoint_info.title + elif shared.cmd_opts.ckpt != shared.default_sd_model_file and shared.cmd_opts.ckpt is not None: + shared.log.warning(f'Load model: path="{shared.cmd_opts.ckpt}" not found') + shared.log.info(f'Available Models: path="{shared.opts.ckpt_dir}" items={len(checkpoints_list)} time={time.time()-t0:.2f}') + checkpoints_list = dict(sorted(checkpoints_list.items(), key=lambda cp: cp[1].filename)) + +def update_model_hashes(): + txt = [] + lst = [ckpt for ckpt in checkpoints_list.values() if ckpt.hash is None] + # shared.log.info(f'Models list: short hash missing for {len(lst)} out of {len(checkpoints_list)} models') + for ckpt in lst: + ckpt.hash = model_hash(ckpt.filename) + # txt.append(f'Calculated short hash: {ckpt.title} {ckpt.hash}') + # txt.append(f'Updated short hashes for {len(lst)} out of {len(checkpoints_list)} models') + lst = [ckpt for ckpt in checkpoints_list.values() if ckpt.sha256 is None or ckpt.shorthash is None] + shared.log.info(f'Models list: hash missing={len(lst)} total={len(checkpoints_list)}') + for ckpt in lst: + ckpt.sha256 = hashes.sha256(ckpt.filename, f"checkpoint/{ckpt.name}") + ckpt.shorthash = ckpt.sha256[0:10] if ckpt.sha256 is not None else None + if ckpt.sha256 is not None: + txt.append(f'Hash: {ckpt.title} {ckpt.shorthash}') + txt.append(f'Updated hashes for {len(lst)} out of {len(checkpoints_list)} models') + txt = '
'.join(txt) + return txt + + +def get_closet_checkpoint_match(s: str): + if s.startswith('https://huggingface.co/'): + s = s.replace('https://huggingface.co/', '') + if s.startswith('huggingface/'): + model_name = s.replace('huggingface/', '') + checkpoint_info = CheckpointInfo(model_name) # create a virutal model info + checkpoint_info.type = 'huggingface' + return checkpoint_info + + # alias search + checkpoint_info = checkpoint_aliases.get(s, None) + if checkpoint_info is not None: + return checkpoint_info + + # models search + found = sorted([info for info in checkpoints_list.values() if os.path.basename(info.title).lower().startswith(s.lower())], key=lambda x: len(x.title)) + if found and len(found) == 1: + return found[0] + + # reference search + """ + found = sorted([info for info in shared.reference_models.values() if os.path.basename(info['path']).lower().startswith(s.lower())], key=lambda x: len(x['path'])) + if found and len(found) == 1: + checkpoint_info = CheckpointInfo(found[0]['path']) # create a virutal model info + checkpoint_info.type = 'huggingface' + return checkpoint_info + """ + + # huggingface search + if shared.opts.sd_checkpoint_autodownload and s.count('/') == 1: + modelloader.hf_login() + found = modelloader.find_diffuser(s, full=True) + shared.log.info(f'HF search: model="{s}" results={found}') + if found is not None and len(found) == 1 and found[0] == s: + checkpoint_info = CheckpointInfo(s) + checkpoint_info.type = 'huggingface' + return checkpoint_info + + # civitai search + if shared.opts.sd_checkpoint_autodownload and s.startswith("https://civitai.com/api/download/models"): + fn = modelloader.download_civit_model_thread(model_name=None, model_url=s, model_path='', model_type='Model', token=None) + if fn is not None: + checkpoint_info = CheckpointInfo(fn) + return checkpoint_info + + return None + + +def model_hash(filename): + """old hash that only looks at a small part of the file and is prone to collisions""" + try: + with open(filename, "rb") as file: + import hashlib + # t0 = time.time() + m = hashlib.sha256() + file.seek(0x100000) + m.update(file.read(0x10000)) + shorthash = m.hexdigest()[0:8] + # t1 = time.time() + # shared.log.debug(f'Calculating short hash: {filename} hash={shorthash} time={(t1-t0):.2f}') + return shorthash + except FileNotFoundError: + return 'NOFILE' + except Exception: + return 'NOHASH' + + +def select_checkpoint(op='model'): + if op == 'dict': + model_checkpoint = shared.opts.sd_model_dict + elif op == 'refiner': + model_checkpoint = shared.opts.data.get('sd_model_refiner', None) + else: + model_checkpoint = shared.opts.sd_model_checkpoint + if model_checkpoint is None or model_checkpoint == 'None': + return None + checkpoint_info = get_closet_checkpoint_match(model_checkpoint) + if checkpoint_info is not None: + shared.log.info(f'Load {op}: select="{checkpoint_info.title if checkpoint_info is not None else None}"') + return checkpoint_info + if len(checkpoints_list) == 0: + shared.log.warning("Cannot generate without a checkpoint") + shared.log.info("Set system paths to use existing folders") + shared.log.info(" or use --models-dir to specify base folder with all models") + shared.log.info(" or use --ckpt-dir to specify folder with sd models") + shared.log.info(" or use --ckpt to force using specific model") + return None + # checkpoint_info = next(iter(checkpoints_list.values())) + if model_checkpoint is not None: + if model_checkpoint != 'model.safetensors' and model_checkpoint != 'stabilityai/stable-diffusion-xl-base-1.0': + shared.log.info(f'Load {op}: search="{model_checkpoint}" not found') + else: + shared.log.info("Selecting first available checkpoint") + # shared.log.warning(f"Loading fallback checkpoint: {checkpoint_info.title}") + # shared.opts.data['sd_model_checkpoint'] = checkpoint_info.title + else: + shared.log.info(f'Load {op}: select="{checkpoint_info.title if checkpoint_info is not None else None}"') + return checkpoint_info + + +def read_metadata_from_safetensors(filename): + global sd_metadata # pylint: disable=global-statement + if sd_metadata is None: + sd_metadata = shared.readfile(sd_metadata_file, lock=True) if os.path.isfile(sd_metadata_file) else {} + res = sd_metadata.get(filename, None) + if res is not None: + return res + if not filename.endswith(".safetensors"): + return {} + if shared.cmd_opts.no_metadata: + return {} + res = {} + # try: + t0 = time.time() + with open(filename, mode="rb") as file: + try: + metadata_len = file.read(8) + metadata_len = int.from_bytes(metadata_len, "little") + json_start = file.read(2) + if metadata_len <= 2 or json_start not in (b'{"', b"{'"): + shared.log.error(f'Model metadata invalid: file="{filename}"') + json_data = json_start + file.read(metadata_len-2) + json_obj = json.loads(json_data) + for k, v in json_obj.get("__metadata__", {}).items(): + if v.startswith("data:"): + v = 'data' + if k == 'format' and v == 'pt': + continue + large = True if len(v) > 2048 else False + if large and k == 'ss_datasets': + continue + if large and k == 'workflow': + continue + if large and k == 'prompt': + continue + if large and k == 'ss_bucket_info': + continue + if v[0:1] == '{': + try: + v = json.loads(v) + if large and k == 'ss_tag_frequency': + v = { i: len(j) for i, j in v.items() } + if large and k == 'sd_merge_models': + scrub_dict(v, ['sd_merge_recipe']) + except Exception: + pass + res[k] = v + except Exception as e: + shared.log.error(f'Model metadata: file="{filename}" {e}') + sd_metadata[filename] = res + global sd_metadata_pending # pylint: disable=global-statement + sd_metadata_pending += 1 + t1 = time.time() + global sd_metadata_timer # pylint: disable=global-statement + sd_metadata_timer += (t1 - t0) + # except Exception as e: + # shared.log.error(f"Error reading metadata from: {filename} {e}") + return res + + +def enable_midas_autodownload(): + """ + Gives the ldm.modules.midas.api.load_model function automatic downloading. + + When the 512-depth-ema model, and other future models like it, is loaded, + it calls midas.api.load_model to load the associated midas depth model. + This function applies a wrapper to download the model to the correct + location automatically. + """ + from urllib import request + import ldm.modules.midas.api + midas_path = os.path.join(paths.models_path, 'midas') + for k, v in ldm.modules.midas.api.ISL_PATHS.items(): + file_name = os.path.basename(v) + ldm.modules.midas.api.ISL_PATHS[k] = os.path.join(midas_path, file_name) + midas_urls = { + "dpt_large": "https://github.com/intel-isl/DPT/releases/download/1_0/dpt_large-midas-2f21e586.pt", + "dpt_hybrid": "https://github.com/intel-isl/DPT/releases/download/1_0/dpt_hybrid-midas-501f0c75.pt", + "midas_v21": "https://github.com/AlexeyAB/MiDaS/releases/download/midas_dpt/midas_v21-f6b98070.pt", + "midas_v21_small": "https://github.com/AlexeyAB/MiDaS/releases/download/midas_dpt/midas_v21_small-70d6b9c8.pt", + } + ldm.modules.midas.api.load_model_inner = ldm.modules.midas.api.load_model + + def load_model_wrapper(model_type): + path = ldm.modules.midas.api.ISL_PATHS[model_type] + if not os.path.exists(path): + if not os.path.exists(midas_path): + os.mkdir(midas_path) + shared.log.info(f"Downloading midas model weights for {model_type} to {path}") + request.urlretrieve(midas_urls[model_type], path) + shared.log.info(f"{model_type} downloaded") + return ldm.modules.midas.api.load_model_inner(model_type) + + ldm.modules.midas.api.load_model = load_model_wrapper + + +def scrub_dict(dict_obj, keys): + for key in list(dict_obj.keys()): + if not isinstance(dict_obj, dict): + continue + if key in keys: + dict_obj.pop(key, None) + elif isinstance(dict_obj[key], dict): + scrub_dict(dict_obj[key], keys) + elif isinstance(dict_obj[key], list): + for item in dict_obj[key]: + scrub_dict(item, keys) + + +def write_metadata(): + global sd_metadata_pending # pylint: disable=global-statement + if sd_metadata_pending == 0: + shared.log.debug(f'Model metadata: file="{sd_metadata_file}" no changes') + return + shared.writefile(sd_metadata, sd_metadata_file) + shared.log.info(f'Model metadata saved: file="{sd_metadata_file}" items={sd_metadata_pending} time={sd_metadata_timer:.2f}') + sd_metadata_pending = 0 diff --git a/modules/sd_detect.py b/modules/sd_detect.py new file mode 100644 index 000000000..31f773607 --- /dev/null +++ b/modules/sd_detect.py @@ -0,0 +1,163 @@ +import os +import time +import torch +import diffusers +from modules import shared, shared_items, devices, errors, model_tools + + +debug_load = os.environ.get('SD_LOAD_DEBUG', None) + + +def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False): + guess = shared.opts.diffusers_pipeline + warn = shared.log.warning if warning else lambda *args, **kwargs: None + size = 0 + pipeline = None + if guess == 'Autodetect': + try: + guess = 'Stable Diffusion XL' if 'XL' in f.upper() else 'Stable Diffusion' + # guess by size + if os.path.isfile(f) and f.endswith('.safetensors'): + size = round(os.path.getsize(f) / 1024 / 1024) + if (size > 0 and size < 128): + warn(f'Model size smaller than expected: {f} size={size} MB') + elif (size >= 316 and size <= 324) or (size >= 156 and size <= 164): # 320 or 160 + warn(f'Model detected as VAE model, but attempting to load as model: {op}={f} size={size} MB') + guess = 'VAE' + elif (size >= 4970 and size <= 4976): # 4973 + guess = 'Stable Diffusion 2' # SD v2 but could be eps or v-prediction + # elif size < 0: # unknown + # guess = 'Stable Diffusion 2B' + elif (size >= 5791 and size <= 5799): # 5795 + if op == 'model': + warn(f'Model detected as SD-XL refiner model, but attempting to load a base model: {op}={f} size={size} MB') + guess = 'Stable Diffusion XL Refiner' + elif (size >= 6611 and size <= 7220): # 6617, HassakuXL is 6776, monkrenRealisticINT_v10 is 7217 + guess = 'Stable Diffusion XL' + elif (size >= 3361 and size <= 3369): # 3368 + guess = 'Stable Diffusion Upscale' + elif (size >= 4891 and size <= 4899): # 4897 + guess = 'Stable Diffusion XL Inpaint' + elif (size >= 9791 and size <= 9799): # 9794 + guess = 'Stable Diffusion XL Instruct' + elif (size > 3138 and size < 3142): #3140 + guess = 'Stable Diffusion XL' + elif (size > 5692 and size < 5698) or (size > 4134 and size < 4138) or (size > 10362 and size < 10366) or (size > 15028 and size < 15228): + guess = 'Stable Diffusion 3' + elif (size > 18414 and size < 18420): # sd35-large aio + guess = 'Stable Diffusion 3' + elif (size > 20000 and size < 40000): + guess = 'FLUX' + # guess by name + """ + if 'LCM_' in f.upper() or 'LCM-' in f.upper() or '_LCM' in f.upper() or '-LCM' in f.upper(): + if shared.backend == shared.Backend.ORIGINAL: + warn(f'Model detected as LCM model, but attempting to load using backend=original: {op}={f} size={size} MB') + guess = 'Latent Consistency Model' + """ + if 'instaflow' in f.lower(): + guess = 'InstaFlow' + if 'segmoe' in f.lower(): + guess = 'SegMoE' + if 'hunyuandit' in f.lower(): + guess = 'HunyuanDiT' + if 'pixart-xl' in f.lower(): + guess = 'PixArt-Alpha' + if 'stable-diffusion-3' in f.lower(): + guess = 'Stable Diffusion 3' + if 'stable-cascade' in f.lower() or 'stablecascade' in f.lower() or 'wuerstchen3' in f.lower() or ('sotediffusion' in f.lower() and "v2" in f.lower()): + if devices.dtype == torch.float16: + warn('Stable Cascade does not support Float16') + guess = 'Stable Cascade' + if 'pixart-sigma' in f.lower(): + guess = 'PixArt-Sigma' + if 'lumina-next' in f.lower(): + guess = 'Lumina-Next' + if 'kolors' in f.lower(): + guess = 'Kolors' + if 'auraflow' in f.lower(): + guess = 'AuraFlow' + if 'cogview' in f.lower(): + guess = 'CogView' + if 'meissonic' in f.lower(): + guess = 'Meissonic' + pipeline = 'custom' + if 'monetico' in f.lower(): + guess = 'Monetico' + pipeline = 'custom' + if 'omnigen' in f.lower(): + guess = 'OmniGen' + pipeline = 'custom' + if 'sd3' in f.lower(): + guess = 'Stable Diffusion 3' + if 'flux' in f.lower(): + guess = 'FLUX' + if size > 11000 and size < 20000: + warn(f'Model detected as FLUX UNET model, but attempting to load a base model: {op}={f} size={size} MB') + # switch for specific variant + if guess == 'Stable Diffusion' and 'inpaint' in f.lower(): + guess = 'Stable Diffusion Inpaint' + elif guess == 'Stable Diffusion' and 'instruct' in f.lower(): + guess = 'Stable Diffusion Instruct' + if guess == 'Stable Diffusion XL' and 'inpaint' in f.lower(): + guess = 'Stable Diffusion XL Inpaint' + elif guess == 'Stable Diffusion XL' and 'instruct' in f.lower(): + guess = 'Stable Diffusion XL Instruct' + # get actual pipeline + pipeline = shared_items.get_pipelines().get(guess, None) if pipeline is None else pipeline + if not quiet: + shared.log.info(f'Autodetect {op}: detect="{guess}" class={getattr(pipeline, "__name__", None)} file="{f}" size={size}MB') + t0 = time.time() + keys = model_tools.get_safetensor_keys(f) + if keys is not None and len(keys) > 0: + modules = model_tools.list_to_dict(keys) + modules = model_tools.remove_entries_after_depth(modules, 3) + t1 = time.time() + shared.log.debug(f'Autodetect modules: {modules} time={t1-t0:.2f}') + except Exception as e: + shared.log.error(f'Autodetect {op}: file="{f}" {e}') + if debug_load: + errors.display(e, f'Load {op}: {f}') + return None, None + else: + try: + size = round(os.path.getsize(f) / 1024 / 1024) + pipeline = shared_items.get_pipelines().get(guess, None) if pipeline is None else pipeline + if not quiet: + shared.log.info(f'Load {op}: detect="{guess}" class={getattr(pipeline, "__name__", None)} file="{f}" size={size}MB') + except Exception as e: + shared.log.error(f'Load {op}: detect="{guess}" file="{f}" {e}') + + if pipeline is None: + shared.log.warning(f'Load {op}: detect="{guess}" file="{f}" size={size} not recognized') + pipeline = diffusers.StableDiffusionPipeline + return pipeline, guess + + +def get_load_config(model_file, model_type, config_type='yaml'): + if config_type == 'yaml': + yaml = os.path.splitext(model_file)[0] + '.yaml' + if os.path.exists(yaml): + return yaml + if model_type == 'Stable Diffusion': + return 'configs/v1-inference.yaml' + if model_type == 'Stable Diffusion XL': + return 'configs/sd_xl_base.yaml' + if model_type == 'Stable Diffusion XL Refiner': + return 'configs/sd_xl_refiner.yaml' + if model_type == 'Stable Diffusion 2': + return None # dont know if its eps or v so let diffusers sort it out + # return 'configs/v2-inference-512-base.yaml' + # return 'configs/v2-inference-768-v.yaml' + elif config_type == 'json': + if not shared.opts.diffuser_cache_config: + return None + if model_type == 'Stable Diffusion': + return 'configs/sd15' + if model_type == 'Stable Diffusion XL': + return 'configs/sdxl' + if model_type == 'Stable Diffusion 3': + return 'configs/sd3' + if model_type == 'FLUX': + return 'configs/flux' + return None diff --git a/modules/sd_models.py b/modules/sd_models.py index 11f601260..3ca081d6c 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -1,16 +1,12 @@ -import re import io import sys -import json import time +import json import copy import inspect import logging import contextlib -import collections import os.path -from os import mkdir -from urllib import request from enum import Enum import diffusers import diffusers.loaders.single_file_utils @@ -18,20 +14,16 @@ from rich import progress # pylint: disable=redefined-builtin import torch import safetensors.torch from omegaconf import OmegaConf -from transformers import logging as transformers_logging from ldm.util import instantiate_from_config -from modules import paths, shared, shared_items, shared_state, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, hashes, sd_models_config, sd_models_compile, sd_hijack_accelerate +from modules import paths, shared, shared_state, modelloader, devices, script_callbacks, sd_vae, sd_unet, errors, sd_models_config, sd_models_compile, sd_hijack_accelerate, sd_detect from modules.timer import Timer from modules.memstats import memory_stats from modules.modeldata import model_data +from modules.sd_checkpoint import CheckpointInfo, select_checkpoint, list_models, checkpoints_list, checkpoint_titles, get_closet_checkpoint_match, model_hash, update_model_hashes, setup_model, write_metadata, read_metadata_from_safetensors # pylint: disable=unused-import -transformers_logging.set_verbosity_error() model_dir = "Stable-diffusion" model_path = os.path.abspath(os.path.join(paths.models_path, model_dir)) -checkpoints_list = {} -checkpoint_aliases = {} -checkpoints_loaded = collections.OrderedDict() sd_metadata_file = os.path.join(paths.data_path, "metadata.json") sd_metadata = None sd_metadata_pending = 0 @@ -40,86 +32,7 @@ debug_move = shared.log.trace if os.environ.get('SD_MOVE_DEBUG', None) is not No debug_load = os.environ.get('SD_LOAD_DEBUG', None) debug_process = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None diffusers_version = int(diffusers.__version__.split('.')[1]) - - -class CheckpointInfo: - def __init__(self, filename, sha=None): - self.name = None - self.hash = sha - self.filename = filename - self.type = '' - relname = filename - app_path = os.path.abspath(paths.script_path) - - def rel(fn, path): - try: - return os.path.relpath(fn, path) - except Exception: - return fn - - if relname.startswith('..'): - relname = os.path.abspath(relname) - if relname.startswith(shared.opts.ckpt_dir): - relname = rel(filename, shared.opts.ckpt_dir) - elif relname.startswith(shared.opts.diffusers_dir): - relname = rel(filename, shared.opts.diffusers_dir) - elif relname.startswith(model_path): - relname = rel(filename, model_path) - elif relname.startswith(paths.script_path): - relname = rel(filename, paths.script_path) - elif relname.startswith(app_path): - relname = rel(filename, app_path) - else: - relname = os.path.abspath(relname) - relname, ext = os.path.splitext(relname) - ext = ext.lower()[1:] - - if os.path.isfile(filename): # ckpt or safetensor - self.name = relname - self.filename = filename - self.sha256 = hashes.sha256_from_cache(self.filename, f"checkpoint/{relname}") - self.type = ext - if 'nf4' in filename: - self.type = 'transformer' - else: # maybe a diffuser - if self.hash is None: - repo = [r for r in modelloader.diffuser_repos if self.filename == r['name']] - else: - repo = [r for r in modelloader.diffuser_repos if self.hash == r['hash']] - if len(repo) == 0: - self.name = filename - self.filename = filename - self.sha256 = None - self.type = 'unknown' - else: - self.name = os.path.join(os.path.basename(shared.opts.diffusers_dir), repo[0]['name']) - self.filename = repo[0]['path'] - self.sha256 = repo[0]['hash'] - self.type = 'diffusers' - - self.shorthash = self.sha256[0:10] if self.sha256 else None - self.title = self.name if self.shorthash is None else f'{self.name} [{self.shorthash}]' - self.path = self.filename - self.model_name = os.path.basename(self.name) - self.metadata = read_metadata_from_safetensors(filename) - # shared.log.debug(f'Checkpoint: type={self.type} name={self.name} filename={self.filename} hash={self.shorthash} title={self.title}') - - def register(self): - checkpoints_list[self.title] = self - for i in [self.name, self.filename, self.shorthash, self.title]: - if i is not None: - checkpoint_aliases[i] = self - - def calculate_shorthash(self): - self.sha256 = hashes.sha256(self.filename, f"checkpoint/{self.name}") - if self.sha256 is None: - return None - self.shorthash = self.sha256[0:10] - if self.title in checkpoints_list: - checkpoints_list.pop(self.title) - self.title = f'{self.name} [{self.shorthash}]' - self.register() - return self.shorthash +checkpoint_tiles = checkpoint_titles # legacy compatibility class NoWatermark: @@ -127,283 +40,6 @@ class NoWatermark: return img -def setup_model(): - list_models() - sd_hijack_accelerate.hijack_hfhub() - # sd_hijack_accelerate.hijack_torch_conv() - if not shared.native: - enable_midas_autodownload() - - -def checkpoint_tiles(use_short=False): # pylint: disable=unused-argument - def convert(name): - return int(name) if name.isdigit() else name.lower() - def alphanumeric_key(key): - return [convert(c) for c in re.split('([0-9]+)', key)] - return sorted([x.title for x in checkpoints_list.values()], key=alphanumeric_key) - - -def list_models(): - t0 = time.time() - global checkpoints_list # pylint: disable=global-statement - checkpoints_list.clear() - checkpoint_aliases.clear() - ext_filter = [".safetensors"] if shared.opts.sd_disable_ckpt or shared.native else [".ckpt", ".safetensors"] - model_list = list(modelloader.load_models(model_path=model_path, model_url=None, command_path=shared.opts.ckpt_dir, ext_filter=ext_filter, download_name=None, ext_blacklist=[".vae.ckpt", ".vae.safetensors"])) - for filename in sorted(model_list, key=str.lower): - checkpoint_info = CheckpointInfo(filename) - if checkpoint_info.name is not None: - checkpoint_info.register() - if shared.native: - for repo in modelloader.load_diffusers_models(clear=True): - checkpoint_info = CheckpointInfo(repo['name'], sha=repo['hash']) - if checkpoint_info.name is not None: - checkpoint_info.register() - if shared.cmd_opts.ckpt is not None: - if not os.path.exists(shared.cmd_opts.ckpt) and not shared.native: - if shared.cmd_opts.ckpt.lower() != "none": - shared.log.warning(f'Load model: path="{shared.cmd_opts.ckpt}" not found') - else: - checkpoint_info = CheckpointInfo(shared.cmd_opts.ckpt) - if checkpoint_info.name is not None: - checkpoint_info.register() - shared.opts.data['sd_model_checkpoint'] = checkpoint_info.title - elif shared.cmd_opts.ckpt != shared.default_sd_model_file and shared.cmd_opts.ckpt is not None: - shared.log.warning(f'Load model: path="{shared.cmd_opts.ckpt}" not found') - shared.log.info(f'Available Models: path="{shared.opts.ckpt_dir}" items={len(checkpoints_list)} time={time.time()-t0:.2f}') - checkpoints_list = dict(sorted(checkpoints_list.items(), key=lambda cp: cp[1].filename)) - - -def update_model_hashes(): - txt = [] - lst = [ckpt for ckpt in checkpoints_list.values() if ckpt.hash is None] - # shared.log.info(f'Models list: short hash missing for {len(lst)} out of {len(checkpoints_list)} models') - for ckpt in lst: - ckpt.hash = model_hash(ckpt.filename) - # txt.append(f'Calculated short hash: {ckpt.title} {ckpt.hash}') - # txt.append(f'Updated short hashes for {len(lst)} out of {len(checkpoints_list)} models') - lst = [ckpt for ckpt in checkpoints_list.values() if ckpt.sha256 is None or ckpt.shorthash is None] - shared.log.info(f'Models list: hash missing={len(lst)} total={len(checkpoints_list)}') - for ckpt in lst: - ckpt.sha256 = hashes.sha256(ckpt.filename, f"checkpoint/{ckpt.name}") - ckpt.shorthash = ckpt.sha256[0:10] if ckpt.sha256 is not None else None - if ckpt.sha256 is not None: - txt.append(f'Hash: {ckpt.title} {ckpt.shorthash}') - txt.append(f'Updated hashes for {len(lst)} out of {len(checkpoints_list)} models') - txt = '
'.join(txt) - return txt - - -def get_closet_checkpoint_match(s: str): - if s.startswith('https://huggingface.co/'): - s = s.replace('https://huggingface.co/', '') - if s.startswith('huggingface/'): - model_name = s.replace('huggingface/', '') - checkpoint_info = CheckpointInfo(model_name) # create a virutal model info - checkpoint_info.type = 'huggingface' - return checkpoint_info - - # alias search - checkpoint_info = checkpoint_aliases.get(s, None) - if checkpoint_info is not None: - return checkpoint_info - - # models search - found = sorted([info for info in checkpoints_list.values() if os.path.basename(info.title).lower().startswith(s.lower())], key=lambda x: len(x.title)) - if found and len(found) == 1: - return found[0] - - # reference search - """ - found = sorted([info for info in shared.reference_models.values() if os.path.basename(info['path']).lower().startswith(s.lower())], key=lambda x: len(x['path'])) - if found and len(found) == 1: - checkpoint_info = CheckpointInfo(found[0]['path']) # create a virutal model info - checkpoint_info.type = 'huggingface' - return checkpoint_info - """ - - # huggingface search - if shared.opts.sd_checkpoint_autodownload and s.count('/') == 1: - modelloader.hf_login() - found = modelloader.find_diffuser(s, full=True) - shared.log.info(f'HF search: model="{s}" results={found}') - if found is not None and len(found) == 1 and found[0] == s: - checkpoint_info = CheckpointInfo(s) - checkpoint_info.type = 'huggingface' - return checkpoint_info - - # civitai search - if shared.opts.sd_checkpoint_autodownload and s.startswith("https://civitai.com/api/download/models"): - fn = modelloader.download_civit_model_thread(model_name=None, model_url=s, model_path='', model_type='Model', token=None) - if fn is not None: - checkpoint_info = CheckpointInfo(fn) - return checkpoint_info - - return None - - -def model_hash(filename): - """old hash that only looks at a small part of the file and is prone to collisions""" - try: - with open(filename, "rb") as file: - import hashlib - # t0 = time.time() - m = hashlib.sha256() - file.seek(0x100000) - m.update(file.read(0x10000)) - shorthash = m.hexdigest()[0:8] - # t1 = time.time() - # shared.log.debug(f'Calculating short hash: {filename} hash={shorthash} time={(t1-t0):.2f}') - return shorthash - except FileNotFoundError: - return 'NOFILE' - except Exception: - return 'NOHASH' - - -def select_checkpoint(op='model'): - if op == 'dict': - model_checkpoint = shared.opts.sd_model_dict - elif op == 'refiner': - model_checkpoint = shared.opts.data.get('sd_model_refiner', None) - else: - model_checkpoint = shared.opts.sd_model_checkpoint - if model_checkpoint is None or model_checkpoint == 'None': - return None - checkpoint_info = get_closet_checkpoint_match(model_checkpoint) - if checkpoint_info is not None: - shared.log.info(f'Load {op}: select="{checkpoint_info.title if checkpoint_info is not None else None}"') - return checkpoint_info - if len(checkpoints_list) == 0: - shared.log.warning("Cannot generate without a checkpoint") - shared.log.info("Set system paths to use existing folders") - shared.log.info(" or use --models-dir to specify base folder with all models") - shared.log.info(" or use --ckpt-dir to specify folder with sd models") - shared.log.info(" or use --ckpt to force using specific model") - return None - # checkpoint_info = next(iter(checkpoints_list.values())) - if model_checkpoint is not None: - if model_checkpoint != 'model.safetensors' and model_checkpoint != 'stabilityai/stable-diffusion-xl-base-1.0': - shared.log.info(f'Load {op}: search="{model_checkpoint}" not found') - else: - shared.log.info("Selecting first available checkpoint") - # shared.log.warning(f"Loading fallback checkpoint: {checkpoint_info.title}") - # shared.opts.data['sd_model_checkpoint'] = checkpoint_info.title - else: - shared.log.info(f'Load {op}: select="{checkpoint_info.title if checkpoint_info is not None else None}"') - return checkpoint_info - - -checkpoint_dict_replacements = { - 'cond_stage_model.transformer.embeddings.': 'cond_stage_model.transformer.text_model.embeddings.', - 'cond_stage_model.transformer.encoder.': 'cond_stage_model.transformer.text_model.encoder.', - 'cond_stage_model.transformer.final_layer_norm.': 'cond_stage_model.transformer.text_model.final_layer_norm.', -} - - -def transform_checkpoint_dict_key(k): - for text, replacement in checkpoint_dict_replacements.items(): - if k.startswith(text): - k = replacement + k[len(text):] - return k - - -def get_state_dict_from_checkpoint(pl_sd): - pl_sd = pl_sd.pop("state_dict", pl_sd) - pl_sd.pop("state_dict", None) - sd = {} - for k, v in pl_sd.items(): - new_key = transform_checkpoint_dict_key(k) - if new_key is not None: - sd[new_key] = v - pl_sd.clear() - pl_sd.update(sd) - return pl_sd - - -def write_metadata(): - global sd_metadata_pending # pylint: disable=global-statement - if sd_metadata_pending == 0: - shared.log.debug(f'Model metadata: file="{sd_metadata_file}" no changes') - return - shared.writefile(sd_metadata, sd_metadata_file) - shared.log.info(f'Model metadata saved: file="{sd_metadata_file}" items={sd_metadata_pending} time={sd_metadata_timer:.2f}') - sd_metadata_pending = 0 - - -def scrub_dict(dict_obj, keys): - for key in list(dict_obj.keys()): - if not isinstance(dict_obj, dict): - continue - if key in keys: - dict_obj.pop(key, None) - elif isinstance(dict_obj[key], dict): - scrub_dict(dict_obj[key], keys) - elif isinstance(dict_obj[key], list): - for item in dict_obj[key]: - scrub_dict(item, keys) - - -def read_metadata_from_safetensors(filename): - global sd_metadata # pylint: disable=global-statement - if sd_metadata is None: - sd_metadata = shared.readfile(sd_metadata_file, lock=True) if os.path.isfile(sd_metadata_file) else {} - res = sd_metadata.get(filename, None) - if res is not None: - return res - if not filename.endswith(".safetensors"): - return {} - if shared.cmd_opts.no_metadata: - return {} - res = {} - # try: - t0 = time.time() - with open(filename, mode="rb") as file: - try: - metadata_len = file.read(8) - metadata_len = int.from_bytes(metadata_len, "little") - json_start = file.read(2) - if metadata_len <= 2 or json_start not in (b'{"', b"{'"): - shared.log.error(f'Model metadata invalid: file="{filename}"') - json_data = json_start + file.read(metadata_len-2) - json_obj = json.loads(json_data) - for k, v in json_obj.get("__metadata__", {}).items(): - if v.startswith("data:"): - v = 'data' - if k == 'format' and v == 'pt': - continue - large = True if len(v) > 2048 else False - if large and k == 'ss_datasets': - continue - if large and k == 'workflow': - continue - if large and k == 'prompt': - continue - if large and k == 'ss_bucket_info': - continue - if v[0:1] == '{': - try: - v = json.loads(v) - if large and k == 'ss_tag_frequency': - v = { i: len(j) for i, j in v.items() } - if large and k == 'sd_merge_models': - scrub_dict(v, ['sd_merge_recipe']) - except Exception: - pass - res[k] = v - except Exception as e: - shared.log.error(f'Model metadata: file="{filename}" {e}') - sd_metadata[filename] = res - global sd_metadata_pending # pylint: disable=global-statement - sd_metadata_pending += 1 - t1 = time.time() - global sd_metadata_timer # pylint: disable=global-statement - sd_metadata_timer += (t1 - t0) - # except Exception as e: - # shared.log.error(f"Error reading metadata from: {filename} {e}") - return res - - def read_state_dict(checkpoint_file, map_location=None, what:str='model'): # pylint: disable=unused-argument if not os.path.isfile(checkpoint_file): shared.log.error(f'Load dict: path="{checkpoint_file}" not a file') @@ -439,36 +75,55 @@ def read_state_dict(checkpoint_file, map_location=None, what:str='model'): # pyl return sd -def get_safetensor_keys(filename): - keys = [] - try: - with safetensors.torch.safe_open(filename, framework="pt", device="cpu") as f: - keys = f.keys() - except Exception as e: - shared.log.error(f'Load dict: path="{filename}" {e}') - return keys +def get_state_dict_from_checkpoint(pl_sd): + checkpoint_dict_replacements = { + 'cond_stage_model.transformer.embeddings.': 'cond_stage_model.transformer.text_model.embeddings.', + 'cond_stage_model.transformer.encoder.': 'cond_stage_model.transformer.text_model.encoder.', + 'cond_stage_model.transformer.final_layer_norm.': 'cond_stage_model.transformer.text_model.final_layer_norm.', + } + + def transform_checkpoint_dict_key(k): + for text, replacement in checkpoint_dict_replacements.items(): + if k.startswith(text): + k = replacement + k[len(text):] + return k + + pl_sd = pl_sd.pop("state_dict", pl_sd) + pl_sd.pop("state_dict", None) + sd = {} + for k, v in pl_sd.items(): + new_key = transform_checkpoint_dict_key(k) + if new_key is not None: + sd[new_key] = v + pl_sd.clear() + pl_sd.update(sd) + return pl_sd def get_checkpoint_state_dict(checkpoint_info: CheckpointInfo, timer): if not os.path.isfile(checkpoint_info.filename): return None + """ if checkpoint_info in checkpoints_loaded: shared.log.info("Load model: cache") checkpoints_loaded.move_to_end(checkpoint_info, last=True) # FIFO -> LRU cache return checkpoints_loaded[checkpoint_info] + """ res = read_state_dict(checkpoint_info.filename, what='model') + """ if shared.opts.sd_checkpoint_cache > 0 and not shared.native: # cache newly loaded model checkpoints_loaded[checkpoint_info] = res # clean up cache if limit is reached while len(checkpoints_loaded) > shared.opts.sd_checkpoint_cache: checkpoints_loaded.popitem(last=False) + """ timer.record("load") return res def load_model_weights(model: torch.nn.Module, checkpoint_info: CheckpointInfo, state_dict, timer): - _pipeline, _model_type = detect_pipeline(checkpoint_info.path, 'model') + _pipeline, _model_type = sd_detect.detect_pipeline(checkpoint_info.path, 'model') shared.log.debug(f'Load model: memory={memory_stats()}') timer.record("hash") if model_data.sd_dict == 'None': @@ -520,41 +175,6 @@ def load_model_weights(model: torch.nn.Module, checkpoint_info: CheckpointInfo, return True -def enable_midas_autodownload(): - """ - Gives the ldm.modules.midas.api.load_model function automatic downloading. - - When the 512-depth-ema model, and other future models like it, is loaded, - it calls midas.api.load_model to load the associated midas depth model. - This function applies a wrapper to download the model to the correct - location automatically. - """ - import ldm.modules.midas.api - midas_path = os.path.join(paths.models_path, 'midas') - for k, v in ldm.modules.midas.api.ISL_PATHS.items(): - file_name = os.path.basename(v) - ldm.modules.midas.api.ISL_PATHS[k] = os.path.join(midas_path, file_name) - midas_urls = { - "dpt_large": "https://github.com/intel-isl/DPT/releases/download/1_0/dpt_large-midas-2f21e586.pt", - "dpt_hybrid": "https://github.com/intel-isl/DPT/releases/download/1_0/dpt_hybrid-midas-501f0c75.pt", - "midas_v21": "https://github.com/AlexeyAB/MiDaS/releases/download/midas_dpt/midas_v21-f6b98070.pt", - "midas_v21_small": "https://github.com/AlexeyAB/MiDaS/releases/download/midas_dpt/midas_v21_small-70d6b9c8.pt", - } - ldm.modules.midas.api.load_model_inner = ldm.modules.midas.api.load_model - - def load_model_wrapper(model_type): - path = ldm.modules.midas.api.ISL_PATHS[model_type] - if not os.path.exists(path): - if not os.path.exists(midas_path): - mkdir(midas_path) - shared.log.info(f"Downloading midas model weights for {model_type} to {path}") - request.urlretrieve(midas_urls[model_type], path) - shared.log.info(f"{model_type} downloaded") - return ldm.modules.midas.api.load_model_inner(model_type) - - ldm.modules.midas.api.load_model = load_model_wrapper - - def repair_config(sd_config): if "use_ema" not in sd_config.model.params: sd_config.model.params.use_ema = False @@ -580,7 +200,6 @@ def change_backend(): unload_model_weights() shared.backend = shared.Backend.ORIGINAL if shared.opts.sd_backend == 'original' else shared.Backend.DIFFUSERS shared.native = shared.backend == shared.Backend.DIFFUSERS - checkpoints_loaded.clear() from modules.sd_samplers import list_samplers list_samplers() list_models() @@ -588,118 +207,6 @@ def change_backend(): refresh_vae_list() -def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False): - guess = shared.opts.diffusers_pipeline - warn = shared.log.warning if warning else lambda *args, **kwargs: None - size = 0 - pipeline = None - if guess == 'Autodetect': - try: - guess = 'Stable Diffusion XL' if 'XL' in f.upper() else 'Stable Diffusion' - # guess by size - if os.path.isfile(f) and f.endswith('.safetensors'): - size = round(os.path.getsize(f) / 1024 / 1024) - if (size > 0 and size < 128): - warn(f'Model size smaller than expected: {f} size={size} MB') - elif (size >= 316 and size <= 324) or (size >= 156 and size <= 164): # 320 or 160 - warn(f'Model detected as VAE model, but attempting to load as model: {op}={f} size={size} MB') - guess = 'VAE' - elif (size >= 4970 and size <= 4976): # 4973 - guess = 'Stable Diffusion 2' # SD v2 but could be eps or v-prediction - # elif size < 0: # unknown - # guess = 'Stable Diffusion 2B' - elif (size >= 5791 and size <= 5799): # 5795 - if op == 'model': - warn(f'Model detected as SD-XL refiner model, but attempting to load a base model: {op}={f} size={size} MB') - guess = 'Stable Diffusion XL Refiner' - elif (size >= 6611 and size <= 7220): # 6617, HassakuXL is 6776, monkrenRealisticINT_v10 is 7217 - guess = 'Stable Diffusion XL' - elif (size >= 3361 and size <= 3369): # 3368 - guess = 'Stable Diffusion Upscale' - elif (size >= 4891 and size <= 4899): # 4897 - guess = 'Stable Diffusion XL Inpaint' - elif (size >= 9791 and size <= 9799): # 9794 - guess = 'Stable Diffusion XL Instruct' - elif (size > 3138 and size < 3142): #3140 - guess = 'Stable Diffusion XL' - elif (size > 5692 and size < 5698) or (size > 4134 and size < 4138) or (size > 10362 and size < 10366) or (size > 15028 and size < 15228): - guess = 'Stable Diffusion 3' - elif (size > 20000 and size < 40000): - guess = 'FLUX' - # guess by name - """ - if 'LCM_' in f.upper() or 'LCM-' in f.upper() or '_LCM' in f.upper() or '-LCM' in f.upper(): - if shared.backend == shared.Backend.ORIGINAL: - warn(f'Model detected as LCM model, but attempting to load using backend=original: {op}={f} size={size} MB') - guess = 'Latent Consistency Model' - """ - if 'instaflow' in f.lower(): - guess = 'InstaFlow' - if 'segmoe' in f.lower(): - guess = 'SegMoE' - if 'hunyuandit' in f.lower(): - guess = 'HunyuanDiT' - if 'pixart-xl' in f.lower(): - guess = 'PixArt-Alpha' - if 'stable-diffusion-3' in f.lower(): - guess = 'Stable Diffusion 3' - if 'stable-cascade' in f.lower() or 'stablecascade' in f.lower() or 'wuerstchen3' in f.lower() or ('sotediffusion' in f.lower() and "v2" in f.lower()): - if devices.dtype == torch.float16: - warn('Stable Cascade does not support Float16') - guess = 'Stable Cascade' - if 'pixart-sigma' in f.lower(): - guess = 'PixArt-Sigma' - if 'lumina-next' in f.lower(): - guess = 'Lumina-Next' - if 'kolors' in f.lower(): - guess = 'Kolors' - if 'auraflow' in f.lower(): - guess = 'AuraFlow' - if 'cogview' in f.lower(): - guess = 'CogView' - if 'meissonic' in f.lower(): - guess = 'Meissonic' - pipeline = 'custom' - if 'omnigen' in f.lower(): - guess = 'OmniGen' - pipeline = 'custom' - if 'flux' in f.lower(): - guess = 'FLUX' - if size > 11000 and size < 20000: - warn(f'Model detected as FLUX UNET model, but attempting to load a base model: {op}={f} size={size} MB') - # switch for specific variant - if guess == 'Stable Diffusion' and 'inpaint' in f.lower(): - guess = 'Stable Diffusion Inpaint' - elif guess == 'Stable Diffusion' and 'instruct' in f.lower(): - guess = 'Stable Diffusion Instruct' - if guess == 'Stable Diffusion XL' and 'inpaint' in f.lower(): - guess = 'Stable Diffusion XL Inpaint' - elif guess == 'Stable Diffusion XL' and 'instruct' in f.lower(): - guess = 'Stable Diffusion XL Instruct' - # get actual pipeline - pipeline = shared_items.get_pipelines().get(guess, None) if pipeline is None else pipeline - if not quiet: - shared.log.info(f'Autodetect {op}: detect="{guess}" class={getattr(pipeline, "__name__", None)} file="{f}" size={size}MB') - except Exception as e: - shared.log.error(f'Autodetect {op}: file="{f}" {e}') - if debug_load: - errors.display(e, f'Load {op}: {f}') - return None, None - else: - try: - size = round(os.path.getsize(f) / 1024 / 1024) - pipeline = shared_items.get_pipelines().get(guess, None) if pipeline is None else pipeline - if not quiet: - shared.log.info(f'Load {op}: detect="{guess}" class={getattr(pipeline, "__name__", None)} file="{f}" size={size}MB') - except Exception as e: - shared.log.error(f'Load {op}: detect="{guess}" file="{f}" {e}') - - if pipeline is None: - shared.log.warning(f'Load {op}: detect="{guess}" file="{f}" size={size} not recognized') - pipeline = diffusers.StableDiffusionPipeline - return pipeline, guess - - def copy_diffuser_options(new_pipe, orig_pipe): new_pipe.sd_checkpoint_info = getattr(orig_pipe, 'sd_checkpoint_info', None) new_pipe.sd_model_checkpoint = getattr(orig_pipe, 'sd_model_checkpoint', None) @@ -858,6 +365,9 @@ def set_diffuser_offload(sd_model, op: str = 'model'): def apply_balanced_offload(sd_model): from accelerate import infer_auto_device_map, dispatch_model from accelerate.hooks import add_hook_to_module, remove_hook_from_module, ModelHook + excluded = ['OmniGenPipeline'] + if sd_model.__class__.__name__ in excluded: + return sd_model class dispatch_from_cpu_hook(ModelHook): def init_hook(self, module): @@ -906,6 +416,8 @@ def apply_balanced_offload(sd_model): devices.torch_gc(fast=True) apply_balanced_offload_to_module(sd_model) + if hasattr(sd_model, "pipe"): + apply_balanced_offload_to_module(sd_model.pipe) if hasattr(sd_model, "prior_pipe"): apply_balanced_offload_to_module(sd_model.prior_pipe) if hasattr(sd_model, "decoder_pipe"): @@ -937,7 +449,6 @@ def move_model(model, device=None, force=False): if hasattr(model.vae, '_hf_hook'): debug_move(f'Model move: to={device} class={model.vae.__class__} fn={fn}') # pylint: disable=protected-access model.vae._hf_hook.execution_device = device # pylint: disable=protected-access - debug_move(f'Model move: device={device} class={model.__class__} accelerate={getattr(model, "has_accelerate", False)} fn={fn}') # pylint: disable=protected-access if hasattr(model, "components"): # accelerate patch for name, m in model.components.items(): if not hasattr(m, "_hf_hook"): # not accelerate hook @@ -956,6 +467,7 @@ def move_model(model, device=None, force=False): if hasattr(model, "device") and devices.normalize_device(model.device) == devices.normalize_device(device): return try: + t0 = time.time() try: model.to(device) if hasattr(model, "prior_pipe"): @@ -977,8 +489,12 @@ def move_model(model, device=None, force=False): pass # ignore model move if sequential offload is enabled else: raise e0 + t1 = time.time() except Exception as e1: + t1 = time.time() shared.log.error(f'Model move: device={device} {e1}') + if os.environ.get('SD_MOVE_DEBUG', None) or (t1-t0) > 0.1: + shared.log.debug(f'Model move: device={device} class={model.__class__.__name__} accelerate={getattr(model, "has_accelerate", False)} fn={fn} time={t1-t0:.2f}') # pylint: disable=protected-access devices.torch_gc() @@ -997,35 +513,6 @@ def move_base(model, device): return R -def get_load_config(model_file, model_type, config_type='yaml'): - if config_type == 'yaml': - yaml = os.path.splitext(model_file)[0] + '.yaml' - if os.path.exists(yaml): - return yaml - if model_type == 'Stable Diffusion': - return 'configs/v1-inference.yaml' - if model_type == 'Stable Diffusion XL': - return 'configs/sd_xl_base.yaml' - if model_type == 'Stable Diffusion XL Refiner': - return 'configs/sd_xl_refiner.yaml' - if model_type == 'Stable Diffusion 2': - return None # dont know if its eps or v so let diffusers sort it out - # return 'configs/v2-inference-512-base.yaml' - # return 'configs/v2-inference-768-v.yaml' - elif config_type == 'json': - if not shared.opts.diffuser_cache_config: - return None - if model_type == 'Stable Diffusion': - return 'configs/sd15' - if model_type == 'Stable Diffusion XL': - return 'configs/sdxl' - if model_type == 'Stable Diffusion 3': - return 'configs/sd3' - if model_type == 'FLUX': - return 'configs/flux' - return None - - def patch_diffuser_config(sd_model, model_file): def load_config(fn, k): model_file = os.path.splitext(fn)[0] @@ -1142,65 +629,65 @@ def load_diffuser_folder(model_type, pipeline, checkpoint_info, diffusers_load_c files = shared.walk_files(checkpoint_info.path, ['.safetensors', '.bin', '.ckpt']) if 'variant' not in diffusers_load_config and any('diffusion_pytorch_model.fp16' in f for f in files): # deal with diffusers lack of variant fallback when loading diffusers_load_config['variant'] = 'fp16' - if model_type is not None and pipeline is not None and 'ONNX' in model_type: # forced pipeline - try: - sd_model = pipeline.from_pretrained(checkpoint_info.path) - except Exception as e: - shared.log.error(f'Load {op}: type=ONNX path="{checkpoint_info.path}" {e}') - if debug_load: - errors.display(e, 'Load') - return None - else: - err1, err2, err3 = None, None, None - if os.path.exists(checkpoint_info.path) and os.path.isdir(checkpoint_info.path): - if os.path.exists(os.path.join(checkpoint_info.path, 'unet', 'diffusion_pytorch_model.bin')): - shared.log.debug(f'Load {op}: type=pickle') - diffusers_load_config['use_safetensors'] = False + if model_type is not None and pipeline is not None and 'ONNX' in model_type: # forced pipeline + try: + sd_model = pipeline.from_pretrained(checkpoint_info.path) + except Exception as e: + shared.log.error(f'Load {op}: type=ONNX path="{checkpoint_info.path}" {e}') if debug_load: - shared.log.debug(f'Load {op}: args={diffusers_load_config}') - try: # 1 - autopipeline, best choice but not all pipelines are available - try: + errors.display(e, 'Load') + return None + else: + err1, err2, err3 = None, None, None + if os.path.exists(checkpoint_info.path) and os.path.isdir(checkpoint_info.path): + if os.path.exists(os.path.join(checkpoint_info.path, 'unet', 'diffusion_pytorch_model.bin')): + shared.log.debug(f'Load {op}: type=pickle') + diffusers_load_config['use_safetensors'] = False + if debug_load: + shared.log.debug(f'Load {op}: args={diffusers_load_config}') + try: # 1 - autopipeline, best choice but not all pipelines are available + try: + sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + sd_model.model_type = sd_model.__class__.__name__ + except ValueError as e: + if 'no variant default' in str(e): + shared.log.warning(f'Load {op}: variant={diffusers_load_config["variant"]} model="{checkpoint_info.path}" using default variant') + diffusers_load_config.pop('variant', None) sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) sd_model.model_type = sd_model.__class__.__name__ - except ValueError as e: - if 'no variant default' in str(e): - shared.log.warning(f'Load {op}: variant={diffusers_load_config["variant"]} model="{checkpoint_info.path}" using default variant') - diffusers_load_config.pop('variant', None) - sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - sd_model.model_type = sd_model.__class__.__name__ - elif 'safetensors found in directory' in str(err1): - shared.log.warning(f'Load {op}: type=pickle') - diffusers_load_config['use_safetensors'] = False - sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - sd_model.model_type = sd_model.__class__.__name__ - else: - raise ValueError from e # reraise - except Exception as e: - err1 = e - if debug_load: - errors.display(e, 'Load AutoPipeline') - # shared.log.error(f'AutoPipeline: {e}') - try: # 2 - diffusion pipeline, works for most non-linked pipelines - if err1 is not None: - sd_model = diffusers.DiffusionPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + elif 'safetensors found in directory' in str(err1): + shared.log.warning(f'Load {op}: type=pickle') + diffusers_load_config['use_safetensors'] = False + sd_model = diffusers.AutoPipelineForText2Image.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) sd_model.model_type = sd_model.__class__.__name__ - except Exception as e: - err2 = e - if debug_load: - errors.display(e, "Load DiffusionPipeline") - # shared.log.error(f'DiffusionPipeline: {e}') - try: # 3 - try basic pipeline just in case - if err2 is not None: - sd_model = diffusers.StableDiffusionPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) - sd_model.model_type = sd_model.__class__.__name__ - except Exception as e: - err3 = e # ignore last error - shared.log.error(f"StableDiffusionPipeline: {e}") - if debug_load: - errors.display(e, "Load StableDiffusionPipeline") - if err3 is not None: - shared.log.error(f'Load {op}: {checkpoint_info.path} auto={err1} diffusion={err2}') - return None + else: + raise ValueError from e # reraise + except Exception as e: + err1 = e + if debug_load: + errors.display(e, 'Load AutoPipeline') + # shared.log.error(f'AutoPipeline: {e}') + try: # 2 - diffusion pipeline, works for most non-linked pipelines + if err1 is not None: + sd_model = diffusers.DiffusionPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + sd_model.model_type = sd_model.__class__.__name__ + except Exception as e: + err2 = e + if debug_load: + errors.display(e, "Load DiffusionPipeline") + # shared.log.error(f'DiffusionPipeline: {e}') + try: # 3 - try basic pipeline just in case + if err2 is not None: + sd_model = diffusers.StableDiffusionPipeline.from_pretrained(checkpoint_info.path, cache_dir=shared.opts.diffusers_dir, **diffusers_load_config) + sd_model.model_type = sd_model.__class__.__name__ + except Exception as e: + err3 = e # ignore last error + shared.log.error(f"StableDiffusionPipeline: {e}") + if debug_load: + errors.display(e, "Load StableDiffusionPipeline") + if err3 is not None: + shared.log.error(f'Load {op}: {checkpoint_info.path} auto={err1} diffusion={err2}') + return None return sd_model @@ -1216,7 +703,7 @@ def load_diffuser_file(model_type, pipeline, checkpoint_info, diffusers_load_con if shared.opts.diffusers_force_zeros: diffusers_load_config['force_zeros_for_empty_prompt '] = shared.opts.diffusers_force_zeros else: - model_config = get_load_config(checkpoint_info.path, model_type, config_type='json') + model_config = sd_detect.get_load_config(checkpoint_info.path, model_type, config_type='json') if model_config is not None: if debug_load: shared.log.debug(f'Load {op}: config="{model_config}"') @@ -1307,11 +794,14 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No return # detect pipeline - pipeline, model_type = detect_pipeline(checkpoint_info.path, op) + pipeline, model_type = sd_detect.detect_pipeline(checkpoint_info.path, op) # preload vae so it can be used as param vae = None sd_vae.loaded_vae_file = None + if model_type is None: + shared.log.error(f'Load {op}: pipeline={shared.opts.diffusers_pipeline} not detected') + return if model_type.startswith('Stable Diffusion') and (op == 'model' or op == 'refiner'): # preload vae for sd models vae_file, vae_source = sd_vae.resolve_vae(checkpoint_info.filename) vae = sd_vae.load_vae_diffusers(checkpoint_info.path, vae_file, vae_source) @@ -1558,6 +1048,17 @@ def clean_diffuser_pipe(pipe): def set_diffuser_pipe(pipe, new_pipe_type): + exclude = [ + 'StableDiffusionReferencePipeline', + 'StableDiffusionAdapterPipeline', + 'AnimateDiffPipeline', + 'AnimateDiffSDXLPipeline', + 'OmniGenPipeline', + 'StableDiffusion3ControlNetPipeline', + 'StableDiffusionXLPuLIDPipeline', + 'InstantIRPipeline', + ] + n = getattr(pipe.__class__, '__name__', '') if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: clean_diffuser_pipe(pipe) @@ -1566,7 +1067,7 @@ def set_diffuser_pipe(pipe, new_pipe_type): return pipe # skip specific pipelines - if n in ['StableDiffusionReferencePipeline', 'StableDiffusionAdapterPipeline', 'AnimateDiffPipeline', 'AnimateDiffSDXLPipeline', 'OmniGenPipeline']: + if n in exclude: return pipe if 'Onnx' in pipe.__class__.__name__: return pipe @@ -1783,7 +1284,7 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None, shared.log.info(f"Model loaded in {timer.summary()}") current_checkpoint_info = None devices.torch_gc(force=True) - shared.log.info(f'Model load finished: {memory_stats()} cached={len(checkpoints_loaded.keys())}') + shared.log.info(f'Model load finished: {memory_stats()}') def reload_text_encoder(initial=False): @@ -1843,7 +1344,7 @@ def reload_model_weights(sd_model=None, info=None, reuse_dict=False, op='model', state_dict = get_checkpoint_state_dict(checkpoint_info, timer) if not shared.native else None checkpoint_config = sd_models_config.find_checkpoint_config(state_dict, checkpoint_info) timer.record("config") - if sd_model is None or checkpoint_config != getattr(sd_model, 'used_config', None): + if sd_model is None or checkpoint_config != getattr(sd_model, 'used_config', None) or force: sd_model = None if not shared.native: load_model(checkpoint_info, already_loaded_state_dict=state_dict, timer=timer, op=op) @@ -1900,9 +1401,10 @@ def disable_offload(sd_model): from accelerate.hooks import remove_hook_from_module if not getattr(sd_model, 'has_accelerate', False): return - for _name, model in sd_model.components.items(): - if isinstance(model, torch.nn.Module): - remove_hook_from_module(model, recurse=True) + if hasattr(sd_model, 'components'): + for _name, model in sd_model.components.items(): + if isinstance(model, torch.nn.Module): + remove_hook_from_module(model, recurse=True) sd_model.has_accelerate = False @@ -1937,82 +1439,6 @@ def unload_model_weights(op='model'): shared.log.debug(f'Unload weights {op}: {memory_stats()}') -def apply_token_merging(sd_model): - current_tome = getattr(sd_model, 'applied_tome', 0) - current_todo = getattr(sd_model, 'applied_todo', 0) - - if shared.opts.token_merging_method == 'ToMe' and shared.opts.tome_ratio > 0: - if current_tome == shared.opts.tome_ratio: - return - if shared.opts.hypertile_unet_enabled and not shared.cmd_opts.experimental: - shared.log.warning('Token merging not supported with HyperTile for UNet') - return - try: - import installer - installer.install('tomesd', 'tomesd', ignore=False) - import tomesd - tomesd.apply_patch( - sd_model, - ratio=shared.opts.tome_ratio, - use_rand=False, # can cause issues with some samplers - merge_attn=True, - merge_crossattn=False, - merge_mlp=False - ) - shared.log.info(f'Applying ToMe: ratio={shared.opts.tome_ratio}') - sd_model.applied_tome = shared.opts.tome_ratio - except Exception: - shared.log.warning(f'Token merging not supported: pipeline={sd_model.__class__.__name__}') - else: - sd_model.applied_tome = 0 - - if shared.opts.token_merging_method == 'ToDo' and shared.opts.todo_ratio > 0: - if current_todo == shared.opts.todo_ratio: - return - if shared.opts.hypertile_unet_enabled and not shared.cmd_opts.experimental: - shared.log.warning('Token merging not supported with HyperTile for UNet') - return - try: - from modules.todo.todo_utils import patch_attention_proc - token_merge_args = { - "ratio": shared.opts.todo_ratio, - "merge_tokens": "keys/values", - "merge_method": "downsample", - "downsample_method": "nearest", - "downsample_factor": 2, - "timestep_threshold_switch": 0.0, - "timestep_threshold_stop": 0.0, - "downsample_factor_level_2": 1, - "ratio_level_2": 0.0, - } - patch_attention_proc(sd_model.unet, token_merge_args=token_merge_args) - shared.log.info(f'Applying ToDo: ratio={shared.opts.todo_ratio}') - sd_model.applied_todo = shared.opts.todo_ratio - except Exception: - shared.log.warning(f'Token merging not supported: pipeline={sd_model.__class__.__name__}') - else: - sd_model.applied_todo = 0 - - -def remove_token_merging(sd_model): - current_tome = getattr(sd_model, 'applied_tome', 0) - current_todo = getattr(sd_model, 'applied_todo', 0) - try: - if current_tome > 0: - import tomesd - tomesd.remove_patch(sd_model) - sd_model.applied_tome = 0 - except Exception: - pass - try: - if current_todo > 0: - from modules.todo.todo_utils import remove_patch - remove_patch(sd_model) - sd_model.applied_todo = 0 - except Exception: - pass - - def path_to_repo(fn: str = ''): if isinstance(fn, CheckpointInfo): fn = fn.name diff --git a/modules/sd_samplers.py b/modules/sd_samplers.py index 82171e0b7..bbc2f360b 100644 --- a/modules/sd_samplers.py +++ b/modules/sd_samplers.py @@ -84,6 +84,8 @@ def create_sampler(name, model): if 'AuraFlow' in model.__class__.__name__: shared.log.warning(f'AuraFlow: sampler="{name}" unsupported') return None + if 'KDiffusion' in model.__class__.__name__: + return None if not hasattr(model, 'scheduler_config'): model.scheduler_config = sampler.sampler.config.copy() if hasattr(sampler.sampler, 'config') else {} model.scheduler = sampler.sampler diff --git a/modules/sd_samplers_common.py b/modules/sd_samplers_common.py index 1b1cd189a..a487fe9b7 100644 --- a/modules/sd_samplers_common.py +++ b/modules/sd_samplers_common.py @@ -44,7 +44,7 @@ def single_sample_to_image(sample, approximation=None): if sample.dtype == torch.bfloat16 and (approximation == 0 or approximation == 1): sample = sample.to(torch.float16) except Exception as e: - warn_once(f'live preview: {e}') + warn_once(f'Preview: {e}') if len(sample.shape) > 4: # likely unknown video latent (e.g. svd) return Image.new(mode="RGB", size=(512, 512)) @@ -82,7 +82,7 @@ def single_sample_to_image(sample, approximation=None): transform = T.ToPILImage() image = transform(x_sample) except Exception as e: - warn_once(f'live preview: {e}') + warn_once(f'Preview: {e}') image = Image.new(mode="RGB", size=(512, 512)) return image diff --git a/modules/sd_vae.py b/modules/sd_vae.py index 52ba77bba..f266f8c38 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -2,7 +2,7 @@ import os import glob from copy import deepcopy import torch -from modules import shared, errors, paths, devices, script_callbacks, sd_models +from modules import shared, errors, paths, devices, script_callbacks, sd_models, sd_detect vae_ignore_keys = {"model_ema.decay", "model_ema.num_updates"} @@ -206,8 +206,8 @@ def load_vae_diffusers(model_file, vae_file=None, vae_source="unknown-source"): diffusers_load_config['variant'] = shared.opts.diffusers_vae_load_variant if shared.opts.diffusers_vae_upcast != 'default': diffusers_load_config['force_upcast'] = True if shared.opts.diffusers_vae_upcast == 'true' else False - _pipeline, model_type = sd_models.detect_pipeline(model_file, 'vae') - vae_config = sd_models.get_load_config(model_file, model_type, config_type='json') + _pipeline, model_type = sd_detect.detect_pipeline(model_file, 'vae') + vae_config = sd_detect.get_load_config(model_file, model_type, config_type='json') if vae_config is not None: diffusers_load_config['config'] = os.path.join(vae_config, 'vae') shared.log.info(f'Load module: type=VAE model="{vae_file}" source={vae_source} config={diffusers_load_config}') diff --git a/modules/sd_vae_approx.py b/modules/sd_vae_approx.py index 2b4399edb..78fe8f08b 100644 --- a/modules/sd_vae_approx.py +++ b/modules/sd_vae_approx.py @@ -46,7 +46,7 @@ def nn_approximation(sample): # Approximate NN sd_vae_approx_model.load_state_dict(approx_weights) sd_vae_approx_model.eval() sd_vae_approx_model.to(device, dtype) - shared.log.debug(f'VAE load: type=approximate model={model_path}') + shared.log.debug(f'VAE load: type=approximate model="{model_path}"') try: in_sample = sample.to(device, dtype).unsqueeze(0) sd_vae_approx_model.to(device, dtype) diff --git a/modules/sd_vae_ostris.py b/modules/sd_vae_ostris.py new file mode 100644 index 000000000..70542e9f5 --- /dev/null +++ b/modules/sd_vae_ostris.py @@ -0,0 +1,41 @@ +import time +import torch +import diffusers +from huggingface_hub import hf_hub_download +from safetensors.torch import load_file +from modules import shared, devices + + +decoder_id = "ostris/vae-kl-f8-d16" +adapter_id = "ostris/16ch-VAE-Adapters" + + +def load_vae(pipe): + if shared.sd_model_type == 'sd': + adapter_file = "16ch-VAE-Adapter-SD15-alpha.safetensors" + elif shared.sd_model_type == 'sdxl': + adapter_file = "16ch-VAE-Adapter-SDXL-alpha_v02.safetensors" + else: + shared.log.error('VAE: type=osiris unsupported model type') + return + t0 = time.time() + ckpt_file = hf_hub_download(adapter_id, adapter_file, cache_dir=shared.opts.hfcache_dir) + ckpt = load_file(ckpt_file) + lora_state_dict = {k: v for k, v in ckpt.items() if "lora" in k} + unet_state_dict = {k.replace("unet_", ""): v for k, v in ckpt.items() if "unet_" in k} + + pipe.unet.conv_in = torch.nn.Conv2d(16, 320, 3, 1, 1) + pipe.unet.conv_out = torch.nn.Conv2d(320, 16, 3, 1, 1) + pipe.unet.load_state_dict(unet_state_dict, strict=False) + pipe.unet.conv_in.to(devices.dtype) + pipe.unet.conv_out.to(devices.dtype) + pipe.unet.config.in_channels = 16 + pipe.unet.config.out_channels = 16 + + pipe.load_lora_weights(lora_state_dict, adapter_name=adapter_id) + # pipe.set_adapters(adapter_names=[adapter_id], adapter_weights=[0.8]) + pipe.fuse_lora(adapter_names=[adapter_id], lora_scale=0.8, fuse_unet=True) + + pipe.vae = diffusers.AutoencoderKL.from_pretrained(decoder_id, torch_dtype=devices.dtype, cache_dir=shared.opts.hfcache_dir) + t1 = time.time() + shared.log.info(f'VAE load: type=osiris decoder="{decoder_id}" adapter="{adapter_id}" time={t1-t0:.2f}s') diff --git a/modules/sd_vae_taesd.py b/modules/sd_vae_taesd.py index 5cd7fab7c..4d213ad48 100644 --- a/modules/sd_vae_taesd.py +++ b/modules/sd_vae_taesd.py @@ -160,11 +160,11 @@ def decode(latents): download_model(model_path) if os.path.exists(model_path): taesd_models[f'{model_class}-decoder'] = TAESD(decoder_path=model_path, encoder_path=None) - shared.log.debug(f'VAE load: type=taesd model={model_path}') + shared.log.debug(f'VAE load: type=taesd model="{model_path}"') vae = taesd_models[f'{model_class}-decoder'] vae.decoder.to(devices.device, dtype) else: - shared.log.error(f'VAE load: type=taesd model={model_path} not found') + shared.log.error(f'VAE load: type=taesd model="{model_path}" not found') return latents if vae is None: return latents @@ -181,10 +181,14 @@ def decode(latents): image = 2.0 * image - 1.0 # typical normalized range except for preview which runs denormalization return image else: - shared.log.error(f'TAESD decode unsupported latent type: {latents.shape}') + if not previous_warnings: + shared.log.error(f'TAESD decode unsupported latent type: {latents.shape}') + previous_warnings = True return latents except Exception as e: - shared.log.error(f'VAE decode taesd: {e}') + if not previous_warnings: + shared.log.error(f'VAE decode taesd: {e}') + previous_warnings = True return latents @@ -204,7 +208,7 @@ def encode(image): model_path = os.path.join(paths.models_path, "TAESD", f"tae{model_class}_encoder.pth") download_model(model_path) if os.path.exists(model_path): - shared.log.debug(f'VAE load: type=taesd model={model_path}') + shared.log.debug(f'VAE load: type=taesd model="{model_path}"') taesd_models[f'{model_class}-encoder'] = TAESD(encoder_path=model_path, decoder_path=None) vae = taesd_models[f'{model_class}-encoder'] vae.encoder.to(devices.device, devices.dtype_vae) diff --git a/modules/shared.py b/modules/shared.py index 4622f1d2c..867651e7a 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -280,12 +280,13 @@ def options_section(section_identifier, options_dict): return options_dict -def list_checkpoint_tiles(): +def list_checkpoint_titles(): import modules.sd_models # pylint: disable=W0621 - return modules.sd_models.checkpoint_tiles() + return modules.sd_models.checkpoint_titles() -default_checkpoint = list_checkpoint_tiles()[0] if len(list_checkpoint_tiles()) > 0 else "model.safetensors" +list_checkpoint_tiles = list_checkpoint_titles # alias for legacy typo +default_checkpoint = list_checkpoint_titles()[0] if len(list_checkpoint_titles()) > 0 else "model.safetensors" def is_url(string): @@ -393,7 +394,7 @@ def get_default_modes(): elif gpu_memory <= 8: cmd_opts.medvram = True default_offload_mode = "model" - log.info(f"Device detect: memory={gpu_memory:.1f} ptimization=medvram") + log.info(f"Device detect: memory={gpu_memory:.1f} optimization=medvram") else: default_offload_mode = "none" log.info(f"Device detect: memory={gpu_memory:.1f} optimization=none") @@ -427,12 +428,12 @@ startup_offload_mode, startup_cross_attention, startup_sdp_options = get_default options_templates.update(options_section(('sd', "Execution & Models"), { "sd_backend": OptionInfo(default_backend, "Execution backend", gr.Radio, {"choices": ["diffusers", "original"] }), - "sd_model_checkpoint": OptionInfo(default_checkpoint, "Base model", DropdownEditable, lambda: {"choices": list_checkpoint_tiles()}, refresh=refresh_checkpoints), - "sd_model_refiner": OptionInfo('None', "Refiner model", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_tiles()}, refresh=refresh_checkpoints), + "sd_model_checkpoint": OptionInfo(default_checkpoint, "Base model", DropdownEditable, lambda: {"choices": list_checkpoint_titles()}, refresh=refresh_checkpoints), + "sd_model_refiner": OptionInfo('None', "Refiner model", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_titles()}, refresh=refresh_checkpoints), "sd_vae": OptionInfo("Automatic", "VAE model", gr.Dropdown, lambda: {"choices": shared_items.sd_vae_items()}, refresh=shared_items.refresh_vae_list), "sd_unet": OptionInfo("None", "UNET model", gr.Dropdown, lambda: {"choices": shared_items.sd_unet_items()}, refresh=shared_items.refresh_unet_list), "sd_text_encoder": OptionInfo('None', "Text encoder model", gr.Dropdown, lambda: {"choices": shared_items.sd_te_items()}, refresh=shared_items.refresh_te_list), - "sd_model_dict": OptionInfo('None', "Use separate base dict", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_tiles()}, refresh=refresh_checkpoints), + "sd_model_dict": OptionInfo('None', "Use separate base dict", gr.Dropdown, lambda: {"choices": ['None'] + list_checkpoint_titles()}, refresh=refresh_checkpoints), "sd_checkpoint_autoload": OptionInfo(True, "Model autoload on start"), "sd_checkpoint_autodownload": OptionInfo(True, "Model auto-download on demand"), "sd_textencoder_cache": OptionInfo(True, "Cache text encoder results", gr.Checkbox, {"visible": False}), @@ -480,6 +481,7 @@ options_templates.update(options_section(('cuda', "Compute Settings"), { "cudnn_benchmark": OptionInfo(False, "Full-depth cuDNN benchmark feature"), "diffusers_fuse_projections": OptionInfo(False, "Fused projections"), "torch_expandable_segments": OptionInfo(False, "Torch expandable segments"), + "cuda_mem_fraction": OptionInfo(0.0, "Torch memory limit", gr.Slider, {"minimum": 0, "maximum": 2.0, "step": 0.05}), "torch_gc_threshold": OptionInfo(80, "Torch memory threshold for GC", gr.Slider, {"minimum": 0, "maximum": 100, "step": 1}), "torch_malloc": OptionInfo("native", "Torch memory allocator", gr.Radio, {"choices": ['native', 'cudaMallocAsync'] }), @@ -820,8 +822,8 @@ options_templates.update(options_section(('postprocessing', "Postprocessing"), { "detailer_conf": OptionInfo(0.6, "Min confidence", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.05, "visible": False}), "detailer_max": OptionInfo(2, "Max detected", gr.Slider, {"minimum": 1, "maximum": 10, "step": 1, "visible": False}), "detailer_iou": OptionInfo(0.5, "Max overlap", gr.Slider, {"minimum": 0, "maximum": 1.0, "step": 0.05, "visible": False}), - "detailer_min_size": OptionInfo(0, "Min object size", gr.Slider, {"minimum": 0, "maximum": 1024, "step": 1, "visible": False}), - "detailer_max_size": OptionInfo(0, "Max object size", gr.Slider, {"minimum": 0, "maximum": 1024, "step": 1, "visible": False}), + "detailer_min_size": OptionInfo(0.0, "Min object size", gr.Slider, {"minimum": 0.1, "maximum": 1, "step": 0.05, "visible": False}), + "detailer_max_size": OptionInfo(1.0, "Max object size", gr.Slider, {"minimum": 0.1, "maximum": 1, "step": 0.05, "visible": False}), "detailer_padding": OptionInfo(20, "Item padding", gr.Slider, {"minimum": 0, "maximum": 100, "step": 1, "visible": False}), "detailer_blur": OptionInfo(10, "Item edge blur", gr.Slider, {"minimum": 0, "maximum": 100, "step": 1, "visible": False}), "detailer_strength": OptionInfo(0.5, "Detailer strength", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01, "visible": False}), @@ -876,8 +878,9 @@ options_templates.update(options_section(('interrogate', "Interrogate"), { })) options_templates.update(options_section(('extra_networks', "Networks"), { - "extra_networks_sep1": OptionInfo("

Extra networks UI

", "", gr.HTML), - "extra_networks": OptionInfo(["All"], "Networks", gr.Dropdown, lambda: {"multiselect":True, "choices": ['All'] + [en.title for en in extra_networks]}), + "extra_networks_sep1": OptionInfo("

Networks UI

", "", gr.HTML), + "extra_networks_show": OptionInfo(True, "UI show on startup"), + "extra_networks": OptionInfo(["All"], "Available networks", gr.Dropdown, lambda: {"multiselect":True, "choices": ['All'] + [en.title for en in extra_networks]}), "extra_networks_sort": OptionInfo("Default", "Sort order", gr.Dropdown, {"choices": ['Default', 'Name [A-Z]', 'Name [Z-A]', 'Date [Newest]', 'Date [Oldest]', 'Size [Largest]', 'Size [Smallest]']}), "extra_networks_view": OptionInfo("gallery", "UI view", gr.Radio, {"choices": ["gallery", "list"]}), "extra_networks_card_cover": OptionInfo("sidebar", "UI position", gr.Radio, {"choices": ["cover", "inline", "sidebar"]}), @@ -887,13 +890,18 @@ options_templates.update(options_section(('extra_networks', "Networks"), { "extra_networks_card_square": OptionInfo(True, "UI disable variable aspect ratio"), "extra_networks_fetch": OptionInfo(True, "UI fetch network info on mouse-over"), "extra_networks_card_fit": OptionInfo("cover", "UI image contain method", gr.Radio, {"choices": ["contain", "cover", "fill"], "visible": False}), - "extra_networks_sep2": OptionInfo("

Extra networks general

", "", gr.HTML), - "extra_network_reference": OptionInfo(False, "Use reference values when available", gr.Checkbox), "extra_network_skip_indexing": OptionInfo(False, "Build info on first access", gr.Checkbox), - "extra_networks_default_multiplier": OptionInfo(1.0, "Default strength", gr.Slider, {"minimum": 0.0, "maximum": 2.0, "step": 0.01}), + + "extra_networks_model_sep": OptionInfo("

Models

", "", gr.HTML), + "extra_network_reference": OptionInfo(False, "Use reference values when available", gr.Checkbox), + "extra_networks_embed_sep": OptionInfo("

Embeddings

", "", gr.HTML), "diffusers_convert_embed": OptionInfo(False, "Auto-convert SD 1.5 embeddings to SDXL ", gr.Checkbox, {"visible": native}), - "extra_networks_sep3": OptionInfo("

Extra networks settings

", "", gr.HTML), + "extra_networks_styles_sep": OptionInfo("

Styles

", "", gr.HTML), "extra_networks_styles": OptionInfo(True, "Show built-in styles"), + "extra_networks_wildcard_sep": OptionInfo("

Wildcards

", "", gr.HTML), + "wildcards_enabled": OptionInfo(True, "Enable file wildcards support"), + "extra_networks_lora_sep": OptionInfo("

LoRA

", "", gr.HTML), + "extra_networks_default_multiplier": OptionInfo(1.0, "Default strength", gr.Slider, {"minimum": 0.0, "maximum": 2.0, "step": 0.01}), "lora_preferred_name": OptionInfo("filename", "LoRA preferred name", gr.Radio, {"choices": ["filename", "alias"]}), "lora_add_hashes_to_infotext": OptionInfo(False, "LoRA add hash info"), "lora_force_diffusers": OptionInfo(False if not cmd_opts.use_openvino else True, "LoRA force loading of all models using Diffusers"), @@ -904,9 +912,9 @@ options_templates.update(options_section(('extra_networks', "Networks"), { "lora_quant": OptionInfo("NF4","LoRA precision in quantized models", gr.Radio, {"choices": ["NF4", "FP4"]}), "lora_functional": OptionInfo(False, "Use Kohya method for handling multiple LoRA", gr.Checkbox, { "visible": False }), "lora_load_gpu": OptionInfo(True if not cmd_opts.lowvram else False, "Load LoRA directly to GPU"), - "hypernetwork_enabled": OptionInfo(False, "Enable Hypernetwork support"), + + "hypernetwork_enabled": OptionInfo(False, "Enable Hypernetwork support", gr.Checkbox, {"visible": False}), "sd_hypernetwork": OptionInfo("None", "Add hypernetwork to prompt", gr.Dropdown, { "choices": ["None"], "visible": False }), - "wildcards_enabled": OptionInfo(True, "Enable file wildcards support"), })) options_templates.update(options_section((None, "Hidden options"), { @@ -1100,7 +1108,7 @@ profiler = None opts = Options() config_filename = cmd_opts.config opts.load(config_filename) -cmd_opts = cmd_args.compatibility_args(opts, cmd_opts) +cmd_opts = cmd_args.settings_args(opts, cmd_opts) if cmd_opts.use_xformers: opts.data['cross_attention_optimization'] = 'xFormers' opts.data['uni_pc_lower_order_final'] = opts.schedulers_use_loworder # compatibility diff --git a/modules/token_merge.py b/modules/token_merge.py new file mode 100644 index 000000000..f97c1fc8e --- /dev/null +++ b/modules/token_merge.py @@ -0,0 +1,77 @@ +from modules import shared + + +def apply_token_merging(sd_model): + current_tome = getattr(sd_model, 'applied_tome', 0) + current_todo = getattr(sd_model, 'applied_todo', 0) + + if shared.opts.token_merging_method == 'ToMe' and shared.opts.tome_ratio > 0: + if current_tome == shared.opts.tome_ratio: + return + if shared.opts.hypertile_unet_enabled and not shared.cmd_opts.experimental: + shared.log.warning('Token merging not supported with HyperTile for UNet') + return + try: + import installer + installer.install('tomesd', 'tomesd', ignore=False) + import tomesd + tomesd.apply_patch( + sd_model, + ratio=shared.opts.tome_ratio, + use_rand=False, # can cause issues with some samplers + merge_attn=True, + merge_crossattn=False, + merge_mlp=False + ) + shared.log.info(f'Applying ToMe: ratio={shared.opts.tome_ratio}') + sd_model.applied_tome = shared.opts.tome_ratio + except Exception: + shared.log.warning(f'Token merging not supported: pipeline={sd_model.__class__.__name__}') + else: + sd_model.applied_tome = 0 + + if shared.opts.token_merging_method == 'ToDo' and shared.opts.todo_ratio > 0: + if current_todo == shared.opts.todo_ratio: + return + if shared.opts.hypertile_unet_enabled and not shared.cmd_opts.experimental: + shared.log.warning('Token merging not supported with HyperTile for UNet') + return + try: + from modules.todo.todo_utils import patch_attention_proc + token_merge_args = { + "ratio": shared.opts.todo_ratio, + "merge_tokens": "keys/values", + "merge_method": "downsample", + "downsample_method": "nearest", + "downsample_factor": 2, + "timestep_threshold_switch": 0.0, + "timestep_threshold_stop": 0.0, + "downsample_factor_level_2": 1, + "ratio_level_2": 0.0, + } + patch_attention_proc(sd_model.unet, token_merge_args=token_merge_args) + shared.log.info(f'Applying ToDo: ratio={shared.opts.todo_ratio}') + sd_model.applied_todo = shared.opts.todo_ratio + except Exception: + shared.log.warning(f'Token merging not supported: pipeline={sd_model.__class__.__name__}') + else: + sd_model.applied_todo = 0 + + +def remove_token_merging(sd_model): + current_tome = getattr(sd_model, 'applied_tome', 0) + current_todo = getattr(sd_model, 'applied_todo', 0) + try: + if current_tome > 0: + import tomesd + tomesd.remove_patch(sd_model) + sd_model.applied_tome = 0 + except Exception: + pass + try: + if current_todo > 0: + from modules.todo.todo_utils import remove_patch + remove_patch(sd_model) + sd_model.applied_todo = 0 + except Exception: + pass diff --git a/modules/ui.py b/modules/ui.py index 8ea565cfc..039ef6487 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -355,13 +355,20 @@ def create_ui(startup_timer = None): ui_onnx.create_ui() with gr.TabItem("Change log", id="change_log", elem_id="system_tab_changelog"): - with open('CHANGELOG.md', 'r', encoding='utf-8') as f: - md = f.read() - gr.Markdown(md) + def get_changelog(): + with open('CHANGELOG.md', 'r', encoding='utf-8') as f: + content = f.read() + content = content.replace('# Change Log for SD.Next', ' ') + return content - with gr.TabItem("Licenses", id="system_licenses", elem_id="system_tab_licenses"): - gr.HTML(shared.html("licenses.html"), elem_id="licenses", elem_classes="licenses") - create_dirty_indicator("tab_licenses", [], interactive=False) + with gr.Column(): + get_changelog_btn = gr.Button(value='Get changelog', elem_id="get_changelog") + with gr.Column(): + _changelog_search = gr.Textbox(label="Search", elem_id="changelog_search") + _changelog_result = gr.HTML(elem_id="changelog_result") + + changelog_markdown = gr.Markdown('', elem_id="changelog_markdown") + get_changelog_btn.click(fn=get_changelog, outputs=[changelog_markdown], show_progress=True) def unload_sd_weights(): modules.sd_models.unload_model_weights(op='model') diff --git a/modules/ui_common.py b/modules/ui_common.py index 5e873355b..9ad87c17f 100644 --- a/modules/ui_common.py +++ b/modules/ui_common.py @@ -319,13 +319,18 @@ def create_output_panel(tabname, preview=True, prompt=None, height=None): return result_gallery, generation_info, html_info, html_info_formatted, html_log -def create_refresh_button(refresh_component, refresh_method, refreshed_args, elem_id, visible: bool = True): +def create_refresh_button(refresh_component, refresh_method, refreshed_args = None, elem_id = None, visible: bool = True): def refresh(): refresh_method() - args = refreshed_args() if callable(refreshed_args) else refreshed_args + if refreshed_args is None: + args = {"choices": refresh_method()} # pylint: disable=unnecessary-lambda-assignment + elif callable(refreshed_args): + args = refreshed_args() + else: + args = refreshed_args for k, v in args.items(): setattr(refresh_component, k, v) - return gr.update(**(args or {})) + return gr.update(**args) refresh_button = ui_components.ToolButton(value=ui_symbols.refresh, elem_id=elem_id, visible=visible) refresh_button.click(fn=refresh, inputs=[], outputs=[refresh_component]) diff --git a/modules/ui_control.py b/modules/ui_control.py index 388e8ede9..fccda2167 100644 --- a/modules/ui_control.py +++ b/modules/ui_control.py @@ -152,7 +152,7 @@ def create_ui(_blocks: gr.Blocks=None): with gr.Row(): override_settings = ui_common.create_override_inputs('control') - with gr.Row(variant='compact', elem_id="control_extra_networks", visible=False) as extra_networks_ui: + with gr.Row(variant='compact', elem_id="control_extra_networks", elem_classes=["extra_networks_root"], visible=False) as extra_networks_ui: from modules import timer, ui_extra_networks extra_networks_ui = ui_extra_networks.create_ui(extra_networks_ui, btn_extra, 'control', skip_indexing=shared.opts.extra_network_skip_indexing) timer.startup.record('ui-networks') @@ -612,6 +612,7 @@ def create_ui(_blocks: gr.Blocks=None): (mask_controls[6], "Mask auto"), # advanced (cfg_scale, "CFG scale"), + (cfg_end, "CFG end"), (clip_skip, "Clip skip"), (image_cfg_scale, "Image CFG scale"), (diffusers_guidance_rescale, "CFG rescale"), diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index 58fb463a9..323f4830f 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -50,6 +50,7 @@ card_list = '''
''' +preview_map = None def init_api(app): @@ -350,6 +351,9 @@ class ExtraNetworksPage: return self.link_preview(preview_file) def update_all_previews(self, items): + global preview_map # pylint: disable=global-statement + if preview_map is None: + preview_map = shared.readfile('html/previews.json', silent=True) t0 = time.time() reference_path = os.path.abspath(os.path.join('models', 'Reference')) possible_paths = list(set([os.path.dirname(item['filename']) for item in items] + [reference_path])) @@ -378,8 +382,14 @@ class ExtraNetworksPage: self.missing_thumbs.append(all_previews[file_idx]) item['preview'] = self.link_preview(all_previews[file_idx]) break + if item.get('preview', None) is None: + found = preview_map.get(base, None) + if found is not None: + item['preview'] = self.link_preview(found) + debug(f'EN mapped-preview: {item["name"]}={found}') if item.get('preview', None) is None: item['preview'] = self.link_preview('html/card-no-preview.png') + debug(f'EN missing-preview: {item["name"]}') self.preview_time += time.time() - t0 diff --git a/modules/ui_img2img.py b/modules/ui_img2img.py index d46ea4dd3..4cb8e4c18 100644 --- a/modules/ui_img2img.py +++ b/modules/ui_img2img.py @@ -41,7 +41,7 @@ def create_ui(): img2img_prompt, img2img_prompt_styles, img2img_negative_prompt, img2img_submit, img2img_reprocess, img2img_paste, img2img_extra_networks_button, img2img_token_counter, img2img_token_button, img2img_negative_token_counter, img2img_negative_token_button = ui_sections.create_toprow(is_img2img=True, id_part="img2img") img2img_prompt_img = gr.File(label="", elem_id="img2img_prompt_image", file_count="single", type="binary", visible=False) - with gr.Row(variant='compact', elem_id="img2img_extra_networks", visible=False) as extra_networks_ui: + with gr.Row(variant='compact', elem_id="img2img_extra_networks", elem_classes=["extra_networks_root"], visible=False) as extra_networks_ui: from modules import ui_extra_networks extra_networks_ui_img2img = ui_extra_networks.create_ui(extra_networks_ui, img2img_extra_networks_button, 'img2img', skip_indexing=shared.opts.extra_network_skip_indexing) timer.startup.record('ui-networks') @@ -263,6 +263,7 @@ def create_ui(): (refiner_start, "Refiner start"), # advanced (cfg_scale, "CFG scale"), + (cfg_end, "CFG end"), (image_cfg_scale, "Image CFG scale"), (clip_skip, "Clip skip"), (diffusers_guidance_rescale, "CFG rescale"), diff --git a/modules/ui_models.py b/modules/ui_models.py index 051ca39a7..e9be428b4 100644 --- a/modules/ui_models.py +++ b/modules/ui_models.py @@ -59,8 +59,8 @@ def create_ui(): with gr.Tab(label="Convert"): with gr.Row(): - model_name = gr.Dropdown(sd_models.checkpoint_tiles(), label="Original model") - create_refresh_button(model_name, sd_models.list_models, lambda: {"choices": sd_models.checkpoint_tiles()}, "refresh_checkpoint_Z") + model_name = gr.Dropdown(sd_models.checkpoint_titles(), label="Original model") + create_refresh_button(model_name, sd_models.list_models, lambda: {"choices": sd_models.checkpoint_titles()}, "refresh_checkpoint_Z") with gr.Row(): custom_name = gr.Textbox(label="Output model name") with gr.Row(): @@ -98,7 +98,7 @@ def create_ui(): with gr.Tab(label="Merge"): def sd_model_choices(): - return ['None'] + sd_models.checkpoint_tiles() + return ['None'] + sd_models.checkpoint_titles() with gr.Row(equal_height=False): with gr.Column(variant='compact'): @@ -213,10 +213,10 @@ def create_ui(): del kwargs['dummy_component'] if kwargs.get("custom_name", None) is None: log.error('Merge: no output model specified') - return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], "No output model specified"] + return [*[gr.Dropdown.update(choices=sd_models.checkpoint_titles()) for _ in range(4)], "No output model specified"] elif kwargs.get("primary_model_name", None) is None or kwargs.get("secondary_model_name", None) is None: log.error('Merge: no models selected') - return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], "No models selected"] + return [*[gr.Dropdown.update(choices=sd_models.checkpoint_titles()) for _ in range(4)], "No models selected"] else: log.debug(f'Merge start: {kwargs}') try: @@ -224,7 +224,7 @@ def create_ui(): except Exception as e: modules.errors.display(e, 'Merge') sd_models.list_models() # to remove the potentially missing models from the list - return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], f"Error merging checkpoints: {e}"] + return [*[gr.Dropdown.update(choices=sd_models.checkpoint_titles()) for _ in range(4)], f"Error merging checkpoints: {e}"] return results def tertiary(mode): diff --git a/modules/ui_txt2img.py b/modules/ui_txt2img.py index 1ae3a8dad..e2886e901 100644 --- a/modules/ui_txt2img.py +++ b/modules/ui_txt2img.py @@ -25,7 +25,7 @@ def create_ui(): txt_prompt_img = gr.File(label="", elem_id="txt2img_prompt_image", file_count="single", type="binary", visible=False) txt_prompt_img.change(fn=modules.images.image_data, inputs=[txt_prompt_img], outputs=[txt2img_prompt, txt_prompt_img]) - with gr.Row(variant='compact', elem_id="txt2img_extra_networks", visible=False) as extra_networks_ui: + with gr.Row(variant='compact', elem_id="txt2img_extra_networks", elem_classes=["extra_networks_root"], visible=False) as extra_networks_ui: from modules import ui_extra_networks extra_networks_ui = ui_extra_networks.create_ui(extra_networks_ui, txt2img_extra_networks_button, 'txt2img', skip_indexing=shared.opts.extra_network_skip_indexing) timer.startup.record('ui-networks') @@ -116,6 +116,7 @@ def create_ui(): (subseed_strength, "Variation strength"), # advanced (cfg_scale, "CFG scale"), + (cfg_end, "CFG end"), (clip_skip, "Clip skip"), (image_cfg_scale, "Image CFG scale"), (diffusers_guidance_rescale, "CFG rescale"), diff --git a/modules/vqa.py b/modules/vqa.py index 7f0f8e17f..a0a5d147c 100644 --- a/modules/vqa.py +++ b/modules/vqa.py @@ -14,6 +14,8 @@ MODELS = { "MS Florence 2 Large": "microsoft/Florence-2-large", # 1.5GB "MiaoshouAI PromptGen 1.5 Base": "MiaoshouAI/Florence-2-base-PromptGen-v1.5@c06a5f02cc6071a5d65ee5d294cf3732d3097540", # 1.1GB "MiaoshouAI PromptGen 1.5 Large": "MiaoshouAI/Florence-2-large-PromptGen-v1.5@28a42440e39c9c32b83f7ae74ec2b3d1540404f0", # 3.3GB + "MiaoshouAI PromptGen 2.0 Base": "MiaoshouAI/Florence-2-base-PromptGen-v2.0", # 1.1GB + "MiaoshouAI PromptGen 2.0 Large": "MiaoshouAI/Florence-2-large-PromptGen-v2.0", # 3.3GB "CogFlorence 2.0 Large": "thwri/CogFlorence-2-Large-Freeze", # 1.6GB "CogFlorence 2.2 Large": "thwri/CogFlorence-2.2-Large", # 1.6GB "Moondream 2": "vikhyatk/moondream2", # 3.7GB @@ -30,8 +32,8 @@ MODELS = { def git(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: - model = transformers.GitForCausalLM.from_pretrained(repo) - processor = transformers.GitProcessor.from_pretrained(repo) + model = transformers.GitForCausalLM.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + processor = transformers.GitProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo model.to(devices.device, devices.dtype) shared.log.debug(f'VQA: class={model.__class__.__name__} processor={processor.__class__} model={repo}') @@ -55,8 +57,8 @@ def git(question: str, image: Image.Image, repo: str = None): def blip(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: - model = transformers.BlipForQuestionAnswering.from_pretrained(repo) - processor = transformers.BlipProcessor.from_pretrained(repo) + model = transformers.BlipForQuestionAnswering.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + processor = transformers.BlipProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo model.to(devices.device, devices.dtype) inputs = processor(image, question, return_tensors="pt") @@ -73,8 +75,8 @@ def blip(question: str, image: Image.Image, repo: str = None): def vilt(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: - model = transformers.ViltForQuestionAnswering.from_pretrained(repo) - processor = transformers.ViltProcessor.from_pretrained(repo) + model = transformers.ViltForQuestionAnswering.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + processor = transformers.ViltProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo model.to(devices.device) shared.log.debug(f'VQA: class={model.__class__.__name__} processor={processor.__class__} model={repo}') @@ -94,8 +96,8 @@ def vilt(question: str, image: Image.Image, repo: str = None): def pix(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: - model = transformers.Pix2StructForConditionalGeneration.from_pretrained(repo) - processor = transformers.Pix2StructProcessor.from_pretrained(repo) + model = transformers.Pix2StructForConditionalGeneration.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) + processor = transformers.Pix2StructProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo model.to(devices.device) shared.log.debug(f'VQA: class={model.__class__.__name__} processor={processor.__class__} model={repo}') @@ -115,8 +117,8 @@ def pix(question: str, image: Image.Image, repo: str = None): def moondream(question: str, image: Image.Image, repo: str = None): global processor, model, loaded # pylint: disable=global-statement if model is None or loaded != repo: - model = transformers.AutoModelForCausalLM.from_pretrained(repo, trust_remote_code=True) # revision = "2024-03-05" - processor = transformers.AutoTokenizer.from_pretrained(repo) # revision = "2024-03-05" + model = transformers.AutoModelForCausalLM.from_pretrained(repo, trust_remote_code=True, cache_dir=shared.opts.hfcache_dir) # revision = "2024-03-05" + processor = transformers.AutoTokenizer.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir) loaded = repo model.eval() model.to(devices.device, devices.dtype) @@ -142,8 +144,8 @@ def florence(question: str, image: Image.Image, repo: str = None, revision: str return R if model is None or loaded != repo: transformers.dynamic_module_utils.get_imports = get_imports - model = transformers.AutoModelForCausalLM.from_pretrained(repo, trust_remote_code=True, revision=revision) - processor = transformers.AutoProcessor.from_pretrained(repo, trust_remote_code=True, revision=revision) + model = transformers.AutoModelForCausalLM.from_pretrained(repo, trust_remote_code=True, revision=revision, cache_dir=shared.opts.hfcache_dir) + processor = transformers.AutoProcessor.from_pretrained(repo, trust_remote_code=True, revision=revision, cache_dir=shared.opts.hfcache_dir) transformers.dynamic_module_utils.get_imports = _get_imports loaded = repo model.eval() diff --git a/modules/zluda_installer.py b/modules/zluda_installer.py index 506652edf..84c130e8d 100644 --- a/modules/zluda_installer.py +++ b/modules/zluda_installer.py @@ -4,7 +4,7 @@ import ctypes import shutil import zipfile import urllib.request -from typing import Optional +from typing import Optional, Union from modules import rocm @@ -15,12 +15,18 @@ DLL_MAPPING = { } HIPSDK_TARGETS = ['rocblas.dll', 'rocsolver.dll', f'hiprtc{"".join([v.zfill(2) for v in rocm.version.split(".")])}.dll'] ZLUDA_TARGETS = ('nvcuda.dll', 'nvml.dll',) +default_agent: Union[rocm.Agent, None] = None def get_path() -> str: return os.path.abspath(os.environ.get('ZLUDA', '.zluda')) +def set_default_agent(agent: rocm.Agent): + global default_agent # pylint: disable=global-statement + default_agent = agent + + def install(zluda_path: os.PathLike) -> None: if os.path.exists(zluda_path): return diff --git a/package.json b/package.json index 2c0fa6b50..bf5a366ea 100644 --- a/package.json +++ b/package.json @@ -7,7 +7,7 @@ "url": "https://github.com/vladmandic/automatic/issues" }, "homepage": "https://github.com/vladmandic/automatic", - "license": "AGPLv3", + "license": "Apache-2.0", "engines": { "node": ">=14.0.0" }, diff --git a/requirements.txt b/requirements.txt index fd8e3ab4a..de0223bd5 100644 --- a/requirements.txt +++ b/requirements.txt @@ -36,11 +36,11 @@ torchsde==0.2.6 antlr4-python3-runtime==4.9.3 requests==2.32.3 tqdm==4.66.5 -accelerate==1.0.0 +accelerate==1.0.1 opencv-contrib-python-headless==4.9.0.80 einops==0.4.1 gradio==3.43.2 -huggingface_hub==0.25.2 +huggingface_hub==0.26.2 numexpr==2.8.8 numpy==1.26.4 numba==0.59.1 @@ -49,8 +49,8 @@ scipy pandas protobuf==4.25.3 pytorch_lightning==1.9.4 -tokenizers==0.20.0 -transformers==4.46.0 +tokenizers==0.20.3 +transformers==4.46.2 urllib3==1.26.19 Pillow==10.4.0 timm==0.9.16 @@ -61,8 +61,8 @@ torchdiffeq dctorch scikit-image seam-carving -open-clip-torch +sentencepiece -# TODO temporary block for torch==2.5.0 -torchvision!=0.20.0 +# block torch!=2.5.0 +torchvision!=0.20.0 diff --git a/scripts/apg.py b/scripts/apg.py index 3325d9333..c7e60c982 100644 --- a/scripts/apg.py +++ b/scripts/apg.py @@ -62,10 +62,13 @@ class Script(scripts.Script): def after(self, p: processing.StableDiffusionProcessing, processed: processing.Processed, eta, momentum, threshold): # pylint: disable=arguments-differ, unused-argument from modules import apg + if self.orig_pipe is None: + return processed # restore pipeline - if shared.sd_model_type == "sdxl": + if shared.sd_model_type == "sdxl" or shared.sd_model_type == "sd": shared.sd_model = self.orig_pipe elif shared.sd_model_type == "sc": shared.sd_model.prior_pipe = self.orig_pipe apg.buffer = None + self.orig_pipe = None return processed diff --git a/scripts/consistory_ext.py b/scripts/consistory_ext.py new file mode 100644 index 000000000..b7454aa75 --- /dev/null +++ b/scripts/consistory_ext.py @@ -0,0 +1,215 @@ +""" +original code from +ported to modules/consistory +- make it non-cuda exclusive +- separate create anchors and create extra +- do not force-load pipeline and unet, use existing model +- uses diffusers==0.25 class definitions, needed quite an update +- forces uses of xformers, converted attention calls to sdp +- unsafe tensor to numpy breaks with bfloat16 +- removed debug print statements +""" +import time +import gradio as gr +import diffusers +from modules import scripts, devices, errors, processing, shared, sd_models, sd_samplers + + +class Script(scripts.Script): + def __init__(self): + super().__init__() + self.anchor_cache_first_stage = None + self.anchor_cache_second_stage = None + + def title(self): + return 'ConsiStory' + + def show(self, is_img2img): + return not is_img2img if shared.native and shared.cmd_opts.experimental else False + + def reset(self): + self.anchor_cache_first_stage = None + self.anchor_cache_second_stage = None + shared.log.debug('ConsiStory reset anchors') + + def ui(self, _is_img2img): # ui elements + with gr.Row(): + gr.HTML('  ConsiStory: Consistent Image Generation
') + with gr.Row(): + gr.HTML('
▪ Anchors are created on first run
▪ Subsequent generate will use anchors and apply to main prompt
▪ Main prompts are separated by newlines') + with gr.Row(): + subject = gr.Textbox(label="Subject", placeholder='short description of a subject', value='') + with gr.Row(): + concepts = gr.Textbox(label="Concept Tokens", placeholder='one or more concepts to extract from subject', value='') + with gr.Row(): + prompts = gr.Textbox(label="Anchor settings", lines=2, placeholder='two scene settings to place subject in', value='') + with gr.Row(): + reset = gr.Button(value="Reset anchors", variant='primary') + reset.click(fn=self.reset, inputs=[], outputs=[]) + with gr.Row(): + dropout = gr.Slider(label="Mask Dropout", minimum=0.0, maximum=1.0, step=0.1, value=0.5) + with gr.Row(): + sampler = gr.Checkbox(label="Override sampler", value=True) + steps = gr.Checkbox(label="Override steps", value=True) + with gr.Row(): + same = gr.Checkbox(label="Same latent", value=False) + queries = gr.Checkbox(label="Share queries", value=True) + with gr.Row(): + sdsa = gr.Checkbox(label="Perform SDSA", value=True) + with gr.Row(): + freeu = gr.Checkbox(label="Enable FreeU", value=False) + freeu_preset = gr.Textbox(label="FreeU preset", value='0.6, 0.4, 1.1, 1.2') + with gr.Row(): + injection = gr.Checkbox(label="Perform Injection", value=False) + alpha = gr.Textbox(label="Alpha preset", value='10, 20, 0.8') + return [subject, concepts, prompts, dropout, sampler, steps, same, queries, sdsa, freeu, freeu_preset, alpha, injection] + + def create_model(self): + diffusers.models.embeddings.PositionNet = diffusers.models.embeddings.GLIGENTextBoundingboxProjection # patch as renamed in https://github.com/huggingface/diffusers/pull/6244/files + import modules.consistory as cs + if shared.sd_model.__class__.__name__ != 'ConsistoryExtendAttnSDXLPipeline': + shared.log.debug('ConsiStory init') + t0 = time.time() + state_dict = shared.sd_model.unet.state_dict() # save existing unet + shared.sd_model = sd_models.switch_pipe(cs.ConsistoryExtendAttnSDXLPipeline, shared.sd_model) + shared.sd_model.unet = cs.ConsistorySDXLUNet2DConditionModel.from_config(shared.sd_model.unet.config) + shared.sd_model.unet.load_state_dict(state_dict) # now load it into new class + shared.sd_model.unet.to(dtype=devices.dtype) + state_dict = None + # sd_models.set_diffuser_options(shared.sd_model) + sd_models.move_model(shared.sd_model, devices.device) + sd_models.move_model(shared.sd_model.unet, devices.device) + t1 = time.time() + shared.log.debug(f'ConsiStory load: model={shared.sd_model.__class__.__name__} time={t1-t0:.2f}') + devices.torch_gc(force=True) + + def set_args(self, p: processing.StableDiffusionProcessing, *args): + subject, concepts, prompts, dropout, sampler, steps, same, queries, sdsa, freeu, freeu_preset, alpha, injection = args # pylint: disable=unused-variable + processing.fix_seed(p) + if sampler: + shared.sd_model.scheduler = diffusers.DDIMScheduler.from_config(shared.sd_model.scheduler.config) + else: + sd_samplers.create_sampler(p.sampler_name, shared.sd_model) + if freeu: + try: + freeu_preset = [float(f.strip()) for f in freeu_preset.split(',')] + except Exception: + freeu_preset = [] + shared.log.warning(f'ConsiStory: freeu="{freeu_preset}" invalid') + if len(freeu) == 4: + shared.sd_model.enable_freeu(s1=freeu[0], s2=freeu[0], b1=freeu[0], b2=freeu[0]) + steps = 50 if steps else p.steps + if injection: + try: + alpha = [a.strip() for a in alpha.split(',')] + if len(alpha) == 3: + alpha = (int(alpha[0]), int(alpha[1]), float(alpha[2])) + except Exception: + alpha=(10, 20, 0.8) + shared.log.warning(f'ConsiStory: alpha="{alpha}" invalid') + else: + alpha=(10, 20, 0.8) + seed = p.seed + concepts = [c.strip() for c in concepts.split(',') if c.strip() != ''] + for c in concepts: + if c not in subject: + shared.log.warning(f'ConsiStory: concept="{c}" not in subject') + subject = f'{subject} {c}' + settings = [p.strip() for p in prompts.split('\n') if p.strip() != ''] + anchors = [f'{subject} {p}' for p in settings] + prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) + prompts = [p.strip() for p in prompt.split('\n') if p.strip() != ''] + for i, prompt in enumerate(prompts): + if subject not in prompt: + prompts[i] = f'{subject} {prompt}' + shared.log.debug(f'ConsiStory args: sampler={shared.sd_model.scheduler.__class__.__name__} steps={steps} sdsa={sdsa} queries={queries} same={same} dropout={dropout} freeu={freeu_preset if freeu else None} alpha={alpha if injection else None}') + return concepts, anchors, prompts, alpha, steps, seed + + def create_anchors(self, anchors, concepts, seed, steps, dropout, same, queries, sdsa, injection, alpha): + import modules.consistory as cs + t0 = time.time() + if len(anchors) == 0: + shared.log.warning('ConsiStory: no anchors') + return [] + shared.log.debug(f'ConsiStory anchors: concepts={concepts} anchors={anchors}') + with devices.inference_context(): + try: + images, self.anchor_cache_first_stage, self.anchor_cache_second_stage = cs.run_anchor_generation( + story_pipeline=shared.sd_model, + prompts=anchors, + concept_token=concepts, + seed=seed, + n_steps=steps, + mask_dropout=dropout, + same_latent=same, + share_queries=queries, + perform_sdsa=sdsa, + inject_range_alpha=alpha, + perform_injection=injection, + ) + except Exception as e: + shared.log.error(f'ConsiStory: {e}') + errors.display(e, 'ConsiStory') + images = [] + devices.torch_gc() + t1 = time.time() + shared.log.debug(f'ConsiStory anchors: images={len(images)} time={t1-t0:.2f}') + return images + + def create_extra(self, prompt, concepts, seed, steps, dropout, same, queries, sdsa, injection, alpha): + import modules.consistory as cs + t0 = time.time() + images = [] + shared.log.debug(f'ConsiStory extra: concepts={concepts} prompt="{prompt}"') + with devices.inference_context(): + try: + images = cs.run_extra_generation( + story_pipeline=shared.sd_model, + prompts=[prompt], + concept_token=concepts, + anchor_cache_first_stage=self.anchor_cache_first_stage, + anchor_cache_second_stage=self.anchor_cache_second_stage, + seed=seed, + n_steps=steps, + mask_dropout=dropout, + same_latent=same, + share_queries=queries, + perform_sdsa=sdsa, + inject_range_alpha=alpha, + perform_injection=injection, + ) + except Exception as e: + shared.log.error(f'ConsiStory: {e}') + errors.display(e, 'ConsiStory') + images = [] + devices.torch_gc() + t1 = time.time() + shared.log.debug(f'ConsiStory extra: images={len(images)} time={t1-t0:.2f}') + return images + + def run(self, p: processing.StableDiffusionProcessing, *args): # pylint: disable=arguments-differ + supported_model_list = ['sdxl'] + if shared.sd_model_type not in supported_model_list and shared.sd_model.__class__.__name__ != 'ConsistoryExtendAttnSDXLPipeline': + shared.log.warning(f'ConsiStory: class={shared.sd_model.__class__.__name__} model={shared.sd_model_type} required={supported_model_list}') + return None + + subject, concepts, prompts, dropout, sampler, steps, same, queries, sdsa, freeu, freeu_preset, alpha, injection = args # pylint: disable=unused-variable + + self.create_model() # create model if not already done + concepts, anchors, prompts, alpha, steps, seed = self.set_args(p, *args) # set arguments + + images = [] + if self.anchor_cache_first_stage is None or self.anchor_cache_second_stage is None: # create anchors if not cached + images = self.create_anchors(anchors, concepts, seed, steps, dropout, same, queries, sdsa, injection, alpha) + + for prompt in prompts: + extra_out_images = self.create_extra(prompt, concepts, seed, steps, dropout, same, queries, sdsa, injection, alpha) + for image in extra_out_images: + images.append(image) + + shared.sd_model.disable_freeu() + processed = processing.Processed(p, images) + return processed + + def after(self, p: processing.StableDiffusionProcessing, processed: processing.Processed, *args): # pylint: disable=arguments-differ, unused-argument + return processed diff --git a/scripts/instantir.py b/scripts/instantir.py new file mode 100644 index 000000000..4c7ce77b7 --- /dev/null +++ b/scripts/instantir.py @@ -0,0 +1,97 @@ +import gradio as gr +import torch +import diffusers +from huggingface_hub import hf_hub_download +from modules import scripts, processing, shared, sd_models, devices, ipadapter + + +class Script(scripts.Script): + def __init__(self): + super().__init__() + self.orig_pipe = None + self.orig_ip_unapply = None + + def title(self): + return 'InstantIR' + + def show(self, is_img2img): + return is_img2img if shared.native else False + + def ui(self, _is_img2img): # ui elements + with gr.Row(): + gr.HTML('  InstantIR: Image Restoration
') + with gr.Row(): + start = gr.Slider(label='Preview start', minimum=0.0, maximum=1.0, step=0.01, value=0.0) + end = gr.Slider(label='Preview end', minimum=0.0, maximum=1.0, step=0.01, value=1.0) + with gr.Row(): + hq = gr.Checkbox(label='HQ init latents', value=False) + multistep = gr.Checkbox(label='Multistep restore', value=False) + adastep = gr.Checkbox(label='Adaptive restore', value=False) + with gr.Row(): + image = gr.Image(label='Override guidance image') + return [start, end, hq, multistep, adastep, image] + + def run(self, p: processing.StableDiffusionProcessing, *args): # pylint: disable=arguments-differ + supported_model_list = ['sdxl'] + if not hasattr(p, 'init_images') or len(p.init_images) == 0: + shared.log.warning('InstantIR: no image') + return None + if shared.sd_model_type not in supported_model_list and shared.sd_model.__class__.__name__ != "InstantIRPipeline": + shared.log.warning(f'InstantIR: class={shared.sd_model.__class__.__name__} model={shared.sd_model_type} required={supported_model_list}') + return None + start, end, hq, multistep, adastep, image = args + from modules import instantir as ir + if shared.sd_model_type == "sdxl": + if shared.sd_model.__class__.__name__ != "InstantIRPipeline": + self.orig_pipe = shared.sd_model + self.orig_ip_unapply = ipadapter.unapply + shared.sd_model = sd_models.switch_pipe(ir.InstantIRPipeline, shared.sd_model) + adapter_file = hf_hub_download('InstantX/InstantIR', subfolder='models', filename='adapter.pt', cache_dir=shared.opts.hfcache_dir) + aggregator_file = hf_hub_download('InstantX/InstantIR', subfolder='models', filename='aggregator.pt', cache_dir=shared.opts.hfcache_dir) + previewer_file = hf_hub_download('InstantX/InstantIR', subfolder='models', filename='previewer_lora_weights.bin', cache_dir=shared.opts.hfcache_dir) + shared.log.debug(f'InstantIR: adapter="{adapter_file}" aggregator="{aggregator_file}" previewer="{previewer_file}"') + ir.load_adapter_to_pipe( + pipe=shared.sd_model, + pretrained_model_path_or_dict=adapter_file, + image_encoder_or_path='facebook/dinov2-large', + use_lcm=False, + use_adaln=True, + ) + shared.sd_model.prepare_previewers(previewer_file) + shared.sd_model.scheduler = diffusers.DDPMScheduler.from_pretrained('stabilityai/stable-diffusion-xl-base-1.0', subfolder="scheduler") + pretrained_state_dict = torch.load(aggregator_file) + shared.sd_model.aggregator.load_state_dict(pretrained_state_dict) + shared.sd_model.aggregator.to(device=devices.device, dtype=devices.dtype) + + shared.log.info(f'InstantIR: class={shared.sd_model.__class__.__name__} start={start} end={end} multistep={multistep} adastep={adastep} hq={hq} cache={shared.opts.hfcache_dir}') + p.sampler_name = 'Default' # ir has its own sampler + p.init() # run init early to take care of resizing + p.task_args['previewer_scheduler'] = ir.LCMSingleStepScheduler.from_config(shared.sd_model.scheduler.config) + p.task_args['image'] = p.init_images + p.task_args['save_preview_row'] = False + p.task_args['init_latents_with_lq'] = not hq + p.task_args['multistep_restore'] = multistep + p.task_args['adastep_restore'] = adastep + p.task_args['preview_start'] = start + p.task_args['preview_end'] = end + p.task_args['ip_adapter_image'] = image + p.extra_generation_params["InstantIR"] = f'Start={start} End={end} HQ={hq} Multistep={multistep} Adastep={adastep}' + ipadapter.unapply = lambda x: x # disable as main processing unloads ipadapter as it thinks its not needed + devices.torch_gc() + + def after(self, p: processing.StableDiffusionProcessing, processed: processing.Processed, *args): # pylint: disable=arguments-differ, unused-argument + # TODO instantir is a mess to unload + """ + if self.orig_pipe is None: + return processed + if hasattr(shared.sd_model, 'aggregator'): + shared.sd_model.aggregator = None + shared.log.debug(f'InstantIR restore: class={shared.sd_model.__class__.__name__}') + shared.sd_model = self.orig_pipe + self.orig_pipe = None + shared.sd_model.unet.register_to_config(encoder_hid_dim_type=None) + ipadapter.unapply = self.orig_ip_unapply + ipadapter.unapply(shared.sd_model) + devices.torch_gc() + """ + return processed diff --git a/scripts/ipadapter.py b/scripts/ipadapter.py index c7fcb3053..60c70b9dc 100644 --- a/scripts/ipadapter.py +++ b/scripts/ipadapter.py @@ -1,7 +1,7 @@ import json from PIL import Image import gradio as gr -from modules import scripts, processing, shared, ipadapter +from modules import scripts, processing, shared, ipadapter, ui_common MAX_ADAPTERS = 4 @@ -60,9 +60,12 @@ class Script(scripts.Script): for i in range(MAX_ADAPTERS): with gr.Accordion(f'Adapter {i+1}', visible=i==0) as unit: with gr.Row(): - adapters.append(gr.Dropdown(label='Adapter', choices=list(ipadapter.ADAPTERS), value='None')) - scales.append(gr.Slider(label='Scale', minimum=0.0, maximum=1.0, step=0.01, value=0.5)) - crops.append(gr.Checkbox(label='Crop', default=False, interactive=True)) + adapter = gr.Dropdown(label='Adapter', choices=list(ipadapter.get_adapters()), value='None') + adapters.append(adapter) + ui_common.create_refresh_button(adapter, ipadapter.get_adapters) + with gr.Row(): + scales.append(gr.Slider(label='Strength', minimum=0.0, maximum=1.0, step=0.01, value=0.5)) + crops.append(gr.Checkbox(label='Crop to portrait', default=False, interactive=True)) with gr.Row(): starts.append(gr.Slider(label='Start', minimum=0.0, maximum=1.0, step=0.1, value=0)) ends.append(gr.Slider(label='End', minimum=0.0, maximum=1.0, step=0.1, value=1)) diff --git a/scripts/ipinstruct.py b/scripts/ipinstruct.py new file mode 100644 index 000000000..4a94197b1 --- /dev/null +++ b/scripts/ipinstruct.py @@ -0,0 +1,111 @@ +""" +Repo: +Models: +adapter: `sd15`=0.35GB `sdxl`=2.12GB `sd3`=1.56GB +encoder: `laion/CLIP-ViT-H-14-laion2B-s32B-b79K`=3.94GB +""" +import os +import importlib +import gradio as gr +from modules import scripts, processing, shared, sd_models, devices + + +repo = 'https://github.com/vladmandic/IP-Instruct' +repo_id = 'CiaraRowles/IP-Adapter-Instruct' +encoder = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K" +folder = os.path.join('repositories', 'ip_instruct') + + +class Script(scripts.Script): + def __init__(self): + super().__init__() + self.orig_pipe = None + self.lib = None + + def title(self): + return 'IP Instruct' + + def show(self, is_img2img): + if shared.cmd_opts.experimental: + return not is_img2img if shared.native else False + else: + return False + + def install(self): + if not os.path.exists(folder): + from installer import clone + clone(repo, folder) + if self.lib is None: + self.lib = importlib.import_module('ip_instruct.ip_adapter') + + + def ui(self, _is_img2img): # ui elements + with gr.Row(): + gr.HTML('  IP Adapter Instruct
') + with gr.Row(): + query = gr.Textbox(lines=1, label='Query', placeholder='use the composition from the image') + with gr.Row(): + image = gr.Image(value=None, label='Image', type='pil', source='upload', width=256, height=256) + with gr.Row(): + strength = gr.Slider(label="Strength", value=1.0, minimum=0, maximum=2.0, step=0.05) + tokens = gr.Slider(label="Tokens", value=4, minimum=1, maximum=32, step=1) + with gr.Row(): + instruct_guidance = gr.Slider(label="Guidance", value=6.0, minimum=1.0, maximum=15.0, step=0.05) + image_guidance = gr.Slider(label="Guidance", value=0.5, minimum=0, maximum=1.0, step=0.05) + return [query, image, strength, tokens, instruct_guidance, image_guidance] + + def run(self, p: processing.StableDiffusionProcessing, query, image, strength, tokens, instruct_guidance, image_guidance): # pylint: disable=arguments-differ + supported_model_list = ['sd', 'sdxl', 'sd3'] + if shared.sd_model_type not in supported_model_list: + shared.log.warning(f'IP-Instruct: class={shared.sd_model.__class__.__name__} model={shared.sd_model_type} required={supported_model_list}') + return None + self.install() + if self.lib is None: + shared.log.error('IP-Instruct: failed to import library') + return None + self.orig_pipe = shared.sd_model + if shared.sd_model_type == 'sdxl': + pipe = self.lib.StableDiffusionXLPipelineExtraCFG + cls = self.lib.IPAdapterInstructSDXL + ckpt = "ip-adapter-instruct-sdxl.bin" + elif shared.sd_model_type == 'sd3': + pipe = self.lib.StableDiffusion3PipelineExtraCFG + cls = self.lib.IPAdapter_sd3_Instruct + ckpt = "ip-adapter-instruct-sd3.bin" + else: + pipe = self.lib.StableDiffusionPipelineCFG + cls = self.lib.IPAdapterInstruct + ckpt = "ip-adapter-instruct-sd15.bin" + + shared.sd_model = sd_models.switch_pipe(pipe, shared.sd_model) + + import huggingface_hub as hf + ip_ckpt = hf.hf_hub_download(repo_id=repo_id, filename=ckpt, cache_dir=shared.opts.hfcache_dir) + ip_model = cls(shared.sd_model, encoder, ip_ckpt, device=devices.device, dtypein=devices.dtype, num_tokens=tokens) + processing.fix_seed(p) + shared.log.debug(f'IP-Instruct: class={shared.sd_model.__class__.__name__} wrapper={ip_model.__class__.__name__} encoder={encoder} adapter={ckpt}') + shared.log.info(f'IP-Instruct: image={image} query="{query}" strength={strength} tokens={tokens} instruct_guidance={instruct_guidance} image_guidance={image_guidance}') + + image_list = ip_model.generate( + query = query, + scale = strength, + instruct_guidance_scale = instruct_guidance, + image_guidance_scale = image_guidance, + + prompt = p.prompt, + pil_image = image, + num_samples = 1, + num_inference_steps = p.steps, + seed = p.seed, + guidance_scale = p.cfg_scale, + auto_scale = False, + simple_cfg_mode = False, + ) + processed = processing.Processed(p, images_list=image_list, seed=p.seed, subseed=p.subseed, index_of_first_image=0) # manually created processed object + # p.extra_generation_params["IPInstruct"] = f'' + return processed + + def after(self, p: processing.StableDiffusionProcessing, processed: processing.Processed, **kwargs): # pylint: disable=unused-argument + if self.orig_pipe is not None: + shared.sd_model = self.orig_pipe + return processed diff --git a/scripts/k_diff.py b/scripts/k_diff.py new file mode 100644 index 000000000..354df5d4b --- /dev/null +++ b/scripts/k_diff.py @@ -0,0 +1,74 @@ +import inspect +import gradio as gr +import diffusers +from modules import scripts, processing, shared, sd_models + + +class Script(scripts.Script): + supported_models = ['sd', 'sdxl'] + orig_pipe = None + + def title(self): + return 'K-Diffusion' + + def show(self, is_img2img): + return not is_img2img if shared.native else False + + def ui(self, _is_img2img): # ui elements + with gr.Row(): + gr.HTML('  K-Diffusion samplers
') + with gr.Row(): + sampler = gr.Dropdown(label="Sampler", choices=self.samplers()) + return [sampler] + + def samplers(self): + samplers = [] + from modules import sd_samplers_kdiffusion + for s in dir(sd_samplers_kdiffusion.k_sampling): + if s.startswith('sample_'): + samplers.append(s.replace('sample_', '')) + return samplers + + def callback(self, d): + _step = d['i'] + + def run(self, p: processing.StableDiffusionProcessing, sampler: str): # pylint: disable=arguments-differ + if sampler is None or len(sampler) == 0: + return None + if shared.sd_model_type not in self.supported_models: + shared.log.warning(f'K-Diffusion: class={shared.sd_model.__class__.__name__} model={shared.sd_model_type} required={self.supported_models}') + return None + cls = None + if shared.sd_model_type == "sd": + cls = diffusers.pipelines.StableDiffusionKDiffusionPipeline + if shared.sd_model_type == "sdxl": + cls = diffusers.pipelines.StableDiffusionXLKDiffusionPipeline + if cls is None: + return None + from modules import sd_samplers_kdiffusion + + sampler_fn = getattr(sd_samplers_kdiffusion.k_sampling, f'sample_{sampler}', None) + if sampler_fn is None: + shared.log.warning(f'K-Diffusion: sampler={sampler} not found') + return None + + self.orig_pipe = shared.sd_model + shared.sd_model = sd_models.switch_pipe(cls, shared.sd_model) + shared.sd_model.sampler = sampler_fn + + params = inspect.signature(shared.sd_model.sampler).parameters.values() + params = {param.name: param.default for param in params if param.default != inspect.Parameter.empty} + # if 'callback' in list(params): + # params['callback'] = self.callback + # if 'disable' in list(params): + # params['disable'] = False + shared.log.info(f'K-diffusion apply: class={shared.sd_model.__class__.__name__} sampler={sampler} params={params}') + p.extra_generation_params["Sampler"] = sampler + + def after(self, p: processing.StableDiffusionProcessing, processed: processing.Processed, sampler): # pylint: disable=arguments-differ, unused-argument + if self.orig_pipe is None: + return processed + if shared.sd_model_type == "sdxl" or shared.sd_model_type == "sd": + shared.sd_model = self.orig_pipe + self.orig_pipe = None + return processed diff --git a/scripts/prompt_enhance.py b/scripts/prompt_enhance.py index 2fb513e76..a02aea608 100644 --- a/scripts/prompt_enhance.py +++ b/scripts/prompt_enhance.py @@ -32,10 +32,10 @@ class Script(scripts.Script): def load(self): if self.tokenizer is None: - self.tokenizer = AutoTokenizer.from_pretrained('gokaygokay/Flux-Prompt-Enhance', cache_dir=shared.opts.diffusers_dir) + self.tokenizer = AutoTokenizer.from_pretrained('gokaygokay/Flux-Prompt-Enhance', cache_dir=shared.opts.hfcache_dir) if self.model is None: shared.log.info(f'Prompt enhance: model="{repo_id}"') - self.model = AutoModelForSeq2SeqLM.from_pretrained('gokaygokay/Flux-Prompt-Enhance', cache_dir=shared.opts.diffusers_dir).to(device=devices.cpu, dtype=devices.dtype) + self.model = AutoModelForSeq2SeqLM.from_pretrained('gokaygokay/Flux-Prompt-Enhance', cache_dir=shared.opts.hfcache_dir).to(device=devices.cpu, dtype=devices.dtype) def enhance(self, prompt, auto_apply: bool = False, temperature: float = 0.7, repetition_penalty: float = 1.2, max_length: int = 128): self.load() diff --git a/scripts/pulid_ext.py b/scripts/pulid_ext.py new file mode 100644 index 000000000..d31c18164 --- /dev/null +++ b/scripts/pulid_ext.py @@ -0,0 +1,206 @@ +import io +import os +import contextlib +import gradio as gr +import numpy as np +from PIL import Image +from modules import shared, devices, errors, scripts, processing, processing_helpers, sd_models + + +debug = os.environ.get('SD_PULID_DEBUG', None) is not None + + +class Script(scripts.Script): + def __init__(self): + self.images = [] + self.pulid = None + self.cache = None + super().__init__() + self.register() # pulid is script with processing override so xyz doesnt execute + + def title(self): + return 'PuLID' + + def show(self, _is_img2img): + return shared.native + + def dependencies(self): + from installer import install, installed + if not installed('insightface', reload=False, quiet=True): + install('insightface', 'insightface', ignore=False) + install('albumentations==1.4.3', 'albumentations', ignore=False, reinstall=True) + install('pydantic==1.10.15', 'pydantic', ignore=False, reinstall=True) + # if not installed('apex', reload=False, quiet=True): + # install('apex', 'apex', ignore=False) + + def register(self): # register xyz grid elements + def apply_field(field): + def fun(p, x, xs): # pylint: disable=unused-argument + setattr(p, field, x) + self.run(p) + return fun + + import sys + xyz_classes = [v for k, v in sys.modules.items() if 'xyz_grid_classes' in k][0] + xyz_classes.axis_options.append(xyz_classes.AxisOption("[PuLID] Strength", float, apply_field("pulid_strength"))) + xyz_classes.axis_options.append(xyz_classes.AxisOption("[PuLID] Zero", int, apply_field("pulid_zero"))) + xyz_classes.axis_options.append(xyz_classes.AxisOption("[PuLID] Ortho", str, apply_field("pulid_ortho"), choices=lambda: ['off', 'v1', 'v2'])) + + def load_images(self, files): + self.images = [] + for file in files or []: + try: + if isinstance(file, str): + from modules.api.api import decode_base64_to_image + image = decode_base64_to_image(file) + elif isinstance(file, Image.Image): + image = file + elif isinstance(file, dict) and 'name' in file: + image = Image.open(file['name']) # _TemporaryFileWrapper from gr.Files + elif hasattr(file, 'name'): + image = Image.open(file.name) # _TemporaryFileWrapper from gr.Files + else: + raise ValueError(f'IP adapter unknown input: {file}') + self.images.append(image) + except Exception as e: + shared.log.warning(f'IP adapter failed to load image: {e}') + return gr.update(value=self.images, visible=len(self.images) > 0) + + # return signature is array of gradio components + def ui(self, _is_img2img): + with gr.Row(): + gr.HTML('  PuLID: Pure and Lightning ID Customization
') + with gr.Row(): + strength = gr.Slider(label = 'Strength', value = 0.8, mininimum = 0, maximum = 1, step = 0.01) + zero = gr.Slider(label = 'Zero', value = 20, mininimum = 0, maximum = 80, step = 1) + with gr.Row(): + sampler = gr.Dropdown(label="Sampler", choices=['dpmpp_sde', 'dpmpp_2m'], value='dpmpp_sde', visible=True) + ortho = gr.Dropdown(label="Ortho", choices=['off', 'v1', 'v2'], value='v2', visible=True) + with gr.Row(): + files = gr.File(label='Input images', file_count='multiple', file_types=['image'], type='file', interactive=True, height=100) + with gr.Row(): + gallery = gr.Gallery(show_label=False, value=[], visible=False, container=False, rows=1) + files.change(fn=self.load_images, inputs=[files], outputs=[gallery]) + return [strength, zero, sampler, ortho, gallery] + + def run(self, p: processing.StableDiffusionProcessing, strength: float = 0.8, zero: int = 20, sampler: str = 'dpmpp_sde', ortho: str = 'v2', gallery: list = []): # pylint: disable=arguments-differ + images = [] + try: + if len(gallery) == 0: + from modules.api.api import decode_base64_to_image + images = getattr(p, 'pulid_images', self.images) + images = [decode_base64_to_image(image) if isinstance(image, str) else image for image in images] + else: + images = [Image.open(f['name']) if isinstance(f, dict) else f for f in gallery] + images = [np.array(image) for image in images] + except Exception as e: + shared.log.error(f'PuLID: failed to load images: {e}') + return None + if len(images) == 0: + shared.log.error('PuLID: no images') + return None + supported_model_list = ['sdxl'] + if shared.sd_model_type not in supported_model_list: + shared.log.error(f'PuLID: class={shared.sd_model.__class__.__name__} model={shared.sd_model_type} required={supported_model_list}') + return None + if self.pulid is None: + self.dependencies() + try: + from modules import pulid # pylint: disable=redefined-outer-name + self.pulid = pulid + # from diffusers import pipelines + # pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING["pilid"] = pulid.StableDiffusionXLPuLIDPipeline + # pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["omnigen"] = pulid.StableDiffusionXLPuLIDPipelineImg2Img + except Exception as e: + shared.log.error(f'PuLID: failed to import library: {e}') + return None + if self.pulid is None: + shared.log.error('PuLID: failed to load PuLID library') + return None + if p.batch_size > 1: + shared.log.warning('PuLID: batch size not supported') + p.batch_size = 1 + + strength = getattr(p, 'pulid_strength', strength) + zero = getattr(p, 'pulid_zero', zero) + ortho = getattr(p, 'pulid_ortho', ortho) + + if shared.sd_model_type == 'sdxl' and not hasattr(shared.sd_model, 'pipe'): + try: + stdout = io.StringIO() + ctx = contextlib.nullcontext if debug else contextlib.redirect_stdout(stdout) + with ctx: + shared.sd_model = self.pulid.StableDiffusionXLPuLIDPipeline( + pipe =shared.sd_model, + device=devices.device, + sampler=sampler, + cache_dir=shared.opts.hfcache_dir, + ) + shared.sd_model.no_recurse = True + sd_models.copy_diffuser_options(shared.sd_model, shared.sd_model.pipe) + sd_models.move_model(shared.sd_model, devices.device) # move pipeline to device + sd_models.set_diffuser_options(shared.sd_model, vae=None, op='model') + devices.torch_gc() + except Exception as e: + shared.log.error(f'PuLID: failed to create pipeline: {e}') + errors.display(e, 'PuLID') + return None + + shared.log.info(f'PuLID: class={shared.sd_model.__class__.__name__} strength={strength} zero={zero} ortho={ortho} sampler={sampler} images={[i.shape for i in images]}') + self.pulid.attention.NUM_ZERO = zero + self.pulid.attention.ORTHO = ortho == 'v1' + self.pulid.attention.ORTHO_v2 = ortho == 'v2' + images = [self.pulid.resize(image, 1024) for image in images] + shared.sd_model.debug_img_list = [] + uncond_id_embedding, id_embedding = shared.sd_model.get_id_embedding(images) + + if debug: # run pipeline directly + shared.state.begin('PuLID') + processing.fix_seed(p) + p.seed = processing_helpers.get_fixed_seed(p.seed) + p.prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) + p.negative_prompt = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) + with devices.inference_context(): + output = shared.sd_model( + prompt=p.prompt, + negative_prompt=p.negative_prompt, + width=p.width, + height=p.height, + seed=p.seed, + num_inference_steps=p.steps, + guidance_scale=p.cfg_scale, + id_embedding=id_embedding, + uncond_id_embedding=uncond_id_embedding, + id_scale=strength, + )[0] + info = processing.create_infotext(p) + processed = processing.Processed(p, [output], info=info) + shared.state.end('PuLID') + else: # let processing run the pipeline + p.task_args['id_embedding'] = id_embedding + p.task_args['uncond_id_embedding'] = uncond_id_embedding + p.task_args['id_scale'] = strength + if len(getattr(p, 'init_images', [])) > 0: + p.task_args['image'] = p.init_images[0] + p.task_args['strength'] = p.denoising_strength + p.extra_generation_params["PuLID"] = f'Strength={strength} Zero={zero} Ortho={ortho}' + if getattr(p, 'xyz', False): # xyz will run its own processing + return None + processed: processing.Processed = processing.process_images(p) # runs processing using main loop + + # interim = [Image.fromarray(img) for img in shared.sd_model.debug_img_list] + # shared.log.debug(f'PuLID: time={t1-t0:.2f}') + return processed + + def after(self, p: processing.StableDiffusionProcessing, processed: processing.Processed, *args): # pylint: disable=unused-argument + if hasattr(shared.sd_model, 'pipe') and shared.sd_model_type == "sdxl": + if hasattr(shared.sd_model, 'app'): + shared.sd_model.app = None + shared.sd_model.ip_adapter = None + shared.sd_model.face_helper = None + shared.sd_model.clip_vision_model = None + shared.sd_model.handler_ante = None + devices.torch_gc(force=True) + shared.sd_model = shared.sd_model.pipe + # shared.log.debug(f'PuLID restore: class={shared.sd_model.__class__.__name__}') + return processed diff --git a/scripts/text2video.py b/scripts/text2video.py index 2c93abf27..8dec9bd0e 100644 --- a/scripts/text2video.py +++ b/scripts/text2video.py @@ -84,7 +84,7 @@ class Script(scripts.Script): shared.log.error(f'Text2Video: failed to find model={model["path"]}') return shared.log.debug(f'Text2Video loading: model={checkpoint}') - shared.opts.sd_model_checkpoint = checkpoint + shared.opts.sd_model_checkpoint = checkpoint.name sd_models.reload_model_weights(op='model') p.ops.append('text2video') diff --git a/scripts/x_adapter.py b/scripts/x_adapter.py index 7c1341701..553a20d30 100644 --- a/scripts/x_adapter.py +++ b/scripts/x_adapter.py @@ -22,7 +22,7 @@ class Script(scripts.Script): with gr.Row(): gr.HTML('  X-Adapter
') with gr.Row(): - model = gr.Dropdown(label='Adapter model', choices=['None'] + sd_models.checkpoint_tiles(), value='None') + model = gr.Dropdown(label='Adapter model', choices=['None'] + sd_models.checkpoint_titles(), value='None') sampler = gr.Dropdown(label='Adapter sampler', choices=[s.name for s in sd_samplers.samplers], value='Default') with gr.Row(): width = gr.Slider(label='Adapter width', minimum=64, maximum=2048, step=8, value=1024) @@ -34,7 +34,7 @@ class Script(scripts.Script): lora = gr.Textbox('', label='Adapter LoRA', default='') return model, sampler, width, height, start, scale, lora - def run(self, p: processing.StableDiffusionProcessing, model, sampler, width, height, start, scale, lora): # pylint: disable=arguments-differ + def run(self, p: processing.StableDiffusionProcessing, model, sampler, width, height, start, scale, lora): # pylint: disable=arguments-differ, unused-argument from modules.xadapter.xadapter_hijacks import PositionNet diffusers.models.embeddings.PositionNet = PositionNet # patch diffusers==0.26 from diffusers==0.20 from modules.xadapter.adapter import Adapter_XL diff --git a/scripts/xyz_grid.py b/scripts/xyz_grid.py index e8f649182..60e608c76 100644 --- a/scripts/xyz_grid.py +++ b/scripts/xyz_grid.py @@ -1,4 +1,5 @@ # xyz grid that shows as selectable script +import os import csv import random from collections import namedtuple @@ -16,6 +17,9 @@ from modules.ui_components import ToolButton import modules.ui_symbols as symbols +debug = shared.log.trace if os.environ.get('SD_XYZ_DEBUG', None) is not None else lambda *args, **kwargs: None + + class Script(scripts.Script): current_axis_options = [] @@ -26,6 +30,7 @@ class Script(scripts.Script): self.current_axis_options = [x for x in axis_options if type(x) == AxisOption or x.is_img2img == is_img2img] with gr.Row(): gr.HTML('  XYZ Grid
') + with gr.Row(): with gr.Column(): with gr.Row(variant='compact'): @@ -43,17 +48,42 @@ class Script(scripts.Script): z_values = gr.Textbox(label="Z values", container=True, lines=1, elem_id=self.elem_id("z_values")) z_values_dropdown = gr.Dropdown(label="Z values", container=True, visible=False, multiselect=True, interactive=True) fill_z_button = ToolButton(value=symbols.fill, elem_id="xyz_grid_fill_z_tool_button", visible=False) + with gr.Row(): with gr.Column(): csv_mode = gr.Checkbox(label='Text inputs', value=False, elem_id=self.elem_id("csv_mode"), container=False) draw_legend = gr.Checkbox(label='Legend', value=True, elem_id=self.elem_id("draw_legend"), container=False) no_fixed_seeds = gr.Checkbox(label='Random seeds', value=False, elem_id=self.elem_id("no_fixed_seeds"), container=False) + include_time = gr.Checkbox(label='Add time info', value=False, elem_id=self.elem_id("include_time"), container=False) + include_text = gr.Checkbox(label='Add text info', value=False, elem_id=self.elem_id("include_text"), container=False) with gr.Column(): - no_grid = gr.Checkbox(label='Skip grid', value=False, elem_id=self.elem_id("no_xyz_grid"), container=False) - include_lone_images = gr.Checkbox(label='Sub-images', value=False, elem_id=self.elem_id("include_lone_images"), container=False) - include_sub_grids = gr.Checkbox(label='Sub-grids', value=False, elem_id=self.elem_id("include_sub_grids"), container=False) + include_grid = gr.Checkbox(label='Include main grid', value=True, elem_id=self.elem_id("no_xyz_grid"), container=False) + include_subgrids = gr.Checkbox(label='Include sub grids', value=False, elem_id=self.elem_id("include_sub_grids"), container=False) + include_images = gr.Checkbox(label='Include images', value=False, elem_id=self.elem_id("include_lone_images"), container=False) + create_video = gr.Checkbox(label='Create video', value=False, elem_id=self.elem_id("xyz_create_video"), container=False) + + with gr.Row(visible=False) as ui_video: + def video_type_change(video_type): + return [ + gr.update(visible=video_type != 'None'), + gr.update(visible=video_type == 'GIF' or video_type == 'PNG'), + gr.update(visible=video_type == 'MP4'), + gr.update(visible=video_type == 'MP4'), + ] + + with gr.Column(): + video_type = gr.Dropdown(label='Video type', choices=['None', 'GIF', 'PNG', 'MP4'], value='None') + with gr.Column(): + video_duration = gr.Slider(label='Duration', minimum=0.25, maximum=300, step=0.25, value=2, visible=False) + video_loop = gr.Checkbox(label='Loop', value=True, visible=False, elem_id="control_video_loop") + video_pad = gr.Slider(label='Pad frames', minimum=0, maximum=24, step=1, value=1, visible=False) + video_interpolate = gr.Slider(label='Interpolate frames', minimum=0, maximum=24, step=1, value=0, visible=False) + video_type.change(fn=video_type_change, inputs=[video_type], outputs=[video_duration, video_loop, video_pad, video_interpolate]) + create_video.change(fn=lambda x: gr.update(visible=x), inputs=[create_video], outputs=[ui_video]) + with gr.Row(): margin_size = gr.Slider(label="Grid margins", minimum=0, maximum=500, value=0, step=2, elem_id=self.elem_id("margin_size")) + with gr.Row(): swap_xy_axes_button = gr.Button(value="Swap X/Y", elem_id="xy_grid_swap_axes_button", variant="secondary") swap_yz_axes_button = gr.Button(value="Swap Y/Z", elem_id="yz_grid_swap_axes_button", variant="secondary") @@ -130,10 +160,25 @@ class Script(scripts.Script): (z_values_dropdown, lambda params:get_dropdown_update_from_params("Z",params)), ) - return [x_type, x_values, x_values_dropdown, y_type, y_values, y_values_dropdown, z_type, z_values, z_values_dropdown, csv_mode, draw_legend, no_fixed_seeds, no_grid, include_lone_images, include_sub_grids, margin_size] + return [ + x_type, x_values, x_values_dropdown, + y_type, y_values, y_values_dropdown, + z_type, z_values, z_values_dropdown, + csv_mode, draw_legend, no_fixed_seeds, + include_grid, include_subgrids, include_images, + include_time, include_text, margin_size, + create_video, video_type, video_duration, video_loop, video_pad, video_interpolate, + ] - def run(self, p, x_type, x_values, x_values_dropdown, y_type, y_values, y_values_dropdown, z_type, z_values, z_values_dropdown, csv_mode, draw_legend, no_fixed_seeds, no_grid, include_lone_images, include_sub_grids, margin_size): # pylint: disable=W0221 - shared.log.debug(f'xyzgrid: x_type={x_type}|x_values={x_values}|x_values_dropdown={x_values_dropdown}|y_type={y_type}|{y_values}={y_values}|{y_values_dropdown}={y_values_dropdown}|z_type={z_type}|z_values={z_values}|z_values_dropdown={z_values_dropdown}|draw_legend={draw_legend}|include_lone_images={include_lone_images}|include_sub_grids={include_sub_grids}|no_grid={no_grid}|margin_size={margin_size}') + def run(self, p, + x_type, x_values, x_values_dropdown, + y_type, y_values, y_values_dropdown, + z_type, z_values, z_values_dropdown, + csv_mode, draw_legend, no_fixed_seeds, + include_grid, include_subgrids, include_images, + include_time, include_text, margin_size, + create_video, video_type, video_duration, video_loop, video_pad, video_interpolate, + ): # pylint: disable=W0221 if not no_fixed_seeds: processing.fix_seed(p) if not shared.opts.return_grid: @@ -258,6 +303,7 @@ class Script(scripts.Script): def cell(x, y, z, ix, iy, iz): if shared.state.interrupted: return processing.Processed(p, [], p.seed, "") + p.xyz = True pc = copy(p) pc.override_settings_restore_afterwards = False pc.styles = pc.styles[:] @@ -307,38 +353,47 @@ class Script(scripts.Script): z_labels=[z_opt.format_value(p, z_opt, z) for z in zs], cell=cell, draw_legend=draw_legend, - include_lone_images=include_lone_images, - include_sub_grids=include_sub_grids, + include_lone_images=include_images, + include_sub_grids=include_subgrids, first_axes_processed=first_axes_processed, second_axes_processed=second_axes_processed, margin_size=margin_size, - no_grid=no_grid, + no_grid=not include_grid, + include_time=include_time, + include_text=include_text, ) if not processed.images: - return processed # It broke, no further handling needed. - z_count = len(zs) - processed.infotexts[:1+z_count] = grid_infotext[:1+z_count] # Set the grid infotexts to the real ones with extra_generation_params (1 main grid + z_count sub-grids) - if not include_lone_images: - # Don't need sub-images anymore, drop from list: - if no_grid and include_sub_grids: - processed.images = processed.images[:z_count] # we don't have the main grid image, and need zero additional sub-images - else: - processed.images = processed.images[:z_count+1] # we either have the main grid image, or need one sub-images - if shared.opts.grid_save: # Auto-save main and sub-grids: - grid_count = z_count + ( 1 if not no_grid and z_count > 1 else 0 ) - for g in range(grid_count): + return processed # something broke, no further handling needed. + # processed.images = (1)*grid + (z > 1 ? z : 0)*subgrids + (x*y*z)*images + have_grid = 1 if include_grid else 0 + have_subgrids = len(zs) if len(zs) > 1 and include_subgrids else 0 + have_images = processed.images[have_grid+have_subgrids:] + processed.infotexts[:have_grid+have_subgrids] = grid_infotext[:have_grid+have_subgrids] # update infotexts with grid and subgrid info + shared.log.debug(f'XYZ grid: grid={have_grid} subgrids={have_subgrids} images={len(have_images)} total={len(processed.images)}') + + if not include_images: # dont need images anymore, drop from list: + processed.images = processed.images[:have_grid+have_subgrids] + debug(f'XYZ grid remove images: total={processed.images}') + + if shared.opts.grid_save and not shared.state.interrupted: # auto-save main and sub-grids: + for g in range(have_grid + have_subgrids): adj_g = g-1 if g > 0 else g info = processed.infotexts[g] prompt = processed.all_prompts[adj_g] seed = processed.all_seeds[adj_g] + debug(f'XYZ grid save grid: i={g+1}') images.save_image(processed.images[g], p.outpath_grids, "grid", info=info, extension=shared.opts.grid_format, prompt=prompt, seed=seed, grid=True, p=processed) - if not include_sub_grids: # Done with sub-grids, drop all related information: - for _sg in range(z_count): + + if not include_subgrids and have_subgrids > 0: # done with sub-grids, drop all related information: + for _sg in range(have_subgrids): del processed.images[1] del processed.all_prompts[1] del processed.all_seeds[1] del processed.infotexts[1] - elif no_grid: - del processed.infotexts[0] + debug(f'XYZ grid remove subgrids: total={processed.images}') + + if create_video and video_type != 'None' and not shared.state.interrupted: + images.save_video(p, filename=None, images=have_images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) + return processed diff --git a/scripts/xyz_grid_classes.py b/scripts/xyz_grid_classes.py index 8a78d2c40..4898c6b73 100644 --- a/scripts/xyz_grid_classes.py +++ b/scripts/xyz_grid_classes.py @@ -1,4 +1,4 @@ -from scripts.xyz_grid_shared import apply_field, apply_task_args, apply_setting, apply_prompt, apply_order, apply_sampler, apply_hr_sampler_name, confirm_samplers, apply_checkpoint, apply_refiner, apply_unet, apply_dict, apply_clip_skip, apply_vae, list_lora, apply_lora, apply_te, apply_styles, apply_upscaler, apply_context, apply_detailer, apply_override, apply_processing, apply_options, apply_seed, format_value_add_label, format_value, format_value_join_list, do_nothing, format_nothing, str_permutations # pylint: disable=no-name-in-module +from scripts.xyz_grid_shared import apply_field, apply_task_args, apply_setting, apply_prompt, apply_order, apply_sampler, apply_hr_sampler_name, confirm_samplers, apply_checkpoint, apply_refiner, apply_unet, apply_dict, apply_clip_skip, apply_vae, list_lora, apply_lora, apply_te, apply_styles, apply_upscaler, apply_context, apply_detailer, apply_override, apply_processing, apply_options, apply_seed, format_value_add_label, format_value, format_value_join_list, do_nothing, format_nothing, str_permutations # pylint: disable=no-name-in-module, unused-import from modules import shared, shared_items, sd_samplers, ipadapter, sd_models, sd_vae, sd_unet @@ -84,13 +84,13 @@ class SharedSettingsStackHelper(object): axis_options = [ AxisOption("Nothing", str, do_nothing, fmt=format_nothing), - AxisOption("[Model] Model", str, apply_checkpoint, cost=1.0, fmt=format_value, choices=lambda: sorted(sd_models.checkpoints_list)), + AxisOption("[Model] Model", str, apply_checkpoint, cost=1.0, fmt=format_value_add_label, choices=lambda: sorted(sd_models.checkpoints_list)), AxisOption("[Model] UNET", str, apply_unet, cost=0.8, choices=lambda: ['None'] + list(sd_unet.unet_dict)), AxisOption("[Model] VAE", str, apply_vae, cost=0.6, choices=lambda: ['None'] + list(sd_vae.vae_dict)), - AxisOption("[Model] Refiner", str, apply_refiner, cost=0.8, fmt=format_value, choices=lambda: ['None'] + sorted(sd_models.checkpoints_list)), + AxisOption("[Model] Refiner", str, apply_refiner, cost=0.8, fmt=format_value_add_label, choices=lambda: ['None'] + sorted(sd_models.checkpoints_list)), AxisOption("[Model] Text encoder", str, apply_te, cost=0.7, choices=shared_items.sd_te_items), - AxisOption("[Model] Dictionary", str, apply_dict, fmt=format_value, cost=0.9, choices=lambda: ['None'] + list(sd_models.checkpoints_list)), - AxisOption("[Prompt] Search & replace", str, apply_prompt, fmt=format_value), + AxisOption("[Model] Dictionary", str, apply_dict, fmt=format_value_add_label, cost=0.9, choices=lambda: ['None'] + list(sd_models.checkpoints_list)), + AxisOption("[Prompt] Search & replace", str, apply_prompt, fmt=format_value_add_label), AxisOption("[Prompt] Prompt order", str_permutations, apply_order, fmt=format_value_join_list), AxisOption("[Network] LoRA", str, apply_lora, cost=0.5, choices=list_lora), AxisOption("[Network] LoRA strength", float, apply_setting('extra_networks_default_multiplier')), @@ -99,7 +99,7 @@ axis_options = [ AxisOption("[Param] Height", int, apply_field("height")), AxisOption("[Param] Seed", int, apply_seed), AxisOption("[Param] Steps", int, apply_field("steps")), - AxisOption("[Param] CFG scale", float, apply_field("cfg_scale")), + AxisOption("[Param] Guidance scale", float, apply_field("cfg_scale")), AxisOption("[Param] Guidance end", float, apply_field("cfg_end")), AxisOption("[Param] Variation seed", int, apply_field("subseed")), AxisOption("[Param] Variation strength", float, apply_field("subseed_strength")), @@ -109,8 +109,8 @@ axis_options = [ AxisOption("[Process] Model args", str, apply_task_args), AxisOption("[Process] Processing args", str, apply_processing), AxisOption("[Process] Server options", str, apply_options), - AxisOptionTxt2Img("[Sampler] Name", str, apply_sampler, fmt=format_value, confirm=confirm_samplers, choices=lambda: [x.name for x in sd_samplers.samplers]), - AxisOptionImg2Img("[Sampler] Name", str, apply_sampler, fmt=format_value, confirm=confirm_samplers, choices=lambda: [x.name for x in sd_samplers.samplers_for_img2img]), + AxisOptionTxt2Img("[Sampler] Name", str, apply_sampler, fmt=format_value_add_label, confirm=confirm_samplers, choices=lambda: [x.name for x in sd_samplers.samplers]), + AxisOptionImg2Img("[Sampler] Name", str, apply_sampler, fmt=format_value_add_label, confirm=confirm_samplers, choices=lambda: [x.name for x in sd_samplers.samplers_for_img2img]), AxisOption("[Sampler] Sigma method", str, apply_setting("schedulers_sigma"), choices=lambda: ['default', 'karras', 'beta', 'exponential']), AxisOption("[Sampler] Timestep spacing", str, apply_setting("schedulers_timestep_spacing"), choices=lambda: ['default', 'linspace', 'leading', 'trailing']), AxisOption("[Sampler] Timestep range", int, apply_setting("schedulers_timesteps_range")), @@ -122,16 +122,16 @@ axis_options = [ AxisOption("[Sampler] eta delta", float, apply_setting("eta_noise_seed_delta")), AxisOption("[Sampler] eta multiplier", float, apply_setting("scheduler_eta")), AxisOption("[Refine] Upscaler", str, apply_field("hr_upscaler"), cost=0.3, choices=lambda: [*shared.latent_upscale_modes, *[x.name for x in shared.sd_upscalers]]), - AxisOption("[Refine] Sampler", str, apply_hr_sampler_name, fmt=format_value, confirm=confirm_samplers, choices=lambda: [x.name for x in sd_samplers.samplers]), + AxisOption("[Refine] Sampler", str, apply_hr_sampler_name, fmt=format_value_add_label, confirm=confirm_samplers, choices=lambda: [x.name for x in sd_samplers.samplers]), AxisOption("[Refine] Denoising strength", float, apply_field("denoising_strength")), AxisOption("[Refine] Hires steps", int, apply_field("hr_second_pass_steps")), - AxisOption("[Refine] CFG scale", float, apply_field("image_cfg_scale")), + AxisOption("[Refine] Guidance scale", float, apply_field("image_cfg_scale")), AxisOption("[Refine] Guidance rescale", float, apply_field("diffusers_guidance_rescale")), AxisOption("[Refine] Refiner start", float, apply_field("refiner_start")), AxisOption("[Refine] Refiner steps", float, apply_field("refiner_steps")), AxisOption("[Postprocess] Upscaler", str, apply_upscaler, cost=0.4, choices=lambda: [x.name for x in shared.sd_upscalers][1:]), AxisOption("[Postprocess] Context", str, apply_context, choices=lambda: ["Add with forward", "Remove with forward", "Add with backward", "Remove with backward"]), - AxisOption("[Postprocess] Detailer", str, apply_detailer, fmt=format_value), + AxisOption("[Postprocess] Detailer", str, apply_detailer, fmt=format_value_add_label), AxisOption("[HDR] Mode", int, apply_field("hdr_mode")), AxisOption("[HDR] Brightness", float, apply_field("hdr_brightness")), AxisOption("[HDR] Color", float, apply_field("hdr_color")), diff --git a/scripts/xyz_grid_draw.py b/scripts/xyz_grid_draw.py index 8f95e696b..80336fa73 100644 --- a/scripts/xyz_grid_draw.py +++ b/scripts/xyz_grid_draw.py @@ -4,23 +4,28 @@ from PIL import Image from modules import shared, images, processing -def draw_xyz_grid(p, xs, ys, zs, x_labels, y_labels, z_labels, cell, draw_legend, include_lone_images, include_sub_grids, first_axes_processed, second_axes_processed, margin_size, no_grid): # pylint: disable=unused-argument - hor_texts = [[images.GridAnnotation(x)] for x in x_labels] - ver_texts = [[images.GridAnnotation(y)] for y in y_labels] - title_texts = [[images.GridAnnotation(z)] for z in z_labels] +def draw_xyz_grid(p, xs, ys, zs, x_labels, y_labels, z_labels, cell, draw_legend, include_lone_images, include_sub_grids, first_axes_processed, second_axes_processed, margin_size, no_grid: False, include_time: False, include_text: False): # pylint: disable=unused-argument + x_texts = [[images.GridAnnotation(x)] for x in x_labels] + y_texts = [[images.GridAnnotation(y)] for y in y_labels] + z_texts = [[images.GridAnnotation(z)] for z in z_labels] list_size = (len(xs) * len(ys) * len(zs)) processed_result = None shared.state.job_count = list_size * p.n_iter t0 = time.time() + i = 0 def process_cell(x, y, z, ix, iy, iz): - nonlocal processed_result + nonlocal processed_result, i + i += 1 + shared.log.debug(f'XYZ grid process: x={ix+1}/{len(xs)} y={iy+1}/{len(ys)} z={iz+1}/{len(zs)} total={i/list_size:.2f}') def index(ix, iy, iz): return ix + iy * len(xs) + iz * len(xs) * len(ys) shared.state.job = 'grid' + p0 = time.time() processed: processing.Processed = cell(x, y, z, ix, iy, iz) + p1 = time.time() if processed_result is None: processed_result = copy(processed) if processed_result is None: @@ -30,13 +35,27 @@ def draw_xyz_grid(p, xs, ys, zs, x_labels, y_labels, z_labels, cell, draw_legend processed_result.all_prompts = [None] * list_size processed_result.all_seeds = [None] * list_size processed_result.infotexts = [None] * list_size + processed_result.time = [0] * list_size processed_result.index_of_first_image = 1 idx = index(ix, iy, iz) if processed is not None and processed.images: processed_result.images[idx] = processed.images[0] + overlay_text = '' + if include_text: + if len(x_labels[ix]) > 0: + overlay_text += f'{x_labels[ix]}\n' + if len(y_labels[iy]) > 0: + overlay_text += f'{y_labels[iy]}\n' + if len(z_labels[iz]) > 0: + overlay_text += f'{z_labels[iz]}\n' + if include_time: + overlay_text += f'Time: {p1 - p0:.2f}' + if len(overlay_text) > 0: + processed_result.images[idx] = images.draw_overlay(processed_result.images[idx], overlay_text) processed_result.all_prompts[idx] = processed.prompt processed_result.all_seeds[idx] = processed.seed processed_result.infotexts[idx] = processed.infotexts[0] + processed_result.time[idx] = round(p1 - p0, 2) else: cell_mode = "P" cell_size = (processed_result.width, processed_result.height) @@ -44,6 +63,7 @@ def draw_xyz_grid(p, xs, ys, zs, x_labels, y_labels, z_labels, cell, draw_legend cell_mode = processed_result.images[0].mode cell_size = processed_result.images[0].size processed_result.images[idx] = Image.new(cell_mode, cell_size) + return if first_axes_processed == 'x': for ix, x in enumerate(xs): @@ -93,7 +113,7 @@ def draw_xyz_grid(p, xs, ys, zs, x_labels, y_labels, z_labels, cell, draw_legend if (not no_grid or include_sub_grids) and images.check_grid_size(to_process): grid = images.image_grid(to_process, rows=len(ys)) if draw_legend: - grid = images.draw_grid_annotations(grid, w, h, hor_texts, ver_texts, margin_size, title=title_texts[i]) + grid = images.draw_grid_annotations(grid, w, h, x_texts, y_texts, margin_size, title=z_texts[i]) processed_result.images.insert(i, grid) processed_result.all_prompts.insert(i, processed_result.all_prompts[idx0]) processed_result.all_seeds.insert(i, processed_result.all_seeds[idx0]) diff --git a/scripts/xyz_grid_on.py b/scripts/xyz_grid_on.py index ac1bc4c2f..f0455f0f1 100644 --- a/scripts/xyz_grid_on.py +++ b/scripts/xyz_grid_on.py @@ -1,4 +1,5 @@ # xyz grid that shows up as alwayson script +import os import csv import random from collections import namedtuple @@ -18,6 +19,7 @@ import modules.ui_symbols as symbols active = False cache = None +debug = shared.log.trace if os.environ.get('SD_XYZ_DEBUG', None) is not None else lambda *args, **kwargs: None class Script(scripts.Script): @@ -35,6 +37,7 @@ class Script(scripts.Script): with gr.Accordion('XYZ Grid', open = False, elem_id='xyz_grid'): with gr.Row(): enabled = gr.Checkbox(label = 'Enabled', value = False) + with gr.Row(): with gr.Column(): with gr.Row(variant='compact'): @@ -52,17 +55,42 @@ class Script(scripts.Script): z_values = gr.Textbox(label="Z values", container=True, lines=1, elem_id=self.elem_id("z_values")) z_values_dropdown = gr.Dropdown(label="Z values", container=True, visible=False, multiselect=True, interactive=True) fill_z_button = ToolButton(value=symbols.fill, elem_id="xyz_grid_fill_z_tool_button", visible=False) + with gr.Row(): with gr.Column(): draw_legend = gr.Checkbox(label='Draw legend', value=True, elem_id=self.elem_id("draw_legend"), container=False) csv_mode = gr.Checkbox(label='Use text inputs', value=False, elem_id=self.elem_id("csv_mode"), container=False) no_fixed_seeds = gr.Checkbox(label='Use random seeds', value=False, elem_id=self.elem_id("no_fixed_seeds"), container=False) + include_time = gr.Checkbox(label='Add time info', value=False, elem_id=self.elem_id("include_time"), container=False) + include_text = gr.Checkbox(label='Add text info', value=False, elem_id=self.elem_id("include_text"), container=False) with gr.Column(): - include_grid = gr.Checkbox(label='Create main grid', value=True, elem_id=self.elem_id("no_xyz_grid"), container=False) - include_subgrids = gr.Checkbox(label='Create partial grids', value=False, elem_id=self.elem_id("include_sub_grids"), container=False) + include_grid = gr.Checkbox(label='Include main grid', value=True, elem_id=self.elem_id("no_xyz_grid"), container=False) + include_subgrids = gr.Checkbox(label='Include sub grids', value=False, elem_id=self.elem_id("include_sub_grids"), container=False) include_images = gr.Checkbox(label='Include images', value=False, elem_id=self.elem_id("include_lone_images"), container=False) + create_video = gr.Checkbox(label='Create video', value=False, elem_id=self.elem_id("xyz_create_video"), container=False) + + with gr.Row(visible=False) as ui_video: + def video_type_change(video_type): + return [ + gr.update(visible=video_type != 'None'), + gr.update(visible=video_type == 'GIF' or video_type == 'PNG'), + gr.update(visible=video_type == 'MP4'), + gr.update(visible=video_type == 'MP4'), + ] + + with gr.Column(): + video_type = gr.Dropdown(label='Video type', choices=['None', 'GIF', 'PNG', 'MP4'], value='None') + with gr.Column(): + video_duration = gr.Slider(label='Duration', minimum=0.25, maximum=300, step=0.25, value=2, visible=False) + video_loop = gr.Checkbox(label='Loop', value=True, visible=False, elem_id="control_video_loop") + video_pad = gr.Slider(label='Pad frames', minimum=0, maximum=24, step=1, value=1, visible=False) + video_interpolate = gr.Slider(label='Interpolate frames', minimum=0, maximum=24, step=1, value=0, visible=False) + video_type.change(fn=video_type_change, inputs=[video_type], outputs=[video_duration, video_loop, video_pad, video_interpolate]) + create_video.change(fn=lambda x: gr.update(visible=x), inputs=[create_video], outputs=[ui_video]) + with gr.Row(): margin_size = gr.Slider(label="Grid margins", minimum=0, maximum=500, value=0, step=2, elem_id=self.elem_id("margin_size")) + with gr.Row(): swap_xy_axes_button = gr.Button(value="Swap X/Y", elem_id="xy_grid_swap_axes_button", variant="secondary") swap_yz_axes_button = gr.Button(value="Swap Y/Z", elem_id="yz_grid_swap_axes_button", variant="secondary") @@ -139,9 +167,27 @@ class Script(scripts.Script): (z_values_dropdown, lambda params:get_dropdown_update_from_params("Z",params)), ) - return [enabled, x_type, x_values, x_values_dropdown, y_type, y_values, y_values_dropdown, z_type, z_values, z_values_dropdown, csv_mode, draw_legend, no_fixed_seeds, include_grid, include_subgrids, include_images, margin_size] + return [ + enabled, + x_type, x_values, x_values_dropdown, + y_type, y_values, y_values_dropdown, + z_type, z_values, z_values_dropdown, + csv_mode, draw_legend, no_fixed_seeds, + include_grid, include_subgrids, include_images, + include_time, include_text, margin_size, + create_video, video_type, video_duration, video_loop, video_pad, video_interpolate, + ] - def process(self, p, enabled, x_type, x_values, x_values_dropdown, y_type, y_values, y_values_dropdown, z_type, z_values, z_values_dropdown, csv_mode, draw_legend, no_fixed_seeds, include_grid, include_subgrids, include_images, margin_size): # pylint: disable=W0221 + def process(self, p, + enabled, + x_type, x_values, x_values_dropdown, + y_type, y_values, y_values_dropdown, + z_type, z_values, z_values_dropdown, + csv_mode, draw_legend, no_fixed_seeds, + include_grid, include_subgrids, include_images, + include_time, include_text, margin_size, + create_video, video_type, video_duration, video_loop, video_pad, video_interpolate, + ): # pylint: disable=W0221 global active, cache # pylint: disable=W0603 cache = None if not enabled or active: @@ -272,6 +318,7 @@ class Script(scripts.Script): def cell(x, y, z, ix, iy, iz): if shared.state.interrupted: return processing.Processed(p, [], p.seed, "") + p.xyz = True pc = copy(p) pc.override_settings_restore_afterwards = False pc.styles = pc.styles[:] @@ -327,33 +374,51 @@ class Script(scripts.Script): second_axes_processed=second_axes_processed, margin_size=margin_size, no_grid=not include_grid, + include_time=include_time, + include_text=include_text, ) if not processed.images: - active = False - return processed # It broke, no further handling needed. - # images stucture: main-grid, sub-grid1, sub-grid2, ..., image-1, image-2, ... - z_count = len(processed.images) - (len(zs) * len(ys) * len(xs)) # how many grids are there: main grid + sub-grids - processed.infotexts[:z_count] = grid_infotext[:z_count] # replace grid info texts - if not include_images: - processed.images = processed.images[:z_count] - if shared.opts.grid_save: # auto-save main and sub-grids: - for i in range(z_count): - info = processed.infotexts[i] - prompt = processed.all_prompts[i] - seed = processed.all_seeds[i] - _fn, _txt, _exif = images.save_image(processed.images[i], p.outpath_grids, "grid", info=info, extension=shared.opts.grid_format, prompt=prompt, seed=seed, grid=True, p=processed) - if not include_subgrids and z_count > 1: # delete sub-grids - for _sg in range(z_count - 1): + return processed # something broke, no further handling needed. + # processed.images = (1)*grid + (z > 1 ? z : 0)*subgrids + (x*y*z)*images + have_grid = 1 if include_grid else 0 + have_subgrids = len(zs) if len(zs) > 1 and include_subgrids else 0 + have_images = processed.images[have_grid+have_subgrids:] + processed.infotexts[:have_grid+have_subgrids] = grid_infotext[:have_grid+have_subgrids] # update infotexts with grid and subgrid info + shared.log.debug(f'XYZ grid: grid={have_grid} subgrids={have_subgrids} images={len(have_images)} total={len(processed.images)}') + + if not include_images: # dont need images anymore, drop from list: + processed.images = processed.images[:have_grid+have_subgrids] + debug(f'XYZ grid remove images: total={processed.images}') + + if shared.opts.grid_save and not shared.state.interrupted: # auto-save main and sub-grids: + for g in range(have_grid + have_subgrids): + adj_g = g-1 if g > 0 else g + info = processed.infotexts[g] + prompt = processed.all_prompts[adj_g] + seed = processed.all_seeds[adj_g] + debug(f'XYZ grid save grid: i={g+1}') + images.save_image(processed.images[g], p.outpath_grids, "grid", info=info, extension=shared.opts.grid_format, prompt=prompt, seed=seed, grid=True, p=processed) + + if not include_subgrids and have_subgrids > 0: # done with sub-grids, drop all related information: + for _sg in range(have_subgrids): del processed.images[1] del processed.all_prompts[1] del processed.all_seeds[1] del processed.infotexts[1] + debug(f'XYZ grid remove subgrids: total={processed.images}') + + if create_video and video_type != 'None' and not shared.state.interrupted: + images.save_video(p, filename=None, images=have_images, video_type=video_type, duration=video_duration, loop=video_loop, pad=video_pad, interpolate=video_interpolate) + p.do_not_save_grid = True p.do_not_save_samples = True active = False cache = processed return processed - def process_images(self, p, enabled, x_type, x_values, x_values_dropdown, y_type, y_values, y_values_dropdown, z_type, z_values, z_values_dropdown, csv_mode, draw_legend, no_fixed_seeds, include_grid, include_subgrids, include_images, margin_size): # pylint: disable=W0221, W0613 + + def process_images(self, p, *args): # pylint: disable=W0221, W0613 + if p.iteration > 0 and cache is not None and len(cache.images) > 0: + cache.images = [] # avoid returning same images multiple items return cache diff --git a/wiki b/wiki index 53def8203..2dba58a69 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 53def8203b6799cbd659327c1af6aa5af9cb9a70 +Subproject commit 2dba58a6962b70e92a077dcda8f178f5e811f175