mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
Merge branch 'dev' into master
This commit is contained in:
@@ -16,6 +16,7 @@ ignore-paths=/usr/lib/.*$,
|
||||
^modules/dml/.*$,
|
||||
^modules/models/diffusion/.*$,
|
||||
^modules/xadapters/.*$,
|
||||
^modules/tcd/.*$,
|
||||
ignore-patterns=
|
||||
ignored-modules=
|
||||
jobs=0
|
||||
|
||||
@@ -56,19 +56,6 @@ For screenshots and informations on other available themes, see [Themes Wiki](ht
|
||||
Supports **SD 1.x** and **SD 2.x** models
|
||||
All other model types such as *SD-XL, LCM, PixArt, Segmind, Kandinsky, etc.* require backend **Diffusers**
|
||||
|
||||
## Control
|
||||
|
||||
**SD.Next** comes with built-in control for all types of text2image, image2image, video2video and batch processing
|
||||
|
||||
*Control interface*:
|
||||

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

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

|
||||
|
||||
## Model support
|
||||
|
||||
Additional models will be added as they become available and there is public interest in them
|
||||
@@ -110,7 +97,6 @@ Also supported are modifiers such as:
|
||||
*InstantID*:
|
||||

|
||||
|
||||
|
||||
> [!IMPORTANT]
|
||||
> - Loading any model other than standard SD 1.x / SD 2.x requires use of backend **Diffusers**
|
||||
> - Loading any other models using **Original** backend is not supported
|
||||
@@ -151,51 +137,89 @@ 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*)
|
||||
|
||||
Below is partial list of all available parameters, run `webui --help` for the full list:
|
||||
List of available parameters, 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
|
||||
--share Enable UI accessible through Gradio site, default: False
|
||||
--insecure Enable extensions tab regardless of other options, default: False
|
||||
--listen Launch web server using public IP address, default: False
|
||||
--auth AUTH Set access authentication like "user:pwd,user:pwd""
|
||||
--autolaunch Open the UI URL in the system's default browser upon launch
|
||||
--docs Mount API docs, default: False
|
||||
--no-hashing Disable hashing of checkpoints, default: False
|
||||
--no-metadata Disable reading of metadata from models, default: False
|
||||
--backend {original,diffusers} force model pipeline type
|
||||
--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:
|
||||
--debug Run installer with debug logging, default: False
|
||||
--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 Run with startup sequence only, 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-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
|
||||
--experimental Allow unsupported versions of libraries, default: False
|
||||
--reinstall Force reinstallation of all requirements, default: False
|
||||
--safe Run in safe mode with no user extensions
|
||||
--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
|
||||
|
||||
Logging options:
|
||||
--log LOG Set log file, default: None
|
||||
--debug Run installer with debug logging, default: False
|
||||
--profile Run profiler, default: False
|
||||
|
||||
## Notes
|
||||
|
||||
### Control
|
||||
|
||||
**SD.Next** comes with built-in control for all types of text2image, image2image, video2video and batch processing
|
||||
|
||||
*Control interface*:
|
||||

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

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

|
||||
|
||||
### **Extensions**
|
||||
|
||||
SD.Next comes with several extensions pre-installed:
|
||||
|
||||
@@ -5,18 +5,14 @@ Main ToDo list can be found at [GitHub projects](https://github.com/users/vladma
|
||||
## Candidates for next release
|
||||
|
||||
- defork
|
||||
- stable cascade: <https://github.com/vladmandic/automatic/wiki/Stable-Cascade>
|
||||
- stable diffusion 3.0
|
||||
- ipadapter masking: <https://github.com/huggingface/diffusers/pull/6847>
|
||||
- init latents: variations, tiling, img2img
|
||||
- x-adapter: <https://github.com/showlab/X-Adapter>
|
||||
- diffusers public callbacks
|
||||
- image2video: pia and vgen pipelines
|
||||
- video2video
|
||||
- async lowvram: <https://github.com/AUTOMATIC1111/stable-diffusion-webui/pull/14855>
|
||||
- init latents: variations, tiling, img2img
|
||||
- diffusers public callbacks
|
||||
- remove builtin: controlnet
|
||||
- remove builtin: image-browser
|
||||
- remove training: ti
|
||||
- remove training: hypernetwork
|
||||
|
||||
## Control missing features
|
||||
|
||||
|
||||
Executable
+57
@@ -0,0 +1,57 @@
|
||||
#!/usr/bin/env python
|
||||
import os
|
||||
import time
|
||||
import base64
|
||||
import logging
|
||||
import argparse
|
||||
import requests
|
||||
import urllib3
|
||||
|
||||
|
||||
sd_url = os.environ.get('SDAPI_URL', "http://127.0.0.1:7860")
|
||||
sd_username = os.environ.get('SDAPI_USR', None)
|
||||
sd_password = os.environ.get('SDAPI_PWD', None)
|
||||
|
||||
|
||||
logging.basicConfig(level = logging.INFO, format = '%(asctime)s %(levelname)s: %(message)s')
|
||||
log = logging.getLogger(__name__)
|
||||
urllib3.disable_warnings(urllib3.exceptions.InsecureRequestWarning)
|
||||
|
||||
|
||||
def auth():
|
||||
if sd_username is not None and sd_password is not None:
|
||||
return requests.auth.HTTPBasicAuth(sd_username, sd_password)
|
||||
return None
|
||||
|
||||
|
||||
def get(endpoint: str, dct: dict = None):
|
||||
req = requests.get(f'{sd_url}{endpoint}', json=dct, timeout=300, verify=False, auth=auth())
|
||||
if req.status_code != 200:
|
||||
return { 'error': req.status_code, 'reason': req.reason, 'url': req.url }
|
||||
else:
|
||||
return req.json()
|
||||
|
||||
|
||||
def post(endpoint: str, dct: dict = None):
|
||||
req = requests.post(f'{sd_url}{endpoint}', json = dct, timeout=300, verify=False, auth=auth())
|
||||
if req.status_code != 200:
|
||||
return { 'error': req.status_code, 'reason': req.reason, 'url': req.url }
|
||||
else:
|
||||
return req.json()
|
||||
|
||||
|
||||
def info(args): # pylint: disable=redefined-outer-name
|
||||
t0 = time.time()
|
||||
with open(args.input, 'rb') as f:
|
||||
content = f.read()
|
||||
data = post('/sdapi/v1/png-info', { 'image': base64.b64encode(content).decode() })
|
||||
t1 = time.time()
|
||||
log.info(f'received: {data} time={t1-t0:.2f}')
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description = 'simple-info')
|
||||
parser.add_argument('--input', required=True, help='input image')
|
||||
args = parser.parse_args()
|
||||
log.info(f'info: {args}')
|
||||
info(args)
|
||||
@@ -0,0 +1,43 @@
|
||||
{
|
||||
"_class_name": "AutoencoderKL",
|
||||
"_diffusers_version": "0.27.0.dev0",
|
||||
"act_fn": "silu",
|
||||
"block_out_channels": [
|
||||
128,
|
||||
256,
|
||||
512,
|
||||
512
|
||||
],
|
||||
"down_block_types": [
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D",
|
||||
"DownEncoderBlock2D"
|
||||
],
|
||||
"force_upcast": true,
|
||||
"in_channels": 3,
|
||||
"latent_channels": 4,
|
||||
"layers_per_block": 2,
|
||||
"norm_num_groups": 32,
|
||||
"out_channels": 3,
|
||||
"sample_size": 1024,
|
||||
"up_block_types": [
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D",
|
||||
"UpDecoderBlock2D"
|
||||
],
|
||||
"latents_mean": [
|
||||
-1.6574,
|
||||
1.886,
|
||||
-1.383,
|
||||
2.5155
|
||||
],
|
||||
"latents_std": [
|
||||
8.4927,
|
||||
5.9022,
|
||||
6.5498,
|
||||
5.2299
|
||||
],
|
||||
"scaling_factor": 0.5
|
||||
}
|
||||
+1
-1
@@ -133,7 +133,7 @@
|
||||
{"id":"","label":"Refiner start","localized":"","hint":"Refiner pass will start when base model is this much complete (set to 0 or 1 to run after full base model run)"},
|
||||
{"id":"","label":"Refiner steps","localized":"","hint":"Number of steps to use for refiner pass"},
|
||||
{"id":"","label":"Secondary CFG Scale","localized":"","hint":"CFG scale used for refiner pass"},
|
||||
{"id":"","label":"Guidance rescale","localized":"","hint":"Rescale CFG generated noise to avoid overexposed images"},
|
||||
{"id":"","label":"Rescale guidance","localized":"","hint":"Rescale CFG generated noise to avoid overexposed images"},
|
||||
{"id":"","label":"Secondary Prompt","localized":"","hint":"Prompt used for both second encoder in base model (if it exists) and for refiner pass (if enabled)"},
|
||||
{"id":"","label":"Secondary negative prompt","localized":"","hint":"Negative prompt used for both second encoder in base model (if it exists) and for refiner pass (if enabled)"},
|
||||
{"id":"","label":"Width","localized":"","hint":"Image width"},
|
||||
|
||||
@@ -169,6 +169,11 @@
|
||||
"desc": "Playground v2 is a diffusion-based text-to-image generative model. The model was trained from scratch by the research team at Playground. Images generated by Playground v2 are favored 2.5 times more than those produced by Stable Diffusion XL, according to Playground’s user study.",
|
||||
"preview": "playgroundai--playground-v2-1024px-aesthetic.jpg"
|
||||
},
|
||||
"Playground v2.5": {
|
||||
"path": "playground-v2.5-1024px-aesthetic.fp16.safetensors@https://huggingface.co/playgroundai/playground-v2.5-1024px-aesthetic/resolve/main/playground-v2.5-1024px-aesthetic.fp16.safetensors?download=true",
|
||||
"desc": "Playground v2.5 is a diffusion-based text-to-image generative model, and a successor to Playground v2. Playground v2.5 is the state-of-the-art open-source model in aesthetic quality. Our user studies demonstrate that our model outperforms SDXL, Playground v2, PixArt-α, DALL-E 3, and Midjourney 5.2.",
|
||||
"preview": "playgroundai--playground-v2-1024px-aesthetic.jpg"
|
||||
},
|
||||
"DeepFloyd IF Medium": {
|
||||
"path": "DeepFloyd/IF-I-M-v1.0",
|
||||
"desc": "DeepFloyd-IF is a pixel-based text-to-image triple-cascaded diffusion model, that can generate pictures with new state-of-the-art for photorealism and language understanding. The result is a highly efficient model that outperforms current state-of-the-art models, achieving a zero-shot FID-30K score of 6.66 on the COCO dataset. It is modular and composed of frozen text mode and three pixel cascaded diffusion modules, each designed to generate images of increasing resolution: 64x64, 256x256, and 1024x1024.",
|
||||
@@ -184,6 +189,14 @@
|
||||
"desc": "Amused is a lightweight text to image model based off of the muse architecture. Amused is particularly useful in applications that require a lightweight and fast model such as generating many images quickly at once.",
|
||||
"preview": "amused--amused-512.jpg"
|
||||
},
|
||||
"KOALA 700M": {
|
||||
"path": "huggingface/etri-vilab/koala-700m-llava-cap",
|
||||
"variant": "fp16",
|
||||
"skip": true,
|
||||
"desc": "Fast text-to-image model, called KOALA, by compressing SDXL's U-Net and distilling knowledge from SDXL into our model. KOALA-700M can generate a 1024x1024 image in less than 1.5 seconds on an NVIDIA 4090 GPU, which is more than 2x faster than SDXL.",
|
||||
"preview": "etri-vilab--koala-700m-llava-cap.jpg"
|
||||
},
|
||||
|
||||
"Tsinghua UniDiffuser": {
|
||||
"path": "thu-ml/unidiffuser-v1",
|
||||
"desc": "UniDiffuser is a unified diffusion framework to fit all distributions relevant to a set of multi-modal data in one transformer. UniDiffuser is able to perform image, text, text-to-image, image-to-text, and image-text pair generation by setting proper timesteps without additional overhead.\nSpecifically, UniDiffuser employs a variation of transformer, called U-ViT, which parameterizes the joint noise prediction network. Other components perform as encoders and decoders of different modalities, including a pretrained image autoencoder from Stable Diffusion, a pretrained image ViT-B/32 CLIP encoder, a pretrained text ViT-L CLIP encoder, and a GPT-2 text decoder finetuned by ourselves.",
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
@font-face { font-family: 'NotoSans'; font-display: swap; font-style: normal; font-weight: 100; src: local('NotoSans'), url('notosans-nerdfont-regular.ttf') }
|
||||
:root { --left-column: 550px; }
|
||||
:root { --left-column: 520px; }
|
||||
a { font-weight: bold; cursor: pointer; }
|
||||
h2 { margin-top: 1em !important; font-size: var(--text-xxl) !important; }
|
||||
footer { display: none; margin-top: 0 !important;}
|
||||
@@ -45,7 +45,7 @@ input[type='color'] { width: 64px; height: 32px; }
|
||||
.gradio-textbox { overflow: visible !important; }
|
||||
.gradio-radio { padding: 0 !important; width: max-content !important; }
|
||||
.gradio-slider { margin-right: var(--spacing-sm) !important; width: max-content !important }
|
||||
.gradio-slider input[type="number"] { width: 6em; font-size: var(--text-xs); height: 16px; text-align: right; }
|
||||
.gradio-slider input[type="number"] { width: 5em; font-size: var(--text-xs); height: 16px; text-align: right; padding: 0; }
|
||||
|
||||
/* custom gradio elements */
|
||||
.accordion-compact { padding: 8px 0px 4px 0px !important; }
|
||||
@@ -107,8 +107,8 @@ div#extras_scale_to_tab div.form{ flex-direction: row; }
|
||||
#mode_img2img .gradio-image>div.fixed-height, #mode_img2img .gradio-image>div.fixed-height img{ height: 480px !important; max-height: 480px !important; min-height: 480px !important; }
|
||||
#img2img_sketch, #img2maskimg, #inpaint_sketch { overflow: overlay !important; resize: auto; background: var(--panel-background-fill); z-index: 5; }
|
||||
.image-buttons button{ min-width: auto; }
|
||||
.infotext { overflow-wrap: break-word; line-height: 1.5em; }
|
||||
.infotext>p { padding-left: 1em; text-indent: -1em; white-space: pre-wrap; }
|
||||
.infotext { overflow-wrap: break-word; line-height: 1.5em; font-size: 0.95em !important; }
|
||||
.infotext > p { padding-left: 1em; text-indent: -1em; white-space: pre-wrap; color: var(--block-info-text-color) !important; }
|
||||
.tooltip { display: block; position: fixed; top: 1em; right: 1em; padding: 0.5em; background: var(--input-background-fill); color: var(--body-text-color); border: 1pt solid var(--button-primary-border-color);
|
||||
width: 22em; min-height: 1.3em; font-size: var(--text-xs); transition: opacity 0.2s ease-in; pointer-events: none; opacity: 0; z-index: 999; }
|
||||
.tooltip-show { opacity: 0.9; }
|
||||
@@ -301,9 +301,10 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt
|
||||
--spacing-xxs: 1px;
|
||||
--spacing-xs: 2px;
|
||||
--spacing-sm: 3px;
|
||||
--spacing-lg: 4px;
|
||||
--spacing-xl: 5px;
|
||||
--spacing-xxl: 6px;
|
||||
--spacing-md: 4px;
|
||||
--spacing-lg: 5px;
|
||||
--spacing-xl: 6px;
|
||||
--spacing-xxl: 7px;
|
||||
}
|
||||
|
||||
@media (hover: none) and (pointer: coarse) { /* Apply different styles for devices with coarse pointers dependant on screen resolution */
|
||||
|
||||
+5
-5
@@ -241,11 +241,11 @@ function clearPrompts(prompt, negative_prompt) {
|
||||
return [prompt, negative_prompt];
|
||||
}
|
||||
|
||||
const promptTokecountUpdateFuncs = {};
|
||||
const promptTokenCountUpdateFuncs = {};
|
||||
|
||||
function recalculatePromptTokens(name) {
|
||||
if (promptTokecountUpdateFuncs[name]) {
|
||||
promptTokecountUpdateFuncs[name]();
|
||||
if (promptTokenCountUpdateFuncs[name]) {
|
||||
promptTokenCountUpdateFuncs[name]();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -330,8 +330,8 @@ onAfterUiUpdate(async () => {
|
||||
if (counter.parentElement === prompt.parentElement) return;
|
||||
prompt.parentElement.insertBefore(counter, prompt);
|
||||
prompt.parentElement.style.position = 'relative';
|
||||
promptTokecountUpdateFuncs[id] = () => { update_token_counter(id_button); };
|
||||
localTextarea.addEventListener('input', promptTokecountUpdateFuncs[id]);
|
||||
promptTokenCountUpdateFuncs[id] = () => { update_token_counter(id_button); };
|
||||
localTextarea.addEventListener('input', promptTokenCountUpdateFuncs[id]);
|
||||
if (!promptsInitialized) log('initPrompts');
|
||||
promptsInitialized = true;
|
||||
}
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 39 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 51 KiB |
+1
-8
@@ -5,7 +5,7 @@ from fastapi import FastAPI, APIRouter, Depends, Request
|
||||
from fastapi.security import HTTPBasic, HTTPBasicCredentials
|
||||
from fastapi.exceptions import HTTPException
|
||||
from modules import errors, shared, postprocessing
|
||||
from modules.api import models, endpoints, script, train, helpers, server, nvml, generate
|
||||
from modules.api import models, endpoints, script, helpers, server, nvml, generate
|
||||
|
||||
|
||||
errors.install()
|
||||
@@ -77,13 +77,6 @@ class Api:
|
||||
self.add_api_route("/sdapi/v1/reload-checkpoint", endpoints.post_reload_checkpoint, methods=["POST"])
|
||||
self.add_api_route("/sdapi/v1/refresh-vae", endpoints.post_refresh_vae, methods=["POST"])
|
||||
|
||||
# train api
|
||||
self.add_api_route("/sdapi/v1/create/embedding", train.post_create_embedding, methods=["POST"], response_model=models.ResCreate)
|
||||
self.add_api_route("/sdapi/v1/create/hypernetwork", train.post_create_hypernetwork, methods=["POST"], response_model=models.ResCreate)
|
||||
self.add_api_route("/sdapi/v1/preprocess", train.post_preprocess, methods=["POST"], response_model=models.ResPreprocess)
|
||||
self.add_api_route("/sdapi/v1/train/embedding", train.post_train_embedding, methods=["POST"], response_model=models.ResTrain)
|
||||
self.add_api_route("/sdapi/v1/train/hypernetwork", train.post_train_hypernetwork, methods=["POST"], response_model=models.ResTrain)
|
||||
|
||||
def add_api_route(self, path: str, endpoint, **kwargs):
|
||||
if (shared.cmd_opts.auth or shared.cmd_opts.auth_file) and shared.cmd_opts.api_only:
|
||||
return self.app.add_api_route(path, endpoint, dependencies=[Depends(self.auth)], **kwargs)
|
||||
|
||||
@@ -146,8 +146,6 @@ def post_pnginfo(req: models.ReqImageInfo):
|
||||
geninfo, items = images.read_info_from_image(image)
|
||||
if geninfo is None:
|
||||
geninfo = ""
|
||||
if items and items['parameters']:
|
||||
del items['parameters']
|
||||
params = generation_parameters_copypaste.parse_generation_parameters(geninfo)
|
||||
script_callbacks.infotext_pasted_callback(geninfo, params)
|
||||
return models.ResImageInfo(info=geninfo, items=items, parameters=params)
|
||||
|
||||
@@ -1,90 +0,0 @@
|
||||
from modules import shared, sd_hijack, devices
|
||||
from modules.api import models
|
||||
from modules.textual_inversion.preprocess import preprocess
|
||||
|
||||
|
||||
def post_create_embedding(args: dict):
|
||||
from modules.textual_inversion.textual_inversion import create_embedding
|
||||
try:
|
||||
shared.state.begin('api-embedding')
|
||||
filename = create_embedding(**args) # create empty embedding
|
||||
sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings() # reload embeddings so new one can be immediately used
|
||||
shared.state.end()
|
||||
return models.CreateResponse(info = f"create embedding filename: {filename}")
|
||||
except AssertionError as e:
|
||||
shared.state.end()
|
||||
return models.TrainResponse(info = f"create embedding error: {e}")
|
||||
|
||||
def post_create_hypernetwork(args: dict):
|
||||
from modules.hypernetworks.hypernetwork import create_hypernetwork
|
||||
try:
|
||||
shared.state.begin('api-hypernetwork')
|
||||
filename = create_hypernetwork(**args) # create empty embedding # pylint: disable=E1111
|
||||
shared.state.end()
|
||||
return models.CreateResponse(info = f"create hypernetwork filename: {filename}")
|
||||
except AssertionError as e:
|
||||
shared.state.end()
|
||||
return models.TrainResponse(info = f"create hypernetwork error: {e}")
|
||||
|
||||
def post_preprocess(args: dict):
|
||||
try:
|
||||
shared.state.begin('api-preprocess')
|
||||
preprocess(**args) # quick operation unless blip/booru interrogation is enabled
|
||||
shared.state.end()
|
||||
return models.PreprocessResponse(info = 'preprocess complete')
|
||||
except KeyError as e:
|
||||
shared.state.end()
|
||||
return models.PreprocessResponse(info = f"preprocess error: invalid token: {e}")
|
||||
except AssertionError as e:
|
||||
shared.state.end()
|
||||
return models.PreprocessResponse(info = f"preprocess error: {e}")
|
||||
except FileNotFoundError as e:
|
||||
shared.state.end()
|
||||
return models.PreprocessResponse(info = f'preprocess error: {e}')
|
||||
|
||||
def post_train_embedding(args: dict):
|
||||
from modules.textual_inversion.textual_inversion import train_embedding
|
||||
try:
|
||||
shared.state.begin('api-embedding')
|
||||
apply_optimizations = False
|
||||
error = None
|
||||
filename = ''
|
||||
if not apply_optimizations:
|
||||
sd_hijack.undo_optimizations()
|
||||
try:
|
||||
_embedding, filename = train_embedding(**args) # can take a long time to complete
|
||||
except Exception as e:
|
||||
error = e
|
||||
finally:
|
||||
if not apply_optimizations:
|
||||
sd_hijack.apply_optimizations()
|
||||
shared.state.end()
|
||||
return models.TrainResponse(info = f"train embedding complete: filename: {filename} error: {error}")
|
||||
except AssertionError as msg:
|
||||
shared.state.end()
|
||||
return models.TrainResponse(info = f"train embedding error: {msg}")
|
||||
|
||||
def post_train_hypernetwork(args: dict):
|
||||
from modules.hypernetworks.hypernetwork import train_hypernetwork
|
||||
try:
|
||||
shared.state.begin('api-hypernetwork')
|
||||
shared.loaded_hypernetworks = []
|
||||
apply_optimizations = False
|
||||
error = None
|
||||
filename = ''
|
||||
if not apply_optimizations:
|
||||
sd_hijack.undo_optimizations()
|
||||
try:
|
||||
_hypernetwork, filename = train_hypernetwork(**args)
|
||||
except Exception as e:
|
||||
error = e
|
||||
finally:
|
||||
shared.sd_model.cond_stage_model.to(devices.device)
|
||||
shared.sd_model.first_stage_model.to(devices.device)
|
||||
if not apply_optimizations:
|
||||
sd_hijack.apply_optimizations()
|
||||
shared.state.end()
|
||||
return models.TrainResponse(info=f"train embedding complete: filename: {filename} error: {error}")
|
||||
except AssertionError:
|
||||
shared.state.end()
|
||||
return models.TrainResponse(info=f"train embedding error: {error}")
|
||||
@@ -27,7 +27,6 @@ group.add_argument("--auth-file", type=str, default=os.environ.get("SD_AUTHFILE"
|
||||
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("--api-log", default=os.environ.get("SD_APILOG", False), action='store_true', help="Enable logging of all API requests, default: %(default)s")
|
||||
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")
|
||||
|
||||
@@ -482,6 +482,8 @@ def control_run(units: List[unit.Unit], inputs, inits, mask, unit_type: str, is_
|
||||
debug(f'Control exec pipeline: task={sd_models.get_diffusers_task(pipe)} class={pipe.__class__}')
|
||||
debug(f'Control exec pipeline: p={vars(p)}')
|
||||
debug(f'Control exec pipeline: args={p.task_args} image={p.task_args.get("image", None)} control={p.task_args.get("control_image", None)} mask={p.task_args.get("mask_image", None) or p.image_mask} ref={p.task_args.get("ref_image", None)}')
|
||||
if sd_models.get_diffusers_task(pipe) != sd_models.DiffusersTaskType.TEXT_2_IMAGE: # force vae back to gpu if not in txt2img mode
|
||||
sd_models.move_model(pipe.vae, devices.device)
|
||||
p.scripts = scripts.scripts_control
|
||||
p.script_args = input_script_args
|
||||
processed = p.scripts.run(p, *input_script_args)
|
||||
@@ -508,6 +510,10 @@ def control_run(units: List[unit.Unit], inputs, inits, mask, unit_type: str, is_
|
||||
output_image = images.resize_image(resize_mode_after, output_image, width_after, height_after, resize_name_after)
|
||||
|
||||
output_images.append(output_image)
|
||||
if shared.opts.include_mask:
|
||||
if processed_image is not None and isinstance(processed_image, Image.Image):
|
||||
output_images.append(processed_image)
|
||||
|
||||
if is_generator:
|
||||
image_txt = f'{output_image.width}x{output_image.height}' if output_image is not None else 'None'
|
||||
if video is not None:
|
||||
|
||||
@@ -4,7 +4,7 @@ from typing import Union
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline, ControlNetModel, StableDiffusionControlNetPipeline, StableDiffusionXLControlNetPipeline
|
||||
from modules.control.units import detect
|
||||
from modules.shared import log, opts, listdir
|
||||
from modules import errors
|
||||
from modules import errors, sd_models
|
||||
|
||||
|
||||
what = 'ControlNet'
|
||||
@@ -193,7 +193,8 @@ class ControlNetPipeline():
|
||||
scheduler=pipeline.scheduler,
|
||||
feature_extractor=getattr(pipeline, 'feature_extractor', None),
|
||||
controlnet=controlnet, # can be a list
|
||||
).to(pipeline.device)
|
||||
)
|
||||
sd_models.move_model(self.pipeline, pipeline.device)
|
||||
elif detect.is_sd15(pipeline):
|
||||
self.pipeline = StableDiffusionControlNetPipeline(
|
||||
vae=pipeline.vae,
|
||||
@@ -205,7 +206,8 @@ class ControlNetPipeline():
|
||||
requires_safety_checker=False,
|
||||
safety_checker=None,
|
||||
controlnet=controlnet, # can be a list
|
||||
).to(pipeline.device)
|
||||
)
|
||||
sd_models.move_model(self.pipeline, pipeline.device)
|
||||
else:
|
||||
log.error(f'Control {what} pipeline: class={pipeline.__class__.__name__} unsupported model type')
|
||||
return
|
||||
|
||||
@@ -4,6 +4,7 @@ import diffusers.utils
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline
|
||||
from modules.shared import log, opts
|
||||
from modules.control.units import detect
|
||||
from modules import sd_models
|
||||
|
||||
|
||||
what = 'Reference'
|
||||
@@ -34,7 +35,8 @@ class ReferencePipeline():
|
||||
unet=pipeline.unet,
|
||||
scheduler=pipeline.scheduler,
|
||||
feature_extractor=getattr(pipeline, 'feature_extractor', None),
|
||||
).to(pipeline.device)
|
||||
)
|
||||
sd_models.move_model(self.pipeline, pipeline.device)
|
||||
elif detect.is_sd15(pipeline):
|
||||
cls = diffusers.utils.get_class_from_dynamic_module('stable_diffusion_reference', module_file='pipeline.py')
|
||||
self.pipeline = cls(
|
||||
@@ -46,7 +48,8 @@ class ReferencePipeline():
|
||||
feature_extractor=getattr(pipeline, 'feature_extractor', None),
|
||||
requires_safety_checker=False,
|
||||
safety_checker=None,
|
||||
).to(pipeline.device)
|
||||
)
|
||||
sd_models.move_model(self.pipeline, pipeline.device)
|
||||
else:
|
||||
log.error(f'Control {what} pipeline: class={pipeline.__class__.__name__} unsupported model type')
|
||||
return
|
||||
|
||||
@@ -3,7 +3,7 @@ import time
|
||||
from typing import Union
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline, T2IAdapter, MultiAdapter, StableDiffusionAdapterPipeline, StableDiffusionXLAdapterPipeline # pylint: disable=unused-import
|
||||
from modules.shared import log
|
||||
from modules import errors
|
||||
from modules import errors, sd_models
|
||||
from modules.control.units import detect
|
||||
|
||||
|
||||
@@ -128,7 +128,8 @@ class AdapterPipeline():
|
||||
scheduler=pipeline.scheduler,
|
||||
feature_extractor=getattr(pipeline, 'feature_extractor', None),
|
||||
adapter=adapter,
|
||||
).to(pipeline.device)
|
||||
)
|
||||
sd_models.move_model(self.pipeline, pipeline.device)
|
||||
elif detect.is_sd15(pipeline):
|
||||
self.pipeline = StableDiffusionAdapterPipeline(
|
||||
vae=pipeline.vae,
|
||||
@@ -140,7 +141,8 @@ class AdapterPipeline():
|
||||
requires_safety_checker=False,
|
||||
safety_checker=None,
|
||||
adapter=adapter,
|
||||
).to(pipeline.device)
|
||||
)
|
||||
sd_models.move_model(self.pipeline, pipeline.device)
|
||||
else:
|
||||
log.error(f'Control {what} pipeline: class={pipeline.__class__.__name__} unsupported model type')
|
||||
return
|
||||
|
||||
@@ -3,7 +3,7 @@ import time
|
||||
from typing import Union
|
||||
from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline
|
||||
from modules.shared import log, opts, listdir
|
||||
from modules import errors
|
||||
from modules import errors, sd_models
|
||||
from modules.control.units.xs_model import ControlNetXSModel
|
||||
from modules.control.units.xs_pipe import StableDiffusionControlNetXSPipeline, StableDiffusionXLControlNetXSPipeline
|
||||
from modules.control.units import detect
|
||||
@@ -126,7 +126,8 @@ class ControlNetXSPipeline():
|
||||
scheduler=pipeline.scheduler,
|
||||
# feature_extractor=getattr(pipeline, 'feature_extractor', None),
|
||||
controlnet=controlnet, # can be a list
|
||||
).to(pipeline.device)
|
||||
)
|
||||
sd_models.move_model(self.pipeline, pipeline.device)
|
||||
elif detect.is_sd15(pipeline):
|
||||
self.pipeline = StableDiffusionControlNetXSPipeline(
|
||||
vae=pipeline.vae,
|
||||
@@ -138,7 +139,8 @@ class ControlNetXSPipeline():
|
||||
requires_safety_checker=False,
|
||||
safety_checker=None,
|
||||
controlnet=controlnet, # can be a list
|
||||
).to(pipeline.device)
|
||||
)
|
||||
sd_models.move_model(self.pipeline, pipeline.device)
|
||||
else:
|
||||
log.error(f'Control {what} pipeline: class={pipeline.__class__.__name__} unsupported model type')
|
||||
return
|
||||
|
||||
@@ -63,7 +63,7 @@ class DeepDanbooru:
|
||||
|
||||
with devices.inference_context(), devices.autocast():
|
||||
x = torch.from_numpy(a).to(devices.device)
|
||||
y = self.model(x)[0].detach().cpu().numpy()
|
||||
y = self.model(x)[0].detach().float().cpu().numpy()
|
||||
|
||||
probability_dict = {}
|
||||
|
||||
|
||||
+187
-130
@@ -6,8 +6,8 @@ import numpy as np
|
||||
import diffusers
|
||||
import huggingface_hub as hf
|
||||
from PIL import Image
|
||||
from modules import processing, shared, devices
|
||||
|
||||
from modules import processing, shared, devices, extra_networks, sd_models, sd_hijack_freeu, script_callbacks, ipadapter
|
||||
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",
|
||||
@@ -18,11 +18,18 @@ FACEID_MODELS = {
|
||||
# "FaceID Portrait v11": "h94/IP-Adapter-FaceID/ip-adapter-faceid-portrait-v11_sd15.bin",
|
||||
# "FaceID XL Plus v2": "h94/IP-Adapter-FaceID/ip-adapter-faceid_sdxl.bin",
|
||||
}
|
||||
faceid_model = None
|
||||
|
||||
faceid_model_weights = None
|
||||
faceid_model_name = None
|
||||
debug = shared.log.trace if os.environ.get("SD_FACE_DEBUG", None) is not None else lambda *args, **kwargs: None
|
||||
|
||||
|
||||
def hijack_load_ip_adapter(self):
|
||||
self.image_proj_model.load_state_dict(faceid_model_weights["image_proj"])
|
||||
ip_layers = torch.nn.ModuleList(self.pipe.unet.attn_processors.values())
|
||||
ip_layers.load_state_dict(faceid_model_weights["ip_adapter"], strict=False)
|
||||
|
||||
|
||||
def face_id(
|
||||
p: processing.StableDiffusionProcessing,
|
||||
app,
|
||||
@@ -33,7 +40,7 @@ def face_id(
|
||||
scale: float,
|
||||
structure: float,
|
||||
):
|
||||
global faceid_model, faceid_model_name # pylint: disable=global-statement
|
||||
global faceid_model_weights, faceid_model_name # pylint: disable=global-statement
|
||||
if source_images is None or len(source_images) == 0:
|
||||
shared.log.warning('FaceID: no input images')
|
||||
return None
|
||||
@@ -53,134 +60,184 @@ def face_id(
|
||||
shared.log.error(f"FaceID incorrect version of ip_adapter: {e}")
|
||||
return None
|
||||
|
||||
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)
|
||||
if model_path is None:
|
||||
shared.log.error(f"FaceID download failed: model={model} file={ip_ckpt}")
|
||||
return None
|
||||
if override:
|
||||
shared.sd_model.scheduler = diffusers.DDIMScheduler(
|
||||
num_train_timesteps=1000,
|
||||
beta_start=0.00085,
|
||||
beta_end=0.012,
|
||||
beta_schedule="scaled_linear",
|
||||
clip_sample=False,
|
||||
set_alpha_to_one=False,
|
||||
steps_offset=1,
|
||||
)
|
||||
shortcut = None
|
||||
if faceid_model is None or faceid_model_name != model or not cache:
|
||||
shared.log.debug(f"FaceID load: model={model} file={ip_ckpt}")
|
||||
if "XL Plus" in model:
|
||||
image_encoder_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K"
|
||||
faceid_model = IPAdapterFaceIDPlusXL(
|
||||
sd_pipe=shared.sd_model,
|
||||
image_encoder_path=image_encoder_path,
|
||||
ip_ckpt=model_path,
|
||||
lora_rank=128,
|
||||
num_tokens=4,
|
||||
device=devices.device,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
elif "XL" in model:
|
||||
faceid_model = IPAdapterFaceIDXL(
|
||||
sd_pipe=shared.sd_model,
|
||||
ip_ckpt=model_path,
|
||||
lora_rank=128,
|
||||
num_tokens=4,
|
||||
device=devices.device,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
elif "Plus" in model:
|
||||
image_encoder_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K"
|
||||
faceid_model = IPAdapterFaceIDPlus(
|
||||
sd_pipe=shared.sd_model,
|
||||
image_encoder_path=image_encoder_path,
|
||||
ip_ckpt=model_path,
|
||||
lora_rank=128,
|
||||
num_tokens=4,
|
||||
device=devices.device,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
elif "Portrait" in model:
|
||||
faceid_model = IPAdapterFaceIDPortrait(
|
||||
sd_pipe=shared.sd_model,
|
||||
ip_ckpt=model_path,
|
||||
num_tokens=16,
|
||||
n_cond=5,
|
||||
device=devices.device,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
else:
|
||||
faceid_model = IPAdapterFaceID(
|
||||
sd_pipe=shared.sd_model,
|
||||
ip_ckpt=model_path,
|
||||
lora_rank=128,
|
||||
num_tokens=4,
|
||||
device=devices.device,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
shortcut = "v2" in model
|
||||
faceid_model_name = model
|
||||
else:
|
||||
shared.log.debug(f"FaceID cached: model={model} file={ip_ckpt}")
|
||||
|
||||
processed_images = []
|
||||
face_embeds = []
|
||||
face_images = []
|
||||
for i, source_image in enumerate(source_images):
|
||||
np_image = cv2.cvtColor(np.array(source_image), cv2.COLOR_RGB2BGR)
|
||||
faces = app.get(np_image)
|
||||
if len(faces) == 0:
|
||||
shared.log.error("FaceID: no faces found")
|
||||
break
|
||||
face_embeds.append(torch.from_numpy(faces[0].normed_embedding).unsqueeze(0))
|
||||
face_images.append(face_align.norm_crop(np_image, landmark=faces[0].kps, image_size=224))
|
||||
shared.log.debug(f'FaceID face: i={i+1} score={faces[0].det_score:.2f} gender={"female" if faces[0].gender==0 else "male"} age={faces[0].age} bbox={faces[0].bbox}')
|
||||
p.extra_generation_params[f"FaceID {i+1}"] = f'{faces[0].det_score:.2f} {"female" if faces[0].gender==0 else "male"} {faces[0].age}y'
|
||||
if len(face_embeds) == 0:
|
||||
shared.log.error("FaceID: no faces found")
|
||||
return None
|
||||
face_embeds = torch.cat(face_embeds, dim=0)
|
||||
|
||||
ip_model_dict = { # main generate dict
|
||||
"num_samples": p.batch_size,
|
||||
"width": p.width,
|
||||
"height": p.height,
|
||||
"num_inference_steps": p.steps,
|
||||
"scale": scale,
|
||||
"guidance_scale": p.cfg_scale,
|
||||
"faceid_embeds": face_embeds.shape, # placeholder
|
||||
}
|
||||
# optional generate dict
|
||||
if shortcut is not None:
|
||||
ip_model_dict["shortcut"] = shortcut
|
||||
if "Plus" in model:
|
||||
ip_model_dict["s_scale"] = structure
|
||||
shared.log.debug(f"FaceID args: {ip_model_dict}")
|
||||
if "Plus" in model:
|
||||
ip_model_dict["face_image"] = face_images
|
||||
ip_model_dict["faceid_embeds"] = face_embeds # overwrite placeholder
|
||||
# run generate
|
||||
faceid_model.set_scale(scale)
|
||||
for i in range(p.n_iter):
|
||||
ip_model_dict.update({
|
||||
"prompt": p.all_prompts[i],
|
||||
"negative_prompt": p.all_negative_prompts[i],
|
||||
"seed": int(p.all_seeds[i]),
|
||||
})
|
||||
debug(f"FaceID: {ip_model_dict}")
|
||||
res = faceid_model.generate(**ip_model_dict)
|
||||
if isinstance(res, list):
|
||||
processed_images += res
|
||||
faceid_model.set_scale(0)
|
||||
faceid_model = None
|
||||
original_load_ip_adapter = None
|
||||
|
||||
if not cache:
|
||||
faceid_model = None
|
||||
faceid_model_name = None
|
||||
devices.torch_gc()
|
||||
try:
|
||||
shared.prompt_styles.apply_styles_to_extra(p)
|
||||
|
||||
if not shared.opts.cuda_compile:
|
||||
sd_models.apply_token_merging(p.sd_model, p.get_token_merging_ratio())
|
||||
sd_hijack_freeu.apply_freeu(p, shared.backend == shared.Backend.ORIGINAL)
|
||||
|
||||
script_callbacks.before_process_callback(p)
|
||||
|
||||
with context_hypertile_vae(p), context_hypertile_unet(p), devices.inference_context():
|
||||
p.init(p.all_prompts, p.all_seeds, p.all_subseeds)
|
||||
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)
|
||||
if model_path is None:
|
||||
shared.log.error(f"FaceID download failed: model={model} file={ip_ckpt}")
|
||||
return None
|
||||
if override:
|
||||
shared.sd_model.scheduler = diffusers.DDIMScheduler(
|
||||
num_train_timesteps=1000,
|
||||
beta_start=0.00085,
|
||||
beta_end=0.012,
|
||||
beta_schedule="scaled_linear",
|
||||
clip_sample=False,
|
||||
set_alpha_to_one=False,
|
||||
steps_offset=1,
|
||||
)
|
||||
if faceid_model_weights is None or faceid_model_name != model or not cache:
|
||||
shared.log.debug(f"FaceID load: model={model} file={ip_ckpt}")
|
||||
faceid_model_weights = torch.load(model_path, map_location="cpu")
|
||||
else:
|
||||
shared.log.debug(f"FaceID cached: model={model} file={ip_ckpt}")
|
||||
|
||||
if "XL Plus" in model:
|
||||
image_encoder_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K"
|
||||
original_load_ip_adapter = IPAdapterFaceIDPlusXL.load_ip_adapter
|
||||
IPAdapterFaceIDPlusXL.load_ip_adapter = hijack_load_ip_adapter
|
||||
faceid_model = IPAdapterFaceIDPlusXL(
|
||||
sd_pipe=shared.sd_model,
|
||||
image_encoder_path=image_encoder_path,
|
||||
ip_ckpt=model_path,
|
||||
lora_rank=128,
|
||||
num_tokens=4,
|
||||
device=devices.device,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
elif "XL" in model:
|
||||
original_load_ip_adapter = IPAdapterFaceIDXL.load_ip_adapter
|
||||
IPAdapterFaceIDXL.load_ip_adapter = hijack_load_ip_adapter
|
||||
faceid_model = IPAdapterFaceIDXL(
|
||||
sd_pipe=shared.sd_model,
|
||||
ip_ckpt=model_path,
|
||||
lora_rank=128,
|
||||
num_tokens=4,
|
||||
device=devices.device,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
elif "Plus" in model:
|
||||
original_load_ip_adapter = IPAdapterFaceIDPlus.load_ip_adapter
|
||||
IPAdapterFaceIDPlus.load_ip_adapter = hijack_load_ip_adapter
|
||||
image_encoder_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K"
|
||||
faceid_model = IPAdapterFaceIDPlus(
|
||||
sd_pipe=shared.sd_model,
|
||||
image_encoder_path=image_encoder_path,
|
||||
ip_ckpt=model_path,
|
||||
lora_rank=128,
|
||||
num_tokens=4,
|
||||
device=devices.device,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
elif "Portrait" in model:
|
||||
original_load_ip_adapter = IPAdapterFaceIDPortrait.load_ip_adapter
|
||||
IPAdapterFaceIDPortrait.load_ip_adapter = hijack_load_ip_adapter
|
||||
faceid_model = IPAdapterFaceIDPortrait(
|
||||
sd_pipe=shared.sd_model,
|
||||
ip_ckpt=model_path,
|
||||
num_tokens=16,
|
||||
n_cond=5,
|
||||
device=devices.device,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
else:
|
||||
original_load_ip_adapter = IPAdapterFaceID.load_ip_adapter
|
||||
IPAdapterFaceID.load_ip_adapter = hijack_load_ip_adapter
|
||||
faceid_model = IPAdapterFaceID(
|
||||
sd_pipe=shared.sd_model,
|
||||
ip_ckpt=model_path,
|
||||
lora_rank=128,
|
||||
num_tokens=4,
|
||||
device=devices.device,
|
||||
torch_dtype=devices.dtype,
|
||||
)
|
||||
|
||||
shortcut = "v2" in model
|
||||
faceid_model_name = model
|
||||
face_embeds = []
|
||||
face_images = []
|
||||
for i, source_image in enumerate(source_images):
|
||||
np_image = cv2.cvtColor(np.array(source_image), cv2.COLOR_RGB2BGR)
|
||||
faces = app.get(np_image)
|
||||
if len(faces) == 0:
|
||||
shared.log.error("FaceID: no faces found")
|
||||
break
|
||||
face_embeds.append(torch.from_numpy(faces[0].normed_embedding).unsqueeze(0))
|
||||
face_images.append(face_align.norm_crop(np_image, landmark=faces[0].kps, image_size=224))
|
||||
shared.log.debug(f'FaceID face: i={i+1} score={faces[0].det_score:.2f} gender={"female" if faces[0].gender==0 else "male"} age={faces[0].age} bbox={faces[0].bbox}')
|
||||
p.extra_generation_params[f"FaceID {i+1}"] = f'{faces[0].det_score:.2f} {"female" if faces[0].gender==0 else "male"} {faces[0].age}y'
|
||||
|
||||
if len(face_embeds) == 0:
|
||||
shared.log.error("FaceID: no faces found")
|
||||
return None
|
||||
|
||||
face_embeds = torch.cat(face_embeds, dim=0)
|
||||
ip_model_dict = { # main generate dict
|
||||
"num_samples": p.batch_size,
|
||||
"width": p.width,
|
||||
"height": p.height,
|
||||
"num_inference_steps": p.steps,
|
||||
"scale": scale,
|
||||
"guidance_scale": p.cfg_scale,
|
||||
"faceid_embeds": face_embeds.shape, # placeholder
|
||||
}
|
||||
|
||||
# optional generate dict
|
||||
if shortcut is not None:
|
||||
ip_model_dict["shortcut"] = shortcut
|
||||
if "Plus" in model:
|
||||
ip_model_dict["s_scale"] = structure
|
||||
shared.log.debug(f"FaceID args: {ip_model_dict}")
|
||||
if "Plus" in model:
|
||||
ip_model_dict["face_image"] = face_images
|
||||
ip_model_dict["faceid_embeds"] = face_embeds # overwrite placeholder
|
||||
faceid_model.set_scale(scale)
|
||||
extra_network_data = None
|
||||
|
||||
for i in range(p.n_iter):
|
||||
p.iteration = i
|
||||
p.prompts = p.all_prompts[i * p.batch_size:(i + 1) * p.batch_size]
|
||||
p.negative_prompts = p.all_negative_prompts[i * p.batch_size:(i + 1) * p.batch_size]
|
||||
p.prompts, extra_network_data = extra_networks.parse_prompts(p.prompts)
|
||||
p.seeds = p.all_seeds[i * p.batch_size:(i + 1) * p.batch_size]
|
||||
if not p.disable_extra_networks:
|
||||
with devices.autocast():
|
||||
extra_networks.activate(p, extra_network_data)
|
||||
ip_model_dict.update({
|
||||
"prompt": p.prompts,
|
||||
"negative_prompt": p.negative_prompts,
|
||||
"seed": int(p.seeds[0]),
|
||||
})
|
||||
debug(f"FaceID: {ip_model_dict}")
|
||||
res = faceid_model.generate(**ip_model_dict)
|
||||
if isinstance(res, list):
|
||||
processed_images += res
|
||||
|
||||
faceid_model.set_scale(0)
|
||||
faceid_model = None
|
||||
|
||||
if not cache:
|
||||
faceid_model_weights = None
|
||||
faceid_model_name = None
|
||||
devices.torch_gc()
|
||||
|
||||
ipadapter.unapply(p.sd_model)
|
||||
if not p.disable_extra_networks:
|
||||
extra_networks.deactivate(p, extra_network_data)
|
||||
|
||||
p.extra_generation_params["IP Adapter"] = f"{basename}:{scale}"
|
||||
finally:
|
||||
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 not shared.opts.cuda_compile:
|
||||
sd_models.apply_token_merging(p.sd_model, 0)
|
||||
script_callbacks.after_process_callback(p)
|
||||
|
||||
p.extra_generation_params["IP Adapter"] = f"{basename}:{scale}"
|
||||
return processed_images
|
||||
|
||||
@@ -41,7 +41,7 @@ def instant_id(p: processing.StableDiffusionProcessing, app, source_images, stre
|
||||
face_adapter = hf.hf_hub_download(repo_id=REPO_ID, filename="ip-adapter.bin")
|
||||
if controlnet_model is None or not cache:
|
||||
controlnet_model = ControlNetModel.from_pretrained(REPO_ID, subfolder="ControlNetModel", torch_dtype=devices.dtype, cache_dir=shared.opts.diffusers_dir)
|
||||
controlnet_model.to(devices.device, devices.dtype)
|
||||
sd_models.move_model(controlnet_model, devices.device)
|
||||
|
||||
processing.process_init(p)
|
||||
|
||||
|
||||
@@ -1,23 +1,13 @@
|
||||
import datetime
|
||||
import html
|
||||
import os
|
||||
from collections import deque
|
||||
import inspect
|
||||
from statistics import stdev, mean
|
||||
from rich import progress
|
||||
import tqdm
|
||||
import torch
|
||||
from torch import einsum
|
||||
from torch.nn.init import normal_, xavier_normal_, xavier_uniform_, kaiming_normal_, kaiming_uniform_, zeros_
|
||||
from einops import rearrange, repeat
|
||||
from ldm.util import default
|
||||
from modules import devices, processing, sd_models, shared, hashes, errors, files_cache
|
||||
import modules.textual_inversion.dataset
|
||||
from modules.textual_inversion import textual_inversion, ti_logging
|
||||
from modules.textual_inversion.learn_schedule import LearnRateScheduler
|
||||
|
||||
|
||||
optimizer_dict = {optim_name : cls_obj for optim_name, cls_obj in inspect.getmembers(torch.optim, inspect.isclass) if optim_name != "Optimizer"}
|
||||
from modules import devices, shared, hashes, errors, files_cache
|
||||
|
||||
|
||||
class HypernetworkModule(torch.nn.Module):
|
||||
@@ -410,341 +400,3 @@ def report_statistics(loss_info:dict):
|
||||
print(recent)
|
||||
except Exception as e:
|
||||
print(e)
|
||||
|
||||
|
||||
def create_hypernetwork(name, enable_sizes, overwrite_old, layer_structure=None, activation_func=None, weight_init=None, add_layer_norm=False, use_dropout=False, dropout_structure=None):
|
||||
# Remove illegal characters from name.
|
||||
name = "".join( x for x in name if (x.isalnum() or x in "._- "))
|
||||
assert name, "Name cannot be empty!"
|
||||
fn = os.path.join(shared.opts.hypernetwork_dir, f"{name}.pt")
|
||||
if not overwrite_old:
|
||||
assert not os.path.exists(fn), f"file {fn} already exists"
|
||||
if type(layer_structure) == str:
|
||||
layer_structure = [float(x.strip()) for x in layer_structure.split(",")]
|
||||
if use_dropout and dropout_structure and type(dropout_structure) == str:
|
||||
dropout_structure = [float(x.strip()) for x in dropout_structure.split(",")]
|
||||
else:
|
||||
dropout_structure = [0] * len(layer_structure)
|
||||
hypernet = modules.hypernetworks.hypernetwork.Hypernetwork(
|
||||
name=name,
|
||||
enable_sizes=[int(x) for x in enable_sizes],
|
||||
layer_structure=layer_structure,
|
||||
activation_func=activation_func,
|
||||
weight_init=weight_init,
|
||||
add_layer_norm=add_layer_norm,
|
||||
use_dropout=use_dropout,
|
||||
dropout_structure=dropout_structure
|
||||
)
|
||||
hypernet.save(fn)
|
||||
shared.reload_hypernetworks()
|
||||
return name
|
||||
|
||||
|
||||
def train_hypernetwork(id_task, hypernetwork_name, learn_rate, batch_size, gradient_step, data_root, log_directory, training_width, training_height, varsize, steps, clip_grad_mode, clip_grad_value, shuffle_tags, tag_drop_out, latent_sampling_method, use_weight, create_image_every, save_hypernetwork_every, template_filename, preview_from_txt2img, preview_prompt, preview_negative_prompt, preview_steps, preview_sampler_index, preview_cfg_scale, preview_seed, preview_width, preview_height): # pylint: disable=unused-argument
|
||||
# images allows training previews to have infotext. Importing it at the top causes a circular import problem.
|
||||
from modules import images, sd_hijack_checkpoint
|
||||
|
||||
save_hypernetwork_every = save_hypernetwork_every or 0
|
||||
create_image_every = create_image_every or 0
|
||||
template_file = textual_inversion.textual_inversion_templates.get(template_filename, None)
|
||||
textual_inversion.validate_train_inputs(hypernetwork_name, learn_rate, batch_size, gradient_step, data_root, template_file, template_filename, steps, save_hypernetwork_every, create_image_every, name="hypernetwork")
|
||||
template_file = template_file.path
|
||||
|
||||
path = shared.hypernetworks.get(hypernetwork_name, None)
|
||||
hypernetwork = Hypernetwork()
|
||||
hypernetwork.load(path)
|
||||
shared.loaded_hypernetworks = [hypernetwork]
|
||||
|
||||
shared.state.job = "train"
|
||||
shared.state.textinfo = "Initializing hypernetwork training..."
|
||||
shared.state.job_count = steps
|
||||
|
||||
hypernetwork_name = hypernetwork_name.rsplit('(', 1)[0]
|
||||
filename = os.path.join(shared.opts.hypernetwork_dir, f'{hypernetwork_name}.pt')
|
||||
|
||||
log_directory = os.path.join(log_directory, datetime.datetime.now().strftime("%Y-%m-%d"), hypernetwork_name)
|
||||
unload = shared.opts.unload_models_when_training
|
||||
|
||||
if save_hypernetwork_every > 0:
|
||||
hypernetwork_dir = os.path.join(log_directory, "hypernetworks")
|
||||
os.makedirs(hypernetwork_dir, exist_ok=True)
|
||||
else:
|
||||
hypernetwork_dir = None
|
||||
|
||||
if create_image_every > 0:
|
||||
images_dir = os.path.join(log_directory, "images")
|
||||
os.makedirs(images_dir, exist_ok=True)
|
||||
else:
|
||||
images_dir = None
|
||||
|
||||
checkpoint = sd_models.select_checkpoint()
|
||||
|
||||
initial_step = hypernetwork.step or 0
|
||||
if initial_step >= steps:
|
||||
shared.state.textinfo = "Model has already been trained beyond specified max steps"
|
||||
return hypernetwork, filename
|
||||
|
||||
scheduler = LearnRateScheduler(learn_rate, steps, initial_step)
|
||||
|
||||
clip_grad = torch.nn.utils.clip_grad_value_ if clip_grad_mode == "value" else torch.nn.utils.clip_grad_norm_ if clip_grad_mode == "norm" else None
|
||||
if clip_grad:
|
||||
clip_grad_sched = LearnRateScheduler(clip_grad_value, steps, initial_step, verbose=False)
|
||||
|
||||
if shared.opts.training_enable_tensorboard:
|
||||
tensorboard_writer = textual_inversion.tensorboard_setup(log_directory)
|
||||
|
||||
# dataset loading may take a while, so input validations and early returns should be done before this
|
||||
shared.state.textinfo = f"Preparing dataset from {html.escape(data_root)}..."
|
||||
|
||||
pin_memory = shared.opts.pin_memory
|
||||
|
||||
ds = modules.textual_inversion.dataset.PersonalizedBase(data_root=data_root, width=training_width, height=training_height, repeats=shared.opts.training_image_repeats_per_epoch, placeholder_token=hypernetwork_name, model=shared.sd_model, cond_model=shared.sd_model.cond_stage_model, device=devices.device, template_file=template_file, include_cond=True, batch_size=batch_size, gradient_step=gradient_step, shuffle_tags=shuffle_tags, tag_drop_out=tag_drop_out, latent_sampling_method=latent_sampling_method, varsize=varsize, use_weight=use_weight)
|
||||
|
||||
if shared.opts.save_training_settings_to_txt:
|
||||
saved_params = dict(
|
||||
model_name=checkpoint.model_name, model_hash=checkpoint.shorthash, num_of_dataset_images=len(ds),
|
||||
**{field: getattr(hypernetwork, field) for field in ['layer_structure', 'activation_func', 'weight_init', 'add_layer_norm', 'use_dropout', ]}
|
||||
)
|
||||
ti_logging.save_settings_to_file(log_directory, {**saved_params, **locals()})
|
||||
|
||||
latent_sampling_method = ds.latent_sampling_method
|
||||
|
||||
dl = modules.textual_inversion.dataset.PersonalizedDataLoader(ds, latent_sampling_method=latent_sampling_method, batch_size=ds.batch_size, pin_memory=pin_memory)
|
||||
|
||||
old_parallel_processing_allowed = shared.parallel_processing_allowed
|
||||
|
||||
if unload:
|
||||
shared.parallel_processing_allowed = False
|
||||
shared.sd_model.cond_stage_model.to(devices.cpu)
|
||||
shared.sd_model.first_stage_model.to(devices.cpu)
|
||||
|
||||
weights = hypernetwork.weights()
|
||||
hypernetwork.train()
|
||||
|
||||
# Here we use optimizer from saved HN, or we can specify as UI option.
|
||||
if hypernetwork.optimizer_name in optimizer_dict:
|
||||
optimizer = optimizer_dict[hypernetwork.optimizer_name](params=weights, lr=scheduler.learn_rate)
|
||||
optimizer_name = hypernetwork.optimizer_name
|
||||
else:
|
||||
print(f"Optimizer type {hypernetwork.optimizer_name} is not defined!")
|
||||
optimizer = torch.optim.AdamW(params=weights, lr=scheduler.learn_rate)
|
||||
optimizer_name = 'AdamW'
|
||||
|
||||
if hypernetwork.optimizer_state_dict: # This line must be changed if Optimizer type can be different from saved optimizer.
|
||||
try:
|
||||
optimizer.load_state_dict(hypernetwork.optimizer_state_dict)
|
||||
except RuntimeError as e:
|
||||
print("Cannot resume from saved optimizer!")
|
||||
print(e)
|
||||
|
||||
scaler = torch.cuda.amp.GradScaler()
|
||||
|
||||
batch_size = ds.batch_size
|
||||
gradient_step = ds.gradient_step
|
||||
# n steps = batch_size * gradient_step * n image processed
|
||||
steps_per_epoch = len(ds) // batch_size // gradient_step
|
||||
max_steps_per_epoch = len(ds) // batch_size - (len(ds) // batch_size) % gradient_step
|
||||
loss_step = 0
|
||||
_loss_step = 0 #internal
|
||||
# size = len(ds.indexes)
|
||||
# loss_dict = defaultdict(lambda : deque(maxlen = 1024))
|
||||
loss_logging = deque(maxlen=len(ds) * 3) # this should be configurable parameter, this is 3 * epoch(dataset size)
|
||||
# losses = torch.zeros((size,))
|
||||
# previous_mean_losses = [0]
|
||||
# previous_mean_loss = 0
|
||||
# print("Mean loss of {} elements".format(size))
|
||||
|
||||
_steps_without_grad = 0
|
||||
|
||||
last_saved_file = "<none>"
|
||||
last_saved_image = "<none>"
|
||||
forced_filename = "<none>"
|
||||
|
||||
pbar = tqdm.tqdm(total=steps - initial_step)
|
||||
try:
|
||||
sd_hijack_checkpoint.add()
|
||||
|
||||
for _i in range((steps-initial_step) * gradient_step):
|
||||
if scheduler.finished:
|
||||
break
|
||||
if shared.state.interrupted:
|
||||
break
|
||||
for j, batch in enumerate(dl):
|
||||
# works as a drop_last=True for gradient accumulation
|
||||
if j == max_steps_per_epoch:
|
||||
break
|
||||
scheduler.apply(optimizer, hypernetwork.step)
|
||||
if scheduler.finished:
|
||||
break
|
||||
if shared.state.interrupted:
|
||||
break
|
||||
|
||||
if clip_grad:
|
||||
clip_grad_sched.step(hypernetwork.step)
|
||||
|
||||
with devices.autocast():
|
||||
x = batch.latent_sample.to(devices.device, non_blocking=pin_memory)
|
||||
if use_weight:
|
||||
w = batch.weight.to(devices.device, non_blocking=pin_memory)
|
||||
if tag_drop_out != 0 or shuffle_tags:
|
||||
shared.sd_model.cond_stage_model.to(devices.device)
|
||||
c = shared.sd_model.cond_stage_model(batch.cond_text).to(devices.device, non_blocking=pin_memory)
|
||||
shared.sd_model.cond_stage_model.to(devices.cpu)
|
||||
else:
|
||||
c = stack_conds(batch.cond).to(devices.device, non_blocking=pin_memory)
|
||||
if use_weight:
|
||||
loss = shared.sd_model.weighted_forward(x, c, w)[0] / gradient_step
|
||||
del w
|
||||
else:
|
||||
loss = shared.sd_model.forward(x, c)[0] / gradient_step
|
||||
del x
|
||||
del c
|
||||
_loss_step += loss.item()
|
||||
|
||||
scaler.scale(loss).backward()
|
||||
# go back until we reach gradient accumulation steps
|
||||
if (j + 1) % gradient_step != 0:
|
||||
continue
|
||||
loss_logging.append(_loss_step)
|
||||
if clip_grad:
|
||||
clip_grad(weights, clip_grad_sched.learn_rate)
|
||||
|
||||
scaler.step(optimizer)
|
||||
scaler.update()
|
||||
hypernetwork.step += 1
|
||||
pbar.update()
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
loss_step = _loss_step
|
||||
_loss_step = 0
|
||||
steps_done = hypernetwork.step + 1
|
||||
epoch_num = hypernetwork.step // steps_per_epoch
|
||||
epoch_step = hypernetwork.step % steps_per_epoch
|
||||
|
||||
description = f"Training hypernetwork [Epoch {epoch_num}: {epoch_step+1}/{steps_per_epoch}]loss: {loss_step:.7f}"
|
||||
pbar.set_description(description)
|
||||
if hypernetwork_dir is not None and steps_done % save_hypernetwork_every == 0:
|
||||
# Before saving, change name to match current checkpoint.
|
||||
hypernetwork_name_every = f'{hypernetwork_name}-{steps_done}'
|
||||
last_saved_file = os.path.join(hypernetwork_dir, f'{hypernetwork_name_every}.pt')
|
||||
hypernetwork.optimizer_name = optimizer_name
|
||||
if shared.opts.save_optimizer_state:
|
||||
hypernetwork.optimizer_state_dict = optimizer.state_dict()
|
||||
save_hypernetwork(hypernetwork, checkpoint, hypernetwork_name, last_saved_file)
|
||||
hypernetwork.optimizer_state_dict = None # dereference it after saving, to save memory.
|
||||
|
||||
|
||||
|
||||
if shared.opts.training_enable_tensorboard:
|
||||
epoch_num = hypernetwork.step // len(ds)
|
||||
epoch_step = hypernetwork.step - (epoch_num * len(ds)) + 1
|
||||
mean_loss = sum(loss_logging) / len(loss_logging)
|
||||
textual_inversion.tensorboard_add(tensorboard_writer, loss=mean_loss, global_step=hypernetwork.step, step=epoch_step, learn_rate=scheduler.learn_rate, epoch_num=epoch_num)
|
||||
|
||||
textual_inversion.write_loss(log_directory, "hypernetwork_loss.csv", hypernetwork.step, steps_per_epoch, {
|
||||
"loss": f"{loss_step:.7f}",
|
||||
"learn_rate": scheduler.learn_rate
|
||||
})
|
||||
|
||||
if images_dir is not None and steps_done % create_image_every == 0:
|
||||
forced_filename = f'{hypernetwork_name}-{steps_done}'
|
||||
last_saved_image = os.path.join(images_dir, forced_filename)
|
||||
hypernetwork.eval()
|
||||
rng_state = torch.get_rng_state()
|
||||
cuda_rng_state = None
|
||||
cuda_rng_state = torch.cuda.get_rng_state_all()
|
||||
shared.sd_model.cond_stage_model.to(devices.device)
|
||||
shared.sd_model.first_stage_model.to(devices.device)
|
||||
|
||||
p = processing.StableDiffusionProcessingTxt2Img(
|
||||
sd_model=shared.sd_model,
|
||||
do_not_save_grid=True,
|
||||
do_not_save_samples=True,
|
||||
)
|
||||
|
||||
p.disable_extra_networks = True
|
||||
|
||||
if preview_from_txt2img:
|
||||
p.prompt = preview_prompt
|
||||
p.negative_prompt = preview_negative_prompt
|
||||
p.steps = preview_steps
|
||||
p.sampler_name = processing.get_sampler_name(preview_sampler_index)
|
||||
p.cfg_scale = preview_cfg_scale
|
||||
p.seed = preview_seed
|
||||
p.width = preview_width
|
||||
p.height = preview_height
|
||||
else:
|
||||
p.prompt = batch.cond_text[0]
|
||||
p.steps = 20
|
||||
p.width = training_width
|
||||
p.height = training_height
|
||||
|
||||
preview_text = p.prompt
|
||||
|
||||
processed = processing.process_images(p)
|
||||
image = processed.images[0] if len(processed.images) > 0 else None
|
||||
|
||||
if unload:
|
||||
shared.sd_model.cond_stage_model.to(devices.cpu)
|
||||
shared.sd_model.first_stage_model.to(devices.cpu)
|
||||
torch.set_rng_state(rng_state)
|
||||
torch.cuda.set_rng_state_all(cuda_rng_state)
|
||||
hypernetwork.train()
|
||||
if image is not None:
|
||||
shared.state.assign_current_image(image)
|
||||
if shared.opts.training_enable_tensorboard and shared.opts.training_tensorboard_save_images:
|
||||
textual_inversion.tensorboard_add_image(tensorboard_writer,
|
||||
f"Validation at epoch {epoch_num}", image,
|
||||
hypernetwork.step)
|
||||
last_saved_image, _last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False)
|
||||
last_saved_image += f", prompt: {preview_text}"
|
||||
|
||||
shared.state.job_no = hypernetwork.step
|
||||
|
||||
shared.state.textinfo = f"""
|
||||
<p>
|
||||
Loss: {loss_step:.7f}<br/>
|
||||
Step: {steps_done}<br/>
|
||||
Last prompt: {html.escape(batch.cond_text[0])}<br/>
|
||||
Last saved hypernetwork: {html.escape(last_saved_file)}<br/>
|
||||
Last saved image: {html.escape(last_saved_image)}<br/>
|
||||
</p>
|
||||
"""
|
||||
except Exception as e:
|
||||
errors.display(e, 'hypernetwork train')
|
||||
finally:
|
||||
pbar.leave = False
|
||||
pbar.close()
|
||||
hypernetwork.eval()
|
||||
#report_statistics(loss_dict)
|
||||
sd_hijack_checkpoint.remove()
|
||||
|
||||
|
||||
|
||||
filename = os.path.join(shared.opts.hypernetwork_dir, f'{hypernetwork_name}.pt')
|
||||
hypernetwork.optimizer_name = optimizer_name
|
||||
if shared.opts.save_optimizer_state:
|
||||
hypernetwork.optimizer_state_dict = optimizer.state_dict()
|
||||
save_hypernetwork(hypernetwork, checkpoint, hypernetwork_name, filename)
|
||||
|
||||
del optimizer
|
||||
hypernetwork.optimizer_state_dict = None # dereference it after saving, to save memory.
|
||||
shared.sd_model.cond_stage_model.to(devices.device)
|
||||
shared.sd_model.first_stage_model.to(devices.device)
|
||||
shared.parallel_processing_allowed = old_parallel_processing_allowed
|
||||
|
||||
return hypernetwork, filename
|
||||
|
||||
def save_hypernetwork(hypernetwork, checkpoint, hypernetwork_name, filename):
|
||||
old_hypernetwork_name = hypernetwork.name
|
||||
old_sd_checkpoint = hypernetwork.sd_checkpoint if hasattr(hypernetwork, "sd_checkpoint") else None
|
||||
old_sd_checkpoint_name = hypernetwork.sd_checkpoint_name if hasattr(hypernetwork, "sd_checkpoint_name") else None
|
||||
try:
|
||||
hypernetwork.sd_checkpoint = checkpoint.shorthash
|
||||
hypernetwork.sd_checkpoint_name = checkpoint.model_name
|
||||
hypernetwork.name = hypernetwork_name
|
||||
hypernetwork.save(filename)
|
||||
except Exception:
|
||||
hypernetwork.sd_checkpoint = old_sd_checkpoint
|
||||
hypernetwork.sd_checkpoint_name = old_sd_checkpoint_name
|
||||
hypernetwork.name = old_hypernetwork_name
|
||||
raise
|
||||
|
||||
@@ -1,31 +0,0 @@
|
||||
import html
|
||||
import gradio as gr
|
||||
import modules.hypernetworks.hypernetwork
|
||||
from modules import devices, sd_hijack, shared
|
||||
|
||||
not_available = ["hardswish", "multiheadattention"]
|
||||
keys = [x for x in modules.hypernetworks.hypernetwork.HypernetworkModule.activation_dict.keys() if x not in not_available]
|
||||
|
||||
|
||||
def create_hypernetwork(name, enable_sizes, overwrite_old, layer_structure=None, activation_func=None, weight_init=None, add_layer_norm=False, use_dropout=False, dropout_structure=None):
|
||||
filename = modules.hypernetworks.hypernetwork.create_hypernetwork(name, enable_sizes, overwrite_old, layer_structure, activation_func, weight_init, add_layer_norm, use_dropout, dropout_structure)
|
||||
return gr.Dropdown.update(choices=sorted(shared.hypernetworks)), f"Created: {filename}", ""
|
||||
|
||||
|
||||
def train_hypernetwork(*args):
|
||||
shared.loaded_hypernetworks = []
|
||||
assert not shared.cmd_opts.lowvram, 'Training models with lowvram is not possible'
|
||||
try:
|
||||
sd_hijack.undo_optimizations()
|
||||
hypernetwork, filename = modules.hypernetworks.hypernetwork.train_hypernetwork(*args)
|
||||
res = f"""
|
||||
Training {'interrupted' if shared.state.interrupted else 'finished'} at {hypernetwork.step} steps.
|
||||
Hypernetwork saved to {html.escape(filename)}
|
||||
"""
|
||||
return res, ""
|
||||
except Exception as e:
|
||||
raise RuntimeError("Hypernetwork error") from e
|
||||
finally:
|
||||
shared.sd_model.cond_stage_model.to(devices.device)
|
||||
shared.sd_model.first_stage_model.to(devices.device)
|
||||
sd_hijack.apply_optimizations()
|
||||
@@ -10,7 +10,7 @@ TODO ipadapter items:
|
||||
import os
|
||||
import time
|
||||
from PIL import Image
|
||||
from modules import processing, shared, devices
|
||||
from modules import processing, shared, devices, sd_models
|
||||
|
||||
|
||||
base_repo = "h94/IP-Adapter"
|
||||
@@ -163,7 +163,7 @@ def apply(pipe, p: processing.StableDiffusionProcessing, adapter_names=[], adapt
|
||||
except Exception as e:
|
||||
shared.log.error(f'IP adapter: failed to load image encoder: {e}')
|
||||
return
|
||||
pipe.image_encoder.to(devices.device)
|
||||
sd_models.move_model(pipe.image_encoder, devices.device)
|
||||
|
||||
# main code
|
||||
t0 = time.time()
|
||||
|
||||
+1
-1
@@ -96,7 +96,7 @@ class SimpleLama:
|
||||
image, mask = prepare_img_and_mask(image, mask, self.device)
|
||||
with devices.inference_context():
|
||||
inpainted = self.model(image, mask)
|
||||
cur_res = inpainted[0].permute(1, 2, 0).detach().cpu().numpy()
|
||||
cur_res = inpainted[0].permute(1, 2, 0).detach().float().cpu().numpy()
|
||||
cur_res = np.clip(cur_res * 255, 0, 255).astype(np.uint8)
|
||||
cur_res = Image.fromarray(cur_res)
|
||||
return cur_res
|
||||
|
||||
+3
-2
@@ -8,7 +8,7 @@ import numpy as np
|
||||
import cv2
|
||||
from PIL import Image, ImageFilter, ImageOps
|
||||
from transformers import SamModel, SamImageProcessor, MaskGenerationPipeline
|
||||
from modules import shared, errors, devices, ui_components, ui_symbols, paths
|
||||
from modules import shared, errors, devices, ui_components, ui_symbols, paths, sd_models
|
||||
from modules.memstats import memory_stats
|
||||
|
||||
|
||||
@@ -480,7 +480,8 @@ def run_lama(input_image: gr.Image, input_mask: gr.Image = None):
|
||||
shared.log.debug(f'Mask LaMa loading: model={modules.lama.LAMA_MODEL_URL}')
|
||||
lama_model = modules.lama.SimpleLama()
|
||||
shared.log.debug(f'Mask LaMa loaded: {memory_stats()}')
|
||||
lama_model.model.to(devices.device)
|
||||
sd_models.move_model(lama_model.model, devices.device)
|
||||
|
||||
result = lama_model(input_image, input_mask)
|
||||
if shared.opts.control_move_processor:
|
||||
lama_model.model.to('cpu')
|
||||
|
||||
+9
-6
@@ -28,12 +28,15 @@ class MemUsageMonitor():
|
||||
|
||||
def reset(self):
|
||||
if not self.disabled:
|
||||
torch.cuda.reset_peak_memory_stats(self.device)
|
||||
self.data['retries'] = 0
|
||||
self.data['oom'] = 0
|
||||
# torch.cuda.reset_accumulated_memory_stats(self.device)
|
||||
# torch.cuda.reset_max_memory_allocated(self.device)
|
||||
# torch.cuda.reset_max_memory_cached(self.device)
|
||||
try:
|
||||
torch.cuda.reset_peak_memory_stats(self.device)
|
||||
self.data['retries'] = 0
|
||||
self.data['oom'] = 0
|
||||
# torch.cuda.reset_accumulated_memory_stats(self.device)
|
||||
# torch.cuda.reset_max_memory_allocated(self.device)
|
||||
# torch.cuda.reset_max_memory_cached(self.device)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def read(self):
|
||||
if not self.disabled:
|
||||
|
||||
@@ -186,14 +186,14 @@ def download_diffusers_model(hub_id: str, cache_dir: str = None, download_config
|
||||
except Exception as e:
|
||||
err = e
|
||||
ok = False
|
||||
debug(f"Diffusers download error: {hub_id} {e}")
|
||||
debug(f'Diffusers download error: id="{hub_id}" {e}')
|
||||
if not ok and 'Repository Not Found' not in str(err):
|
||||
try:
|
||||
download_config.pop('load_connected_pipeline', None)
|
||||
download_config.pop('variant', None)
|
||||
pipeline_dir = hf.snapshot_download(hub_id, **download_config)
|
||||
except Exception as e:
|
||||
debug(f"Diffusers download error: {hub_id} {e}")
|
||||
debug(f'Diffusers download error: id="{hub_id}" {e}')
|
||||
if 'gated' in str(e):
|
||||
shared.log.error(f'Diffusers download error: id="{hub_id}" model access requires login')
|
||||
return None
|
||||
|
||||
@@ -396,6 +396,10 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed:
|
||||
if not p.disable_extra_networks:
|
||||
extra_networks.deactivate(p, extra_network_data)
|
||||
|
||||
if shared.opts.include_mask:
|
||||
if getattr(p, 'image_mask', None) is not None and isinstance(p.image_mask, Image.Image):
|
||||
output_images.append(p.image_mask)
|
||||
|
||||
processed = Processed(
|
||||
p,
|
||||
images_list=output_images,
|
||||
|
||||
@@ -261,14 +261,14 @@ class StableDiffusionProcessingTxt2Img(StableDiffusionProcessing):
|
||||
self.hr_upscale_to_y = self.hr_resize_y
|
||||
self.truncate_x = (self.hr_upscale_to_x - target_w) // 8
|
||||
self.truncate_y = (self.hr_upscale_to_y - target_h) // 8
|
||||
# special case: the user has chosen to do nothing
|
||||
if (self.hr_upscale_to_x == self.width and self.hr_upscale_to_y == self.height) or self.hr_upscaler is None or self.hr_upscaler == 'None':
|
||||
self.is_hr_pass = False
|
||||
return
|
||||
self.is_hr_pass = True
|
||||
hypertile_set(self, hr=True)
|
||||
shared.state.job_count = 2 * self.n_iter
|
||||
shared.log.debug(f'Init hires: upscaler="{self.hr_upscaler}" sampler="{self.hr_sampler_name}" resize={self.hr_resize_x}x{self.hr_resize_y} upscale={self.hr_upscale_to_x}x{self.hr_upscale_to_y}')
|
||||
if shared.backend == shared.Backend.ORIGINAL: # diffusers are handled in processing_diffusers
|
||||
if (self.hr_upscale_to_x == self.width and self.hr_upscale_to_y == self.height) or self.hr_upscaler is None or self.hr_upscaler == 'None': # special case: the user has chosen to do nothing
|
||||
self.is_hr_pass = False
|
||||
return
|
||||
self.is_hr_pass = True
|
||||
hypertile_set(self, hr=True)
|
||||
shared.state.job_count = 2 * self.n_iter
|
||||
shared.log.debug(f'Init hires: upscaler="{self.hr_upscaler}" sampler="{self.hr_sampler_name}" resize={self.hr_resize_x}x{self.hr_resize_y} upscale={self.hr_upscale_to_x}x{self.hr_upscale_to_y}')
|
||||
|
||||
def sample(self, conditioning, unconditional_conditioning, seeds, subseeds, subseed_strength, prompts):
|
||||
from modules import processing_original
|
||||
|
||||
@@ -46,11 +46,11 @@ def soft_clamp_tensor(tensor, threshold=0.8, boundary=4):
|
||||
def center_tensor(tensor, channel_shift=0.0, full_shift=0.0, offset=0.0):
|
||||
if channel_shift == 0 and full_shift == 0 and offset == 0:
|
||||
return tensor
|
||||
debug(f'HDR center: Before Adjustment: Full mean={tensor.mean().item()} Channel means={tensor.mean(dim=(-1, -2)).cpu().numpy()}')
|
||||
debug(f'HDR center: Before Adjustment: Full mean={tensor.mean().item()} Channel means={tensor.mean(dim=(-1, -2)).float().cpu().numpy()}')
|
||||
tensor -= tensor.mean(dim=(-1, -2), keepdim=True) * channel_shift
|
||||
tensor -= tensor.mean() * full_shift - offset
|
||||
debug(f'HDR center: channel-shift={channel_shift} full-shift={full_shift}')
|
||||
debug(f'HDR center: After Adjustment: Full mean={tensor.mean().item()} Channel means={tensor.mean(dim=(-1, -2)).cpu().numpy()}')
|
||||
debug(f'HDR center: After Adjustment: Full mean={tensor.mean().item()} Channel means={tensor.mean(dim=(-1, -2)).float().cpu().numpy()}')
|
||||
return tensor
|
||||
|
||||
|
||||
@@ -122,9 +122,9 @@ def correction_callback(p, timestep, kwargs):
|
||||
for i in range(latents.shape[0]):
|
||||
latents[i] = correction(p, timestep, latents[i])
|
||||
debug(f"Full Mean: {latents[i].mean().item()}")
|
||||
debug(f"Channel Means: {latents[i].mean(dim=(-1, -2), keepdim=True).flatten().cpu().numpy()}")
|
||||
debug(f"Channel Mins: {latents[i].min(-1, keepdim=True)[0].min(-2, keepdim=True)[0].flatten().cpu().numpy()}")
|
||||
debug(f"Channel Maxes: {latents[i].max(-1, keepdim=True)[0].min(-2, keepdim=True)[0].flatten().cpu().numpy()}")
|
||||
debug(f"Channel Means: {latents[i].mean(dim=(-1, -2), keepdim=True).flatten().float().cpu().numpy()}")
|
||||
debug(f"Channel Mins: {latents[i].min(-1, keepdim=True)[0].min(-2, keepdim=True)[0].flatten().float().cpu().numpy()}")
|
||||
debug(f"Channel Maxes: {latents[i].max(-1, keepdim=True)[0].min(-2, keepdim=True)[0].flatten().float().cpu().numpy()}")
|
||||
elif len(latents.shape) == 5 and latents.shape[0] == 1: # probably animatediff
|
||||
latents = latents.squeeze(0).permute(1, 0, 2, 3)
|
||||
for i in range(latents.shape[0]):
|
||||
|
||||
@@ -37,6 +37,15 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
|
||||
for j in range(len(decoded)):
|
||||
images.save_image(decoded[j], path=p.outpath_samples, basename="", seed=p.seeds[i], prompt=p.prompts[i], extension=shared.opts.samples_format, info=info, p=p, suffix=suffix)
|
||||
|
||||
def apply_circular(enable):
|
||||
try:
|
||||
for layer in [layer for layer in shared.sd_model.unet.modules() if type(layer) is torch.nn.Conv2d]:
|
||||
layer.padding_mode = 'circular' if enable else 'zeros'
|
||||
for layer in [layer for layer in shared.sd_model.vae.modules() if type(layer) is torch.nn.Conv2d]:
|
||||
layer.padding_mode = 'circular' if enable else 'zeros'
|
||||
except Exception as e:
|
||||
debug(f"Diffusers tiling failed: {e}")
|
||||
|
||||
def diffusers_callback_legacy(step: int, timestep: int, latents: typing.Union[torch.FloatTensor, np.ndarray]):
|
||||
if isinstance(latents, np.ndarray): # latents from Onnx pipelines is ndarray.
|
||||
latents = torch.from_numpy(latents)
|
||||
@@ -159,6 +168,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
|
||||
|
||||
def set_pipeline_args(model, prompts: list, negative_prompts: list, prompts_2: typing.Optional[list]=None, negative_prompts_2: typing.Optional[list]=None, desc:str='', **kwargs):
|
||||
t0 = time.time()
|
||||
apply_circular(p.tiling)
|
||||
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 = {}
|
||||
@@ -219,7 +229,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
|
||||
if 'latents' in possible and getattr(p, "init_latent", None) is not None:
|
||||
args['latents'] = p.init_latent
|
||||
if 'output_type' in possible:
|
||||
if hasattr(model, 'vae'):
|
||||
if not hasattr(model, 'vae'):
|
||||
args['output_type'] = 'np' # only set latent if model has vae
|
||||
|
||||
# stable cascade
|
||||
@@ -345,7 +355,6 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
|
||||
p.task_args['sag_scale'] = p.sag_scale
|
||||
else:
|
||||
shared.log.warning(f'SAG incompatible scheduler: current={sd_model.scheduler.__class__.__name__} supported={supported}')
|
||||
|
||||
if shared.opts.cuda_compile_backend == "olive-ai":
|
||||
sd_model = olive_check_parameters_changed(p, is_refiner_enabled())
|
||||
if sd_model.__class__.__name__ == "OnnxRawPipeline":
|
||||
@@ -362,12 +371,6 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
|
||||
shared.sd_model = orig_pipeline
|
||||
return results
|
||||
|
||||
if shared.opts.diffusers_move_base:
|
||||
sd_models.move_model(shared.sd_model, devices.device)
|
||||
|
||||
# recompile if a parameter changes
|
||||
sd_models_compile.openvino_recompile_model(p, hires=False, refiner=False)
|
||||
|
||||
# pipeline type is set earlier in processing, but check for sanity
|
||||
is_control = getattr(p, 'is_control', False) is True
|
||||
has_images = len(getattr(p, 'init_images' ,[])) > 0
|
||||
@@ -378,10 +381,14 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
|
||||
if len(getattr(p, 'init_images' ,[])) == 0:
|
||||
p.init_images = [TF.to_pil_image(torch.rand((3, getattr(p, 'height', 512), getattr(p, 'width', 512))))]
|
||||
|
||||
sd_models.move_model(shared.sd_model, devices.device)
|
||||
sd_models_compile.openvino_recompile_model(p, hires=False, refiner=False) # recompile if a parameter changes
|
||||
|
||||
use_refiner_start = is_txt2img() and is_refiner_enabled() and not p.is_hr_pass and p.refiner_start > 0 and p.refiner_start < 1
|
||||
use_denoise_start = not is_txt2img() and p.refiner_start > 0 and p.refiner_start < 1
|
||||
|
||||
shared.sd_model = update_pipeline(shared.sd_model, p)
|
||||
shared.log.info(f'Base: class={shared.sd_model.__class__.__name__}')
|
||||
base_args = set_pipeline_args(
|
||||
model=shared.sd_model,
|
||||
prompts=p.prompts,
|
||||
@@ -442,54 +449,65 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
|
||||
shared.sd_model = orig_pipeline
|
||||
return results
|
||||
|
||||
# optional hires pass
|
||||
if p.enable_hr and getattr(p, 'hr_upscaler', 'None') != 'None' and len(getattr(p, 'init_images', [])) == 0:
|
||||
# optional second pass
|
||||
if p.enable_hr and len(getattr(p, 'init_images', [])) == 0:
|
||||
p.is_hr_pass = True
|
||||
latent_scale_mode = shared.latent_upscale_modes.get(p.hr_upscaler, None) if (hasattr(p, "hr_upscaler") and p.hr_upscaler is not None) else shared.latent_upscale_modes.get(shared.latent_upscale_default_mode, "None")
|
||||
if p.is_hr_pass:
|
||||
p.init_hr()
|
||||
prev_job = shared.state.job
|
||||
if hasattr(p, 'height') and hasattr(p, 'width') and (p.width != p.hr_upscale_to_x or p.height != p.hr_upscale_to_y):
|
||||
|
||||
# upscale
|
||||
if hasattr(p, 'height') and hasattr(p, 'width') and p.hr_upscaler is not None and p.hr_upscaler != 'None':
|
||||
shared.log.info(f'Upscale: upscaler="{p.hr_upscaler}" resize={p.hr_resize_x}x{p.hr_resize_y} upscale={p.hr_upscale_to_x}x{p.hr_upscale_to_y}')
|
||||
p.ops.append('upscale')
|
||||
if shared.opts.save and not p.do_not_save_samples and shared.opts.save_images_before_highres_fix and hasattr(shared.sd_model, 'vae'):
|
||||
save_intermediate(latents=output.images, suffix="-before-hires")
|
||||
shared.state.job = 'upscale'
|
||||
output.images = resize_hires(p, latents=output.images)
|
||||
if (latent_scale_mode is not None or p.hr_force) and p.denoising_strength > 0:
|
||||
p.ops.append('hires')
|
||||
sd_models_compile.openvino_recompile_model(p, hires=True, refiner=False)
|
||||
shared.sd_model = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.IMAGE_2_IMAGE)
|
||||
if shared.sd_model.__class__.__name__ == "OnnxRawPipeline":
|
||||
shared.sd_model = preprocess_onnx_pipeline(p)
|
||||
update_sampler(shared.sd_model, second_pass=True)
|
||||
hires_args = set_pipeline_args(
|
||||
model=shared.sd_model,
|
||||
prompts=[p.refiner_prompt] if len(p.refiner_prompt) > 0 else p.prompts,
|
||||
negative_prompts=[p.refiner_negative] if len(p.refiner_negative) > 0 else p.negative_prompts,
|
||||
prompts_2=[p.refiner_prompt] if len(p.refiner_prompt) > 0 else p.prompts,
|
||||
negative_prompts_2=[p.refiner_negative] if len(p.refiner_negative) > 0 else p.negative_prompts,
|
||||
num_inference_steps=calculate_hires_steps(p),
|
||||
eta=shared.opts.scheduler_eta,
|
||||
guidance_scale=p.image_cfg_scale if p.image_cfg_scale is not None else p.cfg_scale,
|
||||
guidance_rescale=p.diffusers_guidance_rescale,
|
||||
output_type='latent' if hasattr(shared.sd_model, 'vae') else 'np',
|
||||
clip_skip=p.clip_skip,
|
||||
image=output.images,
|
||||
strength=p.denoising_strength,
|
||||
desc='Hires',
|
||||
)
|
||||
shared.state.job = 'hires'
|
||||
shared.state.sampling_steps = hires_args['num_inference_steps']
|
||||
try:
|
||||
sd_models_compile.check_deepcache(enable=True)
|
||||
output = shared.sd_model(**hires_args) # pylint: disable=not-callable
|
||||
if isinstance(output, dict):
|
||||
output = SimpleNamespace(**output)
|
||||
sd_models_compile.check_deepcache(enable=False)
|
||||
sd_models_compile.openvino_post_compile(op="base")
|
||||
except AssertionError as e:
|
||||
shared.log.info(e)
|
||||
p.init_images = []
|
||||
sd_hijack_hypertile.hypertile_set(p, hr=True)
|
||||
|
||||
latent_upscale = shared.latent_upscale_modes.get(p.hr_upscaler, None)
|
||||
if (latent_upscale is not None or p.hr_force) and p.denoising_strength > 0:
|
||||
p.ops.append('hires')
|
||||
sd_models_compile.openvino_recompile_model(p, hires=True, refiner=False)
|
||||
if shared.sd_model.__class__.__name__ == "OnnxRawPipeline":
|
||||
shared.sd_model = preprocess_onnx_pipeline(p)
|
||||
p.hr_force = True
|
||||
|
||||
# hires
|
||||
if p.hr_force:
|
||||
shared.state.job_count = 2 * p.n_iter
|
||||
shared.sd_model = sd_models.set_diffuser_pipe(shared.sd_model, sd_models.DiffusersTaskType.IMAGE_2_IMAGE)
|
||||
update_sampler(shared.sd_model, second_pass=True)
|
||||
shared.log.info(f'HiRes: class={shared.sd_model.__class__.__name__} sampler="{p.hr_sampler_name}"')
|
||||
hires_args = set_pipeline_args(
|
||||
model=shared.sd_model,
|
||||
prompts=[p.refiner_prompt] if len(p.refiner_prompt) > 0 else p.prompts,
|
||||
negative_prompts=[p.refiner_negative] if len(p.refiner_negative) > 0 else p.negative_prompts,
|
||||
prompts_2=[p.refiner_prompt] if len(p.refiner_prompt) > 0 else p.prompts,
|
||||
negative_prompts_2=[p.refiner_negative] if len(p.refiner_negative) > 0 else p.negative_prompts,
|
||||
num_inference_steps=calculate_hires_steps(p),
|
||||
eta=shared.opts.scheduler_eta,
|
||||
guidance_scale=p.image_cfg_scale if p.image_cfg_scale is not None else p.cfg_scale,
|
||||
guidance_rescale=p.diffusers_guidance_rescale,
|
||||
output_type='latent' if hasattr(shared.sd_model, 'vae') else 'np',
|
||||
clip_skip=p.clip_skip,
|
||||
image=output.images,
|
||||
strength=p.denoising_strength,
|
||||
desc='Hires',
|
||||
)
|
||||
shared.state.job = 'hires'
|
||||
shared.state.sampling_steps = hires_args['num_inference_steps']
|
||||
try:
|
||||
sd_models_compile.check_deepcache(enable=True)
|
||||
output = shared.sd_model(**hires_args) # pylint: disable=not-callable
|
||||
if isinstance(output, dict):
|
||||
output = SimpleNamespace(**output)
|
||||
sd_models_compile.check_deepcache(enable=False)
|
||||
sd_models_compile.openvino_post_compile(op="base")
|
||||
except AssertionError as e:
|
||||
shared.log.info(e)
|
||||
p.init_images = []
|
||||
shared.state.job = prev_job
|
||||
shared.state.nextjob()
|
||||
p.is_hr_pass = False
|
||||
@@ -523,6 +541,7 @@ def process_diffusers(p: processing.StableDiffusionProcessing):
|
||||
image = processing_vae.vae_decode(latents=image, model=shared.sd_model, full_quality=p.full_quality, output_type='pil')
|
||||
p.extra_generation_params['Noise level'] = noise_level
|
||||
output_type = 'np'
|
||||
shared.log.info(f'Refiner: class={shared.sd_refiner.__class__.__name__}')
|
||||
refiner_args = set_pipeline_args(
|
||||
model=shared.sd_refiner,
|
||||
prompts=[p.refiner_prompt] if len(p.refiner_prompt) > 0 else p.prompts[i],
|
||||
|
||||
@@ -352,7 +352,7 @@ def resize_hires(p, latents): # input=latents output=pil
|
||||
first_pass_images = processing_vae.vae_decode(latents=latents, model=shared.sd_model, full_quality=p.full_quality, output_type='pil')
|
||||
return first_pass_images
|
||||
latent_upscaler = shared.latent_upscale_modes.get(p.hr_upscaler, None)
|
||||
shared.log.info(f'Hires: upscaler={p.hr_upscaler} width={p.hr_upscale_to_x} height={p.hr_upscale_to_y} images={latents.shape[0]}')
|
||||
# shared.log.info(f'Hires: upscaler={p.hr_upscaler} width={p.hr_upscale_to_x} height={p.hr_upscale_to_y} images={latents.shape[0]}')
|
||||
if latent_upscaler is not None:
|
||||
latents = torch.nn.functional.interpolate(latents, size=(p.hr_upscale_to_y // 8, p.hr_upscale_to_x // 8), mode=latent_upscaler["mode"], antialias=latent_upscaler["antialias"])
|
||||
first_pass_images = processing_vae.vae_decode(latents=latents, model=shared.sd_model, full_quality=p.full_quality, output_type='pil')
|
||||
|
||||
@@ -72,10 +72,10 @@ def create_infotext(p: StableDiffusionProcessing, all_prompts=None, all_seeds=No
|
||||
args["Second pass"] = p.enable_hr
|
||||
args["Hires force"] = p.hr_force
|
||||
args["Hires steps"] = p.hr_second_pass_steps
|
||||
args["Hires upscaler"] = p.hr_upscaler
|
||||
args["Hires upscale"] = p.hr_scale
|
||||
args["Hires resize"] = f"{p.hr_resize_x}x{p.hr_resize_y}"
|
||||
args["Hires size"] = f"{p.hr_upscale_to_x}x{p.hr_upscale_to_y}"
|
||||
args["Hires upscaler"] = p.hr_upscaler if p.hr_upscaler is not None and p.hr_upscaler != 'None' else None
|
||||
args["Hires upscale"] = p.hr_scale if p.hr_upscaler is not None and p.hr_upscaler != 'None' else None
|
||||
args["Hires resize"] = f"{p.hr_resize_x}x{p.hr_resize_y}" if p.hr_upscaler is not None and p.hr_upscaler != 'None' else None
|
||||
args["Hires size"] = f"{p.hr_upscale_to_x}x{p.hr_upscale_to_y}" if p.hr_upscaler is not None and p.hr_upscaler != 'None' else None
|
||||
args["Denoising strength"] = p.denoising_strength
|
||||
args["Hires sampler"] = p.hr_sampler_name
|
||||
args["Image CFG scale"] = p.image_cfg_scale
|
||||
|
||||
@@ -46,9 +46,20 @@ def full_vae_decode(latents, model):
|
||||
model.upcast_vae()
|
||||
if hasattr(model.vae, "post_quant_conv"):
|
||||
latents = latents.to(next(iter(model.vae.post_quant_conv.parameters())).dtype)
|
||||
decoded = model.vae.decode(latents / model.vae.config.scaling_factor, return_dict=False)[0]
|
||||
|
||||
# Delete PyTorch VAE after OpenVINO compile
|
||||
# normalize latents
|
||||
latents_mean = model.vae.config.get("latents_mean", None)
|
||||
latents_std = model.vae.config.get("latents_std", None)
|
||||
scaling_factor = model.vae.config.get("scaling_factor", None)
|
||||
if latents_mean and latents_std:
|
||||
latents_mean = (torch.tensor(latents_mean).view(1, 4, 1, 1).to(latents.device, latents.dtype))
|
||||
latents_std = (torch.tensor(latents_std).view(1, 4, 1, 1).to(latents.device, latents.dtype))
|
||||
latents = latents * latents_std / scaling_factor + latents_mean
|
||||
else:
|
||||
latents = latents / scaling_factor
|
||||
decoded = model.vae.decode(latents, return_dict=False)[0]
|
||||
|
||||
# delete vae after OpenVINO compile
|
||||
if shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx" and shared.compiled_model_state.first_pass_vae:
|
||||
shared.compiled_model_state.first_pass_vae = False
|
||||
if not shared.opts.openvino_disable_memory_cleanup and hasattr(shared.sd_model, "vae"):
|
||||
|
||||
@@ -109,7 +109,6 @@ callback_map = dict(
|
||||
callbacks_after_process=[],
|
||||
callbacks_model_loaded=[],
|
||||
callbacks_ui_tabs=[],
|
||||
callbacks_ui_train_tabs=[],
|
||||
callbacks_ui_settings=[],
|
||||
callbacks_before_image_saved=[],
|
||||
callbacks_image_saved=[],
|
||||
@@ -212,16 +211,6 @@ def ui_tabs_callback():
|
||||
return res
|
||||
|
||||
|
||||
def ui_train_tabs_callback(params: UiTrainTabParams):
|
||||
for c in callback_map['callbacks_ui_train_tabs']:
|
||||
try:
|
||||
t0 = time.time()
|
||||
c.callback(params)
|
||||
timer(t0, c.script, 'ui_train_tabs')
|
||||
except Exception as e:
|
||||
report_exception(e, c, 'callbacks_ui_train_tabs')
|
||||
|
||||
|
||||
def ui_settings_callback():
|
||||
for c in callback_map['callbacks_ui_settings']:
|
||||
try:
|
||||
@@ -434,13 +423,6 @@ def on_ui_tabs(callback):
|
||||
add_callback(callback_map['callbacks_ui_tabs'], callback)
|
||||
|
||||
|
||||
def on_ui_train_tabs(callback):
|
||||
"""register a function to be called when the UI is creating new tabs for the train tab.
|
||||
Create your new tabs with gr.Tab.
|
||||
"""
|
||||
add_callback(callback_map['callbacks_ui_train_tabs'], callback)
|
||||
|
||||
|
||||
def on_ui_settings(callback):
|
||||
"""register a function to be called before UI settings are populated; add your settings
|
||||
by using shared.opts.add_option(shared.OptionInfo(...)) """
|
||||
|
||||
@@ -56,11 +56,11 @@ def apply_optimizations():
|
||||
if can_use_sdp and shared.opts.cross_attention_optimization == "Scaled-Dot-Product":
|
||||
optimization_method = 'sdp'
|
||||
if 'Memory attention' in shared.opts.sdp_options:
|
||||
ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.scaled_dot_product_no_mem_attention_forward
|
||||
ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.sdp_no_mem_attnblock_forward
|
||||
else:
|
||||
ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.scaled_dot_product_attention_forward
|
||||
ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.sdp_attnblock_forward
|
||||
else:
|
||||
ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.scaled_dot_product_no_mem_attention_forward
|
||||
ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.sdp_no_mem_attnblock_forward
|
||||
if shared.xformers_available and shared.opts.cross_attention_optimization == "xFormers":
|
||||
ldm.modules.attention.CrossAttention.forward = sd_hijack_optimizations.xformers_attention_forward
|
||||
ldm.modules.diffusionmodules.model.AttnBlock.forward = sd_hijack_optimizations.xformers_attnblock_forward
|
||||
|
||||
@@ -22,6 +22,7 @@ error_reported = False
|
||||
reset_needed = False
|
||||
skip_hypertile = False
|
||||
|
||||
|
||||
def iterative_closest_divisors(hw:int, aspect_ratio:float) -> tuple[int, int]:
|
||||
"""
|
||||
Finds h and w such that h*w = hw and h/w = aspect_ratio
|
||||
@@ -34,6 +35,7 @@ def iterative_closest_divisors(hw:int, aspect_ratio:float) -> tuple[int, int]:
|
||||
closest_pair = pairs[ratios.index(closest_ratio)] # closest pair of divisors to aspect_ratio
|
||||
return closest_pair
|
||||
|
||||
|
||||
@cache
|
||||
def find_hw_candidates(hw:int, aspect_ratio:float) -> tuple[int, int]:
|
||||
"""
|
||||
@@ -96,7 +98,7 @@ def split_attention(layer: nn.Module, tile_size: int=256, min_tile_size: int=128
|
||||
def self_attn_forward(forward: Callable) -> Callable:
|
||||
@wraps(forward)
|
||||
def wrapper(*args, **kwargs):
|
||||
global height, width, max_h, max_w, reset_needed, error_reported, skip_hypertile # pylint: disable=global-statement
|
||||
global height, width, max_h, max_w, reset_needed, error_reported # pylint: disable=global-statement
|
||||
if skip_hypertile:
|
||||
return forward(*args, **kwargs)
|
||||
x = args[0]
|
||||
@@ -198,7 +200,6 @@ def context_hypertile_vae(p):
|
||||
return split_attention(vae, tile_size=tile_size, min_tile_size=128, swap_size=shared.opts.hypertile_vae_swap_size)
|
||||
|
||||
|
||||
|
||||
def context_hypertile_unet(p):
|
||||
from modules import shared
|
||||
if p.sd_model is None or not shared.opts.hypertile_unet_enabled:
|
||||
|
||||
+126
-57
@@ -323,37 +323,40 @@ def read_metadata_from_safetensors(filename):
|
||||
# try:
|
||||
t0 = time.time()
|
||||
with open(filename, mode="rb") as file:
|
||||
metadata_len = file.read(8)
|
||||
metadata_len = int.from_bytes(metadata_len, "little")
|
||||
json_start = file.read(2)
|
||||
if metadata_len <= 2 or json_start not in (b'{"', b"{'"):
|
||||
shared.log.error(f"Not a valid safetensors file: {filename}")
|
||||
json_data = json_start + file.read(metadata_len-2)
|
||||
json_obj = json.loads(json_data)
|
||||
for k, v in json_obj.get("__metadata__", {}).items():
|
||||
if v.startswith("data:"):
|
||||
v = 'data'
|
||||
if k == 'format' and v == 'pt':
|
||||
continue
|
||||
large = True if len(v) > 2048 else False
|
||||
if large and k == 'ss_datasets':
|
||||
continue
|
||||
if large and k == 'workflow':
|
||||
continue
|
||||
if large and k == 'prompt':
|
||||
continue
|
||||
if large and k == 'ss_bucket_info':
|
||||
continue
|
||||
if v[0:1] == '{':
|
||||
try:
|
||||
v = json.loads(v)
|
||||
if large and k == 'ss_tag_frequency':
|
||||
v = { i: len(j) for i, j in v.items() }
|
||||
if large and k == 'sd_merge_models':
|
||||
scrub_dict(v, ['sd_merge_recipe'])
|
||||
except Exception:
|
||||
pass
|
||||
res[k] = v
|
||||
try:
|
||||
metadata_len = file.read(8)
|
||||
metadata_len = int.from_bytes(metadata_len, "little")
|
||||
json_start = file.read(2)
|
||||
if metadata_len <= 2 or json_start not in (b'{"', b"{'"):
|
||||
shared.log.error(f"Model metadata invalid: fn={filename}")
|
||||
json_data = json_start + file.read(metadata_len-2)
|
||||
json_obj = json.loads(json_data)
|
||||
for k, v in json_obj.get("__metadata__", {}).items():
|
||||
if v.startswith("data:"):
|
||||
v = 'data'
|
||||
if k == 'format' and v == 'pt':
|
||||
continue
|
||||
large = True if len(v) > 2048 else False
|
||||
if large and k == 'ss_datasets':
|
||||
continue
|
||||
if large and k == 'workflow':
|
||||
continue
|
||||
if large and k == 'prompt':
|
||||
continue
|
||||
if large and k == 'ss_bucket_info':
|
||||
continue
|
||||
if v[0:1] == '{':
|
||||
try:
|
||||
v = json.loads(v)
|
||||
if large and k == 'ss_tag_frequency':
|
||||
v = { i: len(j) for i, j in v.items() }
|
||||
if large and k == 'sd_merge_models':
|
||||
scrub_dict(v, ['sd_merge_recipe'])
|
||||
except Exception:
|
||||
pass
|
||||
res[k] = v
|
||||
except Exception as e:
|
||||
shared.log.error(f"Model metadata: fn={filename} {e}")
|
||||
sd_metadata[filename] = res
|
||||
global sd_metadata_pending # pylint: disable=global-statement
|
||||
sd_metadata_pending += 1
|
||||
@@ -553,14 +556,16 @@ def detect_pipeline(f: str, op: str = 'model', warning=True):
|
||||
elif (size >= 316 and size <= 324) or (size >= 156 and size <= 164): # 320 or 160
|
||||
warn(f'Model detected as VAE model, but attempting to load as model: {op}={f} size={size} MB')
|
||||
guess = 'VAE'
|
||||
elif size >= 5351 and size <= 5359: # 5353
|
||||
guess = 'Stable Diffusion' # SD v2
|
||||
elif size >= 4970 and size <= 4976: # 4973
|
||||
guess = 'Stable Diffusion 2' # SD v2 but could be eps or v-prediction
|
||||
# elif size < 0: # unknown
|
||||
# guess = 'Stable Diffusion 2B'
|
||||
elif size >= 5791 and size <= 5799: # 5795
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
warn(f'Model detected as SD-XL refiner model, but attempting to load using backend=original: {op}={f} size={size} MB')
|
||||
if op == 'model':
|
||||
warn(f'Model detected as SD-XL refiner model, but attempting to load a base model: {op}={f} size={size} MB')
|
||||
guess = 'Stable Diffusion XL'
|
||||
guess = 'Stable Diffusion XL Refiner'
|
||||
elif (size >= 6611 and size <= 7220): # 6617, HassakuXL is 6776, monkrenRealisticINT_v10 is 7217
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
warn(f'Model detected as SD-XL base model, but attempting to load using backend=original: {op}={f} size={size} MB')
|
||||
@@ -740,8 +745,16 @@ def set_diffuser_options(sd_model, vae = None, op: str = 'model'):
|
||||
sd_model.unet.to(memory_format=torch.channels_last)
|
||||
|
||||
|
||||
def move_model(model, device=None):
|
||||
if model is not None and not getattr(model, 'has_accelerate', False):
|
||||
def move_model(model, device=None, force=False):
|
||||
if model is not None:
|
||||
if getattr(model, 'vae', None) is not None and get_diffusers_task(model) != DiffusersTaskType.TEXT_2_IMAGE:
|
||||
if device == devices.device: # force vae back to gpu if not in txt2img mode
|
||||
model.vae.to(device)
|
||||
if hasattr(model.vae, '_hf_hook'):
|
||||
debug_move(f'Model move: to={device} class={model.vae.__class__} function={sys._getframe(1).f_code.co_name}') # pylint: disable=protected-access
|
||||
model.vae._hf_hook.execution_device = device # pylint: disable=protected-access
|
||||
if getattr(model, 'has_accelerate', False) and not force:
|
||||
return
|
||||
debug_move(f'Model move: to={device} class={model.__class__} function={sys._getframe(1).f_code.co_name}') # pylint: disable=protected-access
|
||||
try:
|
||||
model.to(device)
|
||||
@@ -750,6 +763,67 @@ def move_model(model, device=None):
|
||||
devices.torch_gc()
|
||||
|
||||
|
||||
def get_load_config(model_file, model_type):
|
||||
yaml = os.path.splitext(model_file)[0] + '.yaml'
|
||||
if os.path.exists(yaml):
|
||||
return yaml
|
||||
elif model_type == 'Stable Diffusion':
|
||||
return 'configs/v1-inference.yaml'
|
||||
elif model_type == 'Stable Diffusion XL':
|
||||
return 'configs/sd_xl_base.yaml'
|
||||
elif model_type == 'Stable Diffusion XL Refiner':
|
||||
return 'configs/sd_xl_refiner.yaml'
|
||||
elif model_type == 'Stable Diffusion 2':
|
||||
return None # dont know if its eps or v so let diffusers sort it out
|
||||
# return 'configs/v2-inference-512-base.yaml'
|
||||
# return 'configs/v2-inference-768-v.yaml'
|
||||
return None
|
||||
|
||||
|
||||
def patch_diffuser_config(sd_model, model_file):
|
||||
def load_config(fn, k):
|
||||
model_file = os.path.splitext(fn)[0]
|
||||
cfg_file = f'{model_file}_{k}.json'
|
||||
try:
|
||||
if os.path.exists(cfg_file):
|
||||
with open(cfg_file, 'r', encoding='utf-8') as f:
|
||||
return json.load(f)
|
||||
cfg_file = f'{os.path.join(paths.sd_configs_path, os.path.basename(model_file))}_{k}.json'
|
||||
if os.path.exists(cfg_file):
|
||||
with open(cfg_file, 'r', encoding='utf-8') as f:
|
||||
return json.load(f)
|
||||
except Exception:
|
||||
pass
|
||||
return {}
|
||||
|
||||
if sd_model is None:
|
||||
return sd_model
|
||||
if hasattr(sd_model, 'unet') and hasattr(sd_model.unet, 'config') and 'inpaint' in model_file.lower():
|
||||
if debug_load:
|
||||
shared.log.debug('Model config patch: type=inpaint')
|
||||
sd_model.unet.config.in_channels = 9
|
||||
if not hasattr(sd_model, '_internal_dict'):
|
||||
return sd_model
|
||||
for c in sd_model._internal_dict.keys(): # pylint: disable=protected-access
|
||||
component = getattr(sd_model, c, None)
|
||||
if hasattr(component, 'config'):
|
||||
if debug_load:
|
||||
shared.log.debug(f'Model config: component={c} config={component.config}')
|
||||
override = load_config(model_file, c)
|
||||
updated = {}
|
||||
for k, v in override.items():
|
||||
if k.startswith('_'):
|
||||
continue
|
||||
if v != component.config.get(k, None):
|
||||
if hasattr(component.config, '__frozen'):
|
||||
component.config.__frozen = False # pylint: disable=protected-access
|
||||
component.config[k] = v
|
||||
updated[k] = v
|
||||
if updated and debug_load:
|
||||
shared.log.debug(f'Model config: component={c} override={updated}')
|
||||
return sd_model
|
||||
|
||||
|
||||
def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=None, op='model'): # pylint: disable=unused-argument
|
||||
if shared.cmd_opts.profile:
|
||||
import cProfile
|
||||
@@ -826,10 +900,13 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
if model_type in ['Stable Cascade']: # forced pipeline
|
||||
# TODO experimental stable cascade
|
||||
try:
|
||||
shared.log.debug(f'StableCascade experimental: args={diffusers_load_config} device={devices.device} dtype={devices.dtype}')
|
||||
diffusers_load_config.pop("vae", None)
|
||||
diffusers_load_config.pop("variant", None)
|
||||
decoder = diffusers.StableCascadeDecoderPipeline.from_pretrained("stabilityai/stable-cascade", cache_dir=shared.opts.diffusers_dir, revision="refs/pr/17", **diffusers_load_config)
|
||||
shared.log.debug(f'StableCascade decoder: scale={decoder.latent_dim_scale}')
|
||||
prior = diffusers.StableCascadePriorPipeline.from_pretrained("stabilityai/stable-cascade-prior", cache_dir=shared.opts.diffusers_dir, **diffusers_load_config)
|
||||
shared.log.debug(f'StableCascade prior: scale={prior.resolution_multiple}')
|
||||
sd_model = diffusers.StableCascadeCombinedPipeline(
|
||||
tokenizer=decoder.tokenizer,
|
||||
text_encoder=decoder.text_encoder,
|
||||
@@ -837,9 +914,12 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
scheduler=decoder.scheduler,
|
||||
vqgan=decoder.vqgan,
|
||||
prior_prior=prior.prior,
|
||||
prior_text_encoder=prior.text_encoder,
|
||||
prior_tokenizer=prior.tokenizer,
|
||||
prior_scheduler=prior.scheduler,
|
||||
feature_extractor=prior.feature_extractor,
|
||||
image_encoder=prior.image_encoder)
|
||||
prior_prior_feature_extractor=prior.feature_extractor,
|
||||
prior_prior_image_encoder=prior.image_encoder)
|
||||
shared.log.debug(f'StableCascade combined: {sd_model.__class__.__name__}')
|
||||
except Exception as e:
|
||||
shared.log.error(f'Diffusers Failed loading {op}: {checkpoint_info.path} {e}')
|
||||
if debug_load:
|
||||
@@ -918,20 +998,8 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
diffusers_load_config['force_zeros_for_empty_prompt '] = shared.opts.diffusers_force_zeros
|
||||
if shared.opts.diffusers_aesthetics_score:
|
||||
diffusers_load_config['requires_aesthetics_score'] = shared.opts.diffusers_aesthetics_score
|
||||
if 'inpainting' in checkpoint_info.path.lower():
|
||||
diffusers_load_config['config_files'] = {
|
||||
'v1': 'configs/v1-inpainting-inference.yaml',
|
||||
'v2': 'configs/v2-inference-768-v.yaml',
|
||||
'xl': 'configs/sd_xl_base.yaml',
|
||||
'xl_refiner': 'configs/sd_xl_refiner.yaml',
|
||||
}
|
||||
else:
|
||||
diffusers_load_config['config_files'] = {
|
||||
'v1': 'configs/v1-inference.yaml',
|
||||
'v2': 'configs/v2-inference-768-v.yaml',
|
||||
'xl': 'configs/sd_xl_base.yaml',
|
||||
'xl_refiner': 'configs/sd_xl_refiner.yaml',
|
||||
}
|
||||
# diffusers_load_config['config_files'] = get_load_config(checkpoint_info.path.lower())
|
||||
diffusers_load_config['original_config_file'] = get_load_config(checkpoint_info.path, model_type)
|
||||
if hasattr(pipeline, 'from_single_file'):
|
||||
diffusers_load_config['use_safetensors'] = True
|
||||
if shared.opts.disable_accelerate:
|
||||
@@ -942,9 +1010,7 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
else:
|
||||
sd_hijack_accelerate.restore_accelerate()
|
||||
sd_model = pipeline.from_single_file(checkpoint_info.path, **diffusers_load_config)
|
||||
if sd_model is not None and hasattr(sd_model, 'unet') and hasattr(sd_model.unet, 'config') and 'inpainting' in checkpoint_info.path.lower():
|
||||
shared.log.debug('Model patch: type=inpaint')
|
||||
sd_model.unet.config.in_channels = 9
|
||||
sd_model = patch_diffuser_config(sd_model, checkpoint_info.path)
|
||||
elif hasattr(pipeline, 'from_ckpt'):
|
||||
sd_model = pipeline.from_ckpt(checkpoint_info.path, **diffusers_load_config)
|
||||
else:
|
||||
@@ -1115,7 +1181,8 @@ def switch_pipe(cls: diffusers.DiffusionPipeline, pipeline: diffusers.DiffusionP
|
||||
unet=pipeline.unet,
|
||||
scheduler=pipeline.scheduler,
|
||||
feature_extractor=getattr(pipeline, 'feature_extractor', None),
|
||||
).to(pipeline.device)
|
||||
)
|
||||
move_model(new_pipe, pipeline.device)
|
||||
switch_mode = 'sdxl'
|
||||
elif 'tokenizer' in possible and hasattr(pipeline, 'tokenizer'):
|
||||
new_pipe = cls(
|
||||
@@ -1127,7 +1194,8 @@ def switch_pipe(cls: diffusers.DiffusionPipeline, pipeline: diffusers.DiffusionP
|
||||
feature_extractor=getattr(pipeline, 'feature_extractor', None),
|
||||
requires_safety_checker=False,
|
||||
safety_checker=None,
|
||||
).to(pipeline.device)
|
||||
)
|
||||
move_model(new_pipe, pipeline.device)
|
||||
switch_mode = 'sd'
|
||||
else:
|
||||
shared.log.error(f'Pipeline switch error: {pipeline.__class__.__name__} unrecognized')
|
||||
@@ -1416,6 +1484,7 @@ def convert_to_faketensors(tensor):
|
||||
return tensor
|
||||
except Exception:
|
||||
pass
|
||||
return tensor
|
||||
|
||||
|
||||
def disable_offload(sd_model):
|
||||
|
||||
@@ -2,6 +2,7 @@ import os
|
||||
import inspect
|
||||
from modules import shared
|
||||
from modules import sd_samplers_common
|
||||
from modules.tcd import TCDScheduler
|
||||
|
||||
|
||||
debug = shared.log.trace if os.environ.get('SD_SAMPLER_DEBUG', None) is not None else lambda *args, **kwargs: None
|
||||
@@ -50,6 +51,7 @@ config = {
|
||||
'PNDM': { 'skip_prk_steps': False, 'set_alpha_to_one': False, 'steps_offset': 0 },
|
||||
'LCM': { 'beta_start': 0.00085, 'beta_end': 0.012, 'beta_schedule': "scaled_linear", 'set_alpha_to_one': True, 'rescale_betas_zero_snr': False, 'thresholding': False },
|
||||
'SA Solver': {'predictor_order': 2, 'corrector_order': 2, 'thresholding': False, 'lower_order_final': True, 'use_karras_sigmas': False, 'timestep_spacing': 'linspace'},
|
||||
'TCD': { 'set_alpha_to_one': True, 'rescale_betas_zero_snr': False, 'beta_schedule': 'scaled_linear' },
|
||||
}
|
||||
|
||||
samplers_data_diffusers = [
|
||||
@@ -70,8 +72,18 @@ samplers_data_diffusers = [
|
||||
sd_samplers_common.SamplerData('Heun', lambda model: DiffusionSampler('Heun', HeunDiscreteScheduler, model), [], {}),
|
||||
sd_samplers_common.SamplerData('LCM', lambda model: DiffusionSampler('LCM', LCMScheduler, model), [], {}),
|
||||
sd_samplers_common.SamplerData('SA Solver', lambda model: DiffusionSampler('SA Solver', SASolverScheduler, model), [], {}),
|
||||
sd_samplers_common.SamplerData('TCD', lambda model: DiffusionSampler('TCD', TCDScheduler, model), [], {}),
|
||||
]
|
||||
|
||||
try: # diffusers==0.27.0
|
||||
from diffusers import EDMDPMSolverMultistepScheduler, EDMEulerScheduler
|
||||
config['DPM++ 2M EDM'] = { 'solver_order': 2, 'solver_type': 'midpoint', 'final_sigmas_type': 'zero' } # 'algorithm_type': 'dpmsolver++'
|
||||
config['Euler EDM'] = { }
|
||||
samplers_data_diffusers.append(sd_samplers_common.SamplerData('DPM++ 2M EDM', lambda model: DiffusionSampler('DPM++ 2M EDM', EDMDPMSolverMultistepScheduler, model), [], {}))
|
||||
samplers_data_diffusers.append(sd_samplers_common.SamplerData('Euler EDM', lambda model: DiffusionSampler('Euler EDM', EDMEulerScheduler, model), [], {}))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
class DiffusionSampler:
|
||||
def __init__(self, name, constructor, model, **kwargs):
|
||||
@@ -126,6 +138,10 @@ class DiffusionSampler:
|
||||
self.config['algorithm_type'] = shared.opts.schedulers_dpm_solver
|
||||
if name == 'DEIS':
|
||||
self.config['algorithm_type'] = 'deis'
|
||||
if 'EDM' in name:
|
||||
del self.config['beta_start']
|
||||
del self.config['beta_end']
|
||||
del self.config['beta_schedule']
|
||||
# validate all config params
|
||||
signature = inspect.signature(constructor, follow_wrapped=True)
|
||||
possible = signature.parameters.keys()
|
||||
@@ -134,5 +150,7 @@ class DiffusionSampler:
|
||||
if key not in possible:
|
||||
shared.log.warning(f'Sampler: sampler="{name}" config={self.config} invalid={key}')
|
||||
del self.config[key]
|
||||
# shared.log.debug(f'Sampler: sampler="{name}" config={self.config}')
|
||||
self.sampler = constructor(**self.config)
|
||||
# shared.log.debug(f'Sampler: class="{self.sampler.__class__.__name__}" config={self.sampler.config}')
|
||||
self.sampler.name = name
|
||||
|
||||
@@ -507,6 +507,7 @@ options_templates.update(options_section(('saving-images', "Image Options"), {
|
||||
"img_max_size_mp": OptionInfo(250, "Maximum image size (MP)", gr.Slider, {"minimum": 100, "maximum": 2000, "step": 1}),
|
||||
"webp_lossless": OptionInfo(False, "WebP lossless compression"),
|
||||
"save_selected_only": OptionInfo(True, "Save only saves selected image"),
|
||||
"include_mask": OptionInfo(False, "Include mask in outputs"),
|
||||
"samples_save_zip": OptionInfo(True, "Create ZIP archive"),
|
||||
|
||||
"image_sep_metadata": OptionInfo("<h2>Metadata/Logging</h2>", "", gr.HTML),
|
||||
|
||||
@@ -42,11 +42,13 @@ def get_pipelines():
|
||||
pipelines = { # note: not all pipelines can be used manually as they require prior pipeline next to decoder pipeline
|
||||
'Autodetect': None,
|
||||
'Stable Diffusion': getattr(diffusers, 'StableDiffusionPipeline', None),
|
||||
'Stable Diffusion 2': getattr(diffusers, 'StableDiffusionPipeline', None),
|
||||
'Stable Diffusion Inpaint': getattr(diffusers, 'StableDiffusionInpaintPipeline', None),
|
||||
'Stable Diffusion Img2Img': getattr(diffusers, 'StableDiffusionImg2ImgPipeline', None),
|
||||
'Stable Diffusion Instruct': getattr(diffusers, 'StableDiffusionInstructPix2PixPipeline', None),
|
||||
'Stable Diffusion Upscale': getattr(diffusers, 'StableDiffusionUpscalePipeline', None),
|
||||
'Stable Diffusion XL': getattr(diffusers, 'StableDiffusionXLPipeline', None),
|
||||
'Stable Diffusion XL Refiner': getattr(diffusers, 'StableDiffusionXLPipeline', None),
|
||||
'Stable Diffusion XL Img2Img': getattr(diffusers, 'StableDiffusionXLImg2ImgPipeline', None),
|
||||
'Stable Diffusion XL Inpaint': getattr(diffusers, 'StableDiffusionXLInpaintPipeline', None),
|
||||
'Stable Diffusion XL Instruct': getattr(diffusers, 'StableDiffusionXLInstructPix2PixPipeline', None),
|
||||
|
||||
@@ -0,0 +1,660 @@
|
||||
# Copied from: https://github.com/jabir-zheng/TCD/blob/main/scheduling_tcd.py
|
||||
# pylint: skip-file
|
||||
|
||||
# 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 TCDSchedulerOutput(BaseOutput):
|
||||
"""
|
||||
Output class for the scheduler's `step` function output.
|
||||
|
||||
Args:
|
||||
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
|
||||
denoising loop.
|
||||
pred_noised_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||
The predicted noised sample `(x_{s})` based on the model output from the current timestep.
|
||||
"""
|
||||
|
||||
prev_sample: torch.FloatTensor
|
||||
pred_noised_sample: 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 TCDScheduler(SchedulerMixin, ConfigMixin):
|
||||
"""
|
||||
`TCDScheduler` incorporates the `Strategic Stochastic Sampling` introduced by the paper `Trajectory Consistency Distillation`,
|
||||
extending the original Multistep Consistency Sampling to enable unrestricted trajectory traversal.
|
||||
|
||||
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.custom_timesteps = False
|
||||
|
||||
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
|
||||
|
||||
def _get_variance(self, timestep, prev_timestep):
|
||||
alpha_prod_t = self.alphas_cumprod[timestep]
|
||||
alpha_prod_t_prev = self.alphas_cumprod[prev_timestep] if prev_timestep >= 0 else self.final_alpha_cumprod
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
beta_prod_t_prev = 1 - alpha_prod_t_prev
|
||||
|
||||
variance = (beta_prod_t_prev / beta_prod_t) * (1 - alpha_prod_t / alpha_prod_t_prev)
|
||||
|
||||
return variance
|
||||
|
||||
# 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: Optional[int] = None,
|
||||
device: Union[str, torch.device] = None,
|
||||
original_inference_steps: Optional[int] = None,
|
||||
timesteps: Optional[List[int]] = None,
|
||||
strength: int = 1.0,
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
Args:
|
||||
num_inference_steps (`int`, *optional*):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model. If used,
|
||||
`timesteps` must be `None`.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
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.
|
||||
timesteps (`List[int]`, *optional*):
|
||||
Custom timesteps used to support arbitrary spacing between timesteps. If `None`, then the default
|
||||
timestep spacing strategy of equal spacing between timesteps on the training/distillation timestep
|
||||
schedule is used. If `timesteps` is passed, `num_inference_steps` must be `None`.
|
||||
"""
|
||||
# 0. Check inputs
|
||||
if num_inference_steps is None and timesteps is None:
|
||||
raise ValueError("Must pass exactly one of `num_inference_steps` or `custom_timesteps`.")
|
||||
|
||||
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`.")
|
||||
|
||||
# 1. Calculate the TCD original training/distillation timestep schedule.
|
||||
original_steps = (
|
||||
original_inference_steps if original_inference_steps is not None else self.config.original_inference_steps
|
||||
)
|
||||
|
||||
if original_steps is not None:
|
||||
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."
|
||||
)
|
||||
# TCD Timesteps Setting
|
||||
# The skipping step parameter k from the paper.
|
||||
k = self.config.num_train_timesteps // original_steps
|
||||
# TCD Training/Distillation Steps Schedule
|
||||
tcd_origin_timesteps = np.asarray(list(range(1, int(original_steps * strength) + 1))) * k - 1
|
||||
else:
|
||||
tcd_origin_timesteps = np.asarray(list(range(0, int(self.config.num_train_timesteps * strength))))
|
||||
|
||||
# 2. Calculate the TCD inference timestep schedule.
|
||||
if timesteps is not None:
|
||||
# 2.1 Handle custom timestep schedules.
|
||||
train_timesteps = set(tcd_origin_timesteps)
|
||||
non_train_timesteps = []
|
||||
for i in range(1, len(timesteps)):
|
||||
if timesteps[i] >= timesteps[i - 1]:
|
||||
raise ValueError("`custom_timesteps` must be in descending order.")
|
||||
|
||||
if timesteps[i] not in train_timesteps:
|
||||
non_train_timesteps.append(timesteps[i])
|
||||
|
||||
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}."
|
||||
)
|
||||
|
||||
# Raise warning if timestep schedule does not start with self.config.num_train_timesteps - 1
|
||||
if strength == 1.0 and timesteps[0] != self.config.num_train_timesteps - 1:
|
||||
logger.warning(
|
||||
f"The first timestep on the custom timestep schedule is {timesteps[0]}, not"
|
||||
f" `self.config.num_train_timesteps - 1`: {self.config.num_train_timesteps - 1}. You may get"
|
||||
f" unexpected results when using this timestep schedule."
|
||||
)
|
||||
|
||||
# Raise warning if custom timestep schedule contains timesteps not on original timestep schedule
|
||||
if non_train_timesteps:
|
||||
logger.warning(
|
||||
f"The custom timestep schedule contains the following timesteps which are not on the original"
|
||||
f" training/distillation timestep schedule: {non_train_timesteps}. You may get unexpected results"
|
||||
f" when using this timestep schedule."
|
||||
)
|
||||
|
||||
# Raise warning if custom timestep schedule is longer than original_steps
|
||||
if original_steps is not None:
|
||||
if len(timesteps) > original_steps:
|
||||
logger.warning(
|
||||
f"The number of timesteps in the custom timestep schedule is {len(timesteps)}, which exceeds the"
|
||||
f" the length of the timestep schedule used for training: {original_steps}. You may get some"
|
||||
f" unexpected results when using this timestep schedule."
|
||||
)
|
||||
else:
|
||||
if len(timesteps) > self.config.num_train_timesteps:
|
||||
logger.warning(
|
||||
f"The number of timesteps in the custom timestep schedule is {len(timesteps)}, which exceeds the"
|
||||
f" the length of the timestep schedule used for training: {self.config.num_train_timesteps}. You may get some"
|
||||
f" unexpected results when using this timestep schedule."
|
||||
)
|
||||
|
||||
timesteps = np.array(timesteps, dtype=np.int64)
|
||||
self.num_inference_steps = len(timesteps)
|
||||
self.custom_timesteps = True
|
||||
|
||||
# Apply strength (e.g. for img2img pipelines) (see StableDiffusionImg2ImgPipeline.get_timesteps)
|
||||
init_timestep = min(int(self.num_inference_steps * strength), self.num_inference_steps)
|
||||
t_start = max(self.num_inference_steps - init_timestep, 0)
|
||||
timesteps = timesteps[t_start * self.order :]
|
||||
# TODO: also reset self.num_inference_steps?
|
||||
else:
|
||||
# 2.2 Create the "standard" TCD inference timestep schedule.
|
||||
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."
|
||||
)
|
||||
|
||||
if original_steps is not None:
|
||||
skipping_step = len(tcd_origin_timesteps) // num_inference_steps
|
||||
|
||||
if skipping_step < 1:
|
||||
raise ValueError(
|
||||
f"The combination of `original_steps x strength`: {original_steps} x {strength} is smaller than `num_inference_steps`: {num_inference_steps}. Make sure to either reduce `num_inference_steps` to a value smaller than {int(original_steps * strength)} or increase `strength` to a value higher than {float(num_inference_steps / original_steps)}."
|
||||
)
|
||||
|
||||
self.num_inference_steps = num_inference_steps
|
||||
|
||||
if original_steps is not None:
|
||||
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."
|
||||
)
|
||||
else:
|
||||
if num_inference_steps > self.config.num_train_timesteps:
|
||||
raise ValueError(
|
||||
f"`num_inference_steps`: {num_inference_steps} cannot be larger than `num_train_timesteps`:"
|
||||
f" {self.config.num_train_timesteps} because the final timestep schedule will be a subset of the"
|
||||
f" `num_train_timesteps`-sized initial timestep schedule."
|
||||
)
|
||||
|
||||
# TCD Inference Steps Schedule
|
||||
tcd_origin_timesteps = tcd_origin_timesteps[::-1].copy()
|
||||
# Select (approximately) evenly spaced indices from tcd_origin_timesteps.
|
||||
inference_indices = np.linspace(0, len(tcd_origin_timesteps), num=num_inference_steps, endpoint=False)
|
||||
inference_indices = np.floor(inference_indices).astype(np.int64)
|
||||
timesteps = tcd_origin_timesteps[inference_indices]
|
||||
|
||||
self.timesteps = torch.from_numpy(timesteps).to(device=device, dtype=torch.long)
|
||||
|
||||
self._step_index = None
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: int,
|
||||
sample: torch.FloatTensor,
|
||||
eta: float,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
return_dict: bool = True,
|
||||
) -> Union[TCDSchedulerOutput, 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 (`int`):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.FloatTensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
eta (`float`):
|
||||
A stochastic parameter (referred to as `gamma` in the paper) used to control the stochasticity in every step.
|
||||
When eta = 0, it represents deterministic sampling, whereas eta = 1 indicates full stochastic sampling.
|
||||
generator (`torch.Generator`, *optional*):
|
||||
A random number generator.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`~schedulers.scheduling_tcd.TCDSchedulerOutput`] or `tuple`.
|
||||
Returns:
|
||||
[`~schedulers.scheduling_utils.TCDSchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`, [`~schedulers.scheduling_tcd.TCDSchedulerOutput`] is returned, otherwise a
|
||||
tuple is returned where the first element is the sample tensor.
|
||||
"""
|
||||
if self.num_inference_steps is None:
|
||||
raise ValueError(
|
||||
"Number of inference steps is 'None', you need to run 'set_timesteps' after creating the scheduler"
|
||||
)
|
||||
|
||||
if self.step_index is None:
|
||||
self._init_step_index(timestep)
|
||||
|
||||
# 1. get previous step value
|
||||
prev_step_index = self.step_index + 1
|
||||
if prev_step_index < len(self.timesteps):
|
||||
prev_timestep = self.timesteps[prev_step_index]
|
||||
else:
|
||||
prev_timestep = torch.tensor(0)
|
||||
|
||||
timestep_s = torch.floor((1 - eta) * prev_timestep).to(dtype=torch.long)
|
||||
|
||||
# 2. compute alphas, betas
|
||||
alpha_prod_t = self.alphas_cumprod[timestep]
|
||||
beta_prod_t = 1 - alpha_prod_t
|
||||
|
||||
alpha_prod_t_prev = self.alphas_cumprod[prev_timestep] if prev_timestep >= 0 else self.final_alpha_cumprod
|
||||
_beta_prod_t_prev = 1 - alpha_prod_t_prev
|
||||
|
||||
alpha_prod_s = self.alphas_cumprod[timestep_s] if timestep_s >= 0 else self.final_alpha_cumprod
|
||||
beta_prod_s = 1 - alpha_prod_s
|
||||
|
||||
# 3. Compute the predicted noised sample x_s based on the model parameterization
|
||||
if self.config.prediction_type == "epsilon": # noise-prediction
|
||||
pred_original_sample = (sample - beta_prod_t.sqrt() * model_output) / alpha_prod_t.sqrt()
|
||||
pred_epsilon = model_output
|
||||
pred_noised_sample = alpha_prod_s.sqrt() * pred_original_sample + beta_prod_s.sqrt() * pred_epsilon
|
||||
elif self.config.prediction_type == "sample": # x-prediction
|
||||
pred_original_sample = model_output
|
||||
pred_epsilon = (sample - alpha_prod_t ** (0.5) * pred_original_sample) / beta_prod_t ** (0.5)
|
||||
pred_noised_sample = alpha_prod_s.sqrt() * pred_original_sample + beta_prod_s.sqrt() * pred_epsilon
|
||||
elif self.config.prediction_type == "v_prediction": # v-prediction
|
||||
pred_original_sample = (alpha_prod_t**0.5) * sample - (beta_prod_t**0.5) * model_output
|
||||
pred_epsilon = (alpha_prod_t**0.5) * model_output + (beta_prod_t**0.5) * sample
|
||||
pred_noised_sample = alpha_prod_s.sqrt() * pred_original_sample + beta_prod_s.sqrt() * pred_epsilon
|
||||
else:
|
||||
raise ValueError(
|
||||
f"prediction_type given as {self.config.prediction_type} must be one of `epsilon`, `sample` or"
|
||||
" `v_prediction` for `TCDScheduler`."
|
||||
)
|
||||
|
||||
# 4. Sample and inject noise z ~ N(0, I) for MultiStep Inference
|
||||
# Noise is not used on the final timestep of the timestep schedule.
|
||||
# This also means that noise is not used for one-step sampling.
|
||||
# Eta (referred to as "gamma" in the paper) was introduced to control the stochasticity in every step.
|
||||
# When eta = 0, it represents deterministic sampling, whereas eta = 1 indicates full stochastic sampling.
|
||||
if eta > 0:
|
||||
if self.step_index != self.num_inference_steps - 1:
|
||||
noise = randn_tensor(
|
||||
model_output.shape, generator=generator, device=model_output.device, dtype=pred_noised_sample.dtype
|
||||
)
|
||||
prev_sample = (alpha_prod_t_prev / alpha_prod_s).sqrt() * pred_noised_sample + (1 - alpha_prod_t_prev / alpha_prod_s).sqrt() * noise
|
||||
else:
|
||||
prev_sample = pred_noised_sample
|
||||
else:
|
||||
prev_sample = pred_noised_sample
|
||||
|
||||
# upon completion increase step index by one
|
||||
self._step_index += 1
|
||||
|
||||
if not return_dict:
|
||||
return (prev_sample, pred_noised_sample)
|
||||
|
||||
return TCDSchedulerOutput(prev_sample=prev_sample, pred_noised_sample=pred_noised_sample)
|
||||
|
||||
# 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
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_ddpm.DDPMScheduler.previous_timestep
|
||||
def previous_timestep(self, timestep):
|
||||
if self.custom_timesteps:
|
||||
index = (self.timesteps == timestep).nonzero(as_tuple=True)[0][0]
|
||||
if index == self.timesteps.shape[0] - 1:
|
||||
prev_t = torch.tensor(-1)
|
||||
else:
|
||||
prev_t = self.timesteps[index + 1]
|
||||
else:
|
||||
num_inference_steps = (
|
||||
self.num_inference_steps if self.num_inference_steps else self.config.num_train_timesteps
|
||||
)
|
||||
prev_t = timestep - self.config.num_train_timesteps // num_inference_steps
|
||||
|
||||
return prev_t
|
||||
@@ -1,337 +0,0 @@
|
||||
import os
|
||||
import cv2
|
||||
import requests
|
||||
import numpy as np
|
||||
from PIL import ImageDraw
|
||||
|
||||
GREEN = "#0F0"
|
||||
BLUE = "#00F"
|
||||
RED = "#F00"
|
||||
|
||||
|
||||
def crop_image(im, settings):
|
||||
""" Intelligently crop an image to the subject matter """
|
||||
|
||||
scale_by = 1
|
||||
if is_landscape(im.width, im.height):
|
||||
scale_by = settings.crop_height / im.height
|
||||
elif is_portrait(im.width, im.height):
|
||||
scale_by = settings.crop_width / im.width
|
||||
elif is_square(im.width, im.height):
|
||||
if is_square(settings.crop_width, settings.crop_height):
|
||||
scale_by = settings.crop_width / im.width
|
||||
elif is_landscape(settings.crop_width, settings.crop_height):
|
||||
scale_by = settings.crop_width / im.width
|
||||
elif is_portrait(settings.crop_width, settings.crop_height):
|
||||
scale_by = settings.crop_height / im.height
|
||||
|
||||
im = im.resize((int(im.width * scale_by), int(im.height * scale_by)))
|
||||
im_debug = im.copy()
|
||||
|
||||
focus = focal_point(im_debug, settings)
|
||||
|
||||
# take the focal point and turn it into crop coordinates that try to center over the focal
|
||||
# point but then get adjusted back into the frame
|
||||
y_half = int(settings.crop_height / 2)
|
||||
x_half = int(settings.crop_width / 2)
|
||||
|
||||
x1 = focus.x - x_half
|
||||
if x1 < 0:
|
||||
x1 = 0
|
||||
elif x1 + settings.crop_width > im.width:
|
||||
x1 = im.width - settings.crop_width
|
||||
|
||||
y1 = focus.y - y_half
|
||||
if y1 < 0:
|
||||
y1 = 0
|
||||
elif y1 + settings.crop_height > im.height:
|
||||
y1 = im.height - settings.crop_height
|
||||
|
||||
x2 = x1 + settings.crop_width
|
||||
y2 = y1 + settings.crop_height
|
||||
|
||||
crop = [x1, y1, x2, y2]
|
||||
|
||||
results = []
|
||||
|
||||
results.append(im.crop(tuple(crop)))
|
||||
|
||||
if settings.annotate_image:
|
||||
d = ImageDraw.Draw(im_debug)
|
||||
rect = list(crop)
|
||||
rect[2] -= 1
|
||||
rect[3] -= 1
|
||||
d.rectangle(rect, outline=GREEN)
|
||||
results.append(im_debug)
|
||||
if settings.destop_view_image:
|
||||
im_debug.show()
|
||||
|
||||
return results
|
||||
|
||||
def focal_point(im, settings):
|
||||
corner_points = image_corner_points(im, settings) if settings.corner_points_weight > 0 else []
|
||||
entropy_points = image_entropy_points(im, settings) if settings.entropy_points_weight > 0 else []
|
||||
face_points = image_face_points(im, settings) if settings.face_points_weight > 0 else []
|
||||
|
||||
pois = []
|
||||
|
||||
weight_pref_total = 0
|
||||
if len(corner_points) > 0:
|
||||
weight_pref_total += settings.corner_points_weight
|
||||
if len(entropy_points) > 0:
|
||||
weight_pref_total += settings.entropy_points_weight
|
||||
if len(face_points) > 0:
|
||||
weight_pref_total += settings.face_points_weight
|
||||
|
||||
corner_centroid = None
|
||||
if len(corner_points) > 0:
|
||||
corner_centroid = centroid(corner_points)
|
||||
corner_centroid.weight = settings.corner_points_weight / weight_pref_total
|
||||
pois.append(corner_centroid)
|
||||
|
||||
entropy_centroid = None
|
||||
if len(entropy_points) > 0:
|
||||
entropy_centroid = centroid(entropy_points)
|
||||
entropy_centroid.weight = settings.entropy_points_weight / weight_pref_total
|
||||
pois.append(entropy_centroid)
|
||||
|
||||
face_centroid = None
|
||||
if len(face_points) > 0:
|
||||
face_centroid = centroid(face_points)
|
||||
face_centroid.weight = settings.face_points_weight / weight_pref_total
|
||||
pois.append(face_centroid)
|
||||
|
||||
average_point = poi_average(pois, settings)
|
||||
|
||||
if settings.annotate_image:
|
||||
d = ImageDraw.Draw(im)
|
||||
max_size = min(im.width, im.height) * 0.07
|
||||
if corner_centroid is not None:
|
||||
color = BLUE
|
||||
box = corner_centroid.bounding(max_size * corner_centroid.weight)
|
||||
d.text((box[0], box[1]-15), f"Edge: {corner_centroid.weight:.02f}", fill=color)
|
||||
d.ellipse(box, outline=color)
|
||||
if len(corner_points) > 1:
|
||||
for f in corner_points:
|
||||
d.rectangle(f.bounding(4), outline=color)
|
||||
if entropy_centroid is not None:
|
||||
color = "#ff0"
|
||||
box = entropy_centroid.bounding(max_size * entropy_centroid.weight)
|
||||
d.text((box[0], box[1]-15), f"Entropy: {entropy_centroid.weight:.02f}", fill=color)
|
||||
d.ellipse(box, outline=color)
|
||||
if len(entropy_points) > 1:
|
||||
for f in entropy_points:
|
||||
d.rectangle(f.bounding(4), outline=color)
|
||||
if face_centroid is not None:
|
||||
color = RED
|
||||
box = face_centroid.bounding(max_size * face_centroid.weight)
|
||||
d.text((box[0], box[1]-15), f"Face: {face_centroid.weight:.02f}", fill=color)
|
||||
d.ellipse(box, outline=color)
|
||||
if len(face_points) > 1:
|
||||
for f in face_points:
|
||||
d.rectangle(f.bounding(4), outline=color)
|
||||
|
||||
d.ellipse(average_point.bounding(max_size), outline=GREEN)
|
||||
|
||||
return average_point
|
||||
|
||||
|
||||
def image_face_points(im, settings):
|
||||
if settings.dnn_model_path is not None:
|
||||
detector = cv2.FaceDetectorYN.create(
|
||||
settings.dnn_model_path,
|
||||
"",
|
||||
(im.width, im.height),
|
||||
0.9, # score threshold
|
||||
0.3, # nms threshold
|
||||
5000 # keep top k before nms
|
||||
)
|
||||
faces = detector.detect(np.array(im))
|
||||
results = []
|
||||
if faces[1] is not None:
|
||||
for face in faces[1]:
|
||||
x = face[0]
|
||||
y = face[1]
|
||||
w = face[2]
|
||||
h = face[3]
|
||||
results.append(
|
||||
PointOfInterest(
|
||||
int(x + (w * 0.5)), # face focus left/right is center
|
||||
int(y + (h * 0.33)), # face focus up/down is close to the top of the head
|
||||
size = w,
|
||||
weight = 1/len(faces[1])
|
||||
)
|
||||
)
|
||||
return results
|
||||
else:
|
||||
np_im = np.array(im)
|
||||
gray = cv2.cvtColor(np_im, cv2.COLOR_BGR2GRAY)
|
||||
|
||||
tries = [
|
||||
[ f'{cv2.data.haarcascades}haarcascade_eye.xml', 0.01 ],
|
||||
[ f'{cv2.data.haarcascades}haarcascade_frontalface_default.xml', 0.05 ],
|
||||
[ f'{cv2.data.haarcascades}haarcascade_profileface.xml', 0.05 ],
|
||||
[ f'{cv2.data.haarcascades}haarcascade_frontalface_alt.xml', 0.05 ],
|
||||
[ f'{cv2.data.haarcascades}haarcascade_frontalface_alt2.xml', 0.05 ],
|
||||
[ f'{cv2.data.haarcascades}haarcascade_frontalface_alt_tree.xml', 0.05 ],
|
||||
[ f'{cv2.data.haarcascades}haarcascade_eye_tree_eyeglasses.xml', 0.05 ],
|
||||
[ f'{cv2.data.haarcascades}haarcascade_upperbody.xml', 0.05 ]
|
||||
]
|
||||
for t in tries:
|
||||
classifier = cv2.CascadeClassifier(t[0])
|
||||
minsize = int(min(im.width, im.height) * t[1]) # at least N percent of the smallest side
|
||||
try:
|
||||
faces = classifier.detectMultiScale(gray, scaleFactor=1.1,
|
||||
minNeighbors=7, minSize=(minsize, minsize), flags=cv2.CASCADE_SCALE_IMAGE)
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
if len(faces) > 0:
|
||||
rects = [[f[0], f[1], f[0] + f[2], f[1] + f[3]] for f in faces]
|
||||
return [PointOfInterest((r[0] +r[2]) // 2, (r[1] + r[3]) // 2, size=abs(r[0]-r[2]), weight=1/len(rects)) for r in rects]
|
||||
return []
|
||||
|
||||
|
||||
def image_corner_points(im, settings): # pylint: disable=unused-argument
|
||||
grayscale = im.convert("L")
|
||||
|
||||
# naive attempt at preventing focal points from collecting at watermarks near the bottom
|
||||
gd = ImageDraw.Draw(grayscale)
|
||||
gd.rectangle([0, im.height*.9, im.width, im.height], fill="#999")
|
||||
|
||||
np_im = np.array(grayscale)
|
||||
|
||||
points = cv2.goodFeaturesToTrack(
|
||||
np_im,
|
||||
maxCorners=100,
|
||||
qualityLevel=0.04,
|
||||
minDistance=min(grayscale.width, grayscale.height)*0.06,
|
||||
useHarrisDetector=False,
|
||||
)
|
||||
|
||||
if points is None:
|
||||
return []
|
||||
|
||||
focal_points = []
|
||||
for point in points:
|
||||
x, y = point.ravel()
|
||||
focal_points.append(PointOfInterest(x, y, size=4, weight=1/len(points)))
|
||||
|
||||
return focal_points
|
||||
|
||||
|
||||
def image_entropy_points(im, settings):
|
||||
landscape = im.height < im.width
|
||||
portrait = im.height > im.width
|
||||
if landscape:
|
||||
move_idx = [0, 2]
|
||||
move_max = im.size[0]
|
||||
elif portrait:
|
||||
move_idx = [1, 3]
|
||||
move_max = im.size[1]
|
||||
else:
|
||||
return []
|
||||
|
||||
e_max = 0
|
||||
crop_current = [0, 0, settings.crop_width, settings.crop_height]
|
||||
crop_best = crop_current
|
||||
while crop_current[move_idx[1]] < move_max:
|
||||
crop = im.crop(tuple(crop_current))
|
||||
e = image_entropy(crop)
|
||||
|
||||
if e > e_max:
|
||||
e_max = e
|
||||
crop_best = list(crop_current)
|
||||
|
||||
crop_current[move_idx[0]] += 4
|
||||
crop_current[move_idx[1]] += 4
|
||||
|
||||
x_mid = int(crop_best[0] + settings.crop_width/2)
|
||||
y_mid = int(crop_best[1] + settings.crop_height/2)
|
||||
|
||||
return [PointOfInterest(x_mid, y_mid, size=25, weight=1.0)]
|
||||
|
||||
|
||||
def image_entropy(im):
|
||||
# greyscale image entropy
|
||||
# band = np.asarray(im.convert("L"))
|
||||
band = np.asarray(im.convert("1"), dtype=np.uint8)
|
||||
hist, _ = np.histogram(band, bins=range(0, 256))
|
||||
hist = hist[hist > 0]
|
||||
return -np.log2(hist / hist.sum()).sum()
|
||||
|
||||
def centroid(pois):
|
||||
x = [poi.x for poi in pois]
|
||||
y = [poi.y for poi in pois]
|
||||
return PointOfInterest(sum(x)/len(pois), sum(y)/len(pois))
|
||||
|
||||
|
||||
def poi_average(pois, settings): # pylint: disable=unused-argument
|
||||
weight = 0.0
|
||||
x = 0.0
|
||||
y = 0.0
|
||||
for poi in pois:
|
||||
weight += poi.weight
|
||||
x += poi.x * poi.weight
|
||||
y += poi.y * poi.weight
|
||||
avg_x = round(weight and x / weight)
|
||||
avg_y = round(weight and y / weight)
|
||||
|
||||
return PointOfInterest(avg_x, avg_y)
|
||||
|
||||
|
||||
def is_landscape(w, h):
|
||||
return w > h
|
||||
|
||||
|
||||
def is_portrait(w, h):
|
||||
return h > w
|
||||
|
||||
|
||||
def is_square(w, h):
|
||||
return w == h
|
||||
|
||||
|
||||
def download_and_cache_models(dirname):
|
||||
download_url = 'https://github.com/opencv/opencv_zoo/blob/91fb0290f50896f38a0ab1e558b74b16bc009428/models/face_detection_yunet/face_detection_yunet_2022mar.onnx?raw=true'
|
||||
model_file_name = 'face_detection_yunet.onnx'
|
||||
if not os.path.exists(dirname):
|
||||
os.makedirs(dirname, exist_ok=True)
|
||||
cache_file = os.path.join(dirname, model_file_name)
|
||||
if not os.path.exists(cache_file):
|
||||
print(f"downloading face detection model from '{download_url}' to '{cache_file}'")
|
||||
response = requests.get(download_url, timeout=60*60*2)
|
||||
with open(cache_file, "wb") as f:
|
||||
f.write(response.content)
|
||||
|
||||
if os.path.exists(cache_file):
|
||||
return cache_file
|
||||
return None
|
||||
|
||||
|
||||
class PointOfInterest:
|
||||
def __init__(self, x, y, weight=1.0, size=10):
|
||||
self.x = x
|
||||
self.y = y
|
||||
self.weight = weight
|
||||
self.size = size
|
||||
|
||||
def bounding(self, size):
|
||||
return [
|
||||
self.x - size//2,
|
||||
self.y - size//2,
|
||||
self.x + size//2,
|
||||
self.y + size//2
|
||||
]
|
||||
|
||||
|
||||
class Settings:
|
||||
def __init__(self, crop_width=512, crop_height=512, corner_points_weight=0.5, entropy_points_weight=0.5, face_points_weight=0.5, annotate_image=False, dnn_model_path=None):
|
||||
self.crop_width = crop_width
|
||||
self.crop_height = crop_height
|
||||
self.corner_points_weight = corner_points_weight
|
||||
self.entropy_points_weight = entropy_points_weight
|
||||
self.face_points_weight = face_points_weight
|
||||
self.annotate_image = annotate_image
|
||||
self.destop_view_image = False
|
||||
self.dnn_model_path = dnn_model_path
|
||||
@@ -1,223 +0,0 @@
|
||||
import os
|
||||
import re
|
||||
import random
|
||||
from collections import defaultdict
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset, DataLoader, Sampler
|
||||
from torchvision import transforms
|
||||
import tqdm
|
||||
from ldm.modules.distributions.distributions import DiagonalGaussianDistribution
|
||||
from modules import devices, shared
|
||||
|
||||
re_numbers_at_start = re.compile(r"^[-\d]+\s*")
|
||||
|
||||
|
||||
class DatasetEntry:
|
||||
def __init__(self, filename=None, filename_text=None, latent_dist=None, latent_sample=None, cond=None, cond_text=None, pixel_values=None, weight=None):
|
||||
self.filename = filename
|
||||
self.filename_text = filename_text
|
||||
self.weight = weight
|
||||
self.latent_dist = latent_dist
|
||||
self.latent_sample = latent_sample
|
||||
self.cond = cond
|
||||
self.cond_text = cond_text
|
||||
self.pixel_values = pixel_values
|
||||
|
||||
|
||||
class PersonalizedBase(Dataset):
|
||||
def __init__(self, data_root, width, height, repeats, flip_p=0.5, placeholder_token="*", model=None, cond_model=None, device=None, template_file=None, include_cond=False, batch_size=1, gradient_step=1, shuffle_tags=False, tag_drop_out=0, latent_sampling_method='once', varsize=False, use_weight=False):
|
||||
re_word = re.compile(shared.opts.dataset_filename_word_regex) if len(shared.opts.dataset_filename_word_regex) > 0 else None
|
||||
|
||||
self.placeholder_token = placeholder_token
|
||||
self.flip = transforms.RandomHorizontalFlip(p=flip_p)
|
||||
self.dataset = []
|
||||
with open(template_file, "r", encoding="utf8") as file:
|
||||
lines = [x.strip() for x in file.readlines()]
|
||||
self.lines = lines
|
||||
|
||||
assert data_root, 'dataset directory not specified'
|
||||
assert os.path.isdir(data_root), "Dataset directory doesn't exist"
|
||||
assert os.listdir(data_root), "Dataset directory is empty"
|
||||
|
||||
self.image_paths = [os.path.join(data_root, file_path) for file_path in os.listdir(data_root)]
|
||||
self.shuffle_tags = shuffle_tags
|
||||
self.tag_drop_out = tag_drop_out
|
||||
groups = defaultdict(list)
|
||||
shared.log.info(f"TI Training: Preparing dataset: {data_root}")
|
||||
for path in tqdm.tqdm(self.image_paths):
|
||||
alpha_channel = None
|
||||
if shared.state.interrupted:
|
||||
raise RuntimeError("interrupted")
|
||||
try:
|
||||
image = Image.open(path)
|
||||
if use_weight and 'A' in image.getbands():
|
||||
alpha_channel = image.getchannel('A')
|
||||
image = image.convert('RGB')
|
||||
if not varsize:
|
||||
image = image.resize((width, height), Image.Resampling.BICUBIC)
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
text_filename = f"{os.path.splitext(path)[0]}.txt"
|
||||
filename = os.path.basename(path)
|
||||
|
||||
if os.path.exists(text_filename):
|
||||
with open(text_filename, "r", encoding="utf8") as file:
|
||||
filename_text = file.read()
|
||||
else:
|
||||
filename_text = os.path.splitext(filename)[0]
|
||||
filename_text = re.sub(re_numbers_at_start, '', filename_text)
|
||||
if re_word:
|
||||
tokens = re_word.findall(filename_text)
|
||||
filename_text = (shared.opts.dataset_filename_join_string or "").join(tokens)
|
||||
|
||||
npimage = np.array(image).astype(np.uint8)
|
||||
npimage = (npimage / 127.5 - 1.0).astype(np.float32)
|
||||
|
||||
torchdata = torch.from_numpy(npimage).permute(2, 0, 1).to(device=device, dtype=torch.float32)
|
||||
latent_sample = None
|
||||
|
||||
with devices.autocast():
|
||||
latent_dist = model.encode_first_stage(torchdata.unsqueeze(dim=0))
|
||||
|
||||
if latent_sampling_method == "deterministic":
|
||||
if isinstance(latent_dist, DiagonalGaussianDistribution):
|
||||
latent_dist.std = torch.exp(0 * latent_dist.logvar)
|
||||
else:
|
||||
latent_sampling_method = "once"
|
||||
latent_sample = model.get_first_stage_encoding(latent_dist).squeeze().to(devices.cpu)
|
||||
|
||||
if use_weight and alpha_channel is not None:
|
||||
channels, *latent_size = latent_sample.shape
|
||||
weight_img = alpha_channel.resize(latent_size)
|
||||
npweight = np.array(weight_img).astype(np.float32)
|
||||
#Repeat for every channel in the latent sample
|
||||
weight = torch.tensor([npweight] * channels).reshape([channels] + latent_size)
|
||||
#Normalize the weight to a minimum of 0 and a mean of 1, that way the loss will be comparable to default.
|
||||
weight -= weight.min()
|
||||
weight /= weight.mean()
|
||||
elif use_weight:
|
||||
#If an image does not have a alpha channel, add a ones weight map anyway so we can stack it later
|
||||
weight = torch.ones(latent_sample.shape)
|
||||
else:
|
||||
weight = None
|
||||
|
||||
if latent_sampling_method == "random":
|
||||
entry = DatasetEntry(filename=path, filename_text=filename_text, latent_dist=latent_dist, weight=weight)
|
||||
else:
|
||||
entry = DatasetEntry(filename=path, filename_text=filename_text, latent_sample=latent_sample, weight=weight)
|
||||
|
||||
if not (self.tag_drop_out != 0 or self.shuffle_tags):
|
||||
entry.cond_text = self.create_text(filename_text)
|
||||
|
||||
if include_cond and not (self.tag_drop_out != 0 or self.shuffle_tags):
|
||||
with devices.autocast():
|
||||
entry.cond = cond_model([entry.cond_text]).to(devices.cpu).squeeze(0)
|
||||
groups[image.size].append(len(self.dataset))
|
||||
self.dataset.append(entry)
|
||||
del torchdata
|
||||
del latent_dist
|
||||
del latent_sample
|
||||
del weight
|
||||
|
||||
self.length = len(self.dataset)
|
||||
self.groups = list(groups.values())
|
||||
assert self.length > 0, "No images have been found in the dataset."
|
||||
self.batch_size = min(batch_size, self.length)
|
||||
self.gradient_step = min(gradient_step, self.length // self.batch_size)
|
||||
self.latent_sampling_method = latent_sampling_method
|
||||
|
||||
def create_text(self, filename_text):
|
||||
text = random.choice(self.lines)
|
||||
tags = filename_text.split(',')
|
||||
if self.tag_drop_out != 0:
|
||||
tags = [t for t in tags if random.random() > self.tag_drop_out]
|
||||
if self.shuffle_tags:
|
||||
random.shuffle(tags)
|
||||
text = text.replace("[filewords]", ','.join(tags))
|
||||
text = text.replace("[name]", self.placeholder_token)
|
||||
return text
|
||||
|
||||
def __len__(self):
|
||||
return self.length
|
||||
|
||||
def __getitem__(self, i):
|
||||
entry = self.dataset[i]
|
||||
if self.tag_drop_out != 0 or self.shuffle_tags:
|
||||
entry.cond_text = self.create_text(entry.filename_text)
|
||||
if self.latent_sampling_method == "random":
|
||||
entry.latent_sample = shared.sd_model.get_first_stage_encoding(entry.latent_dist).to(devices.cpu)
|
||||
return entry
|
||||
|
||||
|
||||
class GroupedBatchSampler(Sampler):
|
||||
def __init__(self, data_source: PersonalizedBase, batch_size: int):
|
||||
super().__init__(data_source)
|
||||
|
||||
n = len(data_source)
|
||||
self.groups = data_source.groups
|
||||
self.len = n_batch = n // batch_size
|
||||
expected = [len(g) / n * n_batch * batch_size for g in data_source.groups]
|
||||
self.base = [int(e) // batch_size for e in expected]
|
||||
self.n_rand_batches = nrb = n_batch - sum(self.base)
|
||||
self.probs = [e%batch_size/nrb/batch_size if nrb>0 else 0 for e in expected]
|
||||
self.batch_size = batch_size
|
||||
|
||||
def __len__(self):
|
||||
return self.len
|
||||
|
||||
def __iter__(self):
|
||||
b = self.batch_size
|
||||
|
||||
for g in self.groups:
|
||||
random.shuffle(g)
|
||||
|
||||
batches = []
|
||||
for g in self.groups:
|
||||
batches.extend(g[i*b:(i+1)*b] for i in range(len(g) // b))
|
||||
for _ in range(self.n_rand_batches):
|
||||
rand_group = random.choices(self.groups, self.probs)[0]
|
||||
batches.append(random.choices(rand_group, k=b))
|
||||
|
||||
random.shuffle(batches)
|
||||
|
||||
yield from batches
|
||||
|
||||
|
||||
class PersonalizedDataLoader(DataLoader):
|
||||
def __init__(self, dataset, latent_sampling_method="once", batch_size=1, pin_memory=False):
|
||||
super(PersonalizedDataLoader, self).__init__(dataset, batch_sampler=GroupedBatchSampler(dataset, batch_size), pin_memory=pin_memory)
|
||||
if latent_sampling_method == "random":
|
||||
self.collate_fn = collate_wrapper_random
|
||||
else:
|
||||
self.collate_fn = collate_wrapper
|
||||
|
||||
|
||||
class BatchLoader:
|
||||
def __init__(self, data):
|
||||
self.cond_text = [entry.cond_text for entry in data]
|
||||
self.cond = [entry.cond for entry in data]
|
||||
self.latent_sample = torch.stack([entry.latent_sample for entry in data]).squeeze(1)
|
||||
if all(entry.weight is not None for entry in data):
|
||||
self.weight = torch.stack([entry.weight for entry in data]).squeeze(1)
|
||||
else:
|
||||
self.weight = None
|
||||
|
||||
def pin_memory(self):
|
||||
self.latent_sample = self.latent_sample.pin_memory()
|
||||
return self
|
||||
|
||||
def collate_wrapper(batch):
|
||||
return BatchLoader(batch)
|
||||
|
||||
class BatchLoaderRandom(BatchLoader):
|
||||
def __init__(self, data):
|
||||
super().__init__(data)
|
||||
|
||||
def pin_memory(self):
|
||||
return self
|
||||
|
||||
def collate_wrapper_random(batch):
|
||||
return BatchLoaderRandom(batch)
|
||||
@@ -8,17 +8,17 @@ from modules.shared import opts
|
||||
|
||||
|
||||
class EmbeddingEncoder(json.JSONEncoder):
|
||||
def default(self, obj):
|
||||
if isinstance(obj, torch.Tensor):
|
||||
return {'TORCHTENSOR': obj.cpu().detach().numpy().tolist()}
|
||||
return json.JSONEncoder.default(self, obj)
|
||||
def default(self, o):
|
||||
if isinstance(o, torch.Tensor):
|
||||
return {'TORCHTENSOR': o.cpu().detach().numpy().tolist()}
|
||||
return json.JSONEncoder.default(self, o)
|
||||
|
||||
|
||||
class EmbeddingDecoder(json.JSONDecoder):
|
||||
def __init__(self, *args, **kwargs):
|
||||
json.JSONDecoder.__init__(self, *args, object_hook=self.object_hook, **kwargs)
|
||||
|
||||
def object_hook(self, d):
|
||||
def object_hook(self, d): # pylint: disable=E0202
|
||||
if 'TORCHTENSOR' in d:
|
||||
return torch.from_numpy(np.array(d['TORCHTENSOR']))
|
||||
return d
|
||||
@@ -41,8 +41,8 @@ def lcg(m=2**32, a=1664525, c=1013904223, seed=0):
|
||||
|
||||
|
||||
def xor_block(block):
|
||||
g = lcg()
|
||||
randblock = np.array([next(g) for _ in range(np.prod(block.shape))]).astype(np.uint8).reshape(block.shape)
|
||||
blk = lcg()
|
||||
randblock = np.array([next(blk) for _ in range(np.prod(block.shape))]).astype(np.uint8).reshape(block.shape)
|
||||
return np.bitwise_xor(block.astype(np.uint8), randblock & 0x0F)
|
||||
|
||||
|
||||
@@ -110,7 +110,7 @@ def crop_black(img, tol=0):
|
||||
|
||||
def extract_image_data_embed(image):
|
||||
d = 3
|
||||
outarr = crop_black(np.array(image.convert('RGB').getdata()).reshape(image.size[1], image.size[0], d).astype(np.uint8)) & 0x0F
|
||||
outarr = crop_black(np.array(image.convert('RGB').getdata()).reshape(image.size[1], image.size[0], d).astype(np.uint8)) & 0x0F # pylint: disable=E1121
|
||||
black_cols = np.where(np.sum(outarr, axis=(0, 2)) == 0)
|
||||
if black_cols[0].shape[0] < 2:
|
||||
return None
|
||||
@@ -182,7 +182,7 @@ if __name__ == '__main__':
|
||||
new_image = Image.new('RGBA', (512, 512), (255, 255, 200, 255))
|
||||
cap_image = caption_image_overlay(new_image, 'title', 'footerLeft', 'footerMid', 'footerRight')
|
||||
|
||||
test_embed = {'string_to_param': {'*': torch.from_numpy(np.random.random((2, 4096)))}}
|
||||
test_embed = {'string_to_param': {'*': torch.from_numpy(np.random.random((2, 4096)))}} # noqa: NPY002
|
||||
|
||||
embedded_image = insert_image_data_embed(cap_image, test_embed)
|
||||
|
||||
|
||||
@@ -1,73 +0,0 @@
|
||||
class LearnScheduleIterator:
|
||||
def __init__(self, learn_rate, max_steps, cur_step=0):
|
||||
"""
|
||||
specify learn_rate as "0.001:100, 0.00001:1000, 1e-5:10000" to have lr of 0.001 until step 100, 0.00001 until 1000, and 1e-5 until 10000
|
||||
"""
|
||||
|
||||
pairs = learn_rate.split(',')
|
||||
self.rates = []
|
||||
self.it = 0
|
||||
self.maxit = 0
|
||||
try:
|
||||
for pair in pairs:
|
||||
if not pair.strip():
|
||||
continue
|
||||
tmp = pair.split(':')
|
||||
if len(tmp) == 2:
|
||||
step = int(tmp[1])
|
||||
if step > cur_step:
|
||||
self.rates.append((float(tmp[0]), min(step, max_steps)))
|
||||
self.maxit += 1
|
||||
if step > max_steps:
|
||||
return
|
||||
elif step == -1:
|
||||
self.rates.append((float(tmp[0]), max_steps))
|
||||
self.maxit += 1
|
||||
return
|
||||
else:
|
||||
self.rates.append((float(tmp[0]), max_steps))
|
||||
self.maxit += 1
|
||||
return
|
||||
assert self.rates
|
||||
except (ValueError, AssertionError) as e:
|
||||
raise RuntimeError('Invalid learning rate schedule. It should be a number or, for example, like "0.001:100, 0.00001:1000, 1e-5:10000" to have lr of 0.001 until step 100, 0.00001 until 1000, and 1e-5 until 10000.') from e
|
||||
|
||||
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
if self.it < self.maxit:
|
||||
self.it += 1
|
||||
return self.rates[self.it - 1]
|
||||
else:
|
||||
raise StopIteration
|
||||
|
||||
|
||||
class LearnRateScheduler:
|
||||
def __init__(self, learn_rate, max_steps, cur_step=0, verbose=True):
|
||||
self.schedules = LearnScheduleIterator(learn_rate, max_steps, cur_step)
|
||||
(self.learn_rate, self.end_step) = next(self.schedules)
|
||||
self.verbose = verbose
|
||||
self.finished = False
|
||||
|
||||
def step(self, step_number):
|
||||
if step_number < self.end_step:
|
||||
return False
|
||||
|
||||
try:
|
||||
(self.learn_rate, self.end_step) = next(self.schedules)
|
||||
except StopIteration:
|
||||
self.finished = True
|
||||
return False
|
||||
return True
|
||||
|
||||
def apply(self, optimizer, step_number):
|
||||
if not self.step(step_number):
|
||||
return
|
||||
|
||||
# if self.verbose:
|
||||
# tqdm.tqdm.write(f'Training at rate of {self.learn_rate} until step {self.end_step}')
|
||||
|
||||
for pg in optimizer.param_groups:
|
||||
pg['lr'] = self.learn_rate
|
||||
@@ -1,217 +0,0 @@
|
||||
import os
|
||||
import math
|
||||
from tqdm import tqdm
|
||||
from PIL import Image, ImageOps
|
||||
from modules import paths, shared, images, deepbooru
|
||||
from modules.textual_inversion import autocrop
|
||||
|
||||
|
||||
def preprocess(id_task, process_src, process_dst, process_width, process_height, preprocess_txt_action, process_keep_original_size=False, process_keep_channels=False, process_flip=False, process_split=False, process_caption_only=False, process_caption=False, process_caption_deepbooru=False, split_threshold=0.5, overlap_ratio=0.2, process_focal_crop=False, process_focal_crop_face_weight=0.9, process_focal_crop_entropy_weight=0.3, process_focal_crop_edges_weight=0.5, process_focal_crop_debug=False, process_multicrop=None, process_multicrop_mindim=None, process_multicrop_maxdim=None, process_multicrop_minarea=None, process_multicrop_maxarea=None, process_multicrop_objective=None, process_multicrop_threshold=None): # pylint: disable=unused-argument
|
||||
try:
|
||||
if process_caption:
|
||||
shared.interrogator.load()
|
||||
|
||||
if process_caption_deepbooru:
|
||||
deepbooru.model.start()
|
||||
|
||||
preprocess_work(process_src, process_dst, process_width, process_height, preprocess_txt_action, process_keep_original_size, process_keep_channels, process_flip, process_split, process_caption, process_caption_deepbooru, process_caption_only, split_threshold, overlap_ratio, process_focal_crop, process_focal_crop_face_weight, process_focal_crop_entropy_weight, process_focal_crop_edges_weight, process_focal_crop_debug, process_multicrop, process_multicrop_mindim, process_multicrop_maxdim, process_multicrop_minarea, process_multicrop_maxarea, process_multicrop_objective, process_multicrop_threshold)
|
||||
|
||||
finally:
|
||||
|
||||
if process_caption:
|
||||
shared.interrogator.send_blip_to_ram()
|
||||
|
||||
if process_caption_deepbooru:
|
||||
deepbooru.model.stop()
|
||||
|
||||
|
||||
class PreprocessParams:
|
||||
src = None
|
||||
dstdir = None
|
||||
subindex = 0
|
||||
flip = False
|
||||
process_caption_only = False
|
||||
process_caption = False
|
||||
process_caption_deepbooru = False
|
||||
preprocess_txt_action = None
|
||||
|
||||
|
||||
def save_pic_with_caption(image, index, params: PreprocessParams, existing_caption=None, existing_caption_filename=None):
|
||||
caption = ""
|
||||
if params.process_caption:
|
||||
caption += shared.interrogator.generate_caption(image)
|
||||
if params.process_caption_deepbooru:
|
||||
if len(caption) > 0:
|
||||
caption += ", "
|
||||
caption += deepbooru.model.tag_multi(image)
|
||||
|
||||
filename_part = params.src
|
||||
filename_part = os.path.splitext(filename_part)[0]
|
||||
filename_part = os.path.basename(filename_part)
|
||||
|
||||
basename = f"{index:05}-{params.subindex}-{filename_part}"
|
||||
if not params.process_caption_only:
|
||||
image.save(os.path.join(params.dstdir, f"{basename}.png"))
|
||||
|
||||
if params.preprocess_txt_action == 'prepend' and existing_caption:
|
||||
caption = f"{existing_caption} {caption}"
|
||||
elif params.preprocess_txt_action == 'append' and existing_caption:
|
||||
caption = f"{caption} {existing_caption}"
|
||||
elif params.preprocess_txt_action == 'copy' and existing_caption:
|
||||
caption = existing_caption
|
||||
caption = caption.strip()
|
||||
if len(caption) > 0:
|
||||
if params.process_caption_only:
|
||||
fn = os.path.join(params.dstdir, f"{filename_part}.txt")
|
||||
elif existing_caption_filename is not None:
|
||||
fn = existing_caption_filename
|
||||
else:
|
||||
fn = os.path.join(params.dstdir, f"{basename}.txt")
|
||||
with open(fn, "w", encoding="utf8") as file:
|
||||
file.write(caption)
|
||||
|
||||
params.subindex += 1
|
||||
|
||||
|
||||
def save_pic(image, index, params, existing_caption=None, existing_caption_filename=None):
|
||||
save_pic_with_caption(image, index, params, existing_caption=existing_caption, existing_caption_filename=existing_caption_filename)
|
||||
if params.flip:
|
||||
save_pic_with_caption(ImageOps.mirror(image), index, params, existing_caption=existing_caption, existing_caption_filename=existing_caption_filename)
|
||||
|
||||
|
||||
def split_pic(image, inverse_xy, width, height, overlap_ratio):
|
||||
if inverse_xy:
|
||||
from_w, from_h = image.height, image.width
|
||||
to_w, to_h = height, width
|
||||
else:
|
||||
from_w, from_h = image.width, image.height
|
||||
to_w, to_h = width, height
|
||||
h = from_h * to_w // from_w
|
||||
if inverse_xy:
|
||||
image = image.resize((h, to_w))
|
||||
else:
|
||||
image = image.resize((to_w, h))
|
||||
|
||||
split_count = math.ceil((h - to_h * overlap_ratio) / (to_h * (1.0 - overlap_ratio)))
|
||||
y_step = (h - to_h) / (split_count - 1)
|
||||
for i in range(split_count):
|
||||
y = int(y_step * i)
|
||||
if inverse_xy:
|
||||
splitted = image.crop((y, 0, y + to_h, to_w))
|
||||
else:
|
||||
splitted = image.crop((0, y, to_w, y + to_h))
|
||||
yield splitted
|
||||
|
||||
# not using torchvision.transforms.CenterCrop because it doesn't allow float regions
|
||||
def center_crop(image: Image, w: int, h: int):
|
||||
iw, ih = image.size
|
||||
if ih / h < iw / w:
|
||||
sw = w * ih / h
|
||||
box = (iw - sw) / 2, 0, iw - (iw - sw) / 2, ih
|
||||
else:
|
||||
sh = h * iw / w
|
||||
box = 0, (ih - sh) / 2, iw, ih - (ih - sh) / 2
|
||||
return image.resize((w, h), Image.Resampling.LANCZOS, box)
|
||||
|
||||
|
||||
def multicrop_pic(image: Image, mindim, maxdim, minarea, maxarea, objective, threshold):
|
||||
iw, ih = image.size
|
||||
err = lambda w, h: 1-(lambda x: x if x < 1 else 1/x)(iw/ih/(w/h)) # pylint: disable=unnecessary-lambda-assignment,unnecessary-direct-lambda-call
|
||||
wh = max(((w, h) for w in range(mindim, maxdim+1, 64) for h in range(mindim, maxdim+1, 64)
|
||||
if minarea <= w * h <= maxarea and err(w, h) <= threshold),
|
||||
key= lambda wh: (wh[0]*wh[1], -err(*wh))[::1 if objective=='Maximize area' else -1],
|
||||
default=None
|
||||
)
|
||||
return wh and center_crop(image, *wh)
|
||||
|
||||
|
||||
def preprocess_work(process_src, process_dst, process_width, process_height, preprocess_txt_action, process_keep_original_size, process_keep_channels, process_flip, process_split, process_caption, process_caption_deepbooru, process_caption_only, split_threshold, overlap_ratio, process_focal_crop, process_focal_crop_face_weight, process_focal_crop_entropy_weight, process_focal_crop_edges_weight, process_focal_crop_debug, process_multicrop, process_multicrop_mindim, process_multicrop_maxdim, process_multicrop_minarea, process_multicrop_maxarea, process_multicrop_objective, process_multicrop_threshold):
|
||||
|
||||
width = process_width
|
||||
height = process_height
|
||||
src = os.path.abspath(process_src)
|
||||
dst = os.path.abspath(process_dst)
|
||||
split_threshold = max(0.0, min(1.0, split_threshold))
|
||||
overlap_ratio = max(0.0, min(0.9, overlap_ratio))
|
||||
assert src != dst, 'same directory specified as source and destination'
|
||||
os.makedirs(dst, exist_ok=True)
|
||||
files = os.listdir(src)
|
||||
shared.state.job = "preprocess"
|
||||
shared.state.textinfo = "Preprocessing..."
|
||||
shared.state.job_count = len(files)
|
||||
params = PreprocessParams()
|
||||
params.dstdir = dst
|
||||
params.flip = process_flip
|
||||
params.process_caption_only = process_caption_only
|
||||
params.process_caption = process_caption
|
||||
params.process_caption_deepbooru = process_caption_deepbooru
|
||||
params.preprocess_txt_action = preprocess_txt_action
|
||||
pbar = tqdm(files)
|
||||
for index, imagefile in enumerate(pbar):
|
||||
params.subindex = 0
|
||||
filename = os.path.join(src, imagefile)
|
||||
try:
|
||||
img = Image.open(filename)
|
||||
img = ImageOps.exif_transpose(img)
|
||||
if not process_keep_channels:
|
||||
img = img.convert("RGB")
|
||||
except Exception:
|
||||
continue
|
||||
|
||||
description = f"Preprocessing image {index + 1}/{len(files)}"
|
||||
pbar.set_description(description)
|
||||
shared.state.textinfo = description
|
||||
params.src = filename
|
||||
existing_caption = None
|
||||
existing_caption_filename = f"{os.path.splitext(filename)[0]}.txt"
|
||||
if os.path.exists(existing_caption_filename):
|
||||
with open(existing_caption_filename, 'r', encoding="utf8") as file:
|
||||
existing_caption = file.read()
|
||||
else:
|
||||
existing_caption_filename = None
|
||||
if shared.state.interrupted:
|
||||
break
|
||||
if img.height > img.width:
|
||||
ratio = (img.width * height) / (img.height * width)
|
||||
inverse_xy = False
|
||||
else:
|
||||
ratio = (img.height * width) / (img.width * height)
|
||||
inverse_xy = True
|
||||
process_default_resize = True
|
||||
if process_split and ratio < 1.0 and ratio <= split_threshold:
|
||||
for splitted in split_pic(img, inverse_xy, width, height, overlap_ratio):
|
||||
save_pic(splitted, index, params, existing_caption=existing_caption, existing_caption_filename=existing_caption_filename)
|
||||
process_default_resize = False
|
||||
if process_focal_crop and img.height != img.width:
|
||||
dnn_model_path = None
|
||||
try:
|
||||
dnn_model_path = autocrop.download_and_cache_models(os.path.join(paths.models_path, "opencv"))
|
||||
except Exception as e:
|
||||
shared.log.error(f"TI unable to load face detection model for auto crop selection. Falling back to lower quality haar method. {e}")
|
||||
autocrop_settings = autocrop.Settings(
|
||||
crop_width = width,
|
||||
crop_height = height,
|
||||
face_points_weight = process_focal_crop_face_weight,
|
||||
entropy_points_weight = process_focal_crop_entropy_weight,
|
||||
corner_points_weight = process_focal_crop_edges_weight,
|
||||
annotate_image = process_focal_crop_debug,
|
||||
dnn_model_path = dnn_model_path,
|
||||
)
|
||||
for focal in autocrop.crop_image(img, autocrop_settings):
|
||||
save_pic(focal, index, params, existing_caption=existing_caption)
|
||||
process_default_resize = False
|
||||
|
||||
if process_multicrop:
|
||||
cropped = multicrop_pic(img, process_multicrop_mindim, process_multicrop_maxdim, process_multicrop_minarea, process_multicrop_maxarea, process_multicrop_objective, process_multicrop_threshold)
|
||||
if cropped is not None:
|
||||
save_pic(cropped, index, params, existing_caption=existing_caption)
|
||||
else:
|
||||
shared.log.error(f"TI skipped {img.width}x{img.height} image {filename} (can't find suitable size within error threshold)")
|
||||
process_default_resize = False
|
||||
if process_keep_original_size:
|
||||
save_pic(img, index, params, existing_caption=existing_caption)
|
||||
process_default_resize = False
|
||||
if process_default_resize:
|
||||
img = images.resize_image(1, img, width, height)
|
||||
save_pic(img, index, params, existing_caption=existing_caption)
|
||||
shared.state.nextjob()
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 478 KiB |
@@ -1,25 +1,17 @@
|
||||
from typing import List, Optional, Union
|
||||
import csv
|
||||
import html
|
||||
from typing import List, Union
|
||||
import os
|
||||
import time
|
||||
from collections import namedtuple
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
import safetensors.torch
|
||||
import numpy as np
|
||||
from PIL import Image, PngImagePlugin
|
||||
from installer import install
|
||||
from modules import shared, devices, processing, sd_models, images, errors
|
||||
import modules.textual_inversion.dataset
|
||||
from modules.textual_inversion.learn_schedule import LearnRateScheduler
|
||||
from modules.textual_inversion.image_embedding import embedding_to_b64, embedding_from_b64, insert_image_data_embed, extract_image_data_embed, caption_image_overlay
|
||||
from modules.textual_inversion.ti_logging import save_settings_to_file
|
||||
from PIL import Image
|
||||
from modules import shared, devices, sd_models, errors
|
||||
from modules.textual_inversion.image_embedding import embedding_from_b64, extract_image_data_embed
|
||||
from modules.files_cache import directory_files, directory_mtime, extension_filter
|
||||
|
||||
|
||||
debug = shared.log.trace if os.environ.get('SD_TI_DEBUG', None) is not None else lambda *args, **kwargs: None
|
||||
debug('Trace: TEXTUAL INVERSION')
|
||||
|
||||
TokenToAdd = namedtuple("TokenToAdd", ["clip_l", "clip_g"])
|
||||
TextualInversionTemplate = namedtuple("TextualInversionTemplate", ["name", "path"])
|
||||
textual_inversion_templates = {}
|
||||
@@ -386,374 +378,3 @@ class EmbeddingDatabase:
|
||||
if tokens[offset:offset + len(ids)] == ids:
|
||||
return embedding, len(ids)
|
||||
return None, None
|
||||
|
||||
|
||||
def create_embedding(name, num_vectors_per_token, overwrite_old, init_text='*'):
|
||||
cond_model = shared.sd_model.cond_stage_model
|
||||
with devices.autocast():
|
||||
cond_model([""]) # will send cond model to GPU if lowvram/medvram is active
|
||||
#cond_model expects at least some text, so we provide '*' as backup.
|
||||
embedded = cond_model.encode_embedding_init_text(init_text or '*', num_vectors_per_token)
|
||||
vec = torch.zeros((num_vectors_per_token, embedded.shape[1]), device=devices.device)
|
||||
#Only copy if we provided an init_text, otherwise keep vectors as zeros
|
||||
if init_text:
|
||||
for i in range(num_vectors_per_token):
|
||||
vec[i] = embedded[i * int(embedded.shape[0]) // num_vectors_per_token]
|
||||
# Remove illegal characters from name.
|
||||
name = "".join( x for x in name if (x.isalnum() or x in "._- "))
|
||||
fn = os.path.join(shared.opts.embeddings_dir, f"{name}.pt")
|
||||
if not overwrite_old and os.path.exists(fn):
|
||||
shared.log.warning(f"Embedding already exists: {fn}")
|
||||
else:
|
||||
embedding = Embedding(vec=vec, name=name, filename=fn)
|
||||
embedding.step = 0
|
||||
embedding.save(fn)
|
||||
shared.log.info(f'Created embedding: {fn} vectors {num_vectors_per_token} init {init_text}')
|
||||
return fn
|
||||
|
||||
|
||||
def write_loss(log_directory, filename, step, epoch_len, values):
|
||||
if shared.opts.training_write_csv_every == 0:
|
||||
return
|
||||
if step % shared.opts.training_write_csv_every != 0:
|
||||
return
|
||||
write_csv_header = False if os.path.exists(os.path.join(log_directory, filename)) else True
|
||||
with open(os.path.join(log_directory, filename), "a+", newline='', encoding='utf-8') as fout:
|
||||
csv_writer = csv.DictWriter(fout, fieldnames=["step", "epoch", "epoch_step", *(values.keys())])
|
||||
if write_csv_header:
|
||||
csv_writer.writeheader()
|
||||
epoch = (step - 1) // epoch_len
|
||||
epoch_step = (step - 1) % epoch_len
|
||||
csv_writer.writerow({
|
||||
"step": step,
|
||||
"epoch": epoch,
|
||||
"epoch_step": epoch_step,
|
||||
**values,
|
||||
})
|
||||
|
||||
|
||||
def tensorboard_setup(log_directory):
|
||||
install('tensorboard')
|
||||
from torch.utils.tensorboard import SummaryWriter
|
||||
os.makedirs(os.path.join(log_directory, "tensorboard"), exist_ok=True)
|
||||
return SummaryWriter(
|
||||
log_dir=os.path.join(log_directory, "tensorboard"),
|
||||
flush_secs=shared.opts.training_tensorboard_flush_every)
|
||||
|
||||
|
||||
def tensorboard_add(tensorboard_writer, loss, global_step, step, learn_rate, epoch_num):
|
||||
tensorboard_add_scaler(tensorboard_writer, "Loss/train", loss, global_step)
|
||||
tensorboard_add_scaler(tensorboard_writer, f"Loss/train/epoch-{epoch_num}", loss, step)
|
||||
tensorboard_add_scaler(tensorboard_writer, "Learn rate/train", learn_rate, global_step)
|
||||
tensorboard_add_scaler(tensorboard_writer, f"Learn rate/train/epoch-{epoch_num}", learn_rate, step)
|
||||
|
||||
|
||||
def tensorboard_add_scaler(tensorboard_writer, tag, value, step):
|
||||
tensorboard_writer.add_scalar(tag=tag, scalar_value=value, global_step=step)
|
||||
|
||||
|
||||
def tensorboard_add_image(tensorboard_writer, tag, pil_image, step):
|
||||
# Convert a pil image to a torch tensor
|
||||
img_tensor = torch.as_tensor(np.array(pil_image, copy=True))
|
||||
img_tensor = img_tensor.view(pil_image.size[1], pil_image.size[0], len(pil_image.getbands()))
|
||||
img_tensor = img_tensor.permute((2, 0, 1))
|
||||
tensorboard_writer.add_image(tag, img_tensor, global_step=step)
|
||||
|
||||
|
||||
def validate_train_inputs(model_name, learn_rate, batch_size, gradient_step, data_root, template_file, template_filename, steps, save_model_every, create_image_every, name="embedding"):
|
||||
assert model_name, f"{name} not selected"
|
||||
assert learn_rate, "Learning rate is empty or 0"
|
||||
assert isinstance(batch_size, int), "Batch size must be integer"
|
||||
assert batch_size > 0, "Batch size must be positive"
|
||||
assert isinstance(gradient_step, int), "Gradient accumulation step must be integer"
|
||||
assert gradient_step > 0, "Gradient accumulation step must be positive"
|
||||
assert data_root, "Dataset directory is empty"
|
||||
assert os.path.isdir(data_root), "Dataset directory doesn't exist"
|
||||
assert os.listdir(data_root), "Dataset directory is empty"
|
||||
assert template_filename, "Prompt template file not selected"
|
||||
assert template_file, f"Prompt template file {template_filename} not found"
|
||||
assert os.path.isfile(template_file.path), f"Prompt template file {template_filename} doesn't exist"
|
||||
assert steps, "Max steps is empty or 0"
|
||||
assert isinstance(steps, int), "Max steps must be integer"
|
||||
assert steps > 0, "Max steps must be positive"
|
||||
assert isinstance(save_model_every, int), "Save {name} must be integer"
|
||||
assert save_model_every >= 0, "Save {name} must be positive or 0"
|
||||
assert isinstance(create_image_every, int), "Create image must be integer"
|
||||
assert create_image_every >= 0, "Create image must be positive or 0"
|
||||
|
||||
|
||||
def train_embedding(id_task, embedding_name, learn_rate, batch_size, gradient_step, data_root, log_directory, training_width, training_height, varsize, steps, clip_grad_mode, clip_grad_value, shuffle_tags, tag_drop_out, latent_sampling_method, use_weight, create_image_every, save_embedding_every, template_filename, save_image_with_stored_embedding, preview_from_txt2img, preview_prompt, preview_negative_prompt, preview_steps, preview_sampler_index, preview_cfg_scale, preview_seed, preview_width, preview_height): # pylint: disable=unused-argument
|
||||
from modules import sd_hijack, sd_hijack_checkpoint
|
||||
|
||||
shared.log.debug(f'train_embedding: embedding_name={embedding_name}|learn_rate={learn_rate}|batch_size={batch_size}|gradient_step={gradient_step}|data_root={data_root}|log_directory={log_directory}|training_width={training_width}|training_height={training_height}|varsize={varsize}|steps={steps}|clip_grad_mode={clip_grad_mode}|clip_grad_value={clip_grad_value}|shuffle_tags={shuffle_tags}|tag_drop_out={tag_drop_out}|latent_sampling_method={latent_sampling_method}|use_weight={use_weight}|create_image_every={create_image_every}|save_embedding_every={save_embedding_every}|template_filename={template_filename}|save_image_with_stored_embedding={save_image_with_stored_embedding}|preview_from_txt2img={preview_from_txt2img}|preview_prompt={preview_prompt}|preview_negative_prompt={preview_negative_prompt}|preview_steps={preview_steps}|preview_sampler_index={preview_sampler_index}|preview_cfg_scale={preview_cfg_scale}|preview_seed={preview_seed}|preview_width={preview_width}|preview_height={preview_height}')
|
||||
save_embedding_every = save_embedding_every or 0
|
||||
create_image_every = create_image_every or 0
|
||||
template_file = textual_inversion_templates.get(template_filename, None)
|
||||
validate_train_inputs(embedding_name, learn_rate, batch_size, gradient_step, data_root, template_file, template_filename, steps, save_embedding_every, create_image_every, name="embedding")
|
||||
if log_directory is None or log_directory == '':
|
||||
log_directory = f"{os.path.join(shared.cmd_opts.data_dir, 'train/log/embeddings')}"
|
||||
template_file = template_file.path
|
||||
|
||||
shared.state.job = "train"
|
||||
shared.state.textinfo = "Initializing textual inversion training..."
|
||||
shared.state.job_count = steps
|
||||
|
||||
filename = os.path.join(shared.opts.embeddings_dir, f'{embedding_name}.pt')
|
||||
|
||||
if log_directory == '':
|
||||
log_directory = f"{os.path.join(shared.cmd_opts.data_dir, 'train/log/embeddings')}"
|
||||
log_directory = os.path.join(log_directory, embedding_name)
|
||||
unload = shared.opts.unload_models_when_training
|
||||
|
||||
if save_embedding_every > 0:
|
||||
embedding_dir = os.path.join(log_directory, "embeddings")
|
||||
os.makedirs(embedding_dir, exist_ok=True)
|
||||
else:
|
||||
embedding_dir = None
|
||||
|
||||
if create_image_every > 0:
|
||||
images_dir = os.path.join(log_directory, "images")
|
||||
os.makedirs(images_dir, exist_ok=True)
|
||||
else:
|
||||
images_dir = None
|
||||
|
||||
if create_image_every > 0 and save_image_with_stored_embedding:
|
||||
images_embeds_dir = os.path.join(log_directory, "image_embeddings")
|
||||
os.makedirs(images_embeds_dir, exist_ok=True)
|
||||
else:
|
||||
images_embeds_dir = None
|
||||
|
||||
hijack = sd_hijack.model_hijack
|
||||
embedding = hijack.embedding_db.word_embeddings[embedding_name]
|
||||
checkpoint = sd_models.select_checkpoint()
|
||||
initial_step = embedding.step or 0
|
||||
if initial_step >= steps:
|
||||
shared.state.textinfo = "Model has already been trained beyond specified max steps"
|
||||
return embedding, filename
|
||||
scheduler = LearnRateScheduler(learn_rate, steps, initial_step)
|
||||
clip_grad = torch.nn.utils.clip_grad_value_ if clip_grad_mode == "value" else \
|
||||
torch.nn.utils.clip_grad_norm_ if clip_grad_mode == "norm" else \
|
||||
None
|
||||
if clip_grad:
|
||||
clip_grad_sched = LearnRateScheduler(clip_grad_value, steps, initial_step, verbose=False)
|
||||
# dataset loading may take a while, so input validations and early returns should be done before this
|
||||
shared.state.textinfo = f"Preparing dataset from {html.escape(data_root)}..."
|
||||
old_parallel_processing_allowed = shared.parallel_processing_allowed
|
||||
|
||||
if shared.opts.training_enable_tensorboard:
|
||||
tensorboard_writer = tensorboard_setup(log_directory)
|
||||
|
||||
pin_memory = shared.opts.pin_memory
|
||||
# init dataset
|
||||
ds = modules.textual_inversion.dataset.PersonalizedBase(data_root=data_root, width=training_width, height=training_height, repeats=shared.opts.training_image_repeats_per_epoch, placeholder_token=embedding_name, model=shared.sd_model, cond_model=shared.sd_model.cond_stage_model, device=devices.device, template_file=template_file, batch_size=batch_size, gradient_step=gradient_step, shuffle_tags=shuffle_tags, tag_drop_out=tag_drop_out, latent_sampling_method=latent_sampling_method, varsize=varsize, use_weight=use_weight)
|
||||
|
||||
if shared.opts.save_training_settings_to_txt:
|
||||
save_settings_to_file(log_directory, {**dict(model_name=checkpoint.model_name, model_hash=checkpoint.shorthash, num_of_dataset_images=len(ds), num_vectors_per_token=len(embedding.vec)), **locals()})
|
||||
latent_sampling_method = ds.latent_sampling_method
|
||||
# init dataloader
|
||||
dl = modules.textual_inversion.dataset.PersonalizedDataLoader(ds, latent_sampling_method=latent_sampling_method, batch_size=ds.batch_size, pin_memory=pin_memory)
|
||||
if unload:
|
||||
shared.parallel_processing_allowed = False
|
||||
shared.sd_model.first_stage_model.to(devices.cpu)
|
||||
|
||||
embedding.vec.requires_grad = True
|
||||
optimizer = torch.optim.AdamW([embedding.vec], lr=scheduler.learn_rate, weight_decay=0.0)
|
||||
if shared.opts.save_optimizer_state:
|
||||
optimizer_state_dict = None
|
||||
if os.path.exists(f"{filename}.optim"):
|
||||
optimizer_saved_dict = torch.load(f"{filename}.optim", map_location='cpu')
|
||||
if embedding.checksum() == optimizer_saved_dict.get('hash', None):
|
||||
optimizer_state_dict = optimizer_saved_dict.get('optimizer_state_dict', None)
|
||||
if optimizer_state_dict is not None:
|
||||
optimizer.load_state_dict(optimizer_state_dict)
|
||||
shared.log.info("Load existing optimizer from checkpoint")
|
||||
else:
|
||||
shared.log.info("No saved optimizer exists in checkpoint")
|
||||
|
||||
scaler = torch.cuda.amp.GradScaler()
|
||||
|
||||
batch_size = ds.batch_size
|
||||
gradient_step = ds.gradient_step
|
||||
# n steps = batch_size * gradient_step * n image processed
|
||||
steps_per_epoch = len(ds) // batch_size // gradient_step
|
||||
max_steps_per_epoch = len(ds) // batch_size - (len(ds) // batch_size) % gradient_step
|
||||
loss_step = 0
|
||||
_loss_step = 0 #internal
|
||||
last_saved_file = "<none>"
|
||||
last_saved_image = "<none>"
|
||||
forced_filename = "<none>"
|
||||
embedding_yet_to_be_embedded = False
|
||||
is_training_inpainting_model = shared.sd_model.model.conditioning_key in {'hybrid', 'concat'}
|
||||
img_c = None
|
||||
|
||||
pbar = tqdm(total=steps - initial_step)
|
||||
try:
|
||||
sd_hijack_checkpoint.add()
|
||||
for _i in range((steps-initial_step) * gradient_step):
|
||||
if scheduler.finished:
|
||||
break
|
||||
if shared.state.interrupted:
|
||||
break
|
||||
for j, batch in enumerate(dl):
|
||||
# works as a drop_last=True for gradient accumulation
|
||||
if j == max_steps_per_epoch:
|
||||
break
|
||||
scheduler.apply(optimizer, embedding.step)
|
||||
if scheduler.finished:
|
||||
break
|
||||
if shared.state.interrupted:
|
||||
break
|
||||
if clip_grad:
|
||||
clip_grad_sched.step(embedding.step)
|
||||
with devices.autocast():
|
||||
x = batch.latent_sample.to(devices.device, non_blocking=pin_memory)
|
||||
if use_weight:
|
||||
w = batch.weight.to(devices.device, non_blocking=pin_memory)
|
||||
c = shared.sd_model.cond_stage_model(batch.cond_text)
|
||||
if is_training_inpainting_model:
|
||||
if img_c is None:
|
||||
img_c = processing.txt2img_image_conditioning(shared.sd_model, c, training_width, training_height)
|
||||
cond = {"c_concat": [img_c], "c_crossattn": [c]}
|
||||
else:
|
||||
cond = c
|
||||
if use_weight:
|
||||
loss = shared.sd_model.weighted_forward(x, cond, w)[0] / gradient_step
|
||||
del w
|
||||
else:
|
||||
loss = shared.sd_model.forward(x, cond)[0] / gradient_step
|
||||
del x
|
||||
_loss_step += loss.item()
|
||||
|
||||
scaler.scale(loss).backward()
|
||||
# go back until we reach gradient accumulation steps
|
||||
if (j + 1) % gradient_step != 0:
|
||||
continue
|
||||
if clip_grad:
|
||||
clip_grad(embedding.vec, clip_grad_sched.learn_rate)
|
||||
|
||||
scaler.step(optimizer)
|
||||
scaler.update()
|
||||
embedding.step += 1
|
||||
pbar.update()
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
loss_step = _loss_step
|
||||
_loss_step = 0
|
||||
steps_done = embedding.step + 1
|
||||
epoch_num = embedding.step // steps_per_epoch
|
||||
|
||||
description = f"Training textual inversion step {embedding.step} loss: {loss_step:.5f} lr: {scheduler.learn_rate:.5f}"
|
||||
pbar.set_description(description)
|
||||
if embedding_dir is not None and steps_done % save_embedding_every == 0:
|
||||
# Before saving, change name to match current checkpoint.
|
||||
embedding_name_every = f'{embedding_name}-{steps_done}'
|
||||
last_saved_file = os.path.join(embedding_dir, f'{embedding_name_every}.pt')
|
||||
save_embedding(embedding, optimizer, checkpoint, embedding_name_every, last_saved_file, remove_cached_checksum=True)
|
||||
embedding_yet_to_be_embedded = True
|
||||
|
||||
write_loss(log_directory, f"{embedding_name}.csv", embedding.step, steps_per_epoch, { "loss": f"{loss_step:.7f}", "learn_rate": scheduler.learn_rate })
|
||||
|
||||
if images_dir is not None and steps_done % create_image_every == 0:
|
||||
forced_filename = f'{embedding_name}-{steps_done}'
|
||||
last_saved_image = os.path.join(images_dir, forced_filename)
|
||||
shared.sd_model.first_stage_model.to(devices.device)
|
||||
|
||||
p = processing.StableDiffusionProcessingTxt2Img(
|
||||
sd_model=shared.sd_model,
|
||||
do_not_save_grid=True,
|
||||
do_not_save_samples=True,
|
||||
do_not_reload_embeddings=True,
|
||||
)
|
||||
|
||||
if preview_from_txt2img:
|
||||
p.prompt = preview_prompt
|
||||
p.negative_prompt = preview_negative_prompt
|
||||
p.steps = preview_steps
|
||||
p.sampler_name = processing.get_sampler_name(preview_sampler_index)
|
||||
p.cfg_scale = preview_cfg_scale
|
||||
p.seed = preview_seed
|
||||
p.width = preview_width
|
||||
p.height = preview_height
|
||||
else:
|
||||
p.prompt = batch.cond_text[0]
|
||||
p.steps = 20
|
||||
p.width = training_width
|
||||
p.height = training_height
|
||||
|
||||
preview_text = p.prompt
|
||||
processed = processing.process_images(p)
|
||||
image = processed.images[0] if len(processed.images) > 0 else None
|
||||
|
||||
if unload:
|
||||
shared.sd_model.first_stage_model.to(devices.cpu)
|
||||
|
||||
if image is not None:
|
||||
shared.state.assign_current_image(image)
|
||||
last_saved_image, _last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False)
|
||||
last_saved_image += f", prompt: {preview_text}"
|
||||
if shared.opts.training_enable_tensorboard and shared.opts.training_tensorboard_save_images:
|
||||
tensorboard_add_image(tensorboard_writer, f"Validation at epoch {epoch_num}", image, embedding.step)
|
||||
|
||||
if save_image_with_stored_embedding and os.path.exists(last_saved_file) and embedding_yet_to_be_embedded:
|
||||
last_saved_image_chunks = os.path.join(images_embeds_dir, f'{embedding_name}-{steps_done}.png')
|
||||
info = PngImagePlugin.PngInfo()
|
||||
data = torch.load(last_saved_file)
|
||||
info.add_text("sd-ti-embedding", embedding_to_b64(data))
|
||||
title = f"<{data.get('name', '???')}>"
|
||||
try:
|
||||
vectorSize = list(data['string_to_param'].values())[0].shape[0]
|
||||
except Exception:
|
||||
vectorSize = '?'
|
||||
checkpoint = sd_models.select_checkpoint()
|
||||
footer_left = checkpoint.model_name
|
||||
footer_mid = f'[{checkpoint.shorthash}]'
|
||||
footer_right = f'{vectorSize}v {steps_done}s'
|
||||
captioned_image = caption_image_overlay(image, title, footer_left, footer_mid, footer_right)
|
||||
captioned_image = insert_image_data_embed(captioned_image, data)
|
||||
captioned_image.save(last_saved_image_chunks, "PNG", pnginfo=info)
|
||||
embedding_yet_to_be_embedded = False
|
||||
|
||||
last_saved_image, _last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False)
|
||||
last_saved_image += f", prompt: {preview_text}"
|
||||
|
||||
shared.state.job_no = embedding.step
|
||||
shared.state.textinfo = f"""
|
||||
<p>
|
||||
Loss: {loss_step:.7f}<br/>
|
||||
Step: {steps_done}<br/>
|
||||
Last prompt: {html.escape(batch.cond_text[0])}<br/>
|
||||
Last saved embedding: {html.escape(last_saved_file)}<br/>
|
||||
Last saved image: {html.escape(last_saved_image)}<br/>
|
||||
</p>
|
||||
"""
|
||||
filename = os.path.join(shared.opts.embeddings_dir, f'{embedding_name}.pt')
|
||||
save_embedding(embedding, optimizer, checkpoint, embedding_name, filename, remove_cached_checksum=True)
|
||||
except Exception as e:
|
||||
errors.display(e, 'embedding train')
|
||||
finally:
|
||||
pbar.leave = False
|
||||
pbar.close()
|
||||
shared.sd_model.first_stage_model.to(devices.device)
|
||||
shared.parallel_processing_allowed = old_parallel_processing_allowed
|
||||
sd_hijack_checkpoint.remove()
|
||||
return embedding, filename
|
||||
|
||||
|
||||
def save_embedding(embedding, optimizer, checkpoint, embedding_name, filename, remove_cached_checksum=True):
|
||||
old_embedding_name = embedding.name
|
||||
old_sd_checkpoint = embedding.sd_checkpoint if hasattr(embedding, "sd_checkpoint") else None
|
||||
old_sd_checkpoint_name = embedding.sd_checkpoint_name if hasattr(embedding, "sd_checkpoint_name") else None
|
||||
old_cached_checksum = embedding.cached_checksum if hasattr(embedding, "cached_checksum") else None
|
||||
try:
|
||||
embedding.sd_checkpoint = checkpoint.shorthash
|
||||
embedding.sd_checkpoint_name = checkpoint.model_name
|
||||
if remove_cached_checksum:
|
||||
embedding.cached_checksum = None
|
||||
embedding.name = embedding_name
|
||||
embedding.optimizer_state_dict = optimizer.state_dict()
|
||||
embedding.save(filename)
|
||||
except Exception:
|
||||
embedding.sd_checkpoint = old_sd_checkpoint
|
||||
embedding.sd_checkpoint_name = old_sd_checkpoint_name
|
||||
embedding.name = old_embedding_name
|
||||
embedding.cached_checksum = old_cached_checksum
|
||||
raise
|
||||
|
||||
@@ -1,22 +0,0 @@
|
||||
import datetime
|
||||
import json
|
||||
import os
|
||||
|
||||
saved_params_shared = {"model_name", "model_hash", "initial_step", "num_of_dataset_images", "learn_rate", "batch_size", "clip_grad_mode", "clip_grad_value", "gradient_step", "data_root", "log_directory", "training_width", "training_height", "steps", "create_image_every", "template_file", "latent_sampling_method"}
|
||||
saved_params_ti = {"embedding_name", "num_vectors_per_token", "save_embedding_every", "save_image_with_stored_embedding"}
|
||||
saved_params_hypernet = {"hypernetwork_name", "layer_structure", "activation_func", "weight_init", "add_layer_norm", "use_dropout", "save_hypernetwork_every"}
|
||||
saved_params_all = saved_params_shared | saved_params_ti | saved_params_hypernet
|
||||
saved_params_previews = {"preview_prompt", "preview_negative_prompt", "preview_steps", "preview_sampler_index", "preview_cfg_scale", "preview_seed", "preview_width", "preview_height"}
|
||||
|
||||
|
||||
def save_settings_to_file(log_directory, all_params):
|
||||
now = datetime.datetime.now()
|
||||
params = {"datetime": now.strftime("%Y-%m-%d %H:%M:%S")}
|
||||
keys = saved_params_all
|
||||
if all_params.get('preview_from_txt2img'):
|
||||
keys = keys | saved_params_previews
|
||||
params.update({k: v for k, v in all_params.items() if k in keys})
|
||||
filename = f"settings-{now.strftime('%Y-%m-%d_%H-%M-%S')}.json"
|
||||
with open(os.path.join(log_directory, filename), "w", encoding='utf-8') as file:
|
||||
print(f'Training settings file: {os.path.join(log_directory, filename)}')
|
||||
json.dump(params, file, indent=2)
|
||||
@@ -1,35 +0,0 @@
|
||||
import html
|
||||
import gradio as gr
|
||||
import modules.textual_inversion.textual_inversion
|
||||
import modules.textual_inversion.preprocess
|
||||
from modules import shared
|
||||
|
||||
|
||||
def create_embedding(name, initialization_text, nvpt, overwrite_old):
|
||||
from modules import sd_hijack
|
||||
filename = modules.textual_inversion.textual_inversion.create_embedding(name, nvpt, overwrite_old, init_text=initialization_text)
|
||||
sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings()
|
||||
return gr.Dropdown.update(choices=sorted(sd_hijack.model_hijack.embedding_db.word_embeddings.keys())), f"Created: {filename}", ""
|
||||
|
||||
|
||||
def preprocess(*args):
|
||||
modules.textual_inversion.preprocess.preprocess(*args)
|
||||
return f"Preprocessing {'interrupted' if shared.state.interrupted else 'finished'}.", ""
|
||||
|
||||
|
||||
def train_embedding(*args):
|
||||
from modules import sd_hijack
|
||||
assert not shared.cmd_opts.lowvram, 'Training models with lowvram not possible'
|
||||
apply_optimizations = False
|
||||
try:
|
||||
if not apply_optimizations:
|
||||
sd_hijack.undo_optimizations()
|
||||
embedding, filename = modules.textual_inversion.textual_inversion.train_embedding(*args)
|
||||
res = f"Training {'interrupted' if shared.state.interrupted else 'finished'} at {embedding.step} steps. Embedding saved to {html.escape(filename)}"
|
||||
return res, ""
|
||||
except Exception as e:
|
||||
shared.log.error(f"Exception in train_embedding: {e}")
|
||||
raise RuntimeError from e
|
||||
finally:
|
||||
if not apply_optimizations:
|
||||
sd_hijack.apply_optimizations()
|
||||
@@ -144,11 +144,6 @@ def create_ui(startup_timer = None):
|
||||
ui_postprocessing.create_ui()
|
||||
timer.startup.record("ui-extras")
|
||||
|
||||
with gr.Blocks(analytics_enabled=False) as train_interface:
|
||||
from modules import ui_train
|
||||
ui_train.create_ui()
|
||||
timer.startup.record("ui-train")
|
||||
|
||||
with gr.Blocks(analytics_enabled=False) as models_interface:
|
||||
from modules import ui_models
|
||||
ui_models.create_ui()
|
||||
@@ -377,7 +372,6 @@ def create_ui(startup_timer = None):
|
||||
interfaces += [(control_interface, "Control", "control")] if control_interface is not None else []
|
||||
interfaces += [(extras_interface, "Process", "process")]
|
||||
interfaces += [(interrogate_interface, "Interrogate", "interrogate")]
|
||||
interfaces += [(train_interface, "Train", "train")]
|
||||
interfaces += [(models_interface, "Models", "models")]
|
||||
interfaces += script_callbacks.ui_tabs_callback()
|
||||
interfaces += [(settings_interface, "System", "system")]
|
||||
|
||||
@@ -38,7 +38,7 @@ def create_ui():
|
||||
modules.scripts.scripts_current = modules.scripts.scripts_img2img
|
||||
modules.scripts.scripts_img2img.initialize_scripts(is_img2img=True)
|
||||
with gr.Blocks(analytics_enabled=False) as _img2img_interface:
|
||||
img2img_prompt, img2img_prompt_styles, img2img_negative_prompt, submit, img2img_paste, img2img_extra_networks_button, img2img_token_counter, img2img_token_button, img2img_negative_token_counter, img2img_negative_token_button = ui_sections.create_toprow(is_img2img=True, id_part="img2img")
|
||||
img2img_prompt, img2img_prompt_styles, img2img_negative_prompt, img2img_submit, img2img_paste, img2img_extra_networks_button, img2img_token_counter, img2img_token_button, img2img_negative_token_counter, img2img_negative_token_button = ui_sections.create_toprow(is_img2img=True, id_part="img2img")
|
||||
img2img_prompt_img = gr.File(label="", elem_id="img2img_prompt_image", file_count="single", type="binary", visible=False)
|
||||
|
||||
with gr.Row(variant='compact', elem_id="img2img_extra_networks", visible=False) as extra_networks_ui:
|
||||
@@ -207,7 +207,8 @@ def create_ui():
|
||||
show_progress=False,
|
||||
)
|
||||
img2img_prompt.submit(**img2img_dict)
|
||||
submit.click(**img2img_dict)
|
||||
img2img_negative_prompt.submit(**img2img_dict)
|
||||
img2img_submit.click(**img2img_dict)
|
||||
dummy_component = gr.Textbox(visible=False, value='dummy')
|
||||
|
||||
interrogate_args = dict(
|
||||
|
||||
+23
-11
@@ -185,6 +185,10 @@ def create_ui():
|
||||
movement = gr.Label(label="Movement", num_top_classes=5)
|
||||
trending = gr.Label(label="Trending", num_top_classes=5)
|
||||
flavor = gr.Label(label="Flavor", num_top_classes=5)
|
||||
with gr.Row():
|
||||
clip_model = gr.Dropdown([], value='ViT-L-14/openai', label='CLIP Model')
|
||||
ui_common.create_refresh_button(clip_model, get_models, lambda: {"choices": get_models()}, 'refresh_interrogate_models')
|
||||
mode = gr.Radio(['best', 'fast', 'classic', 'caption', 'negative'], label='Mode', value='best')
|
||||
with gr.Row():
|
||||
btn_interrogate_img = gr.Button("Interrogate", variant='primary')
|
||||
btn_analyze_img = gr.Button("Analyze", variant='primary')
|
||||
@@ -193,6 +197,9 @@ def create_ui():
|
||||
buttons = parameters_copypaste.create_buttons(["txt2img", "img2img", "extras", "control"])
|
||||
for tabname, button in buttons.items():
|
||||
parameters_copypaste.register_paste_params_button(parameters_copypaste.ParamBinding(paste_button=button, tabname=tabname, source_text_component=prompt, source_image_component=image,))
|
||||
btn_interrogate_img.click(interrogate_image, inputs=[image, clip_model, mode], outputs=prompt)
|
||||
btn_analyze_img.click(analyze_image, inputs=[image, clip_model], outputs=[medium, artist, movement, trending, flavor])
|
||||
btn_unload.click(unload)
|
||||
with gr.Tab("Batch"):
|
||||
with gr.Row():
|
||||
batch_files = gr.File(label="Files", show_label=True, file_count='multiple', file_types=['image'], type='file', interactive=True, height=100)
|
||||
@@ -204,16 +211,21 @@ def create_ui():
|
||||
batch = gr.Text(label="Prompts", lines=10)
|
||||
with gr.Row():
|
||||
write = gr.Checkbox(label='Write prompts to files', value=False)
|
||||
with gr.Row():
|
||||
clip_model = gr.Dropdown([], value='ViT-L-14/openai', label='CLIP Model')
|
||||
ui_common.create_refresh_button(clip_model, get_models, lambda: {"choices": get_models()}, 'refresh_interrogate_models')
|
||||
with gr.Row():
|
||||
btn_interrogate_batch = gr.Button("Interrogate", variant='primary')
|
||||
with gr.Column():
|
||||
with gr.Row():
|
||||
# clip_model = gr.Dropdown(get_models(), value='ViT-L-14/openai', label='CLIP Model')
|
||||
clip_model = gr.Dropdown([], value='ViT-L-14/openai', label='CLIP Model')
|
||||
ui_common.create_refresh_button(clip_model, get_models, lambda: {"choices": get_models()}, 'refresh_interrogate_models')
|
||||
with gr.Row():
|
||||
mode = gr.Radio(['best', 'fast', 'classic', 'caption', 'negative'], label='Mode', value='best')
|
||||
btn_interrogate_img.click(interrogate_image, inputs=[image, clip_model, mode], outputs=prompt)
|
||||
btn_analyze_img.click(analyze_image, inputs=[image, clip_model], outputs=[medium, artist, movement, trending, flavor])
|
||||
btn_interrogate_batch.click(interrogate_batch, inputs=[batch_files, batch_folder, batch_str, clip_model, mode, write], outputs=[batch])
|
||||
btn_unload.click(unload)
|
||||
btn_interrogate_batch.click(interrogate_batch, inputs=[batch_files, batch_folder, batch_str, clip_model, mode, write], outputs=[batch])
|
||||
with gr.Tab("VQA"):
|
||||
from modules import vqa
|
||||
with gr.Row():
|
||||
vqa_image = gr.Image(type='pil', label="Image")
|
||||
with gr.Row():
|
||||
vqa_question = gr.Textbox(label="Question")
|
||||
with gr.Row():
|
||||
vqa_answer = gr.Textbox(label="Answer", lines=3)
|
||||
with gr.Row():
|
||||
vqa_model = gr.Dropdown(list(vqa.MODELS), value='None', label='VQA Model')
|
||||
vqa_submit = gr.Button("Interrogate", variant='primary')
|
||||
vqa_submit.click(vqa.interrogate, inputs=[vqa_question, vqa_image, vqa_model], outputs=[vqa_answer])
|
||||
|
||||
+13
-15
@@ -106,7 +106,7 @@ def create_advanced_inputs(tab):
|
||||
cfg_end = gr.Slider(minimum=0.0, maximum=1.0, step=0.1, label='CFG end', value=1.0, elem_id=f"{tab}_cfg_end")
|
||||
with gr.Row():
|
||||
image_cfg_scale = gr.Slider(minimum=0.0, maximum=30.0, step=0.1, label='Secondary guidance', value=6.0, elem_id=f"{tab}_image_cfg_scale")
|
||||
diffusers_guidance_rescale = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Guidance rescale', value=0.7, elem_id=f"{tab}_image_cfg_rescale", visible=shared.backend == shared.Backend.DIFFUSERS)
|
||||
diffusers_guidance_rescale = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Rescale guidance', value=0.7, elem_id=f"{tab}_image_cfg_rescale", visible=shared.backend == shared.Backend.DIFFUSERS)
|
||||
diffusers_sag_scale = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Attention guidance', value=0.0, elem_id=f"{tab}_image_sag_scale", visible=shared.backend == shared.Backend.DIFFUSERS)
|
||||
with gr.Row():
|
||||
clip_skip = gr.Slider(label='CLIP skip', value=1, minimum=0, maximum=12, step=0.1, elem_id=f"{tab}_clip_skip", interactive=True)
|
||||
@@ -115,7 +115,7 @@ def create_advanced_inputs(tab):
|
||||
with gr.Row(elem_id=f"{tab}_advanced_options"):
|
||||
full_quality = gr.Checkbox(label='Full quality', value=True, elem_id=f"{tab}_full_quality")
|
||||
restore_faces = gr.Checkbox(label='Face restore', value=False, visible=len(shared.face_restorers) > 1, elem_id=f"{tab}_restore_faces")
|
||||
tiling = gr.Checkbox(label='Tiling', value=False, elem_id=f"{tab}_tiling", visible=shared.backend == shared.Backend.ORIGINAL)
|
||||
tiling = gr.Checkbox(label='Tiling', value=False, elem_id=f"{tab}_tiling", visible=True)
|
||||
return cfg_scale, clip_skip, image_cfg_scale, diffusers_guidance_rescale, diffusers_sag_scale, cfg_end, full_quality, restore_faces, tiling
|
||||
|
||||
def create_correction_inputs(tab):
|
||||
@@ -185,24 +185,22 @@ def create_sampler_and_steps_selection(choices, tabname):
|
||||
|
||||
|
||||
def create_hires_inputs(tab):
|
||||
with gr.Accordion(open=False, label="Second pass", elem_id=f"{tab}_second_pass", elem_classes=["small-accordion"]):
|
||||
with gr.Accordion(open=False, label="Refine", elem_id=f"{tab}_second_pass", elem_classes=["small-accordion"]):
|
||||
with gr.Group():
|
||||
with gr.Row(elem_id=f"{tab}_hires_row1"):
|
||||
enable_hr = gr.Checkbox(label='Enable second pass', value=False, elem_id=f"{tab}_enable_hr")
|
||||
with gr.Row(elem_id=f"{tab}_hires_row2"):
|
||||
hr_sampler_index = gr.Dropdown(label='Secondary sampler', elem_id=f"{tab}_sampling_alt", choices=[x.name for x in sd_samplers.samplers], value='Default', type="index")
|
||||
denoising_strength = gr.Slider(minimum=0.0, maximum=0.99, step=0.01, label='Denoising strength', value=0.5, elem_id=f"{tab}_denoising_strength")
|
||||
with gr.Row(elem_id=f"{tab}_hires_finalres", variant="compact"):
|
||||
hr_final_resolution = gr.HTML(value="", elem_id=f"{tab}_hr_finalres", label="Upscaled resolution", interactive=False)
|
||||
with gr.Row(elem_id=f"{tab}_hires_fix_row1", variant="compact"):
|
||||
hr_upscaler = gr.Dropdown(label="Upscaler", elem_id=f"{tab}_hr_upscaler", choices=[*shared.latent_upscale_modes, *[x.name for x in shared.sd_upscalers]], value=shared.latent_upscale_default_mode)
|
||||
hr_force = gr.Checkbox(label='Force Hires', value=False, elem_id=f"{tab}_hr_force")
|
||||
with gr.Row(elem_id=f"{tab}_hires_fix_row2", variant="compact"):
|
||||
hr_second_pass_steps = gr.Slider(minimum=0, maximum=99, step=1, label='Hires steps', elem_id=f"{tab}_steps_alt", value=20)
|
||||
hr_scale = gr.Slider(minimum=1.0, maximum=8.0, step=0.05, label="Upscale by", value=2.0, elem_id=f"{tab}_hr_scale")
|
||||
hr_scale = gr.Slider(minimum=0.1, maximum=8.0, step=0.05, label="Rescale by", value=2.0, elem_id=f"{tab}_hr_scale")
|
||||
with gr.Row(elem_id=f"{tab}_hires_fix_row3", variant="compact"):
|
||||
hr_resize_x = gr.Slider(minimum=0, maximum=4096, step=8, label="Resize width to", value=0, elem_id=f"{tab}_hr_resize_x")
|
||||
hr_resize_y = gr.Slider(minimum=0, maximum=4096, step=8, label="Resize height to", value=0, elem_id=f"{tab}_hr_resize_y")
|
||||
hr_resize_x = gr.Slider(minimum=0, maximum=4096, step=8, label="Width resize", value=0, elem_id=f"{tab}_hr_resize_x")
|
||||
hr_resize_y = gr.Slider(minimum=0, maximum=4096, step=8, label="Height resize", value=0, elem_id=f"{tab}_hr_resize_y")
|
||||
with gr.Row(elem_id=f"{tab}_hires_fix_row2", variant="compact"):
|
||||
hr_force = gr.Checkbox(label='Force HiRes', value=False, elem_id=f"{tab}_hr_force")
|
||||
hr_sampler_index = gr.Dropdown(label='Secondary sampler', elem_id=f"{tab}_sampling_alt", choices=[x.name for x in sd_samplers.samplers], value='Default', type="index")
|
||||
with gr.Row(elem_id=f"{tab}_hires_row2"):
|
||||
hr_second_pass_steps = gr.Slider(minimum=0, maximum=99, step=1, label='HiRes steps', elem_id=f"{tab}_steps_alt", value=20)
|
||||
denoising_strength = gr.Slider(minimum=0.0, maximum=0.99, step=0.01, label='Strength', value=0.5, elem_id=f"{tab}_denoising_strength")
|
||||
with gr.Group(visible=shared.backend == shared.Backend.DIFFUSERS):
|
||||
with gr.Row(elem_id=f"{tab}_refiner_row1", variant="compact"):
|
||||
refiner_start = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, label='Refiner start', value=0.8, elem_id=f"{tab}_refiner_start")
|
||||
@@ -211,7 +209,7 @@ def create_hires_inputs(tab):
|
||||
refiner_prompt = gr.Textbox(value='', label='Secondary prompt', elem_id=f"{tab}_refiner_prompt")
|
||||
with gr.Row(elem_id="txt2img_refiner_row4", variant="compact"):
|
||||
refiner_negative = gr.Textbox(value='', label='Secondary negative prompt', elem_id=f"{tab}_refiner_neg_prompt")
|
||||
return enable_hr, hr_sampler_index, denoising_strength, hr_final_resolution, hr_upscaler, hr_force, hr_second_pass_steps, hr_scale, hr_resize_x, hr_resize_y, refiner_steps, refiner_start, refiner_prompt, refiner_negative
|
||||
return enable_hr, hr_sampler_index, denoising_strength, hr_upscaler, hr_force, hr_second_pass_steps, hr_scale, hr_resize_x, hr_resize_y, refiner_steps, refiner_start, refiner_prompt, refiner_negative
|
||||
|
||||
|
||||
def create_resize_inputs(tab, images, scale_visible=True, mode=None, accordion=True, latent=False):
|
||||
|
||||
@@ -1,379 +0,0 @@
|
||||
import os
|
||||
import gradio as gr
|
||||
from modules import script_callbacks, shared
|
||||
from modules.ui_common import create_refresh_button
|
||||
from modules.ui_sections import create_sampler_inputs
|
||||
from modules.call_queue import wrap_gradio_gpu_call
|
||||
|
||||
|
||||
def create_ui():
|
||||
from modules.textual_inversion import textual_inversion
|
||||
import modules.hypernetworks.ui
|
||||
dummy_component = gr.Label(visible=False)
|
||||
|
||||
with gr.Row(elem_id="train_tab"):
|
||||
with gr.Column(elem_id='train_output_container', scale=1):
|
||||
train_output = gr.Text(elem_id="train_output", value="", show_label=False)
|
||||
gr.Gallery(label='Output', show_label=False, elem_id='train_gallery', columns=1)
|
||||
gr.HTML(elem_id="train_progress", value="")
|
||||
train_outcome = gr.HTML(elem_id="train_error", value="")
|
||||
|
||||
with gr.Row(visible=True) as action_pp:
|
||||
process_run = gr.Button(value="Preprocess", variant='primary')
|
||||
process_stop = gr.Button("Stop")
|
||||
|
||||
with gr.Row(visible=False) as action_ti:
|
||||
ti_train = gr.Button(value="Train embedding", variant='primary')
|
||||
ti_stop = gr.Button(value="Stop")
|
||||
|
||||
with gr.Row(visible=False) as action_hn:
|
||||
hn_train = gr.Button(value="Train hypernetwork", variant='primary')
|
||||
hn_stop = gr.Button(value="Stop")
|
||||
|
||||
with gr.Column(elem_id='train_input_container', scale=3):
|
||||
|
||||
with gr.Tabs(elem_id="train_tabs"):
|
||||
def gr_show(visible=True):
|
||||
return {"visible": visible, "__type__": "update"}
|
||||
|
||||
def train_tab_change(tab):
|
||||
if tab == 'ti':
|
||||
return gr_show(False), gr_show(True), gr_show(False)
|
||||
elif tab == 'hn':
|
||||
return gr_show(False), gr_show(False), gr_show(True)
|
||||
elif tab == 'pr':
|
||||
return gr_show(True), gr_show(False), gr_show(False)
|
||||
else:
|
||||
return gr_show(False), gr_show(False), gr_show(False)
|
||||
|
||||
### preview tab
|
||||
|
||||
with gr.Tab(label="Preview settings", id="train_preview_tab") as tab_preview:
|
||||
tab_preview.select(fn=lambda: train_tab_change('pr'), inputs=[], outputs=[action_pp, action_ti, action_hn])
|
||||
prompt = gr.Textbox(label="Prompt", value="", placeholder="Prompt to be used for previews", lines=2)
|
||||
negative = gr.Textbox(label="Negative prompt", value="", placeholder="Negative prompt to be used for previews", lines=2)
|
||||
steps, sampler_index = create_sampler_inputs('train', accordion=False)
|
||||
cfg_scale = gr.Slider(minimum=0.0, maximum=30.0, step=0.1, label='CFG scale', value=6.0)
|
||||
seed = gr.Number(label='Initial seed', value=-1)
|
||||
with gr.Row():
|
||||
width = gr.Slider(minimum=64, maximum=8192, step=8, label="Width", value=512)
|
||||
height = gr.Slider(minimum=64, maximum=8192, step=8, label="Height", value=512)
|
||||
txt2img_preview_params = [prompt, negative, steps, sampler_index, cfg_scale, seed, width, height]
|
||||
|
||||
### preprocess tab
|
||||
|
||||
with gr.Tab(label="Preprocess images", id="preprocess_images") as tab_preprocess:
|
||||
tab_preprocess.select(fn=lambda: train_tab_change('pp'), inputs=[], outputs=[action_pp, action_ti, action_hn])
|
||||
process_src = gr.Textbox(label='Source directory')
|
||||
process_dst = gr.Textbox(label='Destination directory')
|
||||
with gr.Row():
|
||||
process_width = gr.Slider(minimum=64, maximum=2048, step=8, label="Width", value=512)
|
||||
process_height = gr.Slider(minimum=64, maximum=2048, step=8, label="Height", value=512)
|
||||
preprocess_txt_action = gr.Dropdown(label='Existing caption text action', value="ignore", choices=["ignore", "copy", "prepend", "append"])
|
||||
|
||||
with gr.Box():
|
||||
gr.HTML('<h2>Preprocessing steps</h2>')
|
||||
process_keep_original_size = gr.Checkbox(label='Keep original size')
|
||||
process_keep_channels = gr.Checkbox(label='Keep original image channels')
|
||||
process_flip = gr.Checkbox(label='Create flipped copies')
|
||||
process_split = gr.Checkbox(label='Split oversized images')
|
||||
process_focal_crop = gr.Checkbox(label='Auto focal point crop')
|
||||
process_multicrop = gr.Checkbox(label='Auto-sized crop')
|
||||
process_caption_only = gr.Checkbox(label='Create captions only')
|
||||
process_caption = gr.Checkbox(label='Create BLIP captions')
|
||||
process_caption_deepbooru = gr.Checkbox(label='Create Deepbooru captions')
|
||||
|
||||
with gr.Row(visible=False) as process_split_extra_row:
|
||||
process_split_threshold = gr.Slider(label='Split image threshold', value=0.5, minimum=0.0, maximum=1.0, step=0.05)
|
||||
process_overlap_ratio = gr.Slider(label='Split image overlap ratio', value=0.2, minimum=0.0, maximum=0.9, step=0.05)
|
||||
|
||||
with gr.Row(visible=False) as process_focal_crop_row:
|
||||
process_focal_crop_face_weight = gr.Slider(label='Focal point face weight', value=0.9, minimum=0.0, maximum=1.0, step=0.05)
|
||||
process_focal_crop_entropy_weight = gr.Slider(label='Focal point entropy weight', value=0.15, minimum=0.0, maximum=1.0, step=0.05)
|
||||
process_focal_crop_edges_weight = gr.Slider(label='Focal point edges weight', value=0.5, minimum=0.0, maximum=1.0, step=0.05)
|
||||
process_focal_crop_debug = gr.Checkbox(label='Create debug image')
|
||||
|
||||
with gr.Column(visible=False) as process_multicrop_col:
|
||||
gr.HTML('<h2>Each image is center-cropped with an automatically chosen width and height</h2>')
|
||||
with gr.Row():
|
||||
process_multicrop_mindim = gr.Slider(minimum=64, maximum=2048, step=8, label="Dimension lower bound", value=384)
|
||||
process_multicrop_maxdim = gr.Slider(minimum=64, maximum=2048, step=8, label="Dimension upper bound", value=768)
|
||||
with gr.Row():
|
||||
process_multicrop_minarea = gr.Slider(minimum=64*64, maximum=2048*2048, step=1, label="Area lower bound", value=64*64)
|
||||
process_multicrop_maxarea = gr.Slider(minimum=64*64, maximum=2048*2048, step=1, label="Area upper bound", value=640*640)
|
||||
with gr.Row():
|
||||
process_multicrop_objective = gr.Radio(["Maximize area", "Minimize error"], value="Maximize area", label="Resizing objective")
|
||||
process_multicrop_threshold = gr.Slider(minimum=0, maximum=1, step=0.01, label="Error threshold", value=0.1)
|
||||
|
||||
from modules.textual_inversion import ui
|
||||
process_split.change(fn=lambda show: gr_show(show), inputs=[process_split], outputs=[process_split_extra_row])
|
||||
process_focal_crop.change(fn=lambda show: gr_show(show), inputs=[process_focal_crop], outputs=[process_focal_crop_row])
|
||||
process_multicrop.change(fn=lambda show: gr_show(show), inputs=[process_multicrop], outputs=[process_multicrop_col])
|
||||
process_stop.click(fn=lambda: shared.state.interrupt(), inputs=[], outputs=[])
|
||||
process_run.click(
|
||||
fn=wrap_gradio_gpu_call(ui.preprocess, extra_outputs=[gr.update()]),
|
||||
_js="startTrainMonitor",
|
||||
inputs=[
|
||||
dummy_component,
|
||||
process_src,
|
||||
process_dst,
|
||||
process_width,
|
||||
process_height,
|
||||
preprocess_txt_action,
|
||||
process_keep_original_size,
|
||||
process_keep_channels,
|
||||
process_flip,
|
||||
process_split,
|
||||
process_caption_only,
|
||||
process_caption,
|
||||
process_caption_deepbooru,
|
||||
process_split_threshold,
|
||||
process_overlap_ratio,
|
||||
process_focal_crop,
|
||||
process_focal_crop_face_weight,
|
||||
process_focal_crop_entropy_weight,
|
||||
process_focal_crop_edges_weight,
|
||||
process_focal_crop_debug,
|
||||
process_multicrop,
|
||||
process_multicrop_mindim,
|
||||
process_multicrop_maxdim,
|
||||
process_multicrop_minarea,
|
||||
process_multicrop_maxarea,
|
||||
process_multicrop_objective,
|
||||
process_multicrop_threshold,
|
||||
],
|
||||
outputs=[
|
||||
train_output,
|
||||
train_outcome,
|
||||
],
|
||||
)
|
||||
|
||||
### train embedding tab
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
from modules import sd_hijack
|
||||
with gr.Tab(label="Train embedding", id="train_embedding_tab") as tab_ti:
|
||||
tab_ti.select(fn=lambda: train_tab_change('ti'), inputs=[], outputs=[action_pp, action_ti, action_hn])
|
||||
def get_textual_inversion_template_names():
|
||||
return sorted(textual_inversion.textual_inversion_templates)
|
||||
|
||||
gr.HTML('<h2>Select existing embedding to continue training or create a new one</h2>')
|
||||
with gr.Row():
|
||||
with gr.Column():
|
||||
with gr.Row():
|
||||
ti_name = gr.Dropdown(label='Select embedding', choices=sorted(sd_hijack.model_hijack.embedding_db.word_embeddings.keys()))
|
||||
create_refresh_button(ti_name, sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings, lambda: {"choices": sorted(sd_hijack.model_hijack.embedding_db.word_embeddings.keys())}, "refresh_train_embedding_name")
|
||||
with gr.Column():
|
||||
ti_new_name = gr.Textbox(label="Create emebedding")
|
||||
ti_init_text = gr.Textbox(label="Initialization text", value="*")
|
||||
ti_vectors = gr.Slider(label="Number of vectors per token", minimum=1, maximum=75, step=1, value=1)
|
||||
ti_overwrite = gr.Checkbox(value=False, label="Overwrite Old Embedding")
|
||||
with gr.Row():
|
||||
ti_create = gr.Button(value="Create embedding", variant='secondary')
|
||||
|
||||
with gr.Box():
|
||||
gr.HTML('<h2>Training parameters</h2>')
|
||||
ti_learn_rate = gr.Textbox(label='Embedding Learning rate', placeholder="Embedding Learning rate", value="0.005")
|
||||
with gr.Row():
|
||||
ti_clip_grad_mode = gr.Dropdown(value="disabled", label="Gradient Clipping", choices=["disabled", "value", "norm"])
|
||||
ti_clip_grad_value = gr.Number(label="Gradient clip value", value=0.1)
|
||||
ti_batch_size = gr.Number(label='Batch size', value=1, precision=0)
|
||||
ti_gradient_step = gr.Number(label='Gradient accumulation steps', value=1, precision=0)
|
||||
ti_steps = gr.Number(label='Max steps', value=1000, precision=0)
|
||||
|
||||
with gr.Box():
|
||||
gr.HTML('<h2>Training images</h2>')
|
||||
ti_dataset_directory = gr.Textbox(label='Dataset directory', placeholder="Path to directory with input images")
|
||||
with gr.Row():
|
||||
ti_varsize = gr.Checkbox(label="Do not resize images", value=False)
|
||||
ti_width = gr.Slider(minimum=64, maximum=2048, step=8, label="Width", value=512)
|
||||
ti_height = gr.Slider(minimum=64, maximum=2048, step=8, label="Height", value=512)
|
||||
ti_use_weight = gr.Checkbox(label="Use PNG alpha channel as loss weight", value=False)
|
||||
|
||||
with gr.Box():
|
||||
gr.HTML('<h2>Dataset processing</h2>')
|
||||
with gr.Row():
|
||||
ti_template = gr.Dropdown(label='Prompt template', value="style_filewords.txt", choices=get_textual_inversion_template_names())
|
||||
create_refresh_button(ti_template, textual_inversion.list_textual_inversion_templates, lambda: {"choices": get_textual_inversion_template_names()}, "refrsh_train_template_file")
|
||||
ti_shuffle = gr.Checkbox(label="Shuffle tags", value=False)
|
||||
ti_tag_drop_out = gr.Slider(minimum=0, maximum=1, step=0.1, label="Drop out tags when creating prompts", value=0)
|
||||
ti_latent_sampling_method = gr.Radio(label='Choose latent sampling method', value="once", choices=['once', 'deterministic', 'random'])
|
||||
|
||||
with gr.Box():
|
||||
gr.HTML('<h2>Training outputs</h2>')
|
||||
with gr.Row():
|
||||
ti_create_every = gr.Number(label='Create interim images', value=500, precision=0)
|
||||
ti_save_every = gr.Number(label='Create interim embeddings', value=500, precision=0)
|
||||
ti_save_image_with_stored_embedding = gr.Checkbox(label='Save images with embedding in PNG chunks', value=True)
|
||||
ti_preview_from_txt2img = gr.Checkbox(label='Use current settings for previews', value=False)
|
||||
ti_log_directory = gr.Textbox(label='Log directory', placeholder="Defaults to train/log/embedding", value="")
|
||||
|
||||
ti_stop.click(fn=lambda: shared.state.interrupt(), inputs=[], outputs=[])
|
||||
|
||||
ti_create.click(
|
||||
fn=modules.textual_inversion.ui.create_embedding,
|
||||
inputs=[
|
||||
ti_new_name,
|
||||
ti_init_text,
|
||||
ti_vectors,
|
||||
ti_overwrite,
|
||||
],
|
||||
outputs=[
|
||||
ti_name,
|
||||
train_output,
|
||||
train_outcome,
|
||||
]
|
||||
)
|
||||
|
||||
ti_train.click(
|
||||
fn=wrap_gradio_gpu_call(modules.textual_inversion.ui.train_embedding, extra_outputs=[gr.update()]),
|
||||
_js="startTrainMonitor",
|
||||
inputs=[
|
||||
dummy_component,
|
||||
ti_name,
|
||||
ti_learn_rate,
|
||||
ti_batch_size,
|
||||
ti_gradient_step,
|
||||
ti_dataset_directory,
|
||||
ti_log_directory,
|
||||
ti_width,
|
||||
ti_height,
|
||||
ti_varsize,
|
||||
ti_steps,
|
||||
ti_clip_grad_mode,
|
||||
ti_clip_grad_value,
|
||||
ti_shuffle,
|
||||
ti_tag_drop_out,
|
||||
ti_latent_sampling_method,
|
||||
ti_use_weight,
|
||||
ti_create_every,
|
||||
ti_save_every,
|
||||
ti_template,
|
||||
ti_save_image_with_stored_embedding,
|
||||
ti_preview_from_txt2img,
|
||||
*txt2img_preview_params,
|
||||
],
|
||||
outputs=[
|
||||
train_output,
|
||||
train_outcome,
|
||||
]
|
||||
)
|
||||
|
||||
### train hypernetwork tab
|
||||
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
from modules import sd_hijack
|
||||
with gr.Tab(label="Train hypernetwork", id="train_hypernetwork_tab") as tab_hn:
|
||||
tab_hn.select(fn=lambda: train_tab_change('hn'), inputs=[], outputs=[action_pp, action_ti, action_hn])
|
||||
gr.HTML('<h2>Select existing hypernetwork to continue training or create a new one</h2>')
|
||||
with gr.Row():
|
||||
with gr.Column():
|
||||
with gr.Row():
|
||||
hn_name = gr.Dropdown(label='Hypernetwork', choices=sorted(shared.hypernetworks))
|
||||
create_refresh_button(hn_name, shared.reload_hypernetworks, lambda: {"choices": sorted(shared.hypernetworks)}, "refresh_train_hypernetwork_name")
|
||||
with gr.Column():
|
||||
hn_new_name = gr.Textbox(label="Name")
|
||||
hn_new_sizes = gr.CheckboxGroup(label="Modules", value=["768", "320", "640", "1280"], choices=["768", "1024", "320", "640", "1280"])
|
||||
hn_new_layer_structure = gr.Textbox("1, 2, 1", label="Enter hypernetwork layer structure", placeholder="1st and last digit must be 1. ex:'1, 2, 1'")
|
||||
with gr.Row():
|
||||
hn_new_activation_func = gr.Dropdown(value="linear", label="Select activation function of hypernetwork", choices=modules.hypernetworks.ui.keys)
|
||||
hn_new_initialization_option = gr.Dropdown(value = "Normal", label="Select Layer weights initialization", choices=["Normal", "KaimingUniform", "KaimingNormal", "XavierUniform", "XavierNormal"])
|
||||
hn_new_add_layer_norm = gr.Checkbox(label="Add layer normalization")
|
||||
hn_new_use_dropout = gr.Checkbox(label="Use dropout")
|
||||
hn_new_dropout_structure = gr.Textbox("0, 0, 0", label="Enter hypernetwork Dropout structure", placeholder="1st and last digit must be 0 and values should be between 0 and 1. ex:'0, 0.01, 0'")
|
||||
hn_overwrite = gr.Checkbox(value=False, label="Overwrite Old Hypernetwork")
|
||||
with gr.Row():
|
||||
hn_create = gr.Button(value="Create hypernetwork", variant='secondary')
|
||||
|
||||
with gr.Box():
|
||||
gr.HTML('<h2>Training parameters</h2>')
|
||||
hn_learn_rate = gr.Textbox(label='Hypernetwork Learning rate', placeholder="Hypernetwork Learning rate", value="0.00001")
|
||||
with gr.Row():
|
||||
hn_clip_grad_mode = gr.Dropdown(value="disabled", label="Gradient Clipping", choices=["disabled", "value", "norm"])
|
||||
hn_clip_grad_value = gr.Number(label="Gradient clip value", value=0.1)
|
||||
hn_batch_size = gr.Number(label='Batch size', value=1, precision=0)
|
||||
hn_gradient_step = gr.Number(label='Gradient accumulation steps', value=1, precision=0)
|
||||
hn_steps = gr.Number(label='Max steps', value=1000, precision=0)
|
||||
|
||||
with gr.Box():
|
||||
gr.HTML('<h2>Training images</h2>')
|
||||
hn_dataset_directory = gr.Textbox(label='Dataset directory', placeholder="Path to directory with input images")
|
||||
with gr.Row():
|
||||
hn_varsize = gr.Checkbox(label="Do not resize images", value=False)
|
||||
hn_width = gr.Slider(minimum=64, maximum=2048, step=8, label="Width", value=512)
|
||||
hn_height = gr.Slider(minimum=64, maximum=2048, step=8, label="Height", value=512)
|
||||
hn_use_weight = gr.Checkbox(label="Use PNG alpha channel as loss weight", value=False)
|
||||
|
||||
with gr.Box():
|
||||
gr.HTML('<h2>Dataset processing</h2>')
|
||||
with gr.Row():
|
||||
hn_template = gr.Dropdown(label='Prompt template', value="style_filewords.txt", choices=get_textual_inversion_template_names())
|
||||
create_refresh_button(hn_template, textual_inversion.list_textual_inversion_templates, lambda: {"choices": get_textual_inversion_template_names()}, "refrsh_train_template_file")
|
||||
hn_shuffle_tags = gr.Checkbox(label="Shuffle tags by ',' when creating prompts.", value=False)
|
||||
hn_tag_drop_out = gr.Slider(minimum=0, maximum=1, step=0.1, label="Drop out tags when creating prompts", value=0)
|
||||
hn_latent_sampling_method = gr.Radio(label='Choose latent sampling method', value="once", choices=['once', 'deterministic', 'random'])
|
||||
|
||||
with gr.Box():
|
||||
gr.HTML('<h2>Training outputs</h2>')
|
||||
with gr.Row():
|
||||
hn_create_every = gr.Number(label='Create interim images', value=500, precision=0)
|
||||
hn_save_every = gr.Number(label='Create interim hypernetworks', value=500, precision=0)
|
||||
hn_preview_from_txt2img = gr.Checkbox(label='Use current settings for previews', value=False)
|
||||
hn_log_directory = gr.Textbox(label='Log directory', placeholder="Path to directory where to write outputs", value=f"{os.path.join('cmd_opts.data_dir', 'train/log/embeddings')}")
|
||||
|
||||
hn_stop.click(fn=lambda: shared.state.interrupt(), inputs=[], outputs=[])
|
||||
|
||||
hn_create.click(
|
||||
fn=modules.hypernetworks.ui.create_hypernetwork,
|
||||
inputs=[
|
||||
hn_new_name,
|
||||
hn_new_sizes,
|
||||
hn_overwrite,
|
||||
hn_new_layer_structure,
|
||||
hn_new_activation_func,
|
||||
hn_new_initialization_option,
|
||||
hn_new_add_layer_norm,
|
||||
hn_new_use_dropout,
|
||||
hn_new_dropout_structure
|
||||
],
|
||||
outputs=[
|
||||
hn_name,
|
||||
train_output,
|
||||
train_outcome,
|
||||
]
|
||||
)
|
||||
|
||||
hn_train.click(
|
||||
fn=wrap_gradio_gpu_call(modules.hypernetworks.ui.train_hypernetwork, extra_outputs=[gr.update()]),
|
||||
_js="startTrainMonitor",
|
||||
inputs=[
|
||||
dummy_component,
|
||||
hn_name,
|
||||
hn_learn_rate,
|
||||
hn_batch_size,
|
||||
hn_gradient_step,
|
||||
hn_dataset_directory,
|
||||
hn_log_directory,
|
||||
hn_width,
|
||||
hn_height,
|
||||
hn_varsize,
|
||||
hn_steps,
|
||||
hn_clip_grad_mode,
|
||||
hn_clip_grad_value,
|
||||
hn_shuffle_tags,
|
||||
hn_tag_drop_out,
|
||||
hn_latent_sampling_method,
|
||||
hn_use_weight,
|
||||
hn_create_every,
|
||||
hn_save_every,
|
||||
hn_template,
|
||||
hn_preview_from_txt2img,
|
||||
*txt2img_preview_params,
|
||||
],
|
||||
outputs=[
|
||||
train_output,
|
||||
train_outcome,
|
||||
]
|
||||
)
|
||||
|
||||
params = script_callbacks.UiTrainTabParams(txt2img_preview_params)
|
||||
script_callbacks.ui_train_tabs_callback(params)
|
||||
+2
-11
@@ -47,22 +47,12 @@ def create_ui():
|
||||
seed, reuse_seed, subseed, reuse_subseed, subseed_strength, seed_resize_from_h, seed_resize_from_w = ui_sections.create_seed_inputs('txt2img')
|
||||
cfg_scale, clip_skip, image_cfg_scale, diffusers_guidance_rescale, sag_scale, cfg_end, full_quality, restore_faces, tiling = ui_sections.create_advanced_inputs('txt2img')
|
||||
hdr_mode, hdr_brightness, hdr_color, hdr_sharpen, hdr_clamp, hdr_boundary, hdr_threshold, hdr_maximize, hdr_max_center, hdr_max_boundry, hdr_color_picker, hdr_tint_ratio, = ui_sections.create_correction_inputs('txt2img')
|
||||
enable_hr, hr_sampler_index, denoising_strength, hr_final_resolution, hr_upscaler, hr_force, hr_second_pass_steps, hr_scale, hr_resize_x, hr_resize_y, refiner_steps, refiner_start, refiner_prompt, refiner_negative = ui_sections.create_hires_inputs('txt2img')
|
||||
enable_hr, hr_sampler_index, denoising_strength, hr_upscaler, hr_force, hr_second_pass_steps, hr_scale, hr_resize_x, hr_resize_y, refiner_steps, refiner_start, refiner_prompt, refiner_negative = ui_sections.create_hires_inputs('txt2img')
|
||||
override_settings = ui_common.create_override_inputs('txt2img')
|
||||
|
||||
with gr.Group(elem_id="txt2img_script_container"):
|
||||
txt2img_script_inputs = modules.scripts.scripts_txt2img.setup_ui(parent='txt2img', accordion=True)
|
||||
|
||||
hr_resolution_preview_inputs = [width, height, hr_scale, hr_resize_x, hr_resize_y, hr_upscaler]
|
||||
for preview_input in hr_resolution_preview_inputs:
|
||||
preview_input.change(
|
||||
fn=calc_resolution_hires,
|
||||
_js="onCalcResolutionHires",
|
||||
inputs=hr_resolution_preview_inputs,
|
||||
outputs=[hr_final_resolution],
|
||||
show_progress=False,
|
||||
)
|
||||
|
||||
txt2img_gallery, txt2img_generation_info, txt2img_html_info, _txt2img_html_info_formatted, txt2img_html_log = ui_common.create_output_panel("txt2img", preview=True, prompt=None)
|
||||
ui_common.connect_reuse_seed(seed, reuse_seed, txt2img_generation_info, is_subseed=False)
|
||||
ui_common.connect_reuse_seed(subseed, reuse_subseed, txt2img_generation_info, is_subseed=True)
|
||||
@@ -97,6 +87,7 @@ def create_ui():
|
||||
show_progress=False,
|
||||
)
|
||||
txt2img_prompt.submit(**txt2img_dict)
|
||||
txt2img_negative_prompt.submit(**txt2img_dict)
|
||||
txt2img_submit.click(**txt2img_dict)
|
||||
txt2img_paste_fields = [
|
||||
# prompt
|
||||
|
||||
+126
@@ -0,0 +1,126 @@
|
||||
import torch
|
||||
import transformers
|
||||
from PIL import Image
|
||||
from modules import shared, devices
|
||||
|
||||
|
||||
processor = None
|
||||
model = None
|
||||
loaded: str = None
|
||||
MODELS = {
|
||||
"None": None,
|
||||
"GIT TextCaps Base": "microsoft/git-base-textcaps", # 0.7GB
|
||||
"GIT VQA Base": "microsoft/git-base-vqav2", # 0.7GB
|
||||
"GIT VQA Large": "microsoft/git-large-vqav2", # 1.6GB
|
||||
"BLIP Base": "Salesforce/blip-vqa-base", # 1.5GB
|
||||
"BLIP Large": "Salesforce/blip-vqa-capfilt-large", # 1.5GB
|
||||
"ViLT Base": "dandelin/vilt-b32-finetuned-vqa", # 0.5GB
|
||||
"Pix Textcaps": "google/pix2struct-textcaps-base", # 1.1GB
|
||||
}
|
||||
|
||||
|
||||
def git(question: str, image: Image.Image, repo: str = None):
|
||||
global processor, model, loaded # pylint: disable=global-statement
|
||||
if model is None or loaded != repo:
|
||||
model = transformers.GitForCausalLM.from_pretrained(repo)
|
||||
processor = transformers.GitProcessor.from_pretrained(repo)
|
||||
loaded = repo
|
||||
model.to(devices.device, devices.dtype)
|
||||
shared.log.debug(f'VQA: class={model.__class__.__name__} processor={processor.__class__} model={repo}')
|
||||
|
||||
pixel_values = processor(images=image, return_tensors="pt").pixel_values
|
||||
git_dict = {}
|
||||
git_dict['pixel_values'] = pixel_values.to(devices.device, devices.dtype)
|
||||
if len(question) > 0:
|
||||
input_ids = processor(text=question, add_special_tokens=False).input_ids
|
||||
input_ids = [processor.tokenizer.cls_token_id] + input_ids
|
||||
input_ids = torch.tensor(input_ids).unsqueeze(0)
|
||||
git_dict['input_ids'] = input_ids.to(devices.device)
|
||||
with devices.inference_context():
|
||||
generated_ids = model.generate(**git_dict)
|
||||
response = processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
|
||||
|
||||
model.to(devices.cpu)
|
||||
shared.log.debug(f'VQA: response={response}')
|
||||
return response
|
||||
|
||||
|
||||
def blip(question: str, image: Image.Image, repo: str = None):
|
||||
global processor, model, loaded # pylint: disable=global-statement
|
||||
if model is None or loaded != repo:
|
||||
model = transformers.BlipForQuestionAnswering.from_pretrained(repo)
|
||||
processor = transformers.BlipProcessor.from_pretrained(repo)
|
||||
loaded = repo
|
||||
model.to(devices.device, devices.dtype)
|
||||
inputs = processor(image, question, return_tensors="pt")
|
||||
inputs = inputs.to(devices.device, devices.dtype)
|
||||
with devices.inference_context():
|
||||
outputs = model.generate(**inputs)
|
||||
response = processor.decode(outputs[0], skip_special_tokens=True)
|
||||
|
||||
model.to(devices.cpu)
|
||||
shared.log.debug(f'VQA: response={response}')
|
||||
return response
|
||||
|
||||
|
||||
def vilt(question: str, image: Image.Image, repo: str = None):
|
||||
global processor, model, loaded # pylint: disable=global-statement
|
||||
if model is None or loaded != repo:
|
||||
model = transformers.ViltForQuestionAnswering.from_pretrained(repo)
|
||||
processor = transformers.ViltProcessor.from_pretrained(repo)
|
||||
loaded = repo
|
||||
model.to(devices.device)
|
||||
shared.log.debug(f'VQA: class={model.__class__.__name__} processor={processor.__class__} model={repo}')
|
||||
|
||||
inputs = processor(image, question, return_tensors="pt")
|
||||
inputs = inputs.to(devices.device)
|
||||
with devices.inference_context():
|
||||
outputs = model(**inputs)
|
||||
logits = outputs.logits
|
||||
idx = logits.argmax(-1).item()
|
||||
response = model.config.id2label[idx]
|
||||
|
||||
model.to(devices.cpu)
|
||||
shared.log.debug(f'VQA: response={response}')
|
||||
return response
|
||||
|
||||
|
||||
def pix(question: str, image: Image.Image, repo: str = None):
|
||||
global processor, model, loaded # pylint: disable=global-statement
|
||||
if model is None or loaded != repo:
|
||||
model = transformers.Pix2StructForConditionalGeneration.from_pretrained(repo)
|
||||
processor = transformers.Pix2StructProcessor.from_pretrained(repo)
|
||||
loaded = repo
|
||||
model.to(devices.device)
|
||||
shared.log.debug(f'VQA: class={model.__class__.__name__} processor={processor.__class__} model={repo}')
|
||||
|
||||
if len(question) > 0:
|
||||
inputs = processor(images=image, text=question, return_tensors="pt").to(devices.device)
|
||||
else:
|
||||
inputs = processor(images=image, return_tensors="pt").to(devices.device)
|
||||
with devices.inference_context():
|
||||
outputs = model.generate(**inputs)
|
||||
response = processor.decode(outputs[0], skip_special_tokens=True)
|
||||
|
||||
model.to(devices.cpu)
|
||||
shared.log.debug(f'VQA: response={response}')
|
||||
return response
|
||||
|
||||
|
||||
def interrogate(vqa_question, vqa_image, vqa_model):
|
||||
vqa_model = MODELS.get(vqa_model, None)
|
||||
shared.log.debug(f'VQA: model="{vqa_model}" question={vqa_question} image={vqa_image}')
|
||||
if vqa_image is None:
|
||||
return 'no image provided'
|
||||
if vqa_model is None:
|
||||
return 'no model selected'
|
||||
if 'git' in vqa_model.lower():
|
||||
return git(vqa_question, vqa_image, vqa_model)
|
||||
if 'vilt' in vqa_model.lower():
|
||||
return vilt(vqa_question, vqa_image, vqa_model)
|
||||
if 'blip' in vqa_model.lower():
|
||||
return blip(vqa_question, vqa_image, vqa_model)
|
||||
if 'pix' in vqa_model.lower():
|
||||
return pix(vqa_question, vqa_image, vqa_model)
|
||||
else:
|
||||
return 'unknown model'
|
||||
@@ -16,6 +16,7 @@ exclude = [
|
||||
"modules/control/units/*_pipe.py",
|
||||
"modules/pipelines/*.py",
|
||||
"modules/xadapter/*.py",
|
||||
"modules/tcd/*.py",
|
||||
]
|
||||
[tool.ruff.lint]
|
||||
select = [
|
||||
|
||||
+1
-1
@@ -55,7 +55,7 @@ pandas
|
||||
protobuf==3.20.3
|
||||
pytorch_lightning==1.9.4
|
||||
tokenizers==0.15.2
|
||||
transformers==4.37.2
|
||||
transformers==4.38.1
|
||||
tomesd==0.1.3
|
||||
urllib3==1.26.18
|
||||
Pillow==10.2.0
|
||||
|
||||
@@ -1518,7 +1518,7 @@ class StableDiffusionDiffImg2ImgPipeline(DiffusionPipeline):
|
||||
negative_prompt=None,
|
||||
prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.FloatTensor] = None,
|
||||
lora_scale: Optional[float] = None,
|
||||
lora_scale: Optional[float] = None, # pylint: disable=unused-argument
|
||||
clip_skip: Optional[int] = None,
|
||||
):
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
@@ -1892,11 +1892,11 @@ class Script(scripts.Script):
|
||||
image = gr.Image(label="Image map", show_label=False, type="pil", source="upload", interactive=True, tool="editor", visible=True, image_mode='RGB')
|
||||
return enabled, strength, invert, model, image
|
||||
|
||||
def depthmap(self, image_init: Image.Image, image_map: Image.Image, model: str, strength: float, invert: bool, output_type="tensor"):
|
||||
def depthmap(self, image_init: Image.Image, image_map: Image.Image, model: str, strength: float, invert: bool):
|
||||
global detector # pylint: disable=global-statement
|
||||
from modules.control.proc.dpt import DPTDetector
|
||||
if image_init is None:
|
||||
return None, None
|
||||
return None, None, None
|
||||
image_map = None
|
||||
if image_map is not None:
|
||||
image_map = image_map.resize(image_init.size, Image.Resampling.LANCZOS)
|
||||
@@ -1916,14 +1916,14 @@ class Script(scripts.Script):
|
||||
init_img_hash = hashlib.sha256(image_map.tobytes()).hexdigest()[0:8] # pylint: disable=attribute-defined-outside-init
|
||||
images.save_image(image_map, path=shared.opts.outdir_init_images, basename=None, forced_filename=init_img_hash, suffix="-init-image")
|
||||
else:
|
||||
return None, None
|
||||
if output_type == "tensor":
|
||||
image_map = transforms.ToTensor()(image_map)
|
||||
image_map = image_map.to(devices.device)
|
||||
image_init = 2 * transforms.ToTensor()(image_init) - 1
|
||||
image_init = image_init.unsqueeze(0)
|
||||
image_init = image_init.to(devices.device)
|
||||
return image_init, image_map
|
||||
return None, None, None
|
||||
image_mask = image_map.copy()
|
||||
image_map = transforms.ToTensor()(image_map)
|
||||
image_map = image_map.to(devices.device)
|
||||
image_init = 2 * transforms.ToTensor()(image_init) - 1
|
||||
image_init = image_init.unsqueeze(0)
|
||||
image_init = image_init.to(devices.device)
|
||||
return image_init, image_map, image_mask
|
||||
|
||||
def run(self, p: processing.StableDiffusionProcessingImg2Img, enabled, strength, invert, model, image): # pylint: disable=arguments-differ
|
||||
if not enabled:
|
||||
@@ -1935,7 +1935,7 @@ class Script(scripts.Script):
|
||||
shared.log.error('Differential-diffusion: no input images')
|
||||
return
|
||||
|
||||
image_init, image_map = self.depthmap(p.init_images[0], image, model, strength, invert, output_type="tensor")
|
||||
image_init, image_map, image_mask = self.depthmap(p.init_images[0], image, model, strength, invert)
|
||||
if image_map is None:
|
||||
shared.log.error('Differential-diffusion: no image map')
|
||||
return
|
||||
@@ -1974,6 +1974,7 @@ class Script(scripts.Script):
|
||||
p.task_args['original_image'] = image_init
|
||||
shared.log.debug(f'Differential-diffusion: pipeline={pipe.__class__.__name__} strength={strength} model={model} auto={image is None}')
|
||||
shared.sd_model = pipe
|
||||
sd_models.move_model(pipe.vae, devices.device, force=True)
|
||||
except Exception as e:
|
||||
shared.log.error(f'Differential-diffusion: pipeline creation failed: {e}')
|
||||
errors.display(e, 'Differential-diffusion: pipeline creation failed')
|
||||
@@ -1981,6 +1982,9 @@ class Script(scripts.Script):
|
||||
|
||||
# run pipeline
|
||||
processed: processing.Processed = processing.process_images(p) # runs processing using main loop
|
||||
if shared.opts.include_mask:
|
||||
if image_mask is not None and isinstance(image_mask, Image.Image):
|
||||
processed.images.append(image_mask)
|
||||
|
||||
# restore pipeline and params
|
||||
pipe = None
|
||||
|
||||
@@ -6,7 +6,7 @@ from modules import scripts, processing, shared, images, sd_models, devices
|
||||
|
||||
MODELS = [
|
||||
{ 'name': 'None', 'info': '' },
|
||||
{ 'name': 'PIA', 'url': 'openmmlab/PIA-condition-adapter', 'info': '<a href="https://huggingface.co/docs/diffusers/main/en/api/pipelines/pia" target="_blank">Open MMLab Personalized Image Animator</a>' },
|
||||
# { 'name': 'PIA', 'url': 'openmmlab/PIA-condition-adapter', 'info': '<a href="https://huggingface.co/docs/diffusers/main/en/api/pipelines/pia" target="_blank">Open MMLab Personalized Image Animator</a>' },
|
||||
{ 'name': 'VGen', 'url': 'ali-vilab/i2vgen-xl', 'info': '<a href="https://huggingface.co/ali-vilab/i2vgen-xl" target="_blank">Alibaba VGen</a>' },
|
||||
]
|
||||
|
||||
@@ -16,8 +16,8 @@ class Script(scripts.Script):
|
||||
return 'Image-to-Video'
|
||||
|
||||
def show(self, is_img2img):
|
||||
# return is_img2img if shared.backend == shared.Backend.DIFFUSERS else False
|
||||
return False
|
||||
return is_img2img if shared.backend == shared.Backend.DIFFUSERS else False
|
||||
# return False
|
||||
|
||||
# return signature is array of gradio components
|
||||
def ui(self, _is_img2img):
|
||||
@@ -75,17 +75,17 @@ class Script(scripts.Script):
|
||||
shared.log.debug(f'Image2Video: model={model_name} frames={num_frames}, video={video_type} duration={duration} loop={gif_loop} pad={mp4_pad} interpolate={mp4_interpolate}')
|
||||
p.ops.append('image2video')
|
||||
p.do_not_save_grid = True
|
||||
orig_pipeline = shared.sd_model
|
||||
|
||||
if model_name == 'PIA':
|
||||
if shared.sd_model_type != 'sd':
|
||||
shared.log.error('Image2Video PIA: base model must be SD15')
|
||||
return
|
||||
orig_pipeline = shared.sd_model
|
||||
shared.log.info(f'Image2Video PIA load: model={repo_id}')
|
||||
motion_adapter = diffusers.MotionAdapter.from_pretrained(repo_id)
|
||||
motion_adapter.to(devices.device, devices.dtype)
|
||||
sd_models.move_model(motion_adapter, devices.device)
|
||||
shared.sd_model = sd_models.switch_pipe(diffusers.PIAPipeline, shared.sd_model, { 'motion_adapter': motion_adapter })
|
||||
sd_models.move_model(shared.sd_model, devices.device) # move pipeline to device
|
||||
sd_models.move_model(shared.sd_model, devices.device, force=True) # move pipeline to device
|
||||
if num_frames > 0:
|
||||
p.task_args['num_frames'] = num_frames
|
||||
p.task_args['image'] = p.init_images[0]
|
||||
@@ -101,7 +101,6 @@ class Script(scripts.Script):
|
||||
shared.log.debug(f'Image2Video PIA: args={p.task_args}')
|
||||
processed = processing.process_images(p)
|
||||
shared.sd_model.motion_adapter = None
|
||||
shared.sd_model = orig_pipeline
|
||||
|
||||
if model_name == 'VGen':
|
||||
if not isinstance(shared.sd_model, diffusers.I2VGenXLPipeline):
|
||||
@@ -121,6 +120,7 @@ class Script(scripts.Script):
|
||||
shared.log.debug(f'Image2Video VGen: args={p.task_args}')
|
||||
processed = processing.process_images(p)
|
||||
|
||||
shared.sd_model = orig_pipeline
|
||||
if video_type != 'None' and processed is not None:
|
||||
images.save_video(p, filename=None, images=processed.images, video_type=video_type, duration=duration, loop=gif_loop, pad=mp4_pad, interpolate=mp4_interpolate)
|
||||
return processed
|
||||
|
||||
@@ -16,7 +16,7 @@ class ScriptPostprocessingUpscale(scripts_postprocessing.ScriptPostprocessing):
|
||||
with gr.Row(elem_id="extras_upscale"):
|
||||
with gr.Tabs(elem_id="extras_resize_mode"):
|
||||
with gr.TabItem('Scale by', elem_id="extras_scale_by_tab") as tab_scale_by:
|
||||
upscaling_resize = gr.Slider(minimum=1.0, maximum=8.0, step=0.05, label="Resize", value=2.0, elem_id="extras_upscaling_resize")
|
||||
upscaling_resize = gr.Slider(minimum=0.1, maximum=8.0, step=0.05, label="Resize", value=2.0, elem_id="extras_upscaling_resize")
|
||||
|
||||
with gr.TabItem('Scale to', elem_id="extras_scale_to_tab") as tab_scale_to:
|
||||
with gr.Row():
|
||||
|
||||
@@ -58,7 +58,7 @@ class Script(scripts.Script):
|
||||
adapter.load_state_dict(adapter_dict)
|
||||
try:
|
||||
if adapter is not None:
|
||||
adapter.to(devices.device)
|
||||
sd_models.move_model(adapter, devices.device)
|
||||
except Exception:
|
||||
pass
|
||||
if adapter is None:
|
||||
|
||||
+4
-1
@@ -246,7 +246,7 @@ axis_options = [
|
||||
AxisOption("[Sampler] Sigma tmax", float, apply_field("s_tmax")),
|
||||
AxisOption("[Sampler] Sigma Churn", float, apply_field("s_churn")),
|
||||
AxisOption("[Sampler] Sigma noise", float, apply_field("s_noise")),
|
||||
AxisOption("[Sampler] ETA", float, apply_field("eta")),
|
||||
AxisOption("[Sampler] ETA", float, apply_setting("scheduler_eta")),
|
||||
AxisOption("[Sampler] Solver order", int, apply_setting("schedulers_solver_order")),
|
||||
AxisOption("[Second pass] Upscaler", str, apply_field("hr_upscaler"), choices=lambda: [*shared.latent_upscale_modes, *[x.name for x in shared.sd_upscalers]]),
|
||||
AxisOption("[Second pass] Sampler", str, apply_hr_sampler_name, fmt=format_value, confirm=confirm_samplers, choices=lambda: [x.name for x in sd_samplers.samplers]),
|
||||
@@ -298,6 +298,9 @@ def draw_xyz_grid(p, xs, ys, zs, x_labels, y_labels, z_labels, cell, draw_legend
|
||||
processed: processing.Processed = cell(x, y, z, ix, iy, iz)
|
||||
if processed_result is None:
|
||||
processed_result = copy(processed)
|
||||
if processed_result is None:
|
||||
shared.log.error('XYZ grid: no processing results')
|
||||
return processing.Processed(p, [])
|
||||
processed_result.images = [None] * list_size
|
||||
processed_result.all_prompts = [None] * list_size
|
||||
processed_result.all_seeds = [None] * list_size
|
||||
|
||||
Reference in New Issue
Block a user