Merge branch 'dev' into refactor-prompt
@@ -90,6 +90,8 @@
|
||||
"getENActiveTab": "readonly",
|
||||
"quickApplyStyle": "readonly",
|
||||
"quickSaveStyle": "readonly",
|
||||
"setupExtraNetworks": "readonly",
|
||||
"showNetworks": "readonly",
|
||||
// from python
|
||||
"localization": "readonly",
|
||||
// progressbar.js
|
||||
@@ -112,6 +114,8 @@
|
||||
"idbPut": "readonly",
|
||||
"idbDel": "readonly",
|
||||
"idbAdd": "readonly",
|
||||
// changelog.js
|
||||
"initChangelog": "readonly",
|
||||
// notification.js
|
||||
"sendNotification": "readonly"
|
||||
},
|
||||
|
||||
@@ -44,6 +44,14 @@ tunableop_results*.csv
|
||||
!webui.sh
|
||||
!package.json
|
||||
|
||||
# pyinstaller
|
||||
*.spec
|
||||
build/
|
||||
dist/
|
||||
|
||||
# dynamically generated
|
||||
/repositories/ip-instruct/
|
||||
|
||||
# all dynamic stuff
|
||||
/extensions/**/*
|
||||
/outputs/**/*
|
||||
@@ -59,7 +67,6 @@ tunableop_results*.csv
|
||||
.vscode/
|
||||
.idea/
|
||||
/localizations
|
||||
|
||||
.*/
|
||||
|
||||
# force included
|
||||
@@ -67,3 +74,4 @@ tunableop_results*.csv
|
||||
!/models/VAE-approx/model.pt
|
||||
!/models/Reference
|
||||
!/models/Reference/**/*
|
||||
|
||||
|
||||
@@ -31,6 +31,9 @@ ignore-paths=/usr/lib/.*$,
|
||||
modules/xadapter,
|
||||
modules/meissonic,
|
||||
modules/omnigen,
|
||||
modules/instantir,
|
||||
modules/consistory,
|
||||
modules/pulid/eva_clip,
|
||||
repositories,
|
||||
extensions-builtin/sd-webui-agent-scheduler,
|
||||
extensions-builtin/sd-extension-chainner/nodes,
|
||||
@@ -130,7 +133,8 @@ confidence=HIGH,
|
||||
INFERENCE_FAILURE,
|
||||
UNDEFINED
|
||||
# disable=C,R,W
|
||||
disable=bad-inline-option,
|
||||
disable=abstract-method,
|
||||
bad-inline-option,
|
||||
bare-except,
|
||||
broad-exception-caught,
|
||||
chained-comparison,
|
||||
@@ -174,6 +178,7 @@ disable=bad-inline-option,
|
||||
unnecessary-dict-index-lookup,
|
||||
unnecessary-dunder-call,
|
||||
unnecessary-lambda,
|
||||
unnecessary-lambda-assigment,
|
||||
use-dict-literal,
|
||||
use-symbolic-message-instead,
|
||||
unknown-option-value,
|
||||
|
||||
@@ -26,6 +26,9 @@ exclude = [
|
||||
"modules/xadapter",
|
||||
"modules/meissonic",
|
||||
"modules/omnigen",
|
||||
"modules/instantir",
|
||||
"modules/consistory",
|
||||
"modules/pulid/eva_clip",
|
||||
"repositories",
|
||||
"extensions-builtin/sd-extension-chainner/nodes",
|
||||
"extensions-builtin/sd-webui-agent-scheduler",
|
||||
|
||||
@@ -1,9 +1,118 @@
|
||||
# Change Log for SD.Next
|
||||
|
||||
## Update for 2024-10-25
|
||||
## Update for 2024-11-06
|
||||
|
||||
Improvements:
|
||||
- Model selector:
|
||||
Smaller release just few days after the last one, but with some important fixes and improvements.
|
||||
This release can be considered an LTS release before we kick off the next round of major updates.
|
||||
|
||||
- Docs:
|
||||
- add built-in [changelog](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) search
|
||||
since changelog is the best up-to-date source of info
|
||||
go to system -> changelog and search/highligh/navigate directly in UI!
|
||||
|
||||
- Integrations:
|
||||
- [PuLID](https://github.com/ToTheBeginning/PuLID): Pure and Lightning ID Customization via Contrastive Alignment
|
||||
- advanced method of face transfer with better quality as well as control over identity and appearance
|
||||
try it out, likely the best quality available for sdxl models
|
||||
- select in *scripts -> pulid*
|
||||
- compatible with *sdxl*
|
||||
- can be used in xyz grid
|
||||
- [InstantIR](https://github.com/instantX-research/InstantIR): Blind Image Restoration with Instant Generative Reference
|
||||
- alternative to traditional `img2img` with more control over restoration process
|
||||
- select in *image -> scripts -> instantir*
|
||||
- compatible with *sdxl*
|
||||
- *note*: after used once it cannot be unloaded without reloading base model
|
||||
- [ConsiStory](https://github.com/NVlabs/consistory): Consistent Image Generation
|
||||
- create consistent anchor image and then generate images that are consistent with anchor
|
||||
- select in *scripts -> consistory*
|
||||
- compatible with *sdxl*
|
||||
- *note*: very resource intensive and not compatible with model offloading
|
||||
- *note*: changing default parameters can lead to unexpected results and/or failures
|
||||
- *note*: after used once it cannot be unloaded without reloading base model
|
||||
- [MiaoshouAI PromptGen v2.0](https://huggingface.co/MiaoshouAI/Florence-2-base-PromptGen-v2.0) base and large:
|
||||
- *in process -> visual query*
|
||||
- caption modes:
|
||||
`<GENERATE_TAGS>` generate tags
|
||||
`<CAPTION>`, `<DETAILED_CAPTION>`, `<MORE_DETAILED_CAPTION>` caption image
|
||||
`<ANALYZE>` image composition
|
||||
`<MIXED_CAPTION>`, `<MIXED_CAPTION_PLUS>` detailed caption and tags with optional analyze
|
||||
|
||||
- Model improvements:
|
||||
- SD3: ControlNets:
|
||||
- *InstantX Canny, Pose, Depth, Tile*
|
||||
- *Alimama Inpainting, SoftEdge*
|
||||
- *note*: that just like with FLUX.1 or any large model, ControlNet are also large and can push your system over the limit
|
||||
e.g. SD3 controlnets vary from 1GB to over 4GB in size
|
||||
- SD3: all-in-one safetensors
|
||||
- *examples*: [large](https://civitai.com/models/882666/sd35-large-google-flan?modelVersionId=1003031), [medium](https://civitai.com/models/900327)
|
||||
- *note*: enable *bnb* on-the-fly quantization for even bigger gains
|
||||
|
||||
- Workflow improvements:
|
||||
- XYZ grid:
|
||||
- optional time benchmark info to individual images
|
||||
- optional add params to individual images
|
||||
- create video from generated grid images
|
||||
supports all standard video types and interpolation
|
||||
- UI:
|
||||
- add additional [hotkeys](https://github.com/vladmandic/automatic/wiki/Hotkeys)
|
||||
- add show networks on startup setting
|
||||
- better mapping of networks previews
|
||||
- optimize networks display load
|
||||
- Other:
|
||||
- Installer:
|
||||
- Log `venv` and package search paths
|
||||
- Auto-remove invalid packages from `venv/site-packages`
|
||||
e.g. packages starting with `~` which are left-over due to windows access violation
|
||||
- Requirements: update
|
||||
- Model loader:
|
||||
- Report modules included in safetensors when attempting to load a model
|
||||
- CLI:
|
||||
- refactor command line params
|
||||
run `webui.sh`/`webui.bat` with `--help` to see all options
|
||||
- added `cli/model-metadata.py` to display metadata in any safetensors file
|
||||
- added `cli/model-keys.py` to quicky display content of any safetensors file
|
||||
- Internal:
|
||||
- Repo: move screenshots to GH pages
|
||||
|
||||
- Fixes:
|
||||
- custom watermark add alphablending
|
||||
- detailer min/max size as fractions of image size
|
||||
- ipadapter load on-demand
|
||||
- ipadapter face use correct yolo model
|
||||
- list diffusers remove duplicates
|
||||
- fix legacy extensions access to shared objects
|
||||
- fix diffusers load from folder
|
||||
- fix lora enum logging on windows
|
||||
- fix xyz grid with batch count
|
||||
- fix xyz grid include images
|
||||
- fix xyz skip on interrupted
|
||||
- fix vqa models ignoring hfcache folder setting
|
||||
- fix network height in standard vs modern ui
|
||||
- fix k-diff enum on startup
|
||||
- fix text2video scripts
|
||||
- dont uninstall flash-attn
|
||||
- move downloads of some auxillary models to hfcache instead of models folder
|
||||
|
||||
## Update for 2024-10-29
|
||||
|
||||
### Highlights for 2024-10-29
|
||||
|
||||
- Support for **all SD3.x variants**
|
||||
*SD3.0-Medium, SD3.5-Medium, SD3.5-Large, SD3.0-Large-Turbo*
|
||||
- Allow quantization using `bitsandbytes` on-the-fly during models load
|
||||
Load any variant of SD3.x or FLUX.1 and apply quantization during load without the need for pre-quantized models
|
||||
- Allow for custom model URL in standard model selector
|
||||
Can be used to specify any model from *HuggingFace* or *CivitAI*
|
||||
- Full support for `torch==2.5.1`
|
||||
- New wiki articles: [Gated Access](https://github.com/vladmandic/automatic/wiki/Gated), [Quantization](https://github.com/vladmandic/automatic/wiki/Quantization), [Offloading](https://github.com/vladmandic/automatic/wiki/Offload)
|
||||
|
||||
Plus tons of smaller improvements and cumulative fixes reported since last release
|
||||
|
||||
[README](https://github.com/vladmandic/automatic/blob/master/README.md) | [CHANGELOG](https://github.com/vladmandic/automatic/blob/master/CHANGELOG.md) | [WiKi](https://github.com/vladmandic/automatic/wiki) | [Discord](https://discord.com/invite/sd-next-federal-batch-inspectors-1101998836328697867)
|
||||
|
||||
### Details for 2024-10-29
|
||||
|
||||
- model selector:
|
||||
- change-in-behavior
|
||||
- when typing, it will auto-load model as soon as exactly one match is found
|
||||
- allows entering model that are not on the list which triggers huggingface search
|
||||
@@ -14,16 +123,40 @@ Improvements:
|
||||
e.g. `https://civitai.com/api/download/models/72396?type=Model&format=SafeTensor&size=full&fp=fp16`
|
||||
- auto-search-and-download can be disabled in settings -> models -> auto-download
|
||||
this also disables reference models as they are auto-downloaded on first use as well
|
||||
- SD3 loader enhancements
|
||||
- sd3 enhancements:
|
||||
- allow on-the-fly bnb quantization during load
|
||||
- report when loading incomplete model
|
||||
- handle missing model components
|
||||
- handle missing model components during load
|
||||
- handle component preloading
|
||||
- native lora handler
|
||||
- support for all sd35 variants: *medium/large/large-turbo*
|
||||
- gguf transformer loader (prototype)
|
||||
- OpenVINO: add accuracy option
|
||||
- ZLUDA: guess GPU arch
|
||||
- flux.1 enhancements:
|
||||
- allow on-the-fly bnb quantization during load
|
||||
- samplers:
|
||||
- support for original k-diffusion samplers
|
||||
select via *scripts -> k-diffusion -> sampler*
|
||||
- ipadapter:
|
||||
- list available adapters based on loaded model type
|
||||
- add adapter `ostris consistency` for sd15/sdxl
|
||||
- detailer:
|
||||
- add `[prompt]` to refine/defailer prompts as placeholder referencing original prompt
|
||||
- torch
|
||||
- use `torch==2.5.1` by default on supported platforms
|
||||
- CUDA set device memory limit
|
||||
in *settings -> compute settings -> torch memory limit*
|
||||
default=0 meaning no limit, if set torch will limit memory usage to specified fraction
|
||||
*note*: this is not a hard limit, torch will try to stay under this value
|
||||
- compute backends:
|
||||
- OpenVINO: add accuracy option
|
||||
- ZLUDA: guess GPU arch
|
||||
- major model load refactor
|
||||
- wiki: new articles
|
||||
- [Gated Access Wiki](https://github.com/vladmandic/automatic/wiki/Gated)
|
||||
- [Quantization Wiki](https://github.com/vladmandic/automatic/wiki/Quantization)
|
||||
- [Offloading Wiki](https://github.com/vladmandic/automatic/wiki/Offload)
|
||||
|
||||
Fixes:
|
||||
fixes:
|
||||
- fix send-to-control
|
||||
- fix k-diffusion
|
||||
- fix sd3 img2img and hires
|
||||
|
||||
@@ -24,5 +24,5 @@ abstract: >-
|
||||
generation
|
||||
keywords:
|
||||
- stablediffusion diffusers sdnext
|
||||
license: AGPL-3.0
|
||||
license: Apache-2.0
|
||||
date-released: 2022-12-24
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
GNU AFFERO GENERAL PUBLIC LICENSE
|
||||
Version 3, 19 November 2007
|
||||
|
||||
Copyright (c) 2023 AUTOMATIC1111
|
||||
|
||||
Copyright (C) 2007 Free Software Foundation, Inc. <https://fsf.org/>
|
||||
Everyone is permitted to copy and distribute verbatim copies
|
||||
of this license document, but changing it is not allowed.
|
||||
@@ -635,8 +633,8 @@ the "copyright" line and a pointer to where the full notice is found.
|
||||
Copyright (C) <year> <name of author>
|
||||
|
||||
This program is free software: you can redistribute it and/or modify
|
||||
it under the terms of the GNU Affero General Public License as published by
|
||||
the Free Software Foundation, either version 3 of the License, or
|
||||
it under the terms of the GNU Affero General Public License as published
|
||||
by the Free Software Foundation, either version 3 of the License, or
|
||||
(at your option) any later version.
|
||||
|
||||
This program is distributed in the hope that it will be useful,
|
||||
|
||||
@@ -50,12 +50,13 @@ All individual features are not listed here, instead check [ChangeLog](CHANGELOG
|
||||
<br>
|
||||
|
||||
*Main interface using **StandardUI***:
|
||||

|
||||

|
||||
|
||||
*Main interface using **ModernUI***:
|
||||

|
||||

|
||||

|
||||
|
||||

|
||||

|
||||

|
||||
|
||||
For screenshots and informations on other available themes, see [Themes Wiki](https://github.com/vladmandic/automatic/wiki/Themes)
|
||||
|
||||
@@ -63,12 +64,13 @@ For screenshots and informations on other available themes, see [Themes Wiki](ht
|
||||
|
||||
## Model support
|
||||
|
||||
Additional models will be added as they become available and there is public interest in them
|
||||
Additional models will be added as they become available and there is public interest in them
|
||||
See [models overview](https://github.com/vladmandic/automatic/wiki/Models) for details on each model, including their architecture, complexity and other info
|
||||
|
||||
- [RunwayML Stable Diffusion](https://github.com/Stability-AI/stablediffusion/) 1.x and 2.x *(all variants)*
|
||||
- [StabilityAI Stable Diffusion XL](https://github.com/Stability-AI/generative-models)
|
||||
- [StabilityAI Stable Diffusion 3 Medium](https://stability.ai/news/stable-diffusion-3-medium)
|
||||
- [Stable Diffusion 3.5 Large](https://huggingface.co/stabilityai/stable-diffusion-3.5-large)
|
||||
- [StabilityAI Stable Diffusion](https://stability.ai/news/stable-diffusion-3-medium)
|
||||
- [Stable Diffusion 3.x](https://huggingface.co/stabilityai/stable-diffusion-3.5-large) 3.0 Medium, 3.5 Medium, 3.5 Large, 3.5 Large Turbo
|
||||
- [StabilityAI Stable Video Diffusion](https://huggingface.co/stabilityai/stable-video-diffusion-img2vid) Base, XT 1.0, XT 1.1
|
||||
- [StabilityAI Stable Cascade](https://github.com/Stability-AI/StableCascade) *Full* and *Lite*
|
||||
- [Black Forest Labs FLUX.1](https://blackforestlabs.ai/announcing-black-forest-labs/) Dev, Schnell
|
||||
@@ -136,13 +138,13 @@ Also supported are modifiers such as:
|
||||
## Examples
|
||||
|
||||
*IP Adapters*:
|
||||

|
||||

|
||||
|
||||
*Color grading*:
|
||||

|
||||

|
||||
|
||||
*InstantID*:
|
||||

|
||||

|
||||
|
||||
> [!IMPORTANT]
|
||||
> - Loading any model other than standard SD 1.x / SD 2.x requires use of backend **Diffusers**
|
||||
@@ -173,92 +175,34 @@ Also supported are modifiers such as:
|
||||
|
||||
Once SD.Next is installed, simply run `webui.ps1` or `webui.bat` (*Windows*) or `webui.sh` (*Linux or MacOS*)
|
||||
|
||||
List of available parameters, run `webui --help` for the full & up-to-date list:
|
||||
For list of available command line options, run `webui --help` for the full & up-to-date list
|
||||
|
||||
Server options:
|
||||
--config CONFIG Use specific server configuration file, default: config.json
|
||||
--ui-config UI_CONFIG Use specific UI configuration file, default: ui-config.json
|
||||
--medvram Split model stages and keep only active part in VRAM, default: False
|
||||
--lowvram Split model components and keep only active part in VRAM, default: False
|
||||
--ckpt CKPT Path to model checkpoint to load immediately, default: None
|
||||
--vae VAE Path to VAE checkpoint to load immediately, default: None
|
||||
--data-dir DATA_DIR Base path where all user data is stored, default:
|
||||
--models-dir MODELS_DIR Base path where all models are stored, default: models
|
||||
--allow-code Allow custom script execution, default: False
|
||||
--share Enable UI accessible through Gradio site, default: False
|
||||
--insecure Enable extensions tab regardless of other options, default: False
|
||||
--use-cpu USE_CPU [USE_CPU ...] Force use CPU for specified modules, default: []
|
||||
--listen Launch web server using public IP address, default: False
|
||||
--port PORT Launch web server with given server port, default: 7860
|
||||
--freeze Disable editing settings
|
||||
--auth AUTH Set access authentication like "user:pwd,user:pwd""
|
||||
--auth-file AUTH_FILE Set access authentication using file, default: None
|
||||
--autolaunch Open the UI URL in the system's default browser upon launch
|
||||
--docs Mount API docs, default: False
|
||||
--api-only Run in API only mode without starting UI
|
||||
--api-log Enable logging of all API requests, default: False
|
||||
--device-id DEVICE_ID Select the default CUDA device to use, default: None
|
||||
--cors-origins CORS_ORIGINS Allowed CORS origins as comma-separated list, default: None
|
||||
--cors-regex CORS_REGEX Allowed CORS origins as regular expression, default: None
|
||||
--tls-keyfile TLS_KEYFILE Enable TLS and specify key file, default: None
|
||||
--tls-certfile TLS_CERTFILE Enable TLS and specify cert file, default: None
|
||||
--tls-selfsign Enable TLS with self-signed certificates, default: False
|
||||
--server-name SERVER_NAME Sets hostname of server, default: None
|
||||
--no-hashing Disable hashing of checkpoints, default: False
|
||||
--no-metadata Disable reading of metadata from models, default: False
|
||||
--disable-queue Disable queues, default: False
|
||||
--subpath SUBPATH Customize the URL subpath for usage with reverse proxy
|
||||
--backend {original,diffusers} force model pipeline type
|
||||
--allowed-paths ALLOWED_PATHS [ALLOWED_PATHS ...] add additional paths to paths allowed for web access
|
||||
|
||||
Setup options:
|
||||
--reset Reset main repository to latest version, default: False
|
||||
--upgrade Upgrade main repository to latest version, default: False
|
||||
--requirements Force re-check of requirements, default: False
|
||||
--quick Bypass version checks, default: False
|
||||
--use-directml Use DirectML if no compatible GPU is detected, default: False
|
||||
--use-openvino Use Intel OpenVINO backend, default: False
|
||||
--use-ipex Force use Intel OneAPI XPU backend, default: False
|
||||
--use-cuda Force use nVidia CUDA backend, default: False
|
||||
--use-rocm Force use AMD ROCm backend, default: False
|
||||
--use-zluda Force use ZLUDA, AMD GPUs only, default: False
|
||||
--use-xformers Force use xFormers cross-optimization, default: False
|
||||
--skip-requirements Skips checking and installing requirements, default: False
|
||||
--skip-extensions Skips running individual extension installers, default: False
|
||||
--skip-git Skips running all GIT operations, default: False
|
||||
--skip-torch Skips running Torch checks, default: False
|
||||
--skip-all Skips running all checks, default: False
|
||||
--skip-env Skips setting of env variables during startup, default: False
|
||||
--experimental Allow unsupported versions of libraries, default: False
|
||||
--reinstall Force reinstallation of all requirements, default: False
|
||||
--test Run test only and exit
|
||||
--version Print version information
|
||||
--ignore Ignore any errors and attempt to continue
|
||||
--safe Run in safe mode with no user extensions
|
||||
--uv Use uv as installer, default: False
|
||||
|
||||
Logging options:
|
||||
--log LOG Set log file, default: None
|
||||
--debug Run installer with debug logging, default: False
|
||||
--profile Run profiler, default: False
|
||||
> [!TIP]
|
||||
> All command line options can also be set via env variable
|
||||
> For example `--debug` is same as `set SD_DEBUG=true`
|
||||
|
||||
## Notes
|
||||
|
||||
> [!TIP]
|
||||
> If you don't want to use built-in `venv` support and prefer to run SD.Next in your own environment such as *Docker* container, *Conda* environment or any other virtual environment, you can skip `venv` create/activate and launch SD.Next directly using `python launch.py` (command line flags noted above still apply).
|
||||
|
||||
### Quantization
|
||||
|
||||
**SD.Next** comes with broad quantization support, including support for BitsAndBytes, Optimum.Quanto, TorchAO, NNCF and GGUF
|
||||
See [Quantization Wiki](https://github.com/vladmandic/automatic/wiki/Quantization)
|
||||
|
||||
### Control
|
||||
|
||||
**SD.Next** comes with built-in control for all types of text2image, image2image, video2video and batch processing
|
||||
|
||||
*Control interface*:
|
||||

|
||||

|
||||
|
||||
*Control processors*:
|
||||

|
||||

|
||||
|
||||
*Masking*:
|
||||

|
||||

|
||||
|
||||
### Extensions
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma
|
||||
|
||||
- async lowvram: <https://github.com/AUTOMATIC1111/stable-diffusion-webui/pull/14855>
|
||||
- fp8: <https://github.com/AUTOMATIC1111/stable-diffusion-webui/pull/14031>
|
||||
- ipadapter-negative: https://github.com/huggingface/diffusers/discussions/7167
|
||||
- ipadapter-negative: <https://github.com/huggingface/diffusers/discussions/7167>
|
||||
- include reference styles
|
||||
|
||||
### Missing
|
||||
|
||||
@@ -1,16 +1,43 @@
|
||||
# Stable-Diffusion Productivity Scripts
|
||||
|
||||
Note: All scripts have built-in `--help` parameter that can be used to get more information
|
||||
## API Examples
|
||||
|
||||
<br>
|
||||
### Run Generate
|
||||
|
||||
## Main Scripts
|
||||
- `cli/api-txt2img.py`
|
||||
- `cli/api-img2img.py`
|
||||
- `cli/api-control.py`
|
||||
|
||||
### Generate
|
||||
### Monitor
|
||||
|
||||
- `cli/api-progress.py`
|
||||
|
||||
### Generic
|
||||
|
||||
- `cli/api-json.py`
|
||||
|
||||
### Process
|
||||
|
||||
- `cli/api-info.py`
|
||||
- `cli/api-upscale.py`
|
||||
- `cli/api-vqa.py`
|
||||
- `cli/api-preprocess.py`
|
||||
|
||||
### Other
|
||||
|
||||
- `cli/api-faceid.py`
|
||||
- `cli/api-faces.py`
|
||||
- `cli/api-mask.py`
|
||||
|
||||
### JavaScript
|
||||
|
||||
- `cli/api-txt2img.js`
|
||||
|
||||
## Generate
|
||||
|
||||
Text-to-image with all of the possible parameters
|
||||
Supports upsampling, face restoration and grid creation
|
||||
> python generate.py
|
||||
> python cli/generate.py
|
||||
|
||||
By default uses parameters from `generate.json`
|
||||
|
||||
@@ -20,25 +47,6 @@ Parameters that are not specified will be randomized:
|
||||
- Sampler/Scheduler will be randomly picked from available ones
|
||||
- CFG Scale set to 5-10
|
||||
|
||||
### Train
|
||||
|
||||
Combined pipeline for **embeddings**, **lora**, **lycoris**, **dreambooth** and **hypernetwork**
|
||||
Optionally runs several image processing steps before training:
|
||||
|
||||
- keep original image
|
||||
- detect and extract face
|
||||
- detect and extract body
|
||||
- detect blur
|
||||
- detect dynamic range
|
||||
- attempt to upscale low resolution images
|
||||
- attempt to restore quality of low quality images
|
||||
- automatically generate captions using interrogate
|
||||
- resize image
|
||||
- square image
|
||||
- run image segmentation to remove background
|
||||
|
||||
> python train.py
|
||||
|
||||
<br>
|
||||
|
||||
## Auxiliary Scripts
|
||||
|
||||
@@ -82,6 +82,7 @@ def watermark(params, file):
|
||||
|
||||
exif = get_exif(image)
|
||||
|
||||
wm = None
|
||||
if params.command == 'read':
|
||||
fn = params.input
|
||||
wm = get_watermark(image, params)
|
||||
|
||||
@@ -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()
|
||||
@@ -127,6 +127,8 @@ def load_diffusers(name, network_on_disk, lora_scale=shared.opts.extra_networks_
|
||||
|
||||
|
||||
def load_network(name, network_on_disk) -> network.Network:
|
||||
if not shared.sd_loaded:
|
||||
return None
|
||||
t0 = time.time()
|
||||
cached = lora_cache.get(name, None)
|
||||
if debug:
|
||||
@@ -531,18 +533,14 @@ def network_MultiheadAttention_load_state_dict(self, *args, **kwargs):
|
||||
|
||||
|
||||
def list_available_networks():
|
||||
t0 = time.time()
|
||||
available_networks.clear()
|
||||
available_network_aliases.clear()
|
||||
forbidden_network_aliases.clear()
|
||||
available_network_hash_lookup.clear()
|
||||
forbidden_network_aliases.update({"none": 1, "Addams": 1})
|
||||
directories = []
|
||||
if os.path.exists(shared.cmd_opts.lora_dir):
|
||||
directories.append(shared.cmd_opts.lora_dir)
|
||||
else:
|
||||
if not os.path.exists(shared.cmd_opts.lora_dir):
|
||||
shared.log.warning(f'LoRA directory not found: path="{shared.cmd_opts.lora_dir}"')
|
||||
if os.path.exists(shared.cmd_opts.lyco_dir) and shared.cmd_opts.lyco_dir != shared.cmd_opts.lora_dir:
|
||||
directories.append(shared.cmd_opts.lyco_dir)
|
||||
|
||||
def add_network(filename):
|
||||
if not os.path.isfile(filename):
|
||||
@@ -563,11 +561,12 @@ def list_available_networks():
|
||||
except OSError as e: # should catch FileNotFoundError and PermissionError etc.
|
||||
shared.log.error(f'LoRA: filename="{filename}" {e}')
|
||||
|
||||
candidates = list(files_cache.list_files(*directories, ext_filter=[".pt", ".ckpt", ".safetensors"]))
|
||||
candidates = list(files_cache.list_files(shared.cmd_opts.lora_dir, ext_filter=[".pt", ".ckpt", ".safetensors"]))
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=shared.max_workers) as executor:
|
||||
for fn in candidates:
|
||||
executor.submit(add_network, fn)
|
||||
shared.log.info(f'Available LoRAs: items={len(available_networks)} folders={len(forbidden_network_aliases)}')
|
||||
t1 = time.time()
|
||||
shared.log.info(f'Available LoRAs: path="{shared.cmd_opts.lora_dir}" items={len(available_networks)} folders={len(forbidden_network_aliases)} time={t1 - t0:.2f}')
|
||||
|
||||
|
||||
def infotext_pasted(infotext, params): # pylint: disable=W0613
|
||||
|
||||
@@ -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" <lora:{l.get_alias()}:{shared.opts.extra_networks_default_multiplier}>"),
|
||||
"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" <lora:{l.get_alias()}:{shared.opts.extra_networks_default_multiplier}>"),
|
||||
"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}')
|
||||
|
||||
|
Before Width: | Height: | Size: 438 KiB After Width: | Height: | Size: 101 KiB |
|
Before Width: | Height: | Size: 52 KiB After Width: | Height: | Size: 32 KiB |
|
Before Width: | Height: | Size: 40 KiB |
|
Before Width: | Height: | Size: 62 KiB After Width: | Height: | Size: 36 KiB |
|
Before Width: | Height: | Size: 26 KiB After Width: | Height: | Size: 31 KiB |
|
Before Width: | Height: | Size: 53 KiB After Width: | Height: | Size: 34 KiB |
|
Before Width: | Height: | Size: 50 KiB After Width: | Height: | Size: 34 KiB |
|
Before Width: | Height: | Size: 315 KiB After Width: | Height: | Size: 33 KiB |
|
Before Width: | Height: | Size: 50 KiB After Width: | Height: | Size: 35 KiB |
|
Before Width: | Height: | Size: 47 KiB After Width: | Height: | Size: 38 KiB |
|
Before Width: | Height: | Size: 42 KiB After Width: | Height: | Size: 33 KiB |
@@ -0,0 +1,11 @@
|
||||
{
|
||||
"stabilityai--stable-diffusion-3-medium-diffusers": "models/Reference/stabilityai--stable-diffusion-3.jpg",
|
||||
"stabilityai--stable-diffusion-3.5-medium": "models/Reference/stabilityai--stable-diffusion-3_5.jpg",
|
||||
"stabilityai--stable-diffusion-3.5-large": "models/Reference/stabilityai--stable-diffusion-3_5.jpg",
|
||||
"Disty0--FLUX.1-dev-qint8": "models/Reference/black-forest-labs--FLUX.1-dev.jpg",
|
||||
"Disty0--FLUX.1-dev-qint4": "models/Reference/black-forest-labs--FLUX.1-dev.jpg",
|
||||
"sayakpaul--flux.1-dev-nf4": "models/Reference/black-forest-labs--FLUX.1-dev.jpg",
|
||||
"THUDM--CogVideoX-2b": "models/Reference/THUDM--CogView3-Plus-3B.jpg",
|
||||
"THUDM--CogVideoX-5b": "models/Reference/THUDM--CogView3-Plus-3B.jpg",
|
||||
"THUDM--CogVideoX-5b-I2V": "models/Reference/THUDM--CogView3-Plus-3B.jpg"
|
||||
}
|
||||
@@ -119,11 +119,19 @@
|
||||
"preview": "stabilityai--stable-diffusion-3.jpg",
|
||||
"extras": "sampler: Default, cfg_scale: 7.0"
|
||||
},
|
||||
"StabilityAI Stable Diffusion 3.5 Medium": {
|
||||
"path": "stabilityai/stable-diffusion-3.5-medium",
|
||||
"skip": true,
|
||||
"variant": "fp16",
|
||||
"desc": "Stable Diffusion 3.5 Medium is a Multimodal Diffusion Transformer with improvements (MMDiT-X) text-to-image model that features improved performance in image quality, typography, complex prompt understanding, and resource-efficiency.",
|
||||
"preview": "stabilityai--stable-diffusion-3_5.jpg",
|
||||
"extras": "sampler: Default, cfg_scale: 7.0"
|
||||
},
|
||||
"StabilityAI Stable Diffusion 3.5 Large": {
|
||||
"path": "stabilityai/stable-diffusion-3.5-large",
|
||||
"skip": true,
|
||||
"variant": "fp16",
|
||||
"desc": "Stable Diffusion 3 Medium is a Multimodal Diffusion Transformer (MMDiT) text-to-image model that features greatly improved performance in image quality, typography, complex prompt understanding, and resource-efficiency",
|
||||
"desc": "Stable Diffusion 3.5 Large is a Multimodal Diffusion Transformer (MMDiT) text-to-image model that features improved performance in image quality, typography, complex prompt understanding, and resource-efficiency.",
|
||||
"preview": "stabilityai--stable-diffusion-3_5.jpg",
|
||||
"extras": "sampler: Default, cfg_scale: 7.0"
|
||||
},
|
||||
@@ -131,7 +139,7 @@
|
||||
"path": "stabilityai/stable-diffusion-3.5-large-turbo",
|
||||
"skip": true,
|
||||
"variant": "fp16",
|
||||
"desc": "Stable Diffusion 3 Medium is a Multimodal Diffusion Transformer (MMDiT) text-to-image model that features greatly improved performance in image quality, typography, complex prompt understanding, and resource-efficiency",
|
||||
"desc": "Stable Diffusion 3.5 Large Turbo is a Multimodal Diffusion Transformer (MMDiT) text-to-image model with Adversarial Diffusion Distillation (ADD) that features improved performance in image quality, typography, complex prompt understanding, and resource-efficiency, with a focus on fewer inference steps.",
|
||||
"preview": "stabilityai--stable-diffusion-3_5.jpg",
|
||||
"extras": "sampler: Default, cfg_scale: 7.0"
|
||||
},
|
||||
|
||||
|
Before Width: | Height: | Size: 76 KiB |
|
Before Width: | Height: | Size: 92 KiB |
|
Before Width: | Height: | Size: 104 KiB |
|
Before Width: | Height: | Size: 63 KiB |
|
Before Width: | Height: | Size: 63 KiB |
|
Before Width: | Height: | Size: 196 KiB |
|
Before Width: | Height: | Size: 99 KiB |
|
Before Width: | Height: | Size: 101 KiB |
|
Before Width: | Height: | Size: 66 KiB |
|
Before Width: | Height: | Size: 92 KiB |
|
Before Width: | Height: | Size: 154 KiB |
|
Before Width: | Height: | Size: 162 KiB |
|
Before Width: | Height: | Size: 155 KiB |
|
Before Width: | Height: | Size: 193 KiB |
|
Before Width: | Height: | Size: 154 KiB |
|
Before Width: | Height: | Size: 100 KiB |
|
Before Width: | Height: | Size: 80 KiB |
|
Before Width: | Height: | Size: 65 KiB |
|
Before Width: | Height: | Size: 102 KiB |
@@ -227,9 +227,9 @@ def installed(package, friendly: str = None, reload = False, quiet = False):
|
||||
exact = pkg_version == p[1]
|
||||
if not exact and not quiet:
|
||||
if args.experimental:
|
||||
log.warning(f"Package: {p[0]} {pkg_version} required {p[1]} allowing experimental")
|
||||
log.warning(f"Package: {p[0]} installed={pkg_version} required={p[1]} allowing experimental")
|
||||
else:
|
||||
log.warning(f"Package: {p[0]} {pkg_version} required {p[1]} version mismatch")
|
||||
log.warning(f"Package: {p[0]} installed={pkg_version} required={p[1]} version mismatch")
|
||||
ok = ok and (exact or args.experimental)
|
||||
else:
|
||||
if not quiet:
|
||||
@@ -254,11 +254,12 @@ def uninstall(package, quiet = False):
|
||||
@lru_cache()
|
||||
def pip(arg: str, ignore: bool = False, quiet: bool = False, uv = True):
|
||||
originalArg = arg
|
||||
uv = uv and args.uv
|
||||
pipCmd = "uv pip" if uv else "pip"
|
||||
arg = arg.replace('>=', '==')
|
||||
package = arg.replace("install", "").replace("--upgrade", "").replace("--no-deps", "").replace("--force", "").replace(" ", " ").strip()
|
||||
uv = uv and args.uv and not package.startswith('git+')
|
||||
pipCmd = "uv pip" if uv else "pip"
|
||||
if not quiet and '-r ' not in arg:
|
||||
log.info(f'Install: package="{arg.replace("install", "").replace("--upgrade", "").replace("--no-deps", "").replace("--force", "").replace(" ", " ").strip()}" mode={"uv" if uv else "pip"}')
|
||||
log.info(f'Install: package="{package}" mode={"uv" if uv else "pip"}')
|
||||
env_args = os.environ.get("PIP_EXTRA_ARGS", "")
|
||||
all_args = f'{pip_log}{arg} {env_args}'.strip()
|
||||
if not quiet:
|
||||
@@ -376,10 +377,11 @@ def update(folder, keep_branch = False, rebase = True):
|
||||
else:
|
||||
res = git(f'pull origin {b} {arg}', folder)
|
||||
debug(f'Install update: folder={folder} branch={b} args={arg} {res}')
|
||||
commit = extensions_commit.get(os.path.basename(folder), None)
|
||||
if commit is not None:
|
||||
res = git(f'checkout {commit}', folder)
|
||||
debug(f'Install update: folder={folder} branch={b} args={arg} commit={commit} {res}')
|
||||
if not args.experimental:
|
||||
commit = extensions_commit.get(os.path.basename(folder), None)
|
||||
if commit is not None:
|
||||
res = git(f'checkout {commit}', folder)
|
||||
debug(f'Install update: folder={folder} branch={b} args={arg} commit={commit} {res}')
|
||||
return res
|
||||
|
||||
|
||||
@@ -454,7 +456,7 @@ def check_python(supported_minors=[9, 10, 11, 12], reason=None):
|
||||
|
||||
# check diffusers version
|
||||
def check_diffusers():
|
||||
sha = '435f6b7e47c031f98b8374b1689e1abeb17bfdb6'
|
||||
sha = '0d1d267b12e47b40b0e8f265339c76e0f45f8c49'
|
||||
pkg = pkg_resources.working_set.by_key.get('diffusers', None)
|
||||
minor = int(pkg.version.split('.')[1] if pkg is not None else 0)
|
||||
cur = opts.get('diffusers_version', '') if minor > 0 else ''
|
||||
@@ -489,7 +491,7 @@ def install_cuda():
|
||||
log.info('CUDA: nVidia toolkit detected')
|
||||
install('onnxruntime-gpu', 'onnxruntime-gpu', ignore=True, quiet=True)
|
||||
# return os.environ.get('TORCH_COMMAND', 'torch torchvision --index-url https://download.pytorch.org/whl/cu124')
|
||||
return os.environ.get('TORCH_COMMAND', 'torch==2.4.1+cu124 torchvision==0.19.1+cu124 --index-url https://download.pytorch.org/whl/cu124')
|
||||
return os.environ.get('TORCH_COMMAND', 'torch==2.5.1+cu124 torchvision==0.20.1+cu124 --index-url https://download.pytorch.org/whl/cu124')
|
||||
|
||||
|
||||
def install_rocm_zluda():
|
||||
@@ -546,11 +548,11 @@ def install_rocm_zluda():
|
||||
os.environ['HIP_VISIBLE_DEVICES'] = args.device_id
|
||||
del args.device_id
|
||||
|
||||
log.warning("ZLUDA support: experimental")
|
||||
error = None
|
||||
from modules import zluda_installer
|
||||
zluda_installer.set_default_agent(device)
|
||||
try:
|
||||
if args.reinstall_zluda:
|
||||
if args.reinstall:
|
||||
zluda_installer.uninstall()
|
||||
zluda_path = zluda_installer.get_path()
|
||||
zluda_installer.install(zluda_path)
|
||||
@@ -570,8 +572,10 @@ def install_rocm_zluda():
|
||||
log.info('Using CPU-only torch')
|
||||
torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision')
|
||||
else:
|
||||
if rocm.version is None or float(rocm.version) >= 6.1: # assume the latest if version check fails
|
||||
#torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision --index-url https://download.pytorch.org/whl/rocm6.1')
|
||||
if rocm.version is None or float(rocm.version) > 6.1: # assume the latest if version check fails
|
||||
# torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.5.1+rocm6.2 torchvision==0.20.1+rocm6.2 --index-url https://download.pytorch.org/whl/rocm6.2')
|
||||
torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.4.1+rocm6.1 torchvision==0.19.1+rocm6.1 --index-url https://download.pytorch.org/whl/rocm6.1')
|
||||
elif rocm.version == "6.1": # lock to 2.4.1, older rocm (5.7) uses torch 2.3
|
||||
torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.4.1+rocm6.1 torchvision==0.19.1+rocm6.1 --index-url https://download.pytorch.org/whl/rocm6.1')
|
||||
elif rocm.version == "6.0": # lock to 2.4.1, older rocm (5.7) uses torch 2.3
|
||||
torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.4.1+rocm6.0 torchvision==0.19.1+rocm6.0 --index-url https://download.pytorch.org/whl/rocm6.0')
|
||||
@@ -591,11 +595,11 @@ def install_rocm_zluda():
|
||||
install(ort_package, 'onnxruntime-training')
|
||||
|
||||
if installed("torch") and device is not None:
|
||||
if 'Flash attention' in opts.get('sdp_options'):
|
||||
if 'Flash attention' in opts.get('sdp_options', ''):
|
||||
if not installed('flash-attn'):
|
||||
install(rocm.get_flash_attention_command(device), reinstall=True)
|
||||
elif not args.experimental:
|
||||
uninstall('flash-attn')
|
||||
#elif not args.experimental:
|
||||
# uninstall('flash-attn')
|
||||
|
||||
if device is not None and rocm.version != "6.2" and rocm.version == rocm.version_torch and rocm.get_blaslt_enabled():
|
||||
log.debug(f'ROCm hipBLASLt: arch={device.name} available={device.blaslt_supported}')
|
||||
@@ -730,7 +734,7 @@ def check_torch():
|
||||
else:
|
||||
if args.use_zluda:
|
||||
log.warning("ZLUDA failed to initialize: no HIP SDK found")
|
||||
log.info('Using CPU-only Torch')
|
||||
log.warning('Torch: CPU-only version installed')
|
||||
torch_command = os.environ.get('TORCH_COMMAND', 'torch torchvision')
|
||||
if 'torch' in torch_command and not args.version:
|
||||
install(torch_command, 'torch torchvision', quiet=True)
|
||||
@@ -817,6 +821,7 @@ def install_packages():
|
||||
log.info('Verifying packages')
|
||||
clip_package = os.environ.get('CLIP_PACKAGE', "git+https://github.com/openai/CLIP.git")
|
||||
install(clip_package, 'clip', quiet=True)
|
||||
install('open-clip-torch', no_deps=True, quiet=True)
|
||||
# tensorflow_package = os.environ.get('TENSORFLOW_PACKAGE', 'tensorflow==2.13.0')
|
||||
# tensorflow_package = os.environ.get('TENSORFLOW_PACKAGE', None)
|
||||
# if tensorflow_package is not None:
|
||||
@@ -1134,6 +1139,28 @@ def check_ui(ver):
|
||||
os.chdir(cwd)
|
||||
|
||||
|
||||
def check_venv():
|
||||
import site
|
||||
pkg_path = [os.path.relpath(p) for p in site.getsitepackages() if os.path.exists(p)]
|
||||
log.debug(f'Packages: venv={os.path.relpath(sys.prefix)} site={pkg_path}')
|
||||
for p in pkg_path:
|
||||
invalid = []
|
||||
for f in os.listdir(p):
|
||||
if f.startswith('~'):
|
||||
invalid.append(f)
|
||||
if len(invalid) > 0:
|
||||
log.warning(f'Packages: site="{p}" invalid={invalid} removing')
|
||||
for f in invalid:
|
||||
fn = os.path.join(p, f)
|
||||
try:
|
||||
if os.path.isdir(fn):
|
||||
shutil.rmtree(fn)
|
||||
elif os.path.isfile(fn):
|
||||
os.unlink(fn)
|
||||
except Exception as e:
|
||||
log.error(f'Packages: site={p} invalid={f} error={e}')
|
||||
|
||||
|
||||
# check version of the main repo and optionally upgrade it
|
||||
def check_version(offline=False, reset=True): # pylint: disable=unused-argument
|
||||
if args.skip_all:
|
||||
@@ -1229,39 +1256,44 @@ def check_timestamp():
|
||||
|
||||
|
||||
def add_args(parser):
|
||||
group = parser.add_argument_group('Setup options')
|
||||
group.add_argument('--reset', default = os.environ.get("SD_RESET",False), action='store_true', help = "Reset main repository to latest version, default: %(default)s")
|
||||
group.add_argument('--upgrade', '--update', default = os.environ.get("SD_UPGRADE",False), action='store_true', help = "Upgrade main repository to latest version, default: %(default)s")
|
||||
group.add_argument('--requirements', default = os.environ.get("SD_REQUIREMENTS",False), action='store_true', help = "Force re-check of requirements, default: %(default)s")
|
||||
group.add_argument('--quick', default = os.environ.get("SD_QUICK",False), action='store_true', help = "Bypass version checks, default: %(default)s")
|
||||
group.add_argument('--use-directml', default = os.environ.get("SD_USEDIRECTML",False), action='store_true', help = "Use DirectML if no compatible GPU is detected, default: %(default)s")
|
||||
group.add_argument("--use-openvino", default = os.environ.get("SD_USEOPENVINO",False), action='store_true', help="Use Intel OpenVINO backend, default: %(default)s")
|
||||
group.add_argument("--use-ipex", default = os.environ.get("SD_USEIPEX",False), action='store_true', help="Force use Intel OneAPI XPU backend, default: %(default)s")
|
||||
group.add_argument("--use-cuda", default = os.environ.get("SD_USECUDA",False), action='store_true', help="Force use nVidia CUDA backend, default: %(default)s")
|
||||
group.add_argument("--use-rocm", default = os.environ.get("SD_USEROCM",False), action='store_true', help="Force use AMD ROCm backend, default: %(default)s")
|
||||
group.add_argument('--use-zluda', default=os.environ.get("SD_USEZLUDA", False), action='store_true', help = "Force use ZLUDA, AMD GPUs only, default: %(default)s")
|
||||
group.add_argument("--use-xformers", default = os.environ.get("SD_USEXFORMERS",False), action='store_true', help="Force use xFormers cross-optimization, default: %(default)s")
|
||||
group.add_argument('--skip-requirements', default = os.environ.get("SD_SKIPREQUIREMENTS",False), action='store_true', help = "Skips checking and installing requirements, default: %(default)s")
|
||||
group.add_argument('--skip-extensions', default = os.environ.get("SD_SKIPEXTENSION",False), action='store_true', help = "Skips running individual extension installers, default: %(default)s")
|
||||
group.add_argument('--skip-git', default = os.environ.get("SD_SKIPGIT",False), action='store_true', help = "Skips running all GIT operations, default: %(default)s")
|
||||
group.add_argument('--skip-torch', default = os.environ.get("SD_SKIPTORCH",False), action='store_true', help = "Skips running Torch checks, default: %(default)s")
|
||||
group.add_argument('--skip-all', default = os.environ.get("SD_SKIPALL",False), action='store_true', help = "Skips running all checks, default: %(default)s")
|
||||
group.add_argument('--skip-env', default = os.environ.get("SD_SKIPENV",False), action='store_true', help = "Skips setting of env variables during startup, default: %(default)s")
|
||||
group.add_argument('--experimental', default = os.environ.get("SD_EXPERIMENTAL",False), action='store_true', help = "Allow unsupported versions of libraries, default: %(default)s")
|
||||
group.add_argument('--reinstall', default = os.environ.get("SD_REINSTALL",False), action='store_true', help = "Force reinstallation of all requirements, default: %(default)s")
|
||||
group.add_argument('--reinstall-zluda', default = os.environ.get("SD_REINSTALL_ZLUDA",False), action='store_true', help = "Force reinstallation of ZLUDA, default: %(default)s")
|
||||
group.add_argument('--test', default = os.environ.get("SD_TEST",False), action='store_true', help = "Run test only and exit")
|
||||
group.add_argument('--version', default = False, action='store_true', help = "Print version information")
|
||||
group.add_argument('--ignore', default = os.environ.get("SD_IGNORE",False), action='store_true', help = "Ignore any errors and attempt to continue")
|
||||
group.add_argument('--safe', default = os.environ.get("SD_SAFE",False), action='store_true', help = "Run in safe mode with no user extensions")
|
||||
group.add_argument('--uv', default = os.environ.get("SD_UV",False), action='store_true', help = "Use uv instead of pip to install the packages")
|
||||
group_setup = parser.add_argument_group('Setup')
|
||||
group_setup.add_argument('--reset', default = os.environ.get("SD_RESET",False), action='store_true', help = "Reset main repository to latest version, default: %(default)s")
|
||||
group_setup.add_argument('--upgrade', '--update', default = os.environ.get("SD_UPGRADE",False), action='store_true', help = "Upgrade main repository to latest version, default: %(default)s")
|
||||
group_setup.add_argument('--requirements', default = os.environ.get("SD_REQUIREMENTS",False), action='store_true', help = "Force re-check of requirements, default: %(default)s")
|
||||
group_setup.add_argument('--reinstall', default = os.environ.get("SD_REINSTALL",False), action='store_true', help = "Force reinstallation of all requirements, default: %(default)s")
|
||||
group_setup.add_argument('--uv', default = os.environ.get("SD_UV",False), action='store_true', help = "Use uv instead of pip to install the packages")
|
||||
|
||||
group = parser.add_argument_group('Logging options')
|
||||
group.add_argument("--log", type=str, default=os.environ.get("SD_LOG", None), help="Set log file, default: %(default)s")
|
||||
group.add_argument('--debug', default = os.environ.get("SD_DEBUG",False), action='store_true', help = "Run installer with debug logging, default: %(default)s")
|
||||
group.add_argument("--profile", default=os.environ.get("SD_PROFILE", False), action='store_true', help="Run profiler, default: %(default)s")
|
||||
group.add_argument('--docs', default=os.environ.get("SD_DOCS", False), action='store_true', help = "Mount API docs, default: %(default)s")
|
||||
group.add_argument("--api-log", default=os.environ.get("SD_APILOG", False), action='store_true', help="Enable logging of all API requests, default: %(default)s")
|
||||
group_startup = parser.add_argument_group('Startup')
|
||||
group_startup.add_argument('--quick', default = os.environ.get("SD_QUICK",False), action='store_true', help = "Bypass version checks, default: %(default)s")
|
||||
group_startup.add_argument('--skip-requirements', default = os.environ.get("SD_SKIPREQUIREMENTS",False), action='store_true', help = "Skips checking and installing requirements, default: %(default)s")
|
||||
group_startup.add_argument('--skip-extensions', default = os.environ.get("SD_SKIPEXTENSION",False), action='store_true', help = "Skips running individual extension installers, default: %(default)s")
|
||||
group_startup.add_argument('--skip-git', default = os.environ.get("SD_SKIPGIT",False), action='store_true', help = "Skips running all GIT operations, default: %(default)s")
|
||||
group_startup.add_argument('--skip-torch', default = os.environ.get("SD_SKIPTORCH",False), action='store_true', help = "Skips running Torch checks, default: %(default)s")
|
||||
group_startup.add_argument('--skip-all', default = os.environ.get("SD_SKIPALL",False), action='store_true', help = "Skips running all checks, default: %(default)s")
|
||||
group_startup.add_argument('--skip-env', default = os.environ.get("SD_SKIPENV",False), action='store_true', help = "Skips setting of env variables during startup, default: %(default)s")
|
||||
|
||||
group_compute = parser.add_argument_group('Compute Engine')
|
||||
group_compute.add_argument('--use-directml', default = os.environ.get("SD_USEDIRECTML",False), action='store_true', help = "Use DirectML if no compatible GPU is detected, default: %(default)s")
|
||||
group_compute.add_argument("--use-openvino", default = os.environ.get("SD_USEOPENVINO",False), action='store_true', help="Use Intel OpenVINO backend, default: %(default)s")
|
||||
group_compute.add_argument("--use-ipex", default = os.environ.get("SD_USEIPEX",False), action='store_true', help="Force use Intel OneAPI XPU backend, default: %(default)s")
|
||||
group_compute.add_argument("--use-cuda", default = os.environ.get("SD_USECUDA",False), action='store_true', help="Force use nVidia CUDA backend, default: %(default)s")
|
||||
group_compute.add_argument("--use-rocm", default = os.environ.get("SD_USEROCM",False), action='store_true', help="Force use AMD ROCm backend, default: %(default)s")
|
||||
group_compute.add_argument('--use-zluda', default=os.environ.get("SD_USEZLUDA", False), action='store_true', help = "Force use ZLUDA, AMD GPUs only, default: %(default)s")
|
||||
group_compute.add_argument("--use-xformers", default = os.environ.get("SD_USEXFORMERS",False), action='store_true', help="Force use xFormers cross-optimization, default: %(default)s")
|
||||
|
||||
group_diag = parser.add_argument_group('Diagnostics')
|
||||
group_diag.add_argument('--safe', default = os.environ.get("SD_SAFE",False), action='store_true', help = "Run in safe mode with no user extensions")
|
||||
group_diag.add_argument('--experimental', default = os.environ.get("SD_EXPERIMENTAL",False), action='store_true', help = "Allow unsupported versions of libraries, default: %(default)s")
|
||||
group_diag.add_argument('--test', default = os.environ.get("SD_TEST",False), action='store_true', help = "Run test only and exit")
|
||||
group_diag.add_argument('--version', default = False, action='store_true', help = "Print version information")
|
||||
group_diag.add_argument('--ignore', default = os.environ.get("SD_IGNORE",False), action='store_true', help = "Ignore any errors and attempt to continue")
|
||||
|
||||
group_log = parser.add_argument_group('Logging')
|
||||
group_log.add_argument("--log", type=str, default=os.environ.get("SD_LOG", None), help="Set log file, default: %(default)s")
|
||||
group_log.add_argument('--debug', default = os.environ.get("SD_DEBUG",False), action='store_true', help = "Run installer with debug logging, default: %(default)s")
|
||||
group_log.add_argument("--profile", default=os.environ.get("SD_PROFILE", False), action='store_true', help="Run profiler, default: %(default)s")
|
||||
group_log.add_argument('--docs', default=os.environ.get("SD_DOCS", False), action='store_true', help = "Mount API docs, default: %(default)s")
|
||||
group_log.add_argument("--api-log", default=os.environ.get("SD_APILOG", False), action='store_true', help="Enable logging of all API requests, default: %(default)s")
|
||||
|
||||
|
||||
def parse_args(parser):
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
let changelogElements = [];
|
||||
|
||||
const getAllChildren = (el) => {
|
||||
const elements = [];
|
||||
for (let i = 0; i < el.children.length; i++) {
|
||||
elements.push(el.children[i]);
|
||||
if (el.children[i].children.length) elements.push(...getAllChildren(el.children[i]));
|
||||
}
|
||||
return elements;
|
||||
};
|
||||
|
||||
function getText(el) {
|
||||
let text = '';
|
||||
el.childNodes.forEach((node) => {
|
||||
if (node.nodeType === Node.TEXT_NODE) text += node.nodeValue;
|
||||
});
|
||||
return text.trim();
|
||||
}
|
||||
|
||||
let currentElement = -1;
|
||||
|
||||
function changelogNavigate(found) {
|
||||
const result = gradioApp().getElementById('changelog_result');
|
||||
result.innerHTML = '';
|
||||
const text = document.createElement('p');
|
||||
|
||||
const onPrev = () => {
|
||||
if (currentElement > 0) {
|
||||
currentElement--;
|
||||
found[currentElement].scrollIntoView();
|
||||
text.innerHTML = `   search item ${currentElement + 1} of ${found.length}`;
|
||||
}
|
||||
};
|
||||
const onNext = () => {
|
||||
if (currentElement < found.length - 1) {
|
||||
currentElement++;
|
||||
found[currentElement].scrollIntoView();
|
||||
text.innerHTML = `   search item ${currentElement + 1} of ${found.length}`;
|
||||
}
|
||||
};
|
||||
|
||||
const prev = document.createElement('p');
|
||||
prev.innerHTML = ' ⇦ ';
|
||||
prev.className = 'changelog_arrow';
|
||||
prev.onclick = onPrev;
|
||||
prev.title = 'Search previous';
|
||||
result.appendChild(prev);
|
||||
|
||||
const next = document.createElement('p');
|
||||
next.innerHTML = ' ⇨ ';
|
||||
next.className = 'changelog_arrow';
|
||||
next.title = 'Search next';
|
||||
next.onclick = onNext;
|
||||
result.appendChild(next);
|
||||
|
||||
text.innerHTML = `   found ${found.length} items`;
|
||||
result.appendChild(text);
|
||||
}
|
||||
|
||||
async function initChangelog() {
|
||||
const search = gradioApp().querySelector('#changelog_search > label> textarea');
|
||||
const md = gradioApp().getElementById('changelog_markdown');
|
||||
const searchChangelog = async (e) => {
|
||||
if (changelogElements.length < 100) changelogElements = getAllChildren(md);
|
||||
const found = [];
|
||||
for (const el of changelogElements) {
|
||||
if (search.value.length > 1 && getText(el).toLowerCase().includes(search.value.toLowerCase())) {
|
||||
el.classList.add('changelog_highlight');
|
||||
found.push(el);
|
||||
} else {
|
||||
el.classList.remove('changelog_highlight');
|
||||
}
|
||||
}
|
||||
changelogNavigate(found);
|
||||
};
|
||||
search.addEventListener('keyup', searchChangelog);
|
||||
}
|
||||
@@ -375,6 +375,7 @@ function setupExtraNetworksForTab(tabname) {
|
||||
if (!tabs) return;
|
||||
|
||||
// buttons
|
||||
const btnShow = gradioApp().getElementById(`${tabname}_extra_networks_btn`);
|
||||
const btnRefresh = gradioApp().getElementById(`${tabname}_extra_refresh`);
|
||||
const btnScan = gradioApp().getElementById(`${tabname}_extra_scan`);
|
||||
const btnSave = gradioApp().getElementById(`${tabname}_extra_save`);
|
||||
@@ -448,50 +449,69 @@ function setupExtraNetworksForTab(tabname) {
|
||||
|
||||
// en style
|
||||
if (!en) return;
|
||||
let lastView;
|
||||
let heightInitialized = false;
|
||||
const intersectionObserver = new IntersectionObserver((entries) => {
|
||||
for (const el of Array.from(gradioApp().querySelectorAll('.extra-networks-page'))) {
|
||||
el.style.height = `${window.opts.extra_networks_height}vh`;
|
||||
el.parentElement.style.width = '-webkit-fill-available';
|
||||
if (!heightInitialized) {
|
||||
heightInitialized = true;
|
||||
let h = 0;
|
||||
const target = window.opts.extra_networks_card_cover === 'sidebar' ? 0 : window.opts.extra_networks_height;
|
||||
if (window.opts.theme_type === 'Standard') h = target > 0 ? target : 55;
|
||||
else h = target > 0 ? target : 87;
|
||||
for (const el of Array.from(gradioApp().getElementById(`${tabname}_extra_tabs`).querySelectorAll('.extra-networks-page'))) {
|
||||
if (h > 0) el.style.height = `${h}vh`;
|
||||
el.parentElement.style.width = '-webkit-fill-available';
|
||||
}
|
||||
}
|
||||
if (entries[0].intersectionRatio > 0) {
|
||||
refreshENpage();
|
||||
// sortExtraNetworks('fixed');
|
||||
if (window.opts.extra_networks_card_cover === 'cover') {
|
||||
en.style.transition = '';
|
||||
en.style.zIndex = 100;
|
||||
en.style.top = '13em';
|
||||
en.style.position = 'absolute';
|
||||
en.style.right = 'unset';
|
||||
en.style.width = 'unset';
|
||||
en.style.height = 'unset';
|
||||
gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = 'unset';
|
||||
} else if (window.opts.extra_networks_card_cover === 'sidebar') {
|
||||
en.style.zIndex = 100;
|
||||
en.style.position = 'absolute';
|
||||
en.style.right = '0';
|
||||
en.style.top = '13em';
|
||||
en.style.height = 'auto';
|
||||
en.style.transition = 'width 0.3s ease';
|
||||
en.style.width = `${window.opts.extra_networks_sidebar_width}vw`;
|
||||
gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = `${100 - 2 - window.opts.extra_networks_sidebar_width}vw`;
|
||||
if (lastView !== entries[0].intersectionRatio > 0) {
|
||||
lastView = entries[0].intersectionRatio > 0;
|
||||
if (lastView) {
|
||||
refreshENpage();
|
||||
// sortExtraNetworks('fixed');
|
||||
if (window.opts.extra_networks_card_cover === 'cover') {
|
||||
en.style.position = 'absolute';
|
||||
en.style.height = 'unset';
|
||||
en.style.width = 'unset';
|
||||
en.style.right = 'unset';
|
||||
en.style.top = '13em';
|
||||
en.style.transition = '';
|
||||
en.style.zIndex = 100;
|
||||
gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = 'unset';
|
||||
} else if (window.opts.extra_networks_card_cover === 'sidebar') {
|
||||
en.style.position = 'absolute';
|
||||
en.style.height = 'auto';
|
||||
en.style.width = `${window.opts.extra_networks_sidebar_width}vw`;
|
||||
en.style.right = '0';
|
||||
en.style.top = '13em';
|
||||
en.style.transition = 'width 0.3s ease';
|
||||
en.style.zIndex = 100;
|
||||
gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = `${100 - 2 - window.opts.extra_networks_sidebar_width}vw`;
|
||||
} else {
|
||||
en.style.position = 'relative';
|
||||
en.style.height = 'unset';
|
||||
en.style.width = 'unset';
|
||||
en.style.right = 'unset';
|
||||
en.style.top = 0;
|
||||
en.style.transition = '';
|
||||
en.style.zIndex = 0;
|
||||
gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = 'unset';
|
||||
}
|
||||
} else {
|
||||
en.style.transition = '';
|
||||
en.style.zIndex = 0;
|
||||
en.style.top = 0;
|
||||
en.style.position = 'relative';
|
||||
en.style.right = 'unset';
|
||||
en.style.width = 'unset';
|
||||
en.style.height = 'unset';
|
||||
if (window.opts.extra_networks_card_cover === 'sidebar') en.style.width = 0;
|
||||
gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = 'unset';
|
||||
}
|
||||
} else {
|
||||
if (window.opts.extra_networks_card_cover === 'sidebar') en.style.width = 0;
|
||||
gradioApp().getElementById(`${tabname}_settings`).parentNode.style.width = 'unset';
|
||||
}
|
||||
});
|
||||
intersectionObserver.observe(en); // monitor visibility
|
||||
}
|
||||
|
||||
async function showNetworks() {
|
||||
for (const tabname of ['txt2img', 'img2img', 'control']) {
|
||||
if (window.opts.extra_networks_show) gradioApp().getElementById(`${tabname}_extra_networks_btn`).click();
|
||||
}
|
||||
log('showNetworks');
|
||||
}
|
||||
|
||||
async function setupExtraNetworks() {
|
||||
setupExtraNetworksForTab('txt2img');
|
||||
setupExtraNetworksForTab('img2img');
|
||||
|
||||
@@ -3,7 +3,7 @@ const appStartTime = performance.now();
|
||||
async function preloadImages() {
|
||||
const dark = window.matchMedia && window.matchMedia('(prefers-color-scheme: dark)').matches;
|
||||
const imagePromises = [];
|
||||
const num = Math.floor(10 * Math.random());
|
||||
const num = Math.floor(9.99 * Math.random());
|
||||
const imageUrls = [
|
||||
`file=html/logo-bg-${dark ? 'dark' : 'light'}.jpg`,
|
||||
`file=html/logo-bg-${num}.jpg`,
|
||||
@@ -27,7 +27,7 @@ async function preloadImages() {
|
||||
async function createSplash() {
|
||||
const dark = window.matchMedia && window.matchMedia('(prefers-color-scheme: dark)').matches;
|
||||
log('createSplash', { theme: dark ? 'dark' : 'light' });
|
||||
const num = Math.floor(11 * Math.random());
|
||||
const num = Math.floor(9.99 * Math.random());
|
||||
const splash = `
|
||||
<div id="splash" class="splash" style="background: ${dark ? 'black' : 'white'}">
|
||||
<div class="loading"><div class="loader"></div></div>
|
||||
|
||||
@@ -133,19 +133,23 @@ document.addEventListener('DOMContentLoaded', () => {
|
||||
});
|
||||
|
||||
/**
|
||||
* Add a ctrl+enter as a shortcut to start a generation
|
||||
* Add a listener to the document for keydown events
|
||||
*/
|
||||
document.addEventListener('keydown', (e) => {
|
||||
let handled = false;
|
||||
if (e.key !== undefined) {
|
||||
if ((e.key === 'Enter' && (e.metaKey || e.ctrlKey || e.altKey))) handled = true;
|
||||
} else if (e.keyCode !== undefined) {
|
||||
if ((e.keyCode === 13 && (e.metaKey || e.ctrlKey || e.altKey))) handled = true;
|
||||
}
|
||||
if (handled) {
|
||||
const button = getUICurrentTabContent().querySelector('button[id$=_generate]');
|
||||
if (button) button.click();
|
||||
let elem;
|
||||
if (e.key === 'Escape') elem = getUICurrentTabContent().querySelector('button[id$=_interrupt]');
|
||||
if (e.key === 'Enter' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id$=_generate]');
|
||||
if (e.key === 'Backspace' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id$=_reprocess]');
|
||||
if (e.key === ' ' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id$=_extra_networks_btn]');
|
||||
if (e.key === 's' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id^=save_]');
|
||||
if (e.key === 'Insert' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id^=save_]');
|
||||
if (e.key === 'Delete' && e.ctrlKey) elem = getUICurrentTabContent().querySelector('button[id^=delete_]');
|
||||
// if (e.key === 'm' && e.ctrlKey) elem = gradioApp().getElementById('setting_sd_model_checkpoint');
|
||||
if (elem) {
|
||||
e.preventDefault();
|
||||
log('hotkey', { key: e.key, meta: e.metaKey, ctrl: e.ctrlKey, alt: e.altKey }, elem?.id, elem.nodeName);
|
||||
if (elem.nodeName === 'BUTTON') elem.click();
|
||||
else elem.focus();
|
||||
}
|
||||
});
|
||||
|
||||
|
||||
@@ -38,7 +38,7 @@ td > div > span { overflow-y: auto; max-height: 3em; overflow-x: hidden; }
|
||||
.gradio-button.secondary-down, .gradio-button.secondary-down:hover { box-shadow: 1px 1px 1px rgba(0,0,0,0.25) inset, 0px 0px 3px rgba(0,0,0,0.15) inset; }
|
||||
.gradio-button.secondary-down:hover { background: var(--button-secondary-background-fill-hover); color: var(--button-secondary-text-color-hover); }
|
||||
.gradio-button.tool { max-width: min-content; min-width: min-content !important; font-size: 20px !important; color: var(--body-text-color) !important; align-self: end; margin-bottom: 4px; }
|
||||
.gradio-checkbox { margin: 0.75em 1.5em 0 0; align-self: center; }
|
||||
.gradio-checkbox { margin-right: 1em !important; align-self: center; }
|
||||
.gradio-column { min-width: min(160px, 100%) !important; }
|
||||
.gradio-container { max-width: unset !important; padding: var(--block-label-padding) !important; }
|
||||
.gradio-container .prose a, .gradio-container .prose a:visited{ color: unset; text-decoration: none; }
|
||||
@@ -203,7 +203,7 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt
|
||||
#extensions .date { opacity: 0.85; font-size: var(--text-sm); }
|
||||
|
||||
/* extra networks */
|
||||
#txt2img_extra_networks, #img2img_extra_networks, #control_extra_networks { width: 0; }
|
||||
.extra_networks_root { width: 0; position: absolute; height: auto; right: 0; top: 13em; z-index: 100; } /* default is sidebar view */
|
||||
.extra-networks { background: var(--background-color); padding: var(--block-label-padding); }
|
||||
.extra-networks > div { margin: 0; border-bottom: none !important; gap: 0.3em 0; }
|
||||
.extra-networks .second-line { display: flex; width: -moz-available; width: -webkit-fill-available; gap: 0.3em; box-shadow: var(--input-shadow); }
|
||||
@@ -319,6 +319,13 @@ div:has(>#tab-gallery-folders) { flex-grow: 0 !important; background-color: var(
|
||||
.gallery-sort { background: var(--input-background-fill) !important; margin: 0 !important; padding: 6px !important; }
|
||||
.gallery-sort:hover { background: var(--button-primary-background-fill-hover) !important; }
|
||||
|
||||
/* changelog */
|
||||
#changelog_markdown { max-height: 55vh; margin-top: 1em; }
|
||||
#changelog_result { display: flex; margin-left: 1em; align-items: center; }
|
||||
.changelog_arrow { font-size: 2em; padding: 0.1em; cursor: pointer; height: 1em; background-color: var(--button-secondary-background-fill); }
|
||||
.changelog_arrow:hover { background-color: var(--button-primary-border-color-hover); }
|
||||
.changelog_highlight { background-color: var(--color-warning); }
|
||||
|
||||
/* loader */
|
||||
.splash { position: fixed; top: 0; left: 0; width: 100vw; height: 100vh; z-index: 1000; display: block; text-align: center; }
|
||||
.motd { margin-top: 2em; color: var(--body-text-color-subdued); font-family: monospace; font-variant: all-petite-caps; }
|
||||
|
||||
@@ -16,6 +16,7 @@ async function initStartup() {
|
||||
initImageViewer();
|
||||
initGallery();
|
||||
initiGenerationParams();
|
||||
initChangelog();
|
||||
setupControlUI();
|
||||
|
||||
// reconnect server session
|
||||
@@ -31,6 +32,7 @@ async function initStartup() {
|
||||
removeSplash();
|
||||
|
||||
// post startup tasks that may take longer but are not critical
|
||||
showNetworks();
|
||||
setHints();
|
||||
applyStyles();
|
||||
initIndexDB();
|
||||
|
||||
@@ -164,7 +164,6 @@ def start_server(immediate=True, server=None):
|
||||
module_spec = importlib.util.spec_from_file_location('webui', 'webui.py')
|
||||
server = importlib.util.module_from_spec(module_spec)
|
||||
installer.log.debug(f'Starting module: {server}')
|
||||
get_custom_args()
|
||||
module_spec.loader.exec_module(server)
|
||||
uvicorn = None
|
||||
if args.test:
|
||||
@@ -209,6 +208,8 @@ def main():
|
||||
installer.log.info('Skipping GIT operations')
|
||||
installer.check_version()
|
||||
installer.log.info(f'Platform: {installer.print_dict(installer.get_platform())}')
|
||||
installer.check_venv()
|
||||
installer.log.info(f'Args: {sys.argv[1:]}')
|
||||
if not args.skip_env:
|
||||
installer.set_environment()
|
||||
if args.uv:
|
||||
@@ -246,6 +247,7 @@ def main():
|
||||
installer.log.warning(f'See log file for more details: {installer.log_file}')
|
||||
installer.extensions_preload(parser) # adds additional args from extensions
|
||||
args = installer.parse_args(parser)
|
||||
get_custom_args()
|
||||
|
||||
uv, instance = start_server(immediate=True, server=None)
|
||||
while True:
|
||||
|
||||
@@ -113,10 +113,9 @@ class APIProcess():
|
||||
images = []
|
||||
scores = []
|
||||
with self.queue_lock:
|
||||
yolo.load()
|
||||
faces = yolo.predict(image)
|
||||
faces = yolo.predict('face-yolo8n', image)
|
||||
for face in faces:
|
||||
images.append(encode_pil_to_base64(face.face))
|
||||
images.append(encode_pil_to_base64(face.item))
|
||||
scores.append(face.score)
|
||||
shared.state.end(api=False)
|
||||
return ResFace(images=images, scores=scores)
|
||||
|
||||
@@ -2,104 +2,134 @@ import os
|
||||
import argparse
|
||||
from modules.paths import data_path
|
||||
|
||||
|
||||
parser = argparse.ArgumentParser(description="SD.Next", conflict_handler='resolve', epilog='For other options see UI Settings page', prog='', add_help=True, formatter_class=lambda prog: argparse.HelpFormatter(prog, max_help_position=55, indent_increment=2, width=200))
|
||||
parser._optionals = parser.add_argument_group('Other options') # pylint: disable=protected-access
|
||||
group = parser.add_argument_group('Server options')
|
||||
|
||||
# main server args
|
||||
group.add_argument("--config", type=str, default=os.environ.get("SD_CONFIG", os.path.join(data_path, 'config.json')), help="Use specific server configuration file, default: %(default)s")
|
||||
group.add_argument("--ui-config", type=str, default=os.environ.get("SD_UICONFIG", os.path.join(data_path, 'ui-config.json')), help="Use specific UI configuration file, default: %(default)s")
|
||||
group.add_argument("--medvram", default=os.environ.get("SD_MEDVRAM", False), action='store_true', help="Split model stages and keep only active part in VRAM, default: %(default)s")
|
||||
group.add_argument("--lowvram", default=os.environ.get("SD_LOWVRAM", False), action='store_true', help="Split model components and keep only active part in VRAM, default: %(default)s")
|
||||
group.add_argument("--ckpt", type=str, default=os.environ.get("SD_MODEL", None), help="Path to model checkpoint to load immediately, default: %(default)s")
|
||||
group.add_argument('--vae', type=str, default=os.environ.get("SD_VAE", None), help='Path to VAE checkpoint to load immediately, default: %(default)s')
|
||||
group.add_argument("--data-dir", type=str, default=os.environ.get("SD_DATADIR", ''), help="Base path where all user data is stored, default: %(default)s")
|
||||
group.add_argument("--models-dir", type=str, default=os.environ.get("SD_MODELSDIR", 'models'), help="Base path where all models are stored, default: %(default)s",)
|
||||
group.add_argument("--allow-code", default=os.environ.get("SD_ALLOWCODE", False), action='store_true', help="Allow custom script execution, default: %(default)s")
|
||||
group.add_argument("--share", default=os.environ.get("SD_SHARE", False), action='store_true', help="Enable UI accessible through Gradio site, default: %(default)s")
|
||||
group.add_argument("--insecure", default=os.environ.get("SD_INSECURE", False), action='store_true', help="Enable extensions tab regardless of other options, default: %(default)s")
|
||||
group.add_argument("--use-cpu", nargs='+', default=[], type=str.lower, help="Force use CPU for specified modules, default: %(default)s")
|
||||
group.add_argument("--listen", default=os.environ.get("SD_LISTEN", False), action='store_true', help="Launch web server using public IP address, default: %(default)s")
|
||||
group.add_argument("--port", type=int, default=os.environ.get("SD_PORT", 7860), help="Launch web server with given server port, default: %(default)s")
|
||||
group.add_argument("--freeze", default=os.environ.get("SD_FREEZE", False), action='store_true', help="Disable editing settings")
|
||||
group.add_argument("--auth", type=str, default=os.environ.get("SD_AUTH", None), help='Set access authentication like "user:pwd,user:pwd""')
|
||||
group.add_argument("--auth-file", type=str, default=os.environ.get("SD_AUTHFILE", None), help='Set access authentication using file, default: %(default)s')
|
||||
group.add_argument("--autolaunch", default=os.environ.get("SD_AUTOLAUNCH", False), action='store_true', help="Open the UI URL in the system's default browser upon launch")
|
||||
group.add_argument('--docs', default=os.environ.get("SD_DOCS", False), action='store_true', help = "Mount API docs, default: %(default)s")
|
||||
group.add_argument('--api-only', default=os.environ.get("SD_APIONLY", False), action='store_true', help = "Run in API only mode without starting UI")
|
||||
group.add_argument("--device-id", type=str, default=os.environ.get("SD_DEVICEID", None), help="Select the default CUDA device to use, default: %(default)s")
|
||||
group.add_argument("--cors-origins", type=str, default=os.environ.get("SD_CORSORIGINS", None), help="Allowed CORS origins as comma-separated list, default: %(default)s")
|
||||
group.add_argument("--cors-regex", type=str, default=os.environ.get("SD_CORSREGEX", None), help="Allowed CORS origins as regular expression, default: %(default)s")
|
||||
group.add_argument("--tls-keyfile", type=str, default=os.environ.get("SD_TLSKEYFILE", None), help="Enable TLS and specify key file, default: %(default)s")
|
||||
group.add_argument("--tls-certfile", type=str, default=os.environ.get("SD_TLSCERTFILE", None), help="Enable TLS and specify cert file, default: %(default)s")
|
||||
group.add_argument("--tls-selfsign", action="store_true", default=os.environ.get("SD_TLSSELFSIGN", False), help="Enable TLS with self-signed certificates, default: %(default)s")
|
||||
group.add_argument("--server-name", type=str, default=os.environ.get("SD_SERVERNAME", None), help="Sets hostname of server, default: %(default)s")
|
||||
group.add_argument("--no-hashing", default=os.environ.get("SD_NOHASHING", False), action='store_true', help="Disable hashing of checkpoints, default: %(default)s")
|
||||
group.add_argument("--no-metadata", default=os.environ.get("SD_NOMETADATA", False), action='store_true', help="Disable reading of metadata from models, default: %(default)s")
|
||||
group.add_argument("--profile", default=os.environ.get("SD_PROFILE", False), action='store_true', help="Run profiler, default: %(default)s")
|
||||
group.add_argument("--disable-queue", default=os.environ.get("SD_DISABLEQUEUE", False), action='store_true', help="Disable queues, default: %(default)s")
|
||||
group.add_argument('--debug', default=os.environ.get("SD_DEBUG", False), action='store_true', help = "Run installer with debug logging, default: %(default)s")
|
||||
group.add_argument('--use-directml', default=os.environ.get("SD_USEDIRECTML", False), action='store_true', help = "Use DirectML if no compatible GPU is detected, default: %(default)s")
|
||||
group.add_argument('--use-zluda', default=os.environ.get("SD_USEZLUDA", False), action='store_true', help = "Force use ZLUDA, AMD GPUs only, default: %(default)s")
|
||||
group.add_argument("--use-openvino", default=os.environ.get("SD_USEOPENVINO", False), action='store_true', help="Use Intel OpenVINO backend, default: %(default)s")
|
||||
group.add_argument("--use-ipex", default=os.environ.get("SD_USEIPX", False), action='store_true', help="Force use Intel OneAPI XPU backend, default: %(default)s")
|
||||
group.add_argument("--use-cuda", default=os.environ.get("SD_USECUDA", False), action='store_true', help="Force use nVidia CUDA backend, default: %(default)s")
|
||||
group.add_argument("--use-rocm", default=os.environ.get("SD_USEROCM", False), action='store_true', help="Force use AMD ROCm backend, default: %(default)s")
|
||||
group.add_argument('--subpath', type=str, default=os.environ.get("SD_SUBPATH", None), help='Customize the URL subpath for usage with reverse proxy')
|
||||
group.add_argument('--backend', type=str, default=os.environ.get("SD_BACKEND", None), choices=['original', 'diffusers'], required=False, help='force model pipeline type')
|
||||
group.add_argument('--theme', type=str, default=os.environ.get("SD_THEME", None), help='Override UI theme')
|
||||
|
||||
|
||||
# removed args are added here as hidden in fixed format for compatbility reasons
|
||||
group.add_argument("-f", action='store_true', help=argparse.SUPPRESS) # allows running as root; implemented outside of webui
|
||||
group.add_argument("--ui-settings-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'config.json'))
|
||||
group.add_argument("--ui-config-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'ui-config.json'))
|
||||
group.add_argument("--hide-ui-dir-config", action='store_true', help=argparse.SUPPRESS, default=False)
|
||||
group.add_argument("--theme", type=str, help=argparse.SUPPRESS, default=None)
|
||||
group.add_argument("--disable-console-progressbars", action='store_true', help=argparse.SUPPRESS, default=True)
|
||||
group.add_argument("--disable-safe-unpickle", action='store_true', help=argparse.SUPPRESS, default=True)
|
||||
group.add_argument("--lowram", action='store_true', help=argparse.SUPPRESS)
|
||||
group.add_argument("--disable-extension-access", default=False, action='store_true', help=argparse.SUPPRESS)
|
||||
group.add_argument("--allowed-paths", nargs='+', default=[], type=str, required=False, help="add additional paths to paths allowed for web access")
|
||||
group.add_argument("--api", help=argparse.SUPPRESS, default=True)
|
||||
group.add_argument("--api-auth", type=str, help=argparse.SUPPRESS, default=None)
|
||||
def main_args():
|
||||
# main server args
|
||||
group_config = parser.add_argument_group('Configuration')
|
||||
group_config.add_argument('--backend', type=str, default=os.environ.get("SD_BACKEND", None), choices=['original', 'diffusers'], required=False, help='force model pipeline type')
|
||||
group_config.add_argument("--config", type=str, default=os.environ.get("SD_CONFIG", os.path.join(data_path, 'config.json')), help="Use specific server configuration file, default: %(default)s")
|
||||
group_config.add_argument("--ui-config", type=str, default=os.environ.get("SD_UICONFIG", os.path.join(data_path, 'ui-config.json')), help="Use specific UI configuration file, default: %(default)s")
|
||||
group_config.add_argument("--medvram", default=os.environ.get("SD_MEDVRAM", False), action='store_true', help="Split model stages and keep only active part in VRAM, default: %(default)s")
|
||||
group_config.add_argument("--lowvram", default=os.environ.get("SD_LOWVRAM", False), action='store_true', help="Split model components and keep only active part in VRAM, default: %(default)s")
|
||||
group_config.add_argument("--freeze", default=os.environ.get("SD_FREEZE", False), action='store_true', help="Disable editing settings")
|
||||
|
||||
group_paths = parser.add_argument_group('Paths')
|
||||
group_paths.add_argument("--ckpt", type=str, default=os.environ.get("SD_MODEL", None), help="Path to model checkpoint to load immediately, default: %(default)s")
|
||||
group_paths.add_argument("--data-dir", type=str, default=os.environ.get("SD_DATADIR", ''), help="Base path where all user data is stored, default: %(default)s")
|
||||
group_paths.add_argument("--models-dir", type=str, default=os.environ.get("SD_MODELSDIR", 'models'), help="Base path where all models are stored, default: %(default)s",)
|
||||
|
||||
group_diag = parser.add_argument_group('Diagnostics')
|
||||
group_diag.add_argument("--no-hashing", default=os.environ.get("SD_NOHASHING", False), action='store_true', help="Disable hashing of checkpoints, default: %(default)s")
|
||||
group_diag.add_argument("--no-metadata", default=os.environ.get("SD_NOMETADATA", False), action='store_true', help="Disable reading of metadata from models, default: %(default)s")
|
||||
group_diag.add_argument("--profile", default=os.environ.get("SD_PROFILE", False), action='store_true', help="Run profiler, default: %(default)s")
|
||||
group_diag.add_argument("--disable-queue", default=os.environ.get("SD_DISABLEQUEUE", False), action='store_true', help="Disable queues, default: %(default)s")
|
||||
group_diag.add_argument('--debug', default=os.environ.get("SD_DEBUG", False), action='store_true', help = "Run installer with debug logging, default: %(default)s")
|
||||
|
||||
group_compute = parser.add_argument_group('Compute Engine')
|
||||
group_compute.add_argument('--use-directml', default=os.environ.get("SD_USEDIRECTML", False), action='store_true', help = "Use DirectML if no compatible GPU is detected, default: %(default)s")
|
||||
group_compute.add_argument('--use-zluda', default=os.environ.get("SD_USEZLUDA", False), action='store_true', help = "Force use ZLUDA, AMD GPUs only, default: %(default)s")
|
||||
group_compute.add_argument("--use-openvino", default=os.environ.get("SD_USEOPENVINO", False), action='store_true', help="Use Intel OpenVINO backend, default: %(default)s")
|
||||
group_compute.add_argument("--use-ipex", default=os.environ.get("SD_USEIPX", False), action='store_true', help="Force use Intel OneAPI XPU backend, default: %(default)s")
|
||||
group_compute.add_argument("--use-cuda", default=os.environ.get("SD_USECUDA", False), action='store_true', help="Force use nVidia CUDA backend, default: %(default)s")
|
||||
group_compute.add_argument("--use-rocm", default=os.environ.get("SD_USEROCM", False), action='store_true', help="Force use AMD ROCm backend, default: %(default)s")
|
||||
group_diag.add_argument("--device-id", type=str, default=os.environ.get("SD_DEVICEID", None), help="Select the default CUDA device to use, default: %(default)s")
|
||||
|
||||
group_http = parser.add_argument_group('HTTP')
|
||||
group_http.add_argument('--theme', type=str, default=os.environ.get("SD_THEME", None), help='Override UI theme')
|
||||
group_http.add_argument("--server-name", type=str, default=os.environ.get("SD_SERVERNAME", None), help="Sets hostname of server, default: %(default)s")
|
||||
group_http.add_argument("--tls-keyfile", type=str, default=os.environ.get("SD_TLSKEYFILE", None), help="Enable TLS and specify key file, default: %(default)s")
|
||||
group_http.add_argument("--tls-certfile", type=str, default=os.environ.get("SD_TLSCERTFILE", None), help="Enable TLS and specify cert file, default: %(default)s")
|
||||
group_http.add_argument("--tls-selfsign", action="store_true", default=os.environ.get("SD_TLSSELFSIGN", False), help="Enable TLS with self-signed certificates, default: %(default)s")
|
||||
group_http.add_argument("--cors-origins", type=str, default=os.environ.get("SD_CORSORIGINS", None), help="Allowed CORS origins as comma-separated list, default: %(default)s")
|
||||
group_http.add_argument("--cors-regex", type=str, default=os.environ.get("SD_CORSREGEX", None), help="Allowed CORS origins as regular expression, default: %(default)s")
|
||||
group_http.add_argument('--subpath', type=str, default=os.environ.get("SD_SUBPATH", None), help='Customize the URL subpath for usage with reverse proxy')
|
||||
group_http.add_argument("--autolaunch", default=os.environ.get("SD_AUTOLAUNCH", False), action='store_true', help="Open the UI URL in the system's default browser upon launch")
|
||||
group_http.add_argument('--docs', default=os.environ.get("SD_DOCS", False), action='store_true', help = "Mount API docs, default: %(default)s")
|
||||
group_http.add_argument("--auth", type=str, default=os.environ.get("SD_AUTH", None), help='Set access authentication like "user:pwd,user:pwd""')
|
||||
group_http.add_argument("--auth-file", type=str, default=os.environ.get("SD_AUTHFILE", None), help='Set access authentication using file, default: %(default)s')
|
||||
group_http.add_argument('--api-only', default=os.environ.get("SD_APIONLY", False), action='store_true', help = "Run in API only mode without starting UI")
|
||||
group_http.add_argument("--allowed-paths", nargs='+', default=[], type=str, required=False, help="add additional paths to paths allowed for web access")
|
||||
group_http.add_argument("--share", default=os.environ.get("SD_SHARE", False), action='store_true', help="Enable UI accessible through Gradio site, default: %(default)s")
|
||||
group_http.add_argument("--insecure", default=os.environ.get("SD_INSECURE", False), action='store_true', help="Enable extensions tab regardless of other options, default: %(default)s")
|
||||
group_http.add_argument("--listen", default=os.environ.get("SD_LISTEN", False), action='store_true', help="Launch web server using public IP address, default: %(default)s")
|
||||
group_http.add_argument("--port", type=int, default=os.environ.get("SD_PORT", 7860), help="Launch web server with given server port, default: %(default)s")
|
||||
|
||||
|
||||
def compatibility_args(opts, args):
|
||||
def compatibility_args():
|
||||
group_compat = parser.add_argument_group('Compatibility options')
|
||||
# removed args are added here as hidden in fixed format for compatbility reasons
|
||||
group_compat.add_argument("--allow-code", default=os.environ.get("SD_ALLOWCODE", False), action='store_true', help=argparse.SUPPRESS)
|
||||
group_compat.add_argument("--use-cpu", nargs='+', default=[], type=str.lower, help=argparse.SUPPRESS)
|
||||
group_compat.add_argument("-f", action='store_true', help=argparse.SUPPRESS) # allows running as root; implemented outside of webui
|
||||
group_compat.add_argument('--vae', type=str, default=os.environ.get("SD_VAE", None), help=argparse.SUPPRESS)
|
||||
group_compat.add_argument("--ui-settings-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'config.json'))
|
||||
group_compat.add_argument("--ui-config-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'ui-config.json'))
|
||||
group_compat.add_argument("--hide-ui-dir-config", action='store_true', help=argparse.SUPPRESS, default=False)
|
||||
group_compat.add_argument("--theme", type=str, help=argparse.SUPPRESS, default=None)
|
||||
group_compat.add_argument("--disable-console-progressbars", action='store_true', help=argparse.SUPPRESS, default=True)
|
||||
group_compat.add_argument("--disable-safe-unpickle", action='store_true', help=argparse.SUPPRESS, default=True)
|
||||
group_compat.add_argument("--lowram", action='store_true', help=argparse.SUPPRESS)
|
||||
group_compat.add_argument("--disable-extension-access", default=False, action='store_true', help=argparse.SUPPRESS)
|
||||
group_compat.add_argument("--api", help=argparse.SUPPRESS, default=True)
|
||||
group_compat.add_argument("--api-auth", type=str, help=argparse.SUPPRESS, default=None)
|
||||
|
||||
|
||||
def settings_args(opts, args):
|
||||
group_compat = parser.add_argument_group('Compatibility options')
|
||||
# removed args are added here as hidden in fixed format for compatbility reasons
|
||||
group_compat.add_argument("--allow-code", default=os.environ.get("SD_ALLOWCODE", False), action='store_true', help=argparse.SUPPRESS)
|
||||
group_compat.add_argument("--use-cpu", nargs='+', default=[], type=str.lower, help=argparse.SUPPRESS)
|
||||
group_compat.add_argument("-f", action='store_true', help=argparse.SUPPRESS) # allows running as root; implemented outside of webui
|
||||
group_compat.add_argument('--vae', type=str, default=os.environ.get("SD_VAE", None), help=argparse.SUPPRESS)
|
||||
group_compat.add_argument("--ui-settings-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'config.json'))
|
||||
group_compat.add_argument("--ui-config-file", type=str, help=argparse.SUPPRESS, default=os.path.join(data_path, 'ui-config.json'))
|
||||
group_compat.add_argument("--hide-ui-dir-config", action='store_true', help=argparse.SUPPRESS, default=False)
|
||||
group_compat.add_argument("--theme", type=str, help=argparse.SUPPRESS, default=None)
|
||||
group_compat.add_argument("--disable-console-progressbars", action='store_true', help=argparse.SUPPRESS, default=True)
|
||||
group_compat.add_argument("--disable-safe-unpickle", action='store_true', help=argparse.SUPPRESS, default=True)
|
||||
group_compat.add_argument("--lowram", action='store_true', help=argparse.SUPPRESS)
|
||||
group_compat.add_argument("--disable-extension-access", default=False, action='store_true', help=argparse.SUPPRESS)
|
||||
group_compat.add_argument("--allowed-paths", nargs='+', default=[], type=str, required=False, help="add additional paths to paths allowed for web access")
|
||||
group_compat.add_argument("--api", help=argparse.SUPPRESS, default=True)
|
||||
group_compat.add_argument("--api-auth", type=str, help=argparse.SUPPRESS, default=None)
|
||||
# removed args that have been moved to opts are added here as hidden with default values as defined in opts
|
||||
group.add_argument("--ckpt-dir", type=str, help=argparse.SUPPRESS, default=opts.ckpt_dir)
|
||||
group.add_argument("--vae-dir", type=str, help=argparse.SUPPRESS, default=opts.vae_dir)
|
||||
group.add_argument("--embeddings-dir", type=str, help=argparse.SUPPRESS, default=opts.embeddings_dir)
|
||||
group.add_argument("--embeddings-templates-dir", type=str, help=argparse.SUPPRESS, default=opts.embeddings_templates_dir)
|
||||
group.add_argument("--hypernetwork-dir", type=str, help=argparse.SUPPRESS, default=opts.hypernetwork_dir)
|
||||
group.add_argument("--codeformer-models-path", type=str, help=argparse.SUPPRESS, default=opts.codeformer_models_path)
|
||||
group.add_argument("--gfpgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.gfpgan_models_path)
|
||||
group.add_argument("--esrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.esrgan_models_path)
|
||||
group.add_argument("--bsrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.bsrgan_models_path)
|
||||
group.add_argument("--realesrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.realesrgan_models_path)
|
||||
group.add_argument("--scunet-models-path", help=argparse.SUPPRESS, default=opts.scunet_models_path)
|
||||
group.add_argument("--swinir-models-path", help=argparse.SUPPRESS, default=opts.swinir_models_path)
|
||||
group.add_argument("--ldsr-models-path", help=argparse.SUPPRESS, default=opts.ldsr_models_path)
|
||||
group.add_argument("--clip-models-path", type=str, help=argparse.SUPPRESS, default=opts.clip_models_path)
|
||||
group.add_argument("--opt-channelslast", help=argparse.SUPPRESS, action='store_true', default=opts.opt_channelslast)
|
||||
group.add_argument("--xformers", default=(opts.cross_attention_optimization == "xFormers"), action='store_true', help=argparse.SUPPRESS)
|
||||
group.add_argument("--disable-nan-check", help=argparse.SUPPRESS, action='store_true', default=opts.disable_nan_check)
|
||||
group.add_argument("--rollback-vae", help=argparse.SUPPRESS, default=opts.rollback_vae)
|
||||
group.add_argument("--no-half", help=argparse.SUPPRESS, action='store_true', default=opts.no_half)
|
||||
group.add_argument("--no-half-vae", help=argparse.SUPPRESS, action='store_true', default=opts.no_half_vae)
|
||||
group.add_argument("--precision", help=argparse.SUPPRESS, default=opts.precision)
|
||||
group.add_argument("--sub-quad-q-chunk-size", help=argparse.SUPPRESS, default=opts.sub_quad_q_chunk_size)
|
||||
group.add_argument("--sub-quad-kv-chunk-size", help=argparse.SUPPRESS, default=opts.sub_quad_kv_chunk_size)
|
||||
group.add_argument("--sub-quad-chunk-threshold", help=argparse.SUPPRESS, default=opts.sub_quad_chunk_threshold)
|
||||
group.add_argument("--lora-dir", help=argparse.SUPPRESS, default=opts.lora_dir)
|
||||
group.add_argument("--lyco-dir", help=argparse.SUPPRESS, default=opts.lyco_dir)
|
||||
group.add_argument("--embeddings-dir", help=argparse.SUPPRESS, default=opts.embeddings_dir)
|
||||
group.add_argument("--hypernetwork-dir", help=argparse.SUPPRESS, default=opts.hypernetwork_dir)
|
||||
group.add_argument("--lyco-patch-lora", help=argparse.SUPPRESS, action='store_true', default=False)
|
||||
group.add_argument("--lyco-debug", help=argparse.SUPPRESS, action='store_true', default=False)
|
||||
group.add_argument("--enable-console-prompts", help=argparse.SUPPRESS, action='store_true', default=False)
|
||||
group.add_argument("--safe", help=argparse.SUPPRESS, action='store_true', default=False)
|
||||
group.add_argument("--use-xformers", help=argparse.SUPPRESS, action='store_true', default=False)
|
||||
group_compat.add_argument("--ckpt-dir", type=str, help=argparse.SUPPRESS, default=opts.ckpt_dir)
|
||||
group_compat.add_argument("--vae-dir", type=str, help=argparse.SUPPRESS, default=opts.vae_dir)
|
||||
group_compat.add_argument("--embeddings-dir", type=str, help=argparse.SUPPRESS, default=opts.embeddings_dir)
|
||||
group_compat.add_argument("--embeddings-templates-dir", type=str, help=argparse.SUPPRESS, default=opts.embeddings_templates_dir)
|
||||
group_compat.add_argument("--hypernetwork-dir", type=str, help=argparse.SUPPRESS, default=opts.hypernetwork_dir)
|
||||
group_compat.add_argument("--codeformer-models-path", type=str, help=argparse.SUPPRESS, default=opts.codeformer_models_path)
|
||||
group_compat.add_argument("--gfpgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.gfpgan_models_path)
|
||||
group_compat.add_argument("--esrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.esrgan_models_path)
|
||||
group_compat.add_argument("--bsrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.bsrgan_models_path)
|
||||
group_compat.add_argument("--realesrgan-models-path", type=str, help=argparse.SUPPRESS, default=opts.realesrgan_models_path)
|
||||
group_compat.add_argument("--scunet-models-path", help=argparse.SUPPRESS, default=opts.scunet_models_path)
|
||||
group_compat.add_argument("--swinir-models-path", help=argparse.SUPPRESS, default=opts.swinir_models_path)
|
||||
group_compat.add_argument("--ldsr-models-path", help=argparse.SUPPRESS, default=opts.ldsr_models_path)
|
||||
group_compat.add_argument("--clip-models-path", type=str, help=argparse.SUPPRESS, default=opts.clip_models_path)
|
||||
group_compat.add_argument("--opt-channelslast", help=argparse.SUPPRESS, action='store_true', default=opts.opt_channelslast)
|
||||
group_compat.add_argument("--xformers", default=(opts.cross_attention_optimization == "xFormers"), action='store_true', help=argparse.SUPPRESS)
|
||||
group_compat.add_argument("--disable-nan-check", help=argparse.SUPPRESS, action='store_true', default=opts.disable_nan_check)
|
||||
group_compat.add_argument("--rollback-vae", help=argparse.SUPPRESS, default=opts.rollback_vae)
|
||||
group_compat.add_argument("--no-half", help=argparse.SUPPRESS, action='store_true', default=opts.no_half)
|
||||
group_compat.add_argument("--no-half-vae", help=argparse.SUPPRESS, action='store_true', default=opts.no_half_vae)
|
||||
group_compat.add_argument("--precision", help=argparse.SUPPRESS, default=opts.precision)
|
||||
group_compat.add_argument("--sub-quad-q-chunk-size", help=argparse.SUPPRESS, default=opts.sub_quad_q_chunk_size)
|
||||
group_compat.add_argument("--sub-quad-kv-chunk-size", help=argparse.SUPPRESS, default=opts.sub_quad_kv_chunk_size)
|
||||
group_compat.add_argument("--sub-quad-chunk-threshold", help=argparse.SUPPRESS, default=opts.sub_quad_chunk_threshold)
|
||||
group_compat.add_argument("--lora-dir", help=argparse.SUPPRESS, default=opts.lora_dir)
|
||||
group_compat.add_argument("--lyco-dir", help=argparse.SUPPRESS, default=opts.lyco_dir)
|
||||
group_compat.add_argument("--embeddings-dir", help=argparse.SUPPRESS, default=opts.embeddings_dir)
|
||||
group_compat.add_argument("--hypernetwork-dir", help=argparse.SUPPRESS, default=opts.hypernetwork_dir)
|
||||
group_compat.add_argument("--lyco-patch-lora", help=argparse.SUPPRESS, action='store_true', default=False)
|
||||
group_compat.add_argument("--lyco-debug", help=argparse.SUPPRESS, action='store_true', default=False)
|
||||
group_compat.add_argument("--enable-console-prompts", help=argparse.SUPPRESS, action='store_true', default=False)
|
||||
group_compat.add_argument("--safe", help=argparse.SUPPRESS, action='store_true', default=False)
|
||||
group_compat.add_argument("--use-xformers", help=argparse.SUPPRESS, action='store_true', default=False)
|
||||
|
||||
# removed opts are added here with fixed values for compatibility reasons
|
||||
opts.use_old_emphasis_implementation = False
|
||||
@@ -126,3 +156,7 @@ def compatibility_args(opts, args):
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
|
||||
|
||||
main_args()
|
||||
compatibility_args()
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
"""
|
||||
original code from <https://github.com/NVlabs/consistory>
|
||||
"""
|
||||
from .consistory_pipeline import ConsistoryExtendAttnSDXLPipeline
|
||||
from .consistory_unet_sdxl import ConsistorySDXLUNet2DConditionModel
|
||||
from .consistory_run import run_anchor_generation, run_extra_generation
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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 = {}
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -564,8 +564,8 @@ def control_run(state: str = '',
|
||||
return [], '', '', 'Reference mode without image'
|
||||
elif unit_type == 'controlnet' and has_models:
|
||||
if input_type == 0: # Control only
|
||||
if shared.sd_model_type == 'f1' and 'control_image' not in p.task_args:
|
||||
p.task_args['control_image'] = p.init_images # flux controlnet mandates this
|
||||
if shared.sd_model_type in ['f1', 'sd3'] and 'control_image' not in p.task_args:
|
||||
p.task_args['control_image'] = p.init_images # some controlnets mandate this
|
||||
p.task_args['strength'] = p.denoising_strength
|
||||
elif input_type == 1: # Init image same as control
|
||||
p.task_args['control_image'] = p.init_images # switch image and control_image
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
import time
|
||||
from typing import Union
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline, FluxPipeline, ControlNetModel
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline, FluxPipeline, StableDiffusion3Pipeline, ControlNetModel
|
||||
from modules.control.units import detect
|
||||
from modules.shared import log, opts, listdir
|
||||
from modules import errors, sd_models, devices, model_quant
|
||||
@@ -50,7 +50,6 @@ predefined_sdxl = {
|
||||
'Depth Zoe XL': 'diffusers/controlnet-zoe-depth-sdxl-1.0',
|
||||
'Depth Mid XL': 'diffusers/controlnet-depth-sdxl-1.0-mid',
|
||||
'OpenPose XL': 'thibaud/controlnet-openpose-sdxl-1.0/bin',
|
||||
# 'OpenPose XL': 'thibaud/controlnet-openpose-sdxl-1.0/OpenPoseXL2.safetensors',
|
||||
'Xinsir Union XL': 'xinsir/controlnet-union-sdxl-1.0',
|
||||
'Xinsir OpenPose XL': 'xinsir/controlnet-openpose-sdxl-1.0',
|
||||
'Xinsir Canny XL': 'xinsir/controlnet-canny-sdxl-1.0',
|
||||
@@ -79,11 +78,20 @@ predefined_f1 = {
|
||||
"XLabs-AI Depth": 'XLabs-AI/flux-controlnet-depth-diffusers',
|
||||
"XLabs-AI HED": 'XLabs-AI/flux-controlnet-hed-diffusers'
|
||||
}
|
||||
predefined_sd3 = {
|
||||
"InstantX Canny": 'InstantX/SD3-Controlnet-Canny',
|
||||
"InstantX Pose": 'InstantX/SD3-Controlnet-Pose',
|
||||
"InstantX Depth": 'InstantX/SD3-Controlnet-Depth',
|
||||
"InstantX Tile": 'InstantX/SD3-Controlnet-Tile',
|
||||
"Alimama Inpainting": 'alimama-creative/SD3-Controlnet-Inpainting',
|
||||
"Alimama SoftEdge": 'alimama-creative/SD3-Controlnet-Softedge',
|
||||
}
|
||||
models = {}
|
||||
all_models = {}
|
||||
all_models.update(predefined_sd15)
|
||||
all_models.update(predefined_sdxl)
|
||||
all_models.update(predefined_f1)
|
||||
all_models.update(predefined_sd3)
|
||||
cache_dir = 'models/control/controlnet'
|
||||
|
||||
|
||||
@@ -118,9 +126,11 @@ def list_models(refresh=False):
|
||||
models = ['None'] + list(predefined_sd15) + sorted(find_models())
|
||||
elif modules.shared.sd_model_type == 'f1':
|
||||
models = ['None'] + list(predefined_f1) + sorted(find_models())
|
||||
elif modules.shared.sd_model_type == 'sd3':
|
||||
models = ['None'] + list(predefined_sd3) + sorted(find_models())
|
||||
else:
|
||||
log.warning(f'Control {what} model list failed: unknown model type')
|
||||
models = ['None'] + sorted(predefined_sd15) + sorted(predefined_sdxl) + sorted(find_models())
|
||||
models = ['None'] + sorted(predefined_sd15) + sorted(predefined_sdxl) + sorted(predefined_f1) + sorted(predefined_sd3) + sorted(find_models())
|
||||
debug(f'Control list {what}: path={cache_dir} models={models}')
|
||||
return models
|
||||
|
||||
@@ -151,6 +161,8 @@ class ControlNet():
|
||||
from diffusers import ControlNetModel as model_class # pylint: disable=reimported # sdxl shares same model class
|
||||
elif modules.shared.sd_model_type == 'f1':
|
||||
from diffusers import FluxControlNetModel as model_class
|
||||
elif modules.shared.sd_model_type == 'sd3':
|
||||
from diffusers import SD3ControlNetModel as model_class
|
||||
else:
|
||||
log.error(f'Control {what}: type={modules.shared.sd_model_type} unsupported model')
|
||||
return None
|
||||
@@ -247,7 +259,11 @@ class ControlNet():
|
||||
|
||||
|
||||
class ControlNetPipeline():
|
||||
def __init__(self, controlnet: Union[ControlNetModel, list[ControlNetModel]], pipeline: Union[StableDiffusionXLPipeline, StableDiffusionPipeline, FluxPipeline], dtype = None):
|
||||
def __init__(self,
|
||||
controlnet: Union[ControlNetModel, list[ControlNetModel]],
|
||||
pipeline: Union[StableDiffusionXLPipeline, StableDiffusionPipeline, FluxPipeline, StableDiffusion3Pipeline],
|
||||
dtype = None,
|
||||
):
|
||||
t0 = time.time()
|
||||
self.orig_pipeline = pipeline
|
||||
self.pipeline = None
|
||||
@@ -293,6 +309,20 @@ class ControlNetPipeline():
|
||||
scheduler=pipeline.scheduler,
|
||||
controlnet=controlnet, # can be a list
|
||||
)
|
||||
elif detect.is_sd3(pipeline):
|
||||
from diffusers import StableDiffusion3ControlNetPipeline
|
||||
self.pipeline = StableDiffusion3ControlNetPipeline(
|
||||
vae=pipeline.vae,
|
||||
text_encoder=pipeline.text_encoder,
|
||||
text_encoder_2=pipeline.text_encoder_2,
|
||||
text_encoder_3=pipeline.text_encoder_3,
|
||||
tokenizer=pipeline.tokenizer,
|
||||
tokenizer_2=pipeline.tokenizer_2,
|
||||
tokenizer_3=pipeline.tokenizer_3,
|
||||
transformer=pipeline.transformer,
|
||||
scheduler=pipeline.scheduler,
|
||||
controlnet=controlnet, # can be a list
|
||||
)
|
||||
else:
|
||||
log.error(f'Control {what} pipeline: class={pipeline.__class__.__name__} unsupported model type')
|
||||
return
|
||||
|
||||
@@ -20,5 +20,12 @@ def is_f1(model):
|
||||
if model is None:
|
||||
return False
|
||||
if hasattr(model, '__name__'):
|
||||
return model.__name__ == p.FluxPipeline.__name__
|
||||
return isinstance(model, p.FluxPipeline)
|
||||
return model.__name__ == p.FluxPipeline.__name__ or model.__name__ == p.FluxImg2ImgPipeline.__name__ or model.__name__ == p.FluxInpaintPipeline.__name__
|
||||
return isinstance(model, p.FluxPipeline) or isinstance(model, p.FluxImg2ImgPipeline) or isinstance(model, p.FluxInpaintPipeline)
|
||||
|
||||
def is_sd3(model):
|
||||
if model is None:
|
||||
return False
|
||||
if hasattr(model, '__name__'):
|
||||
return model.__name__ == p.StableDiffusion3Pipeline.__name__ or model.__name__ == p.StableDiffusion3Img2ImgPipeline.__name__ or model.__name__ == p.StableDiffusion3InpaintPipeline.__name__
|
||||
return isinstance(model, p.StableDiffusion3Pipeline) or isinstance(model, p.StableDiffusion3Img2ImgPipeline) or isinstance(model, p.StableDiffusion3InpaintPipeline)
|
||||
|
||||
@@ -4,6 +4,7 @@ import time
|
||||
import contextlib
|
||||
from functools import wraps
|
||||
import torch
|
||||
from modules import rocm
|
||||
from modules.errors import log, display, install as install_traceback
|
||||
from installer import install
|
||||
|
||||
@@ -50,8 +51,8 @@ def has_zluda() -> bool:
|
||||
if not cuda_ok:
|
||||
return False
|
||||
try:
|
||||
device = torch.device("cuda")
|
||||
return torch.cuda.get_device_name(device).endswith("[ZLUDA]")
|
||||
dev = torch.device("cuda")
|
||||
return torch.cuda.get_device_name(dev).endswith("[ZLUDA]")
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
@@ -206,7 +207,7 @@ def torch_gc(force=False, fast=False):
|
||||
force = True
|
||||
if oom > previous_oom:
|
||||
previous_oom = oom
|
||||
log.warning(f'GPU out-of-memory error: {mem}')
|
||||
log.warning(f'Torch GPU out-of-memory error: {mem}')
|
||||
force = True
|
||||
if force:
|
||||
# actual gc
|
||||
@@ -246,13 +247,26 @@ def set_cuda_sync_mode(mode):
|
||||
return
|
||||
try:
|
||||
import ctypes
|
||||
log.info(f'Set cuda sync: mode={mode}')
|
||||
log.info(f'Torch CUDA sync: mode={mode}')
|
||||
torch.cuda.set_device(torch.device(get_optimal_device_name()))
|
||||
ctypes.CDLL('libcudart.so').cudaSetDeviceFlags({'auto': 0, 'spin': 1, 'yield': 2, 'block': 4}[mode])
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def set_cuda_memory_limit():
|
||||
if not cuda_ok or opts.cuda_mem_fraction == 0:
|
||||
return
|
||||
from modules.shared import cmd_opts
|
||||
try:
|
||||
torch_gc(force=True)
|
||||
mem = torch.cuda.get_device_properties(device).total_memory
|
||||
torch.cuda.set_per_process_memory_fraction(float(opts.cuda_mem_fraction), cmd_opts.device_id if cmd_opts.device_id is not None else 0)
|
||||
log.info(f'Torch CUDA memory limit: fraction={opts.cuda_mem_fraction:.2f} limit={round(opts.cuda_mem_fraction * mem / 1024 / 1024)} total={round(mem / 1024 / 1024)}')
|
||||
except Exception as e:
|
||||
log.warning(f'Torch CUDA memory limit: fraction={opts.cuda_mem_fraction:.2f} {e}')
|
||||
|
||||
|
||||
def test_fp16():
|
||||
global fp16_ok # pylint: disable=global-statement
|
||||
if fp16_ok is not None:
|
||||
@@ -283,16 +297,14 @@ def test_bf16():
|
||||
if sys.platform == "darwin" or backend == 'openvino' or backend == 'directml': # override
|
||||
bf16_ok = False
|
||||
return bf16_ok
|
||||
elif backend == 'zluda':
|
||||
device_name = torch.cuda.get_device_name(device)
|
||||
if device_name.startswith("AMD Radeon RX "): # only force AMD
|
||||
device_name = device_name.replace("AMD Radeon RX ", "").split(" ", maxsplit=1)[0]
|
||||
if len(device_name) == 4 and device_name[0] in {"5", "6"}: # RDNA 1 and 2
|
||||
bf16_ok = False
|
||||
return bf16_ok
|
||||
elif backend == 'rocm':
|
||||
gcn_arch = getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000")[3:7]
|
||||
if len(gcn_arch) == 4 and gcn_arch[0:2] == "10": # RDNA 1 and 2
|
||||
elif backend == 'rocm' or backend == 'zluda':
|
||||
agent = None
|
||||
if backend == 'rocm':
|
||||
agent = rocm.Agent(getattr(torch.cuda.get_device_properties(device), "gcnArchName", "gfx0000"))
|
||||
else:
|
||||
from modules.zluda_installer import default_agent
|
||||
agent = default_agent
|
||||
if agent is not None and agent.gfx_version < 0x1100 and agent.arch != rocm.MicroArchitecture.CDNA: # all cards before RDNA 3 except for CDNA cards
|
||||
bf16_ok = False
|
||||
return bf16_ok
|
||||
try:
|
||||
@@ -450,6 +462,7 @@ def set_dtype():
|
||||
|
||||
def set_cuda_params():
|
||||
override_ipex_math()
|
||||
set_cuda_memory_limit()
|
||||
set_cudnn_params()
|
||||
set_sdpa_params()
|
||||
set_dtype()
|
||||
|
||||
@@ -188,7 +188,7 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument
|
||||
_, extension = os.path.splitext(output_modelname)
|
||||
|
||||
if os.path.exists(output_modelname) and not kwargs.get("overwrite", False):
|
||||
return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], f"Model alredy exists: {output_modelname}"]
|
||||
return [*[gr.Dropdown.update(choices=sd_models.checkpoint_titles()) for _ in range(4)], f"Model alredy exists: {output_modelname}"]
|
||||
if extension.lower() == ".safetensors":
|
||||
safetensors.torch.save_file(theta_0, output_modelname, metadata=metadata)
|
||||
else:
|
||||
@@ -202,7 +202,7 @@ def run_modelmerger(id_task, **kwargs): # pylint: disable=unused-argument
|
||||
created_model.calculate_shorthash()
|
||||
devices.torch_gc(force=True)
|
||||
shared.state.end()
|
||||
return [*[gr.Dropdown.update(choices=sd_models.checkpoint_tiles()) for _ in range(4)], f"Model saved to {output_modelname}"]
|
||||
return [*[gr.Dropdown.update(choices=sd_models.checkpoint_titles()) for _ in range(4)], f"Model saved to {output_modelname}"]
|
||||
|
||||
|
||||
def run_modelconvert(model, checkpoint_formats, precision, conv_type, custom_name, unet_conv, text_encoder_conv,
|
||||
|
||||
@@ -6,9 +6,10 @@ import numpy as np
|
||||
import diffusers
|
||||
import huggingface_hub as hf
|
||||
from PIL import Image
|
||||
from modules import processing, shared, devices, extra_networks, sd_models, sd_hijack_freeu, script_callbacks, ipadapter
|
||||
from modules import processing, shared, devices, extra_networks, sd_hijack_freeu, script_callbacks, ipadapter, token_merge
|
||||
from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet
|
||||
|
||||
|
||||
FACEID_MODELS = {
|
||||
"FaceID Base": "h94/IP-Adapter-FaceID/ip-adapter-faceid_sd15.bin",
|
||||
"FaceID Plus v1": "h94/IP-Adapter-FaceID/ip-adapter-faceid-plus_sd15.bin",
|
||||
@@ -69,7 +70,7 @@ def face_id(
|
||||
shared.prompt_styles.apply_styles_to_extra(p)
|
||||
|
||||
if shared.opts.cuda_compile_backend == 'none':
|
||||
sd_models.apply_token_merging(p.sd_model)
|
||||
token_merge.apply_token_merging(p.sd_model)
|
||||
sd_hijack_freeu.apply_freeu(p, not shared.native)
|
||||
|
||||
script_callbacks.before_process_callback(p)
|
||||
@@ -78,7 +79,7 @@ def face_id(
|
||||
ip_ckpt = FACEID_MODELS[model]
|
||||
folder, filename = os.path.split(ip_ckpt)
|
||||
basename, _ext = os.path.splitext(filename)
|
||||
model_path = hf.hf_hub_download(repo_id=folder, filename=filename, cache_dir=shared.opts.diffusers_dir)
|
||||
model_path = hf.hf_hub_download(repo_id=folder, filename=filename, cache_dir=shared.opts.hfcache_dir)
|
||||
if model_path is None:
|
||||
shared.log.error(f'FaceID download failed: model={model} file="{ip_ckpt}"')
|
||||
return None
|
||||
@@ -246,7 +247,7 @@ def face_id(
|
||||
if faceid_model is not None and original_load_ip_adapter is not None:
|
||||
faceid_model.__class__.load_ip_adapter = original_load_ip_adapter
|
||||
if shared.opts.cuda_compile_backend == 'none':
|
||||
sd_models.remove_token_merging(p.sd_model)
|
||||
token_merge.remove_token_merging(p.sd_model)
|
||||
script_callbacks.after_process_callback(p)
|
||||
|
||||
return processed_images
|
||||
|
||||
@@ -32,7 +32,7 @@ def get_app(mp_name):
|
||||
repo_id='vladmandic/insightface-faceanalysis',
|
||||
filename=f'{mp_name}.zip',
|
||||
local_dir_use_symlinks=False,
|
||||
cache_dir=opts.diffusers_dir,
|
||||
cache_dir=opts.hfcache_dir,
|
||||
local_dir=local_dir
|
||||
)
|
||||
if not os.path.exists(extract_dir):
|
||||
|
||||
@@ -189,7 +189,7 @@ def create_override_settings_dict(text_pairs):
|
||||
def connect_paste(button, local_paste_fields, input_comp, override_settings_component, tabname):
|
||||
|
||||
def paste_func(prompt):
|
||||
if prompt is None or len(prompt.strip()) == 0 and not shared.cmd_opts.hide_ui_dir_config:
|
||||
if prompt is None or len(prompt.strip()) == 0:
|
||||
filename = os.path.join(data_path, "params.txt")
|
||||
if os.path.exists(filename):
|
||||
with open(filename, "r", encoding="utf8") as file:
|
||||
|
||||
@@ -10,7 +10,7 @@ import threading
|
||||
import numpy as np
|
||||
import piexif
|
||||
import piexif.helper
|
||||
from PIL import Image, PngImagePlugin, ExifTags
|
||||
from PIL import Image, PngImagePlugin, ExifTags, ImageDraw
|
||||
from modules import sd_samplers, shared, script_callbacks, errors, paths
|
||||
from modules.images_grid import image_grid, get_grid_size, split_grid, combine_grid, check_grid_size, get_font, draw_grid_annotations, draw_prompt_matrix, GridAnnotation, Grid # pylint: disable=unused-import
|
||||
from modules.images_resize import resize_image # pylint: disable=unused-import
|
||||
@@ -361,11 +361,21 @@ def flatten(img, bgcolor):
|
||||
return img.convert('RGB')
|
||||
|
||||
|
||||
def draw_overlay(im, text: str = '', y_offset: int = 0):
|
||||
d = ImageDraw.Draw(im)
|
||||
fontsize = (im.width + im.height) // 50
|
||||
font = get_font(fontsize)
|
||||
d.text((fontsize//2, fontsize//2 + y_offset), text, font=font, fill=shared.opts.font_color)
|
||||
return im
|
||||
|
||||
|
||||
def set_watermark(image, watermark):
|
||||
if shared.opts.image_watermark_position != 'none': # visible watermark
|
||||
wm_image = None
|
||||
try:
|
||||
wm_image = Image.open(shared.opts.image_watermark_image)
|
||||
if wm_image.mode != 'RGBA':
|
||||
wm_image = wm_image.convert('RGBA')
|
||||
except Exception as e:
|
||||
shared.log.warning(f'Set image watermark: fn="{shared.opts.image_watermark_image}" {e}')
|
||||
if wm_image is not None:
|
||||
@@ -384,8 +394,14 @@ def set_watermark(image, watermark):
|
||||
try:
|
||||
for x in range(wm_image.width):
|
||||
for y in range(wm_image.height):
|
||||
r, g, b, _a = wm_image.getpixel((x, y))
|
||||
if not (r == 0 and g == 0 and b == 0):
|
||||
rgba = wm_image.getpixel((x, y))
|
||||
orig = image.getpixel((x+position[0], y+position[1]))
|
||||
# alpha blend
|
||||
a = rgba[3] / 255
|
||||
r = int(rgba[0] * a + orig[0] * (1 - a))
|
||||
g = int(rgba[1] * a + orig[1] * (1 - a))
|
||||
b = int(rgba[2] * a + orig[2] * (1 - a))
|
||||
if not a == 0:
|
||||
image.putpixel((x+position[0], y+position[1]), (r, g, b))
|
||||
shared.log.debug(f'Set image watermark: fn="{shared.opts.image_watermark_image}" image={wm_image} position={position}')
|
||||
except Exception as e:
|
||||
|
||||
@@ -113,7 +113,7 @@ def get_font(fontsize):
|
||||
return ImageFont.truetype("javascript/notosans-nerdfont-regular.ttf", fontsize)
|
||||
|
||||
|
||||
def draw_grid_annotations(im, width, height, hor_texts, ver_texts, margin=0, title=None):
|
||||
def draw_grid_annotations(im, width, height, x_texts, y_texts, margin=0, title=None):
|
||||
def wrap(drawing, text, font, line_length):
|
||||
lines = ['']
|
||||
for word in text.split():
|
||||
@@ -140,15 +140,15 @@ def draw_grid_annotations(im, width, height, hor_texts, ver_texts, margin=0, tit
|
||||
line_spacing = fontsize // 2
|
||||
font = get_font(fontsize)
|
||||
color_inactive = (127, 127, 127)
|
||||
pad_left = 0 if sum([sum([len(line.text) for line in lines]) for lines in ver_texts]) == 0 else width * 3 // 4
|
||||
cols = len(hor_texts)
|
||||
rows = len(ver_texts)
|
||||
pad_left = 0 if sum([sum([len(line.text) for line in lines]) for lines in y_texts]) == 0 else width * 3 // 4
|
||||
cols = len(x_texts)
|
||||
rows = len(y_texts)
|
||||
# assert cols == len(hor_texts), f'bad number of horizontal texts: {len(hor_texts)}; must be {cols}'
|
||||
# assert rows == len(hor_texts), f'bad number of vertical texts: {len(ver_texts)}; must be {rows}'
|
||||
calc_img = Image.new("RGB", (1, 1), shared.opts.grid_background)
|
||||
calc_d = ImageDraw.Draw(calc_img)
|
||||
title_texts = [title] if title else [[GridAnnotation()]]
|
||||
for texts, allowed_width in zip(hor_texts + ver_texts + title_texts, [width] * len(hor_texts) + [pad_left] * len(ver_texts) + [(width+margin)*cols]):
|
||||
for texts, allowed_width in zip(x_texts + y_texts + title_texts, [width] * len(x_texts) + [pad_left] * len(y_texts) + [(width+margin)*cols]):
|
||||
items = [] + texts
|
||||
texts.clear()
|
||||
for line in items:
|
||||
@@ -158,8 +158,8 @@ def draw_grid_annotations(im, width, height, hor_texts, ver_texts, margin=0, tit
|
||||
bbox = calc_d.multiline_textbbox((0, 0), line.text, font=font)
|
||||
line.size = (bbox[2] - bbox[0], bbox[3] - bbox[1])
|
||||
line.allowed_width = allowed_width
|
||||
hor_text_heights = [sum([line.size[1] + line_spacing for line in lines]) - line_spacing for lines in hor_texts]
|
||||
ver_text_heights = [sum([line.size[1] + line_spacing for line in lines]) - line_spacing * len(lines) for lines in ver_texts]
|
||||
hor_text_heights = [sum([line.size[1] + line_spacing for line in lines]) - line_spacing for lines in x_texts]
|
||||
ver_text_heights = [sum([line.size[1] + line_spacing for line in lines]) - line_spacing * len(lines) for lines in y_texts]
|
||||
pad_top = 0 if sum(hor_text_heights) == 0 else max(hor_text_heights) + line_spacing * 2
|
||||
title_pad = 0
|
||||
if title:
|
||||
@@ -178,11 +178,11 @@ def draw_grid_annotations(im, width, height, hor_texts, ver_texts, margin=0, tit
|
||||
for col in range(cols):
|
||||
x = pad_left + (width + margin) * col + width / 2
|
||||
y = (pad_top / 2 - hor_text_heights[col] / 2) + title_pad
|
||||
draw_texts(d, x, y, hor_texts[col], font, fontsize)
|
||||
draw_texts(d, x, y, x_texts[col], font, fontsize)
|
||||
for row in range(rows):
|
||||
x = pad_left / 2
|
||||
y = (pad_top + (height + margin) * row + height / 2 - ver_text_heights[row] / 2) + title_pad
|
||||
draw_texts(d, x, y, ver_texts[row], font, fontsize)
|
||||
draw_texts(d, x, y, y_texts[row], font, fontsize)
|
||||
return result
|
||||
|
||||
|
||||
|
||||
@@ -267,6 +267,7 @@ def img2img(id_task: str, state: str, mode: int,
|
||||
processed = modules.scripts.scripts_img2img.run(p, *args)
|
||||
if processed is None:
|
||||
processed = processing.process_images(p)
|
||||
processed = modules.scripts.scripts_img2img.after(p, processed, *args)
|
||||
p.close()
|
||||
generation_info_js = processed.js() if processed is not None else ''
|
||||
if processed is None:
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -3,8 +3,6 @@ Lightweight IP-Adapter applied to existing pipeline in Diffusers
|
||||
- Downloads image_encoder or first usage (2.5GB)
|
||||
- Introduced via: https://github.com/huggingface/diffusers/pull/5713
|
||||
- IP adapters: https://huggingface.co/h94/IP-Adapter
|
||||
TODO ipadapter items:
|
||||
- SD/SDXL autodetect
|
||||
"""
|
||||
|
||||
import os
|
||||
@@ -14,21 +12,42 @@ from PIL import Image
|
||||
from modules import processing, shared, devices, sd_models
|
||||
|
||||
|
||||
base_repo = "h94/IP-Adapter"
|
||||
clip_repo = "h94/IP-Adapter"
|
||||
clip_loaded = None
|
||||
ADAPTERS = {
|
||||
'None': 'none',
|
||||
'Base': 'ip-adapter_sd15.safetensors',
|
||||
'Base ViT-G': 'ip-adapter_sd15_vit-G.safetensors',
|
||||
'Light': 'ip-adapter_sd15_light.safetensors',
|
||||
'Plus': 'ip-adapter-plus_sd15.safetensors',
|
||||
'Plus Face': 'ip-adapter-plus-face_sd15.safetensors',
|
||||
'Full Face': 'ip-adapter-full-face_sd15.safetensors',
|
||||
'Base SDXL': 'ip-adapter_sdxl.safetensors',
|
||||
'Base ViT-H SDXL': 'ip-adapter_sdxl_vit-h.safetensors',
|
||||
'Plus ViT-H SDXL': 'ip-adapter-plus_sdxl_vit-h.safetensors',
|
||||
'Plus Face ViT-H SDXL': 'ip-adapter-plus-face_sdxl_vit-h.safetensors',
|
||||
ADAPTERS_NONE = {
|
||||
'None': { 'name': 'none', 'repo': 'none', 'subfolder': 'none' },
|
||||
}
|
||||
ADAPTERS_SD15 = {
|
||||
'None': { 'name': 'none', 'repo': 'none', 'subfolder': 'none' },
|
||||
'Base': { 'name': 'ip-adapter_sd15.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' },
|
||||
'Base ViT-G': { 'name': 'ip-adapter_sd15_vit-G.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' },
|
||||
'Light': { 'name': 'ip-adapter_sd15_light.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' },
|
||||
'Plus': { 'name': 'ip-adapter-plus_sd15.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' },
|
||||
'Plus Face': { 'name': 'ip-adapter-plus-face_sd15.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' },
|
||||
'Full Face': { 'name': 'ip-adapter-full-face_sd15.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'models' },
|
||||
'Ostris Composition ViT-H': { 'name': 'ip_plus_composition_sd15.safetensors', 'repo': 'ostris/ip-composition-adapter', 'subfolder': '' },
|
||||
}
|
||||
ADAPTERS_SDXL = {
|
||||
'None': { 'name': 'none', 'repo': 'none', 'subfolder': 'none' },
|
||||
'Base SDXL': { 'name': 'ip-adapter_sdxl.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'sdxl_models' },
|
||||
'Base ViT-H SDXL': { 'name': 'ip-adapter_sdxl_vit-h.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'sdxl_models' },
|
||||
'Plus ViT-H SDXL': { 'name': 'ip-adapter-plus_sdxl_vit-h.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'sdxl_models' },
|
||||
'Plus Face ViT-H SDXL': { 'name': 'ip-adapter-plus-face_sdxl_vit-h.safetensors', 'repo': 'h94/IP-Adapter', 'subfolder': 'sdxl_models' },
|
||||
'Ostris Composition ViT-H SDXL': { 'name': 'ip_plus_composition_sdxl.safetensors', 'repo': 'ostris/ip-composition-adapter', 'subfolder': '' },
|
||||
}
|
||||
ADAPTERS = { **ADAPTERS_SD15, **ADAPTERS_SDXL }
|
||||
ADAPTERS_ALL = { **ADAPTERS_SD15, **ADAPTERS_SDXL }
|
||||
|
||||
|
||||
def get_adapters():
|
||||
global ADAPTERS # pylint: disable=global-statement
|
||||
if shared.sd_model_type == 'sd':
|
||||
ADAPTERS = ADAPTERS_SD15
|
||||
elif shared.sd_model_type == 'sdxl':
|
||||
ADAPTERS = ADAPTERS_SDXL
|
||||
else:
|
||||
ADAPTERS = ADAPTERS_NONE
|
||||
return list(ADAPTERS)
|
||||
|
||||
|
||||
def get_images(input_images):
|
||||
@@ -83,13 +102,12 @@ def crop_images(images, crops):
|
||||
try:
|
||||
for i in range(len(images)):
|
||||
if crops[i]:
|
||||
from shared import yolo # pylint: disable=no-name-in-module
|
||||
yolo.load()
|
||||
from modules.shared import yolo # pylint: disable=no-name-in-module
|
||||
cropped = []
|
||||
for image in images[i]:
|
||||
faces = yolo.predict(image)
|
||||
faces = yolo.predict('face-yolo8n', image)
|
||||
if len(faces) > 0:
|
||||
cropped.append(faces[0].face)
|
||||
cropped.append(faces[0].item)
|
||||
if len(cropped) == len(images[i]):
|
||||
images[i] = cropped
|
||||
else:
|
||||
@@ -117,13 +135,13 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt
|
||||
if hasattr(p, 'ip_adapter_names'):
|
||||
if isinstance(p.ip_adapter_names, str):
|
||||
p.ip_adapter_names = [p.ip_adapter_names]
|
||||
adapters = [ADAPTERS.get(adapter, None) for adapter in p.ip_adapter_names if adapter is not None and adapter.lower() != 'none']
|
||||
adapters = [ADAPTERS_ALL.get(adapter_name, None) for adapter_name in p.ip_adapter_names if adapter_name is not None and adapter_name.lower() != 'none']
|
||||
adapter_names = p.ip_adapter_names
|
||||
else:
|
||||
if isinstance(adapter_names, str):
|
||||
adapter_names = [adapter_names]
|
||||
adapters = [ADAPTERS.get(adapter, None) for adapter in adapter_names]
|
||||
adapters = [adapter for adapter in adapters if adapter is not None and adapter.lower() != 'none']
|
||||
adapters = [adapter for adapter in adapters if adapter is not None and adapter['name'].lower() != 'none']
|
||||
if len(adapters) == 0:
|
||||
unapply(pipe)
|
||||
if hasattr(p, 'ip_adapter_images'):
|
||||
@@ -189,41 +207,48 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt
|
||||
|
||||
for adapter_name in adapter_names:
|
||||
# which clip to use
|
||||
if 'ViT' not in adapter_name:
|
||||
clip_repo = base_repo
|
||||
clip_subfolder = 'models/image_encoder' if shared.sd_model_type == 'sd' else 'sdxl_models/image_encoder' # defaults per model
|
||||
if 'ViT' not in adapter_name: # defaults per model
|
||||
if shared.sd_model_type == 'sd':
|
||||
clip_subfolder = 'models/image_encoder'
|
||||
else:
|
||||
clip_subfolder = 'sdxl_models/image_encoder'
|
||||
elif 'ViT-H' in adapter_name:
|
||||
clip_repo = base_repo
|
||||
clip_subfolder = 'models/image_encoder' # this is vit-h
|
||||
elif 'ViT-G' in adapter_name:
|
||||
clip_repo = base_repo
|
||||
clip_subfolder = 'sdxl_models/image_encoder' # this is vit-g
|
||||
else:
|
||||
shared.log.error(f'IP adapter: unknown model type: {adapter_name}')
|
||||
return False
|
||||
|
||||
# load feature extractor used by ip adapter
|
||||
if pipe.feature_extractor is None:
|
||||
# load feature extractor used by ip adapter
|
||||
if pipe.feature_extractor is None:
|
||||
try:
|
||||
from transformers import CLIPImageProcessor
|
||||
shared.log.debug('IP adapter load: feature extractor')
|
||||
pipe.feature_extractor = CLIPImageProcessor()
|
||||
# load image encoder used by ip adapter
|
||||
if pipe.image_encoder is None or clip_loaded != f'{clip_repo}/{clip_subfolder}':
|
||||
try:
|
||||
from transformers import CLIPVisionModelWithProjection
|
||||
shared.log.debug(f'IP adapter load: image encoder="{clip_repo}/{clip_subfolder}"')
|
||||
pipe.image_encoder = CLIPVisionModelWithProjection.from_pretrained(clip_repo, subfolder=clip_subfolder, torch_dtype=devices.dtype, cache_dir=shared.opts.diffusers_dir, use_safetensors=True)
|
||||
clip_loaded = f'{clip_repo}/{clip_subfolder}'
|
||||
except Exception as e:
|
||||
shared.log.error(f'IP adapter: failed to load image encoder: {e}')
|
||||
return False
|
||||
sd_models.move_model(pipe.image_encoder, devices.device)
|
||||
except Exception as e:
|
||||
shared.log.error(f'IP adapter load: feature extractor {e}')
|
||||
return False
|
||||
|
||||
# load image encoder used by ip adapter
|
||||
if pipe.image_encoder is None or clip_loaded != f'{clip_repo}/{clip_subfolder}':
|
||||
try:
|
||||
from transformers import CLIPVisionModelWithProjection
|
||||
shared.log.debug(f'IP adapter load: image encoder="{clip_repo}/{clip_subfolder}"')
|
||||
pipe.image_encoder = CLIPVisionModelWithProjection.from_pretrained(clip_repo, subfolder=clip_subfolder, torch_dtype=devices.dtype, cache_dir=shared.opts.diffusers_dir, use_safetensors=True)
|
||||
clip_loaded = f'{clip_repo}/{clip_subfolder}'
|
||||
except Exception as e:
|
||||
shared.log.error(f'IP adapter load: image encoder="{clip_repo}/{clip_subfolder}" {e}')
|
||||
return False
|
||||
sd_models.move_model(pipe.image_encoder, devices.device)
|
||||
|
||||
# main code
|
||||
t0 = time.time()
|
||||
ip_subfolder = 'models' if shared.sd_model_type == 'sd' else 'sdxl_models'
|
||||
try:
|
||||
pipe.load_ip_adapter([base_repo], subfolder=[ip_subfolder], weight_name=adapters)
|
||||
t0 = time.time()
|
||||
repos = [adapter['repo'] for adapter in adapters]
|
||||
subfolders = [adapter['subfolder'] for adapter in adapters]
|
||||
names = [adapter['name'] for adapter in adapters]
|
||||
pipe.load_ip_adapter(repos, subfolder=subfolders, weight_name=names)
|
||||
if hasattr(p, 'ip_adapter_layers'):
|
||||
pipe.set_ip_adapter_scale(p.ip_adapter_layers)
|
||||
ip_str = ';'.join(adapter_names) + ':' + json.dumps(p.ip_adapter_layers)
|
||||
@@ -240,5 +265,5 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt
|
||||
t1 = time.time()
|
||||
shared.log.info(f'IP adapter: {ip_str} image={adapter_images} mask={adapter_masks is not None} time={t1-t0:.2f}')
|
||||
except Exception as e:
|
||||
shared.log.error(f'IP adapter failed to load: repo="{base_repo}" folder="{ip_subfolder}" weights={adapters} names={adapter_names} {e}')
|
||||
shared.log.error(f'IP adapter load: adapters={adapter_names} repo={repos} folders={subfolders} names={names} {e}')
|
||||
return True
|
||||
|
||||
@@ -44,6 +44,8 @@ if ".dev" in torch.__version__ or "+git" in torch.__version__:
|
||||
timer.startup.record("torch")
|
||||
|
||||
import transformers # pylint: disable=W0611,C0411
|
||||
from transformers import logging as transformers_logging # pylint: disable=W0611,C0411
|
||||
transformers_logging.set_verbosity_error()
|
||||
timer.startup.record("transformers")
|
||||
|
||||
import accelerate # pylint: disable=W0611,C0411
|
||||
@@ -61,6 +63,9 @@ errors.install([gradio])
|
||||
import pydantic # pylint: disable=W0611,C0411
|
||||
timer.startup.record("pydantic")
|
||||
|
||||
import diffusers.utils.import_utils # pylint: disable=W0611,C0411
|
||||
diffusers.utils.import_utils._k_diffusion_available = True # pylint: disable=protected-access # monkey-patch since we use k-diffusion from git
|
||||
diffusers.utils.import_utils._k_diffusion_version = '0.0.12' # pylint: disable=protected-access
|
||||
import diffusers # pylint: disable=W0611,C0411
|
||||
import diffusers.loaders.single_file # pylint: disable=W0611,C0411
|
||||
import huggingface_hub # pylint: disable=W0611,C0411
|
||||
|
||||
@@ -122,10 +122,12 @@ def quant_flux_bnb(checkpoint_info, transformer, text_encoder_2):
|
||||
bnb_4bit_quant_type=shared.opts.bnb_quantization_type,
|
||||
bnb_4bit_compute_dtype=devices.dtype
|
||||
)
|
||||
if 'Model' in shared.opts.bnb_quantization and transformer is None:
|
||||
if ('Model' in shared.opts.bnb_quantization) and (transformer is None):
|
||||
transformer = diffusers.FluxTransformer2DModel.from_pretrained(repo_id, subfolder="transformer", cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype)
|
||||
shared.log.debug(f'Quantization: module=transformer type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}')
|
||||
if 'Text Encoder' in shared.opts.bnb_quantization and text_encoder_2 is None:
|
||||
if ('Text Encoder' in shared.opts.bnb_quantization) and (text_encoder_2 is None):
|
||||
if repo_id == 'sayakpaul/flux.1-dev-nf4':
|
||||
repo_id = 'black-forest-labs/FLUX.1-dev' # workaround since sayakpaul model is missing model_index.json
|
||||
text_encoder_2 = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_2", cache_dir=cache_dir, quantization_config=bnb_config, torch_dtype=devices.dtype)
|
||||
shared.log.debug(f'Quantization: module=t5 type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}')
|
||||
except Exception as e:
|
||||
@@ -285,25 +287,26 @@ def load_flux(checkpoint_info, diffusers_load_config): # triggered by opts.sd_ch
|
||||
errors.display(e, 'FLUX Quanto:')
|
||||
|
||||
# initialize pipeline with pre-loaded components
|
||||
components = {}
|
||||
transformer, text_encoder_2 = quant_flux_bnb(checkpoint_info, transformer, text_encoder_2)
|
||||
kwargs = {}
|
||||
# transformer, text_encoder_2 = quant_flux_bnb(checkpoint_info, transformer, text_encoder_2)
|
||||
if transformer is not None:
|
||||
components['transformer'] = transformer
|
||||
kwargs['transformer'] = transformer
|
||||
sd_unet.loaded_unet = shared.opts.sd_unet
|
||||
if text_encoder_1 is not None:
|
||||
components['text_encoder'] = text_encoder_1
|
||||
kwargs['text_encoder'] = text_encoder_1
|
||||
model_te.loaded_te = shared.opts.sd_text_encoder
|
||||
if text_encoder_2 is not None:
|
||||
components['text_encoder_2'] = text_encoder_2
|
||||
kwargs['text_encoder_2'] = text_encoder_2
|
||||
model_te.loaded_te = shared.opts.sd_text_encoder
|
||||
if vae is not None:
|
||||
components['vae'] = vae
|
||||
shared.log.debug(f'Load model: type=FLUX preloaded={list(components)}')
|
||||
kwargs['vae'] = vae
|
||||
shared.log.debug(f'Load model: type=FLUX preloaded={list(kwargs)}')
|
||||
if repo_id == 'sayakpaul/flux.1-dev-nf4':
|
||||
repo_id = 'black-forest-labs/FLUX.1-dev' # workaround since sayakpaul model is missing model_index.json
|
||||
for c in components:
|
||||
if components[c].dtype == torch.float32 and devices.dtype != torch.float32:
|
||||
shared.log.warning(f'Load model: type=FLUX component={c} dtype={components[c].dtype} cast dtype={devices.dtype}')
|
||||
components[c] = components[c].to(dtype=devices.dtype)
|
||||
pipe = diffusers.FluxPipeline.from_pretrained(repo_id, cache_dir=shared.opts.diffusers_dir, **components, **diffusers_load_config)
|
||||
for c in kwargs:
|
||||
if kwargs[c].dtype == torch.float32 and devices.dtype != torch.float32:
|
||||
shared.log.warning(f'Load model: type=FLUX component={c} dtype={kwargs[c].dtype} cast dtype={devices.dtype}')
|
||||
kwargs[c] = kwargs[c].to(dtype=devices.dtype)
|
||||
kwargs = model_quant.create_bnb_config(kwargs)
|
||||
pipe = diffusers.FluxPipeline.from_pretrained(repo_id, cache_dir=shared.opts.diffusers_dir, **kwargs, **diffusers_load_config)
|
||||
return pipe
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import sys
|
||||
import diffusers
|
||||
from installer import install, log
|
||||
|
||||
|
||||
@@ -6,6 +7,27 @@ bnb = None
|
||||
quanto = None
|
||||
|
||||
|
||||
def create_bnb_config(kwargs = None):
|
||||
from modules import shared, devices
|
||||
if len(shared.opts.bnb_quantization) > 0:
|
||||
if 'Model' in shared.opts.bnb_quantization and 'transformer' not in (kwargs or {}):
|
||||
load_bnb()
|
||||
bnb_config = diffusers.BitsAndBytesConfig(
|
||||
load_in_8bit=shared.opts.bnb_quantization_type in ['fp8'],
|
||||
load_in_4bit=shared.opts.bnb_quantization_type in ['nf4', 'fp4'],
|
||||
bnb_4bit_quant_storage=shared.opts.bnb_quantization_storage,
|
||||
bnb_4bit_quant_type=shared.opts.bnb_quantization_type,
|
||||
bnb_4bit_compute_dtype=devices.dtype
|
||||
)
|
||||
shared.log.debug(f'Quantization: module=all type=bnb dtype={shared.opts.bnb_quantization_type} storage={shared.opts.bnb_quantization_storage}')
|
||||
if kwargs is None:
|
||||
return bnb_config
|
||||
else:
|
||||
kwargs['quantization_config'] = bnb_config
|
||||
return kwargs
|
||||
return kwargs
|
||||
|
||||
|
||||
def load_bnb(msg='', silent=False):
|
||||
global bnb # pylint: disable=global-statement
|
||||
if bnb is not None:
|
||||
@@ -16,6 +38,8 @@ def load_bnb(msg='', silent=False):
|
||||
try:
|
||||
import bitsandbytes
|
||||
bnb = bitsandbytes
|
||||
diffusers.utils.import_utils._bitsandbytes_available = True # pylint: disable=protected-access
|
||||
diffusers.utils.import_utils._bitsandbytes_version = '0.43.3' # pylint: disable=protected-access
|
||||
return bnb
|
||||
except Exception as e:
|
||||
if len(msg) > 0:
|
||||
@@ -23,6 +47,7 @@ def load_bnb(msg='', silent=False):
|
||||
bnb = None
|
||||
if not silent:
|
||||
raise
|
||||
return None
|
||||
|
||||
|
||||
def load_quanto(msg='', silent=False):
|
||||
@@ -42,6 +67,7 @@ def load_quanto(msg='', silent=False):
|
||||
quanto = None
|
||||
if not silent:
|
||||
raise
|
||||
return None
|
||||
|
||||
|
||||
def get_quant(name):
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
import diffusers
|
||||
import transformers
|
||||
from modules import shared, devices, sd_models, sd_unet, model_te
|
||||
from modules import shared, devices, sd_models, sd_unet, model_te, model_quant, model_tools
|
||||
|
||||
|
||||
def load_overrides(kwargs, cache_dir):
|
||||
@@ -51,8 +51,7 @@ def load_overrides(kwargs, cache_dir):
|
||||
|
||||
def load_quants(kwargs, repo_id, cache_dir):
|
||||
if len(shared.opts.bnb_quantization) > 0:
|
||||
from modules.model_quant import load_bnb
|
||||
load_bnb('Load model: type=SD3')
|
||||
model_quant.load_bnb('Load model: type=SD3')
|
||||
bnb_config = diffusers.BitsAndBytesConfig(
|
||||
load_in_8bit=shared.opts.bnb_quantization_type in ['fp8'],
|
||||
load_in_4bit=shared.opts.bnb_quantization_type in ['nf4', 'fp4'],
|
||||
@@ -70,7 +69,7 @@ def load_quants(kwargs, repo_id, cache_dir):
|
||||
|
||||
|
||||
def load_missing(kwargs, fn, cache_dir):
|
||||
keys = sd_models.get_safetensor_keys(fn)
|
||||
keys = model_tools.get_safetensor_keys(fn)
|
||||
size = os.stat(fn).st_size // 1024 // 1024
|
||||
if size > 15000:
|
||||
repo_id = 'stabilityai/stable-diffusion-3.5-large'
|
||||
@@ -85,6 +84,9 @@ def load_missing(kwargs, fn, cache_dir):
|
||||
if 'text_encoder_3' not in kwargs and 'text_encoder_3' not in keys:
|
||||
kwargs['text_encoder_3'] = transformers.T5EncoderModel.from_pretrained(repo_id, subfolder="text_encoder_3", variant='fp16', cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
shared.log.debug(f'Load model: type=SD3 missing=te3 repo="{repo_id}"')
|
||||
if 'vae' not in kwargs and 'vae' not in keys:
|
||||
kwargs['vae'] = diffusers.AutoencoderKL.from_pretrained(repo_id, subfolder='vae', cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
shared.log.debug(f'Load model: type=SD3 missing=vae repo="{repo_id}"')
|
||||
# if 'transformer' not in kwargs and 'transformer' not in keys:
|
||||
# kwargs['transformer'] = diffusers.SD3Transformer2DModel.from_pretrained(default_repo_id, subfolder="transformer", cache_dir=cache_dir, torch_dtype=devices.dtype)
|
||||
return kwargs
|
||||
@@ -120,13 +122,18 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None):
|
||||
|
||||
kwargs = {}
|
||||
kwargs = load_overrides(kwargs, cache_dir)
|
||||
kwargs = load_quants(kwargs, repo_id, cache_dir)
|
||||
if fn is None or not os.path.exists(fn):
|
||||
kwargs = load_quants(kwargs, repo_id, cache_dir)
|
||||
|
||||
loader = diffusers.StableDiffusion3Pipeline.from_pretrained
|
||||
if fn is not None and os.path.exists(fn):
|
||||
if fn is not None and os.path.exists(fn) and os.path.isfile(fn):
|
||||
if fn.endswith('.safetensors'):
|
||||
loader = diffusers.StableDiffusion3Pipeline.from_single_file
|
||||
kwargs = load_missing(kwargs, fn, cache_dir)
|
||||
# required_modules = model_tools.get_modules(diffusers.StableDiffusion3Pipeline)
|
||||
# have_modules = model_tools.get_safetensor_keys(fn)
|
||||
# loaded_modules = model_tools.load_modules('stabilityai/stable-diffusion-3.5-medium', required_modules)
|
||||
# kwargs = {**kwargs, **loaded_modules}
|
||||
# kwargs = load_missing(kwargs, fn, cache_dir)
|
||||
repo_id = fn
|
||||
elif fn.endswith('.gguf'):
|
||||
kwargs = load_gguf(kwargs, fn)
|
||||
@@ -135,8 +142,9 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None):
|
||||
else:
|
||||
kwargs['variant'] = 'fp16'
|
||||
|
||||
shared.log.debug(f'Load model: type=SD3 preloaded={list(kwargs)}')
|
||||
shared.log.debug(f'Load model: type=SD3 kwargs={list(kwargs)} repo="{repo_id}"')
|
||||
|
||||
kwargs = model_quant.create_bnb_config(kwargs)
|
||||
pipe = loader(
|
||||
repo_id,
|
||||
torch_dtype=devices.dtype,
|
||||
@@ -144,5 +152,5 @@ def load_sd3(checkpoint_info, cache_dir=None, config=None):
|
||||
config=config,
|
||||
**kwargs,
|
||||
)
|
||||
devices.torch_gc()
|
||||
devices.torch_gc(force=True)
|
||||
return pipe
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
import inspect
|
||||
import diffusers
|
||||
import transformers
|
||||
import safetensors.torch
|
||||
from modules import shared, devices, model_quant
|
||||
|
||||
|
||||
def remove_entries_after_depth(d, depth, current_depth=0):
|
||||
if current_depth >= depth:
|
||||
return None
|
||||
if isinstance(d, dict):
|
||||
return {k: remove_entries_after_depth(v, depth, current_depth + 1) for k, v in d.items() if remove_entries_after_depth(v, depth, current_depth + 1) is not None}
|
||||
return d
|
||||
|
||||
|
||||
def list_to_dict(flat_list):
|
||||
result_dict = {}
|
||||
try:
|
||||
for item in flat_list:
|
||||
keys = item.split('.')
|
||||
d = result_dict
|
||||
for key in keys[:-1]:
|
||||
d = d.setdefault(key, {})
|
||||
d[keys[-1]] = None
|
||||
except Exception:
|
||||
pass
|
||||
return result_dict
|
||||
|
||||
|
||||
def get_safetensor_keys(filename):
|
||||
keys = []
|
||||
try:
|
||||
with safetensors.torch.safe_open(filename, framework="pt", device="cpu") as f:
|
||||
keys = f.keys()
|
||||
except Exception:
|
||||
pass
|
||||
return keys
|
||||
|
||||
|
||||
def get_modules(model: callable):
|
||||
signature = inspect.signature(model.__init__, follow_wrapped=True)
|
||||
params = {param.name: param.annotation for param in signature.parameters.values() if param.annotation != inspect._empty and hasattr(param.annotation, 'from_pretrained')} # pylint: disable=protected-access
|
||||
for name, cls in params.items():
|
||||
shared.log.debug(f'Analyze: model={model} module={name} class={cls.__name__} loadable={getattr(cls, "from_pretrained", None)}')
|
||||
return params
|
||||
|
||||
|
||||
def load_modules(repo_id: str, params: dict):
|
||||
cache_dir = shared.opts.hfcache_dir
|
||||
modules = {}
|
||||
for name, cls in params.items():
|
||||
subfolder = None
|
||||
kwargs = {}
|
||||
if cls == diffusers.AutoencoderKL:
|
||||
subfolder = 'vae'
|
||||
if cls == transformers.CLIPTextModel: # clip-vit-l
|
||||
subfolder = 'text_encoder'
|
||||
if cls == transformers.CLIPTextModelWithProjection: # clip-vit-g
|
||||
subfolder = 'text_encoder_2'
|
||||
if cls == transformers.T5EncoderModel: # t5-xxl
|
||||
subfolder = 'text_encoder_3'
|
||||
kwargs['quantization_config'] = model_quant.create_bnb_config()
|
||||
kwargs['variant'] = 'fp16'
|
||||
if cls == diffusers.SD3Transformer2DModel:
|
||||
subfolder = 'transformer'
|
||||
kwargs['quantization_config'] = model_quant.create_bnb_config()
|
||||
if subfolder is None:
|
||||
continue
|
||||
shared.log.debug(f'Load: module={name} class={cls.__name__} repo={repo_id} location={subfolder}')
|
||||
modules[name] = cls.from_pretrained(repo_id, subfolder=subfolder, cache_dir=cache_dir, torch_dtype=devices.dtype, **kwargs)
|
||||
return modules
|
||||
@@ -273,6 +273,7 @@ def load_diffusers_models(clear=True):
|
||||
place = os.path.join(models_path, 'Diffusers')
|
||||
if clear:
|
||||
diffuser_repos.clear()
|
||||
already_found = []
|
||||
try:
|
||||
for folder in os.listdir(place):
|
||||
try:
|
||||
@@ -303,7 +304,11 @@ def load_diffusers_models(clear=True):
|
||||
if (not os.path.exists(index)) and (not os.path.exists(info)) and (not os.path.exists(config)):
|
||||
debug(f'Diffusers skip model no info: {name}')
|
||||
continue
|
||||
if name in already_found:
|
||||
debug(f'Diffusers skip model already found: {name}')
|
||||
continue
|
||||
repo = { 'name': name, 'filename': name, 'friendly': friendly, 'folder': folder, 'path': commit, 'hash': snapshot, 'mtime': mtime, 'model_info': info, 'model_index': index, 'model_config': config }
|
||||
already_found.append(name)
|
||||
diffuser_repos.append(repo)
|
||||
if os.path.exists(os.path.join(folder, 'hidden')):
|
||||
continue
|
||||
@@ -318,7 +323,7 @@ def load_diffusers_models(clear=True):
|
||||
def find_diffuser(name: str, full=False):
|
||||
repo = [r for r in diffuser_repos if name == r['name'] or name == r['friendly'] or name == r['path']]
|
||||
if len(repo) > 0:
|
||||
return repo['name']
|
||||
return [repo[0]['name']]
|
||||
hf_api = hf.HfApi()
|
||||
models = list(hf_api.list_models(model_name=name, library=['diffusers'], full=True, limit=20, sort="downloads", direction=-1))
|
||||
shared.log.debug(f'Searching diffusers models: {name} {len(models) > 0}')
|
||||
|
||||
@@ -184,8 +184,8 @@ def preprocess_pipeline(p):
|
||||
return shared.sd_model
|
||||
|
||||
|
||||
def ORTDiffusionModelPart_to(self, *args, **kwargs):
|
||||
self.parent_model = self.parent_model.to(*args, **kwargs)
|
||||
def ORTPipelinePart_to(self, *args, **kwargs):
|
||||
self.parent_pipeline = self.parent_pipeline.to(*args, **kwargs)
|
||||
return self
|
||||
|
||||
|
||||
@@ -241,9 +241,9 @@ def initialize_onnx():
|
||||
diffusers.ORTStableDiffusionXLPipeline = diffusers.OnnxStableDiffusionXLPipeline # Huggingface model compatibility
|
||||
diffusers.ORTStableDiffusionXLImg2ImgPipeline = diffusers.OnnxStableDiffusionXLImg2ImgPipeline
|
||||
|
||||
optimum.onnxruntime.modeling_diffusion._ORTDiffusionModelPart.to = ORTDiffusionModelPart_to # pylint: disable=protected-access
|
||||
except Exception:
|
||||
pass
|
||||
optimum.onnxruntime.modeling_diffusion.ORTPipelinePart.to = ORTPipelinePart_to # pylint: disable=protected-access
|
||||
except Exception as e:
|
||||
log.debug(f'ONNX failed to initialize XL pipelines: {e}')
|
||||
|
||||
initialized = True
|
||||
|
||||
|
||||
@@ -15,7 +15,7 @@ def create_ui():
|
||||
from modules.ui_common import create_refresh_button
|
||||
from modules.ui_components import DropdownMulti
|
||||
from modules.shared import log, opts, cmd_opts, refresh_checkpoints
|
||||
from modules.sd_models import checkpoint_tiles, get_closet_checkpoint_match
|
||||
from modules.sd_models import checkpoint_titles, get_closet_checkpoint_match
|
||||
from modules.paths import sd_configs_path
|
||||
from .execution_providers import ExecutionProvider, install_execution_provider
|
||||
from .utils import check_diffusers_cache
|
||||
@@ -46,7 +46,7 @@ def create_ui():
|
||||
with gr.TabItem("Manage cache", id="manage_cache"):
|
||||
cache_state_dirname = gr.Textbox(value=None, visible=False)
|
||||
with gr.Row():
|
||||
model_dropdown = gr.Dropdown(label="Model", value="Please select model", choices=checkpoint_tiles())
|
||||
model_dropdown = gr.Dropdown(label="Model", value="Please select model", choices=checkpoint_titles())
|
||||
create_refresh_button(model_dropdown, refresh_checkpoints, {}, "onnx_cache_refresh_diffusers_model")
|
||||
with gr.Row():
|
||||
def remove_cache_onnx_converted(dirname: str):
|
||||
|
||||
@@ -18,14 +18,13 @@ PREDEFINED = [ # <https://huggingface.co/vladmandic/yolo-detailers/tree/main>
|
||||
|
||||
|
||||
class YoloResult:
|
||||
def __init__(self, cls: int, label: str, score: float, box: list[int], mask: Image.Image = None, item: Image.Image = None, size: float = 0, width = 0, height = 0, args = {}):
|
||||
def __init__(self, cls: int, label: str, score: float, box: list[int], mask: Image.Image = None, item: Image.Image = None, width = 0, height = 0, args = {}):
|
||||
self.cls = cls
|
||||
self.label = label
|
||||
self.score = score
|
||||
self.box = box
|
||||
self.mask = mask
|
||||
self.item = item
|
||||
self.size = size
|
||||
self.width = width
|
||||
self.height = height
|
||||
self.args = args
|
||||
@@ -78,6 +77,10 @@ class YoloRestorer(Detailer):
|
||||
) -> list[YoloResult]:
|
||||
|
||||
result = []
|
||||
if isinstance(model, str):
|
||||
model = self.models.get(model, None)
|
||||
if model is None:
|
||||
_, model = self.load(model)
|
||||
if model is None:
|
||||
return result
|
||||
args = {
|
||||
@@ -123,21 +126,24 @@ class YoloRestorer(Detailer):
|
||||
box = box.tolist()
|
||||
mask_image = None
|
||||
w, h = box[2] - box[0], box[3] - box[1]
|
||||
size = w * h / (image.width * image.height)
|
||||
if (min(w, h) > shared.opts.detailer_min_size if shared.opts.detailer_min_size > 0 else True) and (max(w, h) < shared.opts.detailer_max_size if shared.opts.detailer_max_size > 0 else True):
|
||||
x_size, y_size = w/image.width, h/image.height
|
||||
min_size = shared.opts.detailer_min_size if shared.opts.detailer_min_size > 0 and shared.opts.detailer_min_size < 1 else 0
|
||||
max_size = shared.opts.detailer_max_size if shared.opts.detailer_max_size > 0 and shared.opts.detailer_max_size < 1 else 1
|
||||
if x_size >= min_size and y_size >=min_size and x_size <= max_size and y_size <= max_size:
|
||||
if mask:
|
||||
mask_image = image.copy()
|
||||
mask_image = Image.new('L', image.size, 0)
|
||||
draw = ImageDraw.Draw(mask_image)
|
||||
draw.rectangle(box, fill="white", outline=None, width=0)
|
||||
cropped = image.crop(box)
|
||||
result.append(YoloResult(cls=cls, label=label, score=round(score, 2), box=box, mask=mask_image, item=cropped, size=size, width=w, height=h, args=args))
|
||||
result.append(YoloResult(cls=cls, label=label, score=round(score, 2), box=box, mask=mask_image, item=cropped, width=w, height=h, args=args))
|
||||
if len(result) >= shared.opts.detailer_max:
|
||||
break
|
||||
return result
|
||||
|
||||
def load(self, model_name: str = None):
|
||||
from modules import modelloader
|
||||
model = None
|
||||
self.dependencies()
|
||||
if model_name is None:
|
||||
model_name = list(self.list)[0]
|
||||
@@ -150,15 +156,15 @@ class YoloRestorer(Detailer):
|
||||
try:
|
||||
model_file = modelloader.load_file_from_url(url=model_url, model_dir=shared.opts.yolo_dir, file_name=file_name)
|
||||
if model_file is not None:
|
||||
from ultralytics import YOLO # pylint: disable=import-outside-toplevel
|
||||
model = YOLO(model_file)
|
||||
import ultralytics
|
||||
model = ultralytics.YOLO(model_file)
|
||||
classes = list(model.names.values())
|
||||
shared.log.info(f'Load: type=Detailer name="{model_name}" model="{model_file}" classes={classes}')
|
||||
shared.log.info(f'Load: type=Detailer name="{model_name}" model="{model_file}" ultralytics={ultralytics.__version__} classes={classes}')
|
||||
self.models[model_name] = model
|
||||
return model_name, model
|
||||
except Exception as e:
|
||||
shared.log.error(f'Load: type=Detailer name="{model_name}" error="{e}"')
|
||||
return None
|
||||
return None, None
|
||||
|
||||
def restore(self, np_image, p: processing.StableDiffusionProcessing = None):
|
||||
if hasattr(p, 'recursion'):
|
||||
@@ -191,12 +197,20 @@ class YoloRestorer(Detailer):
|
||||
pp = None
|
||||
shared.opts.data['mask_apply_overlay'] = True
|
||||
resolution = 512 if shared.sd_model_type in ['none', 'sd', 'lcm', 'unknown'] else 1024
|
||||
orig_prompt: str = orig_p.get('all_prompts', [''])[0]
|
||||
orig_negative: str = orig_p.get('all_negative_prompts', [''])[0]
|
||||
prompt: str = orig_p.get('refiner_prompt', '')
|
||||
negative: str = orig_p.get('refiner_negative', '')
|
||||
if len(prompt) == 0:
|
||||
prompt = orig_p.get('all_prompts', [''])[0]
|
||||
prompt = orig_prompt
|
||||
else:
|
||||
prompt = prompt.replace('[PROMPT]', orig_prompt)
|
||||
prompt = prompt.replace('[prompt]', orig_prompt)
|
||||
if len(negative) == 0:
|
||||
negative = orig_p.get('all_negative_prompts', [''])[0]
|
||||
negative = orig_negative
|
||||
else:
|
||||
negative = negative.replace('[PROMPT]', orig_negative)
|
||||
negative = negative.replace('[prompt]', orig_negative)
|
||||
prompt_lines = prompt.split('\n')
|
||||
negative_lines = negative.split('\n')
|
||||
prompt = prompt_lines[i % len(prompt_lines)]
|
||||
@@ -315,8 +329,10 @@ class YoloRestorer(Detailer):
|
||||
min_confidence = gr.Slider(label="Min confidence", elem_id=f"{tab}_detailer_conf", value=shared.opts.detailer_conf, minimum=0.0, maximum=1.0, step=0.05)
|
||||
iou = gr.Slider(label="Max overlap", elem_id=f"{tab}_detailer_iou", value=shared.opts.detailer_iou, minimum=0, maximum=1.0, step=0.05)
|
||||
with gr.Row():
|
||||
min_size = gr.Slider(label="Min size", elem_id=f"{tab}_detailer_min_size", value=shared.opts.detailer_min_size, minimum=0, maximum=1024, step=1)
|
||||
max_size = gr.Slider(label="Max size", elem_id=f"{tab}_detailer_max_size", value=shared.opts.detailer_max_size, minimum=0, maximum=1024, step=1)
|
||||
min_size = shared.opts.detailer_min_size if shared.opts.detailer_min_size < 1 else 0.0
|
||||
min_size = gr.Slider(label="Min size", elem_id=f"{tab}_detailer_min_size", value=min_size, minimum=0.0, maximum=1.0, step=0.05)
|
||||
max_size = shared.opts.detailer_max_size if shared.opts.detailer_max_size < 1 and shared.opts.detailer_max_size > 0 else 1.0
|
||||
max_size = gr.Slider(label="Max size", elem_id=f"{tab}_detailer_max_size", value=max_size, minimum=0.0, maximum=1.0, step=0.05)
|
||||
detailers.change(fn=ui_settings_change, inputs=[detailers, classes, strength, padding, blur, min_confidence, max_detected, min_size, max_size, iou], outputs=[])
|
||||
classes.change(fn=ui_settings_change, inputs=[detailers, classes, strength, padding, blur, min_confidence, max_detected, min_size, max_size, iou], outputs=[])
|
||||
strength.change(fn=ui_settings_change, inputs=[detailers, classes, strength, padding, blur, min_confidence, max_detected, min_size, max_size, iou], outputs=[])
|
||||
|
||||
@@ -37,7 +37,6 @@ def run_postprocessing(extras_mode, image, image_folder: List[tempfile.NamedTemp
|
||||
image_ext.append(ext)
|
||||
shared.log.debug(f'Process: mode=batch inputs={len(image_folder)} images={len(image_data)}')
|
||||
elif extras_mode == 2:
|
||||
assert not shared.cmd_opts.hide_ui_dir_config, '--hide-ui-dir-config option must be disabled'
|
||||
assert input_dir, 'input directory not selected'
|
||||
image_list = os.listdir(input_dir)
|
||||
for filename in image_list:
|
||||
|
||||
@@ -4,7 +4,7 @@ import time
|
||||
from contextlib import nullcontext
|
||||
import numpy as np
|
||||
from PIL import Image, ImageOps
|
||||
from modules import shared, devices, errors, images, scripts, memstats, lowvram, script_callbacks, extra_networks, detailer, sd_hijack_freeu, sd_models, sd_vae, processing_helpers, timer, face_restoration
|
||||
from modules import shared, devices, errors, images, scripts, memstats, lowvram, script_callbacks, extra_networks, detailer, sd_hijack_freeu, sd_models, sd_checkpoint, sd_vae, processing_helpers, timer, face_restoration, token_merge
|
||||
from modules.sd_hijack_hypertile import context_hypertile_vae, context_hypertile_unet
|
||||
from modules.processing_class import StableDiffusionProcessing, StableDiffusionProcessingTxt2Img, StableDiffusionProcessingImg2Img, StableDiffusionProcessingControl # pylint: disable=unused-import
|
||||
from modules.processing_info import create_infotext
|
||||
@@ -34,19 +34,20 @@ images_tensor_to_samples = processing_helpers.images_tensor_to_samples
|
||||
|
||||
|
||||
class Processed:
|
||||
def __init__(self, p: StableDiffusionProcessing, images_list, seed=-1, info="", subseed=None, all_prompts=None, all_negative_prompts=None, all_seeds=None, all_subseeds=None, index_of_first_image=0, infotexts=None, comments=""):
|
||||
def __init__(self, p: StableDiffusionProcessing, images_list, seed=-1, info=None, subseed=None, all_prompts=None, all_negative_prompts=None, all_seeds=None, all_subseeds=None, index_of_first_image=0, infotexts=None, comments=""):
|
||||
self.images = images_list
|
||||
self.prompt = p.prompt or ''
|
||||
self.negative_prompt = p.negative_prompt or ''
|
||||
self.seed = seed if seed != -1 else p.seed
|
||||
self.subseed = subseed
|
||||
self.subseed_strength = p.subseed_strength
|
||||
self.info = info
|
||||
self.info = info or create_infotext(p)
|
||||
self.comments = comments or ''
|
||||
self.width = p.width if hasattr(p, 'width') else (self.images[0].width if len(self.images) > 0 else 0)
|
||||
self.height = p.height if hasattr(p, 'height') else (self.images[0].height if len(self.images) > 0 else 0)
|
||||
self.sampler_name = p.sampler_name or ''
|
||||
self.cfg_scale = p.cfg_scale or 0
|
||||
self.cfg_scale = p.cfg_scale if p.cfg_scale > 1 else None
|
||||
self.cfg_end = p.cfg_end if p.cfg_end < 0 else None
|
||||
self.image_cfg_scale = p.image_cfg_scale or 0
|
||||
self.steps = p.steps or 0
|
||||
self.batch_size = max(1, p.batch_size)
|
||||
@@ -79,7 +80,7 @@ class Processed:
|
||||
self.all_negative_prompts = all_negative_prompts or p.all_negative_prompts or [self.negative_prompt]
|
||||
self.all_seeds = all_seeds or p.all_seeds or [self.seed]
|
||||
self.all_subseeds = all_subseeds or p.all_subseeds or [self.subseed]
|
||||
self.infotexts = infotexts or [info]
|
||||
self.infotexts = infotexts or [self.info]
|
||||
|
||||
def js(self):
|
||||
obj = {
|
||||
@@ -96,6 +97,7 @@ class Processed:
|
||||
"height": self.height,
|
||||
"sampler_name": self.sampler_name,
|
||||
"cfg_scale": self.cfg_scale,
|
||||
"cfg_end": self.cfg_end,
|
||||
"steps": self.steps,
|
||||
"batch_size": self.batch_size,
|
||||
"detailer": self.detailer,
|
||||
@@ -136,11 +138,11 @@ def process_images(p: StableDiffusionProcessing) -> Processed:
|
||||
processed = None
|
||||
try:
|
||||
# if no checkpoint override or the override checkpoint can't be found, remove override entry and load opts checkpoint
|
||||
if p.override_settings.get('sd_model_checkpoint', None) is not None and sd_models.checkpoint_aliases.get(p.override_settings.get('sd_model_checkpoint')) is None:
|
||||
if p.override_settings.get('sd_model_checkpoint', None) is not None and sd_checkpoint.checkpoint_aliases.get(p.override_settings.get('sd_model_checkpoint')) is None:
|
||||
shared.log.warning(f"Override not found: checkpoint={p.override_settings.get('sd_model_checkpoint', None)}")
|
||||
p.override_settings.pop('sd_model_checkpoint', None)
|
||||
sd_models.reload_model_weights()
|
||||
if p.override_settings.get('sd_model_refiner', None) is not None and sd_models.checkpoint_aliases.get(p.override_settings.get('sd_model_refiner')) is None:
|
||||
if p.override_settings.get('sd_model_refiner', None) is not None and sd_checkpoint.checkpoint_aliases.get(p.override_settings.get('sd_model_refiner')) is None:
|
||||
shared.log.warning(f"Override not found: refiner={p.override_settings.get('sd_model_refiner', None)}")
|
||||
p.override_settings.pop('sd_model_refiner', None)
|
||||
sd_models.reload_model_weights()
|
||||
@@ -162,7 +164,7 @@ def process_images(p: StableDiffusionProcessing) -> Processed:
|
||||
shared.prompt_styles.apply_styles_to_extra(p)
|
||||
shared.prompt_styles.extract_comments(p)
|
||||
if shared.opts.cuda_compile_backend == 'none':
|
||||
sd_models.apply_token_merging(p.sd_model)
|
||||
token_merge.apply_token_merging(p.sd_model)
|
||||
sd_hijack_freeu.apply_freeu(p, not shared.native)
|
||||
|
||||
if p.width is not None:
|
||||
@@ -205,7 +207,7 @@ def process_images(p: StableDiffusionProcessing) -> Processed:
|
||||
finally:
|
||||
pag.unapply()
|
||||
if shared.opts.cuda_compile_backend == 'none':
|
||||
sd_models.remove_token_merging(p.sd_model)
|
||||
token_merge.remove_token_merging(p.sd_model)
|
||||
|
||||
script_callbacks.after_process_callback(p)
|
||||
|
||||
|
||||
@@ -100,10 +100,11 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2
|
||||
if hasattr(model, "set_progress_bar_config"):
|
||||
model.set_progress_bar_config(bar_format='Progress {rate_fmt}{postfix} {bar} {percentage:3.0f}% {n_fmt}/{total_fmt} {elapsed} {remaining} ' + '\x1b[38;5;71m' + desc, ncols=80, colour='#327fba')
|
||||
args = {}
|
||||
if hasattr(model, 'pipe'): # recurse
|
||||
if hasattr(model, 'pipe') and not hasattr(model, 'no_recurse'): # recurse
|
||||
model = model.pipe
|
||||
signature = inspect.signature(type(model).__call__, follow_wrapped=True)
|
||||
possible = list(signature.parameters)
|
||||
|
||||
debug(f'Diffusers pipeline possible: {possible}')
|
||||
prompts, negative_prompts, prompts_2, negative_prompts_2 = fix_prompts(prompts, negative_prompts, prompts_2, negative_prompts_2)
|
||||
parser = 'Fixed attention'
|
||||
@@ -128,7 +129,7 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2
|
||||
if 'prompt' in possible:
|
||||
if 'OmniGen' in model.__class__.__name__:
|
||||
prompts = [p.replace('|image|', '<|image_1|>') for p in prompts]
|
||||
if hasattr(model, 'text_encoder') and 'prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None:
|
||||
if hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'prompt_embeds' in possible and len(p.prompt_embeds) > 0 and p.prompt_embeds[0] is not None:
|
||||
args['prompt_embeds'] = prompt_parser_diffusers.embedder('prompt_embeds')
|
||||
if 'StableCascade' in model.__class__.__name__ and len(getattr(p, 'negative_pooleds', [])) > 0:
|
||||
args['prompt_embeds_pooled'] = prompt_parser_diffusers.embedder('positive_pooleds').unsqueeze(0)
|
||||
@@ -141,7 +142,7 @@ def set_pipeline_args(p, model, prompts: list, negative_prompts: list, prompts_2
|
||||
else:
|
||||
args['prompt'] = prompts
|
||||
if 'negative_prompt' in possible:
|
||||
if hasattr(model, 'text_encoder') and 'negative_prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None:
|
||||
if hasattr(model, 'text_encoder') and hasattr(model, 'tokenizer') and 'negative_prompt_embeds' in possible and prompt_parser_diffusers.embedder is not None:
|
||||
args['negative_prompt_embeds'] = prompt_parser_diffusers.embedder('negative_prompt_embeds')
|
||||
if 'StableCascade' in model.__class__.__name__ and prompt_parser_diffusers.embedder is not None:
|
||||
args['negative_prompt_embeds_pooled'] = prompt_parser_diffusers.embedder('negative_pooleds').unsqueeze(0)
|
||||
|
||||
@@ -32,7 +32,7 @@ def diffusers_callback_legacy(step: int, timestep: int, latents: typing.Union[to
|
||||
time.sleep(0.1)
|
||||
|
||||
|
||||
def diffusers_callback(pipe, step: int, timestep: int, kwargs: dict):
|
||||
def diffusers_callback(pipe, step: int = 0, timestep: int = 0, kwargs: dict = {}):
|
||||
t0 = time.time()
|
||||
if p is None:
|
||||
return kwargs
|
||||
|
||||