diff --git a/CHANGELOG.md b/CHANGELOG.md
index 746f07d25..76122441a 100644
--- a/CHANGELOG.md
+++ b/CHANGELOG.md
@@ -8,20 +8,16 @@ BLOCKERS:
OPTIONAL:
- pending `diffusers==0.26.0`
- wuerstchen v3 [pr](https://github.com/huggingface/diffusers/pull/6487)
-- animatediff image2video [pr](https://github.com/huggingface/diffusers/pull/6509)
- tiledvae [pr](https://github.com/huggingface/diffusers/pull/1441)
- style aligned [pr](https://github.com/huggingface/diffusers/pull/6489)
- mixture tiling [pr](https://github.com/huggingface/diffusers/tree/main/examples/community#stable-diffusion-mixture-tiling)
-- depth anything [repo](https://depth-anything.github.io/)
- instaflow [pr](https://github.com/huggingface/diffusers/pull/6057)[repo](https://github.com/gnobitab/RectifiedFlow)
- control api
-- photomaker api
-- interrogate api
-- remb api
+- face api
- masking api
- preprocess api
-## Update for 2023-01-27
+## Update for 2023-01-28
Another big release, highlights being:
- A lot more functionality in the **Control** module:
@@ -277,6 +273,9 @@ As of this release, default backend is set to **diffusers** as its more feature
- major internal ui module refactoring
this may cause compatibility issues if an extension is doing a direct import from `ui.py`
in which case, report it so we can add a compatibility layer
+ - major public api refactoring
+ this may cause compatibility issues if an extension is doing a direct import from `api.py` or `models.py`
+ in which case, report it so we can add a compatibility layer
## Update for 2023-12-29
diff --git a/installer.py b/installer.py
index 9b4476637..1f023efa5 100644
--- a/installer.py
+++ b/installer.py
@@ -48,7 +48,7 @@ args = Dot({
git_commit = "unknown"
submodules_commit = {
'sd-webui-controlnet': 'ecd33eb',
- 'stable-diffusion-webui-images-browser': '27fe4a7',
+ # 'stable-diffusion-webui-images-browser': '27fe4a7',
}
# setup console and file logging
diff --git a/modules/api/api.py b/modules/api/api.py
index c1cf45b46..2cce1edf6 100644
--- a/modules/api/api.py
+++ b/modules/api/api.py
@@ -1,101 +1,17 @@
-import io
-import time
-import base64
-from io import BytesIO
-from typing import List, Dict, Any, Optional
+from typing import List, Optional
from threading import Lock
from secrets import compare_digest
from fastapi import FastAPI, APIRouter, Depends, Request
from fastapi.security import HTTPBasic, HTTPBasicCredentials
from fastapi.exceptions import HTTPException
-from PIL import PngImagePlugin,Image
-import requests
-import piexif
-import piexif.helper
-import gradio as gr
-from modules import errors, shared, sd_samplers, deepbooru, sd_hijack, images, scripts, ui, postprocessing, script_callbacks, generation_parameters_copypaste
-from modules.sd_vae import vae_dict
-from modules.api import models
+from modules import errors, shared, sd_samplers, scripts, ui, postprocessing
+from modules.api import models, endpoints, script, train, helpers, server
from modules.processing import StableDiffusionProcessingTxt2Img, StableDiffusionProcessingImg2Img, process_images
-from modules.textual_inversion.textual_inversion import create_embedding, train_embedding
-from modules.textual_inversion.preprocess import preprocess
-from modules.hypernetworks.hypernetwork import create_hypernetwork, train_hypernetwork
-from modules.sd_models import checkpoints_list, unload_model_weights, reload_model_weights
-from modules.sd_models_config import find_checkpoint_config_near_filename
-from modules import devices
+
errors.install()
-def upscaler_to_index(name: str):
- try:
- return [x.name.lower() for x in shared.sd_upscalers].index(name.lower())
- except Exception as e:
- raise HTTPException(status_code=400, detail=f"Invalid upscaler, needs to be one of these: {' , '.join([x.name for x in shared.sd_upscalers])}") from e
-
-def script_name_to_index(name, scripts_list):
- try:
- return [script.title().lower() for script in scripts_list].index(name.lower())
- except Exception as e:
- raise HTTPException(status_code=422, detail=f"Script '{name}' not found") from e
-
-def validate_sampler_name(name):
- config = sd_samplers.all_samplers_map.get(name, None)
- if config is None:
- raise HTTPException(status_code=404, detail="Sampler not found")
- return name
-
-def setUpscalers(req: dict):
- reqDict = vars(req)
- reqDict['extras_upscaler_1'] = reqDict.pop('upscaler_1', None)
- reqDict['extras_upscaler_2'] = reqDict.pop('upscaler_2', None)
- return reqDict
-
-def decode_base64_to_image(encoding):
- if encoding.startswith("data:image/"):
- encoding = encoding.split(";")[1].split(",")[1]
- try:
- image = Image.open(BytesIO(base64.b64decode(encoding)))
- return image
- except Exception as e:
- shared.log.warning(f'API cannot decode image: {e}')
- raise HTTPException(status_code=500, detail="Invalid encoded image") from e
-
-
-def save_image(image, fn, ext):
- # actual save
- parameters = image.info.get('parameters', None)
- image_format = Image.registered_extensions()[f'.{ext}']
- if image_format == 'PNG':
- pnginfo_data = PngImagePlugin.PngInfo()
- for k, v in image.info.items():
- pnginfo_data.add_text(k, str(v))
- image.save(fn, format=image_format, quality=shared.opts.jpeg_quality, pnginfo=pnginfo_data)
- elif image_format == 'JPEG':
- if image.mode == 'RGBA':
- shared.log.warning('Saving RGBA image as JPEG: Alpha channel will be lost')
- image = image.convert("RGB")
- elif image.mode == 'I;16':
- image = image.point(lambda p: p * 0.0038910505836576).convert("L")
- exif_bytes = piexif.dump({ "Exif": { piexif.ExifIFD.UserComment: piexif.helper.UserComment.dump(parameters or "", encoding="unicode") } })
- image.save(fn, format=image_format, quality=shared.opts.jpeg_quality, exif=exif_bytes)
- elif image_format == 'WEBP':
- if image.mode == 'I;16':
- image = image.point(lambda p: p * 0.0038910505836576).convert("RGB")
- exif_bytes = piexif.dump({ "Exif": { piexif.ExifIFD.UserComment: piexif.helper.UserComment.dump(parameters or "", encoding="unicode") } })
- image.save(fn, format=image_format, quality=shared.opts.jpeg_quality, lossless=shared.opts.webp_lossless, exif=exif_bytes)
- else:
- # shared.log.warning(f'Unrecognized image format: {extension} attempting save as {image_format}')
- image.save(fn, format=image_format, quality=shared.opts.jpeg_quality)
-
-
-def encode_pil_to_base64(image):
- with io.BytesIO() as output_bytes:
- save_image(image, output_bytes, shared.opts.samples_format)
- bytes_data = output_bytes.getvalue()
- return base64.b64encode(bytes_data)
-
-
class Api:
def __init__(self, app: FastAPI, queue_lock: Lock):
self.credentials = {}
@@ -112,44 +28,58 @@ class Api:
self.router = APIRouter()
self.app = app
self.queue_lock = queue_lock
- self.add_api_route("/sdapi/v1/txt2img", self.text2imgapi, methods=["POST"], response_model=models.TextToImageResponse)
- self.add_api_route("/sdapi/v1/img2img", self.img2imgapi, methods=["POST"], response_model=models.ImageToImageResponse)
- self.add_api_route("/sdapi/v1/extra-single-image", self.extras_single_image_api, methods=["POST"], response_model=models.ExtrasSingleImageResponse)
- self.add_api_route("/sdapi/v1/extra-batch-images", self.extras_batch_images_api, methods=["POST"], response_model=models.ExtrasBatchImagesResponse)
- self.add_api_route("/sdapi/v1/png-info", self.pnginfoapi, methods=["POST"], response_model=models.PNGInfoResponse)
- self.add_api_route("/sdapi/v1/progress", self.progressapi, methods=["GET"], response_model=models.ProgressResponse)
- self.add_api_route("/sdapi/v1/interrogate", self.interrogateapi, methods=["POST"])
- self.add_api_route("/sdapi/v1/interrupt", self.interruptapi, methods=["POST"])
- self.add_api_route("/sdapi/v1/skip", self.skip, methods=["POST"])
- self.add_api_route("/sdapi/v1/options", self.get_config, methods=["GET"], response_model=models.OptionsModel)
- self.add_api_route("/sdapi/v1/options", self.set_config, methods=["POST"])
- self.add_api_route("/sdapi/v1/cmd-flags", self.get_cmd_flags, methods=["GET"], response_model=models.FlagsModel)
- self.add_api_route("/sdapi/v1/samplers", self.get_samplers, methods=["GET"], response_model=List[models.SamplerItem])
- self.add_api_route("/sdapi/v1/upscalers", self.get_upscalers, methods=["GET"], response_model=List[models.UpscalerItem])
- self.add_api_route("/sdapi/v1/sd-models", self.get_sd_models, methods=["GET"], response_model=List[models.SDModelItem])
- self.add_api_route("/sdapi/v1/hypernetworks", self.get_hypernetworks, methods=["GET"], response_model=List[models.HypernetworkItem])
- self.add_api_route("/sdapi/v1/face-restorers", self.get_face_restorers, methods=["GET"], response_model=List[models.FaceRestorerItem])
- self.add_api_route("/sdapi/v1/prompt-styles", self.get_prompt_styles, methods=["GET"], response_model=List[models.StyleItem])
- self.add_api_route("/sdapi/v1/embeddings", self.get_embeddings, methods=["GET"], response_model=models.EmbeddingsResponse)
- self.add_api_route("/sdapi/v1/refresh-checkpoints", self.refresh_checkpoints, methods=["POST"])
- self.add_api_route("/sdapi/v1/sd-vae", self.get_sd_vaes, methods=["GET"], response_model=List[models.SDVaeItem])
- self.add_api_route("/sdapi/v1/refresh-vae", self.refresh_vaes, methods=["POST"])
- self.add_api_route("/sdapi/v1/create/embedding", self.create_embedding, methods=["POST"], response_model=models.CreateResponse)
- self.add_api_route("/sdapi/v1/create/hypernetwork", self.create_hypernetwork, methods=["POST"], response_model=models.CreateResponse)
- self.add_api_route("/sdapi/v1/preprocess", self.preprocess, methods=["POST"], response_model=models.PreprocessResponse)
- self.add_api_route("/sdapi/v1/train/embedding", self.train_embedding, methods=["POST"], response_model=models.TrainResponse)
- self.add_api_route("/sdapi/v1/train/hypernetwork", self.train_hypernetwork, methods=["POST"], response_model=models.TrainResponse)
- self.add_api_route("/sdapi/v1/shutdown", self.shutdown, methods=["POST"])
- self.add_api_route("/sdapi/v1/memory", self.get_memory, methods=["GET"], response_model=models.MemoryResponse)
- self.add_api_route("/sdapi/v1/unload-checkpoint", self.unloadapi, methods=["POST"])
- self.add_api_route("/sdapi/v1/reload-checkpoint", self.reloadapi, methods=["POST"])
- self.add_api_route("/sdapi/v1/scripts", self.get_scripts_list, methods=["GET"], response_model=models.ScriptsList)
- self.add_api_route("/sdapi/v1/script-info", self.get_script_info, methods=["GET"], response_model=List[models.ScriptInfo])
- self.add_api_route("/sdapi/v1/extensions", self.get_extensions_list, methods=["GET"], response_model=List[models.ExtensionItem])
- self.add_api_route("/sdapi/v1/log", self.get_log_buffer, methods=["GET"], response_model=List)
- self.add_api_route("/sdapi/v1/start", self.session_start, methods=["GET"])
- self.add_api_route("/sdapi/v1/motd", self.get_motd, methods=["GET"], response_model=str)
- self.add_api_route("/sdapi/v1/extra-networks", self.get_extra_networks, methods=["GET"], response_model=List[models.ExtraNetworkItem])
+
+ # server api
+ self.add_api_route("/sdapi/v1/motd", server.get_motd, methods=["GET"], response_model=str)
+ self.add_api_route("/sdapi/v1/log", server.get_log_buffer, methods=["GET"], response_model=List[str])
+ self.add_api_route("/sdapi/v1/start", self.get_session_start, methods=["GET"])
+ self.add_api_route("/sdapi/v1/progress", server.get_progress, methods=["GET"], response_model=models.ResProgress)
+ self.add_api_route("/sdapi/v1/interrupt", server.post_interrupt, methods=["POST"])
+ self.add_api_route("/sdapi/v1/skip", server.post_skip, methods=["POST"])
+ self.add_api_route("/sdapi/v1/shutdown", server.post_shutdown, methods=["POST"])
+ self.add_api_route("/sdapi/v1/memory", server.get_memory, methods=["GET"], response_model=models.ResMemory)
+ self.add_api_route("/sdapi/v1/options", server.get_config, methods=["GET"], response_model=models.OptionsModel)
+ self.add_api_route("/sdapi/v1/options", server.set_config, methods=["POST"])
+ self.add_api_route("/sdapi/v1/cmd-flags", server.get_cmd_flags, methods=["GET"], response_model=models.FlagsModel)
+
+ # core api using locking
+ self.add_api_route("/sdapi/v1/txt2img", self.post_text2img, methods=["POST"], response_model=models.ResTxt2Img)
+ self.add_api_route("/sdapi/v1/img2img", self.post_img2img, methods=["POST"], response_model=models.ResImg2Img)
+ self.add_api_route("/sdapi/v1/extra-single-image", self.extras_single_image_api, methods=["POST"], response_model=models.ResProcessImage)
+ self.add_api_route("/sdapi/v1/extra-batch-images", self.extras_batch_images_api, methods=["POST"], response_model=models.ResProcessBatch)
+
+ # api dealing with optional scripts
+ self.add_api_route("/sdapi/v1/scripts", script.get_scripts_list, methods=["GET"], response_model=models.ResScripts)
+ self.add_api_route("/sdapi/v1/script-info", script.get_script_info, methods=["GET"], response_model=List[models.ItemScript])
+
+ # enumerator api
+ self.add_api_route("/sdapi/v1/interrogate", endpoints.get_interrogate, methods=["GET"], response_model=List[str])
+ self.add_api_route("/sdapi/v1/samplers", endpoints.get_samplers, methods=["GET"], response_model=List[models.ItemSampler])
+ self.add_api_route("/sdapi/v1/upscalers", endpoints.get_upscalers, methods=["GET"], response_model=List[models.ItemUpscaler])
+ self.add_api_route("/sdapi/v1/sd-models", endpoints.get_sd_models, methods=["GET"], response_model=List[models.ItemModel])
+ self.add_api_route("/sdapi/v1/hypernetworks", endpoints.get_hypernetworks, methods=["GET"], response_model=List[models.ItemHypernetwork])
+ self.add_api_route("/sdapi/v1/face-restorers", endpoints.get_face_restorers, methods=["GET"], response_model=List[models.ItemFaceRestorer])
+ self.add_api_route("/sdapi/v1/prompt-styles", endpoints.get_prompt_styles, methods=["GET"], response_model=List[models.ItemStyle])
+ self.add_api_route("/sdapi/v1/embeddings", endpoints.get_embeddings, methods=["GET"], response_model=models.ResEmbeddings)
+ self.add_api_route("/sdapi/v1/sd-vae", endpoints.get_sd_vaes, methods=["GET"], response_model=List[models.ItemVae])
+ self.add_api_route("/sdapi/v1/extensions", endpoints.get_extensions_list, methods=["GET"], response_model=List[models.ItemExtension])
+ self.add_api_route("/sdapi/v1/extra-networks", endpoints.get_extra_networks, methods=["GET"], response_model=List[models.ItemExtraNetwork])
+
+ # functional api
+ self.add_api_route("/sdapi/v1/png-info", endpoints.post_pnginfo, methods=["POST"], response_model=models.ResImageInfo)
+ self.add_api_route("/sdapi/v1/interrogate", endpoints.post_interrogate, methods=["POST"])
+ self.add_api_route("/sdapi/v1/refresh-checkpoints", endpoints.post_refresh_checkpoints, methods=["POST"])
+ self.add_api_route("/sdapi/v1/unload-checkpoint", endpoints.post_unload_checkpoint, methods=["POST"])
+ self.add_api_route("/sdapi/v1/reload-checkpoint", endpoints.post_reload_checkpoint, methods=["POST"])
+ self.add_api_route("/sdapi/v1/refresh-vae", endpoints.post_refresh_vae, methods=["POST"])
+
+ # train api
+ self.add_api_route("/sdapi/v1/create/embedding", train.post_create_embedding, methods=["POST"], response_model=models.ResCreate)
+ self.add_api_route("/sdapi/v1/create/hypernetwork", train.post_create_hypernetwork, methods=["POST"], response_model=models.ResCreate)
+ self.add_api_route("/sdapi/v1/preprocess", train.post_preprocess, methods=["POST"], response_model=models.ResPreprocess)
+ self.add_api_route("/sdapi/v1/train/embedding", train.post_train_embedding, methods=["POST"], response_model=models.ResTrain)
+ self.add_api_route("/sdapi/v1/train/hypernetwork", train.post_train_hypernetwork, methods=["POST"], response_model=models.ResTrain)
+
self.default_script_arg_txt2img = []
self.default_script_arg_img2img = []
@@ -165,103 +95,13 @@ class Api:
return True
raise HTTPException(status_code=401, detail="Unauthorized", headers={"WWW-Authenticate": "Basic"})
- def get_log_buffer(self, req: models.LogRequest = Depends()):
- lines = shared.log.buffer[:req.lines] if req.lines > 0 else shared.log.buffer.copy()
- if req.clear:
- shared.log.buffer.clear()
- return lines
-
- def session_start(self, req: Request, agent: Optional[str] = None):
+ def get_session_start(self, req: Request, agent: Optional[str] = None):
token = req.cookies.get("access-token") or req.cookies.get("access-token-unsecure")
user = self.app.tokens.get(token) if hasattr(self.app, 'tokens') else None
shared.log.info(f'Browser session: user={user} client={req.client.host} agent={agent}')
return {}
- def get_motd(self):
- from installer import get_version
- motd = ''
- ver = get_version()
- if ver.get('updated', None) is not None:
- motd = f"version {ver['hash']} {ver['updated']} {ver['url'].split('/')[-1]}
"
- if shared.opts.motd:
- res = requests.get('https://vladmandic.github.io/automatic/motd', timeout=10)
- if res.status_code == 200:
- msg = (res.text or '').strip()
- shared.log.info(f'MOTD: {msg if len(msg) > 0 else "N/A"}')
- motd += res.text
- return motd
-
- def get_selectable_script(self, script_name, script_runner):
- if script_name is None or script_name == "":
- return None, None
- script_idx = script_name_to_index(script_name, script_runner.selectable_scripts)
- script = script_runner.selectable_scripts[script_idx]
- return script, script_idx
-
- def get_scripts_list(self):
- t2ilist = [script.name for script in scripts.scripts_txt2img.scripts if script.name is not None]
- i2ilist = [script.name for script in scripts.scripts_img2img.scripts if script.name is not None]
- control = [script.name for script in scripts.scripts_control.scripts if script.name is not None]
- return models.ScriptsList(txt2img = t2ilist, img2img = i2ilist, control = control)
-
- def get_script_info(self, script_name: Optional[str] = None):
- res = []
- for script_list in [scripts.scripts_txt2img.scripts, scripts.scripts_img2img.scripts, scripts.scripts_control.scripts]:
- for script in script_list:
- if script.api_info is not None and (script_name is None or script_name == script.api_info.name):
- res.append(script.api_info)
- return res
-
- def get_script(self, script_name, script_runner):
- if script_name is None or script_name == "":
- return None, None
- script_idx = script_name_to_index(script_name, script_runner.scripts)
- return script_runner.scripts[script_idx]
-
- def init_default_script_args(self, script_runner):
- #find max idx from the scripts in runner and generate a none array to init script_args
- last_arg_index = 1
- for script in script_runner.scripts:
- if last_arg_index < script.args_to:
- last_arg_index = script.args_to
- # None everywhere except position 0 to initialize script args
- script_args = [None]*last_arg_index
- script_args[0] = 0
-
- # get default values
- if gr is None:
- return script_args
- with gr.Blocks(): # will throw errors calling ui function without this
- for script in script_runner.scripts:
- if script.ui(script.is_img2img):
- ui_default_values = []
- for elem in script.ui(script.is_img2img):
- ui_default_values.append(elem.value)
- script_args[script.args_from:script.args_to] = ui_default_values
- return script_args
-
- def init_script_args(self, p, request, default_script_args, selectable_scripts, selectable_script_idx, script_runner):
- script_args = default_script_args.copy()
- # position 0 in script_arg is the idx+1 of the selectable script that is going to be run when using scripts.scripts_*2img.run()
- if selectable_scripts:
- script_args[selectable_scripts.args_from:selectable_scripts.args_to] = request.script_args
- script_args[0] = selectable_script_idx + 1
- # Now check for always on scripts
- if request.alwayson_scripts and (len(request.alwayson_scripts) > 0):
- for alwayson_script_name in request.alwayson_scripts.keys():
- alwayson_script = self.get_script(alwayson_script_name, script_runner)
- if alwayson_script is None:
- raise HTTPException(status_code=422, detail=f"Always on script not found: {alwayson_script_name}")
- if not alwayson_script.alwayson:
- raise HTTPException(status_code=422, detail=f"Selectable script cannot be in always on params: {alwayson_script_name}")
- if "args" in request.alwayson_scripts[alwayson_script_name]:
- # min between arg length in scriptrunner and arg length in the request
- for idx in range(0, min((alwayson_script.args_to - alwayson_script.args_from), len(request.alwayson_scripts[alwayson_script_name]["args"]))):
- script_args[alwayson_script.args_from + idx] = request.alwayson_scripts[alwayson_script_name]["args"][idx]
- p.per_script_args[alwayson_script.title()] = request.alwayson_scripts[alwayson_script_name]["args"]
- return script_args
-
- def prepare_img_gen_request(self, request, img_gen_type: str): # pylint: disable=unused-argument
+ def prepare_img_gen_request(self, request):
if hasattr(request, "face_id") and request.face_id and not request.script_name and (not request.alwayson_scripts or "FaceID" not in request.alwayson_scripts.keys()):
request.script_name = "FaceID"
request.script_args = [
@@ -289,7 +129,7 @@ class Api:
if isinstance(args[idx], str) and len(args[idx]) >= 1000:
args[idx] = f""
- def sanitize_img_gen_request(self, request, img_gen_type: str):
+ def sanitize_img_gen_request(self, request):
if hasattr(request, "alwayson_scripts") and request.alwayson_scripts:
for script_name in request.alwayson_scripts.keys():
script_obj = request.alwayson_scripts[script_name]
@@ -300,18 +140,24 @@ class Api:
if hasattr(request, "script_args") and request.script_args:
self.sanitize_args(request.script_args)
- def text2imgapi(self, txt2imgreq: models.StableDiffusionTxt2ImgProcessingAPI):
- self.prepare_img_gen_request(txt2imgreq, "txt2img")
+ def validate_sampler_name(self, name):
+ config = sd_samplers.all_samplers_map.get(name, None)
+ if config is None:
+ raise HTTPException(status_code=404, detail="Sampler not found")
+ return name
+
+ def post_text2img(self, txt2imgreq: models.ReqTxt2Img):
+ self.prepare_img_gen_request(txt2imgreq)
script_runner = scripts.scripts_txt2img
if not script_runner.scripts:
script_runner.initialize_scripts(False)
ui.create_ui(None)
if not self.default_script_arg_txt2img:
- self.default_script_arg_txt2img = self.init_default_script_args(script_runner)
- selectable_scripts, selectable_script_idx = self.get_selectable_script(txt2imgreq.script_name, script_runner)
+ self.default_script_arg_txt2img = script.init_default_script_args(script_runner)
+ selectable_scripts, selectable_script_idx = script.get_selectable_script(txt2imgreq.script_name, script_runner)
populate = txt2imgreq.copy(update={ # Override __init__ params
- "sampler_name": validate_sampler_name(txt2imgreq.sampler_name or txt2imgreq.sampler_index),
+ "sampler_name": self.validate_sampler_name(txt2imgreq.sampler_name or txt2imgreq.sampler_index),
"do_not_save_samples": not txt2imgreq.save_images,
"do_not_save_grid": not txt2imgreq.save_images,
})
@@ -332,7 +178,7 @@ class Api:
p.outpath_grids = shared.opts.outdir_grids or shared.opts.outdir_txt2img_grids
p.outpath_samples = shared.opts.outdir_samples or shared.opts.outdir_txt2img_samples
shared.state.begin('api-txt2img', api=True)
- script_args = self.init_script_args(p, txt2imgreq, self.default_script_arg_txt2img, selectable_scripts, selectable_script_idx, script_runner)
+ script_args = script.init_script_args(p, txt2imgreq, self.default_script_arg_txt2img, selectable_scripts, selectable_script_idx, script_runner)
if selectable_scripts is not None:
processed = scripts.scripts_txt2img.run(p, *script_args) # Need to pass args as list here
else:
@@ -340,28 +186,28 @@ class Api:
processed = process_images(p)
shared.state.end(api=False)
- b64images = list(map(encode_pil_to_base64, processed.images)) if send_images else []
- self.sanitize_img_gen_request(txt2imgreq, "txt2img")
- return models.TextToImageResponse(images=b64images, parameters=vars(txt2imgreq), info=processed.js())
+ b64images = list(map(helpers.encode_pil_to_base64, processed.images)) if send_images else []
+ self.sanitize_img_gen_request(txt2imgreq)
+ return models.ResTxt2Img(images=b64images, parameters=vars(txt2imgreq), info=processed.js())
- def img2imgapi(self, img2imgreq: models.StableDiffusionImg2ImgProcessingAPI):
- self.prepare_img_gen_request(img2imgreq, "img2img")
+ def post_img2img(self, img2imgreq: models.ReqImg2Img):
+ self.prepare_img_gen_request(img2imgreq)
init_images = img2imgreq.init_images
if init_images is None:
raise HTTPException(status_code=404, detail="Init image not found")
mask = img2imgreq.mask
if mask:
- mask = decode_base64_to_image(mask)
+ mask = helpers.decode_base64_to_image(mask)
script_runner = scripts.scripts_img2img
if not script_runner.scripts:
script_runner.initialize_scripts(True)
ui.create_ui(None)
if not self.default_script_arg_img2img:
- self.default_script_arg_img2img = self.init_default_script_args(script_runner)
- selectable_scripts, selectable_script_idx = self.get_selectable_script(img2imgreq.script_name, script_runner)
+ self.default_script_arg_img2img = script.init_default_script_args(script_runner)
+ selectable_scripts, selectable_script_idx = script.get_selectable_script(img2imgreq.script_name, script_runner)
populate = img2imgreq.copy(update={ # Override __init__ params
- "sampler_name": validate_sampler_name(img2imgreq.sampler_name or img2imgreq.sampler_index),
+ "sampler_name": self.validate_sampler_name(img2imgreq.sampler_name or img2imgreq.sampler_index),
"do_not_save_samples": not img2imgreq.save_images,
"do_not_save_grid": not img2imgreq.save_images,
"mask": mask,
@@ -380,12 +226,12 @@ class Api:
with self.queue_lock:
p = StableDiffusionProcessingImg2Img(sd_model=shared.sd_model, **args)
- p.init_images = [decode_base64_to_image(x) for x in init_images]
+ p.init_images = [helpers.decode_base64_to_image(x) for x in init_images]
p.scripts = script_runner
p.outpath_grids = shared.opts.outdir_img2img_grids
p.outpath_samples = shared.opts.outdir_img2img_samples
shared.state.begin('api-img2img', api=True)
- script_args = self.init_script_args(p, img2imgreq, self.default_script_arg_img2img, selectable_scripts, selectable_script_idx, script_runner)
+ script_args = script.init_script_args(p, img2imgreq, self.default_script_arg_img2img, selectable_scripts, selectable_script_idx, script_runner)
if selectable_scripts is not None:
processed = scripts.scripts_img2img.run(p, *script_args) # Need to pass args as list here
else:
@@ -393,329 +239,33 @@ class Api:
processed = process_images(p)
shared.state.end(api=False)
- b64images = list(map(encode_pil_to_base64, processed.images)) if send_images else []
+ b64images = list(map(helpers.encode_pil_to_base64, processed.images)) if send_images else []
if not img2imgreq.include_init_images:
img2imgreq.init_images = None
img2imgreq.mask = None
- self.sanitize_img_gen_request(img2imgreq, "img2img")
- return models.ImageToImageResponse(images=b64images, parameters=vars(img2imgreq), info=processed.js())
+ self.sanitize_img_gen_request(img2imgreq)
+ return models.ResImg2Img(images=b64images, parameters=vars(img2imgreq), info=processed.js())
- def extras_single_image_api(self, req: models.ExtrasSingleImageRequest):
- reqDict = setUpscalers(req)
- reqDict['image'] = decode_base64_to_image(reqDict['image'])
+ def set_upscalers(self, req: dict):
+ reqDict = vars(req)
+ reqDict['extras_upscaler_1'] = reqDict.pop('upscaler_1', None)
+ reqDict['extras_upscaler_2'] = reqDict.pop('upscaler_2', None)
+ return reqDict
+
+ def extras_single_image_api(self, req: models.ReqProcessImage):
+ reqDict = self.set_upscalers(req)
+ reqDict['image'] = helpers.decode_base64_to_image(reqDict['image'])
with self.queue_lock:
result = postprocessing.run_extras(extras_mode=0, image_folder="", input_dir="", output_dir="", save_output=False, **reqDict)
- return models.ExtrasSingleImageResponse(image=encode_pil_to_base64(result[0][0]), html_info=result[1])
+ return models.ResProcessImage(image=helpers.encode_pil_to_base64(result[0][0]), html_info=result[1])
- def extras_batch_images_api(self, req: models.ExtrasBatchImagesRequest):
- reqDict = setUpscalers(req)
+ def extras_batch_images_api(self, req: models.ReqProcessBatch):
+ reqDict = self.set_upscalers(req)
image_list = reqDict.pop('imageList', [])
- image_folder = [decode_base64_to_image(x.data) for x in image_list]
+ image_folder = [helpers.decode_base64_to_image(x.data) for x in image_list]
with self.queue_lock:
result = postprocessing.run_extras(extras_mode=1, image_folder=image_folder, image="", input_dir="", output_dir="", save_output=False, **reqDict)
- return models.ExtrasBatchImagesResponse(images=list(map(encode_pil_to_base64, result[0])), html_info=result[1])
-
- def pnginfoapi(self, req: models.PNGInfoRequest):
- if not req.image.strip():
- return models.PNGInfoResponse(info="")
-
- image = decode_base64_to_image(req.image.strip())
- if image is None:
- return models.PNGInfoResponse(info="")
-
- geninfo, items = images.read_info_from_image(image)
- if geninfo is None:
- geninfo = ""
-
- if items and items['parameters']:
- del items['parameters']
-
- params = generation_parameters_copypaste.parse_generation_parameters(geninfo)
- script_callbacks.infotext_pasted_callback(geninfo, params)
-
- return models.PNGInfoResponse(info=geninfo, items=items, parameters=params)
-
- def progressapi(self, req: models.ProgressRequest = Depends()):
- if shared.state.job_count == 0:
- return models.ProgressResponse(progress=0, eta_relative=0, state=shared.state.dict(), textinfo=shared.state.textinfo)
-
- shared.state.do_set_current_image()
- current_image = None
- if shared.state.current_image and not req.skip_current_image:
- current_image = encode_pil_to_base64(shared.state.current_image)
-
- batch_x = max(shared.state.job_no, 0)
- batch_y = max(shared.state.job_count, 1)
- step_x = max(shared.state.sampling_step, 0)
- step_y = max(shared.state.sampling_steps, 1)
- current = step_y * batch_x + step_x
- total = step_y * batch_y
- progress = current / total if current > 0 and total > 0 else 0
- time_since_start = time.time() - shared.state.time_start
- eta_relative = (time_since_start / progress) - time_since_start if progress > 0 else 0
-
- res = models.ProgressResponse(progress=progress, eta_relative=eta_relative, state=shared.state.dict(), current_image=current_image, textinfo=shared.state.textinfo)
- return res
-
-
- def interrogateapi(self, interrogatereq: models.InterrogateRequest):
- image_b64 = interrogatereq.image
- if image_b64 is None:
- raise HTTPException(status_code=404, detail="Image not found")
- img = decode_base64_to_image(image_b64)
- img = img.convert('RGB')
- with self.queue_lock:
- if interrogatereq.model == "clip":
- processed = shared.interrogator.interrogate(img)
- elif interrogatereq.model == "deepdanbooru":
- processed = deepbooru.model.tag(img)
- else:
- raise HTTPException(status_code=404, detail="Model not found")
- return models.InterrogateResponse(caption=processed)
-
- def interruptapi(self):
- shared.state.interrupt()
- return {}
-
- def unloadapi(self):
- unload_model_weights(op='model')
- unload_model_weights(op='refiner')
- return {}
-
- def reloadapi(self):
- reload_model_weights()
- return {}
-
- def skip(self):
- shared.state.skip()
-
- def get_config(self):
- options = {}
- for k in shared.opts.data.keys():
- if shared.opts.data_labels.get(k) is not None:
- options.update({k: shared.opts.data.get(k, shared.opts.data_labels.get(k).default)})
- else:
- options.update({k: shared.opts.data.get(k, None)})
- if 'sd_lyco' in options:
- del options['sd_lyco']
- if 'sd_lora' in options:
- del options['sd_lora']
- return options
-
- def set_config(self, req: Dict[str, Any]):
- updated = []
- for k, v in req.items():
- updated.append({ k: shared.opts.set(k, v) })
- shared.opts.save(shared.config_filename)
- return { "updated": updated }
-
- def get_cmd_flags(self):
- return vars(shared.cmd_opts)
-
- def get_samplers(self):
- return [{"name": sampler[0], "aliases":sampler[2], "options":sampler[3]} for sampler in sd_samplers.all_samplers]
-
- def get_sd_vaes(self):
- return [{"model_name": x, "filename": vae_dict[x]} for x in vae_dict.keys()]
-
- def get_upscalers(self):
- return [{"name": upscaler.name, "model_name": upscaler.scaler.model_name, "model_path": upscaler.data_path, "model_url": None, "scale": upscaler.scale} for upscaler in shared.sd_upscalers]
-
- def get_sd_models(self):
- return [{"title": x.title, "model_name": x.name, "filename": x.filename, "type": x.type, "hash": x.shorthash, "sha256": x.sha256, "config": find_checkpoint_config_near_filename(x)} for x in checkpoints_list.values()]
-
- def get_hypernetworks(self):
- return [{"name": name, "path": shared.hypernetworks[name]} for name in shared.hypernetworks]
-
- def get_face_restorers(self):
- return [{"name":x.name(), "cmd_dir": getattr(x, "cmd_dir", None)} for x in shared.face_restorers]
-
- def get_prompt_styles(self):
- return [{ 'name': v.name, 'prompt': v.prompt, 'negative_prompt': v.negative_prompt, 'extra': v.extra, 'filename': v.filename, 'preview': v.preview} for v in shared.prompt_styles.styles.values()]
-
- def get_embeddings(self):
- db = sd_hijack.model_hijack.embedding_db
- def convert_embedding(embedding):
- return {"step": embedding.step, "sd_checkpoint": embedding.sd_checkpoint, "sd_checkpoint_name": embedding.sd_checkpoint_name, "shape": embedding.shape, "vectors": embedding.vectors}
-
- def convert_embeddings(embeddings):
- return {embedding.name: convert_embedding(embedding) for embedding in embeddings.values()}
-
- return {"loaded": convert_embeddings(db.word_embeddings), "skipped": convert_embeddings(db.skipped_embeddings)}
-
- def get_extra_networks(self, page: Optional[str] = None, name: Optional[str] = None, filename: Optional[str] = None, title: Optional[str] = None, fullname: Optional[str] = None, hash: Optional[str] = None): # pylint: disable=redefined-builtin
- res = []
- for pg in shared.extra_networks:
- if page is not None and pg.name != page.lower():
- continue
- for item in pg.items:
- if name is not None and item.get('name', '') != name:
- continue
- if title is not None and item.get('title', '') != title:
- continue
- if filename is not None and item.get('filename', '') != filename:
- continue
- if fullname is not None and item.get('fullname', '') != fullname:
- continue
- if hash is not None and (item.get('shorthash', None) or item.get('hash')) != hash:
- continue
- res.append({
- 'name': item.get('name', ''),
- 'type': pg.name,
- 'title': item.get('title', None),
- 'fullname': item.get('fullname', None),
- 'filename': item.get('filename', None),
- 'hash': item.get('shorthash', None) or item.get('hash'),
- "preview": item.get('preview', None),
- })
- return res
-
- def refresh_checkpoints(self):
- return shared.refresh_checkpoints()
-
- def refresh_vaes(self):
- return shared.refresh_vaes()
-
- def create_embedding(self, args: dict):
- try:
- shared.state.begin('api-embedding')
- filename = create_embedding(**args) # create empty embedding
- sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings() # reload embeddings so new one can be immediately used
- shared.state.end()
- return models.CreateResponse(info = f"create embedding filename: {filename}")
- except AssertionError as e:
- shared.state.end()
- return models.TrainResponse(info = f"create embedding error: {e}")
-
- def create_hypernetwork(self, args: dict):
- try:
- shared.state.begin('api-hypernetwork')
- filename = create_hypernetwork(**args) # create empty embedding # pylint: disable=E1111
- shared.state.end()
- return models.CreateResponse(info = f"create hypernetwork filename: {filename}")
- except AssertionError as e:
- shared.state.end()
- return models.TrainResponse(info = f"create hypernetwork error: {e}")
-
- def preprocess(self, args: dict):
- try:
- shared.state.begin('api-preprocess')
- preprocess(**args) # quick operation unless blip/booru interrogation is enabled
- shared.state.end()
- return models.PreprocessResponse(info = 'preprocess complete')
- except KeyError as e:
- shared.state.end()
- return models.PreprocessResponse(info = f"preprocess error: invalid token: {e}")
- except AssertionError as e:
- shared.state.end()
- return models.PreprocessResponse(info = f"preprocess error: {e}")
- except FileNotFoundError as e:
- shared.state.end()
- return models.PreprocessResponse(info = f'preprocess error: {e}')
-
- def train_embedding(self, args: dict):
- try:
- shared.state.begin('api-embedding')
- apply_optimizations = False
- error = None
- filename = ''
- if not apply_optimizations:
- sd_hijack.undo_optimizations()
- try:
- _embedding, filename = train_embedding(**args) # can take a long time to complete
- except Exception as e:
- error = e
- finally:
- if not apply_optimizations:
- sd_hijack.apply_optimizations()
- shared.state.end()
- return models.TrainResponse(info = f"train embedding complete: filename: {filename} error: {error}")
- except AssertionError as msg:
- shared.state.end()
- return models.TrainResponse(info = f"train embedding error: {msg}")
-
- def train_hypernetwork(self, args: dict):
- try:
- shared.state.begin('api-hypernetwork')
- shared.loaded_hypernetworks = []
- apply_optimizations = False
- error = None
- filename = ''
- if not apply_optimizations:
- sd_hijack.undo_optimizations()
- try:
- _hypernetwork, filename = train_hypernetwork(**args)
- except Exception as e:
- error = e
- finally:
- shared.sd_model.cond_stage_model.to(devices.device)
- shared.sd_model.first_stage_model.to(devices.device)
- if not apply_optimizations:
- sd_hijack.apply_optimizations()
- shared.state.end()
- return models.TrainResponse(info=f"train embedding complete: filename: {filename} error: {error}")
- except AssertionError:
- shared.state.end()
- return models.TrainResponse(info=f"train embedding error: {error}")
-
- def shutdown(self):
- shared.log.info('Shutdown request received')
- import sys
- sys.exit(0)
-
- def get_memory(self):
- try:
- import os
- import psutil
- process = psutil.Process(os.getpid())
- res = process.memory_info() # only rss is cross-platform guaranteed so we dont rely on other values
- ram_total = 100 * res.rss / process.memory_percent() # and total memory is calculated as actual value is not cross-platform safe
- ram = { 'free': ram_total - res.rss, 'used': res.rss, 'total': ram_total }
- except Exception as err:
- ram = { 'error': f'{err}' }
- try:
- import torch
- if torch.cuda.is_available():
- s = torch.cuda.mem_get_info()
- system = { 'free': s[0], 'used': s[1] - s[0], 'total': s[1] }
- s = dict(torch.cuda.memory_stats(shared.device))
- allocated = { 'current': s['allocated_bytes.all.current'], 'peak': s['allocated_bytes.all.peak'] }
- reserved = { 'current': s['reserved_bytes.all.current'], 'peak': s['reserved_bytes.all.peak'] }
- active = { 'current': s['active_bytes.all.current'], 'peak': s['active_bytes.all.peak'] }
- inactive = { 'current': s['inactive_split_bytes.all.current'], 'peak': s['inactive_split_bytes.all.peak'] }
- warnings = { 'retries': s['num_alloc_retries'], 'oom': s['num_ooms'] }
- cuda = {
- 'system': system,
- 'active': active,
- 'allocated': allocated,
- 'reserved': reserved,
- 'inactive': inactive,
- 'events': warnings,
- }
- else:
- cuda = { 'error': 'unavailable' }
- except Exception as err:
- cuda = { 'error': f'{err}' }
- return models.MemoryResponse(ram = ram, cuda = cuda)
-
- def get_extensions_list(self):
- from modules import extensions
- extensions.list_extensions()
- ext_list = []
- for ext in extensions.extensions:
- ext: extensions.Extension
- ext.read_info()
- if ext.remote is not None:
- ext_list.append({
- "name": ext.name,
- "remote": ext.remote,
- "branch": ext.branch,
- "commit_hash":ext.commit_hash,
- "commit_date":ext.commit_date,
- "version":ext.version,
- "enabled":ext.enabled
- })
- return ext_list
+ return models.ResProcessBatch(images=list(map(helpers.encode_pil_to_base64, result[0])), html_info=result[1])
def launch(self):
config = {
@@ -727,9 +277,9 @@ class Api:
"http": "auto", # auto, h11, httptools
}
from modules.server import UvicornServer
- server = UvicornServer(self.app, **config)
+ http_server = UvicornServer(self.app, **config)
# from modules.server import HypercornServer
# server = HypercornServer(self.app, **config)
- server.start()
+ http_server.start()
shared.log.info(f'API server: Uvicorn options={config}')
- return server
+ return http_server
diff --git a/modules/api/endpoints.py b/modules/api/endpoints.py
new file mode 100644
index 000000000..8e9386cb1
--- /dev/null
+++ b/modules/api/endpoints.py
@@ -0,0 +1,147 @@
+from typing import Optional
+from fastapi.exceptions import HTTPException
+from modules import shared
+from modules.api import models, helpers
+
+
+
+def get_samplers():
+ from modules import sd_samplers
+ return [{"name": sampler[0], "aliases":sampler[2], "options":sampler[3]} for sampler in sd_samplers.all_samplers]
+
+def get_sd_vaes():
+ from modules.sd_vae import vae_dict
+ return [{"model_name": x, "filename": vae_dict[x]} for x in vae_dict.keys()]
+
+def get_upscalers():
+ return [{"name": upscaler.name, "model_name": upscaler.scaler.model_name, "model_path": upscaler.data_path, "model_url": None, "scale": upscaler.scale} for upscaler in shared.sd_upscalers]
+
+def get_sd_models():
+ from modules import sd_models, sd_models_config
+ return [{"title": x.title, "model_name": x.name, "filename": x.filename, "type": x.type, "hash": x.shorthash, "sha256": x.sha256, "config": sd_models_config.find_checkpoint_config_near_filename(x)} for x in sd_models.checkpoints_list.values()]
+
+def get_hypernetworks():
+ return [{"name": name, "path": shared.hypernetworks[name]} for name in shared.hypernetworks]
+
+def get_face_restorers():
+ return [{"name":x.name(), "cmd_dir": getattr(x, "cmd_dir", None)} for x in shared.face_restorers]
+
+def get_prompt_styles():
+ return [{ 'name': v.name, 'prompt': v.prompt, 'negative_prompt': v.negative_prompt, 'extra': v.extra, 'filename': v.filename, 'preview': v.preview} for v in shared.prompt_styles.styles.values()]
+
+def get_embeddings():
+ from modules import sd_hijack
+ db = sd_hijack.model_hijack.embedding_db
+ def convert_embedding(embedding):
+ return {"step": embedding.step, "sd_checkpoint": embedding.sd_checkpoint, "sd_checkpoint_name": embedding.sd_checkpoint_name, "shape": embedding.shape, "vectors": embedding.vectors}
+
+ def convert_embeddings(embeddings):
+ return {embedding.name: convert_embedding(embedding) for embedding in embeddings.values()}
+
+ return {"loaded": convert_embeddings(db.word_embeddings), "skipped": convert_embeddings(db.skipped_embeddings)}
+
+def get_extra_networks(page: Optional[str] = None, name: Optional[str] = None, filename: Optional[str] = None, title: Optional[str] = None, fullname: Optional[str] = None, hash: Optional[str] = None): # pylint: disable=redefined-builtin
+ res = []
+ for pg in shared.extra_networks:
+ if page is not None and pg.name != page.lower():
+ continue
+ for item in pg.items:
+ if name is not None and item.get('name', '') != name:
+ continue
+ if title is not None and item.get('title', '') != title:
+ continue
+ if filename is not None and item.get('filename', '') != filename:
+ continue
+ if fullname is not None and item.get('fullname', '') != fullname:
+ continue
+ if hash is not None and (item.get('shorthash', None) or item.get('hash')) != hash:
+ continue
+ res.append({
+ 'name': item.get('name', ''),
+ 'type': pg.name,
+ 'title': item.get('title', None),
+ 'fullname': item.get('fullname', None),
+ 'filename': item.get('filename', None),
+ 'hash': item.get('shorthash', None) or item.get('hash'),
+ "preview": item.get('preview', None),
+ })
+ return res
+
+def get_interrogate():
+ from modules.ui_interrogate import get_models
+ return ['clip', 'deepdanbooru'] + get_models()
+
+def post_interrogate(req: models.InterrogateRequest):
+ if req.image is None or len(req.image) < 64:
+ raise HTTPException(status_code=404, detail="Image not found")
+ image = helpers.decode_base64_to_image(req.image)
+ image = image.convert('RGB')
+ if req.model == "clip":
+ caption = shared.interrogator.interrogate(image)
+ return models.InterrogateResponse(caption)
+ elif req.model == "deepdanbooru":
+ from mobules import deepbooru
+ caption = deepbooru.model.tag(image)
+ return models.InterrogateResponse(caption)
+ else:
+ from modules.ui_interrogate import interrogate_image, analyze_image, get_models
+ if req.model not in get_models():
+ raise HTTPException(status_code=404, detail="Model not found")
+ caption = interrogate_image(image, model=req.model, mode=req.mode)
+ if not req.analyze:
+ return models.InterrogateResponse(caption)
+ else:
+ medium, artist, movement, trending, flavor = analyze_image(image, model=req.model)
+ return models.InterrogateResponse(caption, medium, artist, movement, trending, flavor)
+
+def post_unload_checkpoint():
+ from modules import sd_models
+ sd_models.unload_model_weights(op='model')
+ sd_models.unload_model_weights(op='refiner')
+ return {}
+
+def post_reload_checkpoint():
+ from modules import sd_models
+ sd_models.reload_model_weights()
+ return {}
+
+def post_refresh_checkpoints():
+ return shared.refresh_checkpoints()
+
+def post_refresh_vae():
+ return shared.refresh_vaes()
+
+def get_extensions_list():
+ from modules import extensions
+ extensions.list_extensions()
+ ext_list = []
+ for ext in extensions.extensions:
+ ext: extensions.Extension
+ ext.read_info()
+ if ext.remote is not None:
+ ext_list.append({
+ "name": ext.name,
+ "remote": ext.remote,
+ "branch": ext.branch,
+ "commit_hash":ext.commit_hash,
+ "commit_date":ext.commit_date,
+ "version":ext.version,
+ "enabled":ext.enabled
+ })
+ return ext_list
+
+def post_pnginfo(req: models.PNGInfoRequest):
+ from modules import images, script_callbacks, generation_parameters_copypaste
+ if not req.image.strip():
+ return models.PNGInfoResponse(info="")
+ image = helpers.decode_base64_to_image(req.image.strip())
+ if image is None:
+ return models.PNGInfoResponse(info="")
+ geninfo, items = images.read_info_from_image(image)
+ if geninfo is None:
+ geninfo = ""
+ if items and items['parameters']:
+ del items['parameters']
+ params = generation_parameters_copypaste.parse_generation_parameters(geninfo)
+ script_callbacks.infotext_pasted_callback(geninfo, params)
+ return models.PNGInfoResponse(info=geninfo, items=items, parameters=params)
diff --git a/modules/api/helpers.py b/modules/api/helpers.py
new file mode 100644
index 000000000..050f8eb54
--- /dev/null
+++ b/modules/api/helpers.py
@@ -0,0 +1,57 @@
+import io
+import base64
+from PIL import Image, PngImagePlugin
+import piexif
+import piexif.helper
+from fastapi.exceptions import HTTPException
+from modules import shared, images
+
+
+def decode_base64_to_image(encoding):
+ if encoding.startswith("data:image/"):
+ encoding = encoding.split(";")[1].split(",")[1]
+ try:
+ image = Image.open(io.BytesIO(base64.b64decode(encoding)))
+ return image
+ except Exception as e:
+ shared.log.warning(f'API cannot decode image: {e}')
+ raise HTTPException(status_code=500, detail="Invalid encoded image") from e
+
+
+def encode_pil_to_base64(image):
+ with io.BytesIO() as output_bytes:
+ images.save_image(image, output_bytes, shared.opts.samples_format)
+ bytes_data = output_bytes.getvalue()
+ return base64.b64encode(bytes_data)
+
+def upscaler_to_index(name: str):
+ try:
+ return [x.name.lower() for x in shared.sd_upscalers].index(name.lower())
+ except Exception as e:
+ raise HTTPException(status_code=400, detail=f"Invalid upscaler, needs to be one of these: {' , '.join([x.name for x in shared.sd_upscalers])}") from e
+
+def save_image(image, fn, ext):
+ # actual save
+ parameters = image.info.get('parameters', None)
+ image_format = Image.registered_extensions()[f'.{ext}']
+ if image_format == 'PNG':
+ pnginfo_data = PngImagePlugin.PngInfo()
+ for k, v in image.info.items():
+ pnginfo_data.add_text(k, str(v))
+ image.save(fn, format=image_format, quality=shared.opts.jpeg_quality, pnginfo=pnginfo_data)
+ elif image_format == 'JPEG':
+ if image.mode == 'RGBA':
+ shared.log.warning('Saving RGBA image as JPEG: Alpha channel will be lost')
+ image = image.convert("RGB")
+ elif image.mode == 'I;16':
+ image = image.point(lambda p: p * 0.0038910505836576).convert("L")
+ exif_bytes = piexif.dump({ "Exif": { piexif.ExifIFD.UserComment: piexif.helper.UserComment.dump(parameters or "", encoding="unicode") } })
+ image.save(fn, format=image_format, quality=shared.opts.jpeg_quality, exif=exif_bytes)
+ elif image_format == 'WEBP':
+ if image.mode == 'I;16':
+ image = image.point(lambda p: p * 0.0038910505836576).convert("RGB")
+ exif_bytes = piexif.dump({ "Exif": { piexif.ExifIFD.UserComment: piexif.helper.UserComment.dump(parameters or "", encoding="unicode") } })
+ image.save(fn, format=image_format, quality=shared.opts.jpeg_quality, lossless=shared.opts.webp_lossless, exif=exif_bytes)
+ else:
+ # shared.log.warning(f'Unrecognized image format: {extension} attempting save as {image_format}')
+ image.save(fn, format=image_format, quality=shared.opts.jpeg_quality)
diff --git a/modules/middleware.py b/modules/api/middleware.py
similarity index 100%
rename from modules/middleware.py
rename to modules/api/middleware.py
diff --git a/modules/api/models.py b/modules/api/models.py
index dca498473..ad65ff8e2 100644
--- a/modules/api/models.py
+++ b/modules/api/models.py
@@ -91,14 +91,81 @@ class PydanticModelGenerator:
DynamicModel.__config__.allow_mutation = True
return DynamicModel
+### item classes
-class IPAdapterItem(BaseModel):
+class ItemSampler(BaseModel):
+ name: str = Field(title="Name")
+ aliases: List[str] = Field(title="Aliases")
+ options: Dict[str, str] = Field(title="Options")
+
+class ItemVae(BaseModel):
+ model_name: str = Field(title="Model Name")
+ filename: str = Field(title="Filename")
+
+class ItemUpscaler(BaseModel):
+ name: str = Field(title="Name")
+ model_name: Optional[str] = Field(title="Model Name")
+ model_path: Optional[str] = Field(title="Path")
+ model_url: Optional[str] = Field(title="URL")
+ scale: Optional[float] = Field(title="Scale")
+
+class ItemModel(BaseModel):
+ title: str = Field(title="Title")
+ model_name: str = Field(title="Model Name")
+ filename: str = Field(title="Filename")
+ type: str = Field(title="Model type")
+ sha256: Optional[str] = Field(title="SHA256 hash")
+ hash: Optional[str] = Field(title="Short hash")
+ config: Optional[str] = Field(title="Config file")
+
+class ItemHypernetwork(BaseModel):
+ name: str = Field(title="Name")
+ path: Optional[str] = Field(title="Path")
+
+class ItemFaceRestorer(BaseModel):
+ name: str = Field(title="Name")
+ cmd_dir: Optional[str] = Field(title="Path")
+
+class ItemGAN(BaseModel):
+ name: str = Field(title="Name")
+ path: Optional[str] = Field(title="Path")
+ scale: Optional[int] = Field(title="Scale")
+
+class ItemStyle(BaseModel):
+ name: str = Field(title="Name")
+ prompt: Optional[str] = Field(title="Prompt")
+ negative_prompt: Optional[str] = Field(title="Negative Prompt")
+ extra: Optional[str] = Field(title="Extra")
+ filename: Optional[str] = Field(title="Filename")
+ preview: Optional[str] = Field(title="Preview")
+
+class ItemExtraNetwork(BaseModel):
+ name: str = Field(title="Name")
+ type: str = Field(title="Type")
+ title: Optional[str] = Field(title="Title")
+ fullname: Optional[str] = Field(title="Fullname")
+ filename: Optional[str] = Field(title="Filename")
+ hash: Optional[str] = Field(title="Hash")
+ preview: Optional[str] = Field(title="Preview image URL")
+
+class ItemArtist(BaseModel):
+ name: str = Field(title="Name")
+ score: float = Field(title="Score")
+ category: str = Field(title="Category")
+
+class ItemEmbedding(BaseModel):
+ step: Optional[int] = Field(title="Step", description="The number of steps that were used to train this embedding, if available")
+ sd_checkpoint: Optional[str] = Field(title="SD Checkpoint", description="The hash of the checkpoint this embedding was trained on, if available")
+ sd_checkpoint_name: Optional[str] = Field(title="SD Checkpoint Name", description="The name of the checkpoint this embedding was trained on, if available. Note that this is the name that was used by the trainer; for a stable identifier, use `sd_checkpoint` instead")
+ shape: int = Field(title="Shape", description="The length of each individual vector in the embedding")
+ vectors: int = Field(title="Vectors", description="The number of vectors in the embedding")
+
+class ItemIPAdapter(BaseModel):
adapter: str = Field(title="Adapter", default="Base", description="Adapter to use")
image: str = Field(title="Image", default="", description="Adapter image, must be a base64 string containing the image's data.")
scale: float = Field(title="Scale", default=0.5, gt=0, le=1, description="Scale of the adapter image, must be between 0 and 1.")
-
-class FaceIDItem(BaseModel):
+class ItemFaceID(BaseModel):
mode: list[str] = Field(title="Mode", default=["FaceID"], description="The mode to use (available values: FaceID, FaceSwap).")
model: str = Field(title="Model", default="FaceID Base", description="The FaceID model to use.")
image: str = Field(title="Image", default="", description="Source face image, must be a base64 string containing the image's data.")
@@ -109,8 +176,32 @@ class FaceIDItem(BaseModel):
tokens: int = Field(title="Tokens", default=4, ge=1, le=16, description="Amount of tokens to use, must be between 1 and 16.")
cache_model: bool = Field(title="Cache", default=True, description="Should the model be cached?")
+class ScriptArg(BaseModel):
+ label: str = Field(default=None, title="Label", description="Name of the argument in UI")
+ value: Optional[Any] = Field(default=None, title="Value", description="Default value of the argument")
+ minimum: Optional[Any] = Field(default=None, title="Minimum", description="Minimum allowed value for the argumentin UI")
+ maximum: Optional[Any] = Field(default=None, title="Minimum", description="Maximum allowed value for the argumentin UI")
+ step: Optional[Any] = Field(default=None, title="Minimum", description="Step for changing value of the argumentin UI")
+ choices: Optional[Any] = Field(default=None, title="Choices", description="Possible values for the argument")
-StableDiffusionTxt2ImgProcessingAPI = PydanticModelGenerator(
+class ItemScript(BaseModel):
+ name: str = Field(default=None, title="Name", description="Script name")
+ is_alwayson: bool = Field(default=None, title="IsAlwayson", description="Flag specifying whether this script is an alwayson script")
+ is_img2img: bool = Field(default=None, title="IsImg2img", description="Flag specifying whether this script is an img2img script")
+ args: List[ScriptArg] = Field(title="Arguments", description="List of script's arguments")
+
+class ItemExtension(BaseModel):
+ name: str = Field(title="Name", description="Extension name")
+ remote: str = Field(title="Remote", description="Extension Repository URL")
+ branch: str = Field(title="Branch", description="Extension Repository Branch")
+ commit_hash: str = Field(title="Commit Hash", description="Extension Repository Commit Hash")
+ version: str = Field(title="Version", description="Extension Version")
+ commit_date: str = Field(title="Commit Date", description="Extension Repository Commit Date")
+ enabled: bool = Field(title="Enabled", description="Flag specifying whether this extension is enabled")
+
+### request/response classes
+
+ReqTxt2Img = PydanticModelGenerator(
"StableDiffusionProcessingTxt2Img",
StableDiffusionProcessingTxt2Img,
[
@@ -120,12 +211,17 @@ StableDiffusionTxt2ImgProcessingAPI = PydanticModelGenerator(
{"key": "send_images", "type": bool, "default": True},
{"key": "save_images", "type": bool, "default": False},
{"key": "alwayson_scripts", "type": dict, "default": {}},
- {"key": "ip_adapter", "type": Optional[IPAdapterItem], "default": None, "exclude": True},
- {"key": "face_id", "type": Optional[FaceIDItem], "default": None, "exclude": True},
+ {"key": "ip_adapter", "type": Optional[ItemIPAdapter], "default": None, "exclude": True},
+ {"key": "face_id", "type": Optional[ItemFaceID], "default": None, "exclude": True},
]
).generate_model()
-StableDiffusionImg2ImgProcessingAPI = PydanticModelGenerator(
+class ResTxt2Img(BaseModel):
+ images: List[str] = Field(default=None, title="Image", description="The generated image in base64 format.")
+ parameters: dict
+ info: str
+
+ReqImg2Img = PydanticModelGenerator(
"StableDiffusionProcessingImg2Img",
StableDiffusionProcessingImg2Img,
[
@@ -139,22 +235,21 @@ StableDiffusionImg2ImgProcessingAPI = PydanticModelGenerator(
{"key": "send_images", "type": bool, "default": True},
{"key": "save_images", "type": bool, "default": False},
{"key": "alwayson_scripts", "type": dict, "default": {}},
- {"key": "ip_adapter", "type": Optional[IPAdapterItem], "default": None, "exclude": True},
- {"key": "face_id", "type": Optional[FaceIDItem], "default": None, "exclude": True},
+ {"key": "ip_adapter", "type": Optional[ItemIPAdapter], "default": None, "exclude": True},
+ {"key": "face_id", "type": Optional[ItemFaceID], "default": None, "exclude": True},
]
).generate_model()
-class TextToImageResponse(BaseModel):
+class ResImg2Img(BaseModel):
images: List[str] = Field(default=None, title="Image", description="The generated image in base64 format.")
parameters: dict
info: str
-class ImageToImageResponse(BaseModel):
- images: List[str] = Field(default=None, title="Image", description="The generated image in base64 format.")
- parameters: dict
- info: str
+class FileData(BaseModel):
+ data: str = Field(title="File data", description="Base64 representation of the file")
+ name: str = Field(title="File name")
-class ExtrasBaseRequest(BaseModel):
+class ReqProcess(BaseModel):
resize_mode: float = Field(default=0, title="Resize Mode", description="Sets the resize mode: 0 to upscale by upscaling_resize amount, 1 to upscale up to upscaling_resize_h x upscaling_resize_w.")
show_extras_results: bool = Field(default=True, title="Show results", description="Should the backend return the generated image?")
gfpgan_visibility: float = Field(default=0, title="GFPGAN Visibility", ge=0, le=1, allow_inf_nan=False, description="Sets the visibility of GFPGAN, values should be between 0 and 1.")
@@ -169,61 +264,62 @@ class ExtrasBaseRequest(BaseModel):
extras_upscaler_2_visibility: float = Field(default=0, title="Secondary upscaler visibility", ge=0, le=1, allow_inf_nan=False, description="Sets the visibility of secondary upscaler, values should be between 0 and 1.")
upscale_first: bool = Field(default=False, title="Upscale first", description="Should the upscaler run before restoring faces?")
-class ExtraBaseResponse(BaseModel):
+class ResProcess(BaseModel):
html_info: str = Field(title="HTML info", description="A series of HTML tags containing the process info.")
-class ExtrasSingleImageRequest(ExtrasBaseRequest):
+class ReqProcessImage(ReqProcess):
image: str = Field(default="", title="Image", description="Image to work on, must be a Base64 string containing the image's data.")
-class ExtrasSingleImageResponse(ExtraBaseResponse):
+class ResProcessImage(ResProcess):
image: str = Field(default=None, title="Image", description="The generated image in base64 format.")
-class FileData(BaseModel):
- data: str = Field(title="File data", description="Base64 representation of the file")
- name: str = Field(title="File name")
-
-class ExtrasBatchImagesRequest(ExtrasBaseRequest):
+class ReqProcessBatch(ReqProcess):
imageList: List[FileData] = Field(title="Images", description="List of images to work on. Must be Base64 strings")
-class ExtrasBatchImagesResponse(ExtraBaseResponse):
+class ResProcessBatch(ResProcess):
images: List[str] = Field(title="Images", description="The generated images in base64 format.")
-class PNGInfoRequest(BaseModel):
+class ReqImageInfo(BaseModel):
image: str = Field(title="Image", description="The base64 encoded PNG image")
-class PNGInfoResponse(BaseModel):
+class ResImageInfo(BaseModel):
info: str = Field(title="Image info", description="A string with the parameters used to generate the image")
items: dict = Field(title="Items", description="A dictionary containing all the other fields the image had")
parameters: dict = Field(title="Parameters", description="A dictionary with parsed generation info fields")
-class LogRequest(BaseModel):
+class ReqLog(BaseModel):
lines: int = Field(default=100, title="Lines", description="How many lines to return")
clear: bool = Field(default=False, title="Clear", description="Should the log be cleared after returning the lines?")
-class ProgressRequest(BaseModel):
+class ReqProgress(BaseModel):
skip_current_image: bool = Field(default=False, title="Skip current image", description="Skip current image serialization")
-class ProgressResponse(BaseModel):
+class ResProgress(BaseModel):
progress: float = Field(title="Progress", description="The progress with a range of 0 to 1")
eta_relative: float = Field(title="ETA in secs")
state: dict = Field(title="State", description="The current state snapshot")
current_image: str = Field(default=None, title="Current image", description="The current image in base64 format. opts.show_progress_every_n_steps is required for this to work.")
textinfo: str = Field(default=None, title="Info text", description="Info text used by WebUI.")
-class InterrogateRequest(BaseModel):
+class ReqInterrogate(BaseModel):
image: str = Field(default="", title="Image", description="Image to work on, must be a Base64 string containing the image's data.")
model: str = Field(default="clip", title="Model", description="The interrogate model used.")
-class InterrogateResponse(BaseModel):
- caption: str = Field(default=None, title="Caption", description="The generated caption for the image.")
+class ResInterrogate(BaseModel):
+ caption: Optional[str] = Field(default=None, title="Caption", description="The generated caption for the image.")
+ medium: Optional[str] = Field(default=None, title="Medium", description="Image medium.")
+ artist: Optional[str] = Field(default=None, title="Medium", description="Image artist.")
+ movement: Optional[str] = Field(default=None, title="Medium", description="Image movement.")
+ trending: Optional[str] = Field(default=None, title="Medium", description="Image trending.")
+ flavor: Optional[str] = Field(default=None, title="Medium", description="Image flavor.")
-class TrainResponse(BaseModel):
+class ResTrain(BaseModel):
info: str = Field(title="Train info", description="Response string from train embedding or hypernetwork task.")
-class CreateResponse(BaseModel):
+class ResCreate(BaseModel):
info: str = Field(title="Create info", description="Response string from create embedding or hypernetwork task.")
-class PreprocessResponse(BaseModel):
+class ResPreprocess(BaseModel):
info: str = Field(title="Preprocess info", description="Response string from preprocessing task.")
fields = {}
@@ -251,109 +347,15 @@ for key in _options:
FlagsModel = create_model("Flags", **flags)
-class SamplerItem(BaseModel):
- name: str = Field(title="Name")
- aliases: List[str] = Field(title="Aliases")
- options: Dict[str, str] = Field(title="Options")
+class ResEmbeddings(BaseModel):
+ loaded: Dict[str, ItemEmbedding] = Field(title="Loaded", description="Embeddings loaded for the current model")
+ skipped: Dict[str, ItemEmbedding] = Field(title="Skipped", description="Embeddings skipped for the current model (likely due to architecture incompatibility)")
-class SDVaeItem(BaseModel):
- model_name: str = Field(title="Model Name")
- filename: str = Field(title="Filename")
-
-class UpscalerItem(BaseModel):
- name: str = Field(title="Name")
- model_name: Optional[str] = Field(title="Model Name")
- model_path: Optional[str] = Field(title="Path")
- model_url: Optional[str] = Field(title="URL")
- scale: Optional[float] = Field(title="Scale")
-
-class SDModelItem(BaseModel):
- title: str = Field(title="Title")
- model_name: str = Field(title="Model Name")
- filename: str = Field(title="Filename")
- type: str = Field(title="Model type")
- sha256: Optional[str] = Field(title="SHA256 hash")
- hash: Optional[str] = Field(title="Short hash")
- config: Optional[str] = Field(title="Config file")
-
-class HypernetworkItem(BaseModel):
- name: str = Field(title="Name")
- path: Optional[str] = Field(title="Path")
-
-class FaceRestorerItem(BaseModel):
- name: str = Field(title="Name")
- cmd_dir: Optional[str] = Field(title="Path")
-
-class RealesrganItem(BaseModel):
- name: str = Field(title="Name")
- path: Optional[str] = Field(title="Path")
- scale: Optional[int] = Field(title="Scale")
-
-class StyleItem(BaseModel):
- name: str = Field(title="Name")
- prompt: Optional[str] = Field(title="Prompt")
- negative_prompt: Optional[str] = Field(title="Negative Prompt")
- extra: Optional[str] = Field(title="Extra")
- filename: Optional[str] = Field(title="Filename")
- preview: Optional[str] = Field(title="Preview")
-
-class ExtraNetworkItem(BaseModel):
- name: str = Field(title="Name")
- type: str = Field(title="Type")
- title: Optional[str] = Field(title="Title")
- fullname: Optional[str] = Field(title="Fullname")
- filename: Optional[str] = Field(title="Filename")
- hash: Optional[str] = Field(title="Hash")
- preview: Optional[str] = Field(title="Preview image URL")
- # description: Optional[str] = Field(title="Description")
- # info: Optional[str] = Field(title="Information")
- # metadata: Optional[Any] = Field(title="Metadata")
- # local: Optional[str] = Field(title="Local")
-
-class ArtistItem(BaseModel):
- name: str = Field(title="Name")
- score: float = Field(title="Score")
- category: str = Field(title="Category")
-
-class EmbeddingItem(BaseModel):
- step: Optional[int] = Field(title="Step", description="The number of steps that were used to train this embedding, if available")
- sd_checkpoint: Optional[str] = Field(title="SD Checkpoint", description="The hash of the checkpoint this embedding was trained on, if available")
- sd_checkpoint_name: Optional[str] = Field(title="SD Checkpoint Name", description="The name of the checkpoint this embedding was trained on, if available. Note that this is the name that was used by the trainer; for a stable identifier, use `sd_checkpoint` instead")
- shape: int = Field(title="Shape", description="The length of each individual vector in the embedding")
- vectors: int = Field(title="Vectors", description="The number of vectors in the embedding")
-
-class EmbeddingsResponse(BaseModel):
- loaded: Dict[str, EmbeddingItem] = Field(title="Loaded", description="Embeddings loaded for the current model")
- skipped: Dict[str, EmbeddingItem] = Field(title="Skipped", description="Embeddings skipped for the current model (likely due to architecture incompatibility)")
-
-class MemoryResponse(BaseModel):
+class ResMemory(BaseModel):
ram: dict = Field(title="RAM", description="System memory stats")
cuda: dict = Field(title="CUDA", description="nVidia CUDA memory stats")
-class ScriptsList(BaseModel):
+class ResScripts(BaseModel):
txt2img: list = Field(default=None, title="Txt2img", description="Titles of scripts (txt2img)")
img2img: list = Field(default=None, title="Img2img", description="Titles of scripts (img2img)")
control: list = Field(default=None, title="Control", description="Titles of scripts (control)")
-
-class ScriptArg(BaseModel):
- label: str = Field(default=None, title="Label", description="Name of the argument in UI")
- value: Optional[Any] = Field(default=None, title="Value", description="Default value of the argument")
- minimum: Optional[Any] = Field(default=None, title="Minimum", description="Minimum allowed value for the argumentin UI")
- maximum: Optional[Any] = Field(default=None, title="Minimum", description="Maximum allowed value for the argumentin UI")
- step: Optional[Any] = Field(default=None, title="Minimum", description="Step for changing value of the argumentin UI")
- choices: Optional[Any] = Field(default=None, title="Choices", description="Possible values for the argument")
-
-class ScriptInfo(BaseModel):
- name: str = Field(default=None, title="Name", description="Script name")
- is_alwayson: bool = Field(default=None, title="IsAlwayson", description="Flag specifying whether this script is an alwayson script")
- is_img2img: bool = Field(default=None, title="IsImg2img", description="Flag specifying whether this script is an img2img script")
- args: List[ScriptArg] = Field(title="Arguments", description="List of script's arguments")
-
-class ExtensionItem(BaseModel):
- name: str = Field(title="Name", description="Extension name")
- remote: str = Field(title="Remote", description="Extension Repository URL")
- branch: str = Field(title="Branch", description="Extension Repository Branch")
- commit_hash: str = Field(title="Commit Hash", description="Extension Repository Commit Hash")
- version: str = Field(title="Version", description="Extension Version")
- commit_date: str = Field(title="Commit Date", description="Extension Repository Commit Date")
- enabled: bool = Field(title="Enabled", description="Flag specifying whether this extension is enabled")
diff --git a/modules/api/script.py b/modules/api/script.py
new file mode 100644
index 000000000..f516b0662
--- /dev/null
+++ b/modules/api/script.py
@@ -0,0 +1,82 @@
+from typing import Optional
+from fastapi.exceptions import HTTPException
+import gradio as gr
+from modules.api import models
+from modules import scripts
+
+
+def script_name_to_index(name, scripts_list):
+ try:
+ return [script.title().lower() for script in scripts_list].index(name.lower())
+ except Exception as e:
+ raise HTTPException(status_code=422, detail=f"Script '{name}' not found") from e
+
+def get_selectable_script(script_name, script_runner):
+ if script_name is None or script_name == "":
+ return None, None
+ script_idx = script_name_to_index(script_name, script_runner.selectable_scripts)
+ script = script_runner.selectable_scripts[script_idx]
+ return script, script_idx
+
+def get_scripts_list():
+ t2ilist = [script.name for script in scripts.scripts_txt2img.scripts if script.name is not None]
+ i2ilist = [script.name for script in scripts.scripts_img2img.scripts if script.name is not None]
+ control = [script.name for script in scripts.scripts_control.scripts if script.name is not None]
+ return models.ScriptsList(txt2img = t2ilist, img2img = i2ilist, control = control)
+
+def get_script_info(script_name: Optional[str] = None):
+ res = []
+ for script_list in [scripts.scripts_txt2img.scripts, scripts.scripts_img2img.scripts, scripts.scripts_control.scripts]:
+ for script in script_list:
+ if script.api_info is not None and (script_name is None or script_name == script.api_info.name):
+ res.append(script.api_info)
+ return res
+
+def get_script(script_name, script_runner):
+ if script_name is None or script_name == "":
+ return None, None
+ script_idx = script_name_to_index(script_name, script_runner.scripts)
+ return script_runner.scripts[script_idx]
+
+def init_default_script_args(script_runner):
+ #find max idx from the scripts in runner and generate a none array to init script_args
+ last_arg_index = 1
+ for script in script_runner.scripts:
+ if last_arg_index < script.args_to:
+ last_arg_index = script.args_to
+ # None everywhere except position 0 to initialize script args
+ script_args = [None]*last_arg_index
+ script_args[0] = 0
+
+ # get default values
+ if gr is None:
+ return script_args
+ with gr.Blocks(): # will throw errors calling ui function without this
+ for script in script_runner.scripts:
+ if script.ui(script.is_img2img):
+ ui_default_values = []
+ for elem in script.ui(script.is_img2img):
+ ui_default_values.append(elem.value)
+ script_args[script.args_from:script.args_to] = ui_default_values
+ return script_args
+
+def init_script_args(p, request, default_script_args, selectable_scripts, selectable_script_idx, script_runner):
+ script_args = default_script_args.copy()
+ # position 0 in script_arg is the idx+1 of the selectable script that is going to be run when using scripts.scripts_*2img.run()
+ if selectable_scripts:
+ script_args[selectable_scripts.args_from:selectable_scripts.args_to] = request.script_args
+ script_args[0] = selectable_script_idx + 1
+ # Now check for always on scripts
+ if request.alwayson_scripts and (len(request.alwayson_scripts) > 0):
+ for alwayson_script_name in request.alwayson_scripts.keys():
+ alwayson_script = get_script(alwayson_script_name, script_runner)
+ if alwayson_script is None:
+ raise HTTPException(status_code=422, detail=f"Always on script not found: {alwayson_script_name}")
+ if not alwayson_script.alwayson:
+ raise HTTPException(status_code=422, detail=f"Selectable script cannot be in always on params: {alwayson_script_name}")
+ if "args" in request.alwayson_scripts[alwayson_script_name]:
+ # min between arg length in scriptrunner and arg length in the request
+ for idx in range(0, min((alwayson_script.args_to - alwayson_script.args_from), len(request.alwayson_scripts[alwayson_script_name]["args"]))):
+ script_args[alwayson_script.args_from + idx] = request.alwayson_scripts[alwayson_script_name]["args"][idx]
+ p.per_script_args[alwayson_script.title()] = request.alwayson_scripts[alwayson_script_name]["args"]
+ return script_args
diff --git a/modules/api/server.py b/modules/api/server.py
new file mode 100644
index 000000000..d8bfcbb3f
--- /dev/null
+++ b/modules/api/server.py
@@ -0,0 +1,116 @@
+from typing import Any, Dict
+from fastapi import Depends
+from modules import shared
+from modules.api import models, helpers
+
+
+def post_shutdown():
+ shared.log.info('Shutdown request received')
+ import sys
+ sys.exit(0)
+
+def get_motd():
+ import requests
+ from installer import get_version
+ motd = ''
+ ver = get_version()
+ if ver.get('updated', None) is not None:
+ motd = f"version {ver['hash']} {ver['updated']} {ver['url'].split('/')[-1]}
"
+ if shared.opts.motd:
+ res = requests.get('https://vladmandic.github.io/automatic/motd', timeout=10)
+ if res.status_code == 200:
+ msg = (res.text or '').strip()
+ shared.log.info(f'MOTD: {msg if len(msg) > 0 else "N/A"}')
+ motd += res.text
+ return motd
+
+def get_log_buffer(req: models.LogRequest = Depends()):
+ lines = shared.log.buffer[:req.lines] if req.lines > 0 else shared.log.buffer.copy()
+ if req.clear:
+ shared.log.buffer.clear()
+ return lines
+
+def get_config():
+ options = {}
+ for k in shared.opts.data.keys():
+ if shared.opts.data_labels.get(k) is not None:
+ options.update({k: shared.opts.data.get(k, shared.opts.data_labels.get(k).default)})
+ else:
+ options.update({k: shared.opts.data.get(k, None)})
+ if 'sd_lyco' in options:
+ del options['sd_lyco']
+ if 'sd_lora' in options:
+ del options['sd_lora']
+ return options
+
+def set_config(req: Dict[str, Any]):
+ updated = []
+ for k, v in req.items():
+ updated.append({ k: shared.opts.set(k, v) })
+ shared.opts.save(shared.config_filename)
+ return { "updated": updated }
+
+def get_cmd_flags():
+ return vars(shared.cmd_opts)
+
+def get_progress(req: models.ProgressRequest = Depends()):
+ import time
+ if shared.state.job_count == 0:
+ return models.ProgressResponse(progress=0, eta_relative=0, state=shared.state.dict(), textinfo=shared.state.textinfo)
+ shared.state.do_set_current_image()
+ current_image = None
+ if shared.state.current_image and not req.skip_current_image:
+ current_image = helpers.encode_pil_to_base64(shared.state.current_image)
+ batch_x = max(shared.state.job_no, 0)
+ batch_y = max(shared.state.job_count, 1)
+ step_x = max(shared.state.sampling_step, 0)
+ step_y = max(shared.state.sampling_steps, 1)
+ current = step_y * batch_x + step_x
+ total = step_y * batch_y
+ progress = current / total if current > 0 and total > 0 else 0
+ time_since_start = time.time() - shared.state.time_start
+ eta_relative = (time_since_start / progress) - time_since_start if progress > 0 else 0
+ res = models.ProgressResponse(progress=progress, eta_relative=eta_relative, state=shared.state.dict(), current_image=current_image, textinfo=shared.state.textinfo)
+ return res
+
+def post_interrupt():
+ shared.state.interrupt()
+ return {}
+
+def post_skip():
+ shared.state.skip()
+
+def get_memory():
+ try:
+ import os
+ import psutil
+ process = psutil.Process(os.getpid())
+ res = process.memory_info() # only rss is cross-platform guaranteed so we dont rely on other values
+ ram_total = 100 * res.rss / process.memory_percent() # and total memory is calculated as actual value is not cross-platform safe
+ ram = { 'free': ram_total - res.rss, 'used': res.rss, 'total': ram_total }
+ except Exception as err:
+ ram = { 'error': f'{err}' }
+ try:
+ import torch
+ if torch.cuda.is_available():
+ s = torch.cuda.mem_get_info()
+ system = { 'free': s[0], 'used': s[1] - s[0], 'total': s[1] }
+ s = dict(torch.cuda.memory_stats(shared.device))
+ allocated = { 'current': s['allocated_bytes.all.current'], 'peak': s['allocated_bytes.all.peak'] }
+ reserved = { 'current': s['reserved_bytes.all.current'], 'peak': s['reserved_bytes.all.peak'] }
+ active = { 'current': s['active_bytes.all.current'], 'peak': s['active_bytes.all.peak'] }
+ inactive = { 'current': s['inactive_split_bytes.all.current'], 'peak': s['inactive_split_bytes.all.peak'] }
+ warnings = { 'retries': s['num_alloc_retries'], 'oom': s['num_ooms'] }
+ cuda = {
+ 'system': system,
+ 'active': active,
+ 'allocated': allocated,
+ 'reserved': reserved,
+ 'inactive': inactive,
+ 'events': warnings,
+ }
+ else:
+ cuda = { 'error': 'unavailable' }
+ except Exception as err:
+ cuda = { 'error': f'{err}' }
+ return models.MemoryResponse(ram = ram, cuda = cuda)
diff --git a/modules/api/train.py b/modules/api/train.py
new file mode 100644
index 000000000..8fb47fb1d
--- /dev/null
+++ b/modules/api/train.py
@@ -0,0 +1,90 @@
+from modules import shared, sd_hijack, devices
+from modules.api import models
+from modules.textual_inversion.preprocess import preprocess
+
+
+def post_create_embedding(args: dict):
+ from modules.textual_inversion.textual_inversion import create_embedding
+ try:
+ shared.state.begin('api-embedding')
+ filename = create_embedding(**args) # create empty embedding
+ sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings() # reload embeddings so new one can be immediately used
+ shared.state.end()
+ return models.CreateResponse(info = f"create embedding filename: {filename}")
+ except AssertionError as e:
+ shared.state.end()
+ return models.TrainResponse(info = f"create embedding error: {e}")
+
+def post_create_hypernetwork(args: dict):
+ from modules.hypernetworks.hypernetwork import create_hypernetwork
+ try:
+ shared.state.begin('api-hypernetwork')
+ filename = create_hypernetwork(**args) # create empty embedding # pylint: disable=E1111
+ shared.state.end()
+ return models.CreateResponse(info = f"create hypernetwork filename: {filename}")
+ except AssertionError as e:
+ shared.state.end()
+ return models.TrainResponse(info = f"create hypernetwork error: {e}")
+
+def post_preprocess(args: dict):
+ try:
+ shared.state.begin('api-preprocess')
+ preprocess(**args) # quick operation unless blip/booru interrogation is enabled
+ shared.state.end()
+ return models.PreprocessResponse(info = 'preprocess complete')
+ except KeyError as e:
+ shared.state.end()
+ return models.PreprocessResponse(info = f"preprocess error: invalid token: {e}")
+ except AssertionError as e:
+ shared.state.end()
+ return models.PreprocessResponse(info = f"preprocess error: {e}")
+ except FileNotFoundError as e:
+ shared.state.end()
+ return models.PreprocessResponse(info = f'preprocess error: {e}')
+
+def post_train_embedding(args: dict):
+ from modules.textual_inversion.textual_inversion import train_embedding
+ try:
+ shared.state.begin('api-embedding')
+ apply_optimizations = False
+ error = None
+ filename = ''
+ if not apply_optimizations:
+ sd_hijack.undo_optimizations()
+ try:
+ _embedding, filename = train_embedding(**args) # can take a long time to complete
+ except Exception as e:
+ error = e
+ finally:
+ if not apply_optimizations:
+ sd_hijack.apply_optimizations()
+ shared.state.end()
+ return models.TrainResponse(info = f"train embedding complete: filename: {filename} error: {error}")
+ except AssertionError as msg:
+ shared.state.end()
+ return models.TrainResponse(info = f"train embedding error: {msg}")
+
+def post_train_hypernetwork(args: dict):
+ from modules.hypernetworks.hypernetwork import train_hypernetwork
+ try:
+ shared.state.begin('api-hypernetwork')
+ shared.loaded_hypernetworks = []
+ apply_optimizations = False
+ error = None
+ filename = ''
+ if not apply_optimizations:
+ sd_hijack.undo_optimizations()
+ try:
+ _hypernetwork, filename = train_hypernetwork(**args)
+ except Exception as e:
+ error = e
+ finally:
+ shared.sd_model.cond_stage_model.to(devices.device)
+ shared.sd_model.first_stage_model.to(devices.device)
+ if not apply_optimizations:
+ sd_hijack.apply_optimizations()
+ shared.state.end()
+ return models.TrainResponse(info=f"train embedding complete: filename: {filename} error: {error}")
+ except AssertionError:
+ shared.state.end()
+ return models.TrainResponse(info=f"train embedding error: {error}")
diff --git a/modules/shared.py b/modules/shared.py
index 55f075601..854f4e2e5 100644
--- a/modules/shared.py
+++ b/modules/shared.py
@@ -27,6 +27,7 @@ from installer import log as central_logger # pylint: disable=E0611
errors.install([gr])
demo: gr.Blocks = None
+api = None
log = central_logger
progress_print_out = sys.stdout
parser = cmd_args.parser
diff --git a/modules/ui_interrogate.py b/modules/ui_interrogate.py
index 5a036fe6e..1707948c8 100644
--- a/modules/ui_interrogate.py
+++ b/modules/ui_interrogate.py
@@ -1,11 +1,7 @@
import os
-import base64
-from io import BytesIO
import gradio as gr
import torch
from PIL import Image
-from pydantic import BaseModel, Field # pylint: disable=no-name-in-module
-from fastapi.exceptions import HTTPException
import modules.generation_parameters_copypaste as parameters_copypaste
from modules import devices, lowvram, shared, paths, ui_common
@@ -29,7 +25,12 @@ class BatchWriter:
self.file.close()
-def load(clip_model_name):
+def get_models():
+ import open_clip
+ return ['/'.join(x) for x in open_clip.list_pretrained()]
+
+
+def load_interrogator(clip_model_name):
from clip_interrogator import Config, Interrogator
global ci # pylint: disable=global-statement
if ci is None:
@@ -54,23 +55,6 @@ def unload():
devices.torch_gc()
-def image_analysis(image, clip_model_name):
- load(clip_model_name)
- image = image.convert('RGB')
- image_features = ci.image_to_features(image)
- top_mediums = ci.mediums.rank(image_features, 5)
- top_artists = ci.artists.rank(image_features, 5)
- top_movements = ci.movements.rank(image_features, 5)
- top_trendings = ci.trendings.rank(image_features, 5)
- top_flavors = ci.flavors.rank(image_features, 5)
- medium_ranks = dict(zip(top_mediums, ci.similarities(image_features, top_mediums)))
- artist_ranks = dict(zip(top_artists, ci.similarities(image_features, top_artists)))
- movement_ranks = dict(zip(top_movements, ci.similarities(image_features, top_movements)))
- trending_ranks = dict(zip(top_trendings, ci.similarities(image_features, top_trendings)))
- flavor_ranks = dict(zip(top_flavors, ci.similarities(image_features, top_flavors)))
- return medium_ranks, artist_ranks, movement_ranks, trending_ranks, flavor_ranks
-
-
def interrogate(image, mode, caption=None):
shared.log.info(f'Interrogate: image={image} mode={mode} config={ci.config}')
if mode == 'best':
@@ -88,14 +72,14 @@ def interrogate(image, mode, caption=None):
return prompt
-def image_to_prompt(image, mode, clip_model_name):
+def interrogate_image(image, model, mode):
shared.state.begin()
shared.state.job = 'interrogate'
try:
if shared.cmd_opts.lowvram or shared.cmd_opts.medvram:
lowvram.send_everything_to_cpu()
devices.torch_gc()
- load(clip_model_name)
+ load_interrogator(model)
image = image.convert('RGB')
shared.log.info(f'Interrogate: image={image} mode={mode} config={ci.config}')
prompt = interrogate(image, mode)
@@ -106,12 +90,7 @@ def image_to_prompt(image, mode, clip_model_name):
return prompt
-def get_models():
- import open_clip
- return ['/'.join(x) for x in open_clip.list_pretrained()]
-
-
-def batch_process(batch_files, batch_folder, batch_str, mode, clip_model, write):
+def interrogate_batch(batch_files, batch_folder, batch_str, model, mode, write):
files = []
if batch_files is not None:
files += [f.name for f in batch_files]
@@ -122,7 +101,6 @@ def batch_process(batch_files, batch_folder, batch_str, mode, clip_model, write)
if len(files) == 0:
shared.log.error('Interrogate batch no images')
return ''
- shared.log.info(f'Interrogate batch: images={len(files)} mode={mode} config={ci.config}')
shared.state.begin()
shared.state.job = 'batch interrogate'
prompts = []
@@ -130,7 +108,8 @@ def batch_process(batch_files, batch_folder, batch_str, mode, clip_model, write)
if shared.cmd_opts.lowvram or shared.cmd_opts.medvram:
lowvram.send_everything_to_cpu()
devices.torch_gc()
- load(clip_model)
+ load_interrogator(model)
+ shared.log.info(f'Interrogate batch: images={len(files)} mode={mode} config={ci.config}')
captions = []
# first pass: generate captions
for file in files:
@@ -168,6 +147,23 @@ def batch_process(batch_files, batch_folder, batch_str, mode, clip_model, write)
return '\n\n'.join(prompts)
+def analyze_image(image, model):
+ load_interrogator(model)
+ image = image.convert('RGB')
+ image_features = ci.image_to_features(image)
+ top_mediums = ci.mediums.rank(image_features, 5)
+ top_artists = ci.artists.rank(image_features, 5)
+ top_movements = ci.movements.rank(image_features, 5)
+ top_trendings = ci.trendings.rank(image_features, 5)
+ top_flavors = ci.flavors.rank(image_features, 5)
+ medium_ranks = dict(zip(top_mediums, ci.similarities(image_features, top_mediums)))
+ artist_ranks = dict(zip(top_artists, ci.similarities(image_features, top_artists)))
+ movement_ranks = dict(zip(top_movements, ci.similarities(image_features, top_movements)))
+ trending_ranks = dict(zip(top_trendings, ci.similarities(image_features, top_trendings)))
+ flavor_ranks = dict(zip(top_flavors, ci.similarities(image_features, top_flavors)))
+ return medium_ranks, artist_ranks, movement_ranks, trending_ranks, flavor_ranks
+
+
def create_ui():
global low_vram # pylint: disable=global-statement
low_vram = shared.cmd_opts.lowvram or shared.cmd_opts.medvram
@@ -190,9 +186,9 @@ def create_ui():
trending = gr.Label(label="Trending", num_top_classes=5)
flavor = gr.Label(label="Flavor", num_top_classes=5)
with gr.Row():
- interrogate_btn = gr.Button("Interrogate", variant='primary')
- analyze_btn = gr.Button("Analyze", variant='primary')
- unload_btn = gr.Button("Unload")
+ btn_interrogate_img = gr.Button("Interrogate", variant='primary')
+ btn_analyze_img = gr.Button("Analyze", variant='primary')
+ btn_unload = gr.Button("Unload")
with gr.Row():
buttons = parameters_copypaste.create_buttons(["txt2img", "img2img", "extras", "control"])
for tabname, button in buttons.items():
@@ -209,7 +205,7 @@ def create_ui():
with gr.Row():
write = gr.Checkbox(label='Write prompts to files', value=False)
with gr.Row():
- batch_btn = gr.Button("Interrogate", variant='primary')
+ btn_interrogate_batch = gr.Button("Interrogate", variant='primary')
with gr.Column():
with gr.Row():
# clip_model = gr.Dropdown(get_models(), value='ViT-L-14/openai', label='CLIP Model')
@@ -217,54 +213,7 @@ def create_ui():
ui_common.create_refresh_button(clip_model, get_models, lambda: {"choices": get_models()}, 'refresh_interrogate_models')
with gr.Row():
mode = gr.Radio(['best', 'fast', 'classic', 'caption', 'negative'], label='Mode', value='best')
- interrogate_btn.click(image_to_prompt, inputs=[image, mode, clip_model], outputs=prompt)
- analyze_btn.click(image_analysis, inputs=[image, clip_model], outputs=[medium, artist, movement, trending, flavor])
- unload_btn.click(unload)
- batch_btn.click(batch_process, inputs=[batch_files, batch_folder, batch_str, mode, clip_model, write], outputs=[batch])
-
-
-def decode_base64_to_image(encoding):
- if encoding.startswith("data:image/"):
- encoding = encoding.split(";")[1].split(",")[1]
- try:
- image = Image.open(BytesIO(base64.b64decode(encoding)))
- return image
- except Exception as e:
- raise HTTPException(status_code=500, detail="Invalid encoded image") from e
-
-
-# TODO redesign interrogator api
-def mount_interrogator_api(_: gr.Blocks, app):
-
- class InterrogatorAnalyzeRequest(BaseModel):
- image: str = Field(default="", title="Image", description="Image to work on, must be a Base64 string containing the image's data.")
- clip_model_name: str = Field(default="ViT-L-14/openai", title="Model", description="The interrogate model used. See the models endpoint for a list of available models.")
-
- class InterrogatorPromptRequest(InterrogatorAnalyzeRequest):
- mode: str = Field(default="fast", title="Mode", description="The mode used to generate the prompt. Can be one of: best, fast, classic, negative.")
-
- @app.get("/interrogator/models")
- async def api_get_models():
- import open_clip
- return ["/".join(x) for x in open_clip.list_pretrained()]
-
- @app.post("/interrogator/prompt")
- async def api_get_prompt(analyzereq: InterrogatorPromptRequest):
- image_b64 = analyzereq.image
- if image_b64 is None:
- raise HTTPException(status_code=404, detail="Image not found")
- img = decode_base64_to_image(image_b64)
- prompt = image_to_prompt(img, analyzereq.mode, analyzereq.clip_model_name)
- return {"prompt": prompt}
-
- @app.post("/interrogator/analyze")
- async def api_analyze(analyzereq: InterrogatorAnalyzeRequest):
- image_b64 = analyzereq.image
- if image_b64 is None:
- raise HTTPException(status_code=404, detail="Image not found")
-
- img = decode_base64_to_image(image_b64)
- (medium_ranks, artist_ranks, movement_ranks, trending_ranks, flavor_ranks) = image_analysis(img, analyzereq.clip_model_name)
- return {"medium": medium_ranks, "artist": artist_ranks, "movement": movement_ranks, "trending": trending_ranks, "flavor": flavor_ranks}
-
-# script_callbacks.on_app_started(mount_interrogator_api)
+ btn_interrogate_img.click(interrogate_image, inputs=[image, clip_model, mode], outputs=prompt)
+ btn_analyze_img.click(analyze_image, inputs=[image, clip_model], outputs=[medium, artist, movement, trending, flavor])
+ btn_interrogate_batch.click(interrogate_batch, inputs=[batch_files, batch_folder, batch_str, clip_model, mode, write], outputs=[batch])
+ btn_unload.click(unload)
diff --git a/webui.py b/webui.py
index 4504b6ae1..f5e3ee086 100644
--- a/webui.py
+++ b/webui.py
@@ -31,7 +31,7 @@ import modules.upscaler
import modules.textual_inversion.textual_inversion
import modules.hypernetworks.hypernetwork
import modules.script_callbacks
-from modules.middleware import setup_middleware
+from modules.api.middleware import setup_middleware
from modules.shared import cmd_opts, opts
@@ -284,7 +284,7 @@ def start_ui():
timer.startup.record("launch")
modules.progress.setup_progress_api(app)
- create_api(app)
+ shared.api = create_api(app)
timer.startup.record("api")
ui_extra_networks.init_api(app)
@@ -347,12 +347,12 @@ def api_only():
from fastapi import FastAPI
app = FastAPI(**fastapi_args)
setup_middleware(app, cmd_opts)
- api = create_api(app)
- api.wants_restart = False
+ shared.api = create_api(app)
+ shared.api.wants_restart = False
modules.script_callbacks.app_started_callback(None, app)
modules.sd_models.write_metadata()
log.info(f"Startup time: {timer.startup.summary()}")
- server = api.launch()
+ server = shared.api.launch()
return server