Download app.py from HyperHail/C: direct link, hf CLI and curl.
- Browser
- Download file 155 kB
-
https://huggingface.co/spaces/HyperHail/C/resolve/main/app.py
- Command line
-
hf download hf://spaces/HyperHail/C/app.py
-
curl -L -o app.py https://huggingface.co/spaces/HyperHail/C/resolve/main/app.py
155 kB
| import base64 | |
| import binascii | |
| import importlib.util | |
| import json | |
| import os | |
| import re | |
| import runpy | |
| import secrets | |
| import shutil | |
| import sqlite3 | |
| import subprocess | |
| import sys | |
| import tempfile | |
| import threading | |
| import time | |
| import urllib.request | |
| import uuid | |
| import warnings | |
| from asyncio.base_events import BaseEventLoop | |
| from concurrent.futures import ThreadPoolExecutor | |
| from datetime import datetime | |
| from functools import lru_cache | |
| from io import BytesIO | |
| from itertools import islice | |
| from pathlib import Path | |
| from urllib.parse import quote, unquote, urlparse | |
| from zoneinfo import ZoneInfo | |
| ASYNCIO_FD_ERROR = "Invalid file descriptor: -1" | |
| loop_del = BaseEventLoop.__del__ | |
| def close_loop(loop): | |
| try: | |
| loop_del(loop) | |
| except ValueError as error: | |
| if str(error) != ASYNCIO_FD_ERROR: | |
| raise | |
| BaseEventLoop.__del__ = close_loop | |
| import gradio as gr | |
| import numpy as np | |
| import py7zr | |
| import spaces | |
| import torch | |
| from cryptography.exceptions import InvalidTag | |
| from cryptography.hazmat.primitives.ciphers.aead import AESGCM | |
| from cryptography.hazmat.primitives.kdf.scrypt import Scrypt | |
| from fastapi import Body, Depends, HTTPException, Query | |
| from fastapi.middleware.cors import CORSMiddleware | |
| from fastapi.responses import FileResponse, Response | |
| from gradio_client import Client | |
| from gradio.routes import App | |
| from huggingface_hub import ( | |
| batch_bucket_files, | |
| download_bucket_files, | |
| list_bucket_tree, | |
| ) | |
| from PIL import Image, ImageDraw, ImageFont, ImageOps | |
| from PIL.PngImagePlugin import PngInfo | |
| from pydantic import BaseModel, ConfigDict, Field | |
| from starlette.background import BackgroundTask | |
| from workflow_api import ( | |
| ALIGN_MODEL_TYPE, | |
| ALIGN_SCHEDULER, | |
| ANIMA_CLIP, | |
| DETAILER_CROP, | |
| DETAILER_DILATION, | |
| DETAILER_DROP_SIZE, | |
| DETAILER_FEATHER, | |
| DETAILER_GUIDE_SIZE, | |
| DETAILER_MAX_SIZE, | |
| DETAILER_THRESHOLD, | |
| GRID_SIZE, | |
| LATENT_SCALE, | |
| image_metadata, | |
| import_custom_nodes, | |
| is_anima_model, | |
| mask_box, | |
| upscale_size, | |
| ) | |
| os.environ.setdefault("YOLO_CONFIG_DIR", "/tmp/Ultralytics") | |
| DATA_MOUNT = Path("/data") | |
| BUCKET_MOUNT = Path("/CB") | |
| HAS_DATA_MOUNT = (DATA_MOUNT / "img").is_dir() | |
| DEFAULT_DATA_DIR = DATA_MOUNT if HAS_DATA_MOUNT else Path.cwd() / "data" | |
| DATA_DIR = Path(os.environ.get("DATA_DIR", DEFAULT_DATA_DIR)) | |
| LOCAL_MODEL_DIR = Path(os.environ.get("LOCAL_MODEL_DIR", "/tmp/models")) | |
| COMFYUI_PATH = Path(os.environ.get("COMFYUI_PATH", Path.cwd() / "ComfyUI")) | |
| MOUNTED_CUSTOM_NODES_DIR = BUCKET_MOUNT / "custom_nodes" | |
| CUSTOM_NODES_DIR = Path("/tmp/custom_nodes") | |
| OUTPUT_DIR = DATA_DIR / "output" | |
| IMAGE_DIR = DATA_DIR / "img" | |
| STAR_DB = DATA_DIR / "explorer.db" | |
| ARTIST_DB = Path(tempfile.gettempdir()) / "artists.sqlite" | |
| ARTIST_PLACEHOLDER = re.compile(r"\{(?:artist|ar)(\d*)\}", re.I) | |
| TIMEZONE = ZoneInfo("Asia/Singapore") | |
| CPU_DEVICE = torch.device("cpu") | |
| MIB = 1024 * 1024 | |
| SALT_SIZE = 16 | |
| NONCE_SIZE = 12 | |
| SCRYPT_N = 2**14 | |
| FILE_MAGIC = b"EPNG1" | |
| PROXY_ENCRYPTION = b"aes-256-gcm" | |
| PROXY_ENCRYPTION_HEADER = b"x-gradio-comfy-encryption" | |
| PROXY_MAGIC = b"GCV1" | |
| IMAGE_SUFFIXES = (".epng",) | |
| PREVIEW_CACHE_SIZE = 256 | |
| PREVIEW_QUALITY = 70 | |
| PREVIEW_SIZE = 320 | |
| EXPLORER_PAGE_SIZE = 80 | |
| EXPLORER_MAX_PAGE_SIZE = 200 | |
| EXPLORER_DB_TIMEOUT = 30 | |
| IMAGE_KEY_CACHE_SIZE = 512 | |
| DUPLICATE_HASH_SIZE = 16 | |
| DUPLICATE_HASH_DISTANCE = 24 | |
| DUPLICATE_COLOR_DISTANCE = 24 | |
| DUPLICATE_CHECK_EXCLUDED_FOLDERS = {"2026-07-30"} | |
| DEFAULT_RETURN_SCALE = 1 | |
| DEFAULT_BATCH_SIZE = 1 | |
| DEFAULT_UI_BATCH_SIZE = 1 | |
| MAX_BATCH_SIZE = 8 | |
| PING_MODEL_ID = 50 | |
| PING_SIZE = 64 | |
| PING_STEPS = 8 | |
| PING_SAMPLER = "euler" | |
| PING_SCHEDULER = "simple" | |
| STARTUP_ASSET_IDS = { | |
| "checkpoints": (16, 50), | |
| "diffusion_models": (4,6), | |
| "loras": (1, 3, 4, 7, 12, 30, 47, 49, 50), | |
| "ultralytics": (1,), | |
| "upscale_models": (1, 3), | |
| "vae": (2, 3), | |
| "ipadapter": (2,), | |
| "clip_vision": (1,), | |
| } | |
| ENVIRONMENT_START = .2 | |
| REGIONAL_GLOBAL_STRENGTH = .6 | |
| REQUIRED_CUSTOM_NODES = ( | |
| "ComfyUI-Impact-Pack", | |
| "ComfyUI-Impact-Subpack", | |
| "ComfyUI-ppm", | |
| "ComfyUI_IPAdapter_plus", | |
| "RES4LYF", | |
| ) | |
| CUSTOM_NODE_REPOS = { | |
| "RES4LYF": "https://github.com/ClownsharkBatwing/RES4LYF", | |
| } | |
| CUSTOM_NODE_MODULES = { | |
| "ComfyUI-Impact-Pack": ( | |
| "segment_anything", "skimage", "piexif", "transformers", "cv2", | |
| "scipy", "dill", "matplotlib", "sam2", | |
| ), | |
| "ComfyUI-Impact-Subpack": ( | |
| "ultralytics", "numpy", "cv2", "dill", "matplotlib", | |
| ), | |
| } | |
| AREA_PRESETS = { | |
| "full": "a1:e5", | |
| "tl": "a1:c3", | |
| "tc": "b1:d3", | |
| "tr": "c1:e3", | |
| "ml": "a2:c4", | |
| "mc": "b2:d4", | |
| "mr": "c2:e4", | |
| "bl": "a3:c5", | |
| "bc": "b3:d5", | |
| "br": "c3:e5", | |
| "th": "a1:e3", | |
| "mh": "a2:e4", | |
| "bh": "a3:e5", | |
| "lh": "a1:c5", | |
| "ch": "b1:d5", | |
| "rh": "c1:e5", | |
| } | |
| AUTO_LAYOUTS = { | |
| 1: ((.2, 0, .6, 1),), | |
| 2: ((0, 0, .55, 1), (.45, 0, .55, 1)), | |
| 3: ((0, 0, .4, 1), (.3, 0, .4, 1), (.6, 0, .4, 1)), | |
| } | |
| REGIONAL_MODES = ("conditioning", "attention") | |
| PORT = int(os.environ.get("PORT", "7860")) | |
| LOCAL_URL = os.environ.get("LOCAL_URL", f"http://127.0.0.1:{PORT}") | |
| PASSWORD = os.environ.get("pass") | |
| if not PASSWORD: | |
| raise RuntimeError("pass environment variable is required") | |
| BUCKET_ID = "HyperHail/CB" | |
| IMAGE_BUCKET_ID = "HyperHail/C" | |
| IMAGE_BUCKET_PREFIX = "img" | |
| IMAGE_TOKEN = os.environ.get("hh") or False | |
| CIVITAI_TOKEN = os.environ.get("CIVIT_MODEL_READ") | |
| CIVITAI_HOSTS = { | |
| "civitai.com", "www.civitai.com", "civitai.red", "www.civitai.red", | |
| } | |
| MODEL_KINDS = ( | |
| "checkpoints", "diffusion_models", "clip", "clip_vision", "vae", | |
| "loras", "ipadapter", "upscale_models", "ultralytics", | |
| ) | |
| COMFY_KINDS = {"clip": "text_encoders"} | |
| MODEL_SUFFIXES = ( | |
| ".bin", ".ckpt", ".pkl", ".pt", ".pt2", ".pth", ".safetensors", ".sft", | |
| ) | |
| NUMBERED_MODEL_KINDS = ("checkpoints", "loras") | |
| MODEL_NUMBER = re.compile(r"^(\d+)_(.+)$") | |
| ANIMA_PREFIX = "anima" | |
| MODEL_LOCATION_CHOICES = [ | |
| (kind.replace("_", " ").title(), kind) | |
| for kind in MODEL_KINDS | |
| ] | |
| DOWNLOAD_HEADERS = { | |
| "User-Agent": ( | |
| "Mozilla/5.0 (Windows NT 10.0; Win64; x64; rv:140.0) " | |
| "Gecko/20100101 Firefox/140.0" | |
| ), | |
| "Accept": "text/html,application/xhtml+xml,application/xml;q=0.9,*/*;q=0.8", | |
| "Accept-Language": "en-US,en;q=0.5", | |
| "Referer": "https://civitai.com/", | |
| } | |
| DEFAULT_NEGATIVE = ( | |
| "(censored, mosaic censoring, bar censor:1.1), bad quality, worst quality, " | |
| "worst detail, bad anatomy, extra fingers, extra toes, extra legs, 4 toes, " | |
| "6 toes, 4 fingers, 6 fingers, malformed fingers, extra limbs, missing fingers, " | |
| "extra arms, censored, deformed, disfigured, text, (multiple views:1.1)" | |
| ) | |
| DEFAULT_SAMPLER = "euler_ancestral" | |
| DEFAULT_SCHEDULER = "karras" | |
| DEFAULT_CFG = 4 | |
| DEFAULT_STEPS = 30 | |
| DEFAULT_MODEL = "52_novaAnimeXL_ilV190.safetensors" | |
| DEFAULT_VAE = "3_sdxlVAE_sdxlVAE.safetensors" | |
| ANIMA_VAE = "2_qwen_image_vae.safetensors" | |
| NON_ANIMA_HEADER = "__non_anima_header__" | |
| ANIMA_HEADER = "__anima_header__" | |
| MODEL_HEADERS = (NON_ANIMA_HEADER, ANIMA_HEADER) | |
| DEFAULT_LORAS = () | |
| DEFAULT_UPSCALE_METHOD = "bislerp" | |
| DEFAULT_UPSCALE_MODEL = "3_1x-Archivist_Soft.pth" | |
| DEFAULT_UPSCALE_SCALE = 1.1 | |
| DEFAULT_SECOND_SAMPLER = "euler" | |
| DEFAULT_SECOND_SCHEDULER = "karras" | |
| DEFAULT_SECOND_STEPS = 18 | |
| DEFAULT_SECOND_CFG = 5 | |
| DEFAULT_DENOISE = .5 | |
| DEFAULT_DETECTOR = "2_face_yolov9c.pt" | |
| STYLE_IPADAPTER = "2_ip-adapter-plus_sdxl_vit-h.safetensors" | |
| STYLE_CLIP_VISION = "1_CLIP-ViT-H-fp16.safetensors" | |
| STYLE_WEIGHT_TYPE = "style transfer" | |
| STYLE_EMBEDS_SCALING = "V only" | |
| STYLE_SCOPES = ("first", "generation", "all") | |
| STYLE_MODEL_KINDS = ("ipadapter", "clip_vision") | |
| DEFAULT_STYLE_SCOPE = "generation" | |
| DEFAULT_STYLE_WEIGHT = 1 | |
| DEFAULT_STYLE_END = 1 | |
| MAX_STYLE_IMAGES = 4 | |
| STYLE_IMAGE_SIZE = 1024 | |
| MAX_STYLE_IMAGE_SIZE = 20 * MIB | |
| BUILTIN_ASSETS = { | |
| ("ipadapter", STYLE_IPADAPTER): ( | |
| "https://huggingface.co/h94/IP-Adapter/resolve/main/sdxl_models/" | |
| "ip-adapter-plus_sdxl_vit-h.safetensors" | |
| ), | |
| ("clip_vision", STYLE_CLIP_VISION): ( | |
| "https://huggingface.co/h94/IP-Adapter/resolve/main/models/" | |
| "image_encoder/model.safetensors" | |
| ), | |
| } | |
| MATRIX_WIDTH = 1152 | |
| MATRIX_HEIGHT = 896 | |
| MATRIX_SAMPLER = "dpmpp_2m" | |
| MATRIX_SCHEDULER = "karras" | |
| MATRIX_FIRST_STEPS = 24 | |
| MATRIX_FIRST_CFG = 6 | |
| MATRIX_UPSCALE_METHOD = "bislerp" | |
| MATRIX_UPSCALE_SCALE = 1.1 | |
| MATRIX_SECOND_STEPS = 24 | |
| MATRIX_SECOND_CFG = 5 | |
| MATRIX_DENOISE = .35 | |
| MATRIX_CELL_WIDTH = round(MATRIX_WIDTH / 8 * MATRIX_UPSCALE_SCALE) * 8 | |
| MATRIX_CELL_HEIGHT = round(MATRIX_HEIGHT / 8 * MATRIX_UPSCALE_SCALE) * 8 | |
| MATRIX_LABEL_SIZE = 32 | |
| MATRIX_ROW_LABEL_WIDTH = 160 | |
| MATRIX_GRID_SCALE = .25 | |
| MATRIX_GRID_CELL_WIDTH = round(MATRIX_CELL_WIDTH * MATRIX_GRID_SCALE) | |
| MATRIX_GRID_CELL_HEIGHT = round(MATRIX_CELL_HEIGHT * MATRIX_GRID_SCALE) | |
| COMBINED_MATRIX_TYPE = "checkpoint+sampler+scheduler" | |
| MATRIX_TYPES = ( | |
| "checkpoint", | |
| "sampler", | |
| "scheduler", | |
| "sampler+scheduler", | |
| COMBINED_MATRIX_TYPE, | |
| ) | |
| UPSCALE_GPU_DURATION = 5 | |
| MAX_UPSCALE_BYTES = 20 * MIB | |
| MAX_UPSCALE_PIXELS = 4 * 1024 * 1024 | |
| UPSCALE_ASSETS = { | |
| "style.css": ("upscale.css", "text/css"), | |
| "image.js": ("upscale.js", "text/javascript"), | |
| "page.js": ("upscale-page.js", "text/javascript"), | |
| } | |
| SCAN_THREAD_COUNT = 8 | |
| GPU_DURATION = 25 | |
| GPU_ATTEMPTS = 3 | |
| GPU_RETRY_ERRORS = ( | |
| "gpu task aborted", | |
| "uncorrectable ecc error", | |
| ) | |
| jobs = {} | |
| images = {} | |
| remote_models = {kind: {} for kind in MODEL_KINDS} | |
| state = {} | |
| lock = threading.Lock() | |
| matrix_lock = threading.Lock() | |
| pool = ThreadPoolExecutor(max_workers=1) | |
| backup_pool = ThreadPoolExecutor(max_workers=1) | |
| matrix_pool = ThreadPoolExecutor(max_workers=1) | |
| model_pool = ThreadPoolExecutor(max_workers=1) | |
| model_upload_lock = threading.Lock() | |
| matrix_grids = set() | |
| local_client = None | |
| local_client_lock = threading.Lock() | |
| def log(message): | |
| print(message, flush=True) | |
| def retry_gpu(call): | |
| for attempt in range(GPU_ATTEMPTS): | |
| try: | |
| return call() | |
| except Exception as error: | |
| if ( | |
| attempt == GPU_ATTEMPTS - 1 | |
| or not any( | |
| text in str(error).casefold() | |
| for text in GPU_RETRY_ERRORS | |
| ) | |
| ): | |
| raise | |
| delay = 2**attempt | |
| log(f"GPU task failed, retrying in {delay}s") | |
| time.sleep(delay) | |
| def get_local_client(): | |
| global local_client | |
| with local_client_lock: | |
| if local_client is None: | |
| local_client = Client(LOCAL_URL, verbose=False) | |
| return local_client | |
| def copy_artist_database(): | |
| source = next( | |
| ( | |
| path | |
| for path in ( | |
| DATA_MOUNT / "_cache" / "artists.sqlite", | |
| DATA_DIR / "_cache" / "artists.sqlite", | |
| Path("_cache/artists.sqlite"), | |
| Path.cwd() / "_cache" / "artists.sqlite", | |
| Path.cwd().parent / "_cache" / "artists.sqlite", | |
| Path.cwd().parent / "data" / "_cache" / "artists.sqlite", | |
| DATA_MOUNT / "artists.sqlite", | |
| DATA_DIR / "artists.sqlite", | |
| Path("artists.sqlite"), | |
| Path("../artists.sqlite"), | |
| Path.cwd() / "artists.sqlite", | |
| Path.cwd().parent / "artists.sqlite", | |
| ) | |
| if path.is_file() | |
| ), | |
| None, | |
| ) | |
| if source is None: | |
| log("Cannot find artist database") | |
| return | |
| if source.resolve() != ARTIST_DB.resolve(): | |
| temp = ARTIST_DB.with_suffix(".sqlite.part") | |
| shutil.copy2(source, temp) | |
| temp.replace(ARTIST_DB) | |
| try: | |
| with sqlite3.connect(ARTIST_DB) as database: | |
| count = database.execute("SELECT count(*) FROM artists").fetchone()[0] | |
| log(f"Found artist database: {count} tags loaded") | |
| except Exception: | |
| log("Cannot find artist database") | |
| def cleanup_mount(): | |
| DATA_DIR.mkdir(parents=True, exist_ok=True) | |
| copy_artist_database() | |
| if OUTPUT_DIR.exists(): | |
| shutil.rmtree(OUTPUT_DIR) | |
| log("Deleted output directory") | |
| if IMAGE_DIR.is_dir(): | |
| for path in IMAGE_DIR.rglob("*"): | |
| if path.suffix.casefold() == ".7z" and path.is_file(): | |
| path.unlink() | |
| log(f"Deleted img/{path.relative_to(IMAGE_DIR)}") | |
| if HAS_DATA_MOUNT: | |
| for folder in IMAGE_DIR.glob("????-??-??"): | |
| pngs = list(folder.glob("*.png")) | |
| if not pngs: | |
| continue | |
| check_duplicates = ( | |
| folder.name not in DUPLICATE_CHECK_EXCLUDED_FOLDERS | |
| ) | |
| fingerprints = [] | |
| if check_duplicates: | |
| for path in folder.glob("*.epng"): | |
| with Image.open(BytesIO(stored_bytes(path))) as image: | |
| fingerprints.append(image_fingerprint(image)) | |
| for path in pngs: | |
| with Image.open(path) as image: | |
| duplicate = False | |
| if check_duplicates: | |
| fingerprint = image_fingerprint(image) | |
| duplicate = any( | |
| (fingerprint[0] ^ known[0]).bit_count() | |
| <= DUPLICATE_HASH_DISTANCE | |
| and sum( | |
| abs(left - right) | |
| for left, right in zip(fingerprint[1], known[1]) | |
| ) | |
| <= DUPLICATE_COLOR_DISTANCE | |
| for known in fingerprints | |
| ) | |
| if not duplicate: | |
| save_named_image(image, path.with_suffix(".epng")) | |
| path.unlink() | |
| if duplicate: | |
| log(f"Deleted duplicate img/{path.relative_to(IMAGE_DIR)}") | |
| def download(url, target): | |
| target.parent.mkdir(parents=True, exist_ok=True) | |
| temp = target.with_suffix(target.suffix + ".part") | |
| last = -1 | |
| def report(blocks, block_size, total): | |
| nonlocal last | |
| if total > 0: | |
| mark = min(4, blocks * block_size * 4 // total) | |
| if mark > last: | |
| last = mark | |
| log(f"Downloading {target.name}: {mark * 25}%") | |
| urllib.request.urlretrieve(url, temp, report) | |
| temp.replace(target) | |
| log(f"Downloaded {target.name}: {target.stat().st_size // MIB} MiB") | |
| def model_path(kind, name): | |
| return LOCAL_MODEL_DIR / kind / name | |
| def download_assets(models): | |
| downloads = [] | |
| for kind, name in models: | |
| target = model_path(kind, name) | |
| if target.is_file(): | |
| continue | |
| remote = remote_models[kind].get(name) | |
| if remote is None: | |
| url = BUILTIN_ASSETS.get((kind, name)) | |
| if url is None: | |
| raise ValueError(f"Unknown {kind} file: {name}") | |
| download(url, target) | |
| continue | |
| target.parent.mkdir(parents=True, exist_ok=True) | |
| temp = target.with_suffix(target.suffix + ".part") | |
| downloads.append((kind, name, remote, temp, target)) | |
| if not downloads: | |
| return | |
| log(f"Downloading {len(downloads)} assets from {BUCKET_ID}") | |
| for kind, name, _, _, _ in downloads: | |
| log(f"Downloading {kind}/{name}") | |
| download_bucket_files( | |
| BUCKET_ID, | |
| files=[ | |
| (remote, str(temp)) | |
| for _, _, remote, temp, _ in downloads | |
| ], | |
| token=False, | |
| ) | |
| for kind, name, _, temp, target in downloads: | |
| temp.replace(target) | |
| log( | |
| f"Downloaded {kind}/{name}: " | |
| f"{target.stat().st_size // MIB} MiB" | |
| ) | |
| def is_anima_asset(name): | |
| name = name.casefold() | |
| return name.startswith(ANIMA_PREFIX) and not name.startswith("animag") | |
| def index_bucket_models(): | |
| models = {kind: {} for kind in MODEL_KINDS} | |
| items = [ | |
| item | |
| for item in list_bucket_tree(BUCKET_ID, recursive=True, token=False) | |
| if item.type == "file" | |
| and Path(item.path).suffix.casefold() in MODEL_SUFFIXES | |
| ] | |
| counters = { | |
| kind: {False: 0, True: 0} | |
| for kind in NUMBERED_MODEL_KINDS | |
| } | |
| for item in items: | |
| kind, separator, name = item.path.partition("/") | |
| match = MODEL_NUMBER.match(name) | |
| if separator and kind in counters and match: | |
| anima = is_anima_asset(match.group(2)) | |
| counters[kind][anima] = max(counters[kind][anima], int(match.group(1))) | |
| copies = [] | |
| deletes = [] | |
| for item in sorted(items, key=lambda item: item.path.casefold()): | |
| kind, separator, name = item.path.partition("/") | |
| if not separator or kind not in models: | |
| continue | |
| if kind in counters and not MODEL_NUMBER.match(name): | |
| anima = is_anima_asset(name) | |
| counters[kind][anima] += 1 | |
| name = f"{counters[kind][anima]}_{name}" | |
| path = f"{kind}/{name}" | |
| copies.append(("bucket", BUCKET_ID, item.xet_hash, path)) | |
| deletes.append(item.path) | |
| else: | |
| path = item.path | |
| models[kind][name] = path | |
| if copies: | |
| batch_bucket_files( | |
| BUCKET_ID, | |
| copy=copies, | |
| delete=deletes, | |
| token=IMAGE_TOKEN, | |
| ) | |
| added = { | |
| kind: set(models[kind]) - set(remote_models[kind]) | |
| for kind in MODEL_KINDS | |
| } | |
| remote_models.clear() | |
| remote_models.update(models) | |
| return added | |
| def bucket_numbers(kind, anima): | |
| numbers = set() | |
| for item in list_bucket_tree( | |
| BUCKET_ID, | |
| prefix=f"{kind}/", | |
| recursive=True, | |
| token=False, | |
| ): | |
| if item.type != "file" or "/" in item.path.removeprefix(f"{kind}/"): | |
| continue | |
| match = MODEL_NUMBER.match(Path(item.path).name) | |
| if match and is_anima_asset(match.group(2)) == anima: | |
| numbers.add(int(match.group(1))) | |
| return numbers | |
| def bucket_url_filename(response): | |
| name = response.headers.get_filename() | |
| if not name: | |
| name = Path(urlparse(response.geturl()).path).name | |
| name = unquote(name).replace("\\", "/").rsplit("/", 1)[-1].strip() | |
| match = MODEL_NUMBER.match(name) | |
| return match.group(2) if match else name | |
| def upload_bucket_assets(files, url, kind, anima, password): | |
| if not valid_pass(password): | |
| raise gr.Error("Invalid password") | |
| if kind not in MODEL_KINDS: | |
| raise gr.Error("Invalid CB location") | |
| files = files or [] | |
| url = (url or "").strip() | |
| if not files and not url: | |
| raise gr.Error("Select a file or enter a URL") | |
| temp = None | |
| try: | |
| assets = [] | |
| for file in files: | |
| source = Path(file) | |
| match = MODEL_NUMBER.match(source.name) | |
| number, name = (int(match.group(1)), match.group(2)) \ | |
| if match else (None, source.name) | |
| if source.suffix.casefold() not in MODEL_SUFFIXES: | |
| raise gr.Error(f"Unsupported model file: {name}") | |
| if anima and not is_anima_asset(name): | |
| name = f"{ANIMA_PREFIX}_{name}" | |
| assets.append((source, number, name)) | |
| if url: | |
| parsed = urlparse(url) | |
| if parsed.scheme not in ("http", "https") or not parsed.netloc: | |
| raise gr.Error("Enter a valid URL") | |
| request = urllib.request.Request(url, headers=DOWNLOAD_HEADERS) | |
| if CIVITAI_TOKEN and parsed.hostname in CIVITAI_HOSTS: | |
| request.add_unredirected_header( | |
| "Authorization", | |
| f"Bearer {CIVITAI_TOKEN}", | |
| ) | |
| with urllib.request.urlopen(request) as response: | |
| name = bucket_url_filename(response) | |
| suffix = Path(name).suffix.casefold() | |
| if suffix not in MODEL_SUFFIXES: | |
| raise gr.Error("URL did not return a model file") | |
| if anima and not is_anima_asset(name): | |
| name = f"{ANIMA_PREFIX}_{name}" | |
| with tempfile.NamedTemporaryFile( | |
| suffix=suffix, | |
| delete=False, | |
| ) as file: | |
| temp = Path(file.name) | |
| shutil.copyfileobj(response, file) | |
| assets.append((temp, None, name)) | |
| with model_upload_lock: | |
| used = bucket_numbers(kind, anima) | |
| reserved = { | |
| number for _, number, _ in assets | |
| if number is not None and number not in used | |
| } | |
| next_number = max(used, default=0) + 1 | |
| additions = [] | |
| paths = [] | |
| for source, number, name in assets: | |
| if number is not None and number in reserved: | |
| reserved.remove(number) | |
| else: | |
| while next_number in used or next_number in reserved: | |
| next_number += 1 | |
| number = next_number | |
| next_number += 1 | |
| used.add(number) | |
| path = f"{kind}/{number}_{name}" | |
| additions.append((source, path)) | |
| paths.append(path) | |
| batch_bucket_files( | |
| BUCKET_ID, | |
| add=additions, | |
| token=IMAGE_TOKEN, | |
| ) | |
| index_bucket_models() | |
| return "Added " + ", ".join(paths) | |
| finally: | |
| if temp: | |
| temp.unlink(missing_ok=True) | |
| def model_kind(name): | |
| return "diffusion_models" if is_anima_model(name) else "checkpoints" | |
| def vae_name(name): | |
| return ANIMA_VAE if is_anima_model(name) else DEFAULT_VAE | |
| def generation_models(): | |
| names = set(remote_models["checkpoints"]) | { | |
| name | |
| for name in remote_models["diffusion_models"] | |
| if is_anima_model(name) | |
| } | |
| return sorted( | |
| names, | |
| key=lambda name: ( | |
| is_anima_model(name), | |
| int(name.partition("_")[0]), | |
| name.casefold(), | |
| ), | |
| ) | |
| def model_choices(models): | |
| non_anima = [name for name in models if not is_anima_model(name)] | |
| anima = [name for name in models if is_anima_model(name)] | |
| return [ | |
| ("──────── Non-Anima ────────", NON_ANIMA_HEADER), | |
| *[(name, name) for name in non_anima], | |
| ("──────── Anima ────────", ANIMA_HEADER), | |
| *[(name, name) for name in anima], | |
| ] | |
| def upscale_model_choices(models): | |
| return [ | |
| ("None", ""), | |
| *[(name, name) for name in models], | |
| ] | |
| class LoraRequest(BaseModel): | |
| name: str | |
| strength: float = 1 | |
| clip: float = 0 | |
| class RegionRequest(BaseModel): | |
| prompt: str | |
| area: str | |
| strength: float = Field(1, gt=0, le=10) | |
| class DetailerRequest(BaseModel): | |
| detector: str | |
| model: str = "" | |
| prompt: str = "" | |
| negative: str = "" | |
| sampler: str = DEFAULT_SECOND_SAMPLER | |
| scheduler: str = DEFAULT_SECOND_SCHEDULER | |
| steps: int = Field(DEFAULT_SECOND_STEPS, ge=1, le=100) | |
| cfg: float = Field(DEFAULT_SECOND_CFG, ge=0, le=100) | |
| denoise: float = Field(.35, gt=0, le=1) | |
| def default_loras(): | |
| return [ | |
| LoraRequest(name=name, strength=strength, clip=clip) | |
| for name, strength, clip in DEFAULT_LORAS | |
| ] | |
| class DirectRequest(BaseModel): | |
| prompt: str | |
| sillytavern: dict = Field(default_factory=dict) | |
| second_prompt: str = "" | |
| regions: list[RegionRequest] = Field(default_factory=list, max_length=3) | |
| regional_mode: str = REGIONAL_MODES[0] | |
| detailers: list[DetailerRequest] = Field(default_factory=list) | |
| style_images: list[str] = Field(default_factory=list, max_length=MAX_STYLE_IMAGES) | |
| style_scope: str = DEFAULT_STYLE_SCOPE | |
| style_weight: float = Field(DEFAULT_STYLE_WEIGHT, ge=0, le=5) | |
| style_end: float = Field(DEFAULT_STYLE_END, gt=0, le=1) | |
| second_style_images: list[str] = Field( | |
| default_factory=list, | |
| max_length=MAX_STYLE_IMAGES, | |
| ) | |
| second_style_weight: float = Field(DEFAULT_STYLE_WEIGHT, ge=0, le=5) | |
| second_style_end: float = Field(DEFAULT_STYLE_END, gt=0, le=1) | |
| model: str = DEFAULT_MODEL | |
| loras: list[LoraRequest] = Field(default_factory=default_loras) | |
| second_model: str = "" | |
| second_loras: list[LoraRequest] = Field(default_factory=list) | |
| negative: str = DEFAULT_NEGATIVE | |
| second_negative: str = "" | |
| width: int = Field(1152, ge=64, le=2048) | |
| height: int = Field(896, ge=64, le=2048) | |
| batch_size: int = Field(DEFAULT_BATCH_SIZE, ge=1, le=MAX_BATCH_SIZE) | |
| sampler: str = DEFAULT_SAMPLER | |
| scheduler: str = DEFAULT_SCHEDULER | |
| steps: int = Field(DEFAULT_STEPS, ge=1, le=100) | |
| cfg: float = Field(DEFAULT_CFG, ge=0, le=100) | |
| upscale: bool = False | |
| upscale_method: str = DEFAULT_UPSCALE_METHOD | |
| upscale_model: str = DEFAULT_UPSCALE_MODEL | |
| upscale_scale: float = Field(DEFAULT_UPSCALE_SCALE, gt=0) | |
| second_sampler: str = DEFAULT_SECOND_SAMPLER | |
| second_scheduler: str = DEFAULT_SECOND_SCHEDULER | |
| second_steps: int = Field(DEFAULT_SECOND_STEPS, ge=1, le=100) | |
| second_cfg: float = Field(DEFAULT_SECOND_CFG, ge=0, le=100) | |
| denoise: float = Field(DEFAULT_DENOISE, ge=0, le=1) | |
| return_scale: float = Field(DEFAULT_RETURN_SCALE, ge=.01, le=1) | |
| artist_min_posts: int = Field(100, ge=0) | |
| artist_blacklist: str = "" | |
| selected_artists: list[str] = Field(default_factory=list) | |
| def resolve_artist_prompts(request): | |
| prompts = [request.prompt, request.second_prompt] | |
| prompts.extend(d.prompt for d in getattr(request, "detailers", []) if getattr(d, "prompt", None)) | |
| prompts.extend(r.prompt for r in getattr(request, "regions", []) if getattr(r, "prompt", None)) | |
| matches = ARTIST_PLACEHOLDER.findall("\n".join(prompts)) | |
| identifiers = list(dict.fromkeys(value.lstrip("0") or "1" for value in matches)) | |
| if not identifiers: | |
| return | |
| if not ARTIST_DB.is_file(): | |
| copy_artist_database() | |
| if not ARTIST_DB.is_file(): | |
| raise ValueError("Artist database is unavailable") | |
| blacklist = { | |
| artist.strip().casefold() | |
| for artist in request.artist_blacklist.split(",") | |
| if artist.strip() | |
| } | |
| database = sqlite3.connect(ARTIST_DB) | |
| try: | |
| artists = [ | |
| row[0] | |
| for row in database.execute( | |
| "SELECT artist FROM artists WHERE post_count >= ?", | |
| (request.artist_min_posts,), | |
| ) | |
| if row[0].casefold() not in blacklist | |
| ] | |
| finally: | |
| database.close() | |
| if len(artists) < len(identifiers): | |
| raise ValueError("Not enough artists match the post minimum and blacklist") | |
| selected = secrets.SystemRandom().sample(artists, len(identifiers)) | |
| replacements = dict(zip(identifiers, selected)) | |
| def replace(match): | |
| artist = replacements[match.group(1).lstrip("0") or "1"] | |
| return re.sub(r"(?<!\\)[()]", r"\\\g<0>", artist) | |
| request.prompt = ARTIST_PLACEHOLDER.sub(replace, request.prompt) | |
| request.second_prompt = ARTIST_PLACEHOLDER.sub(replace, request.second_prompt) | |
| for d in getattr(request, "detailers", []): | |
| if getattr(d, "prompt", None): | |
| d.prompt = ARTIST_PLACEHOLDER.sub(replace, d.prompt) | |
| for r in getattr(request, "regions", []): | |
| if getattr(r, "prompt", None): | |
| r.prompt = ARTIST_PLACEHOLDER.sub(replace, r.prompt) | |
| request.selected_artists = list(dict.fromkeys(selected)) | |
| class ModelRequest(BaseModel): | |
| model: str = DEFAULT_MODEL | |
| loras: list[LoraRequest] = Field(default_factory=default_loras) | |
| instant_style: bool = False | |
| second_model: str = "" | |
| second_loras: list[LoraRequest] = Field(default_factory=list) | |
| upscale: bool = False | |
| upscale_model: str = DEFAULT_UPSCALE_MODEL | |
| detailers: list[DetailerRequest] = Field(default_factory=list) | |
| class DownloadRequest(BaseModel): | |
| items: list[str] | |
| class StarRequest(BaseModel): | |
| path: str | |
| starred: bool | |
| class MatrixRequest(BaseModel): | |
| model_config = ConfigDict(extra="forbid") | |
| generation: str = Field(pattern=r"^g\d+$") | |
| type: str = "checkpoint" | |
| positive: str | |
| negative: str | |
| model_1: int = Field(0, ge=0) | |
| model_2: int = Field(0, ge=0) | |
| sampler: str = "" | |
| class MatrixCellRequest(BaseModel): | |
| model_config = ConfigDict(extra="forbid") | |
| generation: str = Field(pattern=r"^g\d+$") | |
| positive: str | |
| negative: str | |
| folder: str | |
| first: tuple[str, str] | |
| second: tuple[str, str] | |
| sampler: str = MATRIX_SAMPLER | |
| scheduler: str = MATRIX_SCHEDULER | |
| second_sampler: str = MATRIX_SAMPLER | |
| second_scheduler: str = MATRIX_SCHEDULER | |
| output: str = "" | |
| def ensure_comfy(): | |
| if COMFYUI_PATH.is_dir(): | |
| log(f"Found ComfyUI at {COMFYUI_PATH}") | |
| else: | |
| archive = Path.cwd() / "comfyui.zip" | |
| download( | |
| "https://github.com/Comfy-Org/ComfyUI/archive/refs/heads/master.zip", | |
| archive, | |
| ) | |
| log("Extracting ComfyUI") | |
| shutil.unpack_archive(archive, Path.cwd()) | |
| archive.unlink() | |
| next(Path.cwd().glob("ComfyUI-*")).replace(COMFYUI_PATH) | |
| log(f"Installed ComfyUI at {COMFYUI_PATH}") | |
| subprocess.run( | |
| [ | |
| sys.executable, "-m", "pip", "install", "-q", "-r", | |
| str(COMFYUI_PATH / "requirements.txt"), | |
| ], | |
| check=True, | |
| ) | |
| def ensure_custom_nodes(): | |
| shutil.rmtree(CUSTOM_NODES_DIR, ignore_errors=True) | |
| CUSTOM_NODES_DIR.mkdir(parents=True) | |
| for name in REQUIRED_CUSTOM_NODES: | |
| source = MOUNTED_CUSTOM_NODES_DIR / name | |
| target = CUSTOM_NODES_DIR / name | |
| if source.is_dir(): | |
| target.symlink_to(source, target_is_directory=True) | |
| elif name in CUSTOM_NODE_REPOS: | |
| subprocess.run( | |
| [ | |
| "git", "clone", "-q", "--depth", "1", | |
| CUSTOM_NODE_REPOS[name], str(target), | |
| ], | |
| check=True, | |
| ) | |
| else: | |
| raise FileNotFoundError(f"Missing required custom node: {source}") | |
| modules = CUSTOM_NODE_MODULES.get(name) | |
| if modules and not any(importlib.util.find_spec(item) is None for item in modules): | |
| continue | |
| requirements = target / "requirements.txt" | |
| if requirements.is_file(): | |
| subprocess.run( | |
| [ | |
| sys.executable, "-m", "pip", "install", "-q", "-r", | |
| str(requirements), | |
| ], | |
| check=True, | |
| stdout=subprocess.DEVNULL, | |
| stderr=subprocess.DEVNULL, | |
| ) | |
| def write_detector_whitelist(): | |
| path = ( | |
| COMFYUI_PATH | |
| / "user" | |
| / "default" | |
| / "ComfyUI-Impact-Subpack" | |
| / "model-whitelist.txt" | |
| ) | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| names = set(path.read_text().splitlines()) if path.is_file() else set() | |
| names.update(remote_models["ultralytics"]) | |
| path.write_text("\n".join(sorted(names, key=str.casefold)) + "\n") | |
| def node_input_options(node, name): | |
| spec = node.INPUT_TYPES()["required"][name] | |
| if isinstance(spec, tuple): | |
| if isinstance(spec[0], (list, tuple)): | |
| return list(spec[0]) | |
| if len(spec) > 1 and isinstance(spec[1], dict): | |
| return list(spec[1].get("options", [])) | |
| return [] | |
| def init_comfy(): | |
| if state: | |
| return | |
| log(f"Using data directory {DATA_DIR}") | |
| ensure_comfy() | |
| with ThreadPoolExecutor(max_workers=2) as startup: | |
| model_index = startup.submit(index_bucket_models) | |
| custom_nodes = startup.submit(ensure_custom_nodes) | |
| model_index.result() | |
| custom_nodes.result() | |
| write_detector_whitelist() | |
| log("Importing ComfyUI") | |
| sys.path.insert(0, str(COMFYUI_PATH)) | |
| from comfy.cli_args import args | |
| args.cpu_vae = False | |
| args.disable_pinned_memory = True | |
| import comfy.sd | |
| import comfy.utils | |
| import folder_paths | |
| from nodes import ( | |
| CLIPTextEncode, | |
| CheckpointLoaderSimple, | |
| ConditioningSetMask, | |
| EmptyLatentImage, | |
| KSampler, | |
| LatentUpscale, | |
| LatentUpscaleBy, | |
| LoraLoader, | |
| UNETLoader, | |
| VAEDecode, | |
| VAEEncode, | |
| VAELoader, | |
| ) | |
| from comfy_extras.nodes_align_your_steps import AlignYourStepsScheduler | |
| from comfy_extras.nodes_custom_sampler import ( | |
| BasicScheduler, | |
| KSamplerSelect, | |
| SamplerCustom, | |
| ) | |
| from comfy_extras.nodes_upscale_model import ( | |
| ImageUpscaleWithModel, | |
| UpscaleModelLoader, | |
| ) | |
| for kind in MODEL_KINDS: | |
| local = LOCAL_MODEL_DIR / kind | |
| local.mkdir(parents=True, exist_ok=True) | |
| folder_paths.add_model_folder_path( | |
| COMFY_KINDS.get(kind, kind), | |
| str(local), | |
| is_default=True, | |
| ) | |
| folder_paths.add_model_folder_path( | |
| "ultralytics_bbox", | |
| str(LOCAL_MODEL_DIR / "ultralytics"), | |
| ) | |
| default_sample_options = KSampler.INPUT_TYPES()["required"] | |
| default_samplers = list(default_sample_options["sampler_name"][0]) | |
| default_schedulers = list(default_sample_options["scheduler"][0]) | |
| default_upscale_methods = list( | |
| LatentUpscaleBy.INPUT_TYPES()["required"]["upscale_method"][0] | |
| ) | |
| folder_paths.add_model_folder_path("custom_nodes", str(CUSTOM_NODES_DIR)) | |
| import_custom_nodes() | |
| from nodes import NODE_CLASS_MAPPINGS | |
| custom_samplers = {} | |
| custom_sampler_nodes = {} | |
| for node_name, prefix in ( | |
| ("DynSamplerSelect", "ppm-dyn"), | |
| ("CFGPPSamplerSelect", "ppm-cfgpp"), | |
| ("PPMSamplerSelect", "ppm"), | |
| ): | |
| node = NODE_CLASS_MAPPINGS.get(node_name) | |
| if node is None: | |
| continue | |
| custom_sampler_nodes[node_name] = node() | |
| for name in node_input_options(node, "sampler_name"): | |
| if name not in default_samplers: | |
| custom_samplers[f"{prefix}:{name}"] = (node_name, name) | |
| current_options = KSampler.INPUT_TYPES()["required"] | |
| sampler_sources = { | |
| name: "RES4LYF" | |
| for name in current_options["sampler_name"][0] | |
| if name not in default_samplers | |
| } | |
| ppm_schedulers = { | |
| "ays", "ays+", "ays_30", "ays_30+", "gits", "beta_1_1", | |
| } | |
| scheduler_sources = { | |
| name: ( | |
| "ComfyUI-ppm" if name in ppm_schedulers | |
| else "RES4LYF" if name in {"beta57", "bong_tangent"} | |
| else "Custom node" | |
| ) | |
| for name in current_options["scheduler"][0] | |
| if name not in default_schedulers | |
| } | |
| state.update( | |
| apply_lora=comfy.sd.load_lora_for_models, | |
| chains={}, | |
| checkpoint=CheckpointLoaderSimple(), | |
| clip_type=comfy.sd.CLIPType.STABLE_DIFFUSION, | |
| clips={}, | |
| lora=LoraLoader(), | |
| load_checkpoint=comfy.sd.load_checkpoint_guess_config, | |
| load_clip=comfy.sd.load_clip, | |
| load_torch_file=comfy.utils.load_torch_file, | |
| model_management=comfy.sd.model_management, | |
| encode=CLIPTextEncode(), | |
| mask=ConditioningSetMask(), | |
| folders=folder_paths, | |
| loras={}, | |
| models={}, | |
| vae_loader=VAELoader(), | |
| vaes={}, | |
| latent=EmptyLatentImage(), | |
| sample=KSampler(), | |
| align=AlignYourStepsScheduler(), | |
| basic_scheduler=BasicScheduler(), | |
| sampler_select=KSamplerSelect(), | |
| sample_custom=SamplerCustom(), | |
| decode=VAEDecode(), | |
| vae_encode=VAEEncode(), | |
| upscale=LatentUpscaleBy(), | |
| resize_latent=LatentUpscale(), | |
| upscale_image=ImageUpscaleWithModel(), | |
| upscale_model_loader=UpscaleModelLoader(), | |
| upscale_models={}, | |
| unet=UNETLoader(), | |
| detector_provider=NODE_CLASS_MAPPINGS["UltralyticsDetectorProvider"](), | |
| detectors={}, | |
| face_detailer=NODE_CLASS_MAPPINGS["FaceDetailer"](), | |
| attention_couple=NODE_CLASS_MAPPINGS["AttentionCouplePPM"](), | |
| clip_vision_loader=NODE_CLASS_MAPPINGS["CLIPVisionLoader"](), | |
| style_model_loader=NODE_CLASS_MAPPINGS["IPAdapterModelLoader"](), | |
| style_apply=NODE_CLASS_MAPPINGS["IPAdapterAdvanced"](), | |
| style_pipeline=None, | |
| custom_samplers=custom_samplers, | |
| custom_sampler_nodes=custom_sampler_nodes, | |
| default_samplers=default_samplers, | |
| default_schedulers=default_schedulers, | |
| default_upscale_methods=default_upscale_methods, | |
| sampler_sources=sampler_sources, | |
| scheduler_sources=scheduler_sources, | |
| ) | |
| log("ComfyUI initialization complete") | |
| def output_item(output): | |
| return getattr(output, "result", output)[0] | |
| def sampler_names(): | |
| options = state["sample"].INPUT_TYPES()["required"] | |
| return [*options["sampler_name"][0], *state["custom_samplers"]] | |
| def scheduler_names(): | |
| options = state["sample"].INPUT_TYPES()["required"] | |
| return [*options["scheduler"][0], ALIGN_SCHEDULER] | |
| def select_sampler(model, seed, name): | |
| custom = state["custom_samplers"].get(name) | |
| if custom is None: | |
| return output_item(state["sampler_select"].get_sampler(name)) | |
| node_name, sampler_name = custom | |
| node = state["custom_sampler_nodes"][node_name] | |
| if node_name == "PPMSamplerSelect": | |
| output = node.get_sampler(sampler_name=sampler_name, model=model) | |
| else: | |
| output = node.get_sampler(sampler_name=sampler_name) | |
| return output_item(output) | |
| def run_sampler( | |
| model, | |
| seed, | |
| steps, | |
| cfg, | |
| sampler_name, | |
| scheduler, | |
| positive, | |
| negative, | |
| latent_image, | |
| denoise, | |
| ): | |
| custom_sampler = sampler_name in state["custom_samplers"] | |
| if scheduler != ALIGN_SCHEDULER and not custom_sampler: | |
| return state["sample"].sample( | |
| model=model, | |
| seed=seed, | |
| steps=steps, | |
| cfg=cfg, | |
| sampler_name=sampler_name, | |
| scheduler=scheduler, | |
| positive=positive, | |
| negative=negative, | |
| latent_image=latent_image, | |
| denoise=denoise, | |
| )[0] | |
| if scheduler == ALIGN_SCHEDULER: | |
| output = state["align"].get_sigmas( | |
| ALIGN_MODEL_TYPE, | |
| steps, | |
| denoise, | |
| ) | |
| else: | |
| output = state["basic_scheduler"].get_sigmas( | |
| model, | |
| scheduler, | |
| steps, | |
| denoise, | |
| ) | |
| sigmas = output_item(output) | |
| sampler = select_sampler(model, seed, sampler_name) | |
| output = state["sample_custom"].sample( | |
| model=model, | |
| add_noise=True, | |
| noise_seed=seed, | |
| cfg=cfg, | |
| positive=positive, | |
| negative=negative, | |
| sampler=sampler, | |
| sigmas=sigmas, | |
| latent_image=latent_image, | |
| ) | |
| return output_item(output) | |
| def convert_latent(samples, source_vae, target_vae): | |
| pixels = state["decode"].decode(vae=source_vae, samples=samples)[0] | |
| return state["vae_encode"].encode(vae=target_vae, pixels=pixels)[0] | |
| def load_upscale_model(name): | |
| if name not in state["upscale_models"]: | |
| stage_model("upscale_models", name) | |
| log(f"Loading upscale model {name}") | |
| state["upscale_models"][name] = state[ | |
| "upscale_model_loader" | |
| ].load_model(name)[0] | |
| log(f"Loaded upscale model {name}") | |
| return state["upscale_models"][name] | |
| def load_detector(name): | |
| if name not in state["detectors"]: | |
| stage_model("ultralytics", name) | |
| log(f"Loading detector {name}") | |
| state["detectors"][name] = state["detector_provider"].doit( | |
| f"bbox/{name}" | |
| )[0] | |
| log(f"Loaded detector {name}") | |
| return state["detectors"][name] | |
| def area_box(area): | |
| key = area.strip().lower() | |
| value = AREA_PRESETS.get(key, key) | |
| match = re.fullmatch(r"([a-e])([1-5])(?::([a-e])([1-5]))?", value) | |
| if not match: | |
| raise ValueError(f"Unsupported region area: {area}") | |
| left, top, right, bottom = match.groups() | |
| right = right or left | |
| bottom = bottom or top | |
| x1, x2 = sorted((ord(left) - ord("a"), ord(right) - ord("a"))) | |
| y1, y2 = sorted((int(top) - 1, int(bottom) - 1)) | |
| return ( | |
| x1 / GRID_SIZE, | |
| y1 / GRID_SIZE, | |
| (x2 - x1 + 1) / GRID_SIZE, | |
| (y2 - y1 + 1) / GRID_SIZE, | |
| ) | |
| def prepare_regions(regions): | |
| if not regions: | |
| return [] | |
| layout = AUTO_LAYOUTS[len(regions)] | |
| return [ | |
| ( | |
| region.prompt, | |
| *(layout[index] if region.area.strip().lower() == "auto" | |
| else area_box(region.area)), | |
| region.strength, | |
| ) | |
| for index, region in enumerate(regions) | |
| ] | |
| def region_mask(region, image_width, image_height): | |
| _, x, y, width, height, _ = region | |
| mask_width = image_width // LATENT_SCALE | |
| mask_height = image_height // LATENT_SCALE | |
| x, y, width, height, left, top, right, bottom = mask_box( | |
| x, | |
| y, | |
| width, | |
| height, | |
| mask_width, | |
| mask_height, | |
| ) | |
| mask = torch.zeros((1, mask_height, mask_width)) | |
| area = mask[:, y:y + height, x:x + width] | |
| area.fill_(1) | |
| if left: | |
| area[:, :, :left] *= torch.linspace(1 / left, 1, left) | |
| if top: | |
| area[:, :top, :] *= torch.linspace(1 / top, 1, top).view(1, -1, 1) | |
| if right: | |
| area[:, :, -right:] *= torch.linspace(1, 1 / right, right) | |
| if bottom: | |
| area[:, -bottom:, :] *= torch.linspace(1, 1 / bottom, bottom).view( | |
| 1, | |
| -1, | |
| 1, | |
| ) | |
| return mask | |
| def scale_conditioning(conditioning, strength): | |
| return [ | |
| [item[0], {**item[1], "strength": strength}] | |
| for item in conditioning | |
| ] | |
| def encode_positive( | |
| model, | |
| clip, | |
| prompt, | |
| regions, | |
| image_width, | |
| image_height, | |
| regional_mode, | |
| ): | |
| positive = state["encode"].encode(clip=clip, text=prompt)[0] | |
| if not regions: | |
| return model, positive | |
| if regional_mode == "conditioning": | |
| positive = [ | |
| [ | |
| item[0], | |
| { | |
| **item[1], | |
| "start_percent": ENVIRONMENT_START, | |
| "end_percent": 1, | |
| "strength": REGIONAL_GLOBAL_STRENGTH, | |
| }, | |
| ] | |
| for item in positive | |
| ] | |
| conditionings = [] | |
| masks = [] | |
| for region in regions: | |
| region_prompt, _, _, _, _, strength = region | |
| conditioning = state["encode"].encode( | |
| clip=clip, | |
| text=region_prompt, | |
| )[0] | |
| mask = region_mask(region, image_width, image_height) | |
| if regional_mode == "attention": | |
| conditionings.append(scale_conditioning(conditioning, strength)) | |
| masks.append(mask) | |
| continue | |
| conditioning = state["mask"].append( | |
| conditioning=conditioning, | |
| mask=mask, | |
| set_cond_area="mask bounds", | |
| strength=strength, | |
| )[0] | |
| positive += conditioning | |
| if regional_mode == "conditioning": | |
| return model, positive | |
| inputs = { | |
| "model": model, | |
| "base_cond": scale_conditioning( | |
| positive, | |
| REGIONAL_GLOBAL_STRENGTH, | |
| ), | |
| "base_mask": torch.ones_like(masks[0]), | |
| } | |
| for index, (conditioning, mask) in enumerate( | |
| zip(conditionings, masks), | |
| 1, | |
| ): | |
| inputs[f"cond_{index}"] = conditioning | |
| inputs[f"mask_{index}"] = mask | |
| output = state["attention_couple"].execute(**inputs) | |
| return getattr(output, "result", output)[0], inputs["base_cond"] | |
| def has_style(request): | |
| return bool( | |
| getattr(request, "style_images", []) | |
| or getattr(request, "second_style_images", []) | |
| ) | |
| def wants_style(request): | |
| return has_style(request) or getattr(request, "instant_style", False) | |
| def style_stage_enabled(request, stage): | |
| first = bool(getattr(request, "style_images", [])) | |
| second = bool(getattr(request, "second_style_images", [])) | |
| if stage == "first": | |
| return first | |
| if stage == "second": | |
| return second or ( | |
| first and request.style_scope in ("generation", "all") | |
| ) | |
| return first and request.style_scope == "all" | |
| def decode_style_images(values): | |
| images = [] | |
| for value in values: | |
| encoded = value.partition(",")[2] if value.startswith("data:") else value | |
| if len(encoded) > (MAX_STYLE_IMAGE_SIZE + 2) // 3 * 4: | |
| raise ValueError("InstantStyle image exceeds 20 MiB") | |
| try: | |
| data = base64.b64decode(encoded, validate=True) | |
| except (ValueError, binascii.Error) as error: | |
| raise ValueError("Invalid InstantStyle image") from error | |
| if len(data) > MAX_STYLE_IMAGE_SIZE: | |
| raise ValueError("InstantStyle image exceeds 20 MiB") | |
| try: | |
| with Image.open(BytesIO(data)) as image: | |
| image = ImageOps.fit( | |
| image.convert("RGB"), | |
| (STYLE_IMAGE_SIZE, STYLE_IMAGE_SIZE), | |
| Image.Resampling.LANCZOS, | |
| ) | |
| images.append(np.asarray(image, dtype=np.float32) / 255) | |
| except Exception as error: | |
| raise ValueError("Invalid InstantStyle image") from error | |
| return torch.from_numpy(np.stack(images)) | |
| def load_style_pipeline(): | |
| if state["style_pipeline"] is None: | |
| ipadapter = state["style_model_loader"].load_ipadapter_model( | |
| STYLE_IPADAPTER | |
| )[0] | |
| clip_vision = state["clip_vision_loader"].load_clip( | |
| STYLE_CLIP_VISION | |
| )[0] | |
| state["style_pipeline"] = ipadapter, clip_vision | |
| return state["style_pipeline"] | |
| def apply_style(request, model, image, stage): | |
| if not style_stage_enabled(request, stage): | |
| return model | |
| ipadapter, clip_vision = load_style_pipeline() | |
| second = stage == "second" and request.second_style_images | |
| return state["style_apply"].apply_ipadapter( | |
| model=model, | |
| ipadapter=ipadapter, | |
| clip_vision=clip_vision, | |
| image=image, | |
| weight=request.second_style_weight if second else request.style_weight, | |
| weight_type=STYLE_WEIGHT_TYPE, | |
| combine_embeds="average", | |
| start_at=0, | |
| end_at=request.second_style_end if second else request.style_end, | |
| embeds_scaling=STYLE_EMBEDS_SCALING, | |
| )[0] | |
| def run_detailers( | |
| request, | |
| image, | |
| base_model, | |
| base_clip, | |
| base_vae, | |
| seeds, | |
| style_image, | |
| ): | |
| final_prompt = request.second_prompt or request.prompt \ | |
| if request.upscale else request.prompt | |
| final_negative = request.second_negative or request.negative \ | |
| if request.upscale else request.negative | |
| for detailer, seed in zip(request.detailers, seeds): | |
| if detailer.model: | |
| model, clip = load_chain(detailer.model, []) | |
| vae = load_vae(vae_name(detailer.model)) | |
| prompt, negative_prompt = request.prompt, request.negative | |
| else: | |
| model, clip, vae = base_model, base_clip, base_vae | |
| prompt, negative_prompt = final_prompt, final_negative | |
| model = apply_style(request, model, style_image, "detailer") | |
| positive = state["encode"].encode( | |
| clip=clip, | |
| text=detailer.prompt or prompt, | |
| )[0] | |
| negative = state["encode"].encode( | |
| clip=clip, | |
| text=detailer.negative or negative_prompt, | |
| )[0] | |
| image = state["face_detailer"].doit( | |
| image=image, | |
| model=model, | |
| clip=clip, | |
| vae=vae, | |
| guide_size=DETAILER_GUIDE_SIZE, | |
| guide_size_for=True, | |
| max_size=DETAILER_MAX_SIZE, | |
| seed=seed, | |
| steps=detailer.steps, | |
| cfg=detailer.cfg, | |
| sampler_name=detailer.sampler, | |
| scheduler=detailer.scheduler, | |
| positive=positive, | |
| negative=negative, | |
| denoise=detailer.denoise, | |
| feather=DETAILER_FEATHER, | |
| noise_mask=True, | |
| force_inpaint=True, | |
| bbox_threshold=DETAILER_THRESHOLD, | |
| bbox_dilation=DETAILER_DILATION, | |
| bbox_crop_factor=DETAILER_CROP, | |
| sam_detection_hint="none", | |
| sam_dilation=0, | |
| sam_threshold=.93, | |
| sam_bbox_expansion=0, | |
| sam_mask_hint_threshold=.7, | |
| sam_mask_hint_use_negative="False", | |
| drop_size=DETAILER_DROP_SIZE, | |
| bbox_detector=load_detector(detailer.detector), | |
| wildcard="", | |
| cycle=1, | |
| )[0] | |
| return image | |
| def infer_image( | |
| request, | |
| regions, | |
| latent, | |
| first_seed, | |
| second_seed, | |
| detailer_seeds, | |
| style_image, | |
| second_style_image, | |
| ): | |
| request = DirectRequest.model_validate(request) | |
| with torch.inference_mode(), warnings.catch_warnings(): | |
| warnings.filterwarnings( | |
| "ignore", | |
| message=r"Should have t[ab](?:<=|>=)t[01] but got", | |
| category=UserWarning, | |
| module=r"torchsde\._brownian\.brownian_interval", | |
| ) | |
| first_model = request.model | |
| second_model = request.second_model or first_model | |
| first_vae = load_vae(vae_name(first_model)) | |
| second_vae = load_vae(vae_name(second_model)) | |
| base_model, clip = load_chain(first_model, request.loras) | |
| model = apply_style(request, base_model, style_image, "first") | |
| model, positive = encode_positive( | |
| model, | |
| clip, | |
| request.prompt, | |
| regions, | |
| request.width, | |
| request.height, | |
| request.regional_mode, | |
| ) | |
| negative = state["encode"].encode(clip=clip, text=request.negative)[0] | |
| samples = run_sampler( | |
| model, | |
| first_seed, | |
| request.steps, | |
| request.cfg, | |
| request.sampler, | |
| request.scheduler, | |
| positive, | |
| negative, | |
| latent, | |
| 1, | |
| ) | |
| if request.upscale: | |
| width, height = upscale_size( | |
| request.width, | |
| request.height, | |
| request.upscale_scale, | |
| ) | |
| samples = state["resize_latent"].upscale( | |
| samples=samples, | |
| upscale_method=request.upscale_method, | |
| width=width, | |
| height=height, | |
| crop="disabled", | |
| )[0] | |
| if is_anima_model(first_model) != is_anima_model(second_model): | |
| samples = convert_latent(samples, first_vae, second_vae) | |
| if request.second_model: | |
| base_model, clip = load_chain( | |
| request.second_model, | |
| request.second_loras, | |
| ) | |
| model = apply_style( | |
| request, | |
| base_model, | |
| second_style_image, | |
| "second", | |
| ) | |
| model, positive = encode_positive( | |
| model, | |
| clip, | |
| request.second_prompt or request.prompt, | |
| regions, | |
| width, | |
| height, | |
| request.regional_mode, | |
| ) | |
| negative = state["encode"].encode( | |
| clip=clip, | |
| text=request.second_negative or request.negative, | |
| )[0] | |
| samples = run_sampler( | |
| model, | |
| second_seed, | |
| request.second_steps, | |
| request.second_cfg, | |
| request.second_sampler, | |
| request.second_scheduler, | |
| positive, | |
| negative, | |
| samples, | |
| request.denoise, | |
| ) | |
| image = state["decode"].decode( | |
| vae=second_vae if request.upscale else first_vae, | |
| samples=samples, | |
| )[0] | |
| image = run_detailers( | |
| request, | |
| image, | |
| base_model, | |
| clip, | |
| second_vae if request.upscale else first_vae, | |
| detailer_seeds, | |
| style_image, | |
| ) | |
| if request.upscale and request.upscale_model: | |
| image = state["upscale_image"].upscale( | |
| upscale_model=load_upscale_model(request.upscale_model), | |
| image=image, | |
| )[0] | |
| return image | |
| def infer_upscale(image, model): | |
| if state["model_management"].get_torch_device().type != "cuda": | |
| raise RuntimeError("CUDA GPU is required for upscaling") | |
| with torch.inference_mode(): | |
| try: | |
| return state["upscale_image"].upscale( | |
| upscale_model=model, | |
| image=image, | |
| )[0].cpu() | |
| finally: | |
| model.to(CPU_DEVICE) | |
| def infer_ping(seed): | |
| with torch.inference_mode(): | |
| model, vae, latent, positive, negative = state["ping"] | |
| samples = run_sampler( | |
| model, | |
| seed, | |
| PING_STEPS, | |
| 1, | |
| PING_SAMPLER, | |
| PING_SCHEDULER, | |
| positive, | |
| negative, | |
| latent, | |
| 1, | |
| ) | |
| return state["decode"].decode(vae=vae, samples=samples)[0] | |
| def generate_gpu(request, regions, style_image, second_style_image): | |
| latent = state["latent"].generate( | |
| width=request.width, | |
| height=request.height, | |
| batch_size=request.batch_size, | |
| )[0] | |
| first_seed = secrets.randbits(64) | |
| second_seed = secrets.randbits(64) if request.upscale else None | |
| detailer_seeds = [secrets.randbits(64) for _ in request.detailers] | |
| image = retry_gpu( | |
| lambda: infer_image( | |
| request.model_dump(), | |
| regions, | |
| latent, | |
| first_seed, | |
| second_seed, | |
| detailer_seeds, | |
| style_image, | |
| second_style_image, | |
| ) | |
| ) | |
| return image, first_seed, second_seed, detailer_seeds | |
| def load_model(name): | |
| if name in state["models"]: | |
| return state["models"][name] | |
| stage_model(model_kind(name), name) | |
| if is_anima_model(name): | |
| stage_model("clip", ANIMA_CLIP) | |
| if ANIMA_CLIP not in state["clips"]: | |
| log(f"Loading text encoder {ANIMA_CLIP}") | |
| path = state["folders"].get_full_path_or_raise( | |
| COMFY_KINDS["clip"], ANIMA_CLIP | |
| ) | |
| state["clips"][ANIMA_CLIP] = state["load_clip"]( | |
| [path], | |
| embedding_directory=state["folders"].get_folder_paths( | |
| "embeddings" | |
| ), | |
| clip_type=state["clip_type"], | |
| model_options={"initial_device": CPU_DEVICE}, | |
| ) | |
| log(f"Loading diffusion model {name}") | |
| model = state["unet"].load_unet( | |
| unet_name=name, | |
| weight_dtype="default", | |
| )[0] | |
| state["models"][name] = model, state["clips"][ANIMA_CLIP] | |
| else: | |
| log(f"Loading checkpoint {name}") | |
| path = state["folders"].get_full_path_or_raise("checkpoints", name) | |
| initial_device = state["model_management"].unet_inital_load_device | |
| state["model_management"].unet_inital_load_device = ( | |
| lambda *_: CPU_DEVICE | |
| ) | |
| try: | |
| state["models"][name] = state["load_checkpoint"]( | |
| path, | |
| output_vae=False, | |
| embedding_directory=state["folders"].get_folder_paths( | |
| "embeddings" | |
| ), | |
| te_model_options={"initial_device": CPU_DEVICE}, | |
| )[:2] | |
| finally: | |
| state["model_management"].unet_inital_load_device = initial_device | |
| log(f"Loaded {name}") | |
| return state["models"][name] | |
| def load_lora(name): | |
| if name not in state["loras"]: | |
| stage_model("loras", name) | |
| log(f"Loading LoRA {name}") | |
| path = state["folders"].get_full_path_or_raise("loras", name) | |
| state["loras"][name] = state["load_torch_file"]( | |
| path, | |
| safe_load=True, | |
| return_metadata=True, | |
| ) | |
| log(f"Loaded LoRA {name}") | |
| return state["loras"][name] | |
| def load_vae(name): | |
| if name not in state["vaes"]: | |
| stage_model("vae", name) | |
| log(f"Loading VAE {name}") | |
| state["vaes"][name] = state["vae_loader"].load_vae( | |
| vae_name=name | |
| )[0] | |
| log(f"Loaded VAE {name}") | |
| return state["vaes"][name] | |
| def chain_key(model_name, loras): | |
| return ( | |
| model_name, | |
| tuple((lora.name, lora.strength, lora.clip) for lora in loras), | |
| ) | |
| def load_chain(model_name, loras): | |
| key = chain_key(model_name, loras) | |
| if key in state["chains"]: | |
| return state["chains"][key] | |
| model, clip = load_model(model_name) | |
| for lora in loras: | |
| data, metadata = load_lora(lora.name) | |
| model, clip = state["apply_lora"]( | |
| model, | |
| clip, | |
| data, | |
| lora.strength, | |
| lora.clip, | |
| lora_metadata=metadata, | |
| ) | |
| state["chains"][key] = model, clip | |
| return state["chains"][key] | |
| def stage_model(kind, name): | |
| target = model_path(kind, name) | |
| if not target.is_file(): | |
| stage_models([(kind, name)]) | |
| return target | |
| def stage_models(models): | |
| models = list(dict.fromkeys(models)) | |
| if any(is_anima_model(name) for kind, name in models if kind in ( | |
| "checkpoints", "diffusion_models", | |
| )): | |
| models.append(("clip", ANIMA_CLIP)) | |
| download_assets(list(dict.fromkeys(models))) | |
| def stage_request_models(request): | |
| models = [(model_kind(request.model), request.model)] | |
| models.extend(("loras", lora.name) for lora in request.loras) | |
| models.append(("vae", vae_name(request.model))) | |
| if wants_style(request): | |
| models.extend(( | |
| ("ipadapter", STYLE_IPADAPTER), | |
| ("clip_vision", STYLE_CLIP_VISION), | |
| )) | |
| if request.upscale: | |
| if request.upscale_model: | |
| models.append(("upscale_models", request.upscale_model)) | |
| if request.second_model: | |
| models.append((model_kind(request.second_model), request.second_model)) | |
| models.extend(("loras", lora.name) for lora in request.second_loras) | |
| models.append(("vae", vae_name(request.second_model or request.model))) | |
| for detailer in request.detailers: | |
| models.append(("ultralytics", detailer.detector)) | |
| if detailer.model: | |
| models.extend(( | |
| (model_kind(detailer.model), detailer.model), | |
| ("vae", vae_name(detailer.model)), | |
| )) | |
| stage_models(models) | |
| def request_models_loaded(request): | |
| chains = [chain_key(request.model, request.loras)] | |
| vaes = [vae_name(request.model)] | |
| upscalers = [] | |
| detectors = [] | |
| if request.upscale: | |
| if request.upscale_model: | |
| upscalers.append(request.upscale_model) | |
| if request.second_model: | |
| chains.append(chain_key(request.second_model, request.second_loras)) | |
| vaes.append(vae_name(request.second_model or request.model)) | |
| for detailer in request.detailers: | |
| detectors.append(detailer.detector) | |
| if detailer.model: | |
| chains.append(chain_key(detailer.model, [])) | |
| vaes.append(vae_name(detailer.model)) | |
| return ( | |
| all(key in state["chains"] for key in chains) | |
| and (not wants_style(request) or state["style_pipeline"] is not None) | |
| and all(name in state["vaes"] for name in vaes) | |
| and all(name in state["upscale_models"] for name in upscalers) | |
| and all(name in state["detectors"] for name in detectors) | |
| ) | |
| def unloaded_model_counts(request): | |
| checkpoints = {request.model} | |
| loras = {lora.name for lora in request.loras} | |
| if request.upscale and request.second_model: | |
| checkpoints.add(request.second_model) | |
| loras.update(lora.name for lora in request.second_loras) | |
| checkpoints.update( | |
| detailer.model for detailer in request.detailers if detailer.model | |
| ) | |
| return { | |
| "checkpoints": sum(name not in state["models"] for name in checkpoints), | |
| "loras": sum(name not in state["loras"] for name in loras), | |
| } | |
| def load_request_models(request): | |
| stage_request_models(request) | |
| load_vae(vae_name(request.model)) | |
| load_vae(vae_name(request.second_model or request.model)) | |
| load_chain(request.model, request.loras) | |
| if wants_style(request): | |
| load_style_pipeline() | |
| if request.upscale and request.upscale_model: | |
| load_upscale_model(request.upscale_model) | |
| if request.upscale and request.second_model: | |
| load_chain(request.second_model, request.second_loras) | |
| for detailer in request.detailers: | |
| load_detector(detailer.detector) | |
| if detailer.model: | |
| load_vae(vae_name(detailer.model)) | |
| load_chain(detailer.model, []) | |
| def stored_image_tensor(path): | |
| with Image.open(BytesIO(stored_bytes(path))) as image: | |
| pixels = np.asarray(image.convert("RGB"), dtype=np.float32) / 255 | |
| return torch.from_numpy(pixels).unsqueeze(0) | |
| def tensor_image(image): | |
| return Image.fromarray( | |
| np.clip( | |
| image[0].detach().cpu().numpy() * 255, | |
| 0, | |
| 255, | |
| ).astype(np.uint8) | |
| ) | |
| def generate_comparison_cell(request, latent): | |
| first_name = request.first[1] | |
| second_name = request.second[1] | |
| first_vae = state["vaes"][vae_name(first_name)] | |
| second_vae = state["vaes"][vae_name(second_name)] | |
| first_model, first_clip = state["models"][first_name] | |
| positive = state["encode"].encode( | |
| clip=first_clip, | |
| text=request.positive, | |
| )[0] | |
| negative = state["encode"].encode( | |
| clip=first_clip, | |
| text=request.negative, | |
| )[0] | |
| samples = run_sampler( | |
| first_model, | |
| int(request.generation[1:]), | |
| MATRIX_FIRST_STEPS, | |
| MATRIX_FIRST_CFG, | |
| request.sampler, | |
| request.scheduler, | |
| positive, | |
| negative, | |
| latent, | |
| 1, | |
| ) | |
| samples = state["upscale"].upscale( | |
| samples=samples, | |
| upscale_method=MATRIX_UPSCALE_METHOD, | |
| scale_by=MATRIX_UPSCALE_SCALE, | |
| )[0] | |
| if is_anima_model(first_name) != is_anima_model(second_name): | |
| samples = convert_latent(samples, first_vae, second_vae) | |
| second_model, second_clip = state["models"][second_name] | |
| positive = state["encode"].encode( | |
| clip=second_clip, | |
| text=request.positive, | |
| )[0] | |
| negative = state["encode"].encode( | |
| clip=second_clip, | |
| text=request.negative, | |
| )[0] | |
| samples = run_sampler( | |
| second_model, | |
| int(request.generation[1:]), | |
| MATRIX_SECOND_STEPS, | |
| MATRIX_SECOND_CFG, | |
| request.second_sampler, | |
| request.second_scheduler, | |
| positive, | |
| negative, | |
| samples, | |
| MATRIX_DENOISE, | |
| ) | |
| image = state["decode"].decode( | |
| vae=second_vae, | |
| samples=samples, | |
| )[0] | |
| return image | |
| def infer_matrix_cell(request, pixels=None, latent=None): | |
| with torch.inference_mode(): | |
| first_id, first_name = request.first | |
| second_id, second_name = request.second | |
| first_vae = state["vaes"][vae_name(first_name)] | |
| second_vae = state["vaes"][vae_name(second_name)] | |
| if request.output: | |
| return generate_comparison_cell(request, latent) | |
| if first_id == second_id: | |
| model, clip = state["models"][first_name] | |
| positive = state["encode"].encode( | |
| clip=clip, | |
| text=request.positive, | |
| )[0] | |
| negative = state["encode"].encode( | |
| clip=clip, | |
| text=request.negative, | |
| )[0] | |
| samples = run_sampler( | |
| model, | |
| int(request.generation[1:]), | |
| MATRIX_FIRST_STEPS, | |
| MATRIX_FIRST_CFG, | |
| MATRIX_SAMPLER, | |
| MATRIX_SCHEDULER, | |
| positive, | |
| negative, | |
| latent, | |
| 1, | |
| ) | |
| image = state["decode"].decode( | |
| vae=first_vae, | |
| samples=samples, | |
| )[0] | |
| return image | |
| samples = state["vae_encode"].encode( | |
| vae=first_vae, | |
| pixels=pixels, | |
| )[0] | |
| samples = state["upscale"].upscale( | |
| samples=samples, | |
| upscale_method=MATRIX_UPSCALE_METHOD, | |
| scale_by=MATRIX_UPSCALE_SCALE, | |
| )[0] | |
| if is_anima_model(first_name) != is_anima_model(second_name): | |
| samples = convert_latent(samples, first_vae, second_vae) | |
| model, clip = state["models"][second_name] | |
| positive = state["encode"].encode( | |
| clip=clip, | |
| text=request.positive, | |
| )[0] | |
| negative = state["encode"].encode( | |
| clip=clip, | |
| text=request.negative, | |
| )[0] | |
| samples = run_sampler( | |
| model, | |
| int(request.generation[1:]), | |
| MATRIX_SECOND_STEPS, | |
| MATRIX_SECOND_CFG, | |
| MATRIX_SAMPLER, | |
| MATRIX_SCHEDULER, | |
| positive, | |
| negative, | |
| samples, | |
| MATRIX_DENOISE, | |
| ) | |
| image = state["decode"].decode( | |
| vae=second_vae, | |
| samples=samples, | |
| )[0] | |
| return image | |
| def generate_matrix_cell(body): | |
| request = MatrixCellRequest.model_validate_json(body) | |
| first_id, first_name = request.first | |
| second_id, second_name = request.second | |
| folder = (IMAGE_DIR / request.folder).resolve() | |
| if folder.parent != IMAGE_DIR.resolve(): | |
| raise ValueError("Invalid matrix folder") | |
| with lock: | |
| init_comfy() | |
| models = [ | |
| ("vae", vae_name(first_name)), | |
| ("vae", vae_name(second_name)), | |
| ] | |
| if request.output: | |
| models.extend(( | |
| (model_kind(first_name), first_name), | |
| (model_kind(second_name), second_name), | |
| )) | |
| else: | |
| name = first_name if first_id == second_id else second_name | |
| models.append((model_kind(name), name)) | |
| stage_models(models) | |
| load_vae(vae_name(first_name)) | |
| load_vae(vae_name(second_name)) | |
| if request.output: | |
| output = (folder / request.output).resolve() | |
| if output.parent != folder: | |
| raise ValueError("Invalid matrix output") | |
| if output.is_file(): | |
| return request.output | |
| load_model(first_name) | |
| load_model(second_name) | |
| latent = state["latent"].generate( | |
| width=MATRIX_WIDTH, | |
| height=MATRIX_HEIGHT, | |
| batch_size=1, | |
| )[0] | |
| image = infer_matrix_cell(request, latent=latent) | |
| result = request.output | |
| else: | |
| output = matrix_image_path( | |
| request.folder, | |
| first_id, | |
| second_id, | |
| request.generation, | |
| ) | |
| if output.is_file(): | |
| return ( | |
| first_id | |
| if first_id == second_id | |
| else f"{first_id}-{second_id}" | |
| ) | |
| if first_id == second_id: | |
| load_model(first_name) | |
| latent = state["latent"].generate( | |
| width=MATRIX_WIDTH, | |
| height=MATRIX_HEIGHT, | |
| batch_size=1, | |
| )[0] | |
| image = infer_matrix_cell(request, latent=latent) | |
| result = first_id | |
| else: | |
| diagonal = matrix_image_path( | |
| request.folder, | |
| first_id, | |
| first_id, | |
| request.generation, | |
| ) | |
| pixels = stored_image_tensor(diagonal) | |
| load_model(second_name) | |
| image = infer_matrix_cell(request, pixels) | |
| result = f"{first_id}-{second_id}" | |
| save_named_image(tensor_image(image), output) | |
| return result | |
| def model_options(kind, local): | |
| local = [local] if isinstance(local, str) else local | |
| names = set(local) | remote_models[kind].keys() | |
| return sorted( | |
| (name for name in names if Path(name).suffix.casefold() in MODEL_SUFFIXES), | |
| key=str.casefold, | |
| ) | |
| def generate_images(request, from_api=True): | |
| resolve_artist_prompts(request) | |
| if ( | |
| request.width % 8 | |
| or request.height % 8 | |
| ): | |
| raise ValueError("Width and height must be multiples of 8") | |
| if request.regional_mode not in REGIONAL_MODES: | |
| raise ValueError("Unsupported regional mode") | |
| if request.style_images and request.style_scope not in STYLE_SCOPES: | |
| raise ValueError("Unsupported InstantStyle scope") | |
| if has_style(request): | |
| style_models = [request.model] if request.style_images else [] | |
| if request.upscale and ( | |
| request.second_style_images | |
| or ( | |
| request.style_images and request.style_scope != "first" | |
| ) | |
| ): | |
| style_models.append(request.second_model or request.model) | |
| if request.style_images and request.style_scope == "all": | |
| final_model = ( | |
| request.second_model | |
| if request.upscale and request.second_model | |
| else request.model | |
| ) | |
| style_models.extend( | |
| detailer.model or final_model | |
| for detailer in request.detailers | |
| ) | |
| if any(is_anima_model(name) for name in style_models): | |
| raise ValueError("InstantStyle only supports SDXL models") | |
| regions = prepare_regions(request.regions) | |
| style_image = ( | |
| decode_style_images(request.style_images) | |
| if request.style_images | |
| else None | |
| ) | |
| second_style_image = ( | |
| decode_style_images(request.second_style_images) | |
| if request.second_style_images | |
| else style_image | |
| ) | |
| with lock: | |
| init_comfy() | |
| detailer_options = state["face_detailer"].INPUT_TYPES()["required"] | |
| samplers = [request.sampler] | |
| schedulers = [request.scheduler] | |
| if request.upscale: | |
| samplers.append(request.second_sampler) | |
| schedulers.append(request.second_scheduler) | |
| if any(value not in sampler_names() for value in samplers): | |
| raise ValueError("Unsupported sampler or scheduler") | |
| if any(value not in scheduler_names() for value in schedulers): | |
| raise ValueError("Unsupported sampler or scheduler") | |
| if any( | |
| detailer.sampler not in detailer_options["sampler_name"][0] | |
| or detailer.scheduler not in detailer_options["scheduler"][0] | |
| for detailer in request.detailers | |
| ): | |
| raise ValueError("Unsupported detailer sampler or scheduler") | |
| first_vae = vae_name(request.model) | |
| second_vae = vae_name(request.second_model or request.model) | |
| load_request_models(request) | |
| image, first_seed, second_seed, detailer_seeds = generate_gpu( | |
| request, | |
| regions, | |
| style_image, | |
| second_style_image, | |
| ) | |
| metadata_config = request.model_dump(exclude={"sillytavern"}) | |
| if request.style_images: | |
| metadata_config["style_images"] = [ | |
| f"style-reference-{index}.png" | |
| for index in range(1, len(request.style_images) + 1) | |
| ] | |
| if request.second_style_images: | |
| metadata_config["second_style_images"] = [ | |
| f"second-style-reference-{index}.png" | |
| for index in range(1, len(request.second_style_images) + 1) | |
| ] | |
| metadata = image_metadata( | |
| metadata_config, | |
| [first_seed, second_seed], | |
| detailer_seeds, | |
| [ | |
| vae_name(detailer.model) | |
| if detailer.model | |
| else second_vae if request.upscale else first_vae | |
| for detailer in request.detailers | |
| ], | |
| [first_vae, second_vae], | |
| regions, | |
| ENVIRONMENT_START, | |
| REGIONAL_GLOBAL_STRENGTH, | |
| ) | |
| if request.sillytavern: | |
| metadata["sillytavern"] = json.dumps( | |
| request.sillytavern, | |
| separators=(",", ":"), | |
| ) | |
| results = [ | |
| Image.fromarray( | |
| np.clip(item.detach().numpy() * 255, 0, 255).astype(np.uint8) | |
| ) | |
| for item in image | |
| ] | |
| for result in results: | |
| result.info.update(metadata) | |
| backup_pool.submit(archive_image, result, from_api) | |
| return results | |
| def combine_images(images, width, height): | |
| if len(images) == 1: | |
| return images[0] | |
| vertical = width > height | |
| image_width, image_height = images[0].size | |
| size = ( | |
| (image_width, image_height * len(images)) | |
| if vertical | |
| else (image_width * len(images), image_height) | |
| ) | |
| combined = Image.new(images[0].mode, size) | |
| for index, image in enumerate(images): | |
| combined.paste( | |
| image, | |
| (0, index * image_height) if vertical else (index * image_width, 0), | |
| ) | |
| combined.info.update(images[0].info) | |
| return combined | |
| def scale_image(image, scale): | |
| if scale == 1: | |
| return image | |
| return image.resize( | |
| (round(image.width * scale), round(image.height * scale)), | |
| Image.Resampling.LANCZOS, | |
| ) | |
| def archive_image(image, from_api): | |
| try: | |
| if HAS_DATA_MOUNT: | |
| save_image(image, from_api) | |
| else: | |
| upload_image(image) | |
| except Exception as error: | |
| log(f"Archive failed: {error}") | |
| def image_fingerprint(image): | |
| pixels = np.asarray( | |
| ImageOps.exif_transpose(image).convert("RGB").resize( | |
| (DUPLICATE_HASH_SIZE + 1, DUPLICATE_HASH_SIZE), | |
| Image.Resampling.LANCZOS, | |
| ) | |
| ) | |
| differences = pixels[:, 1:] > pixels[:, :-1] | |
| return ( | |
| int.from_bytes(np.packbits(differences).tobytes()), | |
| tuple(int(value) for value in pixels.mean(axis=(0, 1))), | |
| ) | |
| def image_png_bytes(image): | |
| pnginfo = PngInfo() | |
| for key, value in image.info.items(): | |
| if isinstance(value, str): | |
| pnginfo.add_text(key, value) | |
| return png_bytes(image, pnginfo) | |
| def image_key(salt): | |
| return Scrypt( | |
| salt=salt, | |
| length=32, | |
| n=SCRYPT_N, | |
| r=8, | |
| p=1, | |
| ).derive(PASSWORD.encode()) | |
| def proxy_key(salt): | |
| return Scrypt( | |
| salt=salt, | |
| length=32, | |
| n=SCRYPT_N, | |
| r=8, | |
| p=1, | |
| ).derive(PASSWORD.encode()) | |
| def encrypt_proxy_payload(data): | |
| salt = os.urandom(SALT_SIZE) | |
| nonce = os.urandom(NONCE_SIZE) | |
| return ( | |
| PROXY_MAGIC | |
| + salt | |
| + nonce | |
| + AESGCM(proxy_key(salt)).encrypt(nonce, data, PROXY_MAGIC) | |
| ) | |
| def decrypt_proxy_payload(data): | |
| if len(data) < len(PROXY_MAGIC) + SALT_SIZE + NONCE_SIZE + 16: | |
| raise ValueError("invalid encrypted payload") | |
| if not data.startswith(PROXY_MAGIC): | |
| raise ValueError("invalid encrypted payload") | |
| salt_start = len(PROXY_MAGIC) | |
| nonce_start = salt_start + SALT_SIZE | |
| data_start = nonce_start + NONCE_SIZE | |
| return AESGCM(proxy_key(data[salt_start:nonce_start])).decrypt( | |
| data[nonce_start:data_start], | |
| data[data_start:], | |
| PROXY_MAGIC, | |
| ) | |
| def save_image(image, from_api): | |
| path = IMAGE_DIR / datetime.now(TIMEZONE).date().isoformat() | |
| path.mkdir(parents=True, exist_ok=True) | |
| suffix = "ST.epng" if from_api else "C.epng" | |
| number = max( | |
| ( | |
| int(file.name.removesuffix(suffix)) | |
| for file in path.iterdir() | |
| if file.name.endswith(suffix) | |
| and file.name.removesuffix(suffix).isdigit() | |
| ), | |
| default=0, | |
| ) + 1 | |
| save_named_image(image, path / f"{number}{suffix}") | |
| def upload_image(image): | |
| date = datetime.now(TIMEZONE).date().isoformat() | |
| folder = f"{IMAGE_BUCKET_PREFIX}/{date}" | |
| suffix = "ST.epng" | |
| number = max( | |
| ( | |
| int(Path(item.path).name.removesuffix(suffix)) | |
| for item in list_bucket_tree( | |
| IMAGE_BUCKET_ID, | |
| prefix=f"{folder}/", | |
| recursive=True, | |
| token=IMAGE_TOKEN, | |
| ) | |
| if item.type == "file" | |
| and Path(item.path).name.endswith(suffix) | |
| and Path(item.path).name.removesuffix(suffix).isdigit() | |
| ), | |
| default=0, | |
| ) + 1 | |
| batch_bucket_files( | |
| IMAGE_BUCKET_ID, | |
| add=[(encrypted_image_bytes(image), f"{folder}/{number}{suffix}")], | |
| token=IMAGE_TOKEN, | |
| ) | |
| def encrypted_image_bytes(image): | |
| return encrypt_image_data(image_png_bytes(image)) | |
| def encrypt_image_data(data, salt=None): | |
| salt = salt if salt is not None else os.urandom(SALT_SIZE) | |
| nonce = os.urandom(NONCE_SIZE) | |
| encrypted = AESGCM(image_key(salt)).encrypt(nonce, data, FILE_MAGIC) | |
| return FILE_MAGIC + salt + nonce + encrypted | |
| def save_named_image(image, output): | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| temp = output.with_suffix(".epng.part") | |
| temp.write_bytes(encrypted_image_bytes(image)) | |
| temp.replace(output) | |
| log(f"Saved img/{output.relative_to(IMAGE_DIR)}") | |
| def png_bytes(image, pnginfo=None): | |
| data = BytesIO() | |
| image.save(data, format="PNG", pnginfo=pnginfo) | |
| return data.getvalue() | |
| def stored_path(value): | |
| root = IMAGE_DIR.resolve() | |
| path = (root / value).resolve() | |
| try: | |
| path.relative_to(root) | |
| except ValueError as error: | |
| raise HTTPException(404, "image not found") from error | |
| if not path.is_file() or path.suffix.casefold() not in IMAGE_SUFFIXES: | |
| raise HTTPException(404, "image not found") | |
| return path | |
| def natural_key(path): | |
| return [ | |
| int(part) if part.isdigit() else part.casefold() | |
| for part in re.split(r"(\d+)", path.name) | |
| ] | |
| def star_database(): | |
| STAR_DB.parent.mkdir(parents=True, exist_ok=True) | |
| database = sqlite3.connect(STAR_DB, timeout=EXPLORER_DB_TIMEOUT) | |
| database.execute("PRAGMA journal_mode=WAL") | |
| database.execute( | |
| "CREATE TABLE IF NOT EXISTS stars (path TEXT PRIMARY KEY)" | |
| ) | |
| database.execute( | |
| """CREATE TABLE IF NOT EXISTS image_prompts ( | |
| path TEXT PRIMARY KEY, | |
| modified INTEGER NOT NULL, | |
| size INTEGER NOT NULL, | |
| prompt TEXT NOT NULL, | |
| second_prompt TEXT NOT NULL, | |
| artists TEXT NOT NULL | |
| )""" | |
| ) | |
| database.execute( | |
| """CREATE TABLE IF NOT EXISTS image_previews ( | |
| path TEXT PRIMARY KEY, | |
| modified INTEGER NOT NULL, | |
| size INTEGER NOT NULL, | |
| data BLOB NOT NULL | |
| )""" | |
| ) | |
| return database | |
| def stored_prompt(path, relative, database): | |
| stat = path.stat() | |
| cached = database.execute( | |
| """SELECT prompt, second_prompt, artists | |
| FROM image_prompts | |
| WHERE path = ? AND modified = ? AND size = ?""", | |
| (relative, stat.st_mtime_ns, stat.st_size), | |
| ).fetchone() | |
| if cached is not None: | |
| return cached[0], cached[1], json.loads(cached[2]) | |
| prompt = "" | |
| second_prompt = "" | |
| artists = [] | |
| try: | |
| with Image.open(BytesIO(stored_bytes(path))) as image: | |
| parameters = json.loads(image.info.get("parameters", "{}")) | |
| prompt = str(parameters.get("prompt", "")) | |
| second_prompt = str(parameters.get("second_prompt") or (prompt if parameters.get("upscale") else "")) | |
| selected = parameters.get("selected_artists", []) | |
| if isinstance(selected, list): | |
| artists = list(dict.fromkeys( | |
| artist for artist in selected if isinstance(artist, str) | |
| )) | |
| except (json.JSONDecodeError, OSError, TypeError, ValueError): | |
| pass | |
| database.execute( | |
| """INSERT OR REPLACE INTO image_prompts | |
| (path, modified, size, prompt, second_prompt, artists) | |
| VALUES (?, ?, ?, ?, ?, ?)""", | |
| ( | |
| relative, | |
| stat.st_mtime_ns, | |
| stat.st_size, | |
| prompt, | |
| second_prompt, | |
| json.dumps(artists, separators=(",", ":")), | |
| ), | |
| ) | |
| return prompt, second_prompt, artists | |
| def starred_paths(): | |
| database = star_database() | |
| try: | |
| return {row[0] for row in database.execute("SELECT path FROM stars")} | |
| finally: | |
| database.close() | |
| def set_star(path, starred): | |
| database = star_database() | |
| try: | |
| if starred: | |
| database.execute("INSERT OR IGNORE INTO stars VALUES (?)", (path,)) | |
| else: | |
| database.execute("DELETE FROM stars WHERE path = ?", (path,)) | |
| database.commit() | |
| finally: | |
| database.close() | |
| def stored_bytes(path): | |
| return decrypt_image_data(path.read_bytes()) | |
| def decrypt_image_data(data): | |
| if not data.startswith(FILE_MAGIC): | |
| raise HTTPException(500, "invalid encrypted image") | |
| salt_start = len(FILE_MAGIC) | |
| nonce_start = salt_start + SALT_SIZE | |
| data_start = nonce_start + NONCE_SIZE | |
| return AESGCM(image_key(data[salt_start:nonce_start])).decrypt( | |
| data[nonce_start:data_start], | |
| data[data_start:], | |
| FILE_MAGIC, | |
| ) | |
| def stored_preview(path, modified, size): | |
| relative = path.relative_to(IMAGE_DIR.resolve()).as_posix() | |
| database = star_database() | |
| try: | |
| cached = database.execute( | |
| """SELECT data FROM image_previews | |
| WHERE path = ? AND modified = ? AND size = ?""", | |
| (relative, modified, size), | |
| ).fetchone() | |
| if cached is not None: | |
| return decrypt_image_data(cached[0]) | |
| source = path.read_bytes() | |
| with Image.open(BytesIO(decrypt_image_data(source))) as image: | |
| image.thumbnail((PREVIEW_SIZE, PREVIEW_SIZE)) | |
| data = BytesIO() | |
| image.save( | |
| data, | |
| format="WEBP", | |
| quality=PREVIEW_QUALITY, | |
| method=0, | |
| ) | |
| preview = data.getvalue() | |
| salt = source[len(FILE_MAGIC):len(FILE_MAGIC) + SALT_SIZE] | |
| database.execute( | |
| "INSERT OR REPLACE INTO image_previews VALUES (?, ?, ?, ?)", | |
| (relative, modified, size, encrypt_image_data(preview, salt)), | |
| ) | |
| database.commit() | |
| return preview | |
| finally: | |
| database.close() | |
| def reserve_matrix_grid(folder): | |
| path = IMAGE_DIR / folder | |
| path.mkdir(parents=True, exist_ok=True) | |
| with matrix_lock: | |
| number = 1 | |
| while ( | |
| (path / f"{number}gr.epng").exists() | |
| or (folder, number) in matrix_grids | |
| ): | |
| number += 1 | |
| matrix_grids.add((folder, number)) | |
| return number | |
| def matrix_models(): | |
| with lock: | |
| init_comfy() | |
| inputs = state["checkpoint"].INPUT_TYPES()["required"] | |
| names = model_options("checkpoints", inputs["ckpt_name"][0]) | |
| models = [] | |
| ids = set() | |
| for name in names: | |
| match = re.match(r"(\d+)_", name) | |
| if match is None: | |
| raise ValueError(f"Checkpoint has no numeric prefix: {name}") | |
| model_id = str(int(match.group(1))) | |
| if model_id in ids: | |
| raise ValueError(f"Duplicate checkpoint prefix: {model_id}") | |
| ids.add(model_id) | |
| models.append((model_id, name)) | |
| return sorted(models, key=lambda item: int(item[0])) | |
| def matrix_options(): | |
| with lock: | |
| init_comfy() | |
| options = state["sample"].INPUT_TYPES()["required"] | |
| samplers = list(options["sampler_name"][0]) | |
| schedulers = [*options["scheduler"][0], ALIGN_SCHEDULER] | |
| return samplers, schedulers | |
| def matrix_model(models, model_id): | |
| model_id = str(model_id) | |
| for model in models: | |
| if model[0] == model_id: | |
| return model | |
| raise ValueError(f"Unknown checkpoint number: {model_id}") | |
| def comparison_plan(request): | |
| models = matrix_models() | |
| first = matrix_model(models, request.model_1) | |
| second = matrix_model(models, request.model_2) | |
| samplers, schedulers = matrix_options() | |
| if request.type == "sampler": | |
| return first, second, [""], samplers | |
| if request.type == "scheduler": | |
| sampler = request.sampler or MATRIX_SAMPLER | |
| if sampler not in samplers: | |
| raise ValueError(f"Unsupported sampler: {sampler}") | |
| return first, second, [""], schedulers | |
| return first, second, schedulers, samplers | |
| def combined_plan(): | |
| models = matrix_models() | |
| samplers, schedulers = matrix_options() | |
| options = [ | |
| (sampler, scheduler) | |
| for sampler in samplers | |
| for scheduler in schedulers | |
| ] | |
| rows = [ | |
| (first, second) | |
| for first in options | |
| for second in options | |
| ] | |
| columns = [] | |
| for index, model in enumerate(models): | |
| columns.extend((first, model) for first in models[:index]) | |
| columns.extend( | |
| (model, second) | |
| for second in reversed(models[:index]) | |
| ) | |
| return rows, columns | |
| def matrix_image_path(folder, first_id, second_id, generation): | |
| name = ( | |
| f"{first_id}{generation}.epng" | |
| if first_id == second_id | |
| else f"{first_id}x{second_id}{generation}.epng" | |
| ) | |
| return IMAGE_DIR / folder / name | |
| def create_matrix_grid(folder, number, models, generation): | |
| count = len(models) | |
| width = MATRIX_LABEL_SIZE + MATRIX_CELL_WIDTH * count | |
| height = MATRIX_LABEL_SIZE + MATRIX_CELL_HEIGHT * count | |
| grid = Image.new("RGB", (width, height), "#111") | |
| draw = ImageDraw.Draw(grid) | |
| font = ImageFont.load_default(size=18) | |
| for index, (model_id, _) in enumerate(models): | |
| x = MATRIX_LABEL_SIZE + index * MATRIX_CELL_WIDTH | |
| y = MATRIX_LABEL_SIZE + index * MATRIX_CELL_HEIGHT | |
| draw.text( | |
| (x + MATRIX_CELL_WIDTH // 2, MATRIX_LABEL_SIZE // 2), | |
| model_id, | |
| fill="#6cf", | |
| font=font, | |
| anchor="mm", | |
| ) | |
| draw.text( | |
| (MATRIX_LABEL_SIZE // 2, y + MATRIX_CELL_HEIGHT // 2), | |
| model_id, | |
| fill="#6cf", | |
| font=font, | |
| anchor="mm", | |
| ) | |
| for column, (second_id, _) in enumerate(models): | |
| path = matrix_image_path( | |
| folder, | |
| model_id, | |
| second_id, | |
| generation, | |
| ) | |
| with Image.open(BytesIO(stored_bytes(path))) as source: | |
| image = source.convert("RGB") | |
| cell_x = MATRIX_LABEL_SIZE + column * MATRIX_CELL_WIDTH | |
| grid.paste( | |
| image, | |
| ( | |
| cell_x + (MATRIX_CELL_WIDTH - image.width) // 2, | |
| y + (MATRIX_CELL_HEIGHT - image.height) // 2, | |
| ), | |
| ) | |
| output = IMAGE_DIR / folder / f"{number}gr.epng" | |
| save_named_image(grid, output) | |
| def comparison_image_path(folder, number, row, column): | |
| return IMAGE_DIR / folder / f"{number}-{row + 1}x{column + 1}.epng" | |
| def create_comparison_grid(folder, number, rows, columns): | |
| font = ImageFont.load_default(size=18) | |
| label_width = ( | |
| max( | |
| MATRIX_ROW_LABEL_WIDTH, | |
| *( | |
| round(font.getlength(label)) + MATRIX_LABEL_SIZE | |
| for label in rows | |
| ), | |
| ) | |
| if len(rows) > 1 | |
| else 0 | |
| ) | |
| width = label_width + MATRIX_GRID_CELL_WIDTH * len(columns) | |
| height = MATRIX_LABEL_SIZE + MATRIX_GRID_CELL_HEIGHT * len(rows) | |
| grid = Image.new("RGB", (width, height), "#111") | |
| draw = ImageDraw.Draw(grid) | |
| for column, label in enumerate(columns): | |
| draw.text( | |
| ( | |
| label_width + column * MATRIX_GRID_CELL_WIDTH | |
| + MATRIX_GRID_CELL_WIDTH // 2, | |
| MATRIX_LABEL_SIZE // 2, | |
| ), | |
| label, | |
| fill="#6cf", | |
| font=font, | |
| anchor="mm", | |
| ) | |
| for row, label in enumerate(rows): | |
| y = MATRIX_LABEL_SIZE + row * MATRIX_GRID_CELL_HEIGHT | |
| if label_width: | |
| draw.text( | |
| (label_width // 2, y + MATRIX_GRID_CELL_HEIGHT // 2), | |
| label, | |
| fill="#6cf", | |
| font=font, | |
| anchor="mm", | |
| ) | |
| for column in range(len(columns)): | |
| path = comparison_image_path(folder, number, row, column) | |
| with Image.open(BytesIO(stored_bytes(path))) as source: | |
| image = source.convert("RGB") | |
| image.thumbnail( | |
| (MATRIX_GRID_CELL_WIDTH, MATRIX_GRID_CELL_HEIGHT), | |
| ) | |
| x = label_width + column * MATRIX_GRID_CELL_WIDTH | |
| grid.paste( | |
| image, | |
| ( | |
| x + (MATRIX_GRID_CELL_WIDTH - image.width) // 2, | |
| y + (MATRIX_GRID_CELL_HEIGHT - image.height) // 2, | |
| ), | |
| ) | |
| save_named_image(grid, IMAGE_DIR / folder / f"{number}gr.epng") | |
| def selected_paths(items): | |
| root = IMAGE_DIR.resolve() | |
| selected = {} | |
| for value in items: | |
| path = (root / value).resolve() | |
| try: | |
| path.relative_to(root) | |
| except ValueError as error: | |
| raise HTTPException(404, "image not found") from error | |
| paths = path.iterdir() if path.is_dir() else (stored_path(value),) | |
| for image in paths: | |
| if image.is_file() and image.suffix.casefold() in IMAGE_SUFFIXES: | |
| name = image.relative_to(root).with_suffix(".png").as_posix() | |
| selected[name] = image | |
| if not selected: | |
| raise HTTPException(422, "no images selected") | |
| return selected | |
| def valid_pass(p): | |
| return secrets.compare_digest(str(p), PASSWORD) | |
| def clean_regions(rows): | |
| return [ | |
| RegionRequest( | |
| prompt=str(row[0]).strip(), | |
| area=str(row[1] or "full").strip(), | |
| strength=row[2] if len(row) > 2 and row[2] is not None else 1, | |
| ) | |
| for row in (rows or [])[:4] | |
| if row and str(row[0] or '').strip() | |
| ] | |
| def clean_detailers(rows): | |
| return [ | |
| DetailerRequest( | |
| detector=str(row[0]), | |
| model=str(row[1] or ""), | |
| prompt=str(row[2] or ""), | |
| negative=str(row[3] or ""), | |
| sampler=str(row[4]), | |
| scheduler=str(row[5]), | |
| steps=row[6], | |
| cfg=row[7], | |
| denoise=row[8], | |
| ) | |
| for row in rows or [] | |
| if row and row[0] | |
| ] | |
| def select_model(model, current): | |
| model = current if model in MODEL_HEADERS else model | |
| return model, model | |
| def add_detailer( | |
| rows, | |
| detector, | |
| model, | |
| prompt, | |
| negative, | |
| sampler, | |
| scheduler, | |
| steps, | |
| cfg, | |
| denoise, | |
| ): | |
| if not detector: | |
| return rows | |
| return [ | |
| *(rows or []), | |
| [ | |
| detector, model, prompt, negative, sampler, scheduler, | |
| steps, cfg, denoise, | |
| ], | |
| ] | |
| def encode_style_images(files): | |
| files = files or [] | |
| if len(files) > MAX_STYLE_IMAGES: | |
| raise gr.Error("InstantStyle accepts up to 4 images") | |
| images = [] | |
| for file in files: | |
| data = Path(file).read_bytes() | |
| if len(data) > MAX_STYLE_IMAGE_SIZE: | |
| raise gr.Error("InstantStyle image exceeds 20 MiB") | |
| images.append(base64.b64encode(data).decode()) | |
| return images | |
| def generate( | |
| prompt, | |
| negative, | |
| regions, | |
| regional_mode, | |
| model, | |
| style_images, | |
| style_scope, | |
| style_weight, | |
| style_end, | |
| detailers, | |
| width=1152, | |
| height=896, | |
| batch_size=DEFAULT_UI_BATCH_SIZE, | |
| sampler=DEFAULT_SAMPLER, | |
| scheduler=DEFAULT_SCHEDULER, | |
| steps=DEFAULT_STEPS, | |
| cfg=DEFAULT_CFG, | |
| upscale=False, | |
| upscale_method=DEFAULT_UPSCALE_METHOD, | |
| upscale_model=DEFAULT_UPSCALE_MODEL, | |
| upscale_scale=DEFAULT_UPSCALE_SCALE, | |
| second_model="", | |
| second_sampler=DEFAULT_SECOND_SAMPLER, | |
| second_scheduler=DEFAULT_SECOND_SCHEDULER, | |
| second_steps=DEFAULT_SECOND_STEPS, | |
| second_cfg=DEFAULT_SECOND_CFG, | |
| denoise=DEFAULT_DENOISE, | |
| return_scale=DEFAULT_RETURN_SCALE, | |
| p="", | |
| ): | |
| if not valid_pass(p): | |
| raise gr.Error("Invalid password") | |
| request = DirectRequest( | |
| prompt=prompt, | |
| regions=clean_regions(regions), | |
| regional_mode=regional_mode, | |
| negative=negative, | |
| model=model, | |
| loras=[], | |
| style_images=encode_style_images(style_images), | |
| style_scope=style_scope, | |
| style_weight=style_weight, | |
| style_end=style_end, | |
| detailers=clean_detailers(detailers), | |
| width=width, | |
| height=height, | |
| batch_size=batch_size, | |
| sampler=sampler, | |
| scheduler=scheduler, | |
| steps=steps, | |
| cfg=cfg, | |
| upscale=upscale, | |
| upscale_method=upscale_method, | |
| upscale_model=upscale_model, | |
| upscale_scale=upscale_scale, | |
| second_model=second_model, | |
| second_loras=[], | |
| second_sampler=second_sampler, | |
| second_scheduler=second_scheduler, | |
| second_steps=second_steps, | |
| second_cfg=second_cfg, | |
| denoise=denoise, | |
| return_scale=return_scale, | |
| ) | |
| return [ | |
| scale_image(image, return_scale) | |
| for image in generate_images(request, False) | |
| ] | |
| def ping_image(p=""): | |
| if not valid_pass(p): | |
| raise gr.Error("Invalid password") | |
| with lock: | |
| prefix = f"{PING_MODEL_ID}_" | |
| models = [name for name in generation_models() if name.startswith(prefix)] | |
| if len(models) != 1: | |
| raise gr.Error(f"Expected one checkpoint with prefix {prefix}") | |
| model_name = models[0] | |
| model, clip = load_chain(model_name, []) | |
| vae = load_vae(vae_name(model_name)) | |
| clip.patcher.load_device = CPU_DEVICE | |
| clip.patcher.offload_device = CPU_DEVICE | |
| latent = state["latent"].generate( | |
| width=PING_SIZE, | |
| height=PING_SIZE, | |
| batch_size=1, | |
| )[0] | |
| positive = state["encode"].encode(clip=clip, text="1girl")[0] | |
| negative = state["encode"].encode(clip=clip, text="")[0] | |
| state["ping"] = model, vae, latent, positive, negative | |
| try: | |
| image = infer_ping(secrets.randbits(64)) | |
| finally: | |
| del state["ping"] | |
| return [tensor_image(image)] | |
| def api_generate(body, p=""): | |
| if not valid_pass(p): | |
| raise gr.Error("Invalid password") | |
| request = DirectRequest.model_validate_json(body) | |
| image = combine_images( | |
| generate_images(request), | |
| request.width, | |
| request.height, | |
| ) | |
| temp = tempfile.NamedTemporaryFile(suffix=".epng", delete=False) | |
| temp.write(encrypted_image_bytes(image)) | |
| temp.close() | |
| return temp.name | |
| def api_upscale(body, p=""): | |
| if not valid_pass(p): | |
| raise gr.Error("Invalid password") | |
| model_name = DEFAULT_UPSCALE_MODEL | |
| scale = 1 | |
| try: | |
| payload = json.loads(body) | |
| if isinstance(payload, dict): | |
| body = payload["image"] | |
| model_name = payload.get("model", model_name) | |
| scale = float(payload.get("scale", scale)) | |
| except json.JSONDecodeError: | |
| pass | |
| except (KeyError, TypeError, ValueError) as error: | |
| raise gr.Error("Invalid upscale request") from error | |
| if not 1 <= scale <= 4: | |
| raise gr.Error("Scale must be between 1 and 4") | |
| if len(body) > (MAX_UPSCALE_BYTES + 2) // 3 * 4: | |
| raise gr.Error("Image exceeds 20 MiB") | |
| try: | |
| data = base64.b64decode(body, validate=True) | |
| if len(data) > MAX_UPSCALE_BYTES: | |
| raise ValueError("Image exceeds 20 MiB") | |
| with Image.open(BytesIO(data)) as source: | |
| if source.width * source.height > MAX_UPSCALE_PIXELS: | |
| raise ValueError("Image exceeds 4 megapixels") | |
| if getattr(source, "is_animated", False): | |
| raise ValueError("Animated images are not supported") | |
| image = ImageOps.exif_transpose(source).convert("RGBA") | |
| except (ValueError, OSError, Image.DecompressionBombError) as error: | |
| raise gr.Error(str(error)) from error | |
| pixels = torch.from_numpy( | |
| np.asarray(image.convert("RGB"), dtype=np.float32) / 255 | |
| ).unsqueeze(0) | |
| with lock: | |
| init_comfy() | |
| if model_name not in model_options("upscale_models", []): | |
| raise gr.Error(f"Unknown upscale model: {model_name}") | |
| model = load_upscale_model(model_name) | |
| result = tensor_image(infer_upscale(pixels, model)) | |
| size = round(image.width * scale), round(image.height * scale) | |
| if result.size != size: | |
| result = result.resize(size, Image.Resampling.LANCZOS) | |
| alpha = image.getchannel("A") | |
| if alpha.getextrema() != (255, 255): | |
| result.putalpha(alpha.resize(size, Image.Resampling.LANCZOS)) | |
| result.info.update(image.info) | |
| return base64.b64encode(image_png_bytes(result)).decode() | |
| def api_health(p=""): | |
| if not valid_pass(p): | |
| raise gr.Error("Invalid password") | |
| return {"status": True} | |
| def grouped_options(defaults, values, sources=None): | |
| defaults = set(defaults) | |
| sources = sources or {} | |
| groups = [] | |
| standard = [ | |
| {"value": value, "text": value} | |
| for value in values | |
| if value in defaults | |
| ] | |
| custom = [] | |
| for value in values: | |
| if value in defaults: | |
| continue | |
| source = sources.get(value) | |
| text = value | |
| if value in state.get("custom_samplers", {}): | |
| text = value.partition(":")[2] | |
| source = "ComfyUI-ppm" | |
| custom.append({ | |
| "value": value, | |
| "text": f"{text} [{source or 'Custom node'}]", | |
| }) | |
| if standard: | |
| groups.append({"label": "Default", "options": standard}) | |
| if custom: | |
| groups.append({"label": "Custom", "options": custom}) | |
| return groups | |
| def api_options(p=""): | |
| if not valid_pass(p): | |
| raise gr.Error("Invalid password") | |
| data = object_info() | |
| detailer = state["face_detailer"].INPUT_TYPES()["required"] | |
| samplers = sampler_names() | |
| schedulers = scheduler_names() | |
| scheduler_sources = { | |
| **state["scheduler_sources"], | |
| ALIGN_SCHEDULER: "ComfyUI", | |
| } | |
| return { | |
| "models": data["CheckpointLoaderSimple"]["input"]["required"][ | |
| "ckpt_name" | |
| ][0], | |
| "loras": data["LoraLoader"]["input"]["required"]["lora_name"][0], | |
| "samplers": grouped_options( | |
| state["default_samplers"], | |
| samplers, | |
| state["sampler_sources"], | |
| ), | |
| "schedulers": grouped_options( | |
| state["default_schedulers"], | |
| schedulers, | |
| scheduler_sources, | |
| ), | |
| "detailer-samplers": grouped_options( | |
| state["default_samplers"], | |
| detailer["sampler_name"][0], | |
| state["sampler_sources"], | |
| ), | |
| "detailer-schedulers": grouped_options( | |
| state["default_schedulers"], | |
| detailer["scheduler"][0], | |
| state["scheduler_sources"], | |
| ), | |
| "instant-style-scopes": list(STYLE_SCOPES), | |
| "upscale-methods": grouped_options( | |
| state["default_upscale_methods"], | |
| data["LatentUpscaleBy"]["input"]["required"][ | |
| "upscale_method" | |
| ][0], | |
| ), | |
| "upscale-models": data["UpscaleModelLoader"]["input"]["required"][ | |
| "model_name" | |
| ][0], | |
| "ultralytics": data["UltralyticsDetectorProvider"]["input"][ | |
| "required" | |
| ]["model_name"][0], | |
| } | |
| def api_refresh(p=""): | |
| if not valid_pass(p): | |
| raise gr.Error("Invalid password") | |
| return refresh_models() | |
| def api_load_models(body, p=""): | |
| if not valid_pass(p): | |
| raise gr.Error("Invalid password") | |
| request = ModelRequest.model_validate_json(body) | |
| with lock: | |
| init_comfy() | |
| unloaded = unloaded_model_counts(request) | |
| changed = not request_models_loaded(request) | |
| if changed: | |
| model_pool.submit(load_models, request) | |
| return {"loaded": not changed, "changed": changed, **unloaded} | |
| def load_models(request): | |
| with lock: | |
| if not request_models_loaded(request): | |
| load_request_models(request) | |
| def api_matrix_cell(body, p=""): | |
| if not valid_pass(p): | |
| raise gr.Error("Invalid password") | |
| return generate_matrix_cell(body) | |
| def queued_generate(request): | |
| result = get_local_client().predict( | |
| request.model_dump_json(), | |
| PASSWORD, | |
| api_name="/generate", | |
| ) | |
| return Image.open(BytesIO(stored_bytes(Path(result)))).copy() | |
| def queued_matrix_cell(request): | |
| return retry_gpu( | |
| lambda: get_local_client().predict( | |
| request.model_dump_json(), | |
| PASSWORD, | |
| api_name="/matrix_cell", | |
| ) | |
| ) | |
| def run_checkpoint_matrix(request, folder, number): | |
| models = matrix_models() | |
| for index, model in enumerate(models): | |
| queued_matrix_cell( | |
| MatrixCellRequest( | |
| generation=request.generation, | |
| positive=request.positive, | |
| negative=request.negative, | |
| folder=folder, | |
| first=model, | |
| second=model, | |
| ) | |
| ) | |
| for first in models[:index]: | |
| queued_matrix_cell( | |
| MatrixCellRequest( | |
| generation=request.generation, | |
| positive=request.positive, | |
| negative=request.negative, | |
| folder=folder, | |
| first=first, | |
| second=model, | |
| ) | |
| ) | |
| for second in reversed(models[:index]): | |
| queued_matrix_cell( | |
| MatrixCellRequest( | |
| generation=request.generation, | |
| positive=request.positive, | |
| negative=request.negative, | |
| folder=folder, | |
| first=model, | |
| second=second, | |
| ) | |
| ) | |
| create_matrix_grid(folder, number, models, request.generation) | |
| def run_comparison_matrix(request, folder, number): | |
| first, second, rows, columns = comparison_plan(request) | |
| for row, row_label in enumerate(rows): | |
| for column, column_label in enumerate(columns): | |
| sampler = column_label | |
| scheduler = MATRIX_SCHEDULER | |
| if request.type == "scheduler": | |
| sampler = request.sampler or MATRIX_SAMPLER | |
| scheduler = column_label | |
| elif request.type == "sampler+scheduler": | |
| scheduler = row_label | |
| output = comparison_image_path( | |
| folder, | |
| number, | |
| row, | |
| column, | |
| ).name | |
| queued_matrix_cell( | |
| MatrixCellRequest( | |
| generation=request.generation, | |
| positive=request.positive, | |
| negative=request.negative, | |
| folder=folder, | |
| first=first, | |
| second=second, | |
| sampler=sampler, | |
| scheduler=scheduler, | |
| second_sampler=sampler, | |
| second_scheduler=scheduler, | |
| output=output, | |
| ) | |
| ) | |
| create_comparison_grid(folder, number, rows, columns) | |
| def run_combined_matrix(request, folder, number): | |
| rows, columns = combined_plan() | |
| for column, (first, second) in enumerate(columns): | |
| for row, (first_options, second_options) in enumerate(rows): | |
| first_sampler, first_scheduler = first_options | |
| second_sampler, second_scheduler = second_options | |
| output = comparison_image_path( | |
| folder, | |
| number, | |
| row, | |
| column, | |
| ).name | |
| queued_matrix_cell( | |
| MatrixCellRequest( | |
| generation=request.generation, | |
| positive=request.positive, | |
| negative=request.negative, | |
| folder=folder, | |
| first=first, | |
| second=second, | |
| sampler=first_sampler, | |
| scheduler=first_scheduler, | |
| second_sampler=second_sampler, | |
| second_scheduler=second_scheduler, | |
| output=output, | |
| ) | |
| ) | |
| create_comparison_grid( | |
| folder, | |
| number, | |
| [ | |
| f"{first[0]}+{first[1]}x{second[0]}+{second[1]}" | |
| for first, second in rows | |
| ], | |
| [f"{first[0]}x{second[0]}" for first, second in columns], | |
| ) | |
| def run_matrix(request, folder, number): | |
| try: | |
| if request.type == "checkpoint": | |
| run_checkpoint_matrix(request, folder, number) | |
| elif request.type == COMBINED_MATRIX_TYPE: | |
| run_combined_matrix(request, folder, number) | |
| else: | |
| run_comparison_matrix(request, folder, number) | |
| except Exception as error: | |
| log(f"Matrix failed: {error}") | |
| finally: | |
| with matrix_lock: | |
| matrix_grids.discard((folder, number)) | |
| def linked_value(workflow, value): | |
| if not isinstance(value, list) or len(value) < 2: | |
| return value | |
| node = workflow.get(str(value[0]), {}) | |
| inputs = node.get("inputs", {}) | |
| if node.get("class_type") == "StringConcatenate": | |
| parts = [ | |
| linked_value( | |
| workflow, | |
| inputs.get("string_a", ""), | |
| ), | |
| linked_value( | |
| workflow, | |
| inputs.get("string_b", ""), | |
| ), | |
| ] | |
| return str(inputs.get("delimiter", ",")).join(map(str, parts)) | |
| return inputs.get("text", "") | |
| def workflow_values(workflow): | |
| sampler = next( | |
| ( | |
| node | |
| for node in workflow.values() | |
| if node.get("class_type") == "KSampler" | |
| ), | |
| {}, | |
| ) | |
| inputs = sampler.get("inputs", {}) | |
| latent_id = inputs.get("latent_image", [None])[0] | |
| latent = workflow.get(str(latent_id), {}).get("inputs", {}) | |
| positive_id = inputs.get("positive", [None])[0] | |
| positive = ( | |
| workflow.get(str(positive_id), {}) | |
| .get("inputs", {}) | |
| .get("text", "") | |
| ) | |
| negative_id = inputs.get("negative", [None])[0] | |
| negative = ( | |
| workflow.get(str(negative_id), {}) | |
| .get("inputs", {}) | |
| .get("text", DEFAULT_NEGATIVE) | |
| ) | |
| return ( | |
| linked_value(workflow, positive), | |
| latent.get("width", 1152), | |
| latent.get("height", 896), | |
| inputs.get("steps", 16), | |
| negative, | |
| inputs.get("sampler_name", DEFAULT_SAMPLER), | |
| inputs.get("scheduler", DEFAULT_SCHEDULER), | |
| inputs.get("cfg", DEFAULT_CFG), | |
| latent.get("batch_size", DEFAULT_BATCH_SIZE), | |
| ) | |
| def run_job(job_id, workflow): | |
| try: | |
| ( | |
| prompt, | |
| width, | |
| height, | |
| steps, | |
| negative, | |
| sampler, | |
| scheduler, | |
| cfg, | |
| batch_size, | |
| ) = workflow_values(workflow) | |
| image = queued_generate( | |
| DirectRequest( | |
| prompt=prompt, | |
| width=width, | |
| height=height, | |
| steps=steps, | |
| negative=negative, | |
| sampler=sampler, | |
| scheduler=scheduler, | |
| cfg=cfg, | |
| batch_size=batch_size, | |
| ) | |
| ) | |
| filename = f"{job_id}.png" | |
| images[filename] = png_bytes(image) | |
| jobs[job_id] = { | |
| "outputs": { | |
| "output": { | |
| "images": [ | |
| { | |
| "filename": filename, | |
| "subfolder": "", | |
| "type": "output", | |
| } | |
| ] | |
| } | |
| }, | |
| "status": { | |
| "status_str": "success", | |
| "completed": True, | |
| "messages": [], | |
| }, | |
| } | |
| except Exception as error: | |
| jobs[job_id] = { | |
| "outputs": {}, | |
| "status": { | |
| "status_str": "error", | |
| "completed": True, | |
| "messages": [ | |
| [ | |
| "execution_error", | |
| { | |
| "node_id": "output", | |
| "node_type": "Generate", | |
| "exception_type": type(error).__name__, | |
| "exception_message": str(error), | |
| }, | |
| ] | |
| ], | |
| }, | |
| } | |
| def replace_asgi_headers(headers, replacements): | |
| names = {name for name, _ in replacements} | |
| return [item for item in headers if item[0].lower() not in names] + replacements | |
| class GradioEncryptionMiddleware: | |
| def __init__(self, app): | |
| self.app = app | |
| async def __call__(self, scope, receive, send): | |
| if scope["type"] != "http": | |
| await self.app(scope, receive, send) | |
| return | |
| headers = dict(scope.get("headers", [])) | |
| encrypted = headers.get(PROXY_ENCRYPTION_HEADER) == PROXY_ENCRYPTION | |
| if not encrypted: | |
| await self.app(scope, receive, send) | |
| return | |
| decrypted_receive = receive | |
| if scope.get("method") not in {"GET", "HEAD"}: | |
| chunks = [] | |
| while True: | |
| message = await receive() | |
| if message["type"] == "http.disconnect": | |
| return | |
| chunks.append(message.get("body", b"")) | |
| if not message.get("more_body", False): | |
| break | |
| direct = False | |
| try: | |
| body = decrypt_proxy_payload(b"".join(chunks)) | |
| if scope.get("path", "").startswith("/gradio_api/call/"): | |
| payload = json.loads(body) | |
| if not isinstance(payload, dict) or not isinstance(payload.get("data"), list): | |
| raise ValueError("invalid Gradio payload") | |
| payload["data"].append(PASSWORD) | |
| body = json.dumps(payload, separators=(",", ":")).encode() | |
| else: | |
| direct = True | |
| except (InvalidTag, ValueError): | |
| content = b'{"detail":"Invalid encrypted payload"}' | |
| await send({ | |
| "type": "http.response.start", | |
| "status": 400, | |
| "headers": [ | |
| (b"content-type", b"application/json"), | |
| (b"content-length", str(len(content)).encode()), | |
| ], | |
| }) | |
| await send({"type": "http.response.body", "body": content}) | |
| return | |
| scope = dict(scope) | |
| if direct: | |
| scope["query_string"] = f"p={quote(PASSWORD)}".encode() | |
| scope["headers"] = replace_asgi_headers( | |
| scope.get("headers", []), | |
| [ | |
| (b"content-type", b"application/json"), | |
| (b"content-length", str(len(body)).encode()), | |
| ], | |
| ) | |
| delivered = False | |
| async def decrypted_receive(): | |
| nonlocal delivered | |
| if delivered: | |
| return {"type": "http.request", "body": b"", "more_body": False} | |
| delivered = True | |
| return {"type": "http.request", "body": body, "more_body": False} | |
| start = None | |
| response_chunks = [] | |
| async def encrypted_send(message): | |
| nonlocal start | |
| if message["type"] == "http.response.start": | |
| start = message | |
| return | |
| if message["type"] == "http.response.pathsend": | |
| await send(start) | |
| await send(message) | |
| return | |
| if message["type"] != "http.response.body": | |
| await send(message) | |
| return | |
| response_chunks.append(message.get("body", b"")) | |
| if message.get("more_body", False): | |
| return | |
| content = b"".join(response_chunks) | |
| if not content.startswith(FILE_MAGIC): | |
| content = encrypt_proxy_payload(content) | |
| start = dict(start) | |
| start["headers"] = replace_asgi_headers( | |
| start.get("headers", []), | |
| [ | |
| (PROXY_ENCRYPTION_HEADER, PROXY_ENCRYPTION), | |
| (b"content-length", str(len(content)).encode()), | |
| ], | |
| ) | |
| await send(start) | |
| await send({"type": "http.response.body", "body": content}) | |
| await self.app(scope, decrypted_receive, encrypted_send) | |
| api = App() | |
| api.add_middleware(GradioEncryptionMiddleware) | |
| api.add_middleware( | |
| CORSMiddleware, | |
| allow_origins=["*"], | |
| allow_methods=["*"], | |
| allow_headers=["*"], | |
| ) | |
| def require_pass(p: str = Query(...)): | |
| if not valid_pass(p): | |
| raise HTTPException(401, "invalid password") | |
| def health(): | |
| return {"status": True} | |
| def start_matrix(body: MatrixRequest): | |
| try: | |
| if body.type not in MATRIX_TYPES: | |
| raise ValueError(f"Unsupported matrix type: {body.type}") | |
| if body.type == "checkpoint": | |
| matrix_models() | |
| elif body.type == COMBINED_MATRIX_TYPE: | |
| combined_plan() | |
| else: | |
| comparison_plan(body) | |
| except ValueError as error: | |
| raise HTTPException(422, str(error)) from error | |
| folder = datetime.now(TIMEZONE).date().isoformat() | |
| number = reserve_matrix_grid(folder) | |
| matrix_pool.submit(run_matrix, body, folder, number) | |
| return { | |
| "status": "started", | |
| "grid": f"{folder}/{number}gr.epng", | |
| } | |
| def system_stats(): | |
| return { | |
| "system": {"os": os.name}, | |
| "devices": [], | |
| } | |
| def object_info(): | |
| init_comfy() | |
| options = state["sample"].INPUT_TYPES()["required"] | |
| loras = state["lora"].INPUT_TYPES()["required"] | |
| vaes = state["vae_loader"].INPUT_TYPES()["required"] | |
| upscalers = state["upscale"].INPUT_TYPES()["required"] | |
| upscale_models = state["upscale_model_loader"].INPUT_TYPES()["required"] | |
| model_names = generation_models() | |
| lora_names = model_options("loras", loras["lora_name"][0]) | |
| vae_names = model_options("vae", vaes["vae_name"][0]) | |
| detector_names = model_options( | |
| "ultralytics", | |
| state["folders"].get_filename_list("ultralytics_bbox"), | |
| ) | |
| return { | |
| "KSampler": { | |
| "input": { | |
| "required": { | |
| "sampler_name": [ | |
| options["sampler_name"][0] | |
| ], | |
| "scheduler": [ | |
| [*options["scheduler"][0], ALIGN_SCHEDULER] | |
| ], | |
| } | |
| } | |
| }, | |
| "CheckpointLoaderSimple": { | |
| "input": { | |
| "required": { | |
| "ckpt_name": [ | |
| model_names | |
| ] | |
| } | |
| } | |
| }, | |
| "LoraLoader": { | |
| "input": { | |
| "required": { | |
| "lora_name": [ | |
| lora_names | |
| ] | |
| } | |
| } | |
| }, | |
| "LatentUpscaleBy": { | |
| "input": { | |
| "required": { | |
| "upscale_method": [ | |
| upscalers["upscale_method"][0] | |
| ] | |
| } | |
| } | |
| }, | |
| "UpscaleModelLoader": { | |
| "input": { | |
| "required": { | |
| "model_name": [ | |
| model_options( | |
| "upscale_models", | |
| upscale_models["model_name"][0], | |
| ) | |
| ] | |
| } | |
| } | |
| }, | |
| "UltralyticsDetectorProvider": { | |
| "input": { | |
| "required": { | |
| "model_name": [detector_names] | |
| } | |
| } | |
| }, | |
| "UNETLoader": { | |
| "input": { | |
| "required": { | |
| "unet_name": [[ | |
| name | |
| for name in model_names | |
| if is_anima_model(name) | |
| ]] | |
| } | |
| } | |
| }, | |
| "VAELoader": { | |
| "input": { | |
| "required": { | |
| "vae_name": [ | |
| vae_names | |
| ] | |
| } | |
| } | |
| }, | |
| } | |
| def refresh_models(): | |
| with lock: | |
| init_comfy() | |
| added = index_bucket_models() | |
| return { | |
| kind: sorted(names, key=str.casefold) | |
| for kind, names in added.items() | |
| } | |
| def queue_prompt(body: dict): | |
| workflow = body.get("prompt") | |
| if not isinstance(workflow, dict): | |
| raise HTTPException( | |
| 400, | |
| "prompt must contain a ComfyUI workflow object", | |
| ) | |
| job_id = str(uuid.uuid4()) | |
| pool.submit(run_job, job_id, workflow) | |
| return { | |
| "prompt_id": job_id, | |
| "number": len(jobs), | |
| "node_errors": {}, | |
| } | |
| def history(): | |
| return jobs | |
| def history_item(job_id): | |
| return {job_id: jobs[job_id]} if job_id in jobs else {} | |
| def view(filename: str = Query(...)): | |
| data = images.pop(Path(filename).name, None) | |
| if data is None: | |
| raise HTTPException(404, "image not found") | |
| return Response(content=data, media_type="image/png") | |
| def interrupt(): | |
| return {} | |
| def direct( | |
| body: DirectRequest, | |
| ): | |
| try: | |
| image = queued_generate(body) | |
| except ValueError as error: | |
| raise HTTPException(422, str(error)) from error | |
| image = scale_image(image, body.return_scale) | |
| return Response( | |
| content=encrypted_image_bytes(image), | |
| media_type="application/octet-stream", | |
| headers={"Content-Disposition": 'attachment; filename="image.epng"'}, | |
| ) | |
| def archive( | |
| body: bytes = Body(), | |
| from_api: bool = Query(True), | |
| ): | |
| image = Image.open(BytesIO(body)).copy() | |
| backup_pool.submit(archive_image, image, from_api).result() | |
| return {"status": True} | |
| def explorer(): | |
| return FileResponse(Path.cwd() / "explorer.html") | |
| def explorer_css(): | |
| return FileResponse(Path.cwd() / "explorer.css", media_type="text/css") | |
| def explorer_js(): | |
| return FileResponse( | |
| Path.cwd() / "explorer.js", | |
| media_type="text/javascript", | |
| ) | |
| def upscale_page(): | |
| return FileResponse(Path.cwd() / "upscale.html") | |
| def upscale_asset(name: str): | |
| if name not in UPSCALE_ASSETS: | |
| raise HTTPException(404, "asset not found") | |
| filename, media_type = UPSCALE_ASSETS[name] | |
| return FileResponse(Path.cwd() / filename, media_type=media_type) | |
| def upscale_options(): | |
| return { | |
| "models": model_options("upscale_models", []), | |
| "default": DEFAULT_UPSCALE_MODEL, | |
| } | |
| def upscale_download( | |
| body: bytes = Body(b"", media_type="application/octet-stream"), | |
| path: str = Query(None), | |
| model: str = Query(DEFAULT_UPSCALE_MODEL), | |
| scale: float = Query(1, ge=1, le=4), | |
| ): | |
| data = stored_bytes(stored_path(path)) if path is not None else body | |
| if len(data) > MAX_UPSCALE_BYTES: | |
| raise HTTPException(413, "Image exceeds 20 MiB") | |
| if not data: | |
| raise HTTPException(422, "An image is required") | |
| request = json.dumps({ | |
| "image": base64.b64encode(data).decode(), | |
| "model": model, | |
| "scale": scale, | |
| }) | |
| try: | |
| result = get_local_client().predict( | |
| request, | |
| PASSWORD, | |
| api_name="/upscale", | |
| ) | |
| except Exception as error: | |
| raise HTTPException(502, str(error)) from error | |
| return Response( | |
| content=base64.b64decode(result, validate=True), | |
| media_type="image/png", | |
| headers={ | |
| "Content-Disposition": 'attachment; filename="upscaled.png"', | |
| "Cache-Control": "no-store", | |
| }, | |
| ) | |
| def explorer_folder_has_images(path): | |
| with os.scandir(path) as entries: | |
| return any( | |
| item.is_file() | |
| and Path(item.name).suffix.casefold() in IMAGE_SUFFIXES | |
| for item in entries | |
| ) | |
| def explorer_files(folder): | |
| path = (IMAGE_DIR / folder).resolve() | |
| if path.parent != IMAGE_DIR.resolve() or not path.is_dir(): | |
| raise HTTPException(404, "folder not found") | |
| with os.scandir(path) as entries: | |
| return [ | |
| (Path(folder) / item.name).as_posix() | |
| for item in sorted(entries, key=natural_key, reverse=True) | |
| if Path(item.name).suffix.casefold() in IMAGE_SUFFIXES | |
| and item.is_file() | |
| ] | |
| def image_list( | |
| folder: str = Query(None), | |
| infinite: bool = Query(False), | |
| favorites: bool = Query(False), | |
| search: str = Query(""), | |
| offset: int = Query(0, ge=0), | |
| limit: int = Query(EXPLORER_PAGE_SIZE, ge=1, le=EXPLORER_MAX_PAGE_SIZE), | |
| ): | |
| grouped = infinite or (favorites and folder is None) | |
| root = IMAGE_DIR | |
| if not root.is_dir(): | |
| key = "groups" if grouped else "folders" if folder is None else "images" | |
| return {key: [], "next_offset": None} | |
| if grouped or folder is None: | |
| with os.scandir(root) as entries: | |
| folders = [ | |
| item.name | |
| for item in sorted(entries, key=natural_key, reverse=True) | |
| if item.is_dir(follow_symlinks=False) | |
| and (grouped or explorer_folder_has_images(item.path)) | |
| ] | |
| if not grouped: | |
| return {"folders": [{"name": name} for name in folders]} | |
| else: | |
| folders = [folder] | |
| starred = starred_paths() | |
| if favorites: | |
| favorite_folders = {relative.split("/", 1)[0] for relative in starred} | |
| folders = [name for name in folders if name in favorite_folders] | |
| files = (relative for name in folders for relative in explorer_files(name) | |
| if not favorites or relative in starred) | |
| database = star_database() | |
| try: | |
| query = search.strip().casefold() | |
| if query: | |
| files = (relative for relative in files if ( | |
| query in Path(relative).with_suffix(".png").name.casefold() | |
| or any(query in prompt.casefold() for prompt in stored_prompt( | |
| root / relative, relative, database | |
| )[:2]) | |
| )) | |
| page = list(islice(files, offset, offset + limit + 1)) | |
| images = [] | |
| for relative in page[:limit]: | |
| prompt, second_prompt, artists = stored_prompt( | |
| root / relative, relative, database | |
| ) | |
| stat = (root / relative).stat() | |
| images.append({ | |
| "name": Path(relative).with_suffix(".png").name, | |
| "path": relative, | |
| "version": f"{stat.st_mtime_ns}-{stat.st_size}", | |
| "starred": relative in starred, | |
| "prompt": prompt, | |
| "second_prompt": second_prompt, | |
| "artists": artists, | |
| }) | |
| database.commit() | |
| finally: | |
| database.close() | |
| result = {"next_offset": offset + limit if len(page) > limit else None} | |
| if grouped: | |
| groups = {} | |
| for image in images: | |
| name = image["path"].split("/", 1)[0] | |
| groups.setdefault(name, []).append(image) | |
| result["groups"] = [{"name": name, "images": items} | |
| for name, items in groups.items()] | |
| else: | |
| result["images"] = images | |
| return result | |
| def image_star(body: StarRequest): | |
| path = stored_path(body.path).relative_to(IMAGE_DIR.resolve()).as_posix() | |
| set_star(path, body.starred) | |
| return {"starred": body.starred} | |
| def image_preview(path: str = Query(...)): | |
| target = stored_path(path) | |
| stat = target.stat() | |
| return Response( | |
| content=stored_preview(target, stat.st_mtime_ns, stat.st_size), | |
| media_type="image/webp", | |
| headers={"Cache-Control": "private, max-age=86400"}, | |
| ) | |
| def image_original(path: str = Query(...)): | |
| return Response( | |
| content=stored_bytes(stored_path(path)), | |
| media_type="image/png", | |
| headers={"Cache-Control": "private, max-age=86400"}, | |
| ) | |
| def image_download(body: DownloadRequest): | |
| temp = tempfile.NamedTemporaryFile(suffix=".7z", delete=False) | |
| temp.close() | |
| try: | |
| with py7zr.SevenZipFile( | |
| temp.name, | |
| "w", | |
| password=PASSWORD, | |
| header_encryption=True, | |
| ) as archive_file: | |
| for name, path in selected_paths(body.items).items(): | |
| archive_file.writestr(stored_bytes(path), name) | |
| except Exception: | |
| Path(temp.name).unlink(missing_ok=True) | |
| raise | |
| return FileResponse( | |
| temp.name, | |
| filename="images.7z", | |
| media_type="application/x-7z-compressed", | |
| background=BackgroundTask(Path(temp.name).unlink, missing_ok=True), | |
| ) | |
| def image_delete(body: DownloadRequest): | |
| paths = selected_paths(body.items) | |
| relatives = [ | |
| path.relative_to(IMAGE_DIR.resolve()).as_posix() | |
| for path in paths.values() | |
| ] | |
| for path, relative in zip(paths.values(), relatives): | |
| set_star(relative, False) | |
| path.unlink() | |
| database = star_database() | |
| try: | |
| for table in ("image_prompts", "image_previews"): | |
| database.executemany( | |
| f"DELETE FROM {table} WHERE path = ?", | |
| ((relative,) for relative in relatives), | |
| ) | |
| database.commit() | |
| finally: | |
| database.close() | |
| for value in body.items: | |
| path = (IMAGE_DIR / value).resolve() | |
| if path.is_dir() and not any(path.iterdir()): | |
| path.rmdir() | |
| return {"deleted": len(paths)} | |
| def login(p): | |
| if not valid_pass(p): | |
| raise gr.Error("Invalid password") | |
| return ( | |
| gr.Column(visible=False), | |
| gr.Column(visible=True), | |
| ) | |
| def copy_mounted_asset(kind, name): | |
| target = model_path(kind, name) | |
| if target.is_file(): | |
| return | |
| log(f"Copying {kind}/{name}") | |
| target.parent.mkdir(parents=True, exist_ok=True) | |
| temp = target.with_suffix(target.suffix + ".part") | |
| shutil.copy2(BUCKET_MOUNT / kind / name, temp) | |
| temp.replace(target) | |
| log(f"Copied {kind}/{name}: {target.stat().st_size // MIB} MiB") | |
| def preload_assets(): | |
| assets = [] | |
| for kind, asset_ids in STARTUP_ASSET_IDS.items(): | |
| for asset_id in asset_ids: | |
| prefix = f"{asset_id}_" | |
| names = [ | |
| name | |
| for name in remote_models[kind] | |
| if name.startswith(prefix) | |
| and (kind == "diffusion_models" or not is_anima_model(name)) | |
| ] | |
| if len(names) != 1: | |
| raise RuntimeError(f"Expected one {kind} file with prefix {prefix}") | |
| assets.append((kind, names[0])) | |
| loaders = { | |
| "checkpoints": load_model, | |
| "diffusion_models": load_model, | |
| "loras": load_lora, | |
| "ultralytics": load_detector, | |
| "upscale_models": load_upscale_model, | |
| "vae": load_vae, | |
| } | |
| for kind, name in assets: | |
| log(f"Preloading {kind}/{name}") | |
| if (BUCKET_MOUNT / kind / name).is_file(): | |
| copy_mounted_asset(kind, name) | |
| else: | |
| stage_model(kind, name) | |
| if kind not in STYLE_MODEL_KINDS: | |
| loaders[kind](name) | |
| log(f"Preloaded {kind}/{name}") | |
| load_style_pipeline() | |
| log("Preloaded InstantStyle pipeline") | |
| def preload_startup_assets(): | |
| log("Starting startup asset preload") | |
| started = time.monotonic() | |
| with lock: | |
| preload_assets() | |
| elapsed = time.monotonic() - started | |
| log(f"Finished preloading all startup assets in {elapsed:.1f}s") | |
| def refresh_ui( | |
| first_model, | |
| second_model, | |
| upscale_model, | |
| detailer_model, | |
| detector, | |
| ): | |
| with lock: | |
| index_bucket_models() | |
| models = generation_models() | |
| second = [ | |
| ("Reuse first-pass model", ""), | |
| *model_choices(models), | |
| ] | |
| return ( | |
| gr.Dropdown(choices=model_choices(models), value=first_model), | |
| gr.Dropdown(choices=second, value=second_model), | |
| gr.Dropdown( | |
| choices=upscale_model_choices( | |
| model_options("upscale_models", []), | |
| ), | |
| value=upscale_model, | |
| ), | |
| gr.Dropdown( | |
| choices=[("Reuse final model", ""), *model_choices(models)], | |
| value=detailer_model, | |
| ), | |
| gr.Dropdown( | |
| choices=model_options("ultralytics", []), | |
| value=detector, | |
| ), | |
| ) | |
| cleanup_mount() | |
| if __name__ == "__main__": | |
| for _ in range(SCAN_THREAD_COUNT): | |
| threading.Thread( | |
| target=runpy.run_path, | |
| args=("scan.py",), | |
| kwargs={"run_name": "__main__"}, | |
| daemon=True, | |
| ).start() | |
| init_comfy() | |
| MODEL_NAMES = generation_models() | |
| SAMPLE_OPTIONS = state["sample"].INPUT_TYPES()["required"] | |
| SAMPLER_NAMES = sampler_names() | |
| SCHEDULER_NAMES = scheduler_names() | |
| DETAILER_OPTIONS = state["face_detailer"].INPUT_TYPES()["required"] | |
| DETAILER_SAMPLER_NAMES = DETAILER_OPTIONS["sampler_name"][0] | |
| DETAILER_SCHEDULER_NAMES = DETAILER_OPTIONS["scheduler"][0] | |
| UPSCALE_NAMES = state["upscale"].INPUT_TYPES()["required"]["upscale_method"][0] | |
| UPSCALE_MODEL_NAMES = model_options("upscale_models", []) | |
| ULTRALYTICS_NAMES = model_options("ultralytics", []) | |
| SECOND_MODELS = [ | |
| ("Reuse first-pass model", ""), | |
| *model_choices(MODEL_NAMES), | |
| ] | |
| DETAILER_MODELS = [ | |
| ("Reuse final model", ""), | |
| *model_choices(MODEL_NAMES), | |
| ] | |
| with gr.Blocks(title="Image generation") as demo: | |
| with gr.Column() as login_panel: | |
| pass_input = gr.Textbox( | |
| label="Password", | |
| type="password", | |
| ) | |
| login_button = gr.Button( | |
| "Login", | |
| variant="primary", | |
| ) | |
| with gr.Column(visible=False) as generate_panel: | |
| gr.HTML( | |
| "<model-note>" | |
| "ComfyUI-compatible generation" | |
| "</model-note>" | |
| ) | |
| with gr.Row(elem_id="workspace"): | |
| with gr.Column(scale=4, min_width=360): | |
| with gr.Row(): | |
| model_input = gr.Dropdown( | |
| model_choices(MODEL_NAMES), | |
| value=DEFAULT_MODEL, | |
| label="Model", | |
| scale=8, | |
| ) | |
| model_state = gr.State(DEFAULT_MODEL) | |
| refresh_button = gr.Button("Refresh", scale=1) | |
| with gr.Accordion("Add CB asset", open=False): | |
| model_files_input = gr.File( | |
| label="Files", | |
| file_count="multiple", | |
| file_types=list(MODEL_SUFFIXES), | |
| type="filepath", | |
| ) | |
| model_url_input = gr.Textbox(label="URL") | |
| with gr.Row(): | |
| model_location_input = gr.Dropdown( | |
| MODEL_LOCATION_CHOICES, | |
| value="checkpoints", | |
| label="CB location", | |
| ) | |
| anima_model_input = gr.Checkbox( | |
| False, | |
| label="Anima model", | |
| ) | |
| model_upload_button = gr.Button("Add") | |
| model_upload_status = gr.Textbox( | |
| label="Status", | |
| interactive=False, | |
| ) | |
| prompt_input = gr.Textbox(label="Prompt", lines=6) | |
| negative_input = gr.Textbox( | |
| DEFAULT_NEGATIVE, | |
| label="Negative prompt", | |
| lines=3, | |
| ) | |
| regions_input = gr.Dataframe( | |
| value=[["", "full", 1]], | |
| headers=["Prompt", "Area", "Strength"], | |
| datatype=["str", "str", "number"], | |
| type="array", | |
| row_count=(1, "dynamic"), | |
| column_count=(3, "fixed"), | |
| label="Regions: auto, preset, or grid range, up to 3", | |
| ) | |
| regional_mode_input = gr.Dropdown( | |
| [ | |
| ("Soft conditioning", "conditioning"), | |
| ("Attention Couple (PPM)", "attention"), | |
| ], | |
| value=REGIONAL_MODES[0], | |
| label="Regional method", | |
| ) | |
| with gr.Accordion("InstantStyle", open=False): | |
| gr.HTML( | |
| "<model-note>" | |
| "SDXL and Illustrious only. Add up to four references; " | |
| "their center crops are averaged. Use varied subjects and " | |
| "palettes. The default styles both generation passes but leaves " | |
| "ADetailer focused on anatomy." | |
| "</model-note>" | |
| ) | |
| style_images_input = gr.File( | |
| label="Style references", | |
| file_count="multiple", | |
| file_types=["image"], | |
| type="filepath", | |
| ) | |
| style_scope_input = gr.Dropdown( | |
| [ | |
| ("First and second passes", "generation"), | |
| ("First pass only", "first"), | |
| ("All passes, including ADetailer", "all"), | |
| ], | |
| value=DEFAULT_STYLE_SCOPE, | |
| label="Apply to", | |
| ) | |
| with gr.Row(): | |
| style_weight_input = gr.Slider( | |
| 0, | |
| 5, | |
| DEFAULT_STYLE_WEIGHT, | |
| step=.05, | |
| label="Strength", | |
| ) | |
| style_end_input = gr.Slider( | |
| .05, | |
| 1, | |
| DEFAULT_STYLE_END, | |
| step=.05, | |
| label="End at", | |
| ) | |
| with gr.Accordion("ADetailers", open=False): | |
| detailer_input = gr.Dataframe( | |
| headers=[ | |
| "Detector", "Model", "Prompt", "Negative", | |
| "Sampler", "Scheduler", "Steps", "CFG", "Denoise", | |
| ], | |
| datatype=[ | |
| "str", "str", "str", "str", "str", "str", | |
| "number", "number", "number", | |
| ], | |
| type="array", | |
| row_count=(1, "dynamic"), | |
| column_count=(9, "fixed"), | |
| label="Ordered detail passes", | |
| ) | |
| with gr.Row(): | |
| detailer_detector_add = gr.Dropdown( | |
| ULTRALYTICS_NAMES, | |
| value=DEFAULT_DETECTOR, | |
| label="Detector", | |
| ) | |
| detailer_model_add = gr.Dropdown( | |
| DETAILER_MODELS, | |
| value="", | |
| label="Model", | |
| ) | |
| detailer_prompt_add = gr.Textbox( | |
| label="Prompt override", | |
| placeholder="Blank reuses the main prompt", | |
| lines=2, | |
| ) | |
| detailer_negative_add = gr.Textbox( | |
| label="Negative override", | |
| placeholder="Blank reuses the main negative prompt", | |
| lines=2, | |
| ) | |
| with gr.Row(): | |
| detailer_sampler_add = gr.Dropdown( | |
| DETAILER_SAMPLER_NAMES, | |
| value=DEFAULT_SECOND_SAMPLER, | |
| label="Sampler", | |
| ) | |
| detailer_scheduler_add = gr.Dropdown( | |
| DETAILER_SCHEDULER_NAMES, | |
| value=DEFAULT_SECOND_SCHEDULER, | |
| label="Scheduler", | |
| ) | |
| with gr.Row(): | |
| detailer_steps_add = gr.Number( | |
| DEFAULT_SECOND_STEPS, | |
| label="Steps", | |
| precision=0, | |
| ) | |
| detailer_cfg_add = gr.Number( | |
| DEFAULT_CFG, | |
| label="CFG", | |
| ) | |
| detailer_denoise_add = gr.Number( | |
| .35, | |
| label="Denoise", | |
| ) | |
| detailer_button = gr.Button("Add", scale=1) | |
| with gr.Accordion("First pass", open=True): | |
| with gr.Row(): | |
| width_input = gr.Number(1152, label="Width", precision=0) | |
| height_input = gr.Number(896, label="Height", precision=0) | |
| batch_size_input = gr.Slider( | |
| 1, | |
| MAX_BATCH_SIZE, | |
| DEFAULT_UI_BATCH_SIZE, | |
| step=1, | |
| label="Images", | |
| ) | |
| with gr.Row(): | |
| sampler_input = gr.Dropdown( | |
| SAMPLER_NAMES, | |
| value=DEFAULT_SAMPLER, | |
| label="Sampler", | |
| ) | |
| scheduler_input = gr.Dropdown( | |
| SCHEDULER_NAMES, | |
| value=DEFAULT_SCHEDULER, | |
| label="Scheduler", | |
| ) | |
| with gr.Row(): | |
| steps_input = gr.Slider( | |
| 1, | |
| 100, | |
| DEFAULT_STEPS, | |
| step=1, | |
| label="Steps", | |
| ) | |
| cfg_input = gr.Slider( | |
| 0, | |
| 20, | |
| DEFAULT_CFG, | |
| step=.1, | |
| label="CFG", | |
| ) | |
| upscale_input = gr.Checkbox( | |
| False, | |
| label="Upscale and run a second pass", | |
| ) | |
| with gr.Column(visible=False) as second_panel: | |
| with gr.Accordion("Second pass", open=True): | |
| second_model_input = gr.Dropdown( | |
| SECOND_MODELS, | |
| value="", | |
| label="Second-pass model", | |
| ) | |
| second_model_state = gr.State("") | |
| with gr.Row(): | |
| upscale_method_input = gr.Dropdown( | |
| UPSCALE_NAMES, | |
| value=DEFAULT_UPSCALE_METHOD, | |
| label="Upscale method", | |
| ) | |
| upscale_scale_input = gr.Number( | |
| DEFAULT_UPSCALE_SCALE, | |
| label="Scale by", | |
| minimum=.01, | |
| ) | |
| upscale_model_input = gr.Dropdown( | |
| upscale_model_choices(UPSCALE_MODEL_NAMES), | |
| value=DEFAULT_UPSCALE_MODEL, | |
| label="Upscale model", | |
| ) | |
| with gr.Row(): | |
| second_sampler_input = gr.Dropdown( | |
| SAMPLER_NAMES, | |
| value=DEFAULT_SECOND_SAMPLER, | |
| label="Sampler", | |
| ) | |
| second_scheduler_input = gr.Dropdown( | |
| SCHEDULER_NAMES, | |
| value=DEFAULT_SECOND_SCHEDULER, | |
| label="Scheduler", | |
| ) | |
| with gr.Row(): | |
| second_steps_input = gr.Slider( | |
| 1, | |
| 100, | |
| DEFAULT_SECOND_STEPS, | |
| step=1, | |
| label="Steps", | |
| ) | |
| second_cfg_input = gr.Slider( | |
| 0, | |
| 20, | |
| DEFAULT_SECOND_CFG, | |
| step=.1, | |
| label="CFG", | |
| ) | |
| denoise_input = gr.Slider( | |
| 0, | |
| 1, | |
| DEFAULT_DENOISE, | |
| step=.01, | |
| label="Denoise", | |
| ) | |
| return_scale_input = gr.Slider( | |
| .01, | |
| 1, | |
| DEFAULT_RETURN_SCALE, | |
| step=.01, | |
| label="Return scale", | |
| ) | |
| with gr.Row(): | |
| button = gr.Button("Generate", variant="primary") | |
| ping_button = gr.Button("Ping") | |
| with gr.Column(scale=6, min_width=420): | |
| output = gr.Gallery( | |
| label="Preview", | |
| format="png", | |
| elem_id="output", | |
| columns=2, | |
| ) | |
| api_button = gr.Button(visible=False) | |
| health_api_button = gr.Button(visible=False) | |
| options_api_button = gr.Button(visible=False) | |
| refresh_api_button = gr.Button(visible=False) | |
| load_models_api_button = gr.Button(visible=False) | |
| matrix_api_button = gr.Button(visible=False) | |
| upscale_api_button = gr.Button(visible=False) | |
| api_pass_input = gr.Textbox(visible=False) | |
| request_input = gr.Textbox(visible=False) | |
| api_json_output = gr.JSON(visible=False) | |
| api_file_output = gr.File(visible=False) | |
| matrix_output = gr.Textbox(visible=False) | |
| upscale_output = gr.Textbox(visible=False) | |
| login_button.click( | |
| login, | |
| pass_input, | |
| [login_panel, generate_panel], | |
| queue=False, | |
| api_visibility="private", | |
| ) | |
| button.click( | |
| generate, | |
| [ | |
| prompt_input, | |
| negative_input, | |
| regions_input, | |
| regional_mode_input, | |
| model_input, | |
| style_images_input, | |
| style_scope_input, | |
| style_weight_input, | |
| style_end_input, | |
| detailer_input, | |
| width_input, | |
| height_input, | |
| batch_size_input, | |
| sampler_input, | |
| scheduler_input, | |
| steps_input, | |
| cfg_input, | |
| upscale_input, | |
| upscale_method_input, | |
| upscale_model_input, | |
| upscale_scale_input, | |
| second_model_input, | |
| second_sampler_input, | |
| second_scheduler_input, | |
| second_steps_input, | |
| second_cfg_input, | |
| denoise_input, | |
| return_scale_input, | |
| pass_input, | |
| ], | |
| output, | |
| api_name="ui_generate", | |
| api_visibility="private", | |
| ) | |
| ping_button.click( | |
| ping_image, | |
| pass_input, | |
| output, | |
| api_visibility="private", | |
| ) | |
| upscale_input.change( | |
| lambda enabled: gr.Column(visible=enabled), | |
| upscale_input, | |
| second_panel, | |
| queue=False, | |
| ) | |
| detailer_button.click( | |
| add_detailer, | |
| [ | |
| detailer_input, | |
| detailer_detector_add, | |
| detailer_model_add, | |
| detailer_prompt_add, | |
| detailer_negative_add, | |
| detailer_sampler_add, | |
| detailer_scheduler_add, | |
| detailer_steps_add, | |
| detailer_cfg_add, | |
| detailer_denoise_add, | |
| ], | |
| detailer_input, | |
| queue=False, | |
| ) | |
| model_input.change( | |
| select_model, | |
| [model_input, model_state], | |
| [model_input, model_state], | |
| queue=False, | |
| ) | |
| second_model_input.change( | |
| select_model, | |
| [second_model_input, second_model_state], | |
| [second_model_input, second_model_state], | |
| queue=False, | |
| ) | |
| refresh_button.click( | |
| refresh_ui, | |
| inputs=[ | |
| model_input, | |
| second_model_input, | |
| upscale_model_input, | |
| detailer_model_add, | |
| detailer_detector_add, | |
| ], | |
| outputs=[ | |
| model_input, | |
| second_model_input, | |
| upscale_model_input, | |
| detailer_model_add, | |
| detailer_detector_add, | |
| ], | |
| queue=False, | |
| ) | |
| model_upload_button.click( | |
| upload_bucket_assets, | |
| [ | |
| model_files_input, | |
| model_url_input, | |
| model_location_input, | |
| anima_model_input, | |
| pass_input, | |
| ], | |
| model_upload_status, | |
| queue=False, | |
| api_visibility="private", | |
| ) | |
| api_button.click( | |
| api_generate, | |
| [ | |
| request_input, | |
| api_pass_input, | |
| ], | |
| api_file_output, | |
| api_name="generate", | |
| ) | |
| upscale_api_button.click( | |
| api_upscale, | |
| [request_input, api_pass_input], | |
| upscale_output, | |
| api_name="upscale", | |
| ) | |
| health_api_button.click( | |
| api_health, | |
| api_pass_input, | |
| api_json_output, | |
| api_name="health", | |
| ) | |
| options_api_button.click( | |
| api_options, | |
| api_pass_input, | |
| api_json_output, | |
| api_name="options", | |
| ) | |
| refresh_api_button.click( | |
| api_refresh, | |
| api_pass_input, | |
| api_json_output, | |
| api_name="refresh", | |
| ) | |
| load_models_api_button.click( | |
| api_load_models, | |
| [request_input, api_pass_input], | |
| api_json_output, | |
| api_name="load_models", | |
| queue=False, | |
| ) | |
| matrix_api_button.click( | |
| api_matrix_cell, | |
| [ | |
| request_input, | |
| api_pass_input, | |
| ], | |
| matrix_output, | |
| api_name="matrix_cell", | |
| ) | |
| demo.queue(default_concurrency_limit=1) | |
| if __name__ == "__main__": | |
| demo.launch( | |
| server_name="0.0.0.0", | |
| server_port=PORT, | |
| share=True, | |
| ssr_mode=False, | |
| css_paths="style.css", | |
| head='<link rel="icon" href="data:,">', | |
| _app=api, | |
| prevent_thread_lock=True, | |
| ) | |
| threading.Thread(target=preload_startup_assets, daemon=True).start() | |
| demo.block_thread() | |