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"(?", 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 @spaces.GPU(duration=GPU_DURATION) 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 @spaces.GPU(duration=UPSCALE_GPU_DURATION) 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) @spaces.GPU(duration=3) 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 @spaces.GPU(duration=GPU_DURATION) 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) @lru_cache(maxsize=IMAGE_KEY_CACHE_SIZE) 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, ) @lru_cache(maxsize=PREVIEW_CACHE_SIZE) 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") @api.get( "/health", dependencies=[Depends(require_pass)], ) def health(): return {"status": True} @api.post( "/start", dependencies=[Depends(require_pass)], ) 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", } @api.get( "/system_stats", dependencies=[Depends(require_pass)], ) def system_stats(): return { "system": {"os": os.name}, "devices": [], } @api.get( "/object_info", dependencies=[Depends(require_pass)], ) 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 ] } } }, } @api.post( "/refresh_models", dependencies=[Depends(require_pass)], ) def refresh_models(): with lock: init_comfy() added = index_bucket_models() return { kind: sorted(names, key=str.casefold) for kind, names in added.items() } @api.post( "/prompt", dependencies=[Depends(require_pass)], ) 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": {}, } @api.get( "/history", dependencies=[Depends(require_pass)], ) def history(): return jobs @api.get( "/history/{job_id}", dependencies=[Depends(require_pass)], ) def history_item(job_id): return {job_id: jobs[job_id]} if job_id in jobs else {} @api.get( "/view", dependencies=[Depends(require_pass)], ) 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") @api.post( "/interrupt", dependencies=[Depends(require_pass)], ) def interrupt(): return {} @api.post( "/api/generate", dependencies=[Depends(require_pass)], ) 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"'}, ) @api.post( "/api/archive", dependencies=[Depends(require_pass)], ) 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} @api.get( "/explorer", dependencies=[Depends(require_pass)], ) def explorer(): return FileResponse(Path.cwd() / "explorer.html") @api.get("/explorer.css") def explorer_css(): return FileResponse(Path.cwd() / "explorer.css", media_type="text/css") @api.get("/explorer.js") def explorer_js(): return FileResponse( Path.cwd() / "explorer.js", media_type="text/javascript", ) @api.get("/upscale", dependencies=[Depends(require_pass)]) def upscale_page(): return FileResponse(Path.cwd() / "upscale.html") @api.get("/upscale/{name}") 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) @api.get("/api/upscale/options", dependencies=[Depends(require_pass)]) def upscale_options(): return { "models": model_options("upscale_models", []), "default": DEFAULT_UPSCALE_MODEL, } @api.post("/api/upscale", dependencies=[Depends(require_pass)]) 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() ] @api.get("/api/images", dependencies=[Depends(require_pass)]) 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 @api.post("/api/images/star", dependencies=[Depends(require_pass)]) 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} @api.get( "/api/images/preview", dependencies=[Depends(require_pass)], ) 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"}, ) @api.get( "/api/images/original", dependencies=[Depends(require_pass)], ) 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"}, ) @api.post( "/api/images/download", dependencies=[Depends(require_pass)], ) 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), ) @api.post( "/api/images/delete", dependencies=[Depends(require_pass)], ) 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( "" "ComfyUI-compatible generation" "" ) 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( "" "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." "" ) 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='', _app=api, prevent_thread_lock=True, ) threading.Thread(target=preload_startup_assets, daemon=True).start() demo.block_thread()