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()