diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 000000000..21e56ed85 --- /dev/null +++ b/.dockerignore @@ -0,0 +1,26 @@ +# defaults +.history +.vscode/ +/__pycache__ +/.ruff_cache +/cache +/cache.json +/config.json +/extensions/* +/html/extensions.json +/html/themes.json +/metadata.json +/node_modules +/outputs/* +/package-lock.json +/params.txt +/pnpm-lock.yaml +/styles.csv +/tmp +/ui-config.json +/user.css +/venv +/webui-user.bat +/webui-user.sh +/*.log.* +/*.log diff --git a/.gitignore b/.gitignore index dca4e17ad..6df029445 100644 --- a/.gitignore +++ b/.gitignore @@ -43,6 +43,7 @@ tunableop_results*.csv !webui.bat !webui.sh !package.json +!requirements.txt # pyinstaller *.spec @@ -74,4 +75,3 @@ dist/ !/models/VAE-approx/model.pt !/models/Reference !/models/Reference/**/* - diff --git a/.pylintrc b/.pylintrc index 45869a8c3..59f1cb127 100644 --- a/.pylintrc +++ b/.pylintrc @@ -8,29 +8,31 @@ fail-under=10 ignore=CVS ignore-paths=/usr/lib/.*$, modules/apg, + modules/consistory, modules/control/proc, modules/control/units, modules/ctrlx, - modules/dcsolver, modules/dml, modules/ggml, modules/hidiffusion, modules/hijack, + modules/instantir, modules/intel/ipex, modules/intel/openvino, modules/k-diffusion, modules/ldsr, + modules/meissonic, + modules/omnigen, modules/onnx_impl, modules/pag, modules/prompt_parser_xhinker.py, + modules/pulid/eva_clip, modules/rife, + modules/schedulers, modules/taesd, modules/todo, modules/unipc, - modules/vdm, modules/xadapter, - modules/meissonic, - modules/omnigen, repositories, extensions-builtin/sd-webui-agent-scheduler, extensions-builtin/sd-extension-chainner/nodes, @@ -130,7 +132,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 +177,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..c2d4a6f9a 100644 --- a/.ruff.toml +++ b/.ruff.toml @@ -4,28 +4,30 @@ exclude = [ ".ruff_cache", ".vscode", "modules/apg", + "modules/consistory", "modules/control/proc", "modules/control/units", - "modules/dcsolver", "modules/ggml", "modules/hidiffusion", "modules/hijack", + "modules/instantir", "modules/intel/ipex", "modules/intel/openvino", "modules/k-diffusion", "modules/ldsr", + "modules/meissonic", + "modules/omnigen", "modules/pag", "modules/postprocess/aurasr_arch.py", "modules/prompt_parser_xhinker.py", + "modules/pulid/eva_clip", "modules/rife", + "modules/schedulers", "modules/segmoe", "modules/taesd", "modules/todo", "modules/unipc", - "modules/vdm", "modules/xadapter", - "modules/meissonic", - "modules/omnigen", "repositories", "extensions-builtin/sd-extension-chainner/nodes", "extensions-builtin/sd-webui-agent-scheduler", diff --git a/CHANGELOG.md b/CHANGELOG.md index bb56efec5..795647a7c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,37 +1,143 @@ # Change Log for SD.Next -## Update for 2024-11-02 +## Update for 2024-11-19 -Smaller release just few days after the last one, but with some important fixes and improvements. +### Highlights for 2024-11-19 + +*What's New?* + +First, a massive update to docs including new UI top-level **info** tab with access to [changelog](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) and [wiki](https://github.com/vladmandic/automatic/wiki), many updates and new articles AND full **built-in documentation search** capabilities + +**New integrations**: +- [PuLID](https://github.com/ToTheBeginning/PuLID): Pure and Lightning ID Customization via Contrastive Alignment +- [InstantIR](https://github.com/instantX-research/InstantIR): Blind Image Restoration with Instant Generative Reference +- [ConsiStory](https://github.com/NVlabs/consistory): Consistent Image Generation +- [MiaoshouAI PromptGen v2.0](https://huggingface.co/MiaoshouAI/Florence-2-base-PromptGen-v2.0) VQA captioning + +**Workflow Improvements**: +- Native Docker support +- SD3x & Flux.1: more ControlNets, all-in-one-safetensors, DPM samplers, etc. +- XYZ grid: benchmarking, video creation, etc. +- Enhanced prompt parsing +- UI improvements +- Installer self-healing `venv` + +And quite a few more improvements and fixes since the last update - for full details see changelog... + +[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-11-19 + +- Docs: + - new top-level **info** tab with access to [changelog](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) and [wiki](https://github.com/vladmandic/automatic/wiki) + - UI 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 info -> changelog and search/highligh/navigate directly in UI! + - UI built-in [wiki](https://github.com/vladmandic/automatic/wiki) + go to info -> wiki and search wiki pages directly in UI! + - major [Wiki](https://github.com/vladmandic/automatic/wiki) and [Home](https://github.com/vladmandic/automatic) updates + - updated API swagger docs for at `/docs` +- Integrations: + - [PuLID](https://github.com/ToTheBeginning/PuLID): Pure and Lightning ID Customization via Contrastive Alignment + - advanced method of face id 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* for text-to-image, image-to-image, inpaint, refine, detailer workflows + - can be used in xyz grid + - *note*: this module contains several advanced features on top of original implementation + - [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 + - FlowMatch samplers: + - Applicable to SD 3.x and Flux.1 models + - Complete family: + +- Workflow improvements: + - Native Docker support with pre-defined [Dockerfile](https://github.com/vladmandic/automatic/blob/dev/Dockerfile) + - 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 + - Prompt parser: + - support for prompt scheduling + - renamed parser options: `native`, `xhinker`, `compel`, `a1111`, `fixed` + - parser options are available in xyz grid + - improved caching + - UI: + - better gallery and networks sidebar sizing + - 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 + - Image2image: + - integrated refine/upscale/hires workflow +- 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 + - Scripts: + - More verbose descriptions for all scripts + - 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: + - Auto pipeline switching coveres wrapper classes and nested pipelines + - Full settings validation on load of `config.json` + - Refactor of all params in main processing classes + +- Fixes: + - custom watermark add alphablending + - 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 + - ui css fixes + +## Update for 2024-11-01 + +Smaller release just 3 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! -- 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 -- 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 -- XYZ grid: - - optional per-image time benchmark info -- CLI: - - refactor command line params - run `webui.sh`/`webui.bat` with `--help` to see all options - Other: - Repo: move screenshots to GH pages - Update requirements - Fixes: - - custom watermark add alphablending - detailer min/max size as fractions of image size - ipadapter load on-demand - ipadapter face use correct yolo model @@ -40,10 +146,7 @@ This release can be considered an LTS release before we kick off the next round - fix diffusers load from folder - fix lora enum logging on windows - fix xyz grid with batch count - - fix vqa models ignoring hfcache folder setting - - fix network height in standard vs modern ui - - fix k-diff enum on startup - - move downloads of some auxillary models to hfcache instead of models folder + - move dowwloads of some auxillary models to hfcache instead of models folder ## Update for 2024-10-29 diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 000000000..5f38d6caa --- /dev/null +++ b/Dockerfile @@ -0,0 +1,57 @@ +# SD.Next Dockerfile +# docs: + +# base image +FROM pytorch/pytorch:2.5.1-cuda12.4-cudnn9-runtime + +# metadata +LABEL org.opencontainers.image.vendor="SD.Next" +LABEL org.opencontainers.image.authors="vladmandic" +LABEL org.opencontainers.image.url="https://github.com/vladmandic/automatic/" +LABEL org.opencontainers.image.documentation="https://github.com/vladmandic/automatic/wiki/Docker" +LABEL org.opencontainers.image.source="https://github.com/vladmandic/automatic/" +LABEL org.opencontainers.image.licenses="AGPL-3.0" +LABEL org.opencontainers.image.title="SD.Next" +LABEL org.opencontainers.image.description="SD.Next: Advanced Implementation of Stable Diffusion and other Diffusion-based generative image models" +LABEL org.opencontainers.image.base.name="https://hub.docker.com/pytorch/pytorch:2.5.1-cuda12.4-cudnn9-runtime" +LABEL org.opencontainers.image.version="latest" + +# minimum install +RUN ["apt-get", "-y", "update"] +RUN ["apt-get", "-y", "install", "git", "build-essential", "google-perftools", "curl"] +# optional if full cuda-dev is required by some downstream library +# RUN ["apt-get", "-y", "nvidia-cuda-toolkit"] +RUN ["/usr/sbin/ldconfig"] + +# copy sdnext +COPY . /app +WORKDIR /app + +# stop pip and uv from caching +ENV PIP_NO_CACHE_DIR=true +ENV PIP_ROOT_USER_ACTION=ignore +ENV UV_NO_CACHE=true +# disable model hashing for faster startup +ENV SD_NOHASHING=true +# set data directories +ENV SD_DATADIR="/mnt/data" +ENV SD_MODELSDIR="/mnt/models" +ENV SD_DOCKER=true + +# tcmalloc is not required but it is highly recommended +ENV LD_PRELOAD=libtcmalloc.so.4 +# sdnext will run all necessary pip install ops and then exit +RUN ["python", "/app/launch.py", "--debug", "--uv", "--use-cuda", "--log", "sdnext.log", "--test", "--optional"] +# preinstall additional packages to avoid installation during runtime + +# actually run sdnext +CMD ["python", "launch.py", "--debug", "--skip-all", "--listen", "--quick", "--api-log", "--log", "sdnext.log"] + +# expose port +EXPOSE 7860 + +# healthcheck function +# HEALTHCHECK --interval=60s --timeout=10s --start-period=60s --retries=3 CMD curl --fail http://localhost:7860/sdapi/v1/status || exit 1 + +# stop signal +STOPSIGNAL SIGINT diff --git a/README.md b/README.md index a2caa5cb7..d099496b8 100644 --- a/README.md +++ b/README.md @@ -1,12 +1,12 @@
-SD.Next +SD.Next -**Stable Diffusion implementation with advanced features** +**Image Diffusion implementation with advanced features** -[![Sponsors](https://img.shields.io/static/v1?label=Sponsor&message=%E2%9D%A4&logo=GitHub&color=%23fe8e86)](https://github.com/sponsors/vladmandic) -![Last Commit](https://img.shields.io/github/last-commit/vladmandic/automatic?svg=true) +![Last update](https://img.shields.io/github/last-commit/vladmandic/automatic?svg=true) ![License](https://img.shields.io/github/license/vladmandic/automatic?svg=true) [![Discord](https://img.shields.io/discord/1101998836328697867?logo=Discord&svg=true)](https://discord.gg/VjvR2tabEX) +[![Sponsors](https://img.shields.io/static/v1?label=Sponsor&message=%E2%9D%A4&logo=GitHub&color=%23fe8e86)](https://github.com/sponsors/vladmandic) [Wiki](https://github.com/vladmandic/automatic/wiki) | [Discord](https://discord.gg/VjvR2tabEX) | [Changelog](CHANGELOG.md) @@ -18,45 +18,36 @@ - [SD.Next Features](#sdnext-features) - [Model support](#model-support) - [Platform support](#platform-support) -- [Backend support](#backend-support) -- [Examples](#examples) -- [Install](#install) -- [Notes](#notes) +- [Getting started](#getting-started) ## SD.Next Features All individual features are not listed here, instead check [ChangeLog](CHANGELOG.md) for full list of changes -- Multiple backends! - ▹ **Diffusers | Original** - Multiple UIs! ▹ **Standard | Modern** - Multiple diffusion models! - ▹ **Stable Diffusion 1.5/2.1/XL/3.0/3.5 | LCM | Lightning | Segmind | Kandinsky | Pixart-α | Pixart-Σ | Stable Cascade | FLUX.1 | AuraFlow | Würstchen | Alpha Lumina | Kwai Kolors | aMUSEd | DeepFloyd IF | UniDiffusion | SD-Distilled | BLiP Diffusion | KOALA | SDXS | Hyper-SD | HunyuanDiT | CogView | OmniGen | Meissonic | etc.** - Built-in Control for Text, Image, Batch and video processing! - ▹ **ControlNet | ControlNet XS | Control LLLite | T2I Adapters | IP Adapters** - Multiplatform! - ▹ **Windows | Linux | MacOS with CPU | nVidia | AMD | IntelArc/IPEX | DirectML | OpenVINO | ONNX+Olive | ZLUDA** -- Platform specific autodetection and tuning performed on install + ▹ **Windows | Linux | MacOS | nVidia | AMD | IntelArc/IPEX | DirectML | OpenVINO | ONNX+Olive | ZLUDA** +- Multiple backends! + ▹ **Diffusers | Original** +- Platform specific autodetection and tuning performed on install - Optimized processing with latest `torch` developments with built-in support for `torch.compile` and multiple compile backends: *Triton, ZLUDA, StableFast, DeepCache, OpenVINO, NNCF, IPEX, OneDiff* - Improved prompt parser -- Enhanced *Lora*/*LoCon*/*Lyco* code supporting latest trends in training - Built-in queue management - Enterprise level logging and hardened API - Built in installer with automatic updates and dependency management -- Modernized UI with theme support and number of built-in themes *(dark and light)* -- Mobile compatible +- Mobile compatible
*Main interface using **StandardUI***: -![screenshot-text2image](https://github.com/user-attachments/assets/87ac2813-65c2-45f4-80b8-67b26ccf5cd6) +![screenshot-standardui](https://github.com/user-attachments/assets/cab47fe3-9adb-4d67-aea9-9ee738df5dcc) *Main interface using **ModernUI***: -![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) +![screenshot-modernui](https://github.com/user-attachments/assets/39e3bc9a-a9f7-4cda-ba33-7da8def08032) For screenshots and informations on other available themes, see [Themes Wiki](https://github.com/vladmandic/automatic/wiki/Themes) @@ -65,12 +56,10 @@ 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 -See [models overview](https://github.com/vladmandic/automatic/wiki/Models) for details on each model, including their architecture, complexity and other info +See [models overview](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](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 Diffusion XL](https://github.com/Stability-AI/generative-models), [StabilityAI Stable Diffusion 3.0](https://stability.ai/news/stable-diffusion-3-medium) Medium, [StabilityAI Stable Diffusion 3.5](https://huggingface.co/stabilityai/stable-diffusion-3.5-large) Medium, Large, 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 @@ -84,13 +73,9 @@ See [models overview](https://github.com/vladmandic/automatic/wiki/Models) for d - [CogView 3+](https://huggingface.co/THUDM/CogView3-Plus-3B) - [LCM: Latent Consistency Models](https://github.com/openai/consistency_models) - [aMUSEd](https://huggingface.co/amused/amused-256) 256 and 512 -- [Segmind Vega](https://huggingface.co/segmind/Segmind-Vega) -- [Segmind SSD-1B](https://huggingface.co/segmind/SSD-1B) -- [Segmind SegMoE](https://github.com/segmind/segmoe) *SD and SD-XL* -- [Segmind SD Distilled](https://huggingface.co/blog/sd_distillation) *(all variants)* +- [Segmind Vega](https://huggingface.co/segmind/Segmind-Vega), [Segmind SSD-1B](https://huggingface.co/segmind/SSD-1B), [Segmind SegMoE](https://github.com/segmind/segmoe) *SD and SD-XL*, [Segmind SD Distilled](https://huggingface.co/blog/sd_distillation) *(all variants)* - [Kandinsky](https://github.com/ai-forever/Kandinsky-2) *2.1 and 2.2 and latest 3.0* -- [PixArt-α XL 2](https://github.com/PixArt-alpha/PixArt-alpha) *Medium and Large* -- [PixArt-Σ](https://github.com/PixArt-alpha/PixArt-sigma) +- [PixArt-α XL 2](https://github.com/PixArt-alpha/PixArt-alpha) *Medium and Large*, [PixArt-Σ](https://github.com/PixArt-alpha/PixArt-sigma) - [Warp Wuerstchen](https://huggingface.co/blog/wuertschen) - [Tsinghua UniDiffusion](https://github.com/thu-ml/unidiffuser) - [DeepFloyd IF](https://github.com/deep-floyd/IF) *Medium and Large* @@ -101,15 +86,6 @@ See [models overview](https://github.com/vladmandic/automatic/wiki/Models) for d - [SDXS](https://github.com/IDKiro/sdxs) - [Hyper-SD](https://huggingface.co/ByteDance/Hyper-SD) - -Also supported are modifiers such as: -- **LCM**, **Turbo** and **Lightning** (*adversarial diffusion distillation*) networks -- All **LoRA** types such as LoCon, LyCORIS, HADA, IA3, Lokr, OFT -- **IP-Adapters** for SD 1.5 and SD-XL -- **InstantID**, **FaceSwap**, **FaceID**, **PhotoMerge** -- **AnimateDiff** for SD 1.5 -- **MuLAN** multi-language support - ## Platform support - *nVidia* GPUs using **CUDA** libraries on both *Windows and Linux* @@ -121,6 +97,25 @@ Also supported are modifiers such as: - Any GPU or device compatible with **OpenVINO** libraries on both *Windows and Linux* - *Apple M1/M2* on *OSX* using built-in support in Torch with **MPS** optimizations - *ONNX/Olive* +- *AMD* GPUs on Windows using **ZLUDA** libraries + +## Getting started + +- Get started with **SD.Next** by following the [installation instructions](wiki/Installation) +- For more details, check out [advanced installation](wiki/Advanced-Install) guide +- List and explanation of [command line arguments](wiki/CLI-Arguments) +- Install walkthrough [video](https://www.youtube.com/watch?v=nWTnTyFTuAs) + +> [!TIP] +> And for platform specific information, check out +> [WSL](wiki/WSL) | [Intel Arc](wiki/Intel-ARC) | [DirectML](wiki/DirectML) | [OpenVINO](wiki/OpenVINO) | [ONNX & Olive](wiki/ONNX-Runtime) | [ZLUDA](wiki/ZLUDA) | [AMD ROCm](wiki/AMD-ROCm) | [MacOS](wiki/MacOS-Python.md) | [nVidia](wiki/nVidia) + +> [!WARNING] +> If you run into issues, check out [troubleshooting](wiki/Troubleshooting) and [debugging](wiki/Debug) guides + +> [!TIP] +> All command line options can also be set via env variable +> For example `--debug` is same as `set SD_DEBUG=true` ## Backend support @@ -129,91 +124,11 @@ Also supported are modifiers such as: - **Diffusers**: Based on new [Huggingface Diffusers](https://huggingface.co/docs/diffusers/index) implementation Supports *all* models listed below This backend is set as default for new installations - See [wiki article](https://github.com/vladmandic/automatic/wiki/Diffusers) for more information - **Original**: Based on [LDM](https://github.com/Stability-AI/stablediffusion) reference implementation and significantly expanded on by [A1111](https://github.com/AUTOMATIC1111/stable-diffusion-webui) This backend and is fully compatible with most existing functionality and extensions written for *A1111 SDWebUI* Supports **SD 1.x** and **SD 2.x** models All other model types such as *SD-XL, LCM, Stable Cascade, PixArt, Playground, Segmind, Kandinsky, etc.* require backend **Diffusers** -## Examples - -*IP Adapters*: -![screenshot-ipadapter](https://github.com/user-attachments/assets/92830894-845c-49ec-92d9-18c8a577d04f) - -*Color grading*: -![screenshot-control](https://github.com/user-attachments/assets/cdad2722-ae7c-4c9c-94d6-5ea35a4b1356) - -*InstantID*: -![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** -> - Loading any other models using **Original** backend is not supported -> - Loading manually download model `.safetensors` files is supported for specified models only (typically SD 1.x / SD 2.x / SD-XL models only) -> - For all other model types, use backend **Diffusers** and use built in Model downloader or - select model from Networks -> Models -> Reference list in which case it will be auto-downloaded and loaded - -## Install - -- [Step-by-step install guide](https://github.com/vladmandic/automatic/wiki/Installation) -- [Advanced install notes](https://github.com/vladmandic/automatic/wiki/Advanced-Install) -- [Video: install and use](https://www.youtube.com/watch?v=nWTnTyFTuAs) -- [Common installation errors](https://github.com/vladmandic/automatic/discussions/1627) -- [FAQ](https://github.com/vladmandic/automatic/discussions/1011) - -> [!TIP] -> - If you can't run SD.Next locally, try cloud deployment using [RunDiffusion](https://rundiffusion.com?utm_source=github&utm_medium=referral&utm_campaign=SDNext)! -> - Server can run with or without virtual environment, - Recommended to use `VENV` to avoid library version conflicts with other applications -> - **nVidia/CUDA** / **AMD/ROCm** / **Intel/OneAPI** are auto-detected if present and available, - For any other use case such as **DirectML**, **ONNX/Olive**, **OpenVINO** specify required parameter explicitly - or wrong packages may be installed as installer will assume CPU-only environment -> - Full startup sequence is logged in `sdnext.log`, - so if you encounter any issues, please check it first - -### Run - -Once SD.Next is installed, simply run `webui.ps1` or `webui.bat` (*Windows*) or `webui.sh` (*Linux or MacOS*) - -For list of available command line options, run `webui --help` for the full & up-to-date list - -> [!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](https://github.com/user-attachments/assets/cdad2722-ae7c-4c9c-94d6-5ea35a4b1356) - -*Control processors*: -![screenshot-processors](https://github.com/user-attachments/assets/7bccb82b-366e-4bdb-ae57-cc53fac95d3c) - -*Masking*: -![screenshot-mask](https://github.com/user-attachments/assets/4b057e65-64f0-44ea-93b4-c3b69bc55532) - -### Extensions - -SD.Next comes with several extensions pre-installed: - -- [System Info](https://github.com/vladmandic/sd-extension-system-info) -- [chaiNNer](https://github.com/vladmandic/sd-extension-chainner) -- [RemBg](https://github.com/vladmandic/sd-extension-rembg) -- [Agent Scheduler](https://github.com/ArtVentureX/sd-webui-agent-scheduler) -- [Modern UI](https://github.com/BinaryQuantumSoul/sdnext-modernui) - ### Collab - We'd love to have additional maintainers (with comes with full repo rights). If you're interested, ping us! @@ -242,12 +157,6 @@ This should be fully cross-platform, but we'd really love to have additional con If you're unsure how to use a feature, best place to start is [Wiki](https://github.com/vladmandic/automatic/wiki) and if its not there, check [ChangeLog](CHANGELOG.md) for when feature was first introduced as it will always have a short note on how to use it -- [Wiki](https://github.com/vladmandic/automatic/wiki) -- [ReadMe](README.md) -- [ToDo](TODO.md) -- [ChangeLog](CHANGELOG.md) -- [CLI Tools](cli/README.md) - ### Sponsors
diff --git a/TODO.md b/TODO.md index 2f3f90852..d5cc19cf7 100644 --- a/TODO.md +++ b/TODO.md @@ -4,11 +4,9 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma ## Future Candidates -- async lowvram: -- fp8: +- sd35 ip-adapter +- flux.1 ip-adapter +- flow-match scheudlers: - ipadapter-negative: - include reference styles - -### Missing - - control api scripts compatibility diff --git a/cli/README.md b/cli/README.md deleted file mode 100644 index 70de255b1..000000000 --- a/cli/README.md +++ /dev/null @@ -1,108 +0,0 @@ -# Stable-Diffusion Productivity Scripts - -Note: All scripts have built-in `--help` parameter that can be used to get more information - -
- -## Main Scripts - -### Generate - -Text-to-image with all of the possible parameters -Supports upsampling, face restoration and grid creation -> python generate.py - -By default uses parameters from `generate.json` - -Parameters that are not specified will be randomized: - -- Prompt will be dynamically created from template of random samples: `random.json` -- 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 - -### Benchmark - -> python run-benchmark.py - -### Create Previews - -Create previews for **embeddings**, **lora**, **lycoris**, **dreambooth** and **hypernetwork** - -> python create-previews.py - -## Image Grid - -> python image-grid.py - -### Image Watermark - -Create invisible image watermark and remove existing EXIF tags - -> python image-watermark.py - -### Image Interrogate - -Runs CLiP and Booru image interrogation - -> python image-interrogate.py - -### Palette Extract - -Extract color palette from image(s) - -> python image-palette.py - -### Prompt Ideas - -Generate complex prompt ideas - -> python prompt-ideas.py - -### Prompt Promptist - -Attempts to beautify the provided prompt - -> python prompt-promptist.py - -### Video Extract - -Extract frames from video files - -> python video-extract.py - -
- -## Utility Scripts - -### SDAPI - -Utility module that handles async communication to Automatic API endpoints -Note: Requires SD API - -Can be used to manually execute specific commands: -> python sdapi.py progress -> python sdapi.py interrupt -> python sdapi.py shutdown diff --git a/cli/api-faces.py b/cli/api-detect.py similarity index 80% rename from cli/api-faces.py rename to cli/api-detect.py index 0a98843c0..ca121b220 100755 --- a/cli/api-faces.py +++ b/cli/api-detect.py @@ -44,16 +44,15 @@ def encode(f): def detect(args): # pylint: disable=redefined-outer-name - data = post('/sdapi/v1/faces', { 'image': encode(args.image) }) - for face in zip(data['images'], data['scores']): - log.info(f'Face: score={face[1]}') - image = Image.open(io.BytesIO(base64.b64decode(face[0]))) - image.save(f'/tmp/face_{face[1]}.jpg') + data = post('/sdapi/v1/detect', { 'image': encode(args.image), 'model': args.model }) + for i in range(len(data['images'])): + log.info(f"Item {i}: score={data['scores'][i]} cls={data['classes'][i]} box={data['boxes'][i]} label={data['labels'][i]}") if __name__ == "__main__": parser = argparse.ArgumentParser(description = 'api-faces') parser.add_argument('--image', required=True, help='input image') + parser.add_argument('--model', required=False, default='', help='model') args = parser.parse_args() - log.info(f'api-faces: {args}') + log.info(f'api-detect: {args}') detect(args) diff --git a/cli/api-faceid.py b/cli/api-faceid.py index e656a4a47..18f1d3503 100755 --- a/cli/api-faceid.py +++ b/cli/api-faceid.py @@ -63,7 +63,7 @@ def generate(args): # pylint: disable=redefined-outer-name options['height'] = args.height options['face'] = { 'mode': 'FaceID', - 'ip_model': 'FaceID Base', + 'ip_model': 'FaceID XL', 'source_images': [encode(args.face)], } data = post('/sdapi/v1/txt2img', options) @@ -86,7 +86,7 @@ if __name__ == "__main__": parser = argparse.ArgumentParser(description = 'api-faceid') parser.add_argument('--width', required=False, default=512, help='image width') parser.add_argument('--height', required=False, default=512, help='image height') - parser.add_argument('--face', required=False, help='face image') + parser.add_argument('--face', required=True, help='face image') parser.add_argument('--prompt', required=False, default='', help='prompt text') parser.add_argument('--negative', required=False, default='', help='negative prompt text') parser.add_argument('--steps', required=False, default=20, help='number of steps') @@ -97,20 +97,3 @@ if __name__ == "__main__": args = parser.parse_args() log.info(f'api-faceid: {args}') generate(args) - -""" -request.face.mode, -request.face.source_images, -request.face.ip_model, -request.face.ip_override_sampler, -request.face.ip_cache_model, -request.face.ip_strength, -request.face.ip_structure, -request.face.id_strength, -request.face.id_conditioning, -request.face.id_cache, -request.face.pm_trigger, -request.face.pm_strength, -request.face.pm_start, -request.face.fs_cache -""" diff --git a/cli/api-json.py b/cli/api-json.py index 79c0ebc3b..61e5ec3ce 100755 --- a/cli/api-json.py +++ b/cli/api-json.py @@ -45,7 +45,7 @@ if __name__ == "__main__": log.info(f'api-json: {args}') if os.path.isfile(args.json[0]): with open(args.json[0], 'r', encoding='ascii') as f: - dct = json.load(f) # TODO fails with b64 encoded images inside json due to string encoding + dct = json.load(f) else: dct = json.loads(args.json[0]) res = post(endpoint=args.endpoint[0], payload=dct) diff --git a/cli/api-progress.py b/cli/api-progress.py index 2c90fe95f..00ed618d2 100755 --- a/cli/api-progress.py +++ b/cli/api-progress.py @@ -1,5 +1,9 @@ #!/usr/bin/env python +""" +check progress of last job and shutdown system if timeout reached +""" + import os import time import datetime @@ -16,7 +20,7 @@ opts = Dot({ "timeout": 3600, "frequency": 60, "action": "sudo shutdown now", - "url": "https://127.0.0.1:7860", + "url": "http://127.0.0.1:7860", "user": "", "password": "", }) diff --git a/cli/api-txt2img.js b/cli/api-txt2img.js index 46d09b3a2..8d0e9f5d1 100755 --- a/cli/api-txt2img.js +++ b/cli/api-txt2img.js @@ -20,23 +20,6 @@ const sd_options = { cfg_scale: 6, width: 512, height: 512, - /* - // enable second pass - enable_hr: true, - // second pass: upscale - hr_upscaler: 'SCUNet GAN', - hr_scale: 2.0, - // second pass: hires - hr_force: true, - hr_second_pass_steps: 20, - hr_sampler_name: 'UniPC', - denoising_strength: 0.5, - // second pass: refiner - refiner_steps: 5, - refiner_start: 0.8, - refiner_prompt: '', - refiner_negative: '', - */ // api return options save_images: false, send_images: true, @@ -55,7 +38,7 @@ async function main() { const json = await res.json(); console.log('result:', json.info); for (const i in json.images) { // eslint-disable-line guard-for-in - const f = `/tmp/test-{${i}.jpg`; + const f = `/tmp/test-${i}.jpg`; fs.writeFileSync(f, atob(json.images[i]), 'binary'); console.log('image saved:', f); } diff --git a/cli/api-txt2img.py b/cli/api-txt2img.py index 89d84be80..868b13eee 100755 --- a/cli/api-txt2img.py +++ b/cli/api-txt2img.py @@ -48,7 +48,7 @@ def generate(args): # pylint: disable=redefined-outer-name options['sampler_name'] = args.sampler options['width'] = int(args.width) options['height'] = int(args.height) - if args.faces: + if args.detailer: options['detailer'] = args.detailer options['denoising_strength'] = 0.5 options['hr_sampler_name'] = args.sampler diff --git a/cli/api-upscale.py b/cli/api-upscale.py index 7f188650f..488f2db45 100755 --- a/cli/api-upscale.py +++ b/cli/api-upscale.py @@ -73,7 +73,8 @@ def upscale(args): # pylint: disable=redefined-outer-name if 'image' in data: b64 = data['image'].split(',',1)[0] image = Image.open(io.BytesIO(base64.b64decode(b64))) - image.save(args.output) + if args.output: + image.save(args.output) log.info(f'received: image={image} file={args.output} time={t1-t0:.2f}') else: log.warning(f'no images received: {data}') @@ -82,7 +83,7 @@ def upscale(args): # pylint: disable=redefined-outer-name if __name__ == "__main__": parser = argparse.ArgumentParser(description = 'api-upscale') parser.add_argument('--input', required=True, help='input image') - parser.add_argument('--output', required=True, help='output image') + parser.add_argument('--output', required=False, help='output image') parser.add_argument('--upscaler', required=False, default='Nearest', help='upscaler name') parser.add_argument('--scale', required=False, default=2, help='upscaler scale') args = parser.parse_args() diff --git a/cli/full-test.sh b/cli/full-test.sh new file mode 100755 index 000000000..e410528ad --- /dev/null +++ b/cli/full-test.sh @@ -0,0 +1,29 @@ +#!/usr/bin/env bash + +source venv/bin/activate +echo image-exif +python cli/api-info.py --input html/logo-bg-0.jpg +echo txt2img +python cli/api-txt2img.py --detailer --prompt "girl on a mountain" --seed 42 --sampler DEIS --width 1280 --height 800 --steps 10 +echo img2img +python cli/api-img2img.py --init html/logo-bg-0.jpg --steps 10 +echo inpaint +python cli/api-img2img.py --init html/logo-bg-0.jpg --mask html/logo-dark.png --steps 10 +echo upscale +python cli/api-upscale.py --input html/logo-bg-0.jpg --upscaler "ESRGAN 4x Valar" --scale 4 +echo vqa +python cli/api-vqa.py --input html/logo-bg-0.jpg +echo detailer +python cli/api-detect.py --image html/invoked.jpg +echo faceid +python cli/api-faceid.py --face html/simple-dark.jpg +echo control-txt2img +python cli/api-control.py --prompt "cute robot" +echo control-img2img +python cli/api-control.py --prompt "cute robot" --input html/logo-bg-0.jpg +echo control-ipsadapter +python cli/api-control.py --prompt "cute robot" --ipadapter "Base SDXL:html/logo-bg-0.jpg:0.8" +echo control-preprocess +python cli/api-preprocess.py --input html/logo-bg-0.jpg --model "Zoe Depth" +echo control-controlnet +python cli/api-control.py --prompt "cute robot" --input html/logo-bg-0.jpg --type controlnet --control "Zoe Depth:Xinsir Union XL:0.5" 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/extra_networks_lora.py b/extensions-builtin/Lora/extra_networks_lora.py index 9172d7336..307d8cc13 100644 --- a/extensions-builtin/Lora/extra_networks_lora.py +++ b/extensions-builtin/Lora/extra_networks_lora.py @@ -26,7 +26,9 @@ def get_stepwise(param, step, steps): return v else: return m - return calculate_weight(sorted_positions(param), step, steps) + + stepwise = calculate_weight(sorted_positions(param), step, steps) + return stepwise class ExtraNetworkLora(extra_networks.ExtraNetwork): @@ -145,7 +147,6 @@ class ExtraNetworkLora(extra_networks.ExtraNetwork): if self.active and networks.debug: shared.log.debug(f"Network end: type=LoRA load={networks.timer['load']:.2f} apply={networks.timer['apply']:.2f} restore={networks.timer['restore']:.2f}") if self.errors: - p.comment("Networks with errors: " + ", ".join(f"{k} ({v})" for k, v in self.errors.items())) for k, v in self.errors.items(): shared.log.error(f'LoRA: name="{k}" errors={v}') self.errors.clear() diff --git a/extensions-builtin/Lora/networks.py b/extensions-builtin/Lora/networks.py index 160487e88..db617ee5b 100644 --- a/extensions-builtin/Lora/networks.py +++ b/extensions-builtin/Lora/networks.py @@ -50,44 +50,45 @@ convert_diffusers_name_to_compvis = lora_convert.convert_diffusers_name_to_compv def assign_network_names_to_compvis_modules(sd_model): if sd_model is None: return + sd_model = getattr(shared.sd_model, "pipe", shared.sd_model) # wrapped model compatiblility network_layer_mapping = {} if shared.native: - if hasattr(shared.sd_model, 'text_encoder') and shared.sd_model.text_encoder is not None: - for name, module in shared.sd_model.text_encoder.named_modules(): - prefix = "lora_te1_" if hasattr(shared.sd_model, 'text_encoder_2') else "lora_te_" + if hasattr(sd_model, 'text_encoder') and sd_model.text_encoder is not None: + for name, module in sd_model.text_encoder.named_modules(): + prefix = "lora_te1_" if hasattr(sd_model, 'text_encoder_2') else "lora_te_" network_name = prefix + name.replace(".", "_") network_layer_mapping[network_name] = module module.network_layer_name = network_name - if hasattr(shared.sd_model, 'text_encoder_2'): - for name, module in shared.sd_model.text_encoder_2.named_modules(): + if hasattr(sd_model, 'text_encoder_2'): + for name, module in sd_model.text_encoder_2.named_modules(): network_name = "lora_te2_" + name.replace(".", "_") network_layer_mapping[network_name] = module module.network_layer_name = network_name - if hasattr(shared.sd_model, 'unet'): - for name, module in shared.sd_model.unet.named_modules(): + if hasattr(sd_model, 'unet'): + for name, module in sd_model.unet.named_modules(): network_name = "lora_unet_" + name.replace(".", "_") network_layer_mapping[network_name] = module module.network_layer_name = network_name - if hasattr(shared.sd_model, 'transformer'): - for name, module in shared.sd_model.transformer.named_modules(): + if hasattr(sd_model, 'transformer'): + for name, module in sd_model.transformer.named_modules(): network_name = "lora_transformer_" + name.replace(".", "_") network_layer_mapping[network_name] = module if "norm" in network_name and "linear" not in network_name: continue module.network_layer_name = network_name else: - if not hasattr(shared.sd_model, 'cond_stage_model'): + if not hasattr(sd_model, 'cond_stage_model'): sd_model.network_layer_mapping = {} return - for name, module in shared.sd_model.cond_stage_model.wrapped.named_modules(): + for name, module in sd_model.cond_stage_model.wrapped.named_modules(): network_name = name.replace(".", "_") network_layer_mapping[network_name] = module module.network_layer_name = network_name - for name, module in shared.sd_model.model.named_modules(): + for name, module in sd_model.model.named_modules(): network_name = name.replace(".", "_") network_layer_mapping[network_name] = module module.network_layer_name = network_name - sd_model.network_layer_mapping = network_layer_mapping + shared.sd_model.network_layer_mapping = network_layer_mapping def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_default_multiplier) -> network.Network: @@ -226,6 +227,7 @@ def load_networks(names, te_multipliers=None, unet_multipliers=None, dyn_dims=No loaded_networks.clear() diffuser_loaded.clear() diffuser_scales.clear() + for i, (network_on_disk, name) in enumerate(zip(networks_on_disk, names)): net = None if network_on_disk is not None: @@ -261,14 +263,22 @@ def load_networks(names, te_multipliers=None, unet_multipliers=None, dyn_dims=No while len(lora_cache) > shared.opts.lora_in_memory_limit: name = next(iter(lora_cache)) lora_cache.pop(name, None) + if len(diffuser_loaded) > 0: - shared.log.debug(f'Load network: type=LoRA loaded={diffuser_loaded} scales={diffuser_scales}') - shared.sd_model.set_adapters(adapter_names=diffuser_loaded, adapter_weights=diffuser_scales) - if shared.opts.lora_fuse_diffusers: - shared.sd_model.fuse_lora(adapter_names=diffuser_loaded, lora_scale=1.0, fuse_unet=True, fuse_text_encoder=True) # fuse uses fixed scale since later apply does the scaling - shared.sd_model.unload_lora_weights() + shared.log.debug(f'Load network: type=LoRA loaded={diffuser_loaded} available={shared.sd_model.get_list_adapters()} active={shared.sd_model.get_active_adapters()} scales={diffuser_scales}') + try: + shared.sd_model.set_adapters(adapter_names=diffuser_loaded, adapter_weights=diffuser_scales) + if shared.opts.lora_fuse_diffusers: + shared.sd_model.fuse_lora(adapter_names=diffuser_loaded, lora_scale=1.0, fuse_unet=True, fuse_text_encoder=True) # fuse uses fixed scale since later apply does the scaling + shared.sd_model.unload_lora_weights() + except Exception as e: + shared.log.error(f'Load network: type=LoRA {e}') + if debug: + errors.display(e, 'LoRA') + if len(loaded_networks) > 0 and debug: shared.log.debug(f'Load network: type=LoRA loaded={len(loaded_networks)} cache={list(lora_cache)}') + devices.torch_gc() if recompile_model: 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 71bdbbd9c..4647bd7f8 160000 --- a/extensions-builtin/sdnext-modernui +++ b/extensions-builtin/sdnext-modernui @@ -1 +1 @@ -Subproject commit 71bdbbd9c0a55ccea38cbf6fb01483323ac93676 +Subproject commit 4647bd7f86be9d2783a9ba1f38acaa9bcec942d2 diff --git a/html/swagger.css b/html/swagger.css new file mode 100644 index 000000000..3c24a39ab --- /dev/null +++ b/html/swagger.css @@ -0,0 +1,868 @@ +.opblock { + border-width: 0 !important; +} +.opblock-summary-operation-id { + display: none !important; +} +.swagger-ui .models .json-schema-2020-12:not(.json-schema-2020-12--embedded)>.json-schema-2020-12-head .json-schema-2020-12__title:first-of-type { + font-size: 14px; + color: white; +} + +.swagger-ui .json-schema-2020-12-keyword__name--primary { + color: aqua; +} + +.swagger-ui .json-schema-2020-12-property .json-schema-2020-12__title { + color: aqua; +} + +@media only screen and (prefers-color-scheme: dark) { + + a { color: #8c8cfa; } + + ::-webkit-scrollbar-track-piece { background-color: rgba(255, 255, 255, .2) !important; } + + ::-webkit-scrollbar-track { background-color: rgba(255, 255, 255, .3) !important; } + + ::-webkit-scrollbar-thumb { background-color: rgba(255, 255, 255, .5) !important; } + + embed[type="application/pdf"] { filter: invert(90%); } + + html { + background: #1f1f1f !important; + box-sizing: border-box; + filter: contrast(100%) brightness(100%) saturate(100%); + overflow-y: scroll; + } + + body { + background: #1f1f1f; + background-color: #1f1f1f; + background-image: none !important; + } + + button, input, select, textarea { + background-color: #1f1f1f; + color: #bfbfbf; + } + + font, html { color: #bfbfbf; } + + .swagger-ui, .swagger-ui section h3 { color: #b5bac9; } + + .swagger-ui a { background-color: transparent; } + + .swagger-ui mark { + background-color: #664b00; + color: #bfbfbf; + } + + .swagger-ui legend { color: inherit; } + + .swagger-ui .debug * { outline: #e6da99 solid 1px; } + + .swagger-ui .debug-white * { outline: #fff solid 1px; } + + .swagger-ui .debug-black * { outline: #bfbfbf solid 1px; } + + .swagger-ui .debug-grid { background: url(data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAgAAAAICAYAAADED76LAAAAGXRFWHRTb2Z0d2FyZQBBZG9iZSBJbWFnZVJlYWR5ccllPAAAAyhpVFh0WE1MOmNvbS5hZG9iZS54bXAAAAAAADw/eHBhY2tldCBiZWdpbj0i77u/IiBpZD0iVzVNME1wQ2VoaUh6cmVTek5UY3prYzlkIj8+IDx4OnhtcG1ldGEgeG1sbnM6eD0iYWRvYmU6bnM6bWV0YS8iIHg6eG1wdGs9IkFkb2JlIFhNUCBDb3JlIDUuNi1jMTExIDc5LjE1ODMyNSwgMjAxNS8wOS8xMC0wMToxMDoyMCAgICAgICAgIj4gPHJkZjpSREYgeG1sbnM6cmRmPSJodHRwOi8vd3d3LnczLm9yZy8xOTk5LzAyLzIyLXJkZi1zeW50YXgtbnMjIj4gPHJkZjpEZXNjcmlwdGlvbiByZGY6YWJvdXQ9IiIgeG1sbnM6eG1wTU09Imh0dHA6Ly9ucy5hZG9iZS5jb20veGFwLzEuMC9tbS8iIHhtbG5zOnN0UmVmPSJodHRwOi8vbnMuYWRvYmUuY29tL3hhcC8xLjAvc1R5cGUvUmVzb3VyY2VSZWYjIiB4bWxuczp4bXA9Imh0dHA6Ly9ucy5hZG9iZS5jb20veGFwLzEuMC8iIHhtcE1NOkRvY3VtZW50SUQ9InhtcC5kaWQ6MTRDOTY4N0U2N0VFMTFFNjg2MzZDQjkwNkQ4MjgwMEIiIHhtcE1NOkluc3RhbmNlSUQ9InhtcC5paWQ6MTRDOTY4N0Q2N0VFMTFFNjg2MzZDQjkwNkQ4MjgwMEIiIHhtcDpDcmVhdG9yVG9vbD0iQWRvYmUgUGhvdG9zaG9wIENDIDIwMTUgKE1hY2ludG9zaCkiPiA8eG1wTU06RGVyaXZlZEZyb20gc3RSZWY6aW5zdGFuY2VJRD0ieG1wLmlpZDo3NjcyQkQ3NjY3QzUxMUU2QjJCQ0UyNDA4MTAwMjE3MSIgc3RSZWY6ZG9jdW1lbnRJRD0ieG1wLmRpZDo3NjcyQkQ3NzY3QzUxMUU2QjJCQ0UyNDA4MTAwMjE3MSIvPiA8L3JkZjpEZXNjcmlwdGlvbj4gPC9yZGY6UkRGPiA8L3g6eG1wbWV0YT4gPD94cGFja2V0IGVuZD0iciI/PsBS+GMAAAAjSURBVHjaYvz//z8DLsD4gcGXiYEAGBIKGBne//fFpwAgwAB98AaF2pjlUQAAAABJRU5ErkJggg==) 0 0; } + + .swagger-ui .debug-grid-16 { background: url(data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAABAAAAAQCAYAAAAf8/9hAAAAGXRFWHRTb2Z0d2FyZQBBZG9iZSBJbWFnZVJlYWR5ccllPAAAAyhpVFh0WE1MOmNvbS5hZG9iZS54bXAAAAAAADw/eHBhY2tldCBiZWdpbj0i77u/IiBpZD0iVzVNME1wQ2VoaUh6cmVTek5UY3prYzlkIj8+IDx4OnhtcG1ldGEgeG1sbnM6eD0iYWRvYmU6bnM6bWV0YS8iIHg6eG1wdGs9IkFkb2JlIFhNUCBDb3JlIDUuNi1jMTExIDc5LjE1ODMyNSwgMjAxNS8wOS8xMC0wMToxMDoyMCAgICAgICAgIj4gPHJkZjpSREYgeG1sbnM6cmRmPSJodHRwOi8vd3d3LnczLm9yZy8xOTk5LzAyLzIyLXJkZi1zeW50YXgtbnMjIj4gPHJkZjpEZXNjcmlwdGlvbiByZGY6YWJvdXQ9IiIgeG1sbnM6eG1wTU09Imh0dHA6Ly9ucy5hZG9iZS5jb20veGFwLzEuMC9tbS8iIHhtbG5zOnN0UmVmPSJodHRwOi8vbnMuYWRvYmUuY29tL3hhcC8xLjAvc1R5cGUvUmVzb3VyY2VSZWYjIiB4bWxuczp4bXA9Imh0dHA6Ly9ucy5hZG9iZS5jb20veGFwLzEuMC8iIHhtcE1NOkRvY3VtZW50SUQ9InhtcC5kaWQ6ODYyRjhERDU2N0YyMTFFNjg2MzZDQjkwNkQ4MjgwMEIiIHhtcE1NOkluc3RhbmNlSUQ9InhtcC5paWQ6ODYyRjhERDQ2N0YyMTFFNjg2MzZDQjkwNkQ4MjgwMEIiIHhtcDpDcmVhdG9yVG9vbD0iQWRvYmUgUGhvdG9zaG9wIENDIDIwMTUgKE1hY2ludG9zaCkiPiA8eG1wTU06RGVyaXZlZEZyb20gc3RSZWY6aW5zdGFuY2VJRD0ieG1wLmlpZDo3NjcyQkQ3QTY3QzUxMUU2QjJCQ0UyNDA4MTAwMjE3MSIgc3RSZWY6ZG9jdW1lbnRJRD0ieG1wLmRpZDo3NjcyQkQ3QjY3QzUxMUU2QjJCQ0UyNDA4MTAwMjE3MSIvPiA8L3JkZjpEZXNjcmlwdGlvbj4gPC9yZGY6UkRGPiA8L3g6eG1wbWV0YT4gPD94cGFja2V0IGVuZD0iciI/PvCS01IAAABMSURBVHjaYmR4/5+BFPBfAMFm/MBgx8RAGWCn1AAmSg34Q6kBDKMGMDCwICeMIemF/5QawEipAWwUhwEjMDvbAWlWkvVBwu8vQIABAEwBCph8U6c0AAAAAElFTkSuQmCC) 0 0; } + + .swagger-ui .debug-grid-8-solid { background: url(data:image/jpeg;base64,/9j/4QAYRXhpZgAASUkqAAgAAAAAAAAAAAAAAP/sABFEdWNreQABAAQAAAAAAAD/4QMxaHR0cDovL25zLmFkb2JlLmNvbS94YXAvMS4wLwA8P3hwYWNrZXQgYmVnaW49Iu+7vyIgaWQ9Ilc1TTBNcENlaGlIenJlU3pOVGN6a2M5ZCI/PiA8eDp4bXBtZXRhIHhtbG5zOng9ImFkb2JlOm5zOm1ldGEvIiB4OnhtcHRrPSJBZG9iZSBYTVAgQ29yZSA1LjYtYzExMSA3OS4xNTgzMjUsIDIwMTUvMDkvMTAtMDE6MTA6MjAgICAgICAgICI+IDxyZGY6UkRGIHhtbG5zOnJkZj0iaHR0cDovL3d3dy53My5vcmcvMTk5OS8wMi8yMi1yZGYtc3ludGF4LW5zIyI+IDxyZGY6RGVzY3JpcHRpb24gcmRmOmFib3V0PSIiIHhtbG5zOnhtcD0iaHR0cDovL25zLmFkb2JlLmNvbS94YXAvMS4wLyIgeG1sbnM6eG1wTU09Imh0dHA6Ly9ucy5hZG9iZS5jb20veGFwLzEuMC9tbS8iIHhtbG5zOnN0UmVmPSJodHRwOi8vbnMuYWRvYmUuY29tL3hhcC8xLjAvc1R5cGUvUmVzb3VyY2VSZWYjIiB4bXA6Q3JlYXRvclRvb2w9IkFkb2JlIFBob3Rvc2hvcCBDQyAyMDE1IChNYWNpbnRvc2gpIiB4bXBNTTpJbnN0YW5jZUlEPSJ4bXAuaWlkOkIxMjI0OTczNjdCMzExRTZCMkJDRTI0MDgxMDAyMTcxIiB4bXBNTTpEb2N1bWVudElEPSJ4bXAuZGlkOkIxMjI0OTc0NjdCMzExRTZCMkJDRTI0MDgxMDAyMTcxIj4gPHhtcE1NOkRlcml2ZWRGcm9tIHN0UmVmOmluc3RhbmNlSUQ9InhtcC5paWQ6QjEyMjQ5NzE2N0IzMTFFNkIyQkNFMjQwODEwMDIxNzEiIHN0UmVmOmRvY3VtZW50SUQ9InhtcC5kaWQ6QjEyMjQ5NzI2N0IzMTFFNkIyQkNFMjQwODEwMDIxNzEiLz4gPC9yZGY6RGVzY3JpcHRpb24+IDwvcmRmOlJERj4gPC94OnhtcG1ldGE+IDw/eHBhY2tldCBlbmQ9InIiPz7/7gAOQWRvYmUAZMAAAAAB/9sAhAAbGhopHSlBJiZBQi8vL0JHPz4+P0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHAR0pKTQmND8oKD9HPzU/R0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0dHR0f/wAARCAAIAAgDASIAAhEBAxEB/8QAWQABAQAAAAAAAAAAAAAAAAAAAAYBAQEAAAAAAAAAAAAAAAAAAAIEEAEBAAMBAAAAAAAAAAAAAAABADECA0ERAAEDBQAAAAAAAAAAAAAAAAARITFBUWESIv/aAAwDAQACEQMRAD8AoOnTV1QTD7JJshP3vSM3P//Z) 0 0 #1c1c21; } + + .swagger-ui .debug-grid-16-solid { background: url(data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAABAAAAAQCAIAAACQkWg2AAAAGXRFWHRTb2Z0d2FyZQBBZG9iZSBJbWFnZVJlYWR5ccllPAAAAyhpVFh0WE1MOmNvbS5hZG9iZS54bXAAAAAAADw/eHBhY2tldCBiZWdpbj0i77u/IiBpZD0iVzVNME1wQ2VoaUh6cmVTek5UY3prYzlkIj8+IDx4OnhtcG1ldGEgeG1sbnM6eD0iYWRvYmU6bnM6bWV0YS8iIHg6eG1wdGs9IkFkb2JlIFhNUCBDb3JlIDUuNi1jMTExIDc5LjE1ODMyNSwgMjAxNS8wOS8xMC0wMToxMDoyMCAgICAgICAgIj4gPHJkZjpSREYgeG1sbnM6cmRmPSJodHRwOi8vd3d3LnczLm9yZy8xOTk5LzAyLzIyLXJkZi1zeW50YXgtbnMjIj4gPHJkZjpEZXNjcmlwdGlvbiByZGY6YWJvdXQ9IiIgeG1sbnM6eG1wPSJodHRwOi8vbnMuYWRvYmUuY29tL3hhcC8xLjAvIiB4bWxuczp4bXBNTT0iaHR0cDovL25zLmFkb2JlLmNvbS94YXAvMS4wL21tLyIgeG1sbnM6c3RSZWY9Imh0dHA6Ly9ucy5hZG9iZS5jb20veGFwLzEuMC9zVHlwZS9SZXNvdXJjZVJlZiMiIHhtcDpDcmVhdG9yVG9vbD0iQWRvYmUgUGhvdG9zaG9wIENDIDIwMTUgKE1hY2ludG9zaCkiIHhtcE1NOkluc3RhbmNlSUQ9InhtcC5paWQ6NzY3MkJEN0U2N0M1MTFFNkIyQkNFMjQwODEwMDIxNzEiIHhtcE1NOkRvY3VtZW50SUQ9InhtcC5kaWQ6NzY3MkJEN0Y2N0M1MTFFNkIyQkNFMjQwODEwMDIxNzEiPiA8eG1wTU06RGVyaXZlZEZyb20gc3RSZWY6aW5zdGFuY2VJRD0ieG1wLmlpZDo3NjcyQkQ3QzY3QzUxMUU2QjJCQ0UyNDA4MTAwMjE3MSIgc3RSZWY6ZG9jdW1lbnRJRD0ieG1wLmRpZDo3NjcyQkQ3RDY3QzUxMUU2QjJCQ0UyNDA4MTAwMjE3MSIvPiA8L3JkZjpEZXNjcmlwdGlvbj4gPC9yZGY6UkRGPiA8L3g6eG1wbWV0YT4gPD94cGFja2V0IGVuZD0iciI/Pve6J3kAAAAzSURBVHjaYvz//z8D0UDsMwMjSRoYP5Gq4SPNbRjVMEQ1fCRDg+in/6+J1AJUxsgAEGAA31BAJMS0GYEAAAAASUVORK5CYII=) 0 0 #1c1c21; } + + .swagger-ui .b--black { border-color: #000; } + + .swagger-ui .b--near-black { border-color: #121212; } + + .swagger-ui .b--dark-gray { border-color: #333; } + + .swagger-ui .b--mid-gray { border-color: #545454; } + + .swagger-ui .b--gray { border-color: #787878; } + + .swagger-ui .b--silver { border-color: #999; } + + .swagger-ui .b--light-silver { border-color: #6e6e6e; } + + .swagger-ui .b--moon-gray { border-color: #4d4d4d; } + + .swagger-ui .b--light-gray { border-color: #2b2b2b; } + + .swagger-ui .b--near-white { border-color: #242424; } + + .swagger-ui .b--white { border-color: #1c1c21; } + + .swagger-ui .b--white-90 { border-color: rgba(28, 28, 33, .9); } + + .swagger-ui .b--white-80 { border-color: rgba(28, 28, 33, .8); } + + .swagger-ui .b--white-70 { border-color: rgba(28, 28, 33, .7); } + + .swagger-ui .b--white-60 { border-color: rgba(28, 28, 33, .6); } + + .swagger-ui .b--white-50 { border-color: rgba(28, 28, 33, .5); } + + .swagger-ui .b--white-40 { border-color: rgba(28, 28, 33, .4); } + + .swagger-ui .b--white-30 { border-color: rgba(28, 28, 33, .3); } + + .swagger-ui .b--white-20 { border-color: rgba(28, 28, 33, .2); } + + .swagger-ui .b--white-10 { border-color: rgba(28, 28, 33, .1); } + + .swagger-ui .b--white-05 { border-color: rgba(28, 28, 33, .05); } + + .swagger-ui .b--white-025 { border-color: rgba(28, 28, 33, .024); } + + .swagger-ui .b--white-0125 { border-color: rgba(28, 28, 33, .01); } + + .swagger-ui .b--black-90 { border-color: rgba(0, 0, 0, .9); } + + .swagger-ui .b--black-80 { border-color: rgba(0, 0, 0, .8); } + + .swagger-ui .b--black-70 { border-color: rgba(0, 0, 0, .7); } + + .swagger-ui .b--black-60 { border-color: rgba(0, 0, 0, .6); } + + .swagger-ui .b--black-50 { border-color: rgba(0, 0, 0, .5); } + + .swagger-ui .b--black-40 { border-color: rgba(0, 0, 0, .4); } + + .swagger-ui .b--black-30 { border-color: rgba(0, 0, 0, .3); } + + .swagger-ui .b--black-20 { border-color: rgba(0, 0, 0, .2); } + + .swagger-ui .b--black-10 { border-color: rgba(0, 0, 0, .1); } + + .swagger-ui .b--black-05 { border-color: rgba(0, 0, 0, .05); } + + .swagger-ui .b--black-025 { border-color: rgba(0, 0, 0, .024); } + + .swagger-ui .b--black-0125 { border-color: rgba(0, 0, 0, .01); } + + .swagger-ui .b--dark-red { border-color: #bc2f36; } + + .swagger-ui .b--red { border-color: #c83932; } + + .swagger-ui .b--light-red { border-color: #ab3c2b; } + + .swagger-ui .b--orange { border-color: #cc6e33; } + + .swagger-ui .b--purple { border-color: #5e2ca5; } + + .swagger-ui .b--light-purple { border-color: #672caf; } + + .swagger-ui .b--dark-pink { border-color: #ab2b81; } + + .swagger-ui .b--hot-pink { border-color: #c03086; } + + .swagger-ui .b--pink { border-color: #8f2464; } + + .swagger-ui .b--light-pink { border-color: #721d4d; } + + .swagger-ui .b--dark-green { border-color: #1c6e50; } + + .swagger-ui .b--green { border-color: #279b70; } + + .swagger-ui .b--light-green { border-color: #228762; } + + .swagger-ui .b--navy { border-color: #0d1d35; } + + .swagger-ui .b--dark-blue { border-color: #20497e; } + + .swagger-ui .b--blue { border-color: #4380d0; } + + .swagger-ui .b--light-blue { border-color: #20517e; } + + .swagger-ui .b--lightest-blue { border-color: #143a52; } + + .swagger-ui .b--washed-blue { border-color: #0c312d; } + + .swagger-ui .b--washed-green { border-color: #0f3d2c; } + + .swagger-ui .b--washed-red { border-color: #411010; } + + .swagger-ui .b--transparent { border-color: transparent; } + + .swagger-ui .b--gold, .swagger-ui .b--light-yellow, .swagger-ui .b--washed-yellow, .swagger-ui .b--yellow { border-color: #664b00; } + + .swagger-ui .shadow-1 { box-shadow: rgba(0, 0, 0, .2) 0 0 4px 2px; } + + .swagger-ui .shadow-2 { box-shadow: rgba(0, 0, 0, .2) 0 0 8px 2px; } + + .swagger-ui .shadow-3 { box-shadow: rgba(0, 0, 0, .2) 2px 2px 4px 2px; } + + .swagger-ui .shadow-4 { box-shadow: rgba(0, 0, 0, .2) 2px 2px 8px 0; } + + .swagger-ui .shadow-5 { box-shadow: rgba(0, 0, 0, .2) 4px 4px 8px 0; } + + @media screen and (min-width: 30em) { + .swagger-ui .shadow-1-ns { box-shadow: rgba(0, 0, 0, .2) 0 0 4px 2px; } + + .swagger-ui .shadow-2-ns { box-shadow: rgba(0, 0, 0, .2) 0 0 8px 2px; } + + .swagger-ui .shadow-3-ns { box-shadow: rgba(0, 0, 0, .2) 2px 2px 4px 2px; } + + .swagger-ui .shadow-4-ns { box-shadow: rgba(0, 0, 0, .2) 2px 2px 8px 0; } + + .swagger-ui .shadow-5-ns { box-shadow: rgba(0, 0, 0, .2) 4px 4px 8px 0; } + } + + @media screen and (max-width: 60em) and (min-width: 30em) { + .swagger-ui .shadow-1-m { box-shadow: rgba(0, 0, 0, .2) 0 0 4px 2px; } + + .swagger-ui .shadow-2-m { box-shadow: rgba(0, 0, 0, .2) 0 0 8px 2px; } + + .swagger-ui .shadow-3-m { box-shadow: rgba(0, 0, 0, .2) 2px 2px 4px 2px; } + + .swagger-ui .shadow-4-m { box-shadow: rgba(0, 0, 0, .2) 2px 2px 8px 0; } + + .swagger-ui .shadow-5-m { box-shadow: rgba(0, 0, 0, .2) 4px 4px 8px 0; } + } + + @media screen and (min-width: 60em) { + .swagger-ui .shadow-1-l { box-shadow: rgba(0, 0, 0, .2) 0 0 4px 2px; } + + .swagger-ui .shadow-2-l { box-shadow: rgba(0, 0, 0, .2) 0 0 8px 2px; } + + .swagger-ui .shadow-3-l { box-shadow: rgba(0, 0, 0, .2) 2px 2px 4px 2px; } + + .swagger-ui .shadow-4-l { box-shadow: rgba(0, 0, 0, .2) 2px 2px 8px 0; } + + .swagger-ui .shadow-5-l { box-shadow: rgba(0, 0, 0, .2) 4px 4px 8px 0; } + } + + .swagger-ui .black-05 { color: rgba(191, 191, 191, .05); } + + .swagger-ui .bg-black-05 { background-color: rgba(0, 0, 0, .05); } + + .swagger-ui .black-90, .swagger-ui .hover-black-90:focus, .swagger-ui .hover-black-90:hover { color: rgba(191, 191, 191, .9); } + + .swagger-ui .black-80, .swagger-ui .hover-black-80:focus, .swagger-ui .hover-black-80:hover { color: rgba(191, 191, 191, .8); } + + .swagger-ui .black-70, .swagger-ui .hover-black-70:focus, .swagger-ui .hover-black-70:hover { color: rgba(191, 191, 191, .7); } + + .swagger-ui .black-60, .swagger-ui .hover-black-60:focus, .swagger-ui .hover-black-60:hover { color: rgba(191, 191, 191, .6); } + + .swagger-ui .black-50, .swagger-ui .hover-black-50:focus, .swagger-ui .hover-black-50:hover { color: rgba(191, 191, 191, .5); } + + .swagger-ui .black-40, .swagger-ui .hover-black-40:focus, .swagger-ui .hover-black-40:hover { color: rgba(191, 191, 191, .4); } + + .swagger-ui .black-30, .swagger-ui .hover-black-30:focus, .swagger-ui .hover-black-30:hover { color: rgba(191, 191, 191, .3); } + + .swagger-ui .black-20, .swagger-ui .hover-black-20:focus, .swagger-ui .hover-black-20:hover { color: rgba(191, 191, 191, .2); } + + .swagger-ui .black-10, .swagger-ui .hover-black-10:focus, .swagger-ui .hover-black-10:hover { color: rgba(191, 191, 191, .1); } + + .swagger-ui .hover-white-90:focus, .swagger-ui .hover-white-90:hover, .swagger-ui .white-90 { color: rgba(255, 255, 255, .9); } + + .swagger-ui .hover-white-80:focus, .swagger-ui .hover-white-80:hover, .swagger-ui .white-80 { color: rgba(255, 255, 255, .8); } + + .swagger-ui .hover-white-70:focus, .swagger-ui .hover-white-70:hover, .swagger-ui .white-70 { color: rgba(255, 255, 255, .7); } + + .swagger-ui .hover-white-60:focus, .swagger-ui .hover-white-60:hover, .swagger-ui .white-60 { color: rgba(255, 255, 255, .6); } + + .swagger-ui .hover-white-50:focus, .swagger-ui .hover-white-50:hover, .swagger-ui .white-50 { color: rgba(255, 255, 255, .5); } + + .swagger-ui .hover-white-40:focus, .swagger-ui .hover-white-40:hover, .swagger-ui .white-40 { color: rgba(255, 255, 255, .4); } + + .swagger-ui .hover-white-30:focus, .swagger-ui .hover-white-30:hover, .swagger-ui .white-30 { color: rgba(255, 255, 255, .3); } + + .swagger-ui .hover-white-20:focus, .swagger-ui .hover-white-20:hover, .swagger-ui .white-20 { color: rgba(255, 255, 255, .2); } + + .swagger-ui .hover-white-10:focus, .swagger-ui .hover-white-10:hover, .swagger-ui .white-10 { color: rgba(255, 255, 255, .1); } + + .swagger-ui .hover-moon-gray:focus, .swagger-ui .hover-moon-gray:hover, .swagger-ui .moon-gray { color: #ccc; } + + .swagger-ui .hover-light-gray:focus, .swagger-ui .hover-light-gray:hover, .swagger-ui .light-gray { color: #ededed; } + + .swagger-ui .hover-near-white:focus, .swagger-ui .hover-near-white:hover, .swagger-ui .near-white { color: #f5f5f5; } + + .swagger-ui .dark-red, .swagger-ui .hover-dark-red:focus, .swagger-ui .hover-dark-red:hover { color: #e6999d; } + + .swagger-ui .hover-red:focus, .swagger-ui .hover-red:hover, .swagger-ui .red { color: #e69d99; } + + .swagger-ui .hover-light-red:focus, .swagger-ui .hover-light-red:hover, .swagger-ui .light-red { color: #e6a399; } + + .swagger-ui .hover-orange:focus, .swagger-ui .hover-orange:hover, .swagger-ui .orange { color: #e6b699; } + + .swagger-ui .gold, .swagger-ui .hover-gold:focus, .swagger-ui .hover-gold:hover { color: #e6d099; } + + .swagger-ui .hover-yellow:focus, .swagger-ui .hover-yellow:hover, .swagger-ui .yellow { color: #e6da99; } + + .swagger-ui .hover-light-yellow:focus, .swagger-ui .hover-light-yellow:hover, .swagger-ui .light-yellow { color: #ede6b6; } + + .swagger-ui .hover-purple:focus, .swagger-ui .hover-purple:hover, .swagger-ui .purple { color: #b99ae4; } + + .swagger-ui .hover-light-purple:focus, .swagger-ui .hover-light-purple:hover, .swagger-ui .light-purple { color: #bb99e6; } + + .swagger-ui .dark-pink, .swagger-ui .hover-dark-pink:focus, .swagger-ui .hover-dark-pink:hover { color: #e699cc; } + + .swagger-ui .hot-pink, .swagger-ui .hover-hot-pink:focus, .swagger-ui .hover-hot-pink:hover, .swagger-ui .hover-pink:focus, .swagger-ui .hover-pink:hover, .swagger-ui .pink { color: #e699c7; } + + .swagger-ui .hover-light-pink:focus, .swagger-ui .hover-light-pink:hover, .swagger-ui .light-pink { color: #edb6d5; } + + .swagger-ui .dark-green, .swagger-ui .green, .swagger-ui .hover-dark-green:focus, .swagger-ui .hover-dark-green:hover, .swagger-ui .hover-green:focus, .swagger-ui .hover-green:hover { color: #99e6c9; } + + .swagger-ui .hover-light-green:focus, .swagger-ui .hover-light-green:hover, .swagger-ui .light-green { color: #a1e8ce; } + + .swagger-ui .hover-navy:focus, .swagger-ui .hover-navy:hover, .swagger-ui .navy { color: #99b8e6; } + + .swagger-ui .blue, .swagger-ui .dark-blue, .swagger-ui .hover-blue:focus, .swagger-ui .hover-blue:hover, .swagger-ui .hover-dark-blue:focus, .swagger-ui .hover-dark-blue:hover { color: #99bae6; } + + .swagger-ui .hover-light-blue:focus, .swagger-ui .hover-light-blue:hover, .swagger-ui .light-blue { color: #a9cbea; } + + .swagger-ui .hover-lightest-blue:focus, .swagger-ui .hover-lightest-blue:hover, .swagger-ui .lightest-blue { color: #d6e9f5; } + + .swagger-ui .hover-washed-blue:focus, .swagger-ui .hover-washed-blue:hover, .swagger-ui .washed-blue { color: #f7fdfc; } + + .swagger-ui .hover-washed-green:focus, .swagger-ui .hover-washed-green:hover, .swagger-ui .washed-green { color: #ebfaf4; } + + .swagger-ui .hover-washed-yellow:focus, .swagger-ui .hover-washed-yellow:hover, .swagger-ui .washed-yellow { color: #fbf9ef; } + + .swagger-ui .hover-washed-red:focus, .swagger-ui .hover-washed-red:hover, .swagger-ui .washed-red { color: #f9e7e7; } + + .swagger-ui .color-inherit, .swagger-ui .hover-inherit:focus, .swagger-ui .hover-inherit:hover { color: inherit; } + + .swagger-ui .bg-black-90, .swagger-ui .hover-bg-black-90:focus, .swagger-ui .hover-bg-black-90:hover { background-color: rgba(0, 0, 0, .9); } + + .swagger-ui .bg-black-80, .swagger-ui .hover-bg-black-80:focus, .swagger-ui .hover-bg-black-80:hover { background-color: rgba(0, 0, 0, .8); } + + .swagger-ui .bg-black-70, .swagger-ui .hover-bg-black-70:focus, .swagger-ui .hover-bg-black-70:hover { background-color: rgba(0, 0, 0, .7); } + + .swagger-ui .bg-black-60, .swagger-ui .hover-bg-black-60:focus, .swagger-ui .hover-bg-black-60:hover { background-color: rgba(0, 0, 0, .6); } + + .swagger-ui .bg-black-50, .swagger-ui .hover-bg-black-50:focus, .swagger-ui .hover-bg-black-50:hover { background-color: rgba(0, 0, 0, .5); } + + .swagger-ui .bg-black-40, .swagger-ui .hover-bg-black-40:focus, .swagger-ui .hover-bg-black-40:hover { background-color: rgba(0, 0, 0, .4); } + + .swagger-ui .bg-black-30, .swagger-ui .hover-bg-black-30:focus, .swagger-ui .hover-bg-black-30:hover { background-color: rgba(0, 0, 0, .3); } + + .swagger-ui .bg-black-20, .swagger-ui .hover-bg-black-20:focus, .swagger-ui .hover-bg-black-20:hover { background-color: rgba(0, 0, 0, .2); } + + .swagger-ui .bg-white-90, .swagger-ui .hover-bg-white-90:focus, .swagger-ui .hover-bg-white-90:hover { background-color: rgba(28, 28, 33, .9); } + + .swagger-ui .bg-white-80, .swagger-ui .hover-bg-white-80:focus, .swagger-ui .hover-bg-white-80:hover { background-color: rgba(28, 28, 33, .8); } + + .swagger-ui .bg-white-70, .swagger-ui .hover-bg-white-70:focus, .swagger-ui .hover-bg-white-70:hover { background-color: rgba(28, 28, 33, .7); } + + .swagger-ui .bg-white-60, .swagger-ui .hover-bg-white-60:focus, .swagger-ui .hover-bg-white-60:hover { background-color: rgba(28, 28, 33, .6); } + + .swagger-ui .bg-white-50, .swagger-ui .hover-bg-white-50:focus, .swagger-ui .hover-bg-white-50:hover { background-color: rgba(28, 28, 33, .5); } + + .swagger-ui .bg-white-40, .swagger-ui .hover-bg-white-40:focus, .swagger-ui .hover-bg-white-40:hover { background-color: rgba(28, 28, 33, .4); } + + .swagger-ui .bg-white-30, .swagger-ui .hover-bg-white-30:focus, .swagger-ui .hover-bg-white-30:hover { background-color: rgba(28, 28, 33, .3); } + + .swagger-ui .bg-white-20, .swagger-ui .hover-bg-white-20:focus, .swagger-ui .hover-bg-white-20:hover { background-color: rgba(28, 28, 33, .2); } + + .swagger-ui .bg-black, .swagger-ui .hover-bg-black:focus, .swagger-ui .hover-bg-black:hover { background-color: #000; } + + .swagger-ui .bg-near-black, .swagger-ui .hover-bg-near-black:focus, .swagger-ui .hover-bg-near-black:hover { background-color: #121212; } + + .swagger-ui .bg-dark-gray, .swagger-ui .hover-bg-dark-gray:focus, .swagger-ui .hover-bg-dark-gray:hover { background-color: #333; } + + .swagger-ui .bg-mid-gray, .swagger-ui .hover-bg-mid-gray:focus, .swagger-ui .hover-bg-mid-gray:hover { background-color: #545454; } + + .swagger-ui .bg-gray, .swagger-ui .hover-bg-gray:focus, .swagger-ui .hover-bg-gray:hover { background-color: #787878; } + + .swagger-ui .bg-silver, .swagger-ui .hover-bg-silver:focus, .swagger-ui .hover-bg-silver:hover { background-color: #999; } + + .swagger-ui .bg-white, .swagger-ui .hover-bg-white:focus, .swagger-ui .hover-bg-white:hover { background-color: #1c1c21; } + + .swagger-ui .bg-transparent, .swagger-ui .hover-bg-transparent:focus, .swagger-ui .hover-bg-transparent:hover { background-color: transparent; } + + .swagger-ui .bg-dark-red, .swagger-ui .hover-bg-dark-red:focus, .swagger-ui .hover-bg-dark-red:hover { background-color: #bc2f36; } + + .swagger-ui .bg-red, .swagger-ui .hover-bg-red:focus, .swagger-ui .hover-bg-red:hover { background-color: #c83932; } + + .swagger-ui .bg-light-red, .swagger-ui .hover-bg-light-red:focus, .swagger-ui .hover-bg-light-red:hover { background-color: #ab3c2b; } + + .swagger-ui .bg-orange, .swagger-ui .hover-bg-orange:focus, .swagger-ui .hover-bg-orange:hover { background-color: #cc6e33; } + + .swagger-ui .bg-gold, .swagger-ui .bg-light-yellow, .swagger-ui .bg-washed-yellow, .swagger-ui .bg-yellow, .swagger-ui .hover-bg-gold:focus, .swagger-ui .hover-bg-gold:hover, .swagger-ui .hover-bg-light-yellow:focus, .swagger-ui .hover-bg-light-yellow:hover, .swagger-ui .hover-bg-washed-yellow:focus, .swagger-ui .hover-bg-washed-yellow:hover, .swagger-ui .hover-bg-yellow:focus, .swagger-ui .hover-bg-yellow:hover { background-color: #664b00; } + + .swagger-ui .bg-purple, .swagger-ui .hover-bg-purple:focus, .swagger-ui .hover-bg-purple:hover { background-color: #5e2ca5; } + + .swagger-ui .bg-light-purple, .swagger-ui .hover-bg-light-purple:focus, .swagger-ui .hover-bg-light-purple:hover { background-color: #672caf; } + + .swagger-ui .bg-dark-pink, .swagger-ui .hover-bg-dark-pink:focus, .swagger-ui .hover-bg-dark-pink:hover { background-color: #ab2b81; } + + .swagger-ui .bg-hot-pink, .swagger-ui .hover-bg-hot-pink:focus, .swagger-ui .hover-bg-hot-pink:hover { background-color: #c03086; } + + .swagger-ui .bg-pink, .swagger-ui .hover-bg-pink:focus, .swagger-ui .hover-bg-pink:hover { background-color: #8f2464; } + + .swagger-ui .bg-light-pink, .swagger-ui .hover-bg-light-pink:focus, .swagger-ui .hover-bg-light-pink:hover { background-color: #721d4d; } + + .swagger-ui .bg-dark-green, .swagger-ui .hover-bg-dark-green:focus, .swagger-ui .hover-bg-dark-green:hover { background-color: #1c6e50; } + + .swagger-ui .bg-green, .swagger-ui .hover-bg-green:focus, .swagger-ui .hover-bg-green:hover { background-color: #279b70; } + + .swagger-ui .bg-light-green, .swagger-ui .hover-bg-light-green:focus, .swagger-ui .hover-bg-light-green:hover { background-color: #228762; } + + .swagger-ui .bg-navy, .swagger-ui .hover-bg-navy:focus, .swagger-ui .hover-bg-navy:hover { background-color: #0d1d35; } + + .swagger-ui .bg-dark-blue, .swagger-ui .hover-bg-dark-blue:focus, .swagger-ui .hover-bg-dark-blue:hover { background-color: #20497e; } + + .swagger-ui .bg-blue, .swagger-ui .hover-bg-blue:focus, .swagger-ui .hover-bg-blue:hover { background-color: #4380d0; } + + .swagger-ui .bg-light-blue, .swagger-ui .hover-bg-light-blue:focus, .swagger-ui .hover-bg-light-blue:hover { background-color: #20517e; } + + .swagger-ui .bg-lightest-blue, .swagger-ui .hover-bg-lightest-blue:focus, .swagger-ui .hover-bg-lightest-blue:hover { background-color: #143a52; } + + .swagger-ui .bg-washed-blue, .swagger-ui .hover-bg-washed-blue:focus, .swagger-ui .hover-bg-washed-blue:hover { background-color: #0c312d; } + + .swagger-ui .bg-washed-green, .swagger-ui .hover-bg-washed-green:focus, .swagger-ui .hover-bg-washed-green:hover { background-color: #0f3d2c; } + + .swagger-ui .bg-washed-red, .swagger-ui .hover-bg-washed-red:focus, .swagger-ui .hover-bg-washed-red:hover { background-color: #411010; } + + .swagger-ui .bg-inherit, .swagger-ui .hover-bg-inherit:focus, .swagger-ui .hover-bg-inherit:hover { background-color: inherit; } + + .swagger-ui .shadow-hover { transition: all .5s cubic-bezier(.165, .84, .44, 1) 0s; } + + .swagger-ui .shadow-hover::after { + border-radius: inherit; + box-shadow: rgba(0, 0, 0, .2) 0 0 16px 2px; + content: ""; + height: 100%; + left: 0; + opacity: 0; + position: absolute; + top: 0; + transition: opacity .5s cubic-bezier(.165, .84, .44, 1) 0s; + width: 100%; + z-index: -1; + } + + .swagger-ui .bg-animate, .swagger-ui .bg-animate:focus, .swagger-ui .bg-animate:hover { transition: background-color .15s ease-in-out 0s; } + + .swagger-ui .nested-links a { + color: #99bae6; + transition: color .15s ease-in 0s; + } + + .swagger-ui .nested-links a:focus, .swagger-ui .nested-links a:hover { + color: #a9cbea; + transition: color .15s ease-in 0s; + } + + .swagger-ui .opblock-tag { + border-bottom: 1px solid rgba(58, 64, 80, .3); + color: #b5bac9; + transition: all .2s ease 0s; + } + + .swagger-ui .opblock-tag svg, .swagger-ui section.models h4 svg { transition: all .4s ease 0s; } + + .swagger-ui .opblock { + border: 1px solid #000; + border-radius: 4px; + box-shadow: rgba(0, 0, 0, .19) 0 0 3px; + margin: 0 0 15px; + } + + .swagger-ui .opblock .tab-header .tab-item.active h4 span::after { background: gray; } + + .swagger-ui .opblock.is-open .opblock-summary { border-bottom: 1px solid #000; } + + .swagger-ui .opblock .opblock-section-header { + background: rgba(28, 28, 33, .8); + box-shadow: rgba(0, 0, 0, .1) 0 1px 2px; + } + + .swagger-ui .opblock .opblock-section-header > label > span { padding: 0 10px 0 0; } + + .swagger-ui .opblock .opblock-summary-method { + background: #000; + color: #fff; + text-shadow: rgba(0, 0, 0, .1) 0 1px 0; + } + + .swagger-ui .opblock.opblock-post { + background: rgba(72, 203, 144, .1); + border-color: #48cb90; + } + + .swagger-ui .opblock.opblock-post .opblock-summary-method, .swagger-ui .opblock.opblock-post .tab-header .tab-item.active h4 span::after { background: #48cb90; } + + .swagger-ui .opblock.opblock-post .opblock-summary { border-color: #48cb90; } + + .swagger-ui .opblock.opblock-put { + background: rgba(213, 157, 88, .1); + border-color: #d59d58; + } + + .swagger-ui .opblock.opblock-put .opblock-summary-method, .swagger-ui .opblock.opblock-put .tab-header .tab-item.active h4 span::after { background: #d59d58; } + + .swagger-ui .opblock.opblock-put .opblock-summary { border-color: #d59d58; } + + .swagger-ui .opblock.opblock-delete { + background: rgba(200, 50, 50, .1); + border-color: #c83232; + } + + .swagger-ui .opblock.opblock-delete .opblock-summary-method, .swagger-ui .opblock.opblock-delete .tab-header .tab-item.active h4 span::after { background: #c83232; } + + .swagger-ui .opblock.opblock-delete .opblock-summary { border-color: #c83232; } + + .swagger-ui .opblock.opblock-get { + background: rgba(42, 105, 167, .1); + border-color: #2a69a7; + } + + .swagger-ui .opblock.opblock-get .opblock-summary-method, .swagger-ui .opblock.opblock-get .tab-header .tab-item.active h4 span::after { background: #2a69a7; } + + .swagger-ui .opblock.opblock-get .opblock-summary { border-color: #2a69a7; } + + .swagger-ui .opblock.opblock-patch { + background: rgba(92, 214, 188, .1); + border-color: #5cd6bc; + } + + .swagger-ui .opblock.opblock-patch .opblock-summary-method, .swagger-ui .opblock.opblock-patch .tab-header .tab-item.active h4 span::after { background: #5cd6bc; } + + .swagger-ui .opblock.opblock-patch .opblock-summary { border-color: #5cd6bc; } + + .swagger-ui .opblock.opblock-head { + background: rgba(140, 63, 207, .1); + border-color: #8c3fcf; + } + + .swagger-ui .opblock.opblock-head .opblock-summary-method, .swagger-ui .opblock.opblock-head .tab-header .tab-item.active h4 span::after { background: #8c3fcf; } + + .swagger-ui .opblock.opblock-head .opblock-summary { border-color: #8c3fcf; } + + .swagger-ui .opblock.opblock-options { + background: rgba(36, 89, 143, .1); + border-color: #24598f; + } + + .swagger-ui .opblock.opblock-options .opblock-summary-method, .swagger-ui .opblock.opblock-options .tab-header .tab-item.active h4 span::after { background: #24598f; } + + .swagger-ui .opblock.opblock-options .opblock-summary { border-color: #24598f; } + + .swagger-ui .opblock.opblock-deprecated { + background: rgba(46, 46, 46, .1); + border-color: #2e2e2e; + opacity: .6; + } + + .swagger-ui .opblock.opblock-deprecated .opblock-summary-method, .swagger-ui .opblock.opblock-deprecated .tab-header .tab-item.active h4 span::after { background: #2e2e2e; } + + .swagger-ui .opblock.opblock-deprecated .opblock-summary { border-color: #2e2e2e; } + + .swagger-ui .filter .operation-filter-input { border: 2px solid #2b3446; } + + .swagger-ui .tab li:first-of-type::after { background: rgba(0, 0, 0, .2); } + + .swagger-ui .download-contents { + background: #7c8192; + color: #fff; + } + + .swagger-ui .scheme-container { + background: #1c1c21; + box-shadow: rgba(0, 0, 0, .15) 0 1px 2px 0; + } + + .swagger-ui .loading-container .loading::before { + animation: 1s linear 0s infinite normal none running rotation, .5s ease 0s 1 normal none running opacity; + border-color: rgba(0, 0, 0, .6) rgba(84, 84, 84, .1) rgba(84, 84, 84, .1); + } + + .swagger-ui .response-control-media-type--accept-controller select { border-color: #196619; } + + .swagger-ui .response-control-media-type__accept-message { color: #99e699; } + + .swagger-ui .version-pragma__message code { background-color: #3b3b3b; } + + .swagger-ui .btn { + background: 0 0; + border: 2px solid gray; + box-shadow: rgba(0, 0, 0, .1) 0 1px 2px; + color: #b5bac9; + } + + .swagger-ui .btn:hover { box-shadow: rgba(0, 0, 0, .3) 0 0 5px; } + + .swagger-ui .btn.authorize, .swagger-ui .btn.cancel { + background-color: transparent; + border-color: #a72a2a; + color: #e69999; + } + + .swagger-ui .btn.cancel:hover { + background-color: #a72a2a; + color: #fff; + } + + .swagger-ui .btn.authorize { + border-color: #48cb90; + color: #9ce3c3; + } + + .swagger-ui .btn.authorize svg { fill: #9ce3c3; } + + .btn.authorize.unlocked:hover { + background-color: #48cb90; + color: #fff; + } + + .btn.authorize.unlocked:hover svg { + fill: #fbfbfb; + } + + .swagger-ui .btn.execute { + background-color: #5892d5; + border-color: #5892d5; + color: #fff; + } + + .swagger-ui .copy-to-clipboard { background: #7c8192; } + + .swagger-ui .copy-to-clipboard button { background: url("data:image/svg+xml;charset=utf-8,") 50% center no-repeat; } + + .swagger-ui select { + background: url("data:image/svg+xml;charset=utf-8,") right 10px center/20px no-repeat #212121; + background: url(data:image/svg+xml;base64,PD94bWwgdmVyc2lvbj0iMS4wIiBlbmNvZGluZz0iVVRGLTgiIHN0YW5kYWxvbmU9Im5vIj8+CjxzdmcKICAgeG1sbnM6ZGM9Imh0dHA6Ly9wdXJsLm9yZy9kYy9lbGVtZW50cy8xLjEvIgogICB4bWxuczpjYz0iaHR0cDovL2NyZWF0aXZlY29tbW9ucy5vcmcvbnMjIgogICB4bWxuczpyZGY9Imh0dHA6Ly93d3cudzMub3JnLzE5OTkvMDIvMjItcmRmLXN5bnRheC1ucyMiCiAgIHhtbG5zOnN2Zz0iaHR0cDovL3d3dy53My5vcmcvMjAwMC9zdmciCiAgIHhtbG5zPSJodHRwOi8vd3d3LnczLm9yZy8yMDAwL3N2ZyIKICAgeG1sbnM6c29kaXBvZGk9Imh0dHA6Ly9zb2RpcG9kaS5zb3VyY2Vmb3JnZS5uZXQvRFREL3NvZGlwb2RpLTAuZHRkIgogICB4bWxuczppbmtzY2FwZT0iaHR0cDovL3d3dy5pbmtzY2FwZS5vcmcvbmFtZXNwYWNlcy9pbmtzY2FwZSIKICAgaW5rc2NhcGU6dmVyc2lvbj0iMS4wICg0MDM1YTRmYjQ5LCAyMDIwLTA1LTAxKSIKICAgc29kaXBvZGk6ZG9jbmFtZT0iZG93bmxvYWQuc3ZnIgogICBpZD0ic3ZnNCIKICAgdmVyc2lvbj0iMS4xIgogICB2aWV3Qm94PSIwIDAgMjAgMjAiPgogIDxtZXRhZGF0YQogICAgIGlkPSJtZXRhZGF0YTEwIj4KICAgIDxyZGY6UkRGPgogICAgICA8Y2M6V29yawogICAgICAgICByZGY6YWJvdXQ9IiI+CiAgICAgICAgPGRjOmZvcm1hdD5pbWFnZS9zdmcreG1sPC9kYzpmb3JtYXQ+CiAgICAgICAgPGRjOnR5cGUKICAgICAgICAgICByZGY6cmVzb3VyY2U9Imh0dHA6Ly9wdXJsLm9yZy9kYy9kY21pdHlwZS9TdGlsbEltYWdlIiAvPgogICAgICA8L2NjOldvcms+CiAgICA8L3JkZjpSREY+CiAgPC9tZXRhZGF0YT4KICA8ZGVmcwogICAgIGlkPSJkZWZzOCIgLz4KICA8c29kaXBvZGk6bmFtZWR2aWV3CiAgICAgaW5rc2NhcGU6Y3VycmVudC1sYXllcj0ic3ZnNCIKICAgICBpbmtzY2FwZTp3aW5kb3ctbWF4aW1pemVkPSIxIgogICAgIGlua3NjYXBlOndpbmRvdy15PSItOSIKICAgICBpbmtzY2FwZTp3aW5kb3cteD0iLTkiCiAgICAgaW5rc2NhcGU6Y3k9IjEwIgogICAgIGlua3NjYXBlOmN4PSIxMCIKICAgICBpbmtzY2FwZTp6b29tPSI0MS41IgogICAgIHNob3dncmlkPSJmYWxzZSIKICAgICBpZD0ibmFtZWR2aWV3NiIKICAgICBpbmtzY2FwZTp3aW5kb3ctaGVpZ2h0PSIxMDAxIgogICAgIGlua3NjYXBlOndpbmRvdy13aWR0aD0iMTkyMCIKICAgICBpbmtzY2FwZTpwYWdlc2hhZG93PSIyIgogICAgIGlua3NjYXBlOnBhZ2VvcGFjaXR5PSIwIgogICAgIGd1aWRldG9sZXJhbmNlPSIxMCIKICAgICBncmlkdG9sZXJhbmNlPSIxMCIKICAgICBvYmplY3R0b2xlcmFuY2U9IjEwIgogICAgIGJvcmRlcm9wYWNpdHk9IjEiCiAgICAgYm9yZGVyY29sb3I9IiM2NjY2NjYiCiAgICAgcGFnZWNvbG9yPSIjZmZmZmZmIiAvPgogIDxwYXRoCiAgICAgc3R5bGU9ImZpbGw6I2ZmZmZmZiIKICAgICBpZD0icGF0aDIiCiAgICAgZD0iTTEzLjQxOCA3Ljg1OWEuNjk1LjY5NSAwIDAxLjk3OCAwIC42OC42OCAwIDAxMCAuOTY5bC0zLjkwOCAzLjgzYS42OTcuNjk3IDAgMDEtLjk3OSAwbC0zLjkwOC0zLjgzYS42OC42OCAwIDAxMC0uOTY5LjY5NS42OTUgMCAwMS45NzggMEwxMCAxMWwzLjQxOC0zLjE0MXoiIC8+Cjwvc3ZnPgo=) right 10px center/20px no-repeat #1c1c21; + border: 2px solid #41444e; + } + + .swagger-ui select[multiple] { background: #212121; } + + .swagger-ui button.invalid, .swagger-ui input[type=email].invalid, .swagger-ui input[type=file].invalid, .swagger-ui input[type=password].invalid, .swagger-ui input[type=search].invalid, .swagger-ui input[type=text].invalid, .swagger-ui select.invalid, .swagger-ui textarea.invalid { + background: #390e0e; + border-color: #c83232; + } + + .swagger-ui input[type=email], .swagger-ui input[type=file], .swagger-ui input[type=password], .swagger-ui input[type=search], .swagger-ui input[type=text], .swagger-ui textarea { + background: #1c1c21; + border: 1px solid #404040; + } + + .swagger-ui textarea { + background: rgba(28, 28, 33, .8); + color: #b5bac9; + } + + .swagger-ui input[disabled], .swagger-ui select[disabled] { + background-color: #1f1f1f; + color: #bfbfbf; + } + + .swagger-ui textarea[disabled] { + background-color: #41444e; + color: #fff; + } + + .swagger-ui select[disabled] { border-color: #878787; } + + .swagger-ui textarea:focus { border: 2px solid #2a69a7; } + + .swagger-ui .checkbox input[type=checkbox] + label > .item { + background: #303030; + box-shadow: #303030 0 0 0 2px; + } + + .swagger-ui .checkbox input[type=checkbox]:checked + label > .item { background: url("data:image/svg+xml;charset=utf-8,") 50% center no-repeat #303030; } + + .swagger-ui .dialog-ux .backdrop-ux { background: rgba(0, 0, 0, .8); } + + .swagger-ui .dialog-ux .modal-ux { + background: #1c1c21; + border: 1px solid #2e2e2e; + box-shadow: rgba(0, 0, 0, .2) 0 10px 30px 0; + } + + .swagger-ui .dialog-ux .modal-ux-header .close-modal { background: 0 0; } + + .swagger-ui .model .deprecated span, .swagger-ui .model .deprecated td { color: #bfbfbf !important; } + + .swagger-ui .model-toggle::after { background: url("data:image/svg+xml;charset=utf-8,") 50% center/100% no-repeat; } + + .swagger-ui .model-hint { + background: rgba(0, 0, 0, .7); + color: #ebebeb; + } + + .swagger-ui section.models { border: 1px solid rgba(58, 64, 80, .3); } + + .swagger-ui section.models.is-open h4 { border-bottom: 1px solid rgba(58, 64, 80, .3); } + + .swagger-ui section.models .model-container { background: rgba(0, 0, 0, .05); } + + .swagger-ui section.models .model-container:hover { background: rgba(0, 0, 0, .07); } + + .swagger-ui .model-box { background: rgba(0, 0, 0, .1); } + + .swagger-ui .prop-type { color: #aaaad4; } + + .swagger-ui table thead tr td, .swagger-ui table thead tr th { + border-bottom: 1px solid rgba(58, 64, 80, .2); + color: #b5bac9; + } + + .swagger-ui .parameter__name.required::after { color: rgba(230, 153, 153, .6); } + + .swagger-ui .topbar .download-url-wrapper .select-label { color: #f0f0f0; } + + .swagger-ui .topbar .download-url-wrapper .download-url-button { + background: #63a040; + color: #fff; + } + + .swagger-ui .info .title small { background: #7c8492; } + + .swagger-ui .info .title small.version-stamp { background-color: #7a9b27; } + + .swagger-ui .auth-container .errors { + background-color: #350d0d; + color: #b5bac9; + } + + .swagger-ui .errors-wrapper { + background: rgba(200, 50, 50, .1); + border: 2px solid #c83232; + } + + .swagger-ui .markdown code, .swagger-ui .renderedmarkdown code { + background: rgba(0, 0, 0, .05); + color: #c299e6; + } + + .swagger-ui .model-toggle:after { background: url(data:image/svg+xml;base64,PD94bWwgdmVyc2lvbj0iMS4wIiBlbmNvZGluZz0iVVRGLTgiIHN0YW5kYWxvbmU9Im5vIj8+CjxzdmcKICAgeG1sbnM6ZGM9Imh0dHA6Ly9wdXJsLm9yZy9kYy9lbGVtZW50cy8xLjEvIgogICB4bWxuczpjYz0iaHR0cDovL2NyZWF0aXZlY29tbW9ucy5vcmcvbnMjIgogICB4bWxuczpyZGY9Imh0dHA6Ly93d3cudzMub3JnLzE5OTkvMDIvMjItcmRmLXN5bnRheC1ucyMiCiAgIHhtbG5zOnN2Zz0iaHR0cDovL3d3dy53My5vcmcvMjAwMC9zdmciCiAgIHhtbG5zPSJodHRwOi8vd3d3LnczLm9yZy8yMDAwL3N2ZyIKICAgeG1sbnM6c29kaXBvZGk9Imh0dHA6Ly9zb2RpcG9kaS5zb3VyY2Vmb3JnZS5uZXQvRFREL3NvZGlwb2RpLTAuZHRkIgogICB4bWxuczppbmtzY2FwZT0iaHR0cDovL3d3dy5pbmtzY2FwZS5vcmcvbmFtZXNwYWNlcy9pbmtzY2FwZSIKICAgaW5rc2NhcGU6dmVyc2lvbj0iMS4wICg0MDM1YTRmYjQ5LCAyMDIwLTA1LTAxKSIKICAgc29kaXBvZGk6ZG9jbmFtZT0iZG93bmxvYWQyLnN2ZyIKICAgaWQ9InN2ZzQiCiAgIHZlcnNpb249IjEuMSIKICAgaGVpZ2h0PSIyNCIKICAgd2lkdGg9IjI0Ij4KICA8bWV0YWRhdGEKICAgICBpZD0ibWV0YWRhdGExMCI+CiAgICA8cmRmOlJERj4KICAgICAgPGNjOldvcmsKICAgICAgICAgcmRmOmFib3V0PSIiPgogICAgICAgIDxkYzpmb3JtYXQ+aW1hZ2Uvc3ZnK3htbDwvZGM6Zm9ybWF0PgogICAgICAgIDxkYzp0eXBlCiAgICAgICAgICAgcmRmOnJlc291cmNlPSJodHRwOi8vcHVybC5vcmcvZGMvZGNtaXR5cGUvU3RpbGxJbWFnZSIgLz4KICAgICAgPC9jYzpXb3JrPgogICAgPC9yZGY6UkRGPgogIDwvbWV0YWRhdGE+CiAgPGRlZnMKICAgICBpZD0iZGVmczgiIC8+CiAgPHNvZGlwb2RpOm5hbWVkdmlldwogICAgIGlua3NjYXBlOmN1cnJlbnQtbGF5ZXI9InN2ZzQiCiAgICAgaW5rc2NhcGU6d2luZG93LW1heGltaXplZD0iMSIKICAgICBpbmtzY2FwZTp3aW5kb3cteT0iLTkiCiAgICAgaW5rc2NhcGU6d2luZG93LXg9Ii05IgogICAgIGlua3NjYXBlOmN5PSIxMiIKICAgICBpbmtzY2FwZTpjeD0iMTIiCiAgICAgaW5rc2NhcGU6em9vbT0iMzQuNTgzMzMzIgogICAgIHNob3dncmlkPSJmYWxzZSIKICAgICBpZD0ibmFtZWR2aWV3NiIKICAgICBpbmtzY2FwZTp3aW5kb3ctaGVpZ2h0PSIxMDAxIgogICAgIGlua3NjYXBlOndpbmRvdy13aWR0aD0iMTkyMCIKICAgICBpbmtzY2FwZTpwYWdlc2hhZG93PSIyIgogICAgIGlua3NjYXBlOnBhZ2VvcGFjaXR5PSIwIgogICAgIGd1aWRldG9sZXJhbmNlPSIxMCIKICAgICBncmlkdG9sZXJhbmNlPSIxMCIKICAgICBvYmplY3R0b2xlcmFuY2U9IjEwIgogICAgIGJvcmRlcm9wYWNpdHk9IjEiCiAgICAgYm9yZGVyY29sb3I9IiM2NjY2NjYiCiAgICAgcGFnZWNvbG9yPSIjZmZmZmZmIiAvPgogIDxwYXRoCiAgICAgc3R5bGU9ImZpbGw6I2ZmZmZmZiIKICAgICBpZD0icGF0aDIiCiAgICAgZD0iTTEwIDZMOC41OSA3LjQxIDEzLjE3IDEybC00LjU4IDQuNTlMMTAgMThsNi02eiIgLz4KPC9zdmc+Cg==) 50% no-repeat; } + + /* arrows for each operation and request are now white */ + .arrow, #large-arrow-up { fill: #fff; } + + #unlocked { fill: #fff; } + + ::-webkit-scrollbar-track { background-color: #646464 !important; } + + ::-webkit-scrollbar-thumb { + background-color: #242424 !important; + border: 2px solid #3e4346 !important; + } + + ::-webkit-scrollbar-button:vertical:start:decrement { + background: linear-gradient(130deg, #696969 40%, rgba(255, 0, 0, 0) 41%), linear-gradient(230deg, #696969 40%, transparent 41%), linear-gradient(0deg, #696969 40%, transparent 31%); + background-color: #b6b6b6; + } + + ::-webkit-scrollbar-button:vertical:end:increment { + background: linear-gradient(310deg, #696969 40%, transparent 41%), linear-gradient(50deg, #696969 40%, transparent 41%), linear-gradient(180deg, #696969 40%, transparent 31%); + background-color: #b6b6b6; + } + + ::-webkit-scrollbar-button:horizontal:end:increment { + background: linear-gradient(210deg, #696969 40%, transparent 41%), linear-gradient(330deg, #696969 40%, transparent 41%), linear-gradient(90deg, #696969 30%, transparent 31%); + background-color: #b6b6b6; + } + + ::-webkit-scrollbar-button:horizontal:start:decrement { + background: linear-gradient(30deg, #696969 40%, transparent 41%), linear-gradient(150deg, #696969 40%, transparent 41%), linear-gradient(270deg, #696969 30%, transparent 31%); + background-color: #b6b6b6; + } + + ::-webkit-scrollbar-button, ::-webkit-scrollbar-track-piece { background-color: #3e4346 !important; } + + .swagger-ui .black, .swagger-ui .checkbox, .swagger-ui .dark-gray, .swagger-ui .download-url-wrapper .loading, .swagger-ui .errors-wrapper .errors small, .swagger-ui .fallback, .swagger-ui .filter .loading, .swagger-ui .gray, .swagger-ui .hover-black:focus, .swagger-ui .hover-black:hover, .swagger-ui .hover-dark-gray:focus, .swagger-ui .hover-dark-gray:hover, .swagger-ui .hover-gray:focus, .swagger-ui .hover-gray:hover, .swagger-ui .hover-light-silver:focus, .swagger-ui .hover-light-silver:hover, .swagger-ui .hover-mid-gray:focus, .swagger-ui .hover-mid-gray:hover, .swagger-ui .hover-near-black:focus, .swagger-ui .hover-near-black:hover, .swagger-ui .hover-silver:focus, .swagger-ui .hover-silver:hover, .swagger-ui .light-silver, .swagger-ui .markdown pre, .swagger-ui .mid-gray, .swagger-ui .model .property, .swagger-ui .model .property.primitive, .swagger-ui .model-title, .swagger-ui .near-black, .swagger-ui .parameter__extension, .swagger-ui .parameter__in, .swagger-ui .prop-format, .swagger-ui .renderedmarkdown pre, .swagger-ui .response-col_links .response-undocumented, .swagger-ui .response-col_status .response-undocumented, .swagger-ui .silver, .swagger-ui section.models h4, .swagger-ui section.models h5, .swagger-ui span.token-not-formatted, .swagger-ui span.token-string, .swagger-ui table.headers .header-example, .swagger-ui table.model tr.description, .swagger-ui table.model tr.extension { color: #bfbfbf; } + + .swagger-ui .hover-white:focus, .swagger-ui .hover-white:hover, .swagger-ui .info .title small pre, .swagger-ui .topbar a, .swagger-ui .white { color: #fff; } + + .swagger-ui .bg-black-10, .swagger-ui .hover-bg-black-10:focus, .swagger-ui .hover-bg-black-10:hover, .swagger-ui .stripe-dark:nth-child(2n + 1) { background-color: rgba(0, 0, 0, .1); } + + .swagger-ui .bg-white-10, .swagger-ui .hover-bg-white-10:focus, .swagger-ui .hover-bg-white-10:hover, .swagger-ui .stripe-light:nth-child(2n + 1) { background-color: rgba(28, 28, 33, .1); } + + .swagger-ui .bg-light-silver, .swagger-ui .hover-bg-light-silver:focus, .swagger-ui .hover-bg-light-silver:hover, .swagger-ui .striped--light-silver:nth-child(2n + 1) { background-color: #6e6e6e; } + + .swagger-ui .bg-moon-gray, .swagger-ui .hover-bg-moon-gray:focus, .swagger-ui .hover-bg-moon-gray:hover, .swagger-ui .striped--moon-gray:nth-child(2n + 1) { background-color: #4d4d4d; } + + .swagger-ui .bg-light-gray, .swagger-ui .hover-bg-light-gray:focus, .swagger-ui .hover-bg-light-gray:hover, .swagger-ui .striped--light-gray:nth-child(2n + 1) { background-color: #2b2b2b; } + + .swagger-ui .bg-near-white, .swagger-ui .hover-bg-near-white:focus, .swagger-ui .hover-bg-near-white:hover, .swagger-ui .striped--near-white:nth-child(2n + 1) { background-color: #242424; } + + .swagger-ui .opblock-tag:hover, .swagger-ui section.models h4:hover { background: rgba(0, 0, 0, .02); } + + .swagger-ui .checkbox p, .swagger-ui .dialog-ux .modal-ux-content h4, .swagger-ui .dialog-ux .modal-ux-content p, .swagger-ui .dialog-ux .modal-ux-header h3, .swagger-ui .errors-wrapper .errors h4, .swagger-ui .errors-wrapper hgroup h4, .swagger-ui .info .base-url, .swagger-ui .info .title, .swagger-ui .info h1, .swagger-ui .info h2, .swagger-ui .info h3, .swagger-ui .info h4, .swagger-ui .info h5, .swagger-ui .info li, .swagger-ui .info p, .swagger-ui .info table, .swagger-ui .loading-container .loading::after, .swagger-ui .model, .swagger-ui .opblock .opblock-section-header h4, .swagger-ui .opblock .opblock-section-header > label, .swagger-ui .opblock .opblock-summary-description, .swagger-ui .opblock .opblock-summary-operation-id, .swagger-ui .opblock .opblock-summary-path, .swagger-ui .opblock .opblock-summary-path__deprecated, .swagger-ui .opblock-description-wrapper, .swagger-ui .opblock-description-wrapper h4, .swagger-ui .opblock-description-wrapper p, .swagger-ui .opblock-external-docs-wrapper, .swagger-ui .opblock-external-docs-wrapper h4, .swagger-ui .opblock-external-docs-wrapper p, .swagger-ui .opblock-tag small, .swagger-ui .opblock-title_normal, .swagger-ui .opblock-title_normal h4, .swagger-ui .opblock-title_normal p, .swagger-ui .parameter__name, .swagger-ui .parameter__type, .swagger-ui .response-col_links, .swagger-ui .response-col_status, .swagger-ui .responses-inner h4, .swagger-ui .responses-inner h5, .swagger-ui .scheme-container .schemes > label, .swagger-ui .scopes h2, .swagger-ui .servers > label, .swagger-ui .tab li, .swagger-ui label, .swagger-ui select, .swagger-ui table.headers td { color: #b5bac9; } + + .swagger-ui .download-url-wrapper .failed, .swagger-ui .filter .failed, .swagger-ui .model-deprecated-warning, .swagger-ui .parameter__deprecated, .swagger-ui .parameter__name.required span, .swagger-ui table.model tr.property-row .star { color: #e69999; } + + .swagger-ui .opblock-body pre.microlight, .swagger-ui textarea.curl { + background: #41444e; + border-radius: 4px; + color: #fff; + } + + .swagger-ui .expand-methods svg, .swagger-ui .expand-methods:hover svg { fill: #bfbfbf; } + + .swagger-ui .auth-container, .swagger-ui .dialog-ux .modal-ux-header { border-bottom: 1px solid #2e2e2e; } + + .swagger-ui .topbar .download-url-wrapper .select-label select, .swagger-ui .topbar .download-url-wrapper input[type=text] { border: 2px solid #63a040; } + + .swagger-ui .info a, .swagger-ui .info a:hover, .swagger-ui .scopes h2 a { color: #99bde6; } + + /* Dark Scrollbar */ + ::-webkit-scrollbar { + width: 14px; + height: 14px; + } + + ::-webkit-scrollbar-button { + background-color: #3e4346 !important; + } + + ::-webkit-scrollbar-track { + background-color: #646464 !important; + } + + ::-webkit-scrollbar-track-piece { + background-color: #3e4346 !important; + } + + ::-webkit-scrollbar-thumb { + height: 50px; + background-color: #242424 !important; + border: 2px solid #3e4346 !important; + } + + ::-webkit-scrollbar-corner {} + + ::-webkit-resizer {} + + ::-webkit-scrollbar-button:vertical:start:decrement { + background: + linear-gradient(130deg, #696969 40%, rgba(255, 0, 0, 0) 41%), + linear-gradient(230deg, #696969 40%, rgba(0, 0, 0, 0) 41%), + linear-gradient(0deg, #696969 40%, rgba(0, 0, 0, 0) 31%); + background-color: #b6b6b6; + } + + ::-webkit-scrollbar-button:vertical:end:increment { + background: + linear-gradient(310deg, #696969 40%, rgba(0, 0, 0, 0) 41%), + linear-gradient(50deg, #696969 40%, rgba(0, 0, 0, 0) 41%), + linear-gradient(180deg, #696969 40%, rgba(0, 0, 0, 0) 31%); + background-color: #b6b6b6; + } + + ::-webkit-scrollbar-button:horizontal:end:increment { + background: + linear-gradient(210deg, #696969 40%, rgba(0, 0, 0, 0) 41%), + linear-gradient(330deg, #696969 40%, rgba(0, 0, 0, 0) 41%), + linear-gradient(90deg, #696969 30%, rgba(0, 0, 0, 0) 31%); + background-color: #b6b6b6; + } + + ::-webkit-scrollbar-button:horizontal:start:decrement { + background: + linear-gradient(30deg, #696969 40%, rgba(0, 0, 0, 0) 41%), + linear-gradient(150deg, #696969 40%, rgba(0, 0, 0, 0) 41%), + linear-gradient(270deg, #696969 30%, rgba(0, 0, 0, 0) 31%); + background-color: #b6b6b6; + } +} \ No newline at end of file diff --git a/installer.py b/installer.py index 492961439..6f48e5790 100644 --- a/installer.py +++ b/installer.py @@ -27,7 +27,7 @@ log_file = os.path.join(os.path.dirname(__file__), 'sdnext.log') log_rolled = False first_call = True quick_allowed = True -errors = 0 +errors = [] opts = {} args = Dot({ 'debug': False, @@ -110,6 +110,9 @@ def setup_logging(): "traceback.border": "black", "traceback.border.syntax_error": "black", "inspect.value.border": "black", + "logging.level.info": "blue_violet", + "logging.level.debug": "purple4", + "logging.level.trace": "dark_blue", })) logging.basicConfig(level=logging.ERROR, format='%(asctime)s | %(name)s | %(levelname)s | %(module)s | %(message)s', handlers=[logging.NullHandler()]) # redirect default logger to null pretty_install(console=console) @@ -277,8 +280,7 @@ def pip(arg: str, ignore: bool = False, quiet: bool = False, uv = True): txt = txt.strip() debug(f'Install {pipCmd}: {txt}') if result.returncode != 0 and not ignore: - global errors # pylint: disable=global-statement - errors += 1 + errors.append(f'pip: {package}') log.error(f'Install: {pipCmd}: {arg}') log.debug(f'Install: pip output {txt}') return txt @@ -321,8 +323,7 @@ def git(arg: str, folder: str = None, ignore: bool = False, optional: bool = Fal if result.returncode != 0 and not ignore: if "couldn't find remote ref" in txt: # not a git repo return txt - global errors # pylint: disable=global-statement - errors += 1 + errors.append(f'git: {folder}') log.error(f'Git: {folder} / {arg}') if 'or stash them' in txt: log.error(f'Git local changes detected: check details log="{log_file}"') @@ -377,10 +378,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 @@ -410,14 +412,14 @@ def get_platform(): else: release = platform.release() return { - # 'host': platform.node(), 'arch': platform.machine(), 'cpu': platform.processor(), 'system': platform.system(), 'release': release, - # 'platform': platform.platform(aliased = True, terse = False), - # 'version': platform.version(), 'python': platform.python_version(), + 'docker': os.environ.get('SD_INSTALL_DEBUG', None) is not None, + # 'host': platform.node(), + # 'version': platform.version(), } except Exception as e: return { 'error': e } @@ -455,12 +457,14 @@ def check_python(supported_minors=[9, 10, 11, 12], reason=None): # check diffusers version def check_diffusers(): - sha = '0d1d267b12e47b40b0e8f265339c76e0f45f8c49' + if args.skip_all or args.skip_requirements: + return + sha = '345907f32de71c8ca67f3d9d00e37127192da543' 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 '' if (minor == 0) or (cur != sha): - log.debug(f'Diffusers {"install" if minor == 0 else "upgrade"}: package={pkg} current={cur} target={sha}') + log.info(f'Diffusers {"install" if minor == 0 else "upgrade"}: package={pkg} current={cur} target={sha}') if minor > 0: pip('uninstall --yes diffusers', ignore=True, quiet=True, uv=False) pip(f'install --upgrade git+https://github.com/huggingface/diffusers@{sha}', ignore=False, quiet=True, uv=False) @@ -470,6 +474,8 @@ def check_diffusers(): # check onnx version def check_onnx(): + if args.skip_all or args.skip_requirements: + return if not installed('onnx', quiet=True): install('onnx', 'onnx', ignore=True) if not installed('onnxruntime', quiet=True) and not (installed('onnxruntime-gpu', quiet=True) or installed('onnxruntime-openvino', quiet=True) or installed('onnxruntime-training', quiet=True)): # allow either @@ -477,6 +483,8 @@ def check_onnx(): def check_torchao(): + if args.skip_all or args.skip_requirements: + return if installed('torchao', quiet=True): ver = package_version('torchao') if ver != '0.5.0': @@ -488,14 +496,16 @@ def check_torchao(): def install_cuda(): log.info('CUDA: nVidia toolkit detected') - install('onnxruntime-gpu', 'onnxruntime-gpu', ignore=True, quiet=True) + if not (args.skip_all or args.skip_requirements): + 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.5.1+cu124 torchvision==0.20.1+cu124 --index-url https://download.pytorch.org/whl/cu124') def install_rocm_zluda(): + if args.skip_all or args.skip_requirements: + return None from modules import rocm - if not rocm.is_installed: log.warning('ROCm: could not find ROCm toolkit installed') log.info('Using CPU-only torch') @@ -547,7 +557,6 @@ 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) @@ -595,11 +604,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}') @@ -657,7 +666,7 @@ def install_torch_addons(): triton_command = os.environ.get('TRITON_COMMAND', 'triton') if sys.platform == 'linux' else None if 'xformers' in xformers_package: try: - install(f'--no-deps {xformers_package}', ignore=True) + install(xformers_package, ignore=True, no_deps=True) import torch # pylint: disable=unused-import import xformers # pylint: disable=unused-import except Exception as e: @@ -844,8 +853,7 @@ def run_extension_installer(folder): txt = result.stdout.decode(encoding="utf8", errors="ignore") debug(f'Extension installer: file="{path_installer}" {txt}') if result.returncode != 0: - global errors # pylint: disable=global-statement - errors += 1 + errors.append(f'ext: {os.path.basename(folder)}') if len(result.stderr) > 0: txt = txt + '\n' + result.stderr.decode(encoding="utf8", errors="ignore") log.error(f'Extension installer error: {path_installer}') @@ -986,6 +994,29 @@ def ensure_base_requirements(): install('requests', 'requests', quiet=True) +def install_optional(): + log.info('Installing optional requirements...') + install('basicsr') + install('gfpgan') + install('clean-fid') + install('optimum-quanto', ignore=True) + install('bitsandbytes', ignore=True) + install('pynvml', ignore=True) + install('ultralytics', ignore=True) + install('Cython', ignore=True) + install('insightface', ignore=True) # problematic build + install('nncf==2.7.0', ignore=True, no_deps=True) # requires older pandas + # install('flash-attn', ignore=True) # requires cuda and nvcc to be installed + install('gguf', ignore=True) + try: + import gguf + scripts_dir = os.path.join(os.path.dirname(gguf.__file__), '..', 'scripts') + if os.path.exists(scripts_dir): + os.rename(scripts_dir, scripts_dir + '_gguf') + except Exception: + pass + + def install_requirements(): if args.profile: pr = cProfile.Profile() @@ -995,10 +1026,13 @@ def install_requirements(): if not installed('diffusers', quiet=True): # diffusers are not installed, so run initial installation global quick_allowed # pylint: disable=global-statement quick_allowed = False - log.info('Installing requirements: this may take a while...') + log.info('Install requirements: this may take a while...') pip('install -r requirements.txt') + if args.optional: + quick_allowed = False + install_optional() installed('torch', reload=True) # reload packages cache - log.info('Verifying requirements') + log.info('Install: verifying requirements') with open('requirements.txt', 'r', encoding='utf8') as f: lines = [line.strip() for line in f.readlines() if line.strip() != '' and not line.startswith('#') and line is not None] for line in lines: @@ -1139,6 +1173,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: @@ -1239,6 +1295,7 @@ def add_args(parser): 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('--optional', default = os.environ.get("SD_OPTIONAL",False), action='store_true', help = "Force installation of optional 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_startup = parser.add_argument_group('Startup') diff --git a/javascript/changelog.js b/javascript/changelog.js index 80c9956b8..97dc9daf5 100644 --- a/javascript/changelog.js +++ b/javascript/changelog.js @@ -75,3 +75,10 @@ async function initChangelog() { }; search.addEventListener('keyup', searchChangelog); } + +function wikiSearch(txt) { + log('wikiSearch', txt); + const url = `https://github.com/search?q=repo%3Avladmandic%2Fautomatic+${encodeURIComponent(txt)}&type=wikis`; + // window.open(url, '_blank').focus(); + return txt; +} diff --git a/javascript/extraNetworks.js b/javascript/extraNetworks.js index 9a33baa86..77fe125f3 100644 --- a/javascript/extraNetworks.js +++ b/javascript/extraNetworks.js @@ -481,11 +481,13 @@ function setupExtraNetworksForTab(tabname) { en.style.position = 'absolute'; en.style.height = 'auto'; en.style.width = `${window.opts.extra_networks_sidebar_width}vw`; + en.style.maxWidth = '655px'; 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`; + // gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = `${100 - 2 - window.opts.extra_networks_sidebar_width}vw`; + gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = `calc(100vw - 2em - min(${window.opts.extra_networks_sidebar_width}vw, 655px))`; } else { en.style.position = 'relative'; en.style.height = 'unset'; diff --git a/javascript/sdnext.css b/javascript/sdnext.css index d412d966f..08fae2eb8 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -13,11 +13,12 @@ footer { display: none; margin-top: 0 !important;} table { overflow-x: auto !important; overflow-y: auto !important; } td { border-bottom: none !important; padding: 0 0.5em !important; } tr { border-bottom: none !important; padding: 0 0.5em !important; } +td > div > span { overflow-y: auto; max-height: 3em; overflow-x: hidden; } textarea { overflow-y: auto !important; } span { font-size: var(--text-md) !important; } button { font-size: var(--text-lg) !important; } input[type='color'] { width: 64px; height: 32px; } -td > div > span { overflow-y: auto; max-height: 3em; overflow-x: hidden; } +input::-webkit-outer-spin-button, input::-webkit-inner-spin-button { margin-left: 4px; } /* gradio elements */ .block .padded:not(.gradio-accordion) { padding: 4px 0 0 0 !important; margin-right: 0; min-width: 90px !important; } @@ -326,6 +327,11 @@ div:has(>#tab-gallery-folders) { flex-grow: 0 !important; background-color: var( .changelog_arrow:hover { background-color: var(--button-primary-border-color-hover); } .changelog_highlight { background-color: var(--color-warning); } +/* wiki */ +#wiki_result > div > div { padding: 0.5em; margin-right: 2em; } +#wiki_result li { display: block; } +#wiki_result h3 { background-color: var(--background-fill-primary); margin: 0; padding: 0.3em; margin-bottom: 0.2em; } + /* 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/launch.py b/launch.py index 8379a5d6d..f944a7e54 100755 --- a/launch.py +++ b/launch.py @@ -204,12 +204,13 @@ def main(): installer.check_python() if args.reset: installer.git_reset() - if args.skip_git: + if args.skip_git or args.skip_all: 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: + if not args.skip_env or args.skip_all: installer.set_environment() if args.uv: installer.install("uv", "uv") @@ -239,7 +240,7 @@ def main(): installer.install_extensions() installer.install_requirements() # redo requirements since extensions may change them installer.update_wiki() - if installer.errors == 0: + if len(installer.errors) == 0: installer.log.debug(f'Setup complete without errors: {round(time.time())}') else: installer.log.warning(f'Setup complete with errors: {installer.errors}') @@ -257,9 +258,7 @@ def main(): alive = False requests = 0 if round(time.time()) % 120 == 0: - state = f'job="{instance.state.job}" {instance.state.job_no}/{instance.state.job_count}' if instance.state.job != '' or instance.state.job_no != 0 or instance.state.job_count != 0 else 'idle' - uptime = round(time.time() - instance.state.server_start) - installer.log.debug(f'Server: alive={alive} jobs={instance.state.total_jobs} requests={requests} uptime={uptime} memory={get_memory_stats()} backend={instance.backend} state={state}') + installer.log.debug(f'Server: alive={alive} requests={requests} memory={get_memory_stats()} {instance.state.status()}') if not alive: if uv is not None and uv.wants_restart: installer.log.info('Server restarting...') diff --git a/modules/api/api.py b/modules/api/api.py index 23c0a77f1..f8346995d 100644 --- a/modules/api/api.py +++ b/modules/api/api.py @@ -5,7 +5,7 @@ from fastapi import FastAPI, APIRouter, Depends, Request from fastapi.security import HTTPBasic, HTTPBasicCredentials from fastapi.exceptions import HTTPException from modules import errors, shared, postprocessing -from modules.api import models, endpoints, script, helpers, server, nvml, generate, process, control, gallery +from modules.api import models, endpoints, script, helpers, server, nvml, generate, process, control, gallery, docs errors.install() @@ -23,8 +23,10 @@ class Api: for line in file.readlines(): user, password = line.split(":") self.credentials[user.replace('"', '').strip()] = password.replace('"', '').strip() - self.router = APIRouter() + if shared.cmd_opts.docs: + docs.create_docs(app) + docs.create_redocs(app) self.app = app self.queue_lock = queue_lock self.generate = generate.APIGenerate(queue_lock) @@ -36,6 +38,7 @@ class Api: self.add_api_route("/sdapi/v1/log", server.get_log_buffer, methods=["GET"], response_model=List[str]) self.add_api_route("/sdapi/v1/start", self.get_session_start, methods=["GET"]) self.add_api_route("/sdapi/v1/version", server.get_version, methods=["GET"]) + self.add_api_route("/sdapi/v1/status", server.get_status, methods=["GET"], response_model=models.ResStatus) self.add_api_route("/sdapi/v1/platform", server.get_platform, methods=["GET"]) self.add_api_route("/sdapi/v1/progress", server.get_progress, methods=["GET"], response_model=models.ResProgress) self.add_api_route("/sdapi/v1/interrupt", server.post_interrupt, methods=["POST"]) @@ -55,7 +58,7 @@ class Api: self.add_api_route("/sdapi/v1/extra-batch-images", self.extras_batch_images_api, methods=["POST"], response_model=models.ResProcessBatch) self.add_api_route("/sdapi/v1/preprocess", self.process.post_preprocess, methods=["POST"]) self.add_api_route("/sdapi/v1/mask", self.process.post_mask, methods=["POST"]) - self.add_api_route("/sdapi/v1/faces", self.process.post_face, methods=["POST"]) + self.add_api_route("/sdapi/v1/detect", self.process.post_detect, methods=["POST"]) # api dealing with optional scripts self.add_api_route("/sdapi/v1/scripts", script.get_scripts_list, methods=["GET"], response_model=models.ResScripts) diff --git a/modules/api/control.py b/modules/api/control.py index cf8916095..29c5a77f1 100644 --- a/modules/api/control.py +++ b/modules/api/control.py @@ -31,6 +31,7 @@ ReqControl = models.create_model_from_signature( {"key": "ip_adapter", "type": Optional[List[models.ItemIPAdapter]], "default": None, "exclude": True}, {"key": "face", "type": Optional[models.ItemFace], "default": None, "exclude": True}, {"key": "control", "type": Optional[List[ItemControl]], "default": [], "exclude": True}, + {"key": "extra", "type": Optional[dict], "default": {}, "exclude": True}, ] ) @@ -103,7 +104,7 @@ class APIControl(): args['ip_adapter_scales'].append(ipadapter.scale) args['ip_adapter_starts'].append(ipadapter.start) args['ip_adapter_ends'].append(ipadapter.end) - args['ip_adapter_crops'].append(ipadapter.end) + args['ip_adapter_crops'].append(ipadapter.crop) args['ip_adapter_images'].append([helpers.decode_base64_to_image(x) for x in ipadapter.images]) if ipadapter.masks: args['ip_adapter_masks'].append([helpers.decode_base64_to_image(x) for x in ipadapter.masks]) @@ -159,6 +160,7 @@ class APIControl(): output_processed = [] output_info = '' run.control_set({ 'do_not_save_grid': not req.save_images, 'do_not_save_samples': not req.save_images, **self.prepare_ip_adapter(req) }) + run.control_set(getattr(req, "extra", {})) res = run.control_run(**args) for item in res: if len(item) > 0 and (isinstance(item[0], list) or item[0] is None): # output_images diff --git a/modules/api/docs.py b/modules/api/docs.py new file mode 100644 index 000000000..f384a1328 --- /dev/null +++ b/modules/api/docs.py @@ -0,0 +1,92 @@ +import json +from starlette.responses import HTMLResponse +from fastapi import FastAPI +from fastapi.openapi.docs import get_redoc_html, swagger_ui_default_parameters +from fastapi.encoders import jsonable_encoder + + +def get_swagger_ui_html(*, + openapi_url: str, + title: str, + swagger_js_url: str = "https://cdn.jsdelivr.net/npm/swagger-ui-dist@5/swagger-ui-bundle.js", + swagger_css_url: str = "https://cdn.jsdelivr.net/npm/swagger-ui-dist@5/swagger-ui.css", + swagger_extra_css_url: str = None, + swagger_favicon_url: str = "https://fastapi.tiangolo.com/img/favicon.png", + oauth2_redirect_url: str = None, + init_oauth: dict = None, + swagger_ui_parameters: dict = None, + ) -> HTMLResponse: + current_swagger_ui_parameters = swagger_ui_default_parameters.copy() + if swagger_ui_parameters: + current_swagger_ui_parameters.update(swagger_ui_parameters) + html = f""" + + + + + + {title} + + +
+ + + + + """ + return HTMLResponse(html) + + +def create_docs(app: FastAPI): + swagger_ui_parameters = { + "displayOperationId": True, + "layout": "BaseLayout", + "showExtensions": True, + "showCommonExtensions": True, + "deepLinking": False, + "dom_id": "#swagger-ui", + } + + @app.get("/docs", include_in_schema=True) + async def custom_swagger_html(): + res = get_swagger_ui_html( + title=f'{app.title}: Swagger UI', + openapi_url=app.openapi_url, + swagger_favicon_url='/file=html/favicon.svg', + swagger_ui_parameters=swagger_ui_parameters, + swagger_extra_css_url='file=html/swagger.css', + ) + # res = inject_css(html.content, 'html/swagger.css') + return res + + +def create_redocs(app: FastAPI): + @app.get("/redocs", include_in_schema=True) + async def custom_redoc_html(): + res = get_redoc_html( + title=f'{app.title}: ReDoc', + openapi_url=app.openapi_url, + redoc_favicon_url='/file=html/favicon.svg', + ) + return res diff --git a/modules/api/generate.py b/modules/api/generate.py index aeafa05a0..d940036fc 100644 --- a/modules/api/generate.py +++ b/modules/api/generate.py @@ -40,7 +40,7 @@ class APIGenerate(): sanitize_str(request.script_args) def prepare_face_module(self, request): - if hasattr(request, "face") and request.face and not request.script_name and (not request.alwayson_scripts or "face" not in request.alwayson_scripts.keys()): + if getattr(request, "face", None) is not None and (not request.alwayson_scripts or "face" not in request.alwayson_scripts.keys()): request.script_name = "face" request.script_args = [ request.face.mode, @@ -106,6 +106,8 @@ class APIGenerate(): p.scripts = script_runner p.outpath_grids = shared.opts.outdir_grids or shared.opts.outdir_txt2img_grids p.outpath_samples = shared.opts.outdir_samples or shared.opts.outdir_txt2img_samples + for key, value in getattr(txt2imgreq, "extra", {}).items(): + setattr(p, key, value) shared.state.begin('API TXT', api=True) script_args = script.init_script_args(p, txt2imgreq, self.default_script_arg_txt2img, selectable_scripts, selectable_script_idx, script_runner) if selectable_scripts is not None: @@ -114,7 +116,10 @@ class APIGenerate(): p.script_args = tuple(script_args) # Need to pass args as tuple here processed = process_images(p) shared.state.end(api=False) - b64images = list(map(helpers.encode_pil_to_base64, processed.images)) if send_images else [] + if processed.images is None or len(processed.images) == 0: + b64images = [] + else: + b64images = list(map(helpers.encode_pil_to_base64, processed.images)) if send_images else [] self.sanitize_b64(txt2imgreq) return models.ResTxt2Img(images=b64images, parameters=vars(txt2imgreq), info=processed.js()) @@ -150,6 +155,8 @@ class APIGenerate(): p.scripts = script_runner p.outpath_grids = shared.opts.outdir_img2img_grids p.outpath_samples = shared.opts.outdir_img2img_samples + for key, value in getattr(img2imgreq, "extra", {}).items(): + setattr(p, key, value) shared.state.begin('API-IMG', api=True) script_args = script.init_script_args(p, img2imgreq, self.default_script_arg_img2img, selectable_scripts, selectable_script_idx, script_runner) if selectable_scripts is not None: @@ -158,7 +165,10 @@ class APIGenerate(): p.script_args = tuple(script_args) # Need to pass args as tuple here processed = process_images(p) shared.state.end(api=False) - b64images = list(map(helpers.encode_pil_to_base64, processed.images)) if send_images else [] + if processed.images is None or len(processed.images) == 0: + b64images = [] + else: + b64images = list(map(helpers.encode_pil_to_base64, processed.images)) if send_images else [] if not img2imgreq.include_init_images: img2imgreq.init_images = None img2imgreq.mask = None diff --git a/modules/api/helpers.py b/modules/api/helpers.py index 1678a851e..d9a87537e 100644 --- a/modules/api/helpers.py +++ b/modules/api/helpers.py @@ -14,7 +14,7 @@ def validate_sampler_name(name): return name -def decode_base64_to_image(encoding): +def decode_base64_to_image(encoding, quiet=False): if encoding.startswith("data:image/"): encoding = encoding.split(";")[1].split(",")[1] try: @@ -22,7 +22,9 @@ def decode_base64_to_image(encoding): return image except Exception as e: shared.log.warning(f'API cannot decode image: {e}') - raise HTTPException(status_code=500, detail="Invalid encoded image") from e + if not quiet: + raise HTTPException(status_code=500, detail="Invalid encoded image") from e + return None def encode_pil_to_base64(image): diff --git a/modules/api/middleware.py b/modules/api/middleware.py index 095c5b23d..7eb2c40e8 100644 --- a/modules/api/middleware.py +++ b/modules/api/middleware.py @@ -90,4 +90,4 @@ def setup_middleware(app: FastAPI, cmd_opts): return handle_exception(req, e) app.build_middleware_stack() # rebuild middleware stack on-the-fly - log.debug(f'FastAPI middleware: {[m.__class__.__name__ for m in app.user_middleware]}') + log.debug(f'API middleware: {[m.cls for m in app.user_middleware]}') diff --git a/modules/api/models.py b/modules/api/models.py index 3cf3aade9..e68ebf081 100644 --- a/modules/api/models.py +++ b/modules/api/models.py @@ -11,14 +11,6 @@ API_NOT_ALLOWED = [ "sd_model", "outpath_samples", "outpath_grids", - "sampler_index", - "extra_generation_params", - "overlay_images", - "do_not_reload_embeddings", - "seed_enable_extras", - "prompt_for_display", - "sampler_noise_scheduler_override", - "ddim_discretize" ] class ModelDef(BaseModel): @@ -202,14 +194,17 @@ ReqTxt2Img = PydanticModelGenerator( "StableDiffusionProcessingTxt2Img", StableDiffusionProcessingTxt2Img, [ - {"key": "sampler_index", "type": str, "default": "UniPC"}, - {"key": "script_name", "type": str, "default": None}, + {"key": "sampler_index", "type": int, "default": 0}, + {"key": "sampler_name", "type": str, "default": "UniPC"}, + {"key": "hr_sampler_name", "type": str, "default": "Same as primary"}, + {"key": "script_name", "type": str, "default": "none"}, {"key": "script_args", "type": list, "default": []}, {"key": "send_images", "type": bool, "default": True}, {"key": "save_images", "type": bool, "default": False}, {"key": "alwayson_scripts", "type": dict, "default": {}}, {"key": "ip_adapter", "type": Optional[List[ItemIPAdapter]], "default": None, "exclude": True}, {"key": "face", "type": Optional[ItemFace], "default": None, "exclude": True}, + {"key": "extra", "type": Optional[dict], "default": {}, "exclude": True}, ] ).generate_model() StableDiffusionTxt2ImgProcessingAPI = ReqTxt2Img @@ -223,7 +218,11 @@ ReqImg2Img = PydanticModelGenerator( "StableDiffusionProcessingImg2Img", StableDiffusionProcessingImg2Img, [ - {"key": "sampler_index", "type": str, "default": "UniPC"}, + {"key": "sampler_index", "type": int, "default": 0}, + {"key": "sampler_name", "type": str, "default": "UniPC"}, + {"key": "hr_sampler_name", "type": str, "default": "Same as primary"}, + {"key": "script_name", "type": str, "default": "none"}, + {"key": "script_args", "type": list, "default": []}, {"key": "init_images", "type": list, "default": None}, {"key": "denoising_strength", "type": float, "default": 0.5}, {"key": "mask", "type": str, "default": None}, @@ -235,6 +234,7 @@ ReqImg2Img = PydanticModelGenerator( {"key": "alwayson_scripts", "type": dict, "default": {}}, {"key": "ip_adapter", "type": Optional[List[ItemIPAdapter]], "default": None, "exclude": True}, {"key": "face_id", "type": Optional[ItemFace], "default": None, "exclude": True}, + {"key": "extra", "type": Optional[dict], "default": {}, "exclude": True}, ] ).generate_model() StableDiffusionImg2ImgProcessingAPI = ReqImg2Img @@ -300,6 +300,23 @@ class ResProgress(BaseModel): current_image: str = Field(default=None, title="Current image", description="The current image in base64 format. opts.show_progress_every_n_steps is required for this to work.") textinfo: str = Field(default=None, title="Info text", description="Info text used by WebUI.") +class ResStatus(BaseModel): + status: str = Field(title="Status", description="Current status") + task: str = Field(title="Task", description="Current task") + timestamp: Optional[str] = Field(title="Timestamp", description="Timestamp of the current job") + id: str = Field(title="ID", description="ID of the current task") + job: int = Field(title="Job", description="Current job") + jobs: int = Field(title="Jobs", description="Total jobs") + total: int = Field(title="Total Jobs", description="Total jobs") + step: int = Field(title="Step", description="Current step") + steps: int = Field(title="Steps", description="Total steps") + queued: int = Field(title="Queued", description="Number of queued tasks") + uptime: int = Field(title="Uptime", description="Uptime of the server") + elapsed: Optional[float] = Field(title="Elapsed time") + eta: Optional[float] = Field(title="ETA in secs") + progress: Optional[float] = Field(title="Progress", description="The progress with a range of 0 to 1") + + class ReqInterrogate(BaseModel): image: str = Field(default="", title="Image", description="Image to work on, must be a Base64 string containing the image's data.") clip_model: str = Field(default="", title="CLiP Model", description="The interrogate model used.") diff --git a/modules/api/process.py b/modules/api/process.py index f50d58381..80b19c52e 100644 --- a/modules/api/process.py +++ b/modules/api/process.py @@ -28,8 +28,12 @@ class ReqMask(BaseModel): class ReqFace(BaseModel): image: str = Field(title="Image", description="The base64 encoded image") + model: Optional[str] = Field(title="Model", description="The model to use for detection") class ResFace(BaseModel): + classes: List[int] = Field(title="Class", description="The class of detected item") + labels: List[str] = Field(title="Label", description="The label of detected item") + boxes: List[List[int]] = Field(title="Box", description="The bounding box of detected item") images: List[str] = Field(title="Image", description="The base64 encoded images of detected faces") scores: List[float] = Field(title="Scores", description="The scores of the detected faces") @@ -106,16 +110,22 @@ class APIProcess(): image = encode_pil_to_base64(processed) return ResMask(mask=image) - def post_face(self, req: ReqFace): - from shared import yolo # pylint: disable=no-name-in-module + def post_detect(self, req: ReqFace): + from modules.shared import yolo # pylint: disable=no-name-in-module image = decode_base64_to_image(req.image) shared.state.begin('API-FACE', api=True) images = [] scores = [] + classes = [] + boxes = [] + labels = [] with self.queue_lock: - faces = yolo.predict('face-yolo8n', image) - for face in faces: - images.append(encode_pil_to_base64(face.item)) - scores.append(face.score) + items = yolo.predict(req.model, image) + for item in items: + images.append(encode_pil_to_base64(item.item)) + scores.append(item.score) + classes.append(item.cls) + labels.append(item.label) + boxes.append(item.box) shared.state.end(api=False) - return ResFace(images=images, scores=scores) + return ResFace(classes=classes, labels=labels, scores=scores, boxes=boxes, images=images) diff --git a/modules/api/script.py b/modules/api/script.py index cae59791e..2c0814ef0 100644 --- a/modules/api/script.py +++ b/modules/api/script.py @@ -3,27 +3,39 @@ from fastapi.exceptions import HTTPException import gradio as gr from modules.api import models from modules import scripts +from modules.errors import log def script_name_to_index(name, scripts_list): - try: - return [script.title().lower() for script in scripts_list].index(name.lower()) - except Exception as e: - raise HTTPException(status_code=422, detail=f"Script '{name}' not found") from e + if name is None or len(name) == 0 or name == 'none': + return None + available = [script.title().lower() for script in scripts_list] + if name.lower() in available: + return available.index(name.lower()) + short = [available.split(':')[0] for available in available] + if name.lower() in short: + return short.index(name.lower()) + log.error(f'API: script={name} available={available} not found') + return None + def get_selectable_script(script_name, script_runner): - if script_name is None or script_name == "": + if script_name is None or script_name == "" or script_name == 'none': return None, None script_idx = script_name_to_index(script_name, script_runner.selectable_scripts) + if script_idx is None: + return None, None script = script_runner.selectable_scripts[script_idx] return script, script_idx + def get_scripts_list(): t2ilist = [script.name for script in scripts.scripts_txt2img.scripts if script.name is not None] i2ilist = [script.name for script in scripts.scripts_img2img.scripts if script.name is not None] control = [script.name for script in scripts.scripts_control.scripts if script.name is not None] return models.ResScripts(txt2img = t2ilist, img2img = i2ilist, control = control) + def get_script_info(script_name: Optional[str] = None): res = [] for script_list in [scripts.scripts_txt2img.scripts, scripts.scripts_img2img.scripts, scripts.scripts_control.scripts]: @@ -32,12 +44,16 @@ def get_script_info(script_name: Optional[str] = None): res.append(script.api_info) return res + def get_script(script_name, script_runner): - if script_name is None or script_name == "": + if script_name is None or script_name == "" or script_name == 'none': return None, None script_idx = script_name_to_index(script_name, script_runner.scripts) + if script_idx is None: + return None return script_runner.scripts[script_idx] + def init_default_script_args(script_runner): # find max idx from the scripts in runner and generate a none array to init script_args last_arg_index = 1 @@ -60,6 +76,7 @@ def init_default_script_args(script_runner): script_args[script.args_from:script.args_to] = ui_default_values return script_args + def init_script_args(p, request, default_script_args, selectable_scripts, selectable_script_idx, script_runner): script_args = default_script_args.copy() # position 0 in script_arg is the idx+1 of the selectable script that is going to be run when using scripts.scripts_*2img.run() diff --git a/modules/api/server.py b/modules/api/server.py index 95233dbcd..939e19c86 100644 --- a/modules/api/server.py +++ b/modules/api/server.py @@ -1,3 +1,4 @@ +import time from typing import Any, Dict from fastapi import Depends from modules import shared @@ -66,7 +67,6 @@ def get_cmd_flags(): return vars(shared.cmd_opts) def get_progress(req: models.ReqProgress = Depends()): - import time if shared.state.job_count == 0: return models.ResProgress(progress=0, eta_relative=0, state=shared.state.dict(), textinfo=shared.state.textinfo) shared.state.do_set_current_image() @@ -85,6 +85,9 @@ def get_progress(req: models.ReqProgress = Depends()): res = models.ResProgress(progress=progress, eta_relative=eta_relative, state=shared.state.dict(), current_image=current_image, textinfo=shared.state.textinfo) return res +def get_status(): + return shared.state.status() + def post_interrupt(): shared.state.interrupt() return {} 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 74bab35c4..5d6343c98 100644 --- a/modules/control/run.py +++ b/modules/control/run.py @@ -66,11 +66,11 @@ def control_run(state: str = '', resize_mode_before: int = 0, resize_name_before: str = 'None', resize_context_before: str = 'None', width_before: int = 512, height_before: int = 512, scale_by_before: float = 1.0, selected_scale_tab_before: int = 0, resize_mode_after: int = 0, resize_name_after: str = 'None', resize_context_after: str = 'None', width_after: int = 0, height_after: int = 0, scale_by_after: float = 1.0, selected_scale_tab_after: int = 0, resize_mode_mask: int = 0, resize_name_mask: str = 'None', resize_context_mask: str = 'None', width_mask: int = 0, height_mask: int = 0, scale_by_mask: float = 1.0, selected_scale_tab_mask: int = 0, - denoising_strength: float = 0, batch_count: int = 1, batch_size: int = 1, - enable_hr: bool = False, hr_sampler_index: int = None, hr_denoising_strength: float = 0.3, hr_resize_mode: int = 0, hr_resize_context: str = 'None', hr_upscaler: str = None, hr_force: bool = False, hr_second_pass_steps: int = 20, + denoising_strength: float = 0.3, batch_count: int = 1, batch_size: int = 1, + enable_hr: bool = False, hr_sampler_index: int = None, hr_denoising_strength: float = 0.0, hr_resize_mode: int = 0, hr_resize_context: str = 'None', hr_upscaler: str = None, hr_force: bool = False, hr_second_pass_steps: int = 20, hr_scale: float = 1.0, hr_resize_x: int = 0, hr_resize_y: int = 0, refiner_steps: int = 5, refiner_start: float = 0.0, refiner_prompt: str = '', refiner_negative: str = '', video_skip_frames: int = 0, video_type: str = 'None', video_duration: float = 2.0, video_loop: bool = False, video_pad: int = 0, video_interpolate: int = 0, - *input_script_args + *input_script_args, ): # handle optional initialization via ui for u in units: diff --git a/modules/errors.py b/modules/errors.py index c4d66c351..527884cf1 100644 --- a/modules/errors.py +++ b/modules/errors.py @@ -59,7 +59,7 @@ def exception(suppress=[]): console.print_exception(show_locals=False, max_frames=16, extra_lines=2, suppress=suppress, theme="ansi_dark", word_wrap=False, width=min([console.width, 200])) -def profile(profiler, msg: str): +def profile(profiler, msg: str, n: int = 5): profiler.disable() import io import pstats @@ -83,7 +83,7 @@ def profile(profiler, msg: str): and 'rich' not in x and x.strip() != '' ] - txt = '\n'.join(lines[:min(5, len(lines))]) + txt = '\n'.join(lines[:min(n, len(lines))]) log.debug(f'Profile {msg}: {txt}') diff --git a/modules/extra_networks.py b/modules/extra_networks.py index a574e8469..b464bd349 100644 --- a/modules/extra_networks.py +++ b/modules/extra_networks.py @@ -102,8 +102,9 @@ def activate(p, extra_network_data, step=0): except Exception as e: errors.display(e, f"Activating network: type={extra_network_name}") + p.extra_network_data = extra_network_data if stepwise: - p.extra_network_data = extra_network_data + p.stepwise_lora = True shared.opts.data['lora_functional'] = functional diff --git a/modules/face/__init__.py b/modules/face/__init__.py index aae3cdaa2..c18da6e2e 100644 --- a/modules/face/__init__.py +++ b/modules/face/__init__.py @@ -9,7 +9,7 @@ debug = shared.log.trace if os.environ.get('SD_FACE_DEBUG', None) is not None el class Script(scripts.Script): def title(self): - return 'Face' + return 'Face: Multiple ID Transfers' def show(self, is_img2img): return True if shared.native else False @@ -28,10 +28,10 @@ class Script(scripts.Script): elif hasattr(file, 'name'): image = Image.open(file.name) # _TemporaryFileWrapper from gr.Files else: - raise ValueError(f'PhotoMaker unknown input: {file}') + raise ValueError(f'Face: unknown input: {file}') init_images.append(image) except Exception as e: - shared.log.warning(f'PhotoMaker failed to load image: {e}') + shared.log.warning(f'Face: failed to load image: {e}') return init_images def mode_change(self, mode): @@ -45,7 +45,7 @@ class Script(scripts.Script): # return signature is array of gradio components def ui(self, _is_img2img): with gr.Row(): - gr.HTML("  Face module
") + gr.HTML("  Face: Multiple ID Transfers
") with gr.Row(): mode = gr.Dropdown(label='Mode', choices=['None', 'FaceID', 'FaceSwap', 'InstantID', 'PhotoMaker'], value='None') with gr.Group(visible=False) as cfg_faceid: diff --git a/modules/face/instantid.py b/modules/face/instantid.py index 9c7c16f61..662d17c7d 100644 --- a/modules/face/instantid.py +++ b/modules/face/instantid.py @@ -68,7 +68,7 @@ def instant_id(p: processing.StableDiffusionProcessing, app, source_images, stre processing.process_init(p) p.init(p.all_prompts, p.all_seeds, p.all_subseeds) orig_prompt_attention = shared.opts.prompt_attention - shared.opts.data['prompt_attention'] = 'Fixed attention' # otherwise need to deal with class_tokens_mask + shared.opts.data['prompt_attention'] = 'fixed' # otherwise need to deal with class_tokens_mask p.task_args['image_embeds'] = face_embeds[0].shape # placeholder p.task_args['image'] = face_images[0] p.task_args['controlnet_conditioning_scale'] = float(conditioning) diff --git a/modules/face/photomaker.py b/modules/face/photomaker.py index c8f58b42a..b89f28a10 100644 --- a/modules/face/photomaker.py +++ b/modules/face/photomaker.py @@ -49,7 +49,7 @@ def photo_maker(p: processing.StableDiffusionProcessing, input_images, trigger, shared.sd_model.to(dtype=devices.dtype) orig_prompt_attention = shared.opts.prompt_attention - shared.opts.data['prompt_attention'] = 'Fixed attention' # otherwise need to deal with class_tokens_mask + shared.opts.data['prompt_attention'] = 'fixed' # otherwise need to deal with class_tokens_mask p.task_args['input_id_images'] = input_images p.task_args['start_merge_step'] = int(start * p.steps) p.task_args['prompt'] = p.all_prompts[0] if p.all_prompts is not None else p.prompt diff --git a/modules/images.py b/modules/images.py index e6a334034..910349bef 100644 --- a/modules/images.py +++ b/modules/images.py @@ -40,8 +40,6 @@ def atomically_save_image(): except Exception: shared.log.warning(f'Save: unknown image format: {extension}') image_format = 'JPEG' - if shared.opts.image_watermark_enabled or (shared.opts.image_watermark_position != 'none' and shared.opts.image_watermark_image != ''): - image = set_watermark(image, shared.opts.image_watermark) exifinfo = (exifinfo or "") if shared.opts.image_metadata else "" # additional metadata saved in files if shared.opts.save_txt and len(exifinfo) > 0: @@ -153,6 +151,11 @@ def save_image(image, info = image.info.get(pnginfo_section_name, '') if info is not None: pnginfo[pnginfo_section_name] = info + + wm_text = getattr(p, 'watermark_text', shared.opts.image_watermark) + wm_image = getattr(p, 'watermark_image', shared.opts.image_watermark_image) + image = set_watermark(image, wm_text, wm_image) + params = script_callbacks.ImageSaveParams(image, p, filename, pnginfo) params.filename = namegen.sanitize(filename) dirname = os.path.dirname(params.filename) @@ -361,53 +364,54 @@ def flatten(img, bgcolor): return img.convert('RGB') -def draw_overlay(im, text): +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), text, font=font, fill=shared.opts.font_color) + 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: - if shared.opts.image_watermark_position == 'top/left': - position = (0, 0) - elif shared.opts.image_watermark_position == 'top/right': - position = (image.width - wm_image.width, 0) - elif shared.opts.image_watermark_position == 'bottom/left': - position = (0, image.height - wm_image.height) - elif shared.opts.image_watermark_position == 'bottom/right': - position = (image.width - wm_image.width, image.height - wm_image.height) - elif shared.opts.image_watermark_position == 'center': - position = ((image.width - wm_image.width) // 2, (image.height - wm_image.height) // 2) - else: - position = (random.randint(0, image.width - wm_image.width), random.randint(0, image.height - wm_image.height)) +def set_watermark(image, wm_text: str = None, wm_image: Image.Image = None): + if shared.opts.image_watermark_position != 'none' and wm_image is not None: # visible watermark + if isinstance(wm_image, str): try: - for x in range(wm_image.width): - for y in range(wm_image.height): - 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}') + wm_image = Image.open(wm_image) except Exception as e: shared.log.warning(f'Set image watermark: image={wm_image} {e}') + return image + if isinstance(wm_image, Image.Image): + if wm_image.mode != 'RGBA': + wm_image = wm_image.convert('RGBA') + if shared.opts.image_watermark_position == 'top/left': + position = (0, 0) + elif shared.opts.image_watermark_position == 'top/right': + position = (image.width - wm_image.width, 0) + elif shared.opts.image_watermark_position == 'bottom/left': + position = (0, image.height - wm_image.height) + elif shared.opts.image_watermark_position == 'bottom/right': + position = (image.width - wm_image.width, image.height - wm_image.height) + elif shared.opts.image_watermark_position == 'center': + position = ((image.width - wm_image.width) // 2, (image.height - wm_image.height) // 2) + else: + position = (random.randint(0, image.width - wm_image.width), random.randint(0, image.height - wm_image.height)) + try: + for x in range(wm_image.width): + for y in range(wm_image.height): + 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: image={wm_image} position={position}') + except Exception as e: + shared.log.warning(f'Set image watermark: image={wm_image} {e}') - if shared.opts.image_watermark_enabled: # invisible watermark + if shared.opts.image_watermark_enabled and wm_text is not None: # invisible watermark from imwatermark import WatermarkEncoder wm_type = 'bytes' wm_method = 'dwtDctSvd' @@ -416,16 +420,16 @@ def set_watermark(image, watermark): info = image.info data = np.asarray(image) encoder = WatermarkEncoder() - text = f"{watermark:<{length}}"[:length] + text = f"{wm_text:<{length}}"[:length] bytearr = text.encode(encoding='ascii', errors='ignore') try: encoder.set_watermark(wm_type, bytearr) encoded = encoder.encode(data, wm_method) image = Image.fromarray(encoded) image.info = info - shared.log.debug(f'Set invisible watermark: {watermark} method={wm_method} bits={wm_length}') + shared.log.debug(f'Set invisible watermark: {wm_text} method={wm_method} bits={wm_length}') except Exception as e: - shared.log.warning(f'Set invisible watermark error: {watermark} method={wm_method} bits={wm_length} {e}') + shared.log.warning(f'Set invisible watermark error: {wm_text} method={wm_method} bits={wm_length} {e}') return image diff --git a/modules/images_namegen.py b/modules/images_namegen.py index d88f85a77..bc58f728a 100644 --- a/modules/images_namegen.py +++ b/modules/images_namegen.py @@ -34,10 +34,10 @@ class FilenameGenerator: 'timestamp': lambda self: getattr(self.p, "job_timestamp", shared.state.job_timestamp), 'job_timestamp': lambda self: getattr(self.p, "job_timestamp", shared.state.job_timestamp), - 'model': lambda self: shared.sd_model.sd_checkpoint_info.title if shared.sd_loaded else '', - 'model_shortname': lambda self: shared.sd_model.sd_checkpoint_info.model_name if shared.sd_loaded else '', - 'model_name': lambda self: shared.sd_model.sd_checkpoint_info.model_name if shared.sd_loaded else '', - 'model_hash': lambda self: shared.sd_model.sd_checkpoint_info.shorthash if shared.sd_loaded else '', + 'model': lambda self: shared.sd_model.sd_checkpoint_info.title if shared.sd_loaded and getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None else '', + 'model_shortname': lambda self: shared.sd_model.sd_checkpoint_info.model_name if shared.sd_loaded and getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None else '', + 'model_name': lambda self: shared.sd_model.sd_checkpoint_info.model_name if shared.sd_loaded and getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None else '', + 'model_hash': lambda self: shared.sd_model.sd_checkpoint_info.shorthash if shared.sd_loaded and getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None else '', 'prompt': lambda self: self.prompt_full(), 'prompt_no_styles': lambda self: self.prompt_no_style(), diff --git a/modules/img2img.py b/modules/img2img.py index faf65161a..8274386cc 100644 --- a/modules/img2img.py +++ b/modules/img2img.py @@ -137,6 +137,7 @@ def img2img(id_task: str, state: str, mode: int, inpaint_full_res, inpaint_full_res_padding, inpainting_mask_invert, img2img_batch_files, img2img_batch_input_dir, img2img_batch_output_dir, img2img_batch_inpaint_mask_dir, hdr_mode, hdr_brightness, hdr_color, hdr_sharpen, hdr_clamp, hdr_boundary, hdr_threshold, hdr_maximize, hdr_max_center, hdr_max_boundry, hdr_color_picker, hdr_tint_ratio, + enable_hr, hr_sampler_index, hr_denoising_strength, hr_resize_mode, hr_resize_context, hr_upscaler, hr_force, hr_second_pass_steps, hr_scale, hr_resize_x, hr_resize_y, refiner_steps, hr_refiner_start, refiner_prompt, refiner_negative, override_settings_texts, *args): # pylint: disable=unused-argument @@ -214,7 +215,6 @@ def img2img(id_task: str, state: str, mode: int, subseed_strength=subseed_strength, seed_resize_from_h=seed_resize_from_h, seed_resize_from_w=seed_resize_from_w, - seed_enable_extras=True, sampler_name = processing.get_sampler_name(sampler_index, img=True), batch_size=batch_size, n_iter=n_iter, @@ -247,6 +247,23 @@ def img2img(id_task: str, state: str, mode: int, inpainting_mask_invert=inpainting_mask_invert, hdr_mode=hdr_mode, hdr_brightness=hdr_brightness, hdr_color=hdr_color, hdr_sharpen=hdr_sharpen, hdr_clamp=hdr_clamp, hdr_boundary=hdr_boundary, hdr_threshold=hdr_threshold, hdr_maximize=hdr_maximize, hdr_max_center=hdr_max_center, hdr_max_boundry=hdr_max_boundry, hdr_color_picker=hdr_color_picker, hdr_tint_ratio=hdr_tint_ratio, + # refiner + enable_hr=enable_hr, + hr_denoising_strength=hr_denoising_strength, + hr_scale=hr_scale, + hr_resize_mode=hr_resize_mode, + hr_resize_context=hr_resize_context, + hr_upscaler=hr_upscaler, + hr_force=hr_force, + hr_second_pass_steps=hr_second_pass_steps, + hr_resize_x=hr_resize_x, + hr_resize_y=hr_resize_y, + hr_sampler_name = processing.get_sampler_name(hr_sampler_index), + refiner_steps=refiner_steps, + hr_refiner_start=hr_refiner_start, + refiner_prompt=refiner_prompt, + refiner_negative=refiner_negative, + # override override_settings=override_settings, ) p.scripts = modules.scripts.scripts_img2img @@ -267,6 +284,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/loader.py b/modules/loader.py index 0711c2906..cd51cc8eb 100644 --- a/modules/loader.py +++ b/modules/loader.py @@ -126,4 +126,5 @@ except ImportError: except ImportError: pass # shrug... -errors.log.info(f'System packages: {get_packages()}') +errors.log.info(f'Torch: torch=={torch.__version__} torchvision=={torchvision.__version__}') +errors.log.info(f'Packages: diffusers=={diffusers.__version__} transformers=={transformers.__version__} accelerate=={accelerate.__version__} gradio=={gradio.__version__}') diff --git a/modules/model_quant.py b/modules/model_quant.py index 547a3d7ae..68bdfa7b2 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -32,14 +32,14 @@ def load_bnb(msg='', silent=False): global bnb # pylint: disable=global-statement if bnb is not None: return bnb - fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access - log.debug(f'Quantization: type=bitsandbytes fn={fn}') # pylint: disable=protected-access install('bitsandbytes', quiet=True) 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 + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + log.debug(f'Quantization: type=bitsandbytes version={bnb.__version__} fn={fn}') # pylint: disable=protected-access return bnb except Exception as e: if len(msg) > 0: @@ -54,12 +54,12 @@ def load_quanto(msg='', silent=False): global quanto # pylint: disable=global-statement if quanto is not None: return quanto - fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access - log.debug(f'Quantization: type=quanto fn={fn}') # pylint: disable=protected-access install('optimum-quanto', quiet=True) try: from optimum import quanto as optimum_quanto # pylint: disable=no-name-in-module quanto = optimum_quanto + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access + log.debug(f'Quantization: type=quanto version={quanto.__version__} fn={fn}') # pylint: disable=protected-access return quanto except Exception as e: if len(msg) > 0: diff --git a/modules/model_tools.py b/modules/model_tools.py index 1212da244..1d016a19e 100644 --- a/modules/model_tools.py +++ b/modules/model_tools.py @@ -5,13 +5,35 @@ 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 as e: - shared.log.error(f'Load dict: path="{filename}" {e}') + except Exception: + pass return keys diff --git a/modules/postprocess/yolo.py b/modules/postprocess/yolo.py index 316984207..f42b6bb9f 100644 --- a/modules/postprocess/yolo.py +++ b/modules/postprocess/yolo.py @@ -18,7 +18,7 @@ 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 @@ -29,6 +29,9 @@ class YoloResult: self.height = height self.args = args + def __str__(self): + return f'cls={self.cls} label={self.label} score={self.score} box={self.box} mask={self.mask} item={self.item} size={self.width}x{self.height} args={self.args}' + class YoloRestorer(Detailer): def __init__(self): @@ -76,11 +79,15 @@ class YoloRestorer(Detailer): offload: bool = shared.opts.detailer_unload, ) -> list[YoloResult]: + if model is None or (isinstance(model, str) and len(model) == 0): + model = 'yolo11m' result = [] if isinstance(model, str): - model = self.models.get(model, None) - if model is None: + cached = self.models.get(model, None) + if cached is None: _, model = self.load(model) + else: + model = cached if model is None: return result args = { @@ -136,7 +143,8 @@ class YoloRestorer(Detailer): 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, width=w, height=h, args=args)) + res = YoloResult(cls=cls, label=label, score=round(score, 2), box=box, mask=mask_image, item=cropped, width=w, height=h, args=args) + result.append(res) if len(result) >= shared.opts.detailer_max: break return result @@ -156,10 +164,10 @@ 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: @@ -194,7 +202,6 @@ class YoloRestorer(Detailer): shared.log.info(f'Detailer: model="{name}" no items detected') continue - 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] @@ -330,9 +337,9 @@ class YoloRestorer(Detailer): 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 = 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.1, maximum=1.0, step=0.05) - max_size = shared.opts.detailer_min_size if shared.opts.detailer_min_size < 1 and shared.opts.detailer_min_size > 0 else 1.0 - max_size = gr.Slider(label="Max size", elem_id=f"{tab}_detailer_max_size", value=max_size, minimum=0.1, maximum=1.0, step=0.05) + 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/processing.py b/modules/processing.py index 99d0cb351..0d557e64e 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -34,14 +34,14 @@ 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) @@ -80,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 = { diff --git a/modules/processing_args.py b/modules/processing_args.py index 1ea91fb08..ff766ec04 100644 --- a/modules/processing_args.py +++ b/modules/processing_args.py @@ -9,6 +9,7 @@ import numpy as np from modules import shared, errors, sd_models, processing, processing_vae, processing_helpers, sd_hijack_hypertile, prompt_parser_diffusers, timer from modules.processing_callbacks import diffusers_callback_legacy, diffusers_callback, set_callbacks_p from modules.processing_helpers import resize_hires, fix_prompts, calculate_base_steps, calculate_hires_steps, calculate_refiner_steps, get_generator, set_latents, apply_circular # pylint: disable=unused-import +from modules.api import helpers debug = shared.log.trace if os.environ.get('SD_DIFFUSERS_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -18,7 +19,9 @@ def task_specific_kwargs(p, model): task_args = {} is_img2img_model = bool('Zero123' in shared.sd_model.__class__.__name__) if len(getattr(p, 'init_images', [])) > 0: - p.init_images = [p.convert('RGB') for p in p.init_images] + if isinstance(p.init_images[0], str): + p.init_images = [helpers.decode_base64_to_image(i, quiet=True) for i in p.init_images] + p.init_images = [i.convert('RGB') if i.mode != 'RGB' else i for i in p.init_images] if sd_models.get_diffusers_task(model) == sd_models.DiffusersTaskType.TEXT_2_IMAGE or len(getattr(p, 'init_images', [])) == 0 and not is_img2img_model: p.ops.append('txt2img') if hasattr(p, 'width') and hasattr(p, 'height'): @@ -27,7 +30,7 @@ def task_specific_kwargs(p, model): 'height': 8 * math.ceil(p.height / 8), } elif (sd_models.get_diffusers_task(model) == sd_models.DiffusersTaskType.IMAGE_2_IMAGE or is_img2img_model) and len(getattr(p, 'init_images', [])) > 0: - if shared.sd_model_type == 'sdxl': + if shared.sd_model_type == 'sdxl' and hasattr(model, 'register_to_config'): model.register_to_config(requires_aesthetics_score = False) p.ops.append('img2img') task_args = { @@ -55,7 +58,7 @@ def task_specific_kwargs(p, model): 'strength': p.denoising_strength, } elif (sd_models.get_diffusers_task(model) == sd_models.DiffusersTaskType.INPAINTING or is_img2img_model) and len(getattr(p, 'init_images', [])) > 0: - if shared.sd_model_type == 'sdxl': + if shared.sd_model_type == 'sdxl' and hasattr(model, 'register_to_config'): model.register_to_config(requires_aesthetics_score = False) if p.detailer: p.ops.append('detailer') @@ -100,62 +103,64 @@ 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' steps = kwargs.get("num_inference_steps", None) or len(getattr(p, 'timesteps', ['1'])) clip_skip = kwargs.pop("clip_skip", 1) - # prompt_parser_diffusers.fix_position_ids(model) - if shared.opts.prompt_attention != 'Fixed attention' and 'Onnx' not in model.__class__.__name__ and ( + parser = 'fixed' + if shared.opts.prompt_attention != 'fixed' and 'Onnx' not in model.__class__.__name__ and ( 'StableDiffusion' in model.__class__.__name__ or 'StableCascade' in model.__class__.__name__ or 'Flux' in model.__class__.__name__ ): try: - prompt_parser_diffusers.encode_prompts(model, p, prompts, negative_prompts, steps=steps, clip_skip=clip_skip) + prompt_parser_diffusers.embedder = prompt_parser_diffusers.PromptEmbedder(prompts, negative_prompts, steps, clip_skip, p) parser = shared.opts.prompt_attention except Exception as e: shared.log.error(f'Prompt parser encode: {e}') if os.environ.get('SD_PROMPT_DEBUG', None) is not None: errors.display(e, 'Prompt parser encode') timer.process.record('encode', reset=False) + else: + prompt_parser_diffusers.embedder = None 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 len(p.prompt_embeds) > 0 and p.prompt_embeds[0] is not None: - args['prompt_embeds'] = p.prompt_embeds[0] + if hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'prompt_embeds' in possible and prompt_parser_diffusers.embedder 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'] = p.positive_pooleds[0].unsqueeze(0) - elif 'XL' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0: - args['pooled_prompt_embeds'] = p.positive_pooleds[0] - elif 'StableDiffusion3' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0: - args['pooled_prompt_embeds'] = p.positive_pooleds[0] - elif 'Flux' in model.__class__.__name__ and len(getattr(p, 'positive_pooleds', [])) > 0: - args['pooled_prompt_embeds'] = p.positive_pooleds[0] + args['prompt_embeds_pooled'] = prompt_parser_diffusers.embedder('positive_pooleds').unsqueeze(0) + elif 'XL' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None: + args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds') + elif 'StableDiffusion3' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None: + args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds') + elif 'Flux' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None: + args['pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('positive_pooleds') else: args['prompt'] = prompts if 'negative_prompt' in possible: - if hasattr(model, 'text_encoder') and 'negative_prompt_embeds' in possible and len(p.negative_embeds) > 0 and p.negative_embeds[0] is not None: - args['negative_prompt_embeds'] = p.negative_embeds[0] - if 'StableCascade' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0: - args['negative_prompt_embeds_pooled'] = p.negative_pooleds[0].unsqueeze(0) - if 'XL' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0: - args['negative_pooled_prompt_embeds'] = p.negative_pooleds[0] - if 'StableDiffusion3' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0: - args['negative_pooled_prompt_embeds'] = p.negative_pooleds[0] + 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) + if 'XL' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None: + args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds') + if 'StableDiffusion3' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None: + args['negative_pooled_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_pooleds') else: if 'PixArtSigmaPipeline' in model.__class__.__name__: # pixart-sigma pipeline throws list-of-list for negative prompt args['negative_prompt'] = negative_prompts[0] else: args['negative_prompt'] = negative_prompts - if 'clip_skip' in possible and parser == 'Fixed attention': + if 'clip_skip' in possible and parser == 'fixed': if clip_skip == 1: pass # clip_skip = None else: @@ -180,6 +185,8 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2 if hasattr(model, 'scheduler') and hasattr(model.scheduler, 'noise_sampler_seed') and hasattr(model.scheduler, 'noise_sampler'): model.scheduler.noise_sampler = None # noise needs to be reset instead of using cached values model.scheduler.noise_sampler_seed = p.seeds # some schedulers have internal noise generator and do not use pipeline generator + if 'seed' in possible: + args['seed'] = p.seed if 'noise_sampler_seed' in possible: args['noise_sampler_seed'] = p.seeds if 'guidance_scale' in possible: diff --git a/modules/processing_callbacks.py b/modules/processing_callbacks.py index 47c8e8827..52ea3e575 100644 --- a/modules/processing_callbacks.py +++ b/modules/processing_callbacks.py @@ -3,8 +3,7 @@ import os import time import torch import numpy as np -from modules import shared, processing_correction, extra_networks, timer - +from modules import shared, processing_correction, extra_networks, timer, prompt_parser_diffusers p = None debug_callback = shared.log.trace if os.environ.get('SD_CALLBACK_DEBUG', None) is not None else lambda *args, **kwargs: None @@ -14,6 +13,19 @@ def set_callbacks_p(processing): global p # pylint: disable=global-statement p = processing +def prompt_callback(step, kwargs): + if prompt_parser_diffusers.embedder is None or 'prompt_embeds' not in kwargs: + return kwargs + try: + prompt_embeds = prompt_parser_diffusers.embedder('prompt_embeds', step + 1) + negative_prompt_embeds = prompt_parser_diffusers.embedder('negative_prompt_embeds', step + 1) + if p.cfg_scale > 1: # Perform guidance + prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0) # Combined embeds + assert prompt_embeds.shape == kwargs['prompt_embeds'].shape, f"prompt_embed shape mismatch {kwargs['prompt_embeds'].shape} {prompt_embeds.shape}" + kwargs['prompt_embeds'] = prompt_embeds + except Exception as e: + debug_callback(f"Callback: {e}") + return kwargs def diffusers_callback_legacy(step: int, timestep: int, latents: typing.Union[torch.FloatTensor, np.ndarray]): if p is None: @@ -33,7 +45,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 @@ -49,7 +61,7 @@ def diffusers_callback(pipe, step: int, timestep: int, kwargs: dict): if shared.state.interrupted or shared.state.skipped: raise AssertionError('Interrupted...') time.sleep(0.1) - if hasattr(p, "extra_network_data"): + if hasattr(p, "stepwise_lora"): extra_networks.activate(p, p.extra_network_data, step=step) if latents is None: return kwargs @@ -67,14 +79,7 @@ def diffusers_callback(pipe, step: int, timestep: int, kwargs: dict): pipe.set_ip_adapter_scale(ip_adapter_scales) if step != getattr(pipe, 'num_timesteps', 0): kwargs = processing_correction.correction_callback(p, timestep, kwargs) - if p.scheduled_prompt and 'prompt_embeds' in kwargs and 'negative_prompt_embeds' in kwargs: - try: - i = (step + 1) % len(p.prompt_embeds) - kwargs["prompt_embeds"] = p.prompt_embeds[i][0:1].expand(kwargs["prompt_embeds"].shape) - j = (step + 1) % len(p.negative_embeds) - kwargs["negative_prompt_embeds"] = p.negative_embeds[j][0:1].expand(kwargs["negative_prompt_embeds"].shape) - except Exception as e: - shared.log.debug(f"Callback: {e}") + kwargs = prompt_callback(step, kwargs) # monkey patch for diffusers callback issues if step == int(getattr(pipe, 'num_timesteps', 100) * p.cfg_end) and 'prompt_embeds' in kwargs and 'negative_prompt_embeds' in kwargs: if "PAG" in shared.sd_model.__class__.__name__: pipe._guidance_scale = 1.001 if pipe._guidance_scale > 1 else pipe._guidance_scale # pylint: disable=protected-access diff --git a/modules/processing_class.py b/modules/processing_class.py index 9265ea3cf..79f51576f 100644 --- a/modules/processing_class.py +++ b/modules/processing_class.py @@ -17,53 +17,44 @@ debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None @dataclass(repr=False) class StableDiffusionProcessing: - """ - The first set of paramaters: sd_models -> do_not_reload_embeddings represent the minimum required to create a StableDiffusionProcessing - """ def __init__(self, - sd_model=None, - outpath_samples=None, - outpath_grids=None, + sd_model=None, # pylint: disable=unused-argument # local instance of sd_model + # base params prompt: str = "", - styles: List[str] = None, + negative_prompt: str = "", seed: int = -1, subseed: int = -1, subseed_strength: float = 0, seed_resize_from_h: int = -1, seed_resize_from_w: int = -1, - seed_enable_extras: bool = True, - sampler_name: str = None, - hr_sampler_name: str = None, batch_size: int = 1, n_iter: int = 1, steps: int = 50, - cfg_scale: float = 7.0, - image_cfg_scale: float = None, clip_skip: int = 1, width: int = 512, height: int = 512, - full_quality: bool = True, - detailer: bool = False, - restore_faces: bool = False, - tiling: bool = False, - hidiffusion: bool = False, - do_not_save_samples: bool = False, - do_not_save_grid: bool = False, - extra_generation_params: Dict[Any, Any] = None, - overlay_images: Any = None, - negative_prompt: str = None, + # samplers + sampler_index: int = None, # pylint: disable=unused-argument # used only to set sampler_name + sampler_name: str = None, + hr_sampler_name: str = None, eta: float = None, - do_not_reload_embeddings: bool = False, - denoising_strength: float = 0, + # guidance + cfg_scale: float = 7.0, + cfg_end: float = 1, diffusers_guidance_rescale: float = 0.7, pag_scale: float = 0.0, pag_adaptive: float = 0.5, - cfg_end: float = 1, - resize_mode: int = 0, - resize_name: str = 'None', - resize_context: str = 'None', - scale_by: float = 0, - selected_scale_tab: int = 0, + # styles + styles: List[str] = [], + # vae + tiling: bool = False, + full_quality: bool = True, + # other + hidiffusion: bool = False, + do_not_reload_embeddings: bool = False, + detailer: bool = False, + restore_faces: bool = False, + # hdr corrections hdr_mode: int = 0, hdr_brightness: float = 0, hdr_color: float = 0, @@ -76,92 +67,196 @@ class StableDiffusionProcessing: hdr_max_boundry: float = 1.0, hdr_color_picker: str = None, hdr_tint_ratio: float = 0, - override_settings: Dict[str, Any] = None, + # img2img + init_images: list = None, + resize_mode: int = 0, + resize_name: str = 'None', + resize_context: str = 'None', + denoising_strength: float = 0.3, + image_cfg_scale: float = None, + initial_noise_multiplier: float = None, # pylint: disable=unused-argument # a1111 compatibility + scale_by: float = 1, + selected_scale_tab: int = 0, # pylint: disable=unused-argument # a1111 compatibility + # inpaint + mask: Any = None, + latent_mask: Any = None, + mask_for_overlay: Any = None, + mask_blur: int = 4, + paste_to: Any = None, + inpainting_fill: int = 0, + inpaint_full_res: bool = False, + inpaint_full_res_padding: int = 0, + inpainting_mask_invert: int = 0, + overlay_images: Any = None, + # refiner + enable_hr: bool = False, + firstphase_width: int = 0, + firstphase_height: int = 0, + hr_scale: float = 2.0, + hr_force: bool = False, + hr_resize_mode: int = 0, + hr_resize_context: str = 'None', + hr_upscaler: str = None, + hr_second_pass_steps: int = 0, + hr_resize_x: int = 0, + hr_resize_y: int = 0, + hr_denoising_strength: float = 0.0, + refiner_steps: int = 5, + refiner_start: float = 0, + refiner_prompt: str = '', + refiner_negative: str = '', + hr_refiner_start: float = 0, + # save options + outpath_samples=None, + outpath_grids=None, + do_not_save_samples: bool = False, + do_not_save_grid: bool = False, + # scripts + script_args: list = [], + # overrides + override_settings: Dict[str, Any] = {}, override_settings_restore_afterwards: bool = True, - sampler_index: int = None, - script_args: list = None - ): # pylint: disable=unused-argument + # metadata + extra_generation_params: Dict[Any, Any] = {}, + ): + # extra args set by processing loop + self.task_args = {} + + # state items self.state: str = '' + self.ops = [] self.skip = [] - self.outpath_samples: str = outpath_samples - self.outpath_grids: str = outpath_grids - self.prompt: str = prompt - self.prompt_for_display: str = None - self.negative_prompt: str = (negative_prompt or "") - self.styles: list = styles or [] - self.seed: int = seed - self.subseed: int = subseed - self.subseed_strength: float = subseed_strength - self.seed_resize_from_h: int = seed_resize_from_h - self.seed_resize_from_w: int = seed_resize_from_w - self.sampler_name: str = sampler_name - self.hr_sampler_name: str = hr_sampler_name if hr_sampler_name != 'Same as primary' else sampler_name - self.batch_size: int = batch_size - self.n_iter: int = n_iter - self.steps: int = steps - self.hr_second_pass_steps = 0 - self.cfg_scale: float = cfg_scale - self.scale_by: float = scale_by + self.color_corrections = [] + self.is_control = False + self.is_hr_pass = False + self.is_refiner_pass = False + self.is_api = False + self.scheduled_prompt = False + self.prompt_embeds = [] + self.positive_pooleds = [] + self.negative_embeds = [] + self.negative_pooleds = [] + self.disable_extra_networks = False + self.iteration = 0 + + # initializers + self.prompt = prompt + self.seed = seed + self.subseed = subseed + self.subseed_strength = subseed_strength + self.seed_resize_from_h = seed_resize_from_h + self.seed_resize_from_w = seed_resize_from_w + self.batch_size = batch_size + self.n_iter = n_iter + self.steps = steps + self.clip_skip = clip_skip + self.width = width + self.height = height + self.negative_prompt = negative_prompt + self.styles = styles + self.tiling = tiling + self.full_quality = full_quality + self.hidiffusion = hidiffusion + self.do_not_reload_embeddings = do_not_reload_embeddings + self.detailer = detailer + self.restore_faces = restore_faces + self.init_images = init_images + self.resize_mode = resize_mode + self.resize_name = resize_name + self.resize_context = resize_context + self.denoising_strength = denoising_strength self.image_cfg_scale = image_cfg_scale + self.scale_by = scale_by + self.mask = mask + self.image_mask = mask # TODO duplciate mask params + self.latent_mask = latent_mask + self.mask_blur = mask_blur + self.inpainting_fill = inpainting_fill + self.inpaint_full_res_padding = inpaint_full_res_padding + self.inpainting_mask_invert = inpainting_mask_invert + self.overlay_images = overlay_images + self.enable_hr = enable_hr + self.firstphase_width = firstphase_width + self.firstphase_height = firstphase_height + self.hr_scale = hr_scale + self.hr_force = hr_force + self.hr_resize_mode = hr_resize_mode + self.hr_resize_context = hr_resize_context + self.hr_upscaler = hr_upscaler + self.hr_second_pass_steps = hr_second_pass_steps + self.hr_resize_x = hr_resize_x + self.hr_resize_y = hr_resize_y + self.hr_upscale_to_x = hr_resize_x + self.hr_upscale_to_y = hr_resize_y + self.hr_denoising_strength = hr_denoising_strength + self.refiner_steps = refiner_steps + self.refiner_start = refiner_start + self.refiner_prompt = refiner_prompt + self.refiner_negative = refiner_negative + self.hr_refiner_start = hr_refiner_start + self.outpath_samples = outpath_samples + self.outpath_grids = outpath_grids + self.do_not_save_samples = do_not_save_samples + self.do_not_save_grid = do_not_save_grid + self.override_settings_restore_afterwards = override_settings_restore_afterwards + self.extra_generation_params = extra_generation_params + self.eta = eta + self.cfg_scale = cfg_scale + self.cfg_end = cfg_end self.diffusers_guidance_rescale = diffusers_guidance_rescale self.pag_scale = pag_scale self.pag_adaptive = pag_adaptive - self.cfg_end = cfg_end - self.width: int = width - self.height: int = height - self.full_quality: bool = full_quality - self.detailer: bool = detailer - self.restore_faces: bool = restore_faces - self.tiling: bool = tiling - self.hidiffusion: bool = hidiffusion - self.do_not_save_samples: bool = do_not_save_samples - self.do_not_save_grid: bool = do_not_save_grid - self.extra_generation_params: dict = extra_generation_params or {} - self.overlay_images = overlay_images - self.eta = eta - self.do_not_reload_embeddings = do_not_reload_embeddings - self.paste_to = None - self.color_corrections = None - self.denoising_strength: float = denoising_strength + self.selected_scale_tab = selected_scale_tab + self.mask_for_overlay = mask_for_overlay + self.paste_to = paste_to + self.init_latent = None + + # special handled items + if firstphase_width != 0 or firstphase_height != 0: + self.hr_upscale_to_x = self.width + self.hr_upscale_to_y = self.height + self.width = firstphase_width + self.height = firstphase_height + self.sampler_name = sampler_name or processing_helpers.get_sampler_name(sampler_index, img=True) + self.hr_sampler_name: str = hr_sampler_name if hr_sampler_name != 'Same as primary' else self.sampler_name self.override_settings = {k: v for k, v in (override_settings or {}).items() if k not in shared.restricted_opts} - self.override_settings_restore_afterwards = override_settings_restore_afterwards - self.is_using_inpainting_conditioning = False # a111 compatibility - self.disable_extra_networks = False - # self.scripts = scripts.ScriptRunner() # set via property - # self.script_args = script_args or [] # set via property - self.per_script_args = {} + self.inpaint_full_res = inpaint_full_res if isinstance(inpaint_full_res, bool) else self.inpaint_full_res + self.inpaint_full_res = inpaint_full_res != 0 if isinstance(inpaint_full_res, int) else self.inpaint_full_res + + # null items initialized later self.all_prompts = None self.all_negative_prompts = None self.all_seeds = None self.all_subseeds = None - self.clip_skip = clip_skip + + # a1111 compatibility items shared.opts.data['clip_skip'] = int(self.clip_skip) # for compatibility with a1111 sd_hijack_clip - self.iteration = 0 - self.is_control = False - self.is_hr_pass = False - self.is_refiner_pass = False - self.hr_force = False - self.enable_hr = None - self.hr_scale = None - self.hr_upscaler = None - self.hr_resize_mode = 0 - self.hr_resize_context = 'None' - self.hr_resize_x = 0 - self.hr_resize_y = 0 - self.hr_upscale_to_x = 0 - self.hr_upscale_to_y = 0 + self.seed_enable_extras: bool = True + self.is_using_inpainting_conditioning = False # a111 compatibility + self.batch_index = 0 + self.refiner_switch_at = 0 + self.hr_prompt = '' + self.all_hr_prompts = [] + self.hr_negative_prompt = '' + self.all_hr_negative_prompts = [] self.truncate_x = 0 self.truncate_y = 0 - self.applied_old_hires_behavior_to = None - self.refiner_steps = 5 - self.refiner_start = 0 - self.refiner_prompt = '' - self.refiner_negative = '' - self.ops = [] - self.resize_mode: int = resize_mode - self.resize_name: str = resize_name - self.resize_context: str = resize_context + self.comments = {} + self.sampler = None + self.nmask = None + self.initial_noise_multiplier = initial_noise_multiplier or shared.opts.initial_noise_multiplier + self.image_conditioning = None + self.prompt_for_display: str = None + + # scripts + self.scripts_value: scripts.ScriptRunner = field(default=None, init=False) + self.script_args_value: list = field(default=None, init=False) + self.scripts_setup_complete: bool = field(default=False, init=False) + self.script_args = script_args + self.per_script_args = {} + + # settings to processing self.ddim_discretize = shared.opts.ddim_discretize self.s_min_uncond = shared.opts.s_min_uncond self.s_churn = shared.opts.s_churn @@ -171,18 +266,7 @@ class StableDiffusionProcessing: self.s_tmin = shared.opts.s_tmin self.s_tmax = float('inf') # not representable as a standard ui option self.task_args = {} - # a1111 compatibility items - self.batch_index = 0 - self.refiner_switch_at = 0 - self.hr_prompt = '' - self.all_hr_prompts = [] - self.hr_negative_prompt = '' - self.all_hr_negative_prompts = [] - self.comments = {} - self.is_api = False - self.scripts_value: scripts.ScriptRunner = field(default=None, init=False) - self.script_args_value: list = field(default=None, init=False) - self.scripts_setup_complete: bool = field(default=False, init=False) + # ip adapter self.ip_adapter_names = [] self.ip_adapter_scales = [0.0] @@ -190,6 +274,7 @@ class StableDiffusionProcessing: self.ip_adapter_starts = [0.0] self.ip_adapter_ends = [1.0] self.ip_adapter_crops = [] + # hdr self.hdr_mode=hdr_mode self.hdr_brightness=hdr_brightness @@ -203,7 +288,10 @@ class StableDiffusionProcessing: self.hdr_max_boundry=hdr_max_boundry self.hdr_color_picker=hdr_color_picker self.hdr_tint_ratio=hdr_tint_ratio + # globals + self.embedder = None + self.override = None self.scheduled_prompt: bool = False self.prompt_embeds = [] self.positive_pooleds = [] @@ -252,57 +340,9 @@ class StableDiffusionProcessing: class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing): - - def __init__(self, - enable_hr: bool = False, - denoising_strength: float = 0.75, - firstphase_width: int = 0, - firstphase_height: int = 0, - hr_scale: float = 2.0, - hr_force: bool = False, - hr_resize_mode: int = 0, - hr_resize_context: str = 'None', - hr_upscaler: str = None, - hr_second_pass_steps: int = 0, - hr_resize_x: int = 0, - hr_resize_y: int = 0, - refiner_steps: int = 5, - refiner_start: float = 0, - refiner_prompt: str = '', - refiner_negative: str = '', - **kwargs - ): - + def __init__(self, **kwargs): + debug(f'Process init: mode={self.__class__.__name__} kwargs={kwargs}') # pylint: disable=protected-access super().__init__(**kwargs) - self.reprocess = {} - self.enable_hr = enable_hr - self.denoising_strength = denoising_strength - self.hr_scale = hr_scale - self.hr_upscaler = hr_upscaler - self.hr_resize_mode = hr_resize_mode - self.hr_resize_context = hr_resize_context - self.hr_force = hr_force - self.hr_second_pass_steps = hr_second_pass_steps - self.hr_resize_x = hr_resize_x - self.hr_resize_y = hr_resize_y - self.hr_upscale_to_x = hr_resize_x - self.hr_upscale_to_y = hr_resize_y - if firstphase_width != 0 or firstphase_height != 0: - self.hr_upscale_to_x = self.width - self.hr_upscale_to_y = self.height - self.width = firstphase_width - self.height = firstphase_height - self.truncate_x = 0 - self.truncate_y = 0 - self.applied_old_hires_behavior_to = None - self.refiner_steps = refiner_steps - self.refiner_start = refiner_start - self.refiner_prompt = refiner_prompt - self.refiner_negative = refiner_negative - self.sampler = None - self.scripts = None - self.script_args = [] - def init(self, all_prompts=None, all_seeds=None, all_subseeds=None): if shared.native: @@ -360,41 +400,9 @@ class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing): class StableDiffusionProcessingImg2Img(StableDiffusionProcessing): - - def __init__(self, init_images: list = None, resize_mode: int = 0, resize_name: str = 'None', resize_context: str = 'None', denoising_strength: float = 0.3, image_cfg_scale: float = None, mask: Any = None, mask_blur: int = 4, inpainting_fill: int = 0, inpaint_full_res: bool = False, inpaint_full_res_padding: int = 0, inpainting_mask_invert: int = 0, initial_noise_multiplier: float = None, scale_by: float = 1, refiner_steps: int = 5, refiner_start: float = 0, refiner_prompt: str = '', refiner_negative: str = '', **kwargs): + def __init__(self, **kwargs): + debug(f'Process init: mode={self.__class__.__name__} kwargs={kwargs}') # pylint: disable=protected-access super().__init__(**kwargs) - self.init_images = init_images - self.resize_mode: int = resize_mode - self.resize_name: str = resize_name - self.resize_context: str = resize_context - self.denoising_strength: float = denoising_strength - self.hr_denoising_strength: float = denoising_strength - self.image_cfg_scale: float = image_cfg_scale - self.init_latent = None - self.image_mask = mask - self.latent_mask = None - self.mask_for_overlay = None - self.mask_blur_x = mask_blur # a1111 compatibility item - self.mask_blur_y = mask_blur # a1111 compatibility item - self.mask_blur = mask_blur - self.inpainting_fill = inpainting_fill - self.inpaint_full_res = inpaint_full_res - self.inpaint_full_res_padding = inpaint_full_res_padding - self.inpainting_mask_invert = inpainting_mask_invert - self.initial_noise_multiplier = shared.opts.initial_noise_multiplier if initial_noise_multiplier is None else initial_noise_multiplier - self.mask = None - self.nmask = None - self.image_conditioning = None - self.refiner_steps = refiner_steps - self.refiner_start = refiner_start - self.refiner_prompt = refiner_prompt - self.refiner_negative = refiner_negative - self.enable_hr = None - self.is_batch = False - self.scale_by = scale_by - self.sampler = None - self.scripts = None - self.script_args = [] def init(self, all_prompts=None, all_seeds=None, all_subseeds=None): if hasattr(self, 'init_images') and self.init_images is not None and len(self.init_images) > 0: @@ -485,7 +493,7 @@ class StableDiffusionProcessingImg2Img(StableDiffusionProcessing): image = images.resize_image(self.resize_mode, image, self.width, self.height, upscaler_name=self.resize_name, context=self.resize_context) self.width = image.width self.height = image.height - if self.image_mask is not None and shared.opts.mask_apply_overlay: + if self.image_mask is not None and shared.opts.mask_apply_overlay and not hasattr(self, 'xyz'): image_masked = Image.new('RGBa', (image.width, image.height)) image_to_paste = image.convert("RGBA").convert("RGBa") image_to_mask = ImageOps.invert(self.mask_for_overlay.convert('L')) if self.mask_for_overlay is not None else None @@ -544,47 +552,8 @@ class StableDiffusionProcessingImg2Img(StableDiffusionProcessing): class StableDiffusionProcessingControl(StableDiffusionProcessingImg2Img): def __init__(self, **kwargs): + debug(f'Process init: mode={self.__class__.__name__} kwargs={kwargs}') # pylint: disable=protected-access super().__init__(**kwargs) - self.strength = None - self.adapter_conditioning_scale = None - self.adapter_conditioning_factor = None - self.guess_mode = None - self.controlnet_conditioning_scale = None - self.control_guidance_start = None - self.control_guidance_end = None - self.control_mode = None - self.reference_attn = None - self.reference_adain = None - self.attention_auto_machine_weight = None - self.gn_auto_machine_weight = None - self.style_fidelity = None - self.ref_image = None - self.image = None - self.query_weight = None - self.adain_weight = None - self.adapter_conditioning_factor = 1.0 - self.attention = 'Attention' - self.fidelity = 0.5 - self.mask_image = None - self.override = None - self.resize_mode_before = None - self.resize_name_before = None - self.width_before = None - self.height_before = None - self.scale_by_before = None - self.selected_scale_tab_before = None - self.resize_mode_after = None - self.resize_name_after = None - self.width_after = None - self.height_after = None - self.scale_by_after = None - self.selected_scale_tab_after = None - self.resize_mode_mask = None - self.resize_name_mask = None - self.width_mask = None - self.height_mask = None - self.scale_by_mask = None - self.selected_scale_tab_mask = None def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts): # abstract pass diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 12fa4bc53..2164134b1 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -71,7 +71,7 @@ def process_base(p: processing.StableDiffusionProcessing): guidance_rescale=p.diffusers_guidance_rescale, denoising_start=0 if use_refiner_start else p.refiner_start if use_denoise_start else None, denoising_end=p.refiner_start if use_refiner_start else 1 if use_denoise_start else None, - output_type='latent' if hasattr(shared.sd_model, 'vae') else 'np', + output_type='latent', clip_skip=p.clip_skip, desc='Base', ) @@ -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 @@ -151,13 +152,14 @@ def process_hires(p: processing.StableDiffusionProcessing, output): p.is_hr_pass = True if hasattr(p, 'init_hr'): p.init_hr(p.hr_scale, p.hr_upscaler, force=p.hr_force) - else: # fake hires for img2img - p.hr_scale = p.scale_by - p.hr_upscaler = p.resize_name - p.hr_resize_mode = p.resize_mode - p.hr_resize_context = p.resize_context - p.hr_upscale_to_x = p.width - p.hr_upscale_to_y = p.height + else: + if not p.is_hr_pass: # fake hires for img2img if not actual hr pass + p.hr_scale = p.scale_by + p.hr_upscaler = p.resize_name + p.hr_resize_mode = p.resize_mode + p.hr_resize_context = p.resize_context + p.hr_upscale_to_x = p.width * p.hr_scale if p.hr_resize_x == 0 else p.hr_resize_x + p.hr_upscale_to_y = p.height * p.hr_scale if p.hr_resize_y == 0 else p.hr_resize_y prev_job = shared.state.job # hires runs on original pipeline @@ -175,7 +177,8 @@ def process_hires(p: processing.StableDiffusionProcessing, output): sd_hijack_hypertile.hypertile_set(p, hr=True) latent_upscale = shared.latent_upscale_modes.get(p.hr_upscaler, None) - if (latent_upscale is not None or p.hr_force) and getattr(p, 'hr_denoising_strength', p.denoising_strength) > 0: + strength = p.hr_denoising_strength if p.hr_denoising_strength > 0 else p.denoising_strength + if (latent_upscale is not None or p.hr_force) and strength > 0: p.ops.append('hires') sd_models_compile.openvino_recompile_model(p, hires=True, refiner=False) if shared.sd_model.__class__.__name__ == "OnnxRawPipeline": @@ -183,8 +186,7 @@ def process_hires(p: processing.StableDiffusionProcessing, output): p.hr_force = True # hires - p.denoising_strength = getattr(p, 'hr_denoising_strength', p.denoising_strength) - if p.hr_force and p.denoising_strength == 0: + if p.hr_force and strength == 0: shared.log.warning('HiRes skip: denoising=0') p.hr_force = False if p.hr_force: @@ -202,9 +204,9 @@ def process_hires(p: processing.StableDiffusionProcessing, output): sd_models.move_model(shared.sd_model.unet, devices.device) if hasattr(shared.sd_model, 'transformer'): sd_models.move_model(shared.sd_model.transformer, devices.device) - orig_denoise = p.denoising_strength - p.denoising_strength = getattr(p, 'hr_denoising_strength', p.denoising_strength) update_sampler(p, shared.sd_model, second_pass=True) + orig_denoise = p.denoising_strength + p.denoising_strength = strength hires_args = set_pipeline_args( p=p, model=shared.sd_model, @@ -216,10 +218,10 @@ def process_hires(p: processing.StableDiffusionProcessing, output): eta=shared.opts.scheduler_eta, guidance_scale=p.image_cfg_scale if p.image_cfg_scale is not None else p.cfg_scale, guidance_rescale=p.diffusers_guidance_rescale, - output_type='latent' if hasattr(shared.sd_model, 'vae') else 'np', + output_type='latent', clip_skip=p.clip_skip, image=output.images, - strength=p.denoising_strength, + strength=strength, desc='Hires', ) shared.state.job = 'HiRes' @@ -277,7 +279,7 @@ def process_refine(p: processing.StableDiffusionProcessing, output): for i in range(len(output.images)): image = output.images[i] noise_level = round(350 * p.denoising_strength) - output_type='latent' if hasattr(shared.sd_refiner, 'vae') else 'np' + output_type='latent' if 'Upscale' in shared.sd_refiner.__class__.__name__ or 'Flux' in shared.sd_refiner.__class__.__name__: image = processing_vae.vae_decode(latents=image, model=shared.sd_model, full_quality=p.full_quality, output_type='pil', width=p.width, height=p.height) p.extra_generation_params['Noise level'] = noise_level @@ -345,7 +347,11 @@ def process_decode(p: processing.StableDiffusionProcessing, output): if not hasattr(output, 'images') and hasattr(output, 'frames'): shared.log.debug(f'Generated: frames={len(output.frames[0])}') output.images = output.frames[0] - if hasattr(shared.sd_model, "vae") and output.images is not None and len(output.images) > 0: + model = shared.sd_model if not is_refiner_enabled(p) else shared.sd_refiner + if not hasattr(model, 'vae'): + if hasattr(model, 'pipe') and hasattr(model.pipe, 'vae'): + model = model.pipe + if hasattr(model, "vae") and output.images is not None and len(output.images) > 0: if p.hr_resize_mode > 0 and (p.hr_upscaler != 'None' or p.hr_resize_mode == 5): width = max(getattr(p, 'width', 0), getattr(p, 'hr_upscale_to_x', 0)) height = max(getattr(p, 'height', 0), getattr(p, 'hr_upscale_to_y', 0)) @@ -354,7 +360,7 @@ def process_decode(p: processing.StableDiffusionProcessing, output): height = getattr(p, 'height', 0) results = processing_vae.vae_decode( latents = output.images, - model = shared.sd_model if not is_refiner_enabled(p) else shared.sd_refiner, + model = model, full_quality = p.full_quality, width = width, height = height, diff --git a/modules/processing_helpers.py b/modules/processing_helpers.py index ef83834e5..ec7fbf048 100644 --- a/modules/processing_helpers.py +++ b/modules/processing_helpers.py @@ -47,16 +47,23 @@ def apply_overlay(image: Image, paste_loc, index, overlays): return image debug(f'Apply overlay: image={image} loc={paste_loc} index={index} overlays={overlays}') overlay = overlays[index] - if paste_loc is not None: - x, y, w, h = paste_loc - if image.width != w or image.height != h or x != 0 or y != 0: - base_image = Image.new('RGBA', (overlay.width, overlay.height)) - image = images.resize_image(2, image, w, h) - base_image.paste(image, (x, y)) - image = base_image - image = image.convert('RGBA') - image.alpha_composite(overlay) - image = image.convert('RGB') + if not isinstance(image, Image.Image) or not isinstance(overlay, Image.Image): + return image + try: + if paste_loc is not None and (isinstance(paste_loc, tuple) or isinstance(paste_loc, list)): + x, y, w, h = paste_loc + if x is None or y is None or w is None or h is None: + return image + if image.width != w or image.height != h or x != 0 or y != 0: + base_image = Image.new('RGBA', (overlay.width, overlay.height)) + image = images.resize_image(2, image, w, h) + base_image.paste(image, (x, y)) + image = base_image + image = image.convert('RGBA') + image.alpha_composite(overlay) + image = image.convert('RGB') + except Exception as e: + shared.log.error(f'Apply overlay: {e}') return image diff --git a/modules/processing_info.py b/modules/processing_info.py index e798211b1..714ebf35f 100644 --- a/modules/processing_info.py +++ b/modules/processing_info.py @@ -4,6 +4,7 @@ from modules import shared, sd_samplers_common, sd_vae, generation_parameters_co from modules.processing_class import StableDiffusionProcessing +debug = shared.log.trace if os.environ.get('SD_PROCESS_DEBUG', None) is not None else lambda *args, **kwargs: None if not shared.native: from modules import sd_hijack else: @@ -39,30 +40,34 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No ops.reverse() args = { # basic + "Size": f"{p.width}x{p.height}" if hasattr(p, 'width') and hasattr(p, 'height') else None, + "Sampler": p.sampler_name if p.sampler_name != 'Default' else None, "Steps": p.steps, "Seed": all_seeds[index], - "Sampler": p.sampler_name if p.sampler_name != 'Default' else None, + "Seed resize from": None if p.seed_resize_from_w == 0 or p.seed_resize_from_h == 0 else f"{p.seed_resize_from_w}x{p.seed_resize_from_h}", "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, + "Clip skip": p.clip_skip if p.clip_skip > 1 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.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', - "Seed resize from": None if p.seed_resize_from_w == 0 or p.seed_resize_from_h == 0 else f"{p.seed_resize_from_w}x{p.seed_resize_from_h}", - "Clip skip": p.clip_skip if p.clip_skip > 1 else None, "Prompt2": p.refiner_prompt if len(p.refiner_prompt) > 0 else None, "Negative2": p.refiner_negative if len(p.refiner_negative) > 0 else None, "Styles": "; ".join(p.styles) if p.styles is not None and len(p.styles) > 0 else None, - "Tiling": p.tiling if p.tiling else None, # sdnext - "Backend": 'Diffusers' if shared.native else 'Original', "App": 'SD.Next', "Version": git_commit, + "Backend": 'Diffusers' if shared.native else 'Original', + "Pipeline": 'LDM', + "Parser": shared.opts.prompt_attention.split()[0], "Comment": comment, "Operations": '; '.join(ops).replace('"', '') if len(p.ops) > 0 else 'none', } + if shared.opts.add_model_name_to_info and getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None: + args["Model"] = shared.sd_model.sd_checkpoint_info.model_name.replace(',', '').replace(':', '') + if shared.opts.add_model_hash_to_info and getattr(shared.sd_model, 'sd_model_hash', None) is not None: + args["Model hash"] = shared.sd_model.sd_model_hash # native if grid is None and (p.n_iter > 1 or p.batch_size > 1) and index >= 0: args['Index'] = f'{p.iteration + 1}x{index + 1}' @@ -165,7 +170,9 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No if isinstance(v, str): if len(v) == 0 or v == '0x0': del args[k] + debug(f'Infotext: args={args}') params_text = ", ".join([k if k == v else f'{k}: {generation_parameters_copypaste.quote(v)}' for k, v in args.items()]) negative_prompt_text = f"\nNegative prompt: {all_negative_prompts[index]}" if all_negative_prompts[index] else "" infotext = f"{all_prompts[index]}{negative_prompt_text}\n{params_text}".strip() + debug(f'Infotext: "{infotext}"') return infotext diff --git a/modules/processing_vae.py b/modules/processing_vae.py index 75347f416..3c0357c81 100644 --- a/modules/processing_vae.py +++ b/modules/processing_vae.py @@ -35,7 +35,9 @@ def create_latents(image, p, dtype=None, device=None): def full_vae_decode(latents, model): t0 = time.time() - if not hasattr(model, 'vae'): + if not hasattr(model, 'vae') and hasattr(model, 'pipe'): + model = model.pipe + if model is None or not hasattr(model, 'vae'): shared.log.error('VAE not found in model') return [] if debug: @@ -147,6 +149,9 @@ def taesd_vae_encode(image): def vae_decode(latents, model, output_type='np', full_quality=True, width=None, height=None): t0 = time.time() + model = model or shared.sd_model + if not hasattr(model, 'vae') and hasattr(model, 'pipe'): + model = model.pipe if latents is None or not torch.is_tensor(latents): # already decoded return latents prev_job = shared.state.job @@ -169,8 +174,8 @@ def vae_decode(latents, model, output_type='np', full_quality=True, width=None, if latents.shape[-1] <= 4: # not a latent, likely an image decoded = latents.float().cpu().numpy() - elif full_quality and hasattr(shared.sd_model, "vae"): - decoded = full_vae_decode(latents=latents, model=shared.sd_model) + elif full_quality and hasattr(model, "vae"): + decoded = full_vae_decode(latents=latents, model=model) else: decoded = taesd_vae_decode(latents=latents) @@ -195,6 +200,8 @@ def vae_decode(latents, model, output_type='np', full_quality=True, width=None, def vae_encode(image, model, full_quality=True): # pylint: disable=unused-variable if shared.state.interrupted or shared.state.skipped: return [] + if not hasattr(model, 'vae') and hasattr(model, 'pipe'): + model = model.pipe if not hasattr(model, 'vae'): shared.log.error('VAE not found in model') return [] diff --git a/modules/progress.py b/modules/progress.py index abd6d906d..d18d1ee9f 100644 --- a/modules/progress.py +++ b/modules/progress.py @@ -73,7 +73,6 @@ def progressapi(req: ProgressRequest): elapsed = time.time() - shared.state.time_start if shared.state.time_start is not None else 0 predicted = elapsed / progress if progress > 0 else None eta = predicted - elapsed if predicted is not None else None - # shared.log.debug(f'Progress: step={step_x}:{step_y} batch={batch_x}:{batch_y} current={current} total={total} progress={progress} elapsed={elapsed} eta={eta}') id_live_preview = req.id_live_preview live_preview = None shared.state.set_current_image() diff --git a/modules/prompt_parser.py b/modules/prompt_parser.py index 2a71d5053..3a1288097 100644 --- a/modules/prompt_parser.py +++ b/modules/prompt_parser.py @@ -308,11 +308,11 @@ def parse_prompt_attention(text): res = [] round_brackets = [] square_brackets = [] - if opts.prompt_attention == 'Fixed attention': + if opts.prompt_attention == 'fixed': res = [[text, 1.0]] debug(f'Prompt: parser="{opts.prompt_attention}" {res}') return res - elif opts.prompt_attention == 'Compel parser': + elif opts.prompt_attention == 'compel': conjunction = Compel.parse_prompt_string(text) if conjunction is None or conjunction.prompts is None or conjunction.prompts is None or len(conjunction.prompts[0].children) == 0: return [["", 1.0]] @@ -321,7 +321,7 @@ def parse_prompt_attention(text): res.append([frag.text, frag.weight]) debug(f'Prompt: parser="{opts.prompt_attention}" {res}') return res - elif opts.prompt_attention == 'A1111 parser': + elif opts.prompt_attention == 'a1111': re_attention = re_attention_v1 whitespace = '' else: @@ -360,7 +360,7 @@ def parse_prompt_attention(text): for i, part in enumerate(parts): if i > 0: res.append(["BREAK", -1]) - if opts.prompt_attention == 'Full parser': + if opts.prompt_attention == 'native': part = re_clean.sub("", part) part = re_whitespace.sub(" ", part).strip() if len(part) == 0: @@ -392,15 +392,15 @@ if __name__ == "__main__": log.info(f'Schedules: {all_schedules}') for schedule in all_schedules: log.info(f'Schedule: {schedule[0]}') - opts.data['prompt_attention'] = 'Fixed attention' + opts.data['prompt_attention'] = 'fixed' output_list = parse_prompt_attention(schedule[1]) log.info(f' Fixed: {output_list}') - opts.data['prompt_attention'] = 'Compel parser' + opts.data['prompt_attention'] = 'compel' output_list = parse_prompt_attention(schedule[1]) log.info(f' Compel: {output_list}') - opts.data['prompt_attention'] = 'A1111 parser' + opts.data['prompt_attention'] = 'a1111' output_list = parse_prompt_attention(schedule[1]) log.info(f' A1111: {output_list}') - opts.data['prompt_attention'] = 'Full parser' + opts.data['prompt_attention'] = 'native' log.info = parse_prompt_attention(schedule[1]) log.info(f' Full: {output_list}') diff --git a/modules/prompt_parser_diffusers.py b/modules/prompt_parser_diffusers.py index cc814f379..234272907 100644 --- a/modules/prompt_parser_diffusers.py +++ b/modules/prompt_parser_diffusers.py @@ -2,6 +2,7 @@ import os import math import time import typing +from collections import OrderedDict import torch from compel.embeddings_provider import BaseTextualInversionManager, EmbeddingsProvider from transformers import PreTrainedTokenizer @@ -14,7 +15,182 @@ debug('Trace: PROMPT') orig_encode_token_ids_to_embeddings = EmbeddingsProvider._encode_token_ids_to_embeddings # pylint: disable=protected-access token_dict = None # used by helper get_tokens token_type = None # used by helper get_tokens -cache = {} +cache = OrderedDict() +embedder = None + + +def prompt_compatible(pipe = None): + pipe = pipe or shared.sd_model + if ( + 'StableDiffusion' not in pipe.__class__.__name__ and + 'DemoFusion' not in pipe.__class__.__name__ and + 'StableCascade' not in pipe.__class__.__name__ and + 'Flux' not in pipe.__class__.__name__ + ): + shared.log.warning(f"Prompt parser not supported: {pipe.__class__.__name__}") + return False + return True + + +def prepare_model(pipe = None): + pipe = pipe or shared.sd_model + if not hasattr(pipe, "text_encoder") and hasattr(shared.sd_model, "pipe"): + pipe = pipe.pipe + if not hasattr(pipe, "text_encoder"): + return None + if shared.opts.diffusers_offload_mode == "balanced": + pipe = sd_models.apply_balanced_offload(pipe) + elif hasattr(pipe, "maybe_free_model_hooks"): + pipe.maybe_free_model_hooks() + devices.torch_gc() + return pipe + + +class PromptEmbedder: + def __init__(self, prompts, negative_prompts, steps, clip_skip, p): + t0 = time.time() + self.prompts = prompts + self.negative_prompts = negative_prompts + self.batchsize = len(self.prompts) + self.attention = None + self.allsame = self.compare_prompts() # collapses batched prompts to single prompt if possible + self.steps = steps + self.clip_skip = clip_skip + # All embeds are nested lists, outer list batch length, inner schedule length + self.prompt_embeds = [[]] * self.batchsize + self.positive_pooleds = [[]] * self.batchsize + self.negative_prompt_embeds = [[]] * self.batchsize + self.negative_pooleds = [[]] * self.batchsize + self.positive_schedule = None + self.negative_schedule = None + self.scheduled_prompt = False + earlyout = self.checkcache(p) + if earlyout: + return + pipe = prepare_model(p.sd_model) + if pipe is None: + shared.log.error("Prompt encode: cannot find text encoder in model") + return + # per prompt in batch + for batchidx, (prompt, negative_prompt) in enumerate(zip(self.prompts, self.negative_prompts)): + self.prepare_schedule(prompt, negative_prompt) + if self.scheduled_prompt: + self.scheduled_encode(pipe, batchidx) + else: + self.encode(pipe, prompt, negative_prompt, batchidx) + self.checkcache(p) + debug(f"Prompt encode: time={(time.time() - t0):.3f}") + + def checkcache(self, p): + if shared.opts.sd_textencoder_cache_size == 0: + return False + if self.attention != shared.opts.prompt_attention: + debug(f"Prompt change: parser={shared.opts.prompt_attention}") + cache.clear() + return False + + def flatten(xss): + return [x for xs in xss for x in xs] + + # unpack EN data in case of TE LoRA + en_data = p.extra_network_data + en_data = [idx.items for item in en_data.values() for idx in item] + effective_batch = 1 if self.allsame else self.batchsize + key = str([self.prompts, self.negative_prompts, effective_batch, self.clip_skip, self.steps, en_data]) + item = cache.get(key) + if not item: + if not any(flatten(emb) for emb in [self.prompt_embeds, + self.negative_prompt_embeds, + self.positive_pooleds, + self.negative_pooleds]): + return False + else: + cache[key] = {'prompt_embeds': self.prompt_embeds, + 'negative_prompt_embeds': self.negative_prompt_embeds, + 'positive_pooleds': self.positive_pooleds, + 'negative_pooleds': self.negative_pooleds, + } + debug(f"Prompt cache: add={key}") + while len(cache) > int(shared.opts.sd_textencoder_cache_size): + cache.popitem(last=False) + if item: + self.__dict__.update(cache[key]) + cache.move_to_end(key) + if self.allsame and len(self.prompt_embeds) < self.batchsize: + self.prompt_embeds = [self.prompt_embeds[0]] * self.batchsize + self.positive_pooleds = [self.positive_pooleds[0]] * self.batchsize + self.negative_prompt_embeds = [self.negative_prompt_embeds[0]] * self.batchsize + self.negative_pooleds = [self.negative_pooleds[0]] * self.batchsize + debug(f"Prompt cache: get={key}") + return True + + def compare_prompts(self): + same = (self.prompts == [self.prompts[0]] * len(self.prompts) and self.negative_prompts == [self.negative_prompts[0]] * len(self.negative_prompts)) + if same: + self.prompts = [self.prompts[0]] + self.negative_prompts = [self.negative_prompts[0]] + return same + + def prepare_schedule(self, prompt, negative_prompt): + self.positive_schedule, scheduled = get_prompt_schedule(prompt, self.steps) + self.negative_schedule, neg_scheduled = get_prompt_schedule(negative_prompt, self.steps) + self.scheduled_prompt = scheduled or neg_scheduled + debug(f"Prompt schedule: positive={self.positive_schedule} negative={self.negative_schedule} scheduled={scheduled}") + + def scheduled_encode(self, pipe, batchidx): + prompt_dict = {} # index cache + for i in range(max(len(self.positive_schedule), len(self.negative_schedule))): + positive_prompt = self.positive_schedule[i % len(self.positive_schedule)] + negative_prompt = self.negative_schedule[i % len(self.negative_schedule)] + # skip repeated scheduled subprompts + idx = prompt_dict.get(positive_prompt+negative_prompt) + if idx is not None: + self.extend_embeds(batchidx, idx) + continue + self.encode(pipe, positive_prompt, negative_prompt, batchidx) + prompt_dict[positive_prompt+negative_prompt] = i + + def extend_embeds(self, batchidx, idx): # Extends scheduled prompt via index + if len(self.prompt_embeds[batchidx]) > 0: + self.prompt_embeds[batchidx].append(self.prompt_embeds[batchidx][idx]) + if len(self.negative_prompt_embeds[batchidx]) > 0: + self.negative_prompt_embeds[batchidx].append(self.negative_prompt_embeds[batchidx][idx]) + if len(self.positive_pooleds[batchidx]) > 0: + self.positive_pooleds[batchidx].append(self.positive_pooleds[batchidx][idx]) + if len(self.negative_pooleds[batchidx]) > 0: + self.negative_pooleds[batchidx].append(self.negative_pooleds[batchidx][idx]) + + def encode(self, pipe, positive_prompt, negative_prompt, batchidx): + self.attention = shared.opts.prompt_attention + if self.attention == "xhinker" or 'Flux' in pipe.__class__.__name__: + prompt_embed, positive_pooled, negative_embed, negative_pooled = get_xhinker_text_embeddings(pipe, positive_prompt, negative_prompt, self.clip_skip) + else: + prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, positive_prompt, negative_prompt, self.clip_skip) + if prompt_embed is not None: + self.prompt_embeds[batchidx].append(prompt_embed) + if negative_embed is not None: + self.negative_prompt_embeds[batchidx].append(negative_embed) + if positive_pooled is not None: + self.positive_pooleds[batchidx].append(positive_pooled) + if negative_pooled is not None: + self.negative_pooleds[batchidx].append(negative_pooled) + + if debug_enabled: + get_tokens(pipe, 'positive', positive_prompt) + get_tokens(pipe, 'negative', negative_prompt) + pipe = prepare_model() + + def __call__(self, key, step=0): + batch = getattr(self, key) + res = [] + for i in range(self.batchsize): + if len(batch[i]) == 0: # if asking for a null key, ie pooled on SD1.5 + return None + try: + res.append(batch[i][step]) + except IndexError: + res.append(batch[i][0]) # if not scheduled, return default + return torch.cat(res) def compel_hijack(self, token_ids: torch.Tensor, attention_mask: typing.Optional[torch.Tensor] = None) -> torch.Tensor: @@ -59,9 +235,9 @@ def insert_parser_highjack(pipename): debug("Load Standard Parser hijack") - insert_parser_highjack("Initialize") + # from https://github.com/damian0815/compel/blob/main/src/compel/diffusers_textual_inversion_manager.py class DiffusersTextualInversionManager(BaseTextualInversionManager): def __init__(self, pipe, tokenizer): @@ -108,12 +284,6 @@ class DiffusersTextualInversionManager(BaseTextualInversionManager): def get_prompt_schedule(prompt, steps): t0 = time.time() - if shared.native: - # TODO prompt scheduling - # prompt schedule returns array of prompts which would require that each prompt is fed to the model per-step - # prompt scheduling should instead interpolate between each prompt in schedule - # this temporarily disables prompt scheduling - return [prompt], False temp = [] schedule = prompt_parser.get_learned_conditioning_prompt_schedules([prompt], steps)[0] if all(x == schedule[0] for x in schedule): @@ -126,25 +296,25 @@ def get_prompt_schedule(prompt, steps): return temp, len(schedule) > 1 -def get_tokens(msg, prompt): +def get_tokens(pipe, msg, prompt): global token_dict, token_type # pylint: disable=global-statement if not shared.native: - return - if shared.sd_loaded and hasattr(shared.sd_model, 'tokenizer') and shared.sd_model.tokenizer is not None: + return 0 + if shared.sd_loaded and hasattr(pipe, 'tokenizer') and pipe.tokenizer is not None: if token_dict is None or token_type != shared.sd_model_type: token_type = shared.sd_model_type - fn = shared.sd_model.tokenizer.name_or_path + fn = pipe.tokenizer.name_or_path if fn.endswith('tokenizer'): - fn = os.path.join(shared.sd_model.tokenizer.name_or_path, 'vocab.json') + fn = os.path.join(pipe.tokenizer.name_or_path, 'vocab.json') else: - fn = os.path.join(shared.sd_model.tokenizer.name_or_path, 'tokenizer', 'vocab.json') + fn = os.path.join(pipe.tokenizer.name_or_path, 'tokenizer', 'vocab.json') token_dict = shared.readfile(fn, silent=True) - for k, v in shared.sd_model.tokenizer.added_tokens_decoder.items(): + for k, v in pipe.tokenizer.added_tokens_decoder.items(): token_dict[str(v)] = k shared.log.debug(f'Tokenizer: words={len(token_dict)} file="{fn}"') - has_bos_token = shared.sd_model.tokenizer.bos_token_id is not None - has_eos_token = shared.sd_model.tokenizer.eos_token_id is not None - ids = shared.sd_model.tokenizer(prompt) + has_bos_token = pipe.tokenizer.bos_token_id is not None + has_eos_token = pipe.tokenizer.eos_token_id is not None + ids = pipe.tokenizer(prompt) ids = getattr(ids, 'input_ids', []) tokens = [] for i in ids: @@ -155,118 +325,7 @@ def get_tokens(msg, prompt): tokens.append(f'UNK_{i}') token_count = len(ids) - int(has_bos_token) - int(has_eos_token) debug(f'Prompt tokenizer: type={msg} tokens={token_count} {tokens}') - - -def encode_prompts(pipe, p, prompts: list, negative_prompts: list, steps: int, clip_skip: typing.Optional[int] = None): - params_match = prompts == cache.get('prompts', None) and negative_prompts == cache.get('negative_prompts', None) and clip_skip == cache.get('clip_skip', None) and steps == cache.get('steps', None) - if ( - 'StableDiffusion' not in pipe.__class__.__name__ and - 'DemoFusion' not in pipe.__class__.__name__ and - 'StableCascade' not in pipe.__class__.__name__ and - 'Flux' not in pipe.__class__.__name__ - ): - shared.log.warning(f"Prompt parser not supported: {pipe.__class__.__name__}") - return - elif shared.opts.sd_textencoder_cache and cache.get('model_type', None) == shared.sd_model_type and params_match: - p.prompt_embeds = cache.get('prompt_embeds', None) - p.positive_pooleds = cache.get('positive_pooleds', None) - p.negative_embeds = cache.get('negative_embeds', None) - p.negative_pooleds = cache.get('negative_pooleds', None) - p.scheduled_prompt = cache.get('scheduled_prompt', None) - debug("Prompt encode: cached") - return - else: - t0 = time.time() - if shared.opts.diffusers_offload_mode == "balanced": - pipe = sd_models.apply_balanced_offload(pipe) - elif hasattr(pipe, "maybe_free_model_hooks"): - pipe.maybe_free_model_hooks() - devices.torch_gc() - - prompt_embeds, positive_pooleds, negative_embeds, negative_pooleds = [], [], [], [] - last_prompt, last_negative = None, None - for prompt, negative in zip(prompts, negative_prompts): - prompt_embed, positive_pooled, negative_embed, negative_pooled = None, None, None, None - if last_prompt == prompt and last_negative == negative: - prompt_embeds.append(prompt_embeds[-1]) - negative_embeds.append(negative_embeds[-1]) - if len(positive_pooleds) > 0: - positive_pooleds.append(positive_pooleds[-1]) - if len(negative_pooleds) > 0: - negative_pooleds.append(negative_pooleds[-1]) - continue - positive_schedule, scheduled = get_prompt_schedule(prompt, steps) - negative_schedule, neg_scheduled = get_prompt_schedule(negative, steps) - p.scheduled_prompt = scheduled or neg_scheduled - p.prompt_embeds = [] - p.positive_pooleds = [] - p.negative_embeds = [] - p.negative_pooleds = [] - - for i in range(max(len(positive_schedule), len(negative_schedule))): - positive_prompt = positive_schedule[i % len(positive_schedule)] - negative_prompt = negative_schedule[i % len(negative_schedule)] - if shared.opts.prompt_attention == "xhinker parser" or 'Flux' in pipe.__class__.__name__: - prompt_embed, positive_pooled, negative_embed, negative_pooled = get_xhinker_text_embeddings(pipe, positive_prompt, negative_prompt, clip_skip) - else: - prompt_embed, positive_pooled, negative_embed, negative_pooled = get_weighted_text_embeddings(pipe, positive_prompt, negative_prompt, clip_skip) - if prompt_embed is not None: - prompt_embeds.append(prompt_embed) - if negative_embed is not None: - negative_embeds.append(negative_embed) - if positive_pooled is not None: - positive_pooleds.append(positive_pooled) - if negative_pooled is not None: - negative_pooleds.append(negative_pooled) - last_prompt, last_negative = prompt, negative - # TODO prompt scheduling - # interpolation should happen here and then we can re-enable prompt scheduling - # ive tried simple torch.mean and its not good-enough - - def fix_length(embeds): - max_len = max([e.shape[1] for e in embeds if e is not None]) - for i, e in enumerate(embeds): - if e is not None and e.shape[1] < max_len: - expanded = torch.zeros((e.shape[0], max_len, e.shape[2]), device=e.device, dtype=e.dtype) - expanded[:, :e.shape[1], :] = e - embeds[i] = expanded - return torch.cat(embeds, dim=0).to(devices.device, dtype=devices.dtype) - - if len(prompt_embeds) > 0: - p.prompt_embeds.append(fix_length(prompt_embeds)) - if len(negative_embeds) > 0: - p.negative_embeds.append(fix_length(negative_embeds)) - if len(positive_pooleds) > 0: - p.positive_pooleds.append(fix_length(positive_pooleds)) - if len(negative_pooleds) > 0: - p.negative_pooleds.append(fix_length(negative_pooleds)) - - if shared.opts.sd_textencoder_cache and p.batch_size == 1: - cache.update({ - 'prompt_embeds': p.prompt_embeds, - 'negative_embeds': p.negative_embeds, - 'positive_pooleds': p.positive_pooleds, - 'negative_pooleds': p.negative_pooleds, - 'scheduled_prompt': p.scheduled_prompt, - 'prompts': prompts, - 'negative_prompts': negative_prompts, - 'clip_skip': clip_skip, - 'steps': steps, - 'model_type': shared.sd_model_type - }) - else: - cache.clear() - if debug_enabled: - get_tokens('positive', prompts[0]) - get_tokens('negative', negative_prompts[0]) - if shared.opts.diffusers_offload_mode == "balanced": - pipe = sd_models.apply_balanced_offload(pipe) - elif hasattr(pipe, "maybe_free_model_hooks"): - # text encoder will stay in the vram and cause oom, send everything back to cpu before continuing - pipe.maybe_free_model_hooks() - debug(f"Prompt encode: time={(time.time() - t0):.3f}") - devices.torch_gc() - return + return token_count def normalize_prompt(pairs: list): @@ -286,14 +345,20 @@ def normalize_prompt(pairs: list): return pairs -def get_prompts_with_weights(prompt: str): +def get_prompts_with_weights(pipe, prompt: str): t0 = time.time() - manager = DiffusersTextualInversionManager(shared.sd_model, shared.sd_model.tokenizer or shared.sd_model.tokenizer_2) - prompt = manager.maybe_convert_prompt(prompt, shared.sd_model.tokenizer or shared.sd_model.tokenizer_2) + manager = DiffusersTextualInversionManager(pipe, pipe.tokenizer or pipe.tokenizer_2) + prompt = manager.maybe_convert_prompt(prompt, pipe.tokenizer or pipe.tokenizer_2) texts_and_weights = prompt_parser.parse_prompt_attention(prompt) if shared.opts.prompt_mean_norm: texts_and_weights = normalize_prompt(texts_and_weights) texts, text_weights = zip(*texts_and_weights) + if debug_enabled: + all_tokens = 0 + for text in texts: + tokens = get_tokens(pipe, 'section', text) + all_tokens += tokens + debug(f'Prompt tokenizer: parser={shared.opts.prompt_attention} tokens={all_tokens}') debug(f'Prompt: weights={texts_and_weights} time={(time.time() - t0):.3f}') return texts, text_weights @@ -354,7 +419,8 @@ def pad_to_same_length(pipe, embeds, empty_embedding_providers=None): embeds[i] = embed return embeds -def split_prompts(prompt, SD3 = False): + +def split_prompts(pipe, prompt, SD3 = False): if prompt.find("TE2:") != -1: prompt, prompt2 = prompt.split("TE2:") else: @@ -372,7 +438,7 @@ def split_prompts(prompt, SD3 = False): prompt3 = " " if prompt3.strip() == "" else prompt3.strip() if SD3 and prompt3 != " ": - ps, _ws = get_prompts_with_weights(prompt3) + ps, _ws = get_prompts_with_weights(pipe, prompt3) prompt3 = " ".join(ps) return prompt, prompt2, prompt3 @@ -380,15 +446,15 @@ def split_prompts(prompt, SD3 = False): def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None): device = devices.device SD3 = hasattr(pipe, 'text_encoder_3') - prompt, prompt_2, prompt_3 = split_prompts(prompt, SD3) - neg_prompt, neg_prompt_2, neg_prompt_3 = split_prompts(neg_prompt, SD3) + prompt, prompt_2, prompt_3 = split_prompts(pipe, prompt, SD3) + neg_prompt, neg_prompt_2, neg_prompt_3 = split_prompts(pipe, neg_prompt, SD3) if prompt != prompt_2: - ps = [get_prompts_with_weights(p) for p in [prompt, prompt_2]] - ns = [get_prompts_with_weights(p) for p in [neg_prompt, neg_prompt_2]] + ps = [get_prompts_with_weights(pipe, p) for p in [prompt, prompt_2]] + ns = [get_prompts_with_weights(pipe, p) for p in [neg_prompt, neg_prompt_2]] else: - ps = 2 * [get_prompts_with_weights(prompt)] - ns = 2 * [get_prompts_with_weights(neg_prompt)] + ps = 2 * [get_prompts_with_weights(pipe, prompt)] + ns = 2 * [get_prompts_with_weights(pipe, neg_prompt)] positives, positive_weights = zip(*ps) negatives, negative_weights = zip(*ns) @@ -434,7 +500,7 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c # negative prompt has no keywords embed, ntokens = embedding_providers[i].get_embeddings_for_weighted_prompt_fragments(text_batch=[negatives[i]], fragment_weights_batch=[negative_weights[i]], device=device, should_return_tokens=True) negative_prompt_embeds.append(embed) - debug(f'Prompt: unpadded shape={prompt_embeds[0].shape} TE{i+1} ptokens={torch.count_nonzero(ptokens)} ntokens={torch.count_nonzero(ntokens)} time={(time.time() - t0):.3f}') + debug(f'Prompt: unpadded={prompt_embeds[0].shape} TE{i+1} ptokens={torch.count_nonzero(ptokens)} ntokens={torch.count_nonzero(ntokens)} time={(time.time() - t0):.3f}') if SD3: t0 = time.time() pooled_prompt_embeds.append(embedding_providers[0].get_pooled_embeddings(texts=positives[0] if len(positives[0]) == 1 else [" ".join(positives[0])], device=device)) @@ -443,7 +509,7 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c negative_pooled_prompt_embeds.append(embedding_providers[1].get_pooled_embeddings(texts=negatives[-1] if len(negatives[-1]) == 1 else [" ".join(negatives[-1])], device=device)) pooled_prompt_embeds = torch.cat(pooled_prompt_embeds, dim=-1) negative_pooled_prompt_embeds = torch.cat(negative_pooled_prompt_embeds, dim=-1) - debug(f'Prompt: pooled shape={pooled_prompt_embeds[0].shape} time={(time.time() - t0):.3f}') + debug(f'Prompt: pooled={pooled_prompt_embeds[0].shape} time={(time.time() - t0):.3f}') elif prompt_embeds[-1].shape[-1] > 768: t0 = time.time() if shared.opts.diffusers_pooled == "weighted": @@ -503,8 +569,8 @@ def get_weighted_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", c def get_xhinker_text_embeddings(pipe, prompt: str = "", neg_prompt: str = "", clip_skip: int = None): is_sd3 = hasattr(pipe, 'text_encoder_3') - prompt, prompt_2, _prompt_3 = split_prompts(prompt, is_sd3) - neg_prompt, neg_prompt_2, _neg_prompt_3 = split_prompts(neg_prompt, is_sd3) + prompt, prompt_2, _prompt_3 = split_prompts(pipe, prompt, is_sd3) + neg_prompt, neg_prompt_2, _neg_prompt_3 = split_prompts(pipe, neg_prompt, is_sd3) try: prompt = pipe.maybe_convert_prompt(prompt, pipe.tokenizer) neg_prompt = pipe.maybe_convert_prompt(neg_prompt, pipe.tokenizer) diff --git a/modules/prompt_parser_xhinker.py b/modules/prompt_parser_xhinker.py index 6a8acf8c6..c0ddc9bc7 100644 --- a/modules/prompt_parser_xhinker.py +++ b/modules/prompt_parser_xhinker.py @@ -1305,7 +1305,7 @@ def get_weighted_text_embeddings_sd3( # ---------------------- get neg t5 embeddings ------------------------- neg_prompt_tokens_3 = torch.tensor([neg_prompt_tokens_3], dtype=torch.long) - t5_neg_prompt_embeds = pipe.text_encoder_3(neg_prompt_tokens_3.to(pipe.pipe.text_encoder_3.device))[0].squeeze(0) + t5_neg_prompt_embeds = pipe.text_encoder_3(neg_prompt_tokens_3.to(pipe.text_encoder_3.device))[0].squeeze(0) t5_neg_prompt_embeds = t5_neg_prompt_embeds.to(device=pipe.text_encoder_3.device) # add weight to neg t5 embeddings diff --git a/modules/pulid/__init__.py b/modules/pulid/__init__.py new file mode 100644 index 000000000..dcee2d7b9 --- /dev/null +++ b/modules/pulid/__init__.py @@ -0,0 +1,11 @@ +""" +Credit and original implementation: +""" + +import os +import sys +sys.path.append(os.path.dirname(__file__)) +from pulid_sdxl import StableDiffusionXLPuLIDPipeline, StableDiffusionXLPuLIDPipelineImage, StableDiffusionXLPuLIDPipelineInpaint +from pulid_utils import resize_numpy_image_long as resize +import attention_processor as attention +import pulid_sampling as sampling diff --git a/modules/pulid/attention_processor.py b/modules/pulid/attention_processor.py new file mode 100644 index 000000000..fa9e4ff82 --- /dev/null +++ b/modules/pulid/attention_processor.py @@ -0,0 +1,418 @@ +# 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_cat = torch.cat((id_embedding, zero_tensor), dim=1) + id_key = self.id_to_k(id_cat).to(query.dtype) + id_value = self.id_to_v(id_cat).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..834d5aa94 --- /dev/null +++ b/modules/pulid/encoders_transformer.py @@ -0,0 +1,250 @@ +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 + + +class IDEncoder(nn.Module): + def __init__(self, width=1280, context_dim=2048, num_token=5): + super().__init__() + self.num_token = num_token + self.context_dim = context_dim + h1 = min((context_dim * num_token) // 4, 1024) + h2 = min((context_dim * num_token) // 2, 1024) + self.body = nn.Sequential( + nn.Linear(width, h1), + nn.LayerNorm(h1), + nn.LeakyReLU(), + nn.Linear(h1, h2), + nn.LayerNorm(h2), + nn.LeakyReLU(), + nn.Linear(h2, context_dim * num_token), + ) + + 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, context_dim), + ), + ) + setattr( + self, + f'mapping_patch_{i}', + nn.Sequential( + nn.Linear(1024, 1024), + nn.LayerNorm(1024), + nn.LeakyReLU(), + nn.Linear(1024, 1024), + nn.LayerNorm(1024), + nn.LeakyReLU(), + nn.Linear(1024, context_dim), + ), + ) + + def forward(self, x, y): + # x shape [N, C] + x = self.body(x) + x = x.reshape(-1, self.num_token, self.context_dim) + + hidden_states = () + for i, emb in enumerate(y): + hidden_state = getattr(self, f'mapping_{i}')(emb[:, :1]) + getattr(self, f'mapping_patch_{i}')( + emb[:, 1:] + ).mean(dim=1, keepdim=True) + hidden_states += (hidden_state,) + hidden_states = torch.cat(hidden_states, dim=1) + + return torch.cat([x, hidden_states], dim=1) 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..bb87c540c --- /dev/null +++ b/modules/pulid/eva_clip/pretrained.py @@ -0,0 +1,331 @@ +import hashlib +import os +import urllib +import warnings +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("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_sampling.py b/modules/pulid/pulid_sampling.py new file mode 100644 index 000000000..6a2ef31f3 --- /dev/null +++ b/modules/pulid/pulid_sampling.py @@ -0,0 +1,594 @@ +import math +from scipy import integrate +import torch +from torch import nn +from torchdiffeq import odeint +import torchsde +from tqdm.auto import trange + + +def append_zero(x): + return torch.cat([x, x.new_zeros([1])]) + + +def get_sigmas_karras(n, sigma_min, sigma_max, rho=7., device='cpu'): + """Constructs the noise schedule of Karras et al. (2022).""" + ramp = torch.linspace(0, 1, n) + min_inv_rho = sigma_min ** (1 / rho) + max_inv_rho = sigma_max ** (1 / rho) + sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho + return append_zero(sigmas).to(device) + + +def get_sigmas_exponential(n, sigma_min, sigma_max, device='cpu'): + """Constructs an exponential noise schedule.""" + sigmas = torch.linspace(math.log(sigma_max), math.log(sigma_min), n, device=device).exp() + return append_zero(sigmas) + + +def get_sigmas_polyexponential(n, sigma_min, sigma_max, rho=1., device='cpu'): + """Constructs an polynomial in log sigma noise schedule.""" + ramp = torch.linspace(1, 0, n, device=device) ** rho + sigmas = torch.exp(ramp * (math.log(sigma_max) - math.log(sigma_min)) + math.log(sigma_min)) + return append_zero(sigmas) + + +def get_sigmas_vp(n, beta_d=19.9, beta_min=0.1, eps_s=1e-3, device='cpu'): + """Constructs a continuous VP noise schedule.""" + t = torch.linspace(1, eps_s, n, device=device) + sigmas = torch.sqrt(torch.exp(beta_d * t ** 2 / 2 + beta_min * t) - 1) + return append_zero(sigmas) + + +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') + return x[(...,) + (None,) * dims_to_append] + + +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.): + """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. + 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 + + +def default_noise_sampler(x): + return lambda sigma, sigma_next: torch.randn_like(x) + + +def inpaint_mask(x, i, steps, mask_args): + noised_original = mask_args["latent"].clone().to(x) + latent_mask = mask_args["latent_mask"].to(x) + if i < steps: + noised_original += mask_args["noise"].to(x) * mask_args["sigmas"][i+1].to(x) + x = (latent_mask * x) + ((1 - latent_mask) * noised_original.to(x)) + return x + + +class BatchedBrownianTree: + """A wrapper around torchsde.BrownianTree that enables batches of entropy.""" + + def __init__(self, x, t0, t1, seed=None, **kwargs): + 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 + 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) + 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): + 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) + + 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_euler(model, x, sigmas, extra_args=None, callback=None, disable=None, s_churn=0., s_tmin=0., s_tmax=float('inf'), s_noise=1., mask_args=None): + """Implements Algorithm 2 (Euler steps) from Karras et al. (2022).""" + extra_args = {} if extra_args is None else extra_args + s_in = x.new_ones([x.shape[0]]) + for i in trange(len(sigmas) - 1, disable=disable): + gamma = min(s_churn / (len(sigmas) - 1), 2 ** 0.5 - 1) if s_tmin <= sigmas[i] <= s_tmax else 0. + eps = torch.randn_like(x) * s_noise + sigma_hat = sigmas[i] * (gamma + 1) + if gamma > 0: + x = x + eps * (sigma_hat ** 2 - sigmas[i] ** 2) ** 0.5 + denoised = model(x, sigma_hat * s_in, **extra_args) + d = to_d(x, sigma_hat, denoised) + if callback is not None: + callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigma_hat, 'denoised': denoised}) + dt = sigmas[i + 1] - sigma_hat + # Euler method + x = x + (d * dt).to(x.dtype) + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) + return x + + +@torch.no_grad() +def sample_euler_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, mask_args=None): + """Ancestral sampling with Euler method steps.""" + extra_args = {} if extra_args is None else extra_args + noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler + s_in = x.new_ones([x.shape[0]]) + for i in trange(len(sigmas) - 1, disable=disable): + denoised = model(x, sigmas[i] * s_in, **extra_args) + sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta) + if callback is not None: + callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised}) + d = to_d(x, sigmas[i], denoised) + # Euler method + dt = sigma_down - sigmas[i] + x = x + (d * dt).to(x.dtype) + if sigmas[i + 1] > 0: + x = x + (noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up).to(x.dtype) + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) + return x + + +def linear_multistep_coeff(order, t, i, j): + if order - 1 > i: + raise ValueError(f'Order {order} too high for step {i}') + def fn(tau): + prod = 1. + for k in range(order): + if j == k: + continue + prod *= (tau - t[i - k]) / (t[i - j] - t[i - k]) + return prod + return integrate.quad(fn, t[i], t[i + 1], epsrel=1e-4)[0] + + +@torch.no_grad() +def log_likelihood(model, x, sigma_min, sigma_max, extra_args=None, atol=1e-4, rtol=1e-4): + extra_args = {} if extra_args is None else extra_args + s_in = x.new_ones([x.shape[0]]) + v = torch.randint_like(x, 2) * 2 - 1 + fevals = 0 + def ode_fn(sigma, x): + nonlocal fevals + with torch.enable_grad(): + x = x[0].detach().requires_grad_() + denoised = model(x, sigma * s_in, **extra_args) + d = to_d(x, sigma, denoised) + fevals += 1 + grad = torch.autograd.grad((d * v).sum(), x)[0] + d_ll = (v * grad).flatten(1).sum(1) + return d.detach(), d_ll + x_min = x, x.new_zeros([x.shape[0]]) + t = x.new_tensor([sigma_min, sigma_max]) + sol = odeint(ode_fn, x_min, t, atol=atol, rtol=rtol, method='dopri5') + latent, delta_ll = sol[0][-1], sol[1][-1] + ll_prior = torch.distributions.Normal(0, sigma_max).log_prob(latent).flatten(1).sum(1) + return ll_prior + delta_ll, {'fevals': fevals} + + +class PIDStepSizeController: + """A PID controller for ODE adaptive step size control.""" + def __init__(self, h, pcoeff, icoeff, dcoeff, order=1, accept_safety=0.81, eps=1e-8): + self.h = h + self.b1 = (pcoeff + icoeff + dcoeff) / order + self.b2 = -(pcoeff + 2 * dcoeff) / order + self.b3 = dcoeff / order + self.accept_safety = accept_safety + self.eps = eps + self.errs = [] + + def limiter(self, x): + return 1 + math.atan(x - 1) + + def propose_step(self, error): + inv_error = 1 / (float(error) + self.eps) + if not self.errs: + self.errs = [inv_error, inv_error, inv_error] + self.errs[0] = inv_error + factor = self.errs[0] ** self.b1 * self.errs[1] ** self.b2 * self.errs[2] ** self.b3 + factor = self.limiter(factor) + accept = factor >= self.accept_safety + if accept: + self.errs[2] = self.errs[1] + self.errs[1] = self.errs[0] + self.h *= factor + return accept + + +class DPMSolver(nn.Module): + """DPM-Solver. See https://arxiv.org/abs/2206.00927.""" + + def __init__(self, model, extra_args=None, eps_callback=None, info_callback=None): + super().__init__() + self.model = model + self.extra_args = {} if extra_args is None else extra_args + self.eps_callback = eps_callback + self.info_callback = info_callback + + def t(self, sigma): + return -sigma.log() + + def sigma(self, t): + return t.neg().exp() + + def eps(self, eps_cache, key, x, t, *args, **kwargs): + if key in eps_cache: + return eps_cache[key], eps_cache + sigma = self.sigma(t) * x.new_ones([x.shape[0]]) + eps = (x - self.model(x, sigma, *args, **self.extra_args, **kwargs)) / self.sigma(t) + if self.eps_callback is not None: + self.eps_callback() + return eps, {key: eps, **eps_cache} + + def dpm_solver_1_step(self, x, t, t_next, eps_cache=None): + eps_cache = {} if eps_cache is None else eps_cache + h = t_next - t + eps, eps_cache = self.eps(eps_cache, 'eps', x, t) + x_1 = x - self.sigma(t_next) * h.expm1() * eps + return x_1, eps_cache + + def dpm_solver_2_step(self, x, t, t_next, r1=1 / 2, eps_cache=None): + eps_cache = {} if eps_cache is None else eps_cache + h = t_next - t + eps, eps_cache = self.eps(eps_cache, 'eps', x, t) + s1 = t + r1 * h + u1 = x - self.sigma(s1) * (r1 * h).expm1() * eps + eps_r1, eps_cache = self.eps(eps_cache, 'eps_r1', u1, s1) + x_2 = x - self.sigma(t_next) * h.expm1() * eps - self.sigma(t_next) / (2 * r1) * h.expm1() * (eps_r1 - eps) + return x_2, eps_cache + + def dpm_solver_3_step(self, x, t, t_next, r1=1 / 3, r2=2 / 3, eps_cache=None): + eps_cache = {} if eps_cache is None else eps_cache + h = t_next - t + eps, eps_cache = self.eps(eps_cache, 'eps', x, t) + s1 = t + r1 * h + s2 = t + r2 * h + u1 = x - self.sigma(s1) * (r1 * h).expm1() * eps + eps_r1, eps_cache = self.eps(eps_cache, 'eps_r1', u1, s1) + u2 = x - self.sigma(s2) * (r2 * h).expm1() * eps - self.sigma(s2) * (r2 / r1) * ((r2 * h).expm1() / (r2 * h) - 1) * (eps_r1 - eps) + eps_r2, eps_cache = self.eps(eps_cache, 'eps_r2', u2, s2) + x_3 = x - self.sigma(t_next) * h.expm1() * eps - self.sigma(t_next) / r2 * (h.expm1() / h - 1) * (eps_r2 - eps) + return x_3, eps_cache + + def dpm_solver_fast(self, x, t_start, t_end, nfe, eta=0., s_noise=1., noise_sampler=None): + noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler + if not t_end > t_start and eta: + raise ValueError('eta must be 0 for reverse sampling') + + m = math.floor(nfe / 3) + 1 + ts = torch.linspace(t_start, t_end, m + 1, device=x.device) + + if nfe % 3 == 0: + orders = [3] * (m - 2) + [2, 1] + else: + orders = [3] * (m - 1) + [nfe % 3] + + for i in range(len(orders)): + eps_cache = {} + t, t_next = ts[i], ts[i + 1] + if eta: + sd, su = get_ancestral_step(self.sigma(t), self.sigma(t_next), eta) + t_next_ = torch.minimum(t_end, self.t(sd)) + su = (self.sigma(t_next) ** 2 - self.sigma(t_next_) ** 2) ** 0.5 + else: + t_next_, su = t_next, 0. + + eps, eps_cache = self.eps(eps_cache, 'eps', x, t) + denoised = x - self.sigma(t) * eps + if self.info_callback is not None: + self.info_callback({'x': x, 'i': i, 't': ts[i], 't_up': t, 'denoised': denoised}) + + if orders[i] == 1: + x, eps_cache = self.dpm_solver_1_step(x, t, t_next_, eps_cache=eps_cache) + elif orders[i] == 2: + x, eps_cache = self.dpm_solver_2_step(x, t, t_next_, eps_cache=eps_cache) + else: + x, eps_cache = self.dpm_solver_3_step(x, t, t_next_, eps_cache=eps_cache) + + x = x + su * s_noise * noise_sampler(self.sigma(t), self.sigma(t_next)) + + return x + + def dpm_solver_adaptive(self, x, t_start, t_end, order=3, rtol=0.05, atol=0.0078, h_init=0.05, pcoeff=0., icoeff=1., dcoeff=0., accept_safety=0.81, eta=0., s_noise=1., noise_sampler=None): + noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler + if order not in {2, 3}: + raise ValueError('order should be 2 or 3') + forward = t_end > t_start + if not forward and eta: + raise ValueError('eta must be 0 for reverse sampling') + h_init = abs(h_init) * (1 if forward else -1) + atol = torch.tensor(atol) + rtol = torch.tensor(rtol) + s = t_start + x_prev = x + accept = True + pid = PIDStepSizeController(h_init, pcoeff, icoeff, dcoeff, 1.5 if eta else order, accept_safety) + info = {'steps': 0, 'nfe': 0, 'n_accept': 0, 'n_reject': 0} + + while s < t_end - 1e-5 if forward else s > t_end + 1e-5: + eps_cache = {} + t = torch.minimum(t_end, s + pid.h) if forward else torch.maximum(t_end, s + pid.h) + if eta: + sd, su = get_ancestral_step(self.sigma(s), self.sigma(t), eta) + t_ = torch.minimum(t_end, self.t(sd)) + su = (self.sigma(t) ** 2 - self.sigma(t_) ** 2) ** 0.5 + else: + t_, su = t, 0. + + eps, eps_cache = self.eps(eps_cache, 'eps', x, s) + denoised = x - self.sigma(s) * eps + + if order == 2: + x_low, eps_cache = self.dpm_solver_1_step(x, s, t_, eps_cache=eps_cache) + x_high, eps_cache = self.dpm_solver_2_step(x, s, t_, eps_cache=eps_cache) + else: + x_low, eps_cache = self.dpm_solver_2_step(x, s, t_, r1=1 / 3, eps_cache=eps_cache) + x_high, eps_cache = self.dpm_solver_3_step(x, s, t_, eps_cache=eps_cache) + delta = torch.maximum(atol, rtol * torch.maximum(x_low.abs(), x_prev.abs())) + error = torch.linalg.norm((x_low - x_high) / delta) / x.numel() ** 0.5 + accept = pid.propose_step(error) + if accept: + x_prev = x_low + x = x_high + su * s_noise * noise_sampler(self.sigma(s), self.sigma(t)) + s = t + info['n_accept'] += 1 + else: + info['n_reject'] += 1 + info['nfe'] += order + info['steps'] += 1 + + if self.info_callback is not None: + self.info_callback({'x': x, 'i': info['steps'] - 1, 't': s, 't_up': s, 'denoised': denoised, 'error': error, 'h': pid.h, **info}) + + return x, info + + +@torch.no_grad() +def sample_dpmpp_2s_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, mask_args=None): + """Ancestral sampling with DPM-Solver++(2S) second-order steps.""" + extra_args = {} if extra_args is None else extra_args + noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler + s_in = x.new_ones([x.shape[0]]) + sigma_fn = lambda t: t.neg().exp() # pylint: disable=C3001 + t_fn = lambda sigma: sigma.log().neg() # pylint: disable=C3001 + + for i in trange(len(sigmas) - 1, disable=disable): + denoised = model(x, sigmas[i] * s_in, **extra_args) + sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta) + if callback is not None: + callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised}) + if sigma_down == 0: + # Euler method + d = to_d(x, sigmas[i], denoised) + dt = sigma_down - sigmas[i] + x = x + d * dt + else: + # DPM-Solver++(2S) + t, t_next = t_fn(sigmas[i]), t_fn(sigma_down) + r = 1 / 2 + h = t_next - t + s = t + r * h + x_2 = (sigma_fn(s) / sigma_fn(t)) * x - (-h * r).expm1() * denoised + denoised_2 = model(x_2, sigma_fn(s) * s_in, **extra_args) + x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_2 + # Noise addition + if sigmas[i + 1] > 0: + x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) + return x + + +@torch.no_grad() +def sample_dpmpp_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, r=1 / 2, mask_args=None): + """DPM-Solver++ (stochastic).""" + sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() + noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max) 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=C3001 + t_fn = lambda sigma: sigma.log().neg() # pylint: disable=C3001 + + 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 + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) + return x + + +@torch.no_grad() +def sample_dpmpp_2m(model, x, sigmas, extra_args=None, callback=None, disable=None, mask_args=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=C3001 + t_fn = lambda sigma: sigma.log().neg() # pylint: disable=C3001 + 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 + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) + return x + + +@torch.no_grad() +def sample_dpmpp_2m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, solver_type='midpoint', mask_args=None): + """DPM-Solver++(2M) SDE.""" + + if solver_type not in {'heun', 'midpoint'}: + raise ValueError('solver_type must be \'heun\' or \'midpoint\'') + + sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() + noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max) 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]]) + + old_denoised = None + h_last = 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}) + if sigmas[i + 1] == 0: + # Denoising step + x = denoised + else: + # DPM-Solver++(2M) SDE + t, s = -sigmas[i].log(), -sigmas[i + 1].log() + h = s - t + eta_h = eta * h + + x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + (-h - eta_h).expm1().neg() * denoised + + if old_denoised is not None: + r = h_last / h + if solver_type == 'heun': + x = x + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * (1 / r) * (denoised - old_denoised) + elif solver_type == 'midpoint': + x = x + 0.5 * (-h - eta_h).expm1().neg() * (1 / r) * (denoised - old_denoised) + + if eta: + x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise + + old_denoised = denoised + h_last = h + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) + return x + + +@torch.no_grad() +def sample_dpmpp_3m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, mask_args=None): + """DPM-Solver++(3M) SDE.""" + + sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max() + noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max) 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]]) + + denoised_1, denoised_2 = None, None + h_1, h_2 = None, 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}) + if sigmas[i + 1] == 0: + # Denoising step + x = denoised + else: + t, s = -sigmas[i].log(), -sigmas[i + 1].log() + h = s - t + h_eta = h * (eta + 1) + + x = torch.exp(-h_eta) * x + (-h_eta).expm1().neg() * denoised + + if h_2 is not None: + r0 = h_1 / h + r1 = h_2 / h + d1_0 = (denoised - denoised_1) / r0 + d1_1 = (denoised_1 - denoised_2) / r1 + d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1) + d2 = (d1_0 - d1_1) / (r0 + r1) + phi_2 = h_eta.neg().expm1() / h_eta + 1 + phi_3 = phi_2 / h_eta - 0.5 + x = x + phi_2 * d1 - phi_3 * d2 + elif h_1 is not None: + r = h_1 / h + d = (denoised - denoised_1) / r + phi_2 = h_eta.neg().expm1() / h_eta + 1 + x = x + phi_2 * d + + if eta: + x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * h * eta).expm1().neg().sqrt() * s_noise + + denoised_1, denoised_2 = denoised, denoised_1 + h_1, h_2 = h, h_1 + if mask_args is not None: + x = inpaint_mask(x, i, len(sigmas) - 2, mask_args) + return x diff --git a/modules/pulid/pulid_sdxl.py b/modules/pulid/pulid_sdxl.py new file mode 100644 index 000000000..01ca660a7 --- /dev/null +++ b/modules/pulid/pulid_sdxl.py @@ -0,0 +1,450 @@ +from typing import Union +import os +import cv2 +import insightface +import numpy as np +import torch +import torch.nn as nn +from PIL import Image +from diffusers import StableDiffusionXLPipeline, StableDiffusionXLImg2ImgPipeline, StableDiffusionXLInpaintPipeline +from diffusers.pipelines.stable_diffusion_xl.pipeline_output import StableDiffusionXLPipelineOutput + +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, IDEncoder +from modules.errors import log + + +debug = log.trace if os.environ.get('SD_PULID_DEBUG', None) is not None else lambda *args, **kwargs: None + + +class StableDiffusionXLPuLIDPipeline: + def __init__(self, + pipe: Union[StableDiffusionXLPipeline, StableDiffusionXLImg2ImgPipeline, StableDiffusionXLInpaintPipeline], + device: torch.device, + dtype: torch.dtype=None, + providers: list=None, + offload: bool=True, + sampler=None, + cache_dir=None, + sdp: bool=True, + version: str='v1.1', + ): + super().__init__() + self.device = device + self.dtype = dtype or torch.float16 + self.pipe = pipe + self.cache_dir = cache_dir + self.offload = offload + self.sdp = sdp + self.version = version + self.folder = 'models--ToTheBeginning--PuLID' + debug(f'PulID init: device={self.device} dtype={self.dtype} dir={self.cache_dir} offload={self.offload} sdp={self.sdp} version={self.version}') + + # self.pipe.scheduler = DPMSolverMultistepScheduler.from_config(self.pipe.scheduler.config) + self.hack_unet_attn_layers(self.pipe.unet) + if self.version == 'v1.1': + self.id_adapter = IDFormer().to(self.device, self.dtype) + else: + self.id_adapter = IDEncoder().to(self.device, self.dtype) + debug(f'PulID load: adapter={self.id_adapter.__class__.__name__}') + self.providers = providers or ['CUDAExecutionProvider', 'CPUExecutionProvider'] + debug(f'PulID load: providers={self.providers}') + + # 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 = init_parsing_model(model_name='bisenet', device=self.device) + debug(f'PulID load: facehelper={self.face_helper.__class__.__name__}') + + # clip-vit backbone + eva_precision = 'fp16' if self.dtype == torch.float16 or self.dtype == torch.bfloat16 else 'fp32' + eva_model, _, _ = create_model_and_transforms('EVA02-CLIP-L-14-336', 'eva_clip', force_custom_clip=True, precision=eva_precision, device=self.device) + self.clip_vision_model = eva_model.visual.to(dtype=self.dtype) + debug(f'PulID load: evaclip={self.clip_vision_model.__class__.__name__} precision={eva_precision}') + 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 + local_dir = os.path.join(self.cache_dir, self.folder, 'models', 'antelopev2') + _loc = snapshot_download('DIAMONIK7777/antelopev2', local_dir=local_dir) + self.app = FaceAnalysis( + name='antelopev2', + root=os.path.join(self.cache_dir, self.folder), + providers=self.providers, + ) + debug(f'PulID load: faceanalysis={_loc}') + 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) + debug(f'PulID load: handler={self.handler_ante.__class__.__name__}') + + 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 + + # default scheduler + if sampler is not None: + self.sampler = sampler + else: + from modules.pulid import sampling + self.sampler = sampling.sample_dpmpp_sde + + @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): + if self.sdp: + from attention_processor import AttnProcessor2_0 as AttnProcessor + from attention_processor import IDAttnProcessor2_0 as IDAttnProcessor + else: + from attention_processor import AttnProcessor + from attention_processor import IDAttnProcessor + 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, unet.dtype) + else: + id_adapter_attn_procs[name] = AttnProcessor() + debug(f'PulID attention: cls={IDAttnProcessor} std={AttnProcessor} len={len(id_adapter_attn_procs.keys())}') + unet.set_attn_processor(id_adapter_attn_procs) + self.id_adapter_attn_layers = nn.ModuleList(unet.attn_processors.values()) + + def load_pretrain(self): + if self.version == 'v1.1': + ckpt_path = hf_hub_download('guozinan/PuLID', 'pulid_v1.1.safetensors', local_dir=os.path.join(self.cache_dir, self.folder)) + state_dict = load_file(ckpt_path) + else: + ckpt_path = hf_hub_download('guozinan/PuLID', 'pulid_v1.bin', local_dir=os.path.join(self.cache_dir, self.folder)) + state_dict = torch.load(ckpt_path, map_location="cpu") + debug(f'PulID load: fn="{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.to(self.dtype) + + for module in state_dict_dict: + 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 = [] + self.face_helper.face_det.to(self.device) + self.clip_vision_model.to(self.device) + 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: + 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).to(self.dtype) + 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) + + self.id_adapter.to(self.device) + id_uncond = torch.zeros_like(id_cond_list[0]).to(self.dtype) + 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]).to(self.dtype)) + + id_cond = torch.stack(id_cond_list, dim=1).to(self.dtype) + 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).to(self.dtype) + id_embedding = self.id_adapter(id_cond, id_vit_hidden) + uncond_id_embedding = self.id_adapter(id_uncond, id_vit_hidden_uncond) + + if self.offload: + self.face_helper.face_det.to('cpu') + self.id_adapter.to('cpu') + self.clip_vision_model.to('cpu') + + # return id_embedding + debug(f'PulID embedding: cond={id_embedding.shape} uncond={uncond_id_embedding.shape}') + 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_sampling + pulid_sampling.trange = functools.partial(trange_orig, bar_format=bar_format, ncols=ncols, colour=colour) + + def sample(self, x, sigma, **extra_args): + t = self.timestep(sigma) + x_ddim_space = x / (sigma[:, None, None, None] ** 2 + self.sigma_data**2) ** 0.5 + cfg_scale = extra_args['cfg_scale'] + # debug(f'PulID sample start: step={self.step+1} x={x.shape} dtype={x.dtype} timestep={t.item()} sigma={sigma.shape} cfg={cfg_scale} args={extra_args.keys()}') + 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 }) + # debug(f'PulID sample end: step={self.step} x={latent.shape} dtype={x.dtype} min={torch.amin(latent)} max={torch.amax(latent)}') + return latent + + def init_latent(self, seed, size, image, mask_image, strength, width, height): # pylint: disable=unused-argument + # standard txt2img will full noise + noise = torch.randn((size[0], 4, size[1] // 8, size[2] // 8), device="cpu", generator=torch.manual_seed(seed)) + noise = noise.to(dtype=self.pipe.unet.dtype, device=self.device) + if strength > 0 and image is not None: + image = self.pipe.image_processor.preprocess(image) + if mask_image is not None: # Inpaint + latents = self.pipe.prepare_latents(1, # batch_size, + self.pipe.vae.config.latent_channels, # num_channels_latents + height, + width, + noise.dtype, + noise.device, + None, # generator + latents=None, + image=image, + timestep=1000, + is_strength_max=False, + add_noise=False, + return_noise=False, + return_image_latents=False, + ) + latents = latents[0] + debug(f'PulID noise: op=inpaint latent={latents.shape} image={image} mask={mask_image} dtype={latents.dtype}') + else: # img2img + latents = self.pipe.prepare_latents(image, + None, # timestep (not needed) + 1, # batch_size + 1, # num_images_per_prompt + noise.dtype, + noise.device, + None, # generator + False, # add_noise + ) + debug(f'PulID noise: op=img2img latent={latents.shape} image={image} dtype={latents.dtype}') + else: + latents = torch.zeros_like(noise) + debug(f'PulID noise: op=txt2img latent={latents.shape} dtype={latents.dtype}') + return latents, noise + + 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, + mask_image: np.ndarray=None, + strength: float=0.3, + id_embedding=None, + uncond_id_embedding=None, + id_scale: float=1.0, + output_type: str='pil', + callback_on_step_end=None, + ): + debug(f'PulID call: width={width} height={height} cfg={guidance_scale} steps={num_inference_steps} seed={seed} strength={strength} id_scale={id_scale} output={output_type}') + self.step = 0 # pylint: disable=attribute-defined-outside-init + self.callback_on_step_end = callback_on_step_end # pylint: disable=attribute-defined-outside-init + if isinstance(image, list) and len(image) > 0 and isinstance(image[0], Image.Image): + if image[0].width != width or image[0].height != height: # override width/height if different + width, height = image[0].width, image[0].height + size = (1, height, width) + # sigmas + sigmas = self.get_sigmas_karras(num_inference_steps).to(self.device) + if image is not None and strength > 0: + _, num_inference_steps = self.pipe.get_timesteps(num_inference_steps, strength, self.device, None) # denoising_start disabled + sigmas = sigmas[-(num_inference_steps + 1):].to(self.device) # shorten sigmas in i2i + debug(f'PulID sigmas: sigmas={sigmas.shape} dtype={sigmas.dtype}') + + # latents + latent, noise = self.init_latent(seed, size, image, mask_image, strength, width, height) + noisy_latent = latent + noise * sigmas[0].to(noise) + debug(f'PulID noisy: latent={noisy_latent.shape} dtype={noisy_latent.dtype}') + + ( + 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}, + ), + ) + if mask_image is not None: + latent_mask = torch.Tensor(np.asarray(mask_image.convert("L").resize((noisy_latent.shape[-1], noisy_latent.shape[-2])))).reshape((noisy_latent.shape[-2], noisy_latent.shape[-1])) + latent_mask /= latent_mask.max() + mask_args = dict( + latent=latent, + latent_mask=latent_mask, + noise=noise, + sigmas=sigmas, + ) + else: + mask_args = None + + # actual sampling loop + latents = self.sampler(self.sample, noisy_latent, sigmas, extra_args=sampler_kwargs, disable=False, mask_args=mask_args) + + # process output + latents = latents.to(dtype=self.pipe.vae.dtype, device=self.device) + debug(f'PulID output: latent={latents.shape} dtype={latents.dtype}') + if output_type == 'latent': + images = self.pipe.image_processor.postprocess(latents, output_type='latent') + elif output_type == 'np': + images = self.pipe.image_processor.postprocess(latents, output_type='np') + else: + latents = latents / self.pipe.vae.config.scaling_factor + images = self.pipe.vae.decode(latents).sample + images = self.pipe.image_processor.postprocess(images, output_type='pil') + debug(f'PulID output: type={type(images)} images={images.shape if hasattr(images, "shape") else images}') + return StableDiffusionXLPipelineOutput(images) + + +class StableDiffusionXLPuLIDPipelineImage(StableDiffusionXLPuLIDPipeline): + def __init__(self, pipe: StableDiffusionXLPipeline, device: torch.device, sampler=None, cache_dir=None): # pylint: disable=useless-parent-delegation + super().__init__(pipe, device, sampler, cache_dir) + # we dont do anything special here, just having different class so task-type can be detected/assigned + + +class StableDiffusionXLPuLIDPipelineInpaint(StableDiffusionXLPuLIDPipeline): + def __init__(self, pipe: StableDiffusionXLPipeline, device: torch.device, sampler=None, cache_dir=None): # pylint: disable=useless-parent-delegation + super().__init__(pipe, device, sampler, cache_dir) + # we dont do anything special here, just having different class so task-type can be detected/assigned diff --git a/modules/pulid/pulid_utils.py b/modules/pulid/pulid_utils.py new file mode 100644 index 000000000..fd7338b5e --- /dev/null +++ b/modules/pulid/pulid_utils.py @@ -0,0 +1,161 @@ +import importlib +import math +import os +import random + +import cv2 +import numpy as np +import torch +from torchvision.utils import make_grid +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 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 diff --git a/modules/dcsolver/__init__.py b/modules/schedulers/scheduler_dc.py similarity index 100% rename from modules/dcsolver/__init__.py rename to modules/schedulers/scheduler_dc.py diff --git a/modules/schedulers/scheduler_dpm_flowmatch.py b/modules/schedulers/scheduler_dpm_flowmatch.py new file mode 100644 index 000000000..83573105e --- /dev/null +++ b/modules/schedulers/scheduler_dpm_flowmatch.py @@ -0,0 +1,897 @@ +# Credits: @ukaprch + +import math +from dataclasses import dataclass +from typing import List, Optional, Tuple, Union + +import numpy as np +import torch +import torchsde + +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 + +class BatchedBrownianTree: + """A wrapper around torchsde.BrownianTree that enables batches of entropy.""" + + def __init__(self, x, t0, t1, seed=None, **kwargs): + 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 + self.trees = [ + torchsde.BrownianInterval( + t0=t0, + t1=t1, + size=w0.shape, + dtype=w0.dtype, + device=w0.device, + entropy=s, + tol=1e-6, + pool_size=24, + halfway_tree=True, + ) + 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) + 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): + 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) + + 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() + +@dataclass +class FlowMatchDPMSolverMultistepSchedulerOutput(BaseOutput): + """ + Output class for the scheduler's `step` function output. + + Args: + prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images): + Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the + denoising loop. + """ + + prev_sample: torch.FloatTensor + +class FlowMatchDPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin): + """ + `DPMSolverMultistepScheduler` is a fast dedicated high-order solver for diffusion ODEs. + + This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic + methods the library implements for all schedulers such as loading and saving. + + Args: + num_train_timesteps (`int`, defaults to 1000): + The number of diffusion steps to train the model. + solver_order (`int`, defaults to 2): + The DPMSolver order which can be `2` or `3`. It is recommended to use `solver_order=2` for guided + sampling, and `solver_order=3` for unconditional sampling. + 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`. + algorithm_type (`str`, defaults to `dpmsolver++2M`): + Algorithm type for the solver; can be `dpmsolver2`, `dpmsolver2A`, `dpmsolver++2M`, `dpmsolver++2S`, `dpmsolver++sde`, `dpmsolver++2Msde`, + or `dpmsolver++3Msde`. + solver_type (`str`, defaults to `midpoint`): + Solver type for the second-order solver; can be `midpoint` or `heun`. The solver type slightly affects the + sample quality, especially for a small number of steps. It is recommended to use `midpoint` solvers. + sigma_schedule (`str`, *optional*, defaults to None): Sigma schedule to compute the `sigmas`. Optionally, we use + the schedule "karras" introduced in the EDM paper (https://arxiv.org/abs/2206.00364). Other acceptable values are + "exponential". The exponential schedule was incorporated in this model: https://huggingface.co/stabilityai/cosxl. + Other acceptable values are "lambdas". The uniform-logSNR for step sizes proposed by Lu's DPM-Solver in the + noise schedule during the sampling process. The sigmas and time steps are determined according to a sequence of `lambda(t)`. + use_noise_sampler for BrownianTreeNoiseSampler (only valid for `dpmsolver++2S`, `dpmsolver++sde`, `dpmsolver++2Msde`, + or `dpmsolver++3Msde`): A noise sampler backed by a torchsde increasing the stability of convergence. Default strategy + (random noise) has it jumping all over the place, but Brownian sampling is more stable. Utilizes the model generation seed provided. + midpoint_ratio (`float`, *optional*, range: 0.4 to 0.6, default=0.5): Only valid for (`dpmsolver++sde`, `dpmsolver++2S`). + Higher values may result in smoothing, more vivid colors and less noise at the expense of more detail and effect. + s_noise (`float`, *optional*, defaults to 1.0): Sigma noise strength: range 0 - 1.1 (only valid for `dpmsolver++2S`, `dpmsolver++sde`, + `dpmsolver++2Msde`, or `dpmsolver++3Msde`). The amount of additional noise to counteract loss of detail during sampling. A + reasonable range is [1.000, 1.011]. Defaults to 1.0 from the original implementation. + use_SD35_sigmas: (`bool` defaults to False for FLUX and True for SD3). Based on original interpretation of using beta values for determining sigmas. + use_dynamic_shifting (`bool` defaults to False for SD3 and True for FLUX). When `True`, shift is ignored. + shift (`float`, defaults to 3.0): The shift value for the timestep schedule for SD3 when not using dynamic shifting + The remaining args are specific to Flux's dynamic shifting based on resolution + """ + + _compatibles = [] + order = 1 + + @register_to_config + def __init__( + self, + num_train_timesteps: int = 1000, + solver_order: int = 2, + thresholding: Optional[bool] = False, + dynamic_thresholding_ratio: float = 0.995, + sample_max_value: Optional[float] = 1.0, + algorithm_type: str = "dpmsolver++2M", + solver_type: str = "midpoint", + sigma_schedule: Optional[str] = None, + shift: float = 3.0, + midpoint_ratio: Optional[float] = 0.5, + s_noise: Optional[float] = 1.0, + use_noise_sampler: Optional[bool] = True, + use_SD35_sigmas: Optional[bool] = False, + use_dynamic_shifting=False, + base_shift: Optional[float] = 0.5, + max_shift: Optional[float] = 1.15, + base_image_seq_len: Optional[int] = 256, + max_image_seq_len: Optional[int] = 4096, + ): + # settings for DPM-Solver + if algorithm_type not in ["dpmsolver2", "dpmsolver2A", "dpmsolver++2M", "dpmsolver++2S", "dpmsolver++sde", "dpmsolver++2Msde", "dpmsolver++3Msde"]: + raise NotImplementedError(f"{algorithm_type} is not implemented for {self.__class__}") + + if solver_type not in ["midpoint", "heun"]: + raise NotImplementedError(f"{solver_type} is not implemented for {self.__class__}") + + # setable values + timesteps = np.linspace(1, num_train_timesteps, num_train_timesteps, dtype=np.float32)[::-1].copy() + timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32) + + sigmas = timesteps / num_train_timesteps + if not use_dynamic_shifting: + # when use_dynamic_shifting is True, we apply the timestep shifting on the fly based on the image resolution + sigmas = shift * sigmas / (1 + (shift - 1) * sigmas) + + self.timesteps = sigmas * num_train_timesteps + self.h_last = None + self.h_1 = None + self.h_2 = None + self.noise_sampler = None + self._step_index = None + self._begin_index = None + self.sigmas = sigmas.to("cpu") # to avoid too much CPU/GPU communication + self.model_outputs = [None] * solver_order + self.sigma_min = self.sigmas[-1].item() + self.sigma_max = self.sigmas[0].item() + + @property + def step_index(self): + """ + The index counter for current timestep. It will increase 1 after each scheduler step. + """ + return self._step_index + + @property + def begin_index(self): + """ + The index for the first timestep. It should be set from pipeline with `set_begin_index` method. + """ + return self._begin_index + + def set_begin_index(self, begin_index: int = 0): + """ + Sets the begin index for the scheduler. This function should be run from pipeline before the inference. + + Args: + begin_index (`int`): + The begin index for the scheduler. + """ + self._begin_index = begin_index + + def time_shift(self, mu: float, sigma: float, t: torch.Tensor): + return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma) + + def set_timesteps(self, + num_inference_steps: int = None, + device: Union[str, torch.device] = None, + sigmas: Optional[List[float]] = None, + mu: Optional[float] = 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. + """ + if self.config.use_dynamic_shifting and mu is None: + raise ValueError(" you have a pass a value for `mu` when `use_dynamic_shifting` is set to be `True`") + + if sigmas is None: + self.use_SD35_sigmas = True + self.num_inference_steps = num_inference_steps + sigmas1 = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps, dtype=np.float64) + beta_start = 0.00085 + beta_end = 0.012 + betas = torch.linspace(beta_start**0.5, beta_end**0.5, self.config.num_train_timesteps, dtype=torch.float64) ** 2 + alphas = 1.0 - betas + alphas_cumprod = torch.cumprod(alphas, dim=0) + sigmas = np.array(((1 - alphas_cumprod) / alphas_cumprod) ** 0.5) + del alphas_cumprod + del alphas + del betas + elif self.use_SD35_sigmas: + num_inference_steps = len(sigmas) + self.num_inference_steps = num_inference_steps + sigmas1 = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps, dtype=np.float64) + beta_start = 0.00085 + beta_end = 0.012 + betas = torch.linspace(beta_start**0.5, beta_end**0.5, self.config.num_train_timesteps, dtype=torch.float64) ** 2 + alphas = 1.0 - betas + alphas_cumprod = torch.cumprod(alphas, dim=0) + sigmas = np.array(((1 - alphas_cumprod) / alphas_cumprod) ** 0.5) + del alphas_cumprod + del alphas + del betas + else: + num_inference_steps = len(sigmas) + self.num_inference_steps = num_inference_steps + + if self.config.sigma_schedule == "exponential": + if self.use_SD35_sigmas: + sigmas = np.flip(sigmas).copy() + sigma_min = sigmas[-1] + sigma_max = sigmas[0] + sigmas = self._convert_to_exponential(sigma_min, sigma_max, num_inference_steps=num_inference_steps) + OldRange = sigma_max - sigma_min + NewRange = 1.0 - sigma_min + sigmas = (((sigmas - sigma_min) * NewRange) / OldRange) + sigma_min + del sigmas1 + else: + sigma_min = sigmas[-1] + sigma_max = sigmas[0] + sigmas = self._convert_to_exponential(sigma_min, sigma_max, num_inference_steps=num_inference_steps) + elif self.config.sigma_schedule == "karras": + if self.use_SD35_sigmas: + sigmas = np.flip(sigmas).copy() + sigma_min = sigmas[-1] + sigma_max = sigmas[0] + sigmas = self._convert_to_karras(sigma_min, sigma_max, num_inference_steps=num_inference_steps) + OldRange = sigma_max - sigma_min + NewRange = 1.0 - sigma_min + sigmas = (((sigmas - sigma_min) * NewRange) / OldRange) + sigma_min + del sigmas1 + else: + sigma_min = sigmas[-1] + sigma_max = sigmas[0] + sigmas = self._convert_to_karras(sigma_min, sigma_max, num_inference_steps=num_inference_steps) + sigmas = torch.from_numpy(sigmas).to(dtype=torch.float64, device=device) + elif self.config.sigma_schedule == "lambdas": + if self.use_SD35_sigmas: + log_sigmas = np.log(sigmas) + lambdas = np.flip(log_sigmas.copy()) + lambdas = self._convert_to_lu(in_lambdas=lambdas, num_inference_steps=num_inference_steps) + sigmas = np.exp(lambdas) + sigma_min = sigmas[-1] + sigma_max = sigmas[0] + OldRange = sigma_max - sigma_min + NewRange = 1.0 - sigma_min + sigmas = (((sigmas - sigma_min) * NewRange) / OldRange) + sigma_min + del sigmas1 + del lambdas + del log_sigmas + else: + log_sigmas = np.log(sigmas) + lambdas = log_sigmas.copy() + lambdas = self._convert_to_lu(in_lambdas=lambdas, num_inference_steps=num_inference_steps) + sigmas = np.exp(lambdas) + del lambdas + del log_sigmas + sigmas = torch.from_numpy(sigmas).to(dtype=torch.float64, device=device) + else: + if self.use_SD35_sigmas: + sigmas = np.flip(sigmas).copy() + sigma_min = sigmas[-1] + sigmas = np.linspace(1.0, sigma_min, num_inference_steps) + del sigmas1 + sigmas = torch.from_numpy(sigmas).to(dtype=torch.float64, device=device) + + if self.config.use_dynamic_shifting: + sigmas = self.time_shift(mu, 1.0, sigmas) + else: + sigmas = self.config.shift * sigmas / (1 + (self.config.shift - 1) * sigmas) + + timesteps = sigmas * self.config.num_train_timesteps + self.timesteps = timesteps.to(device=device) + self.sigmas = torch.cat([sigmas, torch.zeros(1, device=sigmas.device)]) + self.h_last = None + self.h_1 = None + self.h_2 = None + self.noise_sampler = None + self.model_outputs = [None] * self.config.solver_order + self._step_index = None + self._begin_index = None + + # Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler._threshold_sample + def _threshold_sample(self, sample: torch.Tensor) -> torch.Tensor: + """ + "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 _convert_to_lu(self, in_lambdas: torch.Tensor, num_inference_steps) -> torch.Tensor: + """Constructs the noise schedule of Lu et al. (2022).""" + + lambda_min: float = in_lambdas[-1].item() + lambda_max: float = in_lambdas[0].item() + + rho = 1.0 # 1.0 is the value used in the paper + ramp = np.linspace(0, 1, num_inference_steps) + min_inv_rho = lambda_min ** (1 / rho) + max_inv_rho = lambda_max ** (1 / rho) + lambdas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho + return lambdas + + # Copied from diffusers.schedulers.scheduling_euler_discrete.EulerDiscreteScheduler._convert_to_karras + def _convert_to_karras(self, sigma_min, sigma_max, num_inference_steps) -> torch.Tensor: + rho = 7.0 # 7.0 is the value used in the paper + ramp = np.linspace(0, 1, num_inference_steps) + min_inv_rho = sigma_min ** (1 / rho) + max_inv_rho = sigma_max ** (1 / rho) + sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho + return sigmas + + def _convert_to_exponential(self, sigma_min, sigma_max, num_inference_steps) -> torch.Tensor: + sigmas = torch.linspace(math.log(sigma_max), math.log(sigma_min), num_inference_steps).exp() + return sigmas + + def convert_model_output( + self, + model_output: torch.Tensor, + sample: torch.Tensor = None, + *args, + **kwargs, + ) -> torch.Tensor: + """ + Convert the model output to the corresponding type the DPMSolver/DPMSolver++ algorithm needs. DPM-Solver is + designed to discretize an integral of the noise prediction model, and DPM-Solver++ is designed to discretize an + integral of the data prediction model. + + + + The algorithm and model type are decoupled. You can use either DPMSolver or DPMSolver++ for both noise + prediction and data prediction models. + + + + Args: + model_output (`torch.Tensor`): + The direct output from the learned diffusion model. + sample (`torch.Tensor`): + A current instance of a sample created by the diffusion process. + + Returns: + `torch.Tensor`: + The converted model output. + """ + timestep = args[0] if len(args) > 0 else kwargs.pop("timestep", None) + if sample is None: + if len(args) > 1: + sample = args[1] + else: + raise ValueError("missing `sample` as a required keyward argument") + + # Flow Match needs to solve an integral of the data prediction model. + sigma = self.sigmas[self.step_index] + x0_pred = sample - sigma * model_output + + if self.config.thresholding: + x0_pred = self._threshold_sample(x0_pred) + + return x0_pred + + def index_for_timestep(self, timestep, schedule_timesteps=None): + if schedule_timesteps is None: + schedule_timesteps = self.timesteps + + indices = (schedule_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) + pos = 1 if len(indices) > 1 else 0 + + return indices[pos].item() + + def _init_step_index(self, timestep): + if self.begin_index is None: + if isinstance(timestep, torch.Tensor): + timestep = timestep.to(self.timesteps.device) + self._step_index = self.index_for_timestep(timestep) + else: + self._step_index = self._begin_index + + def step( + self, + model_output: torch.FloatTensor, + timestep: Union[float, torch.FloatTensor], + sample: torch.FloatTensor, + generator: Optional[torch.Generator] = None, + variance_noise: Optional[torch.FloatTensor] = None, + return_dict: bool = True, + ) -> Union[FlowMatchDPMSolverMultistepSchedulerOutput, Tuple]: + """ + Predict the sample from the previous timestep by reversing the SDE. This function propagates the sample with + the multistep DPMSolver. + + Args: + model_output (`torch.Tensor`): + The direct output from learned diffusion model. + timestep (`int`): + The current discrete timestep in the diffusion chain. + sample (`torch.Tensor`): + A current instance of a sample created by the diffusion process. + generator (`torch.Generator`, *optional*): + A random number generator. + variance_noise (`torch.Tensor`): + Alternative to generating noise with `generator` by directly providing the noise for the variance + itself. Useful for methods such as [`LEdits++`]. + return_dict (`bool`): + Whether or not to return a [`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`. + + Returns: + [`~schedulers.scheduling_utils.SchedulerOutput`] or `tuple`: + If return_dict is `True`, [`~schedulers.scheduling_utils.SchedulerOutput`] is returned, otherwise a + tuple is returned where the first element is the sample tensor. + + """ + if self.num_inference_steps is None: + raise ValueError( + "Number of inference steps is 'None', you need to run 'set_timesteps' after creating the scheduler" + ) + + if self.step_index is None: + self._init_step_index(timestep) + + if self.config.algorithm_type in ["dpmsolver2", "dpmsolver2A"]: + pass + else: + model_output = self.convert_model_output(model_output, sample=sample) + for i in range(self.config.solver_order - 1): + self.model_outputs[i] = self.model_outputs[i + 1] + self.model_outputs[-1] = model_output + + # Upcast to avoid precision issues when computing prev_sample + if sample.dtype != model_output.dtype: + sample = sample.to(model_output.dtype) + + if self.config.algorithm_type in ["dpmsolver2A", "dpmsolver++2S", "dpmsolver++sde", "dpmsolver++2Msde", "dpmsolver++3Msde"] and variance_noise is None: + # Create a noise sampler if it hasn't been created yet + if self.config.use_noise_sampler: + if self.noise_sampler is None: + min_sigma, max_sigma = self.sigmas[self.sigmas > 0].min(), self.sigmas.max() + self.noise_sampler = BrownianTreeNoiseSampler(sample, min_sigma, max_sigma, generator) + else: + noise = randn_tensor(model_output.shape, generator=generator, device=model_output.device, dtype=model_output.dtype) + elif self.config.algorithm_type in ["dpmsolver2A", "dpmsolver++2S", "dpmsolver++sde", "dpmsolver++2Msde", "dpmsolver++3Msde"]: + noise = variance_noise.to(device=model_output.device, dtype=model_output.dtype) + else: + noise = None + + def sigma_fn(_t: torch.Tensor) -> torch.Tensor: + return _t.neg().exp() + def t_fn(_sigma: torch.Tensor) -> torch.Tensor: + return _sigma.log().neg() + sigma = self.sigmas[self.step_index] + sigma_next = self.sigmas[self.step_index + 1] + sigma_prev = self.sigmas[self.step_index - 1] + if self.config.algorithm_type == "dpmsolver2": + if self.config.solver_order == 2: + if sigma_next == 0: + # Euler method + model_output = sample - sigma * model_output + d = (sample - model_output) / sigma + dt = sigma_next - sigma + sample = sample + d * dt + else: + # DPM-Solver2 + sigma_mid = sigma.log().lerp(sigma_next.log(), 0.5).exp() + + #using epsilon for new model output: + pred_original_sample = sample - sigma * model_output + # 2. Convert to an ODE derivative for 1st order + d = (sample - pred_original_sample) / sigma + # 3. delta timestep + dt = sigma_mid - sigma + x_2 = sample + d * dt + + #using epsilon for new model output: + denoised_2 = x_2 - sigma_mid * model_output + # 2. Convert to an ODE derivative for 2nd order + d = (x_2 - denoised_2) / sigma_mid + + # 3. delta timestep + dt = sigma_next - sigma + sample = sample + d * dt + + del pred_original_sample + del denoised_2 + del x_2 + del d + elif self.config.algorithm_type == "dpmsolver2A": + if self.config.solver_order == 2: + # get ancestral step + sigma_from = sigma + sigma_to = sigma_next + su = min(sigma_to, (sigma_to**2 * (sigma_from**2 - sigma_to**2) / sigma_from**2) ** 0.5) + sd = (sigma_to**2 - su**2) ** 0.5 + if sd == 0: + # Euler method + model_output = sample - sigma * model_output + d = (sample - model_output) / sigma + dt = sd - sigma + sample = sample + d * dt + else: + # DPM-Solver2A + sigma_mid = sigma.log().lerp(sd.log(), 0.5).exp() + + #using epsilon for new model output: + model_output = sample - sigma * model_output + # 2. Convert to an ODE derivative for 1st order + d = (sample - model_output) / sigma + dt = sd - sigma + sample = sample + d * dt + + #using epsilon for new model output: + pred_original_sample = sample - sigma * model_output + # 2. Convert to an ODE derivative for 1st order + d = (sample - pred_original_sample) / sigma + # 3. delta timestep + dt_1 = sigma_mid - sigma + x_2 = sample + d * dt_1 + + #using epsilon for new model output: + denoised_2 = x_2 - sigma_mid * model_output + # 2. Convert to an ODE derivative for 2nd order + d_2 = (x_2 - denoised_2) / sigma_mid + + # 3. delta timestep + dt_2 = sd - sigma_mid + sample = sample + d_2 * dt_2 + + if self.config.use_noise_sampler: + sample = sample + self.noise_sampler(sigma, sigma_next) * self.config.s_noise * su + else: + sample = sample + noise * self.config.s_noise * su + + del pred_original_sample + del denoised_2 + del x_2 + del d + elif self.config.algorithm_type == "dpmsolver++2M": + if self.config.solver_order == 2: + t, t_next = t_fn(sigma), t_fn(sigma_next) + h = t_next - t + if self.model_outputs[-2] is None or sigma_next == 0: + sample = (sigma_fn(t_next) / sigma_fn(t)) * sample - (-h).expm1() * model_output + else: + # DPM-Solver++(2M) + h_last = t - t_fn(sigma_prev) + r = h_last / h + denoised_d = (1 + 1 / (2 * r)) * model_output - (1 / (2 * r)) * self.model_outputs[-2] + sample = (sigma_fn(t_next) / sigma_fn(t)) * sample - (-h).expm1() * denoised_d + del denoised_d + elif self.config.algorithm_type == "dpmsolver++2S": + if self.config.solver_order == 2: + # get ancestral step + sigma_from = sigma + sigma_to = sigma_next + su = min(sigma_to, (sigma_to**2 * (sigma_from**2 - sigma_to**2) / sigma_from**2) ** 0.5) + sd = (sigma_to**2 - su**2) ** 0.5 + if sd == 0: + # Euler method + d = (sample - model_output) / sigma + dt = sd - sigma + sample = sample + d * dt + else: + # DPM-Solver++(2S) + t, t_next = t_fn(sigma), t_fn(sd) + r = self.config.midpoint_ratio + h = t_next - t + s = t + r * h + + # Euler method + d = (sample - model_output) / sigma + dt = sd - sigma + sample = sample + d * dt + + x_2 = (sigma_fn(s) / sigma_fn(t)) * sample - (-h * r).expm1() * model_output + + #using epsilon for new model output: + denoised_2 = x_2 - sigma_fn(s) * model_output + # 2. Convert to an ODE derivative for 2nd order + d = (x_2 - denoised_2) / sigma_fn(s) + dt = sd - sigma_next + sample = sample + d * dt + + del x_2 + del denoised_2 + del d + # Noise addition + if sigma_next > 0: + if self.config.use_noise_sampler: + sample = sample + self.noise_sampler(sigma, sigma_next) * self.config.s_noise * su + else: + sample = sample + noise * self.config.s_noise * su + elif self.config.algorithm_type == "dpmsolver++sde": + if self.config.solver_order == 2: + if sigma_next == 0: + # Euler method + d = (sample - model_output) / sigma + dt = sigma_next - sigma + sample = sample + d * dt + else: + # DPM-Solver++(SDE) + t, t_next = t_fn(sigma), t_fn(sigma_next) + r = self.config.midpoint_ratio + h = t_next - t + s = t + r * h + + # Euler method + d = (sample - model_output) / sigma + dt = sigma_next - sigma + sample = sample + d * dt + + # Step 1 + # get ancestral step + sigma_from = sigma_fn(t) + sigma_to = sigma_fn(s) + su = min(sigma_to, (sigma_to**2 * (sigma_from**2 - sigma_to**2) / sigma_from**2) ** 0.5) + sd = (sigma_to**2 - su**2) ** 0.5 + + # Euler method + d = (sample - model_output) / sigma + dt = sd - sigma + sample = sample + d * dt + + s_ = t_fn(sd) + x_2 = (sigma_fn(s_) / sigma_fn(t)) * sample - (t - s_).expm1() * model_output + if self.config.use_noise_sampler: + x_2 = x_2 + self.noise_sampler(sigma_fn(t), sigma_fn(s)) * self.config.s_noise * su + else: + x_2 = x_2 + noise * self.config.s_noise * su + + # Step 2 + # get ancestral step + sigma_from = sigma_fn(t) + sigma_to = sigma_fn(t_next) + su = min(sigma_to, (sigma_to**2 * (sigma_from**2 - sigma_to**2) / sigma_from**2) ** 0.5) + sd = (sigma_to**2 - su**2) ** 0.5 + + #using epsilon for new model output: + denoised_2 = x_2 - sigma_fn(s) * model_output + # 2. Convert to an ODE derivative for 2nd order + d = (x_2 - denoised_2) / sigma_fn(s) + dt = sd - sigma_next + sample = sample + d * dt + + if self.config.use_noise_sampler: + sample = sample + self.noise_sampler(sigma_fn(t), sigma_fn(t_next)) * self.config.s_noise * su + else: + sample = sample + noise * self.config.s_noise * su + del x_2 + del denoised_2 + del d + elif self.config.algorithm_type == "dpmsolver++2Msde": + if self.config.solver_order == 2: + if sigma_next == 0: + sample = model_output + else: + # DPM-Solver++(2M) SDE + t, s = -sigma.log(), -sigma_next.log() + h = s - t + eta_h = h * 1 + + # 3. Delta timestep + dt = sigma_next - sigma + sample = sample + model_output * dt + + sample = sigma_next / sigma * (-eta_h).exp() * sample + (-h - eta_h).expm1().neg() * model_output + + if self.model_outputs[-2] is not None: + r = self.h_last / h + if self.solver_type == 'heun': + sample = sample + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * (1 / r) * (model_output - self.model_outputs[-2]) + elif self.solver_type == 'midpoint': + sample = sample + 0.5 * (-h - eta_h).expm1().neg() * (1 / r) * (model_output - self.model_outputs[-2]) + + if self.config.use_noise_sampler: + sample = sample + self.noise_sampler(sigma, sigma_next) * sigma_next * (-2 * eta_h).expm1().neg().sqrt() * self.config.s_noise + else: + sample = sample + noise * sigma_next * (-2 * eta_h).expm1().neg().sqrt() * self.config.s_noise + + self.h_last = h + elif self.config.algorithm_type == "dpmsolver++3Msde": + if self.config.solver_order == 3: + if sigma_next == 0: + sample = model_output + else: + # DPM-Solver++(3M) SDE + t, s = -sigma.log(), -sigma_next.log() + h = s - t + h_eta = h * 2 + + # 3. Delta timestep + dt = sigma_next - sigma + sample = sample + model_output * dt + + sample = torch.exp(-h_eta) * sample + (-h_eta).expm1().neg() * model_output + + if self.h_2 is not None: + r0 = self.h_1 / h + r1 = self.h_2 / h + d1_0 = (model_output - self.model_outputs[-2]) / r0 + d1_1 = (self.model_outputs[-2] - self.model_outputs[-3]) / r1 + d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1) + d2 = (d1_0 - d1_1) / (r0 + r1) + phi_2 = h_eta.neg().expm1() / h_eta + 1 + phi_3 = phi_2 / h_eta - 0.5 + sample = sample + phi_2 * d1 - phi_3 * d2 + del d1_0 + del d1_1 + del d1 + del d2 + del phi_2 + del phi_3 + elif self.h_1 is not None: + r = self.h_1 / h + d = (model_output - self.model_outputs[-2]) / r + phi_2 = h_eta.neg().expm1() / h_eta + 1 + sample = sample + phi_2 * d + del d + del phi_2 + + if self.config.use_noise_sampler: + sample = sample + self.noise_sampler(sigma, sigma_next) * sigma_next * (-2 * h).expm1().neg().sqrt() * self.config.s_noise + else: + sample = sample + noise * sigma_next * (-2 * h).expm1().neg().sqrt() * self.config.s_noise + + self.h_2 = self.h_1 + self.h_1 = h + if not self.config.use_noise_sampler and noise is not None: + del noise + prev_sample = sample + + # Cast sample back to expected dtype + prev_sample = prev_sample.to(model_output.dtype) + + # upon completion increase step index by one + self._step_index += 1 + + torch.cuda.empty_cache() + + if not return_dict: + return (prev_sample,) + + return FlowMatchDPMSolverMultistepSchedulerOutput(prev_sample=prev_sample) + + def scale_model_input(self, sample: torch.Tensor, *args, **kwargs) -> torch.Tensor: + """ + Ensures interchangeability with schedulers that need to scale the denoising model input depending on the + current timestep. + + Args: + sample (`torch.Tensor`): + The input sample. + + Returns: + `torch.Tensor`: + A scaled input sample. + """ + return sample + + def scale_noise( + self, + sample: torch.FloatTensor, + timestep: Union[float, torch.FloatTensor], + noise: Optional[torch.FloatTensor] = None, + ) -> torch.FloatTensor: + """ + Forward process in flow-matching + + Args: + sample (`torch.FloatTensor`): + The input sample. + timestep (`int`, *optional*): + The current timestep in the diffusion chain. + + Returns: + `torch.FloatTensor`: + A scaled input sample. + """ + # Make sure sigmas and timesteps have the same device and dtype as original_samples + sigmas = self.sigmas.to(device=sample.device, dtype=sample.dtype) + + if sample.device.type == "mps" and torch.is_floating_point(timestep): + # mps does not support float64 + schedule_timesteps = self.timesteps.to(sample.device, dtype=torch.float32) + timestep = timestep.to(sample.device, dtype=torch.float32) + else: + schedule_timesteps = self.timesteps.to(sample.device) + timestep = timestep.to(sample.device) + + # self.begin_index is None when scheduler is used for training, or pipeline does not implement set_begin_index + if self.begin_index is None: + step_indices = [self.index_for_timestep(t, schedule_timesteps) for t in timestep] + elif self.step_index is not None: + # add_noise is called after first denoising step (for inpainting) + step_indices = [self.step_index] * timestep.shape[0] + else: + # add noise is called before first denoising step to create initial latent(img2img) + step_indices = [self.begin_index] * timestep.shape[0] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < len(sample.shape): + sigma = sigma.unsqueeze(-1) + + sample = sigma * noise + (1.0 - sigma) * sample + + return sample + + def __len__(self): + return self.config.num_train_timesteps diff --git a/modules/tcd/__init__.py b/modules/schedulers/scheduler_tcd.py similarity index 100% rename from modules/tcd/__init__.py rename to modules/schedulers/scheduler_tcd.py diff --git a/modules/vdm/__init__.py b/modules/schedulers/scheduler_vdm.py similarity index 100% rename from modules/vdm/__init__.py rename to modules/schedulers/scheduler_vdm.py diff --git a/modules/scripts.py b/modules/scripts.py index 8a67d0a50..cf2cf25b9 100644 --- a/modules/scripts.py +++ b/modules/scripts.py @@ -352,7 +352,9 @@ class ScriptRunner: self.selectable_scripts.clear() auto_processing_scripts = scripts_auto_postprocessing.create_auto_preprocessing_script_data() - for script_class, path, _basedir, _script_module in auto_processing_scripts + scripts_data: + all_scripts = auto_processing_scripts + scripts_data + sorted_scripts = sorted(all_scripts, key=lambda x: x.script_class().title().lower()) + for script_class, path, _basedir, _script_module in sorted_scripts: try: script = script_class() script.filename = path diff --git a/modules/sd_checkpoint.py b/modules/sd_checkpoint.py index e1787246f..afc5842e4 100644 --- a/modules/sd_checkpoint.py +++ b/modules/sd_checkpoint.py @@ -295,13 +295,7 @@ def read_metadata_from_safetensors(filename): 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': + if large and k in ['ss_datasets', 'workflow', 'prompt', 'ss_bucket_info', 'sd_metadata_file']: continue if v[0:1] == '{': try: diff --git a/modules/sd_detect.py b/modules/sd_detect.py index 148e788d5..31f773607 100644 --- a/modules/sd_detect.py +++ b/modules/sd_detect.py @@ -1,7 +1,8 @@ import os +import time import torch import diffusers -from modules import shared, shared_items, devices, errors +from modules import shared, shared_items, devices, errors, model_tools debug_load = os.environ.get('SD_LOAD_DEBUG', None) @@ -81,6 +82,9 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False): 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' @@ -103,6 +107,13 @@ def detect_pipeline(f: str, op: str = 'model', warning=True, quiet=False): 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: diff --git a/modules/sd_models.py b/modules/sd_models.py index 7743ce27d..e34f7c6f9 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -291,10 +291,13 @@ def set_diffuser_options(sd_model, vae = None, op: str = 'model', offload=True): def set_accelerate_to_module(model): - for k in model._internal_dict.keys(): # pylint: disable=protected-access - component = getattr(model, k, None) - if isinstance(component, torch.nn.Module): - component.has_accelerate = True + if hasattr(model, "pipe"): + set_accelerate_to_module(model.pipe) + if hasattr(model, "_internal_dict"): + for k in model._internal_dict.keys(): # pylint: disable=protected-access + component = getattr(model, k, None) + if isinstance(component, torch.nn.Module): + component.has_accelerate = True def set_accelerate(sd_model): @@ -397,7 +400,13 @@ def apply_balanced_offload(sd_model): return module def apply_balanced_offload_to_module(pipe): - for module_name in pipe._internal_dict.keys(): # pylint: disable=protected-access + if hasattr(pipe, "pipe"): + apply_balanced_offload_to_module(pipe.pipe) + if hasattr(pipe, "_internal_dict"): + keys = pipe._internal_dict.keys() # pylint: disable=protected-access + else: + keys = get_signature(shared.sd_model).keys() + for module_name in keys: # pylint: disable=protected-access module = getattr(pipe, module_name, None) if isinstance(module, torch.nn.Module): checkpoint_name = pipe.sd_checkpoint_info.name if getattr(pipe, "sd_checkpoint_info", None) is not None else None @@ -416,6 +425,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"): @@ -440,6 +451,9 @@ def move_model(model, device=None, force=False): devices.torch_gc() return + if hasattr(model, 'pipe'): + move_model(model.pipe, device, force) + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access if getattr(model, 'vae', None) is not None and get_diffusers_task(model) != DiffusersTaskType.TEXT_2_IMAGE: if device == devices.device and model.vae.device.type != "meta": # force vae back to gpu if not in txt2img mode @@ -467,7 +481,8 @@ def move_model(model, device=None, force=False): try: t0 = time.time() try: - model.to(device) + if hasattr(model, 'to'): + model.to(device) if hasattr(model, "prior_pipe"): model.prior_pipe.to(device) except Exception as e0: @@ -477,7 +492,8 @@ def move_model(model, device=None, force=False): if hasattr(component, 'modules'): for module in component.modules(): try: - module.to(device) + if hasattr(module, 'to'): + module.to(device) except Exception as e2: if 'Cannot copy out of meta tensor' in str(e2): if os.environ.get('SD_MOVE_DEBUG', None): @@ -771,11 +787,11 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No if shared.opts.data.get('sd_model_checkpoint', '') == 'model.safetensors' or shared.opts.data.get('sd_model_checkpoint', '') == '': shared.opts.data['sd_model_checkpoint'] = "stabilityai/stable-diffusion-xl-base-1.0" - if op == 'model' or op == 'dict': - if (model_data.sd_model is not None) and (checkpoint_info is not None) and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model + if (op == 'model' or op == 'dict'): + if (model_data.sd_model is not None) and (checkpoint_info is not None) and (getattr(model_data.sd_model, 'sd_checkpoint_info', None) is not None) and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model return else: - if (model_data.sd_refiner is not None) and (checkpoint_info is not None) and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model + if (model_data.sd_refiner is not None) and (checkpoint_info is not None) and (getattr(model_data.sd_refiner, 'sd_checkpoint_info', None) is not None) and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model return sd_model = None @@ -797,6 +813,9 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No # 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) @@ -862,8 +881,9 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No sd_model.embedding_db.load_textual_inversion_embeddings(force_reload=True) timer.record("embeddings") - from modules.prompt_parser_diffusers import insert_parser_highjack - insert_parser_highjack(sd_model.__class__.__name__) + from modules import prompt_parser_diffusers + prompt_parser_diffusers.insert_parser_highjack(sd_model.__class__.__name__) + prompt_parser_diffusers.cache.clear() set_diffuser_options(sd_model, vae, op, offload=False) if shared.opts.nncf_compress_weights and not ('Model' in shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx"): @@ -874,7 +894,8 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No set_diffuser_offload(sd_model, op) if op == 'model' and not (os.path.isdir(checkpoint_info.path) or checkpoint_info.type == 'huggingface'): - sd_vae.apply_vae_config(shared.sd_model.sd_checkpoint_info.filename, vae_file, sd_model) + if getattr(shared.sd_model, 'sd_checkpoint_info', None) is not None: + sd_vae.apply_vae_config(shared.sd_model.sd_checkpoint_info.filename, vae_file, sd_model) if op == 'refiner' and shared.opts.diffusers_move_refiner: shared.log.debug('Moving refiner model to CPU') move_model(sd_model, devices.cpu) @@ -1049,6 +1070,7 @@ def set_diffuser_pipe(pipe, new_pipe_type): 'AnimateDiffSDXLPipeline', 'OmniGenPipeline', 'StableDiffusion3ControlNetPipeline', + 'InstantIRPipeline', ] n = getattr(pipe.__class__, '__name__', '') @@ -1059,12 +1081,15 @@ def set_diffuser_pipe(pipe, new_pipe_type): return pipe # skip specific pipelines + cls = pipe.__class__.__name__ if n in exclude: return pipe - if 'Onnx' in pipe.__class__.__name__: + if 'Onnx' in cls: return pipe - if new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE or new_pipe_type == DiffusersTaskType.INPAINTING: # in some cases we want to reset the pipeline as they dont have their own variants + new_pipe = None + # in some cases we want to reset the pipeline to parent as they dont have their own variants + if new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE or new_pipe_type == DiffusersTaskType.INPAINTING: if n == 'StableDiffusionPAGPipeline': pipe = switch_pipe(diffusers.StableDiffusionPipeline, pipe) if n == 'StableDiffusionXLPAGPipeline': @@ -1080,19 +1105,38 @@ def set_diffuser_pipe(pipe, new_pipe_type): image_encoder = getattr(pipe, "image_encoder", None) feature_extractor = getattr(pipe, "feature_extractor", None) - try: - if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: - new_pipe = diffusers.AutoPipelineForText2Image.from_pipe(pipe) - elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE: - new_pipe = diffusers.AutoPipelineForImage2Image.from_pipe(pipe) - elif new_pipe_type == DiffusersTaskType.INPAINTING: - new_pipe = diffusers.AutoPipelineForInpainting.from_pipe(pipe) + if new_pipe is None: + if hasattr(pipe, 'config'): # real pipeline which can be auto-switched + try: + if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: + new_pipe = diffusers.AutoPipelineForText2Image.from_pipe(pipe) + elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE: + new_pipe = diffusers.AutoPipelineForImage2Image.from_pipe(pipe) + elif new_pipe_type == DiffusersTaskType.INPAINTING: + new_pipe = diffusers.AutoPipelineForInpainting.from_pipe(pipe) + else: + shared.log.error(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls}') + return pipe + except Exception as e: # pylint: disable=unused-variable + shared.log.warning(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls} {e}') + return pipe else: - shared.log.error(f'Pipeline class change failed: type={new_pipe_type} pipeline={pipe.__class__.__name__}') - return pipe - except Exception as e: # pylint: disable=unused-variable - shared.log.warning(f'Pipeline class change failed: type={new_pipe_type} pipeline={pipe.__class__.__name__} {e}') - return pipe + try: # maybe a wrapper pipeline so just change the class + if new_pipe_type == DiffusersTaskType.TEXT_2_IMAGE: + pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_TEXT2IMAGE_PIPELINES_MAPPING, cls) # pylint: disable=protected-access + new_pipe = pipe + elif new_pipe_type == DiffusersTaskType.IMAGE_2_IMAGE: + pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING, cls) # pylint: disable=protected-access + new_pipe = pipe + elif new_pipe_type == DiffusersTaskType.INPAINTING: + pipe.__class__ = diffusers.pipelines.auto_pipeline._get_task_class(diffusers.pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING, cls) # pylint: disable=protected-access + new_pipe = pipe + else: + shared.log.error(f'Pipeline class change failed: type={new_pipe_type} pipeline={cls}') + return pipe + except Exception as e: # pylint: disable=unused-variable + shared.log.warning(f'Pipeline class set failed: type={new_pipe_type} pipeline={cls} {e}') + return pipe # if pipe.__class__ == new_pipe.__class__: # return pipe @@ -1108,10 +1152,14 @@ def set_diffuser_pipe(pipe, new_pipe_type): new_pipe.is_sdxl = getattr(pipe, 'is_sdxl', False) # a1111 compatibility item new_pipe.is_sd2 = getattr(pipe, 'is_sd2', False) new_pipe.is_sd1 = getattr(pipe, 'is_sd1', True) - if hasattr(new_pipe, "watermark"): + if hasattr(new_pipe, 'watermark'): new_pipe.watermark = NoWatermark() + + if hasattr(new_pipe, 'pipe'): # also handle nested pipelines + new_pipe.pipe = set_diffuser_pipe(new_pipe.pipe, new_pipe_type) + fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access - shared.log.debug(f"Pipeline class change: original={pipe.__class__.__name__} target={new_pipe.__class__.__name__} device={pipe.device} fn={fn}") # pylint: disable=protected-access + shared.log.debug(f"Pipeline class change: original={cls} target={new_pipe.__class__.__name__} device={pipe.device} fn={fn}") # pylint: disable=protected-access pipe = new_pipe return pipe @@ -1140,6 +1188,9 @@ def set_diffusers_attention(pipe): else: module.set_attn_processor(attention) + # if hasattr(pipe, 'pipe'): + # set_diffusers_attention(pipe.pipe) + if 'ControlNet' in pipe.__class__.__name__: # do not replace attention in ControlNet pipelines return shared.log.debug(f'Setting model: attention="{shared.opts.cross_attention_optimization}"') @@ -1180,10 +1231,10 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None, if checkpoint_info is None: return if op == 'model' or op == 'dict': - if model_data.sd_model is not None and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model + if (model_data.sd_model is not None) and (getattr(model_data.sd_model, 'sd_checkpoint_info', None) is not None) and (checkpoint_info.hash == model_data.sd_model.sd_checkpoint_info.hash): # trying to load the same model return else: - if model_data.sd_refiner is not None and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model + if (model_data.sd_refiner is not None) and (getattr(model_data.sd_refiner, 'sd_checkpoint_info', None) is not None) and (checkpoint_info.hash == model_data.sd_refiner.sd_checkpoint_info.hash): # trying to load the same model return shared.log.debug(f'Load {op}: name={checkpoint_info.filename} dict={already_loaded_state_dict is not None}') if timer is None: @@ -1192,12 +1243,12 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None, if op == 'model' or op == 'dict': if model_data.sd_model is not None: sd_hijack.model_hijack.undo_hijack(model_data.sd_model) - current_checkpoint_info = model_data.sd_model.sd_checkpoint_info + current_checkpoint_info = getattr(model_data.sd_model, 'sd_checkpoint_info', None) unload_model_weights(op=op) else: if model_data.sd_refiner is not None: sd_hijack.model_hijack.undo_hijack(model_data.sd_refiner) - current_checkpoint_info = model_data.sd_refiner.sd_checkpoint_info + current_checkpoint_info = getattr(model_data.sd_refiner, 'sd_checkpoint_info', None) unload_model_weights(op=op) if not shared.native: @@ -1226,15 +1277,6 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None, sd_model = instantiate_from_config(sd_config.model) else: with contextlib.redirect_stdout(stdout): - """ - try: - clip_is_included_into_sd = sd1_clip_weight in state_dict or sd2_clip_weight in state_dict - with sd_disable_initialization.DisableInitialization(disable_clip=clip_is_included_into_sd): - sd_model = instantiate_from_config(sd_config.model) - except Exception as e: - shared.log.error(f'LDM: instantiate from config: {e}') - sd_model = instantiate_from_config(sd_config.model) - """ sd_model = instantiate_from_config(sd_config.model) for line in stdout.getvalue().splitlines(): if len(line) > 0: diff --git a/modules/sd_samplers.py b/modules/sd_samplers.py index bbc2f360b..e560744dd 100644 --- a/modules/sd_samplers.py +++ b/modules/sd_samplers.py @@ -47,6 +47,10 @@ def visible_sampler_names(): def create_sampler(name, model): + try: + current = model.scheduler.__class__.__name__ + except Exception: + current = None if name == 'Default' and hasattr(model, 'scheduler'): if getattr(model, "default_scheduler", None) is not None: model.scheduler = copy.deepcopy(model.default_scheduler) @@ -54,12 +58,13 @@ def create_sampler(name, model): model.prior_pipe.scheduler = copy.deepcopy(model.default_scheduler) model.prior_pipe.scheduler.config.clip_sample = False config = {k: v for k, v in model.scheduler.config.items() if not k.startswith('_')} - shared.log.debug(f'Sampler: sampler=default class={model.scheduler.__class__.__name__}: {config}') + shared.log.debug(f'Sampler: sampler=default class={current}: {config}') return model.scheduler config = find_sampler_config(name) if config is None or config.constructor is None: # shared.log.warning(f'Sampler: sampler="{name}" not found') return None + sampler = None if not shared.native: sampler = config.constructor(model) sampler.config = config @@ -68,24 +73,18 @@ def create_sampler(name, model): shared.log.debug(f'Sampler: sampler="{name}" config={config.options}') return sampler elif shared.native: - sampler = config.constructor(model) - if 'Flux' in model.__class__.__name__: - if 'base_image_seq_len' not in sampler.sampler.config or 'max_image_seq_len' not in sampler.sampler.config or 'base_shift' not in sampler.sampler.config or 'max_shift' not in sampler.sampler.config: - shared.log.warning(f'FLUX: sampler="{name}" unsupported') - # sampler.sampler.register_to_config(base_image_seq_len=256, max_image_seq_len=4096, base_shift=0.5, max_shift=1.15) - return None - if 'Lumina' in model.__class__.__name__: - shared.log.warning(f'AlphaVLLM-Lumina: sampler="{name}" unsupported') - return None - if 'StableDiffusion3' in model.__class__.__name__: - if sampler.name != 'Heun FlowMatch': - return None - return None - if 'AuraFlow' in model.__class__.__name__: - shared.log.warning(f'AuraFlow: sampler="{name}" unsupported') - return None + FlowModels = ['Flux', 'StableDiffusion3', 'Lumina', 'AuraFlow'] if 'KDiffusion' in model.__class__.__name__: return None + if any(x in model.__class__.__name__ for x in FlowModels) and 'FlowMatch' not in name: + shared.log.warning(f'Sampler: default={current} target="{name}" class={model.__class__.__name__} linear scheduler unsupported') + return None + if not any(x in model.__class__.__name__ for x in FlowModels) and 'FlowMatch' in name: + shared.log.warning(f'Sampler: default={current} target="{name}" class={model.__class__.__name__} flow-match scheduler unsupported') + return None + sampler = config.constructor(model) + if sampler is None: + sampler = config.constructor(model) 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_diffusers.py b/modules/sd_samplers_diffusers.py index c7eedf59b..370cb767b 100644 --- a/modules/sd_samplers_diffusers.py +++ b/modules/sd_samplers_diffusers.py @@ -1,16 +1,16 @@ import os -import copy import re +import copy import inspect +import diffusers from modules import shared, errors from modules import sd_samplers_common -from modules.tcd import TCDScheduler -from modules.dcsolver import DCSolverMultistepScheduler -from modules.vdm import VDMScheduler + debug = shared.log.trace if os.environ.get('SD_SAMPLER_DEBUG', None) is not None else lambda *args, **kwargs: None debug('Trace: SAMPLER') + try: from diffusers import ( CMStochasticIterativeScheduler, @@ -44,7 +44,15 @@ try: KDPM2AncestralDiscreteScheduler, ) except Exception as e: - import diffusers + shared.log.error(f'Diffusers import error: version={diffusers.__version__} error: {e}') + if os.environ.get('SD_SAMPLER_DEBUG', None) is not None: + errors.display(e, 'Samplers') +try: + from modules.schedulers.scheduler_tcd import TCDScheduler # pylint: disable=ungrouped-imports + from modules.schedulers.scheduler_dc import DCSolverMultistepScheduler # pylint: disable=ungrouped-imports + from modules.schedulers.scheduler_vdm import VDMScheduler # pylint: disable=ungrouped-imports + from modules.schedulers.scheduler_dpm_flowmatch import FlowMatchDPMSolverMultistepScheduler # pylint: disable=ungrouped-imports +except Exception as e: shared.log.error(f'Diffusers import error: version={diffusers.__version__} error: {e}') if os.environ.get('SD_SAMPLER_DEBUG', None) is not None: errors.display(e, 'Samplers') @@ -63,33 +71,40 @@ config = { 'Euler EDM': { 'sigma_schedule': "karras" }, 'Euler FlowMatch': { 'timestep_spacing': "linspace", 'shift': 1, 'use_dynamic_shifting': False }, - 'DPM++': { 'solver_order': 2, 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'final_sigmas_type': 'sigma_min' }, - 'DPM++ 1S': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 1 }, - 'DPM++ 2M': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 2 }, - 'DPM++ 3M': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 3 }, - 'DPM++ 2M SDE': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "sde-dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 2 }, + 'DPM++': { 'solver_order': 2, 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'final_sigmas_type': 'sigma_min' }, + 'DPM++ 1S': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'use_lu_lambdas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 1 }, + 'DPM++ 2M': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'use_lu_lambdas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 2 }, + 'DPM++ 3M': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'use_lu_lambdas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 3 }, + 'DPM++ 2M SDE': { 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "sde-dpmsolver++", 'solver_type': "midpoint", 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'use_lu_lambdas': False, 'final_sigmas_type': 'zero', 'timestep_spacing': 'linspace', 'solver_order': 2 }, 'DPM++ 2M EDM': { 'solver_order': 2, 'solver_type': 'midpoint', 'final_sigmas_type': 'zero', 'algorithm_type': 'dpmsolver++' }, 'DPM++ Cosine': { 'solver_order': 2, 'sigma_schedule': "exponential", 'prediction_type': "v-prediction" }, - 'DPM SDE': { 'use_karras_sigmas': False, 'noise_sampler_seed': None, 'timestep_spacing': 'linspace', 'steps_offset': 0 }, + 'DPM SDE': { 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'noise_sampler_seed': None, 'timestep_spacing': 'linspace', 'steps_offset': 0, }, - 'Heun': { 'use_beta_sigmas': False, 'use_karras_sigmas': False, 'timestep_spacing': 'linspace' }, + 'DPM2 FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver2', 'use_noise_sampler': True }, + 'DPM2a FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver2A', 'use_noise_sampler': True }, + 'DPM2++ 2M FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver++2M', 'use_noise_sampler': True }, + 'DPM2++ 2S FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver++2S', 'use_noise_sampler': True }, + 'DPM2++ SDE FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver++sde', 'use_noise_sampler': True }, + 'DPM2++ 2M SDE FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 2, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver++2Msde', 'use_noise_sampler': True }, + 'DPM2++ 3M SDE FlowMatch': { 'shift': 1, 'use_dynamic_shifting': False, 'solver_order': 3, 'sigma_schedule': None, 'use_SD35_sigmas': False, 'algorithm_type': 'dpmsolver++3Msde', 'use_noise_sampler': True }, + + 'Heun': { 'use_beta_sigmas': False, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'timestep_spacing': 'linspace' }, 'Heun FlowMatch': { 'timestep_spacing': "linspace", 'shift': 1 }, - 'DEIS': { 'solver_order': 2, 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "deis", 'solver_type': "logrho", 'lower_order_final': True, 'timestep_spacing': 'linspace' }, - 'SA Solver': {'predictor_order': 2, 'corrector_order': 2, 'thresholding': False, 'lower_order_final': True, 'use_karras_sigmas': False, 'timestep_spacing': 'linspace'}, + 'DEIS': { 'solver_order': 2, 'thresholding': False, 'sample_max_value': 1.0, 'algorithm_type': "deis", 'solver_type': "logrho", 'lower_order_final': True, 'timestep_spacing': 'linspace', 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False }, + 'SA Solver': {'predictor_order': 2, 'corrector_order': 2, 'thresholding': False, 'lower_order_final': True, 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'timestep_spacing': 'linspace'}, 'DC Solver': { 'beta_start': 0.0001, 'beta_end': 0.02, 'solver_order': 2, 'prediction_type': "epsilon", 'thresholding': False, 'solver_type': 'bh2', 'lower_order_final': True, 'dc_order': 2, 'disable_corrector': [0] }, 'VDM Solver': { 'clip_sample_range': 2.0, }, - 'LCM': { 'beta_start': 0.00085, 'beta_end': 0.012, 'beta_schedule': "scaled_linear", 'set_alpha_to_one': True, 'rescale_betas_zero_snr': False, 'thresholding': False, 'timestep_spacing': 'linspace' }, 'TCD': { 'set_alpha_to_one': True, 'rescale_betas_zero_snr': False, 'beta_schedule': 'scaled_linear' }, 'PNDM': { 'skip_prk_steps': False, 'set_alpha_to_one': False, 'steps_offset': 0, 'timestep_spacing': 'linspace' }, 'IPNDM': { }, 'DDPM': { 'variance_type': "fixed_small", 'clip_sample': False, 'thresholding': False, 'clip_sample_range': 1.0, 'sample_max_value': 1.0, 'timestep_spacing': 'linspace', 'rescale_betas_zero_snr': False }, - 'LMSD': { 'use_karras_sigmas': False, 'timestep_spacing': 'linspace', 'steps_offset': 0 }, - 'KDPM2': { 'steps_offset': 0, 'timestep_spacing': 'linspace' }, - 'KDPM2 a': { 'steps_offset': 0, 'timestep_spacing': 'linspace' }, - 'CMSI': { }, #{ 'sigma_min': 0.002, 'sigma_max': 80.0, 'sigma_data': 0.5, 's_noise': 1.0, 'rho': 7.0, 'clip_denoised': True }, + 'LMSD': { 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'timestep_spacing': 'linspace', 'steps_offset': 0 }, + 'KDPM2': { 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'steps_offset': 0, 'timestep_spacing': 'linspace' }, + 'KDPM2 a': { 'use_karras_sigmas': False, 'use_exponential_sigmas': False, 'use_beta_sigmas': False, 'steps_offset': 0, 'timestep_spacing': 'linspace' }, + 'CMSI': { }, } samplers_data_diffusers = [ @@ -112,6 +127,14 @@ samplers_data_diffusers = [ sd_samplers_common.SamplerData('DPM++ Cosine', lambda model: DiffusionSampler('DPM++ 2M EDM', CosineDPMSolverMultistepScheduler, model), [], {}), sd_samplers_common.SamplerData('DPM SDE', lambda model: DiffusionSampler('DPM SDE', DPMSolverSDEScheduler, model), [], {}), + sd_samplers_common.SamplerData('DPM2 FlowMatch', lambda model: DiffusionSampler('DPM2 FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}), + sd_samplers_common.SamplerData('DPM2a FlowMatch', lambda model: DiffusionSampler('DPM2a FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}), + sd_samplers_common.SamplerData('DPM2++ 2M FlowMatch', lambda model: DiffusionSampler('DPM2++ 2M FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}), + sd_samplers_common.SamplerData('DPM2++ 2S FlowMatch', lambda model: DiffusionSampler('DPM2++ 2S FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}), + sd_samplers_common.SamplerData('DPM2++ SDE FlowMatch', lambda model: DiffusionSampler('DPM2++ SDE FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}), + sd_samplers_common.SamplerData('DPM2++ 2M SDE FlowMatch', lambda model: DiffusionSampler('DPM2++ 2M SDE FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}), + sd_samplers_common.SamplerData('DPM2++ 3M SDE FlowMatch', lambda model: DiffusionSampler('DPM2++ 3M SDE FlowMatch', FlowMatchDPMSolverMultistepScheduler, model), [], {}), + sd_samplers_common.SamplerData('Heun', lambda model: DiffusionSampler('Heun', HeunDiscreteScheduler, model), [], {}), sd_samplers_common.SamplerData('Heun FlowMatch', lambda model: DiffusionSampler('Heun FlowMatch', FlowMatchHeunDiscreteScheduler, model), [], {}), @@ -183,6 +206,10 @@ class DiffusionSampler: self.config['use_karras_sigmas'] = shared.opts.schedulers_sigma == 'karras' if 'use_exponential_sigmas' in self.config: self.config['use_exponential_sigmas'] = shared.opts.schedulers_sigma == 'exponential' + if 'use_lu_lambdas' in self.config: + self.config['use_lu_lambdas'] = shared.opts.schedulers_sigma == 'lambdas' + if 'sigma_schedule' in self.config: + self.config['sigma_schedule'] = shared.opts.schedulers_sigma if shared.opts.schedulers_sigma != 'default' else None else: pass # timesteps are set using set_timesteps in set_pipeline_args @@ -199,9 +226,18 @@ class DiffusionSampler: if 'beta_end' in self.config and shared.opts.schedulers_beta_end > 0: self.config['beta_end'] = shared.opts.schedulers_beta_end if 'shift' in self.config: - self.config['shift'] = shared.opts.schedulers_shift + if shared.opts.schedulers_shift == 0: + if 'StableDiffusion3' in model.__class__.__name__: + self.config['shift'] = 3 + if 'Flux' in model.__class__.__name__: + self.config['shift'] = 1 + else: + self.config['shift'] = shared.opts.schedulers_shift if 'use_dynamic_shifting' in self.config: - self.config['use_dynamic_shifting'] = shared.opts.schedulers_dynamic_shift + if 'Flux' in model.__class__.__name__: + self.config['use_dynamic_shifting'] = shared.opts.schedulers_dynamic_shift + if 'use_SD35_sigmas' in self.config: + self.config['use_SD35_sigmas'] = 'StableDiffusion3' in model.__class__.__name__ if 'rescale_betas_zero_snr' in self.config: self.config['rescale_betas_zero_snr'] = shared.opts.schedulers_rescale_betas if 'timestep_spacing' in self.config and shared.opts.schedulers_timestep_spacing != 'default' and shared.opts.schedulers_timestep_spacing is not None: diff --git a/modules/sd_vae.py b/modules/sd_vae.py index f266f8c38..95ac05c93 100644 --- a/modules/sd_vae.py +++ b/modules/sd_vae.py @@ -289,7 +289,7 @@ def reload_vae_weights(sd_model=None, vae_file=unspecified): if vae_file is not None: shared.log.info(f"VAE weights loaded: {vae_file}") else: - if hasattr(sd_model, "vae") and hasattr(sd_model, "sd_checkpoint_info"): + if hasattr(sd_model, "vae") and getattr(sd_model, "sd_checkpoint_info", None) is not None: vae = load_vae_diffusers(sd_model.sd_checkpoint_info.filename, vae_file, vae_source) if vae is not None: if not hasattr(sd_model, 'original_vae'): diff --git a/modules/shared.py b/modules/shared.py index c58110214..0bd893c9e 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -1,3 +1,4 @@ +from functools import lru_cache import io import os import sys @@ -83,6 +84,8 @@ console = Console(log_time=True, log_time_format='%H:%M:%S-%f') dir_timestamps = {} dir_cache = {} max_workers = 8 +if os.environ.get("SD_HFCACHEDIR", None) is not None: + hfcache_dir = os.environ.get("SD_HFCACHEDIR") if os.environ.get("HF_HUB_CACHE", None) is not None: hfcache_dir = os.environ.get("HF_HUB_CACHE") elif os.environ.get("HF_HUB", None) is not None: @@ -249,6 +252,7 @@ class OptionInfo: self.comment_before = comment_before # HTML text that will be added after label in UI self.comment_after = comment_after # HTML text that will be added before label in UI self.submit = submit + self.exclude = ['sd_model_checkpoint', 'sd_model_refiner', 'sd_vae', 'sd_unet', 'sd_text_encoder', 'sd_model_dict'] def needs_reload_ui(self): return self @@ -273,6 +277,42 @@ class OptionInfo: self.comment_after += " (requires restart)" return self + def validate(self, opt, value): + if opt in self.exclude: + return True + args = self.component_args if self.component_args is not None else {} + if callable(args): + try: + args = args() + except Exception: + args = {} + choices = args.get("choices", []) + if callable(choices): + try: + choices = choices() + except Exception: + choices = [] + if len(choices) > 0: + if not isinstance(value, list): + value = [value] + for v in value: + if v not in choices: + log.warning(f'Setting validation: "{opt}"="{v}" default="{self.default}" choices={choices}') + return False + minimum = args.get("minimum", None) + maximum = args.get("maximum", None) + if (minimum is not None and value < minimum) or (maximum is not None and value > maximum): + log.error(f'Setting validation: "{opt}"={value} default={self.default} minimum={minimum} maximum={maximum}') + return False + return True + + def __str__(self) -> str: + args = self.component_args if self.component_args is not None else {} + if callable(args): + args = args() + choices = args.get("choices", []) + return f'OptionInfo: label="{self.label}" section="{self.section}" component="{self.component}" default="{self.default}" refresh="{self.refresh is not None}" change="{self.onchange is not None}" args={args} choices={choices}' + def options_section(section_identifier, options_dict): for v in options_dict.values(): @@ -436,17 +476,14 @@ options_templates.update(options_section(('sd', "Execution & Models"), { "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"), + "sd_textencoder_cache": OptionInfo(True, "Cache text encoder results", gr.Checkbox, {"visible": False}), + "sd_textencoder_cache_size": OptionInfo(4, "Text encoder results LRU cache size", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1}), "stream_load": OptionInfo(False, "Load models using stream loading method", gr.Checkbox, {"visible": not native }), - "model_reuse_dict": OptionInfo(False, "Reuse loaded model dictionary", gr.Checkbox, {"visible": False}), "prompt_mean_norm": OptionInfo(False, "Prompt attention normalization", gr.Checkbox), "comma_padding_backtrack": OptionInfo(20, "Prompt padding", gr.Slider, {"minimum": 0, "maximum": 74, "step": 1, "visible": not native }), - "prompt_attention": OptionInfo("Full parser", "Prompt attention parser", gr.Radio, {"choices": ["Full parser", "Compel parser", "xhinker parser", "A1111 parser", "Fixed attention"] }), + "prompt_attention": OptionInfo("native", "Prompt attention parser", gr.Radio, {"choices": ["native", "compel", "xhinker", "a1111", "fixed"] }), "latent_history": OptionInfo(16, "Latent history size", gr.Slider, {"minimum": 1, "maximum": 100, "step": 1}), "sd_checkpoint_cache": OptionInfo(0, "Cached models", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1, "visible": not native }), - "sd_vae_checkpoint_cache": OptionInfo(0, "Cached VAEs", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1, "visible": False}), - "sd_disable_ckpt": OptionInfo(False, "Disallow models in ckpt format", gr.Checkbox, {"visible": False}), - "diffusers_version": OptionInfo("", "Diffusers version", gr.Textbox, {"visible": False}), })) options_templates.update(options_section(('cuda', "Compute Settings"), { @@ -460,8 +497,7 @@ options_templates.update(options_section(('cuda', "Compute Settings"), { "upcast_sampling": OptionInfo(False if sys.platform != "darwin" else True, "Upcast sampling"), "upcast_attn": OptionInfo(False, "Upcast attention layer"), "cuda_cast_unet": OptionInfo(False, "Fixed UNet precision"), - "disable_nan_check": OptionInfo(True, "Disable NaN check", gr.Checkbox, {"visible": False}), - "nan_skip": OptionInfo(False, "Skip Generation if NaN found in latents", gr.Checkbox, {"visible": True}), + "nan_skip": OptionInfo(False, "Skip Generation if NaN found in latents", gr.Checkbox), "rollback_vae": OptionInfo(False, "Attempt VAE roll back for NaN values"), "cross_attention_sep": OptionInfo("

Cross Attention

", "", gr.HTML), @@ -534,7 +570,6 @@ options_templates.update(options_section(('diffusers', "Diffusers Settings"), { "diffusers_eval": OptionInfo(True, "Force model eval"), "diffusers_to_gpu": OptionInfo(False, "Load model directly to GPU"), "disable_accelerate": OptionInfo(False, "Disable accelerate"), - "diffusers_force_zeros": OptionInfo(False, "Force zeros for prompts when empty", gr.Checkbox, {"visible": False}), "diffusers_pooled": OptionInfo("default", "Diffusers SDXL pooled embeds", gr.Radio, {"choices": ['default', 'weighted']}), "diffusers_zeros_prompt_pad": OptionInfo(False, "Use zeros for prompt padding", gr.Checkbox), "huggingface_token": OptionInfo('', 'HuggingFace token'), @@ -615,9 +650,7 @@ options_templates.update(options_section(('system-paths', "System Paths"), { "vae_dir": OptionInfo(os.path.join(paths.models_path, 'VAE'), "Folder with VAE files", folder=True), "unet_dir": OptionInfo(os.path.join(paths.models_path, 'UNET'), "Folder with UNET files", folder=True), "te_dir": OptionInfo(os.path.join(paths.models_path, 'Text-encoder'), "Folder with Text encoder files", folder=True), - "sd_lora": OptionInfo("", "Add LoRA to prompt", gr.Textbox, {"visible": False}), "lora_dir": OptionInfo(os.path.join(paths.models_path, 'Lora'), "Folder with LoRA network(s)", folder=True), - "lyco_dir": OptionInfo(os.path.join(paths.models_path, 'LyCORIS'), "Folder with LyCORIS network(s)", gr.Text, {"visible": False}), "styles_dir": OptionInfo(os.path.join(paths.data_path, 'styles.csv'), "File or Folder with user-defined styles", folder=True), "wildcards_dir": OptionInfo(os.path.join(paths.models_path, 'wildcards'), "Folder with user-defined wildcards", folder=True), "embeddings_dir": OptionInfo(os.path.join(paths.models_path, 'embeddings'), "Folder with textual inversion embeddings", folder=True), @@ -693,12 +726,10 @@ options_templates.update(options_section(('saving-paths', "Image Naming & Paths" "saving_sep_images": OptionInfo("

Save options

", "", gr.HTML), "save_images_add_number": OptionInfo(True, "Numbered filenames", component_args=hide_dirs), "use_original_name_batch": OptionInfo(True, "Batch uses original name"), - "use_upscaler_name_as_suffix": OptionInfo(True, "Use upscaler as suffix", gr.Checkbox, {"visible": False}), "save_to_dirs": OptionInfo(False, "Save images to a subdirectory"), "directories_filename_pattern": OptionInfo("[date]", "Directory name pattern", component_args=hide_dirs), "samples_filename_pattern": OptionInfo("[seq]-[model_name]-[prompt_words]", "Images filename pattern", component_args=hide_dirs), "directories_max_prompt_words": OptionInfo(8, "Max words per pattern", gr.Slider, {"minimum": 1, "maximum": 99, "step": 1, **hide_dirs}), - "use_save_to_dirs_for_ui": OptionInfo(False, "Save images to a subdirectory when using Save button", gr.Checkbox, {"visible": False}), "outdir_sep_dirs": OptionInfo("

Folders

", "", gr.HTML), "outdir_samples": OptionInfo("", "Images folder", component_args=hide_dirs, folder=True), @@ -711,8 +742,6 @@ options_templates.update(options_section(('saving-paths', "Image Naming & Paths" "outdir_init_images": OptionInfo("outputs/init-images", "Folder for init images", component_args=hide_dirs, folder=True), "outdir_sep_grids": OptionInfo("

Grids

", "", gr.HTML), - "grid_extended_filename": OptionInfo(True, "Add extended info to filename when saving grid", gr.Checkbox, {"visible": False}), - "grid_save_to_dirs": OptionInfo(False, "Save grids to a subdirectory", gr.Checkbox, {"visible": False}), "outdir_grids": OptionInfo("", "Grids folder", component_args=hide_dirs, folder=True), "outdir_txt2img_grids": OptionInfo("outputs/grids", 'Folder for txt2img grids', component_args=hide_dirs, folder=True), "outdir_img2img_grids": OptionInfo("outputs/grids", 'Folder for img2img grids', component_args=hide_dirs, folder=True), @@ -725,7 +754,6 @@ options_templates.update(options_section(('ui', "User Interface Options"), { "gradio_theme": OptionInfo("black-teal", "UI theme", gr.Dropdown, lambda: {"choices": theme.list_themes()}, refresh=theme.refresh_themes), "autolaunch": OptionInfo(False, "Autolaunch browser upon startup"), "font_size": OptionInfo(14, "Font size", gr.Slider, {"minimum": 8, "maximum": 32, "step": 1, "visible": True}), - "tooltips": OptionInfo("UI Tooltips", "UI tooltips", gr.Radio, {"choices": ["None", "Browser default", "UI tooltips"], "visible": False}), "aspect_ratios": OptionInfo("1:1, 4:3, 3:2, 16:9, 16:10, 21:9, 2:3, 3:4, 9:16, 10:16, 9:21", "Allowed aspect ratios"), "motd": OptionInfo(True, "Show MOTD"), "compact_view": OptionInfo(False, "Compact view"), @@ -735,22 +763,14 @@ options_templates.update(options_section(('ui', "User Interface Options"), { "disable_weights_auto_swap": OptionInfo(True, "Do not change selected model when reading generation parameters"), "send_seed": OptionInfo(True, "Send seed when sending prompt or image to other interface"), "send_size": OptionInfo(True, "Send size when sending prompt or image to another interface"), - "keyedit_precision_attention": OptionInfo(0.1, "Ctrl+up/down precision when editing (attention:1.1)", gr.Slider, {"minimum": 0.01, "maximum": 0.2, "step": 0.001, "visible": False}), - "keyedit_precision_extra": OptionInfo(0.05, "Ctrl+up/down precision when editing ", gr.Slider, {"minimum": 0.01, "maximum": 0.2, "step": 0.001, "visible": False}), - "keyedit_delimiters": OptionInfo(r".,\/!?%^*;:{}=`~()", "Ctrl+up/down word delimiters", gr.Textbox, { "visible": False }), "quicksettings_list": OptionInfo(["sd_model_checkpoint"], "Quicksettings list", gr.Dropdown, lambda: {"multiselect":True, "choices": list(opts.data_labels.keys())}), - "ui_scripts_reorder": OptionInfo("", "UI scripts order", gr.Textbox, { "visible": False }), })) options_templates.update(options_section(('live-preview', "Live Previews"), { - "show_progressbar": OptionInfo(True, "Show progressbar", gr.Checkbox, {"visible": False}), - "live_previews_enable": OptionInfo(True, "Show live previews", gr.Checkbox, {"visible": False}), - "show_progress_grid": OptionInfo(True, "Show previews as a grid", gr.Checkbox, {"visible": False}), "notification_audio_enable": OptionInfo(False, "Play a notification upon completion"), "notification_audio_path": OptionInfo("html/notification.mp3","Path to notification sound", component_args=hide_dirs, folder=True), "show_progress_every_n_steps": OptionInfo(1, "Live preview display period", gr.Slider, {"minimum": 0, "maximum": 32, "step": 1}), "show_progress_type": OptionInfo("Approximate", "Live preview method", gr.Radio, {"choices": ["Simple", "Approximate", "TAESD", "Full VAE"]}), - "live_preview_content": OptionInfo("Combined", "Live preview subject", gr.Radio, {"choices": ["Combined", "Prompt", "Negative prompt"], "visible": False}), "live_preview_refresh_period": OptionInfo(500, "Progress update period", gr.Slider, {"minimum": 0, "maximum": 5000, "step": 25}), "live_preview_taesd_layers": OptionInfo(3, "TAESD decode layers", gr.Slider, {"minimum": 1, "maximum": 3, "step": 1}), "logmonitor_show": OptionInfo(True, "Show log view"), @@ -781,7 +801,7 @@ options_templates.update(options_section(('sampler-params', "Sampler Settings"), 'schedulers_beta_start': OptionInfo(0, "Beta start", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.00001, "visible": native}), 'schedulers_beta_end': OptionInfo(0, "Beta end", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.00001, "visible": native}), 'schedulers_timesteps_range': OptionInfo(1000, "Timesteps range", gr.Slider, {"minimum": 250, "maximum": 4000, "step": 1, "visible": native}), - 'schedulers_shift': OptionInfo(1, "Sampler shift", gr.Slider, {"minimum": 0.1, "maximum": 10, "step": 0.1, "visible": native}), + 'schedulers_shift': OptionInfo(0, "Sampler shift", gr.Slider, {"minimum": 0.1, "maximum": 10, "step": 0.1, "visible": native}), 'schedulers_dynamic_shift': OptionInfo(True, "Sampler dynamic shift", gr.Checkbox, {"visible": native}), # managed from ui.py for backend original k-diffusion @@ -797,8 +817,6 @@ options_templates.update(options_section(('sampler-params', "Sampler Settings"), 'uni_pc_variant': OptionInfo("bh2", "UniPC variant", gr.Radio, {"choices": ["bh1", "bh2", "vary_coeff"], "visible": not native}), 'uni_pc_skip_type': OptionInfo("time_uniform", "UniPC skip type", gr.Radio, {"choices": ["time_uniform", "time_quadratic", "logSNR"], "visible": not native}), "ddim_discretize": OptionInfo('uniform', "DDIM discretize img2img", gr.Radio, {"choices": ['uniform', 'quad'], "visible": not native}), - "pad_cond_uncond": OptionInfo(True, "Pad prompt and negative prompt to be same length", gr.Checkbox, {"visible": False}), - "batch_cond_uncond": OptionInfo(True, "Do conditional and unconditional denoising in one batch", gr.Checkbox, {"visible": False}), })) options_templates.update(options_section(('postprocessing', "Postprocessing"), { @@ -808,12 +826,10 @@ options_templates.update(options_section(('postprocessing', "Postprocessing"), { "postprocessing_sep_img2img": OptionInfo("

Img2Img & Inpainting

", "", gr.HTML), "img2img_color_correction": OptionInfo(False, "Apply color correction"), "mask_apply_overlay": OptionInfo(True, "Apply mask as overlay"), - "img2img_fix_steps": OptionInfo(False, "For image processing do exact number of steps as specified", gr.Checkbox, { "visible": False }), "img2img_background_color": OptionInfo("#ffffff", "Image transparent color fill", gr.ColorPicker, {}), "inpainting_mask_weight": OptionInfo(1.0, "Inpainting conditioning mask strength", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.01}), - "initial_noise_multiplier": OptionInfo(1.0, "Noise multiplier for image processing", gr.Slider, {"minimum": 0.1, "maximum": 1.5, "step": 0.01}), - "img2img_extra_noise": OptionInfo(0.0, "Extra noise multiplier for img2img", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.01}), - "CLIP_stop_at_last_layers": OptionInfo(1, "Clip skip", gr.Slider, {"minimum": 1, "maximum": 8, "step": 1, "visible": False}), + "initial_noise_multiplier": OptionInfo(1.0, "Noise multiplier for image processing", gr.Slider, {"minimum": 0.1, "maximum": 1.5, "step": 0.01, "visible": not native}), + "img2img_extra_noise": OptionInfo(0.0, "Extra noise multiplier for img2img", gr.Slider, {"minimum": 0.0, "maximum": 1.0, "step": 0.01, "visible": not native}), # "postprocessing_sep_detailer": OptionInfo("

Detailer

", "", gr.HTML), "detailer_model": OptionInfo("Detailer", "Detailer model", gr.Radio, lambda: {"choices": [x.name() for x in detailers], "visible": False}), @@ -835,7 +851,6 @@ options_templates.update(options_section(('postprocessing', "Postprocessing"), { "postprocessing_sep_upscalers": OptionInfo("

Upscaling

", "", gr.HTML), "upscaler_unload": OptionInfo(False, "Unload upscaler after processing"), - "upscaler_for_img2img": OptionInfo("None", "Default upscaler for image resize operations", gr.Dropdown, lambda: {"choices": [x.name for x in sd_upscalers], "visible": False}, refresh=refresh_upscalers), "upscaler_tile_size": OptionInfo(192, "Upscaler tile size", gr.Slider, {"minimum": 0, "maximum": 512, "step": 16}), "upscaler_tile_overlap": OptionInfo(8, "Upscaler tile overlap", gr.Slider, {"minimum": 0, "maximum": 64, "step": 1}), })) @@ -846,28 +861,12 @@ options_templates.update(options_section(('control', "Control Options"), { "control_unload_processor": OptionInfo(False, "Processor unload after use"), })) -options_templates.update(options_section(('interrogate', "Interrogate"), { # "Training" section disabled so just a placeholder - "unload_models_when_training": OptionInfo(False, "Move VAE and CLIP to RAM when training", gr.Checkbox, { "visible": False }), - "pin_memory": OptionInfo(True, "Pin training dataset to memory", gr.Checkbox, { "visible": False }), - "save_optimizer_state": OptionInfo(False, "Save resumable optimizer state when training", gr.Checkbox, { "visible": False }), - "save_training_settings_to_txt": OptionInfo(True, "Save training settings to a text file", gr.Checkbox, { "visible": False }), - "dataset_filename_word_regex": OptionInfo("", "Filename word regex", gr.Textbox, { "visible": False }), - "dataset_filename_join_string": OptionInfo(" ", "Filename join string", gr.Textbox, { "visible": False }), - "embeddings_templates_dir": OptionInfo("", "Embeddings train templates directory", gr.Textbox, { "visible": False }), - "training_image_repeats_per_epoch": OptionInfo(1, "Image repeats per epoch", gr.Slider, {"minimum": 1, "maximum": 100, "step": 1, "visible": False }), - "training_write_csv_every": OptionInfo(0, "Save loss CSV file every n steps", gr.Number, { "visible": False }), - "training_enable_tensorboard": OptionInfo(False, "Enable tensorboard logging", gr.Checkbox, { "visible": False }), - "training_tensorboard_save_images": OptionInfo(False, "Save generated images within tensorboard", gr.Checkbox, { "visible": False }), - "training_tensorboard_flush_every": OptionInfo(120, "Tensorboard flush period", gr.Number, { "visible": False }), -})) - options_templates.update(options_section(('interrogate', "Interrogate"), { "interrogate_keep_models_in_memory": OptionInfo(False, "Interrogate: keep models in VRAM"), "interrogate_return_ranks": OptionInfo(True, "Interrogate: include ranks of model tags matches in results"), "interrogate_clip_num_beams": OptionInfo(1, "Interrogate: num_beams for BLIP", gr.Slider, {"minimum": 1, "maximum": 16, "step": 1}), "interrogate_clip_min_length": OptionInfo(32, "Interrogate: minimum description length", gr.Slider, {"minimum": 1, "maximum": 128, "step": 1}), "interrogate_clip_max_length": OptionInfo(192, "Interrogate: maximum description length", gr.Slider, {"minimum": 1, "maximum": 256, "step": 1}), - "interrogate_clip_dict_limit": OptionInfo(2048, "CLIP: maximum number of lines in text file", gr.Slider, { "visible": False }), "interrogate_clip_skip_categories": OptionInfo(["artists", "movements", "flavors"], "Interrogate: skip categories", gr.CheckboxGroup, lambda: {"choices": modules.interrogate.category_types()}, refresh=modules.interrogate.category_types), "interrogate_deepbooru_score_threshold": OptionInfo(0.65, "Interrogate: deepbooru score threshold", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01}), "deepbooru_sort_alpha": OptionInfo(False, "Interrogate: deepbooru sort alphabetically"), @@ -888,7 +887,6 @@ options_templates.update(options_section(('extra_networks', "Networks"), { "extra_networks_card_size": OptionInfo(160, "UI card size (px)", gr.Slider, {"minimum": 20, "maximum": 2000, "step": 1}), "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_network_skip_indexing": OptionInfo(False, "Build info on first access", gr.Checkbox), "extra_networks_model_sep": OptionInfo("

Models

", "", gr.HTML), @@ -909,17 +907,59 @@ options_templates.update(options_section(('extra_networks', "Networks"), { "lora_apply_tags": OptionInfo(0, "LoRA auto-apply tags", gr.Slider, {"minimum": -1, "maximum": 32, "step": 1}), "lora_in_memory_limit": OptionInfo(0, "LoRA memory cache", gr.Slider, {"minimum": 0, "maximum": 24, "step": 1}), "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", gr.Checkbox, {"visible": False}), - "sd_hypernetwork": OptionInfo("None", "Add hypernetwork to prompt", gr.Dropdown, { "choices": ["None"], "visible": False }), +options_templates.update(options_section((None, "Internal options"), { + "diffusers_version": OptionInfo("", "Diffusers version", gr.Textbox, {"visible": False}), + "disabled_extensions": OptionInfo([], "Disable these extensions"), + "sd_checkpoint_hash": OptionInfo("", "SHA256 hash of the current checkpoint"), + "tooltips": OptionInfo("UI Tooltips", "UI tooltips", gr.Radio, {"choices": ["None", "Browser default", "UI tooltips"], "visible": False}), })) options_templates.update(options_section((None, "Hidden options"), { - "disabled_extensions": OptionInfo([], "Disable these extensions"), + "batch_cond_uncond": OptionInfo(True, "Do conditional and unconditional denoising in one batch", gr.Checkbox, {"visible": False}), + "CLIP_stop_at_last_layers": OptionInfo(1, "Clip skip", gr.Slider, {"minimum": 1, "maximum": 8, "step": 1, "visible": False}), + "dataset_filename_join_string": OptionInfo(" ", "Filename join string", gr.Textbox, { "visible": False }), + "dataset_filename_word_regex": OptionInfo("", "Filename word regex", gr.Textbox, { "visible": False }), + "diffusers_force_zeros": OptionInfo(False, "Force zeros for prompts when empty", gr.Checkbox, {"visible": False}), "disable_all_extensions": OptionInfo("none", "Disable all extensions (preserves the list of disabled extensions)", gr.Radio, {"choices": ["none", "user", "all"]}), - "sd_checkpoint_hash": OptionInfo("", "SHA256 hash of the current checkpoint"), + "disable_nan_check": OptionInfo(True, "Disable NaN check", gr.Checkbox, {"visible": False}), + "embeddings_templates_dir": OptionInfo("", "Embeddings train templates directory", gr.Textbox, { "visible": False }), + "extra_networks_card_fit": OptionInfo("cover", "UI image contain method", gr.Radio, {"choices": ["contain", "cover", "fill"], "visible": False}), + "grid_extended_filename": OptionInfo(True, "Add extended info to filename when saving grid", gr.Checkbox, {"visible": False}), + "grid_save_to_dirs": OptionInfo(False, "Save grids to a subdirectory", gr.Checkbox, {"visible": False}), + "hypernetwork_enabled": OptionInfo(False, "Enable Hypernetwork support", gr.Checkbox, {"visible": False}), + "img2img_fix_steps": OptionInfo(False, "For image processing do exact number of steps as specified", gr.Checkbox, { "visible": False }), + "interrogate_clip_dict_limit": OptionInfo(2048, "CLIP: maximum number of lines in text file", gr.Slider, { "visible": False }), + "keyedit_delimiters": OptionInfo(r".,\/!?%^*;:{}=`~()", "Ctrl+up/down word delimiters", gr.Textbox, { "visible": False }), + "keyedit_precision_attention": OptionInfo(0.1, "Ctrl+up/down precision when editing (attention:1.1)", gr.Slider, {"minimum": 0.01, "maximum": 0.2, "step": 0.001, "visible": False}), + "keyedit_precision_extra": OptionInfo(0.05, "Ctrl+up/down precision when editing ", gr.Slider, {"minimum": 0.01, "maximum": 0.2, "step": 0.001, "visible": False}), + "live_preview_content": OptionInfo("Combined", "Live preview subject", gr.Radio, {"choices": ["Combined", "Prompt", "Negative prompt"], "visible": False}), + "live_previews_enable": OptionInfo(True, "Show live previews", gr.Checkbox, {"visible": False}), + "lora_functional": OptionInfo(False, "Use Kohya method for handling multiple LoRA", gr.Checkbox, { "visible": False }), + "lyco_dir": OptionInfo(os.path.join(paths.models_path, 'LyCORIS'), "Folder with LyCORIS network(s)", gr.Text, {"visible": False}), + "model_reuse_dict": OptionInfo(False, "Reuse loaded model dictionary", gr.Checkbox, {"visible": False}), + "pad_cond_uncond": OptionInfo(True, "Pad prompt and negative prompt to be same length", gr.Checkbox, {"visible": False}), + "pin_memory": OptionInfo(True, "Pin training dataset to memory", gr.Checkbox, { "visible": False }), + "save_optimizer_state": OptionInfo(False, "Save resumable optimizer state when training", gr.Checkbox, { "visible": False }), + "save_training_settings_to_txt": OptionInfo(True, "Save training settings to a text file", gr.Checkbox, { "visible": False }), + "sd_disable_ckpt": OptionInfo(False, "Disallow models in ckpt format", gr.Checkbox, {"visible": False}), + "sd_hypernetwork": OptionInfo("None", "Add hypernetwork to prompt", gr.Dropdown, { "choices": ["None"], "visible": False }), + "sd_lora": OptionInfo("", "Add LoRA to prompt", gr.Textbox, {"visible": False}), + "sd_vae_checkpoint_cache": OptionInfo(0, "Cached VAEs", gr.Slider, {"minimum": 0, "maximum": 10, "step": 1, "visible": False}), + "show_progress_grid": OptionInfo(True, "Show previews as a grid", gr.Checkbox, {"visible": False}), + "show_progressbar": OptionInfo(True, "Show progressbar", gr.Checkbox, {"visible": False}), + "training_enable_tensorboard": OptionInfo(False, "Enable tensorboard logging", gr.Checkbox, { "visible": False }), + "training_image_repeats_per_epoch": OptionInfo(1, "Image repeats per epoch", gr.Slider, {"minimum": 1, "maximum": 100, "step": 1, "visible": False }), + "training_tensorboard_flush_every": OptionInfo(120, "Tensorboard flush period", gr.Number, { "visible": False }), + "training_tensorboard_save_images": OptionInfo(False, "Save generated images within tensorboard", gr.Checkbox, { "visible": False }), + "training_write_csv_every": OptionInfo(0, "Save loss CSV file every n steps", gr.Number, { "visible": False }), + "ui_scripts_reorder": OptionInfo("", "UI scripts order", gr.Textbox, { "visible": False }), + "unload_models_when_training": OptionInfo(False, "Move VAE and CLIP to RAM when training", gr.Checkbox, { "visible": False }), + "upscaler_for_img2img": OptionInfo("None", "Default upscaler for image resize operations", gr.Dropdown, lambda: {"choices": [x.name for x in sd_upscalers], "visible": False}, refresh=refresh_upscalers), + "use_save_to_dirs_for_ui": OptionInfo(False, "Save images to a subdirectory when using Save button", gr.Checkbox, {"visible": False}), + "use_upscaler_name_as_suffix": OptionInfo(True, "Use upscaler as suffix", gr.Checkbox, {"visible": False}), })) options_templates.update() @@ -993,7 +1033,7 @@ class Options: if filename is None: filename = self.filename if cmd_opts.freeze: - log.warning(f'Settings saving is disabled: {filename}') + log.warning(f'Setting: fn="{filename}" save disabled') return try: # output = json.dumps(self.data, indent=2) @@ -1001,12 +1041,12 @@ class Options: unused_settings = [] if os.environ.get('SD_CONFIG_DEBUG', None) is not None: - log.debug('Config: user settings') + log.debug('Settings: user') for k, v in self.data.items(): log.trace(f' Config: item={k} value={v} default={self.data_labels[k].default if k in self.data_labels else None}') - log.debug('Config: default settings') + log.debug('Settings: defaults') for k in self.data_labels.keys(): - log.trace(f' Config: item={k} default={self.data_labels[k].default}') + log.trace(f' Setting: item={k} default={self.data_labels[k].default}') for k, v in self.data.items(): if k in self.data_labels: @@ -1021,9 +1061,9 @@ class Options: unused_settings.append(k) writefile(diff, filename, silent=silent) if len(unused_settings) > 0: - log.debug(f"Unused settings: {unused_settings}") + log.debug(f"Settings: unused={unused_settings}") except Exception as err: - log.error(f'Save settings failed: {filename} {err}') + log.error(f'Settings: fn="{filename}" {err}') def save(self, filename=None, silent=False): threading.Thread(target=self.save_atomic, args=(filename, silent)).start() @@ -1039,7 +1079,7 @@ class Options: if filename is None: filename = self.filename if not os.path.isfile(filename): - log.debug(f'Created default config: {filename}') + log.debug(f'Settings: fn="{filename}" created') self.save(filename) return self.data = readfile(filename, lock=True) @@ -1047,13 +1087,17 @@ class Options: self.data['quicksettings_list'] = [i.strip() for i in self.data.get('quicksettings').split(',')] unknown_settings = [] for k, v in self.data.items(): - info = self.data_labels.get(k, None) + info: OptionInfo = self.data_labels.get(k, None) + if info is not None: + if not info.validate(k, v): + self.data[k] = info.default if info is not None and not self.same_type(info.default, v): - log.error(f"Error: bad setting value: {k}: {v} ({type(v).__name__}; expected {type(info.default).__name__})") + log.warning(f"Setting validation: {k}={v} ({type(v).__name__} expected={type(info.default).__name__})") + self.data[k] = info.default if info is None and k not in compatibility_opts and not k.startswith('uiux_'): unknown_settings.append(k) if len(unknown_settings) > 0: - log.debug(f"Unknown settings: {unknown_settings}") + log.warning(f"Setting validation: unknown={unknown_settings}") def onchange(self, key, func, call=True): item = self.data_labels.get(key) @@ -1111,7 +1155,7 @@ 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 -opts.data['uni_pc_order'] = opts.schedulers_solver_order # compatibility +opts.data['uni_pc_order'] = max(2, opts.schedulers_solver_order) # compatibility log.info(f'Engine: backend={backend} compute={devices.backend} device={devices.get_optimal_device_name()} attention="{opts.cross_attention_optimization}" mode={devices.inference_context.__name__}') if not native: log.warning('Backend=original is in maintainance-only mode') @@ -1228,6 +1272,7 @@ def html(filename): return "" +@lru_cache(maxsize=1) def get_version(): version = None if version is None: diff --git a/modules/shared_state.py b/modules/shared_state.py index 067fb21eb..9947dcb70 100644 --- a/modules/shared_state.py +++ b/modules/shared_state.py @@ -62,6 +62,38 @@ class State: } return obj + def status(self): + from modules import progress + from modules.api import models + res = models.ResStatus( + task=self.job, + id=progress.current_task or '', + job=max(self.job_no, 0), + jobs=max(self.frame_count, self.job_count, self.job_no), + total=self.total_jobs, + timestamp=self.job_timestamp if self.job != '' else None, + step=self.sampling_step, + steps=self.sampling_steps, + queued=len(progress.pending_tasks), + status='unknown', + uptime = round(time.time() - self.server_start) + ) + res.step = res.steps * res.job + res.step + res.steps = res.steps * res.jobs + res.progress = round(min(1, abs(res.step / res.steps) if res.steps > 0 else 0), 2) + res.elapsed = round(time.time() - self.time_start, 2) if self.time_start is not None else None + predicted = round(res.elapsed / res.progress, 2) if res.progress > 0 and res.elapsed is not None else None + res.eta = round(predicted - res.elapsed, 2) if predicted is not None else None + if self.paused: + res.status = 'paused' + elif self.interrupted: + res.status = 'interrupted' + elif self.skipped: + res.status = 'skipped' + else: + res.status = 'running' if self.job != '' else 'idle' + return res + def begin(self, title="", api=None): import modules.devices self.total_jobs += 1 diff --git a/modules/styles.py b/modules/styles.py index de9ef43c4..0dd48eb7f 100644 --- a/modules/styles.py +++ b/modules/styles.py @@ -112,6 +112,8 @@ def apply_wildcards_to_prompt(prompt, all_wildcards, seed=-1, silent=False): def get_reference_style(): + if getattr(shared.sd_model, 'sd_checkpoint_info', None) is None: + return None name = shared.sd_model.sd_checkpoint_info.name name = name.replace('\\', '/').replace('Diffusers/', '') for k, v in shared.reference_models.items(): @@ -251,27 +253,33 @@ class StyleDatabase: return found[0] if len(found) > 0 else self.no_style def get_style_prompts(self, styles): - if styles is None or not isinstance(styles, list): + if styles is None: + return [] + if not isinstance(styles, list): shared.log.error(f'Styles invalid: {styles}') return [] return [self.find_style(x).prompt for x in styles] def get_negative_style_prompts(self, styles): - if styles is None or not isinstance(styles, list): + if styles is None: + return [] + if not isinstance(styles, list): shared.log.error(f'Styles invalid: {styles}') return [] return [self.find_style(x).negative_prompt for x in styles] def apply_styles_to_prompts(self, prompts, negatives, styles, seeds): - if styles is None or not isinstance(styles, list): + if styles is None: + return prompts, negatives + if not isinstance(styles, list): shared.log.error(f'Styles invalid styles: {styles}') - return prompts + return prompts, negatives if prompts is None or not isinstance(prompts, list): shared.log.error(f'Styles invalid prompts: {prompts}') - return prompts + return prompts, negatives if seeds is None or not isinstance(prompts, list): shared.log.error(f'Styles invalid seeds: {seeds}') - return prompts + return prompts, negatives parsed_positive = [] parsed_negative = [] for i in range(len(prompts)): @@ -286,7 +294,9 @@ class StyleDatabase: return parsed_positive, parsed_negative def apply_styles_to_prompt(self, prompt, styles): - if styles is None or not isinstance(styles, list): + if styles is None: + return prompt + if not isinstance(styles, list): shared.log.error(f'Styles invalid: {styles}') return prompt prompt = apply_styles_to_prompt(prompt, [self.find_style(x).prompt for x in styles]) @@ -294,7 +304,9 @@ class StyleDatabase: return prompt def apply_negative_styles_to_prompt(self, prompt, styles): - if styles is None or not isinstance(styles, list): + if styles is None: + return prompt + if not isinstance(styles, list): shared.log.error(f'Styles invalid: {styles}') return prompt prompt = apply_styles_to_prompt(prompt, [self.find_style(x).negative_prompt for x in styles]) @@ -302,6 +314,8 @@ class StyleDatabase: return prompt def apply_styles_to_extra(self, p): + if p.styles is None: + return if p.styles is None or not isinstance(p.styles, list): shared.log.error(f'Styles invalid: {p.styles}') return diff --git a/modules/txt2img.py b/modules/txt2img.py index 38cde0aca..2f0e2f4b3 100644 --- a/modules/txt2img.py +++ b/modules/txt2img.py @@ -49,7 +49,6 @@ def txt2img(id_task, state, subseed_strength=subseed_strength, seed_resize_from_h=seed_resize_from_h, seed_resize_from_w=seed_resize_from_w, - seed_enable_extras=True, sampler_name = processing.get_sampler_name(sampler_index), hr_sampler_name = processing.get_sampler_name(hr_sampler_index), batch_size=batch_size, diff --git a/modules/ui.py b/modules/ui.py index 039ef6487..4490bbf8c 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -354,22 +354,6 @@ def create_ui(startup_timer = None): from modules.onnx_impl import ui as ui_onnx ui_onnx.create_ui() - with gr.TabItem("Change log", id="change_log", elem_id="system_tab_changelog"): - 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.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') modules.sd_models.unload_model_weights(op='refiner') @@ -390,6 +374,16 @@ def create_ui(startup_timer = None): timer.startup.record("ui-settings") + with gr.Blocks(analytics_enabled=False) as info_interface: + with gr.Tabs(elem_id="tabs_info"): + with gr.TabItem("Change log", id="change_log", elem_id="system_tab_changelog"): + from modules import ui_docs + ui_docs.create_ui_logs() + + with gr.TabItem("Wiki", id="wiki", elem_id="system_tab_wiki"): + from modules import ui_docs + ui_docs.create_ui_wiki() + interfaces = [] interfaces += [(txt2img_interface, "Text", "txt2img")] interfaces += [(img2img_interface, "Image", "img2img")] @@ -399,6 +393,7 @@ def create_ui(startup_timer = None): interfaces += [(models_interface, "Models", "models")] interfaces += script_callbacks.ui_tabs_callback() interfaces += [(settings_interface, "System", "system")] + interfaces += [(info_interface, "Info", "info")] from modules import ui_extensions extensions_interface = ui_extensions.create_ui() diff --git a/modules/ui_common.py b/modules/ui_common.py index 9ad87c17f..9c4bb5cdc 100644 --- a/modules/ui_common.py +++ b/modules/ui_common.py @@ -246,7 +246,7 @@ def create_output_panel(tabname, preview=True, prompt=None, height=None): # columns are for <576px, <768px, <992px, <1200px, <1400px, >1400px result_gallery = gr.Gallery(value=[], label='Output', show_label=False, show_download_button=True, allow_preview=True, container=False, preview=preview, - columns=5, object_fit='scale-down', height=height, + columns=4, object_fit='scale-down', height=height, elem_id=f"{tabname}_gallery", ) if prompt is not None: diff --git a/modules/ui_docs.py b/modules/ui_docs.py new file mode 100644 index 000000000..08159beb3 --- /dev/null +++ b/modules/ui_docs.py @@ -0,0 +1,66 @@ +import gradio as gr +from modules import ui_symbols, ui_components + + +def create_ui_logs(): + 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.Column(): + get_changelog_btn = gr.Button(value='Get changelog', elem_id="get_changelog") + gr.HTML('  Open GitHub Changelog') + with gr.Column(): + _changelog_search = gr.Textbox(label="Search Changelog", 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 create_ui_wiki(): + def search_github(search_term): + import requests + from urllib.parse import quote + from installer import install + + install('beautifulsoup4') + from bs4 import BeautifulSoup + + url = f'https://github.com/search?q=repo%3Avladmandic%2Fautomatic+{quote(search_term)}&type=wikis' + res = requests.get(url, timeout=10) + if res.status_code == 200: + html = res.content + soup = BeautifulSoup(html, 'html.parser') + + # remove header links + tags = soup.find_all(attrs={"data-hovercard-url": "/vladmandic/automatic/hovercard"}) + for tag in tags: + tag.extract() + + # replace relative links with full links + tags = soup.find_all('a') + for tag in tags: + if tag.has_attr('href') and tag['href'].startswith('/'): + tag['href'] = 'https://github.com' + tag['href'] + + # find result only + result = soup.find(attrs={"data-testid": "results-list"}) + if result is None: + return 'No results found' + html = str(result) + return html + else: + return f'Error: {res.status_code}' + + with gr.Row(): + gr.HTML('  Open GitHub Wiki') + with gr.Row(): + wiki_search = gr.Textbox(label="Search Wiki Pages", elem_id="wiki_search") + wiki_search_btn = ui_components.ToolButton(value=ui_symbols.search, label="Search", elem_id="wiki_search_btn") + with gr.Row(): + wiki_result = gr.HTML(elem_id="wiki_result", value='') + wiki_search.submit(_js="wikiSearch", fn=search_github, inputs=[wiki_search], outputs=[wiki_result]) + wiki_search_btn.click(_js="wikiSearch", fn=search_github, inputs=[wiki_search], outputs=[wiki_result]) diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index 323f4830f..f6e6cee97 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -148,7 +148,8 @@ class ExtraNetworksPage: if self.title == 'Model': return opt = xyz_grid.AxisOption(f"[Network] {self.title}", str, add_prompt, choices=lambda: [x["name"] for x in self.items]) - xyz_grid.axis_options.append(opt) + if opt not in xyz_grid.axis_options: + xyz_grid.axis_options.append(opt) def link_preview(self, filename): quoted_filename = urllib.parse.quote(filename.replace('\\', '/')) diff --git a/modules/ui_img2img.py b/modules/ui_img2img.py index 4cb8e4c18..22c89dac8 100644 --- a/modules/ui_img2img.py +++ b/modules/ui_img2img.py @@ -131,6 +131,7 @@ def create_ui(): full_quality, tiling, hidiffusion, cfg_scale, clip_skip, image_cfg_scale, diffusers_guidance_rescale, pag_scale, pag_adaptive, cfg_end = ui_sections.create_advanced_inputs('img2img') hdr_mode, hdr_brightness, hdr_color, hdr_sharpen, hdr_clamp, hdr_boundary, hdr_threshold, hdr_maximize, hdr_max_center, hdr_max_boundry, hdr_color_picker, hdr_tint_ratio = ui_sections.create_correction_inputs('img2img') + enable_hr, hr_sampler_index, hr_denoising_strength, hr_resize_mode, hr_resize_context, hr_upscaler, hr_force, hr_second_pass_steps, hr_scale, hr_resize_x, hr_resize_y, refiner_steps, hr_refiner_start, refiner_prompt, refiner_negative = ui_sections.create_hires_inputs('txt2img') detailer = shared.yolo.ui('img2img') # with gr.Group(elem_id="inpaint_controls", visible=False) as inpaint_controls: @@ -192,6 +193,7 @@ def create_ui(): inpaint_full_res, inpaint_full_res_padding, inpainting_mask_invert, img2img_batch_files, img2img_batch_input_dir, img2img_batch_output_dir, img2img_batch_inpaint_mask_dir, hdr_mode, hdr_brightness, hdr_color, hdr_sharpen, hdr_clamp, hdr_boundary, hdr_threshold, hdr_maximize, hdr_max_center, hdr_max_boundry, hdr_color_picker, hdr_tint_ratio, + enable_hr, hr_sampler_index, hr_denoising_strength, hr_resize_mode, hr_resize_context, hr_upscaler, hr_force, hr_second_pass_steps, hr_scale, hr_resize_x, hr_resize_y, refiner_steps, hr_refiner_start, refiner_prompt, refiner_negative, override_settings, ] img2img_dict = dict( diff --git a/modules/ui_models.py b/modules/ui_models.py index e9be428b4..624c3849d 100644 --- a/modules/ui_models.py +++ b/modules/ui_models.py @@ -453,7 +453,7 @@ def create_ui(): if tag is not None and len(tag) > 0: url += f'&tag={tag}' r = req(url) - log.debug(f'CivitAI search: name="{name}" tag={tag or "none"} url="{url}" status={r.status_code}') + log.debug(f'CivitAI search: type={model_type} name="{name}" tag={tag or "none"} url="{url}" status={r.status_code}') if r.status_code != 200: log.warning(f'CivitAI search: name="{name}" tag={tag} status={r.status_code}') return [], gr.update(visible=False, value=[]), gr.update(visible=False, value=None), gr.update(visible=False, value=None) diff --git a/modules/ui_sections.py b/modules/ui_sections.py index 391b5c609..7951a9227 100644 --- a/modules/ui_sections.py +++ b/modules/ui_sections.py @@ -276,7 +276,7 @@ def create_sampler_options(tabname): else: # shared.native with gr.Row(elem_classes=['flex-break']): - sampler_sigma = gr.Dropdown(label='Sigma method', elem_id=f"{tabname}_sampler_sigma", choices=['default', 'karras', 'beta', 'exponential'], value=shared.opts.schedulers_sigma, type='value') + sampler_sigma = gr.Dropdown(label='Sigma method', elem_id=f"{tabname}_sampler_sigma", choices=['default', 'karras', 'beta', 'exponential', 'lambdas'], value=shared.opts.schedulers_sigma, type='value') sampler_spacing = gr.Dropdown(label='Timestep spacing', elem_id=f"{tabname}_sampler_spacing", choices=['default', 'linspace', 'leading', 'trailing'], value=shared.opts.schedulers_timestep_spacing, type='value') with gr.Row(elem_classes=['flex-break']): sampler_beta = gr.Dropdown(label='Beta schedule', elem_id=f"{tabname}_sampler_beta", choices=['default', 'linear', 'scaled', 'cosine'], value=shared.opts.schedulers_beta_schedule, type='value') diff --git a/modules/unipc/sampler.py b/modules/unipc/sampler.py index bcc4eed76..b5e116d61 100644 --- a/modules/unipc/sampler.py +++ b/modules/unipc/sampler.py @@ -186,6 +186,6 @@ class UniPCSampler(object): ) uni_pc = UniPC(model_fn, self.noise_schedule, predict_x0=True, thresholding=False, variant=shared.opts.uni_pc_variant, condition=conditioning, unconditional_condition=unconditional_conditioning, before_sample=self.before_sample, after_sample=self.after_sample, after_update=self.after_update) - x = uni_pc.sample(img, steps=S, skip_type=shared.opts.uni_pc_skip_type, method="multistep", order=shared.opts.schedulers_solver_order, lower_order_final=shared.opts.schedulers_use_loworder) + x = uni_pc.sample(img, steps=S, skip_type=shared.opts.uni_pc_skip_type, method="multistep", order=shared.opts.uni_pc_order, lower_order_final=shared.opts.uni_pc_lower_order_final) return x.to(device), None diff --git a/modules/vqa.py b/modules/vqa.py index 64ba83696..ee4197a5e 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 @@ -154,8 +156,8 @@ def florence(question: str, image: Image.Image, repo: str = None, revision: str task = question.split('>', 1)[0] + '>' else: task = '' - question = task + question - inputs = processor(text=question, images=image, return_tensors="pt") + # question = task + question + inputs = processor(text=task, images=image, return_tensors="pt") input_ids = inputs['input_ids'].to(devices.device) pixel_values = inputs['pixel_values'].to(devices.device, devices.dtype) with devices.inference_context(): diff --git a/modules/zluda_hijacks.py b/modules/zluda_hijacks.py index ea906e7c0..0f42a5448 100644 --- a/modules/zluda_hijacks.py +++ b/modules/zluda_hijacks.py @@ -19,6 +19,11 @@ def fft_ifftn(input: torch.Tensor, *args, **kwargs) -> torch.Tensor: # pylint: d return _fft_ifftn(input.cpu(), *args, **kwargs).to(input.device) +_fft_rfftn = torch.fft.rfftn +def fft_rfftn(input: torch.Tensor, *args, **kwargs) -> torch.Tensor: # pylint: disable=redefined-builtin + return _fft_rfftn(input.cpu(), *args, **kwargs).to(input.device) + + def jit_script(f, *_, **__): # experiment / provide dummy graph f.graph = torch._C.Graph() # pylint: disable=protected-access return f @@ -29,4 +34,5 @@ def do_hijack(): torch.topk = topk torch.fft.fftn = fft_fftn torch.fft.ifftn = fft_ifftn + torch.fft.rfftn = fft_rfftn torch.jit.script = jit_script diff --git a/package.json b/package.json index bf5a366ea..b30d3f87d 100644 --- a/package.json +++ b/package.json @@ -16,11 +16,13 @@ "url": "git+https://github.com/vladmandic/automatic.git" }, "scripts": { + "venv": "source venv/bin/activate", "start": "python launch.py --debug --experimental", "ruff": "ruff check", "eslint": "eslint javascript/ extensions-builtin/sdnext-modernui/javascript/", "pylint": "pylint *.py modules/ extensions-builtin/", - "lint": "npm run eslint; npm run ruff; npm run pylint" + "lint": "npm run eslint; npm run ruff; npm run pylint", + "test": "cli/test.sh" }, "devDependencies": { "esbuild": "^0.18.15" diff --git a/requirements.txt b/requirements.txt index c6fcf97ec..12a9f85cb 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,4 +1,7 @@ +# required for python 3.12 setuptools==69.5.1 + +# standard patch-ng anyio addict @@ -27,6 +30,8 @@ ruff pylint invisible-watermark pi-heif + +# versioned safetensors==0.4.5 tensordict==0.1.2 peft==0.13.1 @@ -36,7 +41,7 @@ torchsde==0.2.6 antlr4-python3-runtime==4.9.3 requests==2.32.3 tqdm==4.66.5 -accelerate==1.0.1 +accelerate==1.1.1 opencv-contrib-python-headless==4.9.0.80 einops==0.4.1 gradio==3.43.2 @@ -44,19 +49,21 @@ huggingface_hub==0.26.2 numexpr==2.8.8 numpy==1.26.4 numba==0.59.1 -blendmodes -scipy -pandas protobuf==4.25.3 pytorch_lightning==1.9.4 -tokenizers==0.20.1 -transformers==4.46.1 +tokenizers==0.20.3 +transformers==4.46.2 urllib3==1.26.19 Pillow==10.4.0 timm==0.9.16 pydantic==1.10.15 pyparsing==3.1.4 typing-extensions==4.12.2 + +# additional +blendmodes +scipy +pandas torchdiffeq dctorch scikit-image diff --git a/scripts/animatediff.py b/scripts/animatediff.py index 09f9e33a9..4c50f9cf6 100644 --- a/scripts/animatediff.py +++ b/scripts/animatediff.py @@ -189,7 +189,7 @@ def set_free_noise(frames): class Script(scripts.Script): def title(self): - return 'AnimateDiff' + return 'Video AnimateDiff' def show(self, is_img2img): # return scripts.AlwaysVisible if shared.native else False @@ -258,7 +258,7 @@ class Script(scripts.Script): shared.log.debug(f'AnimateDiff args: {p.task_args}') set_prompt(p) orig_prompt_attention = shared.opts.prompt_attention - shared.opts.data['prompt_attention'] = 'Fixed attention' + shared.opts.data['prompt_attention'] = 'fixed' processed: processing.Processed = processing.process_images(p) # runs processing using main loop shared.opts.data['prompt_attention'] = orig_prompt_attention devices.torch_gc() diff --git a/scripts/apg.py b/scripts/apg.py index c7e60c982..6d0ec107e 100644 --- a/scripts/apg.py +++ b/scripts/apg.py @@ -2,6 +2,9 @@ import gradio as gr from modules import scripts, processing, shared, sd_models +registered = False + + class Script(scripts.Script): def __init__(self): super().__init__() @@ -9,14 +12,14 @@ class Script(scripts.Script): self.register() def title(self): - return 'APG' + return 'APG: Adaptive Projected Guidance' 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('  APG: Adaptive projected guidance
') + gr.HTML('  APG: Adaptive Projected Guidance
') with gr.Row(): eta = gr.Slider(label="ETA", value=1.0, minimum=0, maximum=2.0, step=0.05) momentum = gr.Slider(label="Momentum", value=-0.50, minimum=-1.0, maximum=1.0, step=0.05) @@ -24,6 +27,10 @@ class Script(scripts.Script): return [eta, momentum, threshold] def register(self): # register xyz grid elements + global registered # pylint: disable=global-statement + if registered: + return + registered = True def apply_field(field): def fun(p, x, xs): # pylint: disable=unused-argument setattr(p, field, x) @@ -32,9 +39,14 @@ class Script(scripts.Script): 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("[APG] ETA", float, apply_field("apg_eta"))) - xyz_classes.axis_options.append(xyz_classes.AxisOption("[APG] Momentum", float, apply_field("apg_momentum"))) - xyz_classes.axis_options.append(xyz_classes.AxisOption("[APG] Threshold", float, apply_field("apg_threshold"))) + options = [ + xyz_classes.AxisOption("[APG] ETA", float, apply_field("apg_eta")), + xyz_classes.AxisOption("[APG] Momentum", float, apply_field("apg_momentum")), + xyz_classes.AxisOption("[APG] Threshold", float, apply_field("apg_threshold")), + ] + for option in options: + if option not in xyz_classes.axis_options: + xyz_classes.axis_options.append(option) def run(self, p: processing.StableDiffusionProcessing, eta = 0.0, momentum = 0.0, threshold = 0.0): # pylint: disable=arguments-differ supported_model_list = ['sd', 'sdxl', 'sc'] diff --git a/scripts/blipdiffusion.py b/scripts/blipdiffusion.py index 39d8974e1..0acc80929 100644 --- a/scripts/blipdiffusion.py +++ b/scripts/blipdiffusion.py @@ -2,19 +2,16 @@ import gradio as gr from modules import scripts, processing, shared, sd_models -title = 'BLIP Diffusion' - - class Script(scripts.Script): def title(self): - return title + return 'BLIP Diffusion: Controllable Generation and Editing' def show(self, is_img2img): return is_img2img if shared.native else False def ui(self, _is_img2img): with gr.Row(): - gr.HTML('  BLIP Diffusion
') + gr.HTML('  BLIP Diffusion: Controllable Generation and Editing
') with gr.Row(): source_subject = gr.Textbox(value='', label='Source subject') with gr.Row(): @@ -26,7 +23,7 @@ class Script(scripts.Script): def run(self, p: processing.StableDiffusionProcessing, source_subject, target_subject, prompt_strength): # pylint: disable=arguments-differ, unused-argument c = shared.sd_model.__class__.__name__ if shared.sd_loaded else '' if c != 'BlipDiffusionPipeline': - shared.log.error(f'{title}: model selected={c} required=BLIPDiffusion') + shared.log.error(f'BLIP: model selected={c} required=BLIPDiffusion') return None if hasattr(p, 'init_images') and len(p.init_images) > 0: p.task_args['reference_image'] = p.init_images[0] @@ -41,5 +38,5 @@ class Script(scripts.Script): processed = processing.process_images(p) return processed else: - shared.log.error(f'{title}: no init_images') + shared.log.error('BLIP: no init_images') return None diff --git a/scripts/cogvideo.py b/scripts/cogvideo.py index a4a3141d4..7f2c7225e 100644 --- a/scripts/cogvideo.py +++ b/scripts/cogvideo.py @@ -22,7 +22,7 @@ debug = (os.environ.get('SD_LOAD_DEBUG', None) is not None) or (os.environ.get(' class Script(scripts.Script): def title(self): - return 'CogVideoX' + return 'Video CogVideoX' def show(self, is_img2img): return shared.native diff --git a/scripts/consistory_ext.py b/scripts/consistory_ext.py new file mode 100644 index 000000000..45de1ea6c --- /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: Consistent Image Generation' + + 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/ctrlx.py b/scripts/ctrlx.py index acfdd5d6e..f372e2a94 100644 --- a/scripts/ctrlx.py +++ b/scripts/ctrlx.py @@ -7,14 +7,14 @@ from modules import shared, scripts, processing, processing_helpers, sd_models, class Script(scripts.Script): def title(self): - return 'Ctrl-X' + return 'Ctrl-X: Controlling Structure and Appearance' def show(self, is_img2img): return shared.native def ui(self, _is_img2img): with gr.Row(): - gr.HTML('  Ctrl-X
') + gr.HTML('  Ctrl-X: Controlling Structure and Appearance
') with gr.Accordion(label='Structure', open=True): with gr.Row(): struct_prompt = gr.Textbox(label='Prompt', value='', rows=1) @@ -49,7 +49,7 @@ class Script(scripts.Script): from modules.ctrlx.utils import get_self_recurrence_schedule orig_prompt_attention = shared.opts.prompt_attention - shared.opts.data['prompt_attention'] = 'Fixed attention' + shared.opts.data['prompt_attention'] = 'fixed' shared.sd_model = sd_models.switch_pipe(CtrlXStableDiffusionXLPipeline, shared.sd_model) shared.sd_model.restore_pipeline = self.restore diff --git a/scripts/demofusion.py b/scripts/demofusion.py index f7cdfe543..6625c0c79 100644 --- a/scripts/demofusion.py +++ b/scripts/demofusion.py @@ -1221,7 +1221,7 @@ class DemoFusionSDXLPipeline(DiffusionPipeline, FromSingleFileMixin, LoraLoaderM class Script(scripts.Script): def title(self): - return 'DemoFusion' + return 'DemoFusion: High-Resolution Image Generation' def show(self, is_img2img): return not is_img2img if shared.native else False @@ -1229,7 +1229,7 @@ class Script(scripts.Script): # return signature is array of gradio components def ui(self, _is_img2img): with gr.Row(): - gr.HTML('  DemoFusion
') + gr.HTML('  DemoFusion: High-Resolution Image Generation
') with gr.Row(): cosine_scale_1 = gr.Slider(minimum=0, maximum=5, step=0.1, value=3, label="Cosine scale 1") cosine_scale_2 = gr.Slider(minimum=0, maximum=5, step=0.1, value=1, label="Cosine scale 2") diff --git a/scripts/differential_diffusion.py b/scripts/differential_diffusion.py index 705242987..da4ae0e2e 100644 --- a/scripts/differential_diffusion.py +++ b/scripts/differential_diffusion.py @@ -1858,14 +1858,14 @@ MODELS = { class Script(scripts.Script): def title(self): - return 'Differential diffusion' + return 'Differential diffusion: Individual Pixel Strength' def show(self, is_img2img): return is_img2img if shared.native else False def ui(self, _is_img2img): with gr.Row(): - gr.HTML('  Differential diffusion
Select a model for auto-preprocess or upload an image map
') + gr.HTML('  Differential diffusion: Individual Pixel Strength
Select a model for auto-preprocess or upload an image map
') with gr.Row(): enabled = gr.Checkbox(label='Enabled', value=True) invert = gr.Checkbox(label='Mask invert', value=False) diff --git a/scripts/hdr.py b/scripts/hdr.py index 788c0add2..9afc3673b 100644 --- a/scripts/hdr.py +++ b/scripts/hdr.py @@ -11,14 +11,14 @@ from modules.shared import opts, state class Script(scripts.Script): def title(self): - return "HDR" + return "HDR: High Dynamic Range" def show(self, is_img2img): return True def ui(self, is_img2img): with gr.Row(): - gr.HTML("  High Dynamic Range
") + gr.HTML("  HDR: High Dynamic Range
") with gr.Row(): save_hdr = gr.Checkbox(label="Save HDR image", value=True) hdr_range = gr.Slider(minimum=0, maximum=1, step=0.05, value=0.65, label='HDR range') diff --git a/scripts/image2video.py b/scripts/image2video.py index 332972a6d..876ed3193 100644 --- a/scripts/image2video.py +++ b/scripts/image2video.py @@ -13,7 +13,7 @@ MODELS = [ class Script(scripts.Script): def title(self): - return 'Image-to-Video' + return 'Video VGen Image-to-Video' def show(self, is_img2img): return is_img2img if shared.native else False @@ -102,6 +102,7 @@ class Script(scripts.Script): processed = processing.process_images(p) shared.sd_model.motion_adapter = None + processed = None if model_name == 'VGen': if not isinstance(shared.sd_model, diffusers.I2VGenXLPipeline): shared.log.info(f'Image2Video VGen load: model={repo_id}') diff --git a/scripts/instantir.py b/scripts/instantir.py new file mode 100644 index 000000000..5eb7d503a --- /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: Image Restoration' + + 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/k_diff.py b/scripts/k_diff.py index 757f2ea2d..92b43149d 100644 --- a/scripts/k_diff.py +++ b/scripts/k_diff.py @@ -7,17 +7,16 @@ from modules import scripts, processing, shared, sd_models class Script(scripts.Script): supported_models = ['sd', 'sdxl'] orig_pipe = None - library = None def title(self): - return 'K-Diffusion' + return 'K-Diffusion Samplers' 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
') + gr.HTML('  K-Diffusion Samplers
') with gr.Row(): sampler = gr.Dropdown(label="Sampler", choices=self.samplers()) return [sampler] @@ -34,6 +33,8 @@ class Script(scripts.Script): _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 @@ -44,12 +45,16 @@ class Script(scripts.Script): 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) - sampler = 'sample_' + sampler - - sampling = getattr(self.library, "sampling", None) - shared.sd_model.sampler = getattr(sampling, sampler) + 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} diff --git a/scripts/layerdiffuse.py b/scripts/layerdiffuse.py index a1e15aa8b..ecf7da1d3 100644 --- a/scripts/layerdiffuse.py +++ b/scripts/layerdiffuse.py @@ -5,7 +5,7 @@ from modules import shared, scripts, sd_models class Script(scripts.Script): def title(self): - return 'LayerDiffuse' + return 'LayerDiffuse: Transparent Image' def show(self, is_img2img): return True if shared.native else False @@ -40,7 +40,7 @@ class Script(scripts.Script): def ui(self, _is_img2img): with gr.Row(): gr.HTML(""" -   LayerDiffuse

+   LayerDiffuse: Transparent Image

- Click Apply to model to apply LayerDiffuse to current model
- Click Reload model to remove LayerDiffuse from current model

""") diff --git a/scripts/ledits.py b/scripts/ledits.py index 1a0e929f0..b75c6ff6f 100644 --- a/scripts/ledits.py +++ b/scripts/ledits.py @@ -5,7 +5,7 @@ from modules import scripts, processing, shared, devices, sd_models class Script(scripts.Script): def title(self): - return 'LEdits++' + return 'LEdits: Limitless Image Editing' def show(self, is_img2img): return is_img2img if shared.native else False @@ -13,7 +13,7 @@ class Script(scripts.Script): # return signature is array of gradio components def ui(self, _is_img2img): with gr.Row(): - gr.HTML('  LEdits++
') + gr.HTML('  LEdits++: Limitless Image Editing
') with gr.Row(): edit_start = gr.Slider(label='Edit start', minimum=0.0, maximum=1.0, step=0.01, value=0.1) edit_stop = gr.Slider(label='Edit stop', minimum=0.0, maximum=1.0, step=0.01, value=1.0) @@ -44,7 +44,7 @@ class Script(scripts.Script): orig_offload = shared.opts.diffusers_model_cpu_offload orig_prompt_attention = shared.opts.prompt_attention shared.opts.data['diffusers_model_cpu_offload'] = False - shared.opts.data['prompt_attention'] = 'Fixed attention' + shared.opts.data['prompt_attention'] = 'fixed' # shared.sd_model.maybe_free_model_hooks() # ledits is not compatible with offloading # shared.sd_model.has_accelerate = False sd_models.move_model(shared.sd_model, devices.device, force=True) diff --git a/scripts/lut.py b/scripts/lut.py index 3d240f291..573222161 100644 --- a/scripts/lut.py +++ b/scripts/lut.py @@ -17,7 +17,7 @@ class Script(scripts.Script): def ui(self, _is_img2img): with gr.Row(): - gr.HTML("  Color grading
") + gr.HTML("  LUT Color grading
") with gr.Row(): original = gr.Checkbox(label='Include original image', value=True) with gr.Row(): diff --git a/scripts/mixture_tiling.py b/scripts/mixture_tiling.py index 5dcaf0156..5b5aab9db 100644 --- a/scripts/mixture_tiling.py +++ b/scripts/mixture_tiling.py @@ -26,14 +26,14 @@ def check_dependencies(): class Script(scripts.Script): def title(self): - return 'Mixture tiling' + return 'Mixture Tiling: Scene Composition' def show(self, is_img2img): return not is_img2img if shared.native else False def ui(self, _is_img2img): with gr.Row(): - gr.HTML('  Mixture tiling
') + gr.HTML('  Mixture Tiling: Scene Composition
') with gr.Row(): gr.HTML('  Separated prompts using new lines
  Number of prompts must matcxh X*Y
') with gr.Row(): @@ -66,7 +66,7 @@ class Script(scripts.Script): shared.sd_model = orig_pipeline return sd_models.set_diffuser_options(shared.sd_model) - shared.opts.data['prompt_attention'] = 'Fixed attention' # this pipeline is not compatible with embeds + shared.opts.data['prompt_attention'] = 'fixed' # this pipeline is not compatible with embeds shared.sd_model.to(torch.float32) # this pipeline unet is not compatible with fp16 processing.fix_seed(p) # set pipeline specific params, note that standard params are applied when applicable diff --git a/scripts/mulan.py b/scripts/mulan.py index 4b80a7c87..829ce1463 100644 --- a/scripts/mulan.py +++ b/scripts/mulan.py @@ -46,14 +46,14 @@ text_encoder_path = None class Script(scripts.Script): def title(self): - return 'MuLan' + return 'MuLan: Multi Language Prompts' def show(self, is_img2img): return True if shared.native else False def ui(self, _is_img2img): with gr.Row(): - gr.HTML('  MuLan
') + gr.HTML('  MuLan: Multi Language Prompts
') with gr.Row(): selected_encoder = gr.Dropdown(label='Encoder', choices=ENCODERS, value=ENCODERS[0]) return [selected_encoder] @@ -87,7 +87,7 @@ class Script(scripts.Script): # mulan only works with single image, single prompt and in fixed attention p.batch_size = 1 p.n_iter = 1 - shared.opts.prompt_attention = 'Fixed attention' + shared.opts.prompt_attention = 'fixed' if isinstance(p.prompt, list): p.prompt = p.prompt[0] p.task_args['prompt'] = p.prompt diff --git a/scripts/pulid_ext.py b/scripts/pulid_ext.py new file mode 100644 index 000000000..40a21e3f7 --- /dev/null +++ b/scripts/pulid_ext.py @@ -0,0 +1,254 @@ +import io +import os +import time +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 +direct = False +registered = False +uploaded_images = [] + + +class Script(scripts.Script): + def __init__(self): + self.pulid = None + self.cache = None + self.preprocess = 0 + super().__init__() + self.register() # pulid is script with processing override so xyz doesnt execute + + def title(self): + return 'PuLID: ID Customization' + + 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) + + def register(self): # register xyz grid elements + global registered # pylint: disable=global-statement + if registered: + return + registered = True + 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] + options = [ + xyz_classes.AxisOption("[PuLID] Strength", float, apply_field("pulid_strength")), + xyz_classes.AxisOption("[PuLID] Zero", int, apply_field("pulid_zero")), + xyz_classes.AxisOption("[PuLID] Ortho", str, apply_field("pulid_ortho"), choices=lambda: ['off', 'v1', 'v2']), + ] + for option in options: + if option not in xyz_classes.axis_options: + xyz_classes.axis_options.append(option) + + + def load_images(self, files): + uploaded_images.clear() + 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}') + uploaded_images.append(image) + except Exception as e: + shared.log.warning(f'IP adapter failed to load image: {e}') + return gr.update(value=uploaded_images, visible=len(uploaded_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", value='dpmpp_sde', choices=['dpmpp_2m', 'dpmpp_2m_sde', 'dpmpp_2s_ancestral', 'dpmpp_3m_sde', 'dpmpp_sde', 'euler', 'euler_ancestral']) + ortho = gr.Dropdown(label="Ortho", choices=['off', 'v1', 'v2'], value='v2') + with gr.Row(): + version = gr.Dropdown(label="Version", value='v1.1', choices=['v1.0', 'v1.1']) + with gr.Row(): + restore = gr.Checkbox(label='Restore pipe on end', value=False) + offload = gr.Checkbox(label='Offload face module', value=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, restore, offload, version] + + def run( + self, + p: processing.StableDiffusionProcessing, + strength: float = 0.8, + zero: int = 20, + sampler: str = 'dpmpp_sde', + ortho: str = 'v2', + gallery: list = [], + restore: bool = False, + offload: bool = True, + version: str = 'v1.1' + ): # pylint: disable=arguments-differ, unused-argument + images = [] + try: + if len(gallery) == 0: + from modules.api.api import decode_base64_to_image + images = getattr(p, 'pulid_images', uploaded_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["pulid"] = pulid.StableDiffusionXLPuLIDPipeline + pipelines.auto_pipeline.AUTO_IMAGE2IMAGE_PIPELINES_MAPPING["pulid"] = pulid.StableDiffusionXLPuLIDPipelineImage + pipelines.auto_pipeline.AUTO_INPAINT_PIPELINES_MAPPING["pulid"] = pulid.StableDiffusionXLPuLIDPipelineInpaint + 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 + + sdp = shared.opts.cross_attention_optimization == "Scaled-Dot-Product" + strength = getattr(p, 'pulid_strength', strength) + zero = getattr(p, 'pulid_zero', zero) + ortho = getattr(p, 'pulid_ortho', ortho) + sampler = getattr(p, 'pulid_sampler', sampler) + sampler_fn = getattr(self.pulid.sampling, f'sample_{sampler}', None) + if sampler_fn is None: + sampler_fn = self.pulid.sampling.sample_dpmpp_2m_sde + + 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, + dtype=devices.dtype, + providers=devices.onnx, + offload=offload, + version=version, + sdp=sdp, + 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') + # shared.sd_model.hack_unet_attn_layers(shared.sd_model.pipe.unet) # reapply attention layers + devices.torch_gc() + except Exception as e: + shared.log.error(f'PuLID: failed to create pipeline: {e}') + errors.display(e, 'PuLID') + return None + + shared.sd_model.sampler = sampler_fn + shared.log.info(f'PuLID: class={shared.sd_model.__class__.__name__} version="{version}" sdp={sdp} strength={strength} zero={zero} ortho={ortho} sampler={sampler_fn} images={[i.shape for i in images]} offload={offload}') + 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 = [] + + # get id embedding used for attention + t0 = time.time() + uncond_id_embedding, id_embedding = shared.sd_model.get_id_embedding(images) + if offload: + devices.torch_gc() + t1 = time.time() + self.preprocess = t1-t0 + + p.seed = processing_helpers.get_fixed_seed(p.seed) + if direct: # run pipeline directly + shared.state.begin('PuLID') + processing.fix_seed(p) + 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 + p.extra_generation_params["PuLID"] = f'Strength={strength} Zero={zero} Ortho={ortho}' + p.extra_generation_params["Sampler"] = sampler + 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 + _strength, _zero, _sampler, _ortho, _gallery, restore, _offload, _version = args + if hasattr(shared.sd_model, 'pipe') and shared.sd_model_type == "sdxl": + restore = getattr(p, 'pulid_restore', restore) + if restore: + 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 + shared.sd_model = shared.sd_model.pipe + devices.torch_gc(force=True) + shared.log.debug(f'PuLID complete: class={shared.sd_model.__class__.__name__} preprocess={self.preprocess:.2f} pipe={"restore" if restore else "cache"}') + return processed diff --git a/scripts/regional_prompting.py b/scripts/regional_prompting.py index cecef747d..08b84dd94 100644 --- a/scripts/regional_prompting.py +++ b/scripts/regional_prompting.py @@ -64,7 +64,7 @@ class Script(scripts.Script): shared.sd_model = orig_pipeline return sd_models.set_diffuser_options(shared.sd_model) - shared.opts.data['prompt_attention'] = 'Fixed attention' # this pipeline is not compatible with embeds + shared.opts.data['prompt_attention'] = 'fixed' # this pipeline is not compatible with embeds processing.fix_seed(p) # set pipeline specific params, note that standard params are applied when applicable rp_args = { diff --git a/scripts/resadapter.py b/scripts/resadapter.py index cbd0bf671..58162f9ab 100644 --- a/scripts/resadapter.py +++ b/scripts/resadapter.py @@ -19,7 +19,7 @@ models = { class Script(scripts.Script): def title(self): - return 'ResAdapter' + return 'ResAdapter: Domain Consistent Resolution' def show(self, is_img2img): return not is_img2img if shared.native else False @@ -27,7 +27,7 @@ class Script(scripts.Script): # return signature is array of gradio components def ui(self, _is_img2img): with gr.Row(): - gr.HTML('  ResAdapter
') + gr.HTML('  ResAdapter: Domain Consistent Resolution
') with gr.Row(): model = gr.Dropdown(label="Model", choices=list(models), value="None") weight = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label="Weight", value=1.0) diff --git a/scripts/sd_upscale.py b/scripts/sd_upscale.py index 9f21c5645..9c5a72204 100644 --- a/scripts/sd_upscale.py +++ b/scripts/sd_upscale.py @@ -9,7 +9,7 @@ from modules.shared import opts, state, log class Script(scripts.Script): def title(self): - return "SD upscale" + return "SD Upscale" def show(self, is_img2img): return is_img2img diff --git a/scripts/stablevideodiffusion.py b/scripts/stablevideodiffusion.py index 585871edc..cbf2ce003 100644 --- a/scripts/stablevideodiffusion.py +++ b/scripts/stablevideodiffusion.py @@ -16,7 +16,7 @@ models = { class Script(scripts.Script): def title(self): - return 'Stable Video Diffusion' + return 'Video: SVD' def show(self, is_img2img): return is_img2img if shared.native else False diff --git a/scripts/t_gate.py b/scripts/t_gate.py index 3bd51445d..3808a796d 100644 --- a/scripts/t_gate.py +++ b/scripts/t_gate.py @@ -5,7 +5,7 @@ from installer import install class Script(scripts.Script): def title(self): - return 'T-Gate' + return 'T-Gate: Accelerate via Gating Attention' def show(self, is_img2img): return not is_img2img if shared.native else False @@ -13,7 +13,7 @@ class Script(scripts.Script): # return signature is array of gradio components def ui(self, _is_img2img): with gr.Row(): - gr.HTML('  T-Gate
') + gr.HTML('  T-Gate: Accelerate via Gating Attention
') with gr.Row(): enabled = gr.Checkbox(label="Enabled", value=True) with gr.Row(): diff --git a/scripts/text2video.py b/scripts/text2video.py index 2c93abf27..dc4c44cac 100644 --- a/scripts/text2video.py +++ b/scripts/text2video.py @@ -23,7 +23,7 @@ MODELS = [ class Script(scripts.Script): def title(self): - return 'Text-to-Video' + return 'Video: ModelScope' def show(self, is_img2img): return not is_img2img if shared.native else False @@ -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 553a20d30..c67eca18b 100644 --- a/scripts/x_adapter.py +++ b/scripts/x_adapter.py @@ -107,7 +107,7 @@ class Script(scripts.Script): pipe.to(device=devices.device, dtype=devices.dtype) except Exception: pass - shared.opts.data['prompt_attention'] = 'Fixed attention' + shared.opts.data['prompt_attention'] = 'fixed' prompt = shared.prompt_styles.apply_styles_to_prompt(p.prompt, p.styles) negative = shared.prompt_styles.apply_negative_styles_to_prompt(p.negative_prompt, p.styles) p.task_args['prompt'] = prompt diff --git a/scripts/xyz_grid.py b/scripts/xyz_grid.py index e8ecb5bd4..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,18 +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(): 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") @@ -131,9 +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, include_grid, include_subgrids, include_images, include_time, 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, include_grid, include_subgrids, include_images, include_time, margin_size): # pylint: disable=W0221 + 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[:] @@ -314,31 +360,40 @@ class Script(scripts.Script): margin_size=margin_size, no_grid=not include_grid, include_time=include_time, + include_text=include_text, ) if not processed.images: return processed # something 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_images: # dont need sub-images anymore, drop from list: - if not include_grid and include_subgrids: - 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 include_grid and z_count > 1 else 0 ) - for g in range(grid_count): + # 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: # 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 include_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 202482157..84a11daff 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 @@ -37,6 +37,7 @@ class SharedSettingsStackHelper(object): sd_text_encoder = None extra_networks_default_multiplier = None disable_weights_auto_swap = None + prompt_attention = None def __enter__(self): #Save overridden settings so they can be restored later. @@ -52,6 +53,7 @@ class SharedSettingsStackHelper(object): self.sd_text_encoder = shared.opts.sd_text_encoder self.extra_networks_default_multiplier = shared.opts.extra_networks_default_multiplier self.disable_weights_auto_swap = shared.opts.disable_weights_auto_swap + self.prompt_attention = shared.opts.prompt_attention shared.opts.data["disable_weights_auto_swap"] = False def __exit__(self, exc_type, exc_value, tb): @@ -62,6 +64,7 @@ class SharedSettingsStackHelper(object): shared.opts.data["tome_ratio"] = self.tome_ratio shared.opts.data["todo_ratio"] = self.todo_ratio shared.opts.data["extra_networks_default_multiplier"] = self.extra_networks_default_multiplier + shared.opts.data["prompt_attention"] = self.prompt_attention if self.sd_model_checkpoint != shared.opts.sd_model_checkpoint: shared.opts.data["sd_model_checkpoint"] = self.sd_model_checkpoint sd_models.reload_model_weights(op='model') @@ -92,6 +95,7 @@ axis_options = [ 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("[Prompt] Prompt parser", str, apply_setting("prompt_attention"), choices=lambda: ["native", "compel", "xhinker", "a1111", "fixed"]), AxisOption("[Network] LoRA", str, apply_lora, cost=0.5, choices=list_lora), AxisOption("[Network] LoRA strength", float, apply_setting('extra_networks_default_multiplier')), AxisOption("[Network] Styles", str, apply_styles, choices=lambda: [s.name for s in shared.prompt_styles.styles.values()]), @@ -111,7 +115,7 @@ axis_options = [ AxisOption("[Process] Server options", str, apply_options), 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] Sigma method", str, apply_setting("schedulers_sigma"), choices=lambda: ['default', 'karras', 'beta', 'exponential', 'lambdas']), 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")), AxisOption("[Sampler] Solver order", int, apply_setting("schedulers_solver_order")), @@ -132,6 +136,7 @@ axis_options = [ 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_add_label), + AxisOption("[Postprocess] Detailer strength", str, apply_field("detailer_strength")), 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 9a9f2246c..80336fa73 100644 --- a/scripts/xyz_grid_draw.py +++ b/scripts/xyz_grid_draw.py @@ -4,7 +4,7 @@ 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: False, include_time: False): # pylint: disable=unused-argument +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] @@ -40,8 +40,18 @@ def draw_xyz_grid(p, xs, ys, zs, x_labels, y_labels, z_labels, cell, draw_legend 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: - processed_result.images[idx] = images.draw_overlay(processed_result.images[idx], f'time: {p1 - p0:.2f}') + 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] diff --git a/scripts/xyz_grid_on.py b/scripts/xyz_grid_on.py index cdef01e60..202a2cfc4 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,18 +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='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") @@ -140,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, include_time, 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, include_time, 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: @@ -273,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[:] @@ -329,62 +375,41 @@ class Script(scripts.Script): 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): - del processed.images[1] - del processed.all_prompts[1] - del processed.all_seeds[1] - del processed.infotexts[1] - p.do_not_save_grid = True - p.do_not_save_samples = True - active = False - cache = processed - return processed - """ + 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 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_images: # dont need sub-images anymore, drop from list: - if not include_grid and include_subgrids: - 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 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: # auto-save main and sub-grids: - grid_count = z_count + ( 1 if include_grid and z_count > 1 else 0 ) - for g in range(grid_count): + 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: # 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 include_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) p.do_not_save_grid = True p.do_not_save_samples = True @@ -394,6 +419,14 @@ class Script(scripts.Script): 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 + if hasattr(cache, 'used'): + cache.images.clear() + cache.used = False + elif cache is not None and len(cache.images) > 0: + cache.used = True + p.restore_faces = False + p.detailer = False + p.color_corrections = None + p.scripts = None + return cache + return None diff --git a/webui.py b/webui.py index 32e7ed75b..9684ef7c8 100644 --- a/webui.py +++ b/webui.py @@ -51,13 +51,10 @@ fastapi_args = { "version": f'0.0.{git_commit}', "title": "SD.Next", "description": "SD.Next", - "docs_url": "/docs" if cmd_opts.docs else None, - "redoc_url": "/redocs" if cmd_opts.docs else None, - "swagger_ui_parameters": { - "displayOperationId": True, - "showCommonExtensions": True, - "deepLinking": False, - } + "docs_url": None, + "redoc_url": None, + # "docs_url": "/docs" if cmd_opts.docs else None, # custom handler in api.py + # "redoc_url": "/redocs" if cmd_opts.docs else None, } import modules.sd_hijack @@ -110,7 +107,7 @@ def initialize(): yolo.initialize() timer.startup.record("detailer") - log.debug('Load extensions') + log.info('Load extensions') t_timer, t_total = modules.scripts.load_scripts() timer.startup.record("extensions") timer.startup.records["extensions"] = t_total # scripts can reset the time @@ -179,7 +176,7 @@ def load_model(): def create_api(app): - log.debug('Creating API') + log.debug('API initialize') from modules.api.api import Api api = Api(app, queue_lock) return api @@ -231,7 +228,7 @@ def start_common(): def start_ui(): - log.debug('Creating UI') + log.info('UI start') modules.script_callbacks.before_ui_callback() timer.startup.record("before-ui") shared.demo = modules.ui.create_ui(timer.startup) @@ -277,7 +274,6 @@ def start_ui(): max_threads=64, show_api=False, quiet=True, - # favicon_path='html/logo.ico', favicon_path='html/favicon.svg', allowed_paths=allowed_paths, app_kwargs=fastapi_args, diff --git a/wiki b/wiki index b36c2e1a4..713906e92 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit b36c2e1a4cb85338f20061639d4130255d10bf48 +Subproject commit 713906e920e02607ea04951858aabeff7ce641f2