From b79d8ea6c3f66706340a474bba22f09800b0915f Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Sun, 28 Jan 2024 10:25:26 -0500 Subject: [PATCH] cleanup api model defs --- javascript/control.js | 6 ++++-- modules/api/api.py | 4 +++- modules/api/endpoints.py | 18 +++++++++--------- modules/api/models.py | 12 ++++++++++++ modules/api/nvml.py | 20 -------------------- modules/api/server.py | 10 +++++----- modules/scripts.py | 2 +- modules/ui_control.py | 8 ++++---- webui.py | 2 -- 9 files changed, 38 insertions(+), 44 deletions(-) diff --git a/javascript/control.js b/javascript/control.js index 56cc5dce6..2f218703a 100644 --- a/javascript/control.js +++ b/javascript/control.js @@ -29,8 +29,10 @@ async function setupControlUI() { const intersectionObserver = new IntersectionObserver((entries) => { if (entries[0].intersectionRatio > 0) { const tab = gradioApp().querySelector('#control-tabs > .tab-nav > .selected')?.innerText.toLowerCase() || ''; // selected tab name - const btn = gradioApp().getElementById(`refresh_${tab}_models`); - if (btn) btn.click(); + for (let i = 0; i < 10; i += 1) { + const btn = gradioApp().getElementById(`refresh_${tab}_models_${i}`); + if (btn) btn.click(); + } } }); intersectionObserver.observe(el); // monitor visibility of tab diff --git a/modules/api/api.py b/modules/api/api.py index 2cce1edf6..7334b247a 100644 --- a/modules/api/api.py +++ b/modules/api/api.py @@ -5,7 +5,7 @@ from fastapi import FastAPI, APIRouter, Depends, Request from fastapi.security import HTTPBasic, HTTPBasicCredentials from fastapi.exceptions import HTTPException from modules import errors, shared, sd_samplers, scripts, ui, postprocessing -from modules.api import models, endpoints, script, train, helpers, server +from modules.api import models, endpoints, script, train, helpers, server, nvml from modules.processing import StableDiffusionProcessingTxt2Img, StableDiffusionProcessingImg2Img, process_images @@ -41,6 +41,8 @@ class Api: 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) + app.add_api_route("/sdapi/v1/nvml", nvml.get_nvml, methods=["GET"], response_model=List[models.ResNVML]) + # core api using locking self.add_api_route("/sdapi/v1/txt2img", self.post_text2img, methods=["POST"], response_model=models.ResTxt2Img) diff --git a/modules/api/endpoints.py b/modules/api/endpoints.py index 8e9386cb1..23d5e48ef 100644 --- a/modules/api/endpoints.py +++ b/modules/api/endpoints.py @@ -71,28 +71,28 @@ def get_interrogate(): from modules.ui_interrogate import get_models return ['clip', 'deepdanbooru'] + get_models() -def post_interrogate(req: models.InterrogateRequest): +def post_interrogate(req: models.ReqInterrogate): 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) + return models.ResInterrogate(caption) elif req.model == "deepdanbooru": from mobules import deepbooru caption = deepbooru.model.tag(image) - return models.InterrogateResponse(caption) + return models.ResInterrogate(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) + return models.ResInterrogate(caption) else: medium, artist, movement, trending, flavor = analyze_image(image, model=req.model) - return models.InterrogateResponse(caption, medium, artist, movement, trending, flavor) + return models.ResInterrogate(caption, medium, artist, movement, trending, flavor) def post_unload_checkpoint(): from modules import sd_models @@ -130,13 +130,13 @@ def get_extensions_list(): }) return ext_list -def post_pnginfo(req: models.PNGInfoRequest): +def post_pnginfo(req: models.ReqImageInfo): from modules import images, script_callbacks, generation_parameters_copypaste if not req.image.strip(): - return models.PNGInfoResponse(info="") + return models.ResImageInfo(info="") image = helpers.decode_base64_to_image(req.image.strip()) if image is None: - return models.PNGInfoResponse(info="") + return models.ResImageInfo(info="") geninfo, items = images.read_info_from_image(image) if geninfo is None: geninfo = "" @@ -144,4 +144,4 @@ def post_pnginfo(req: models.PNGInfoRequest): 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) + return models.ResImageInfo(info=geninfo, items=items, parameters=params) diff --git a/modules/api/models.py b/modules/api/models.py index ad65ff8e2..e257e396b 100644 --- a/modules/api/models.py +++ b/modules/api/models.py @@ -215,6 +215,7 @@ ReqTxt2Img = PydanticModelGenerator( {"key": "face_id", "type": Optional[ItemFaceID], "default": None, "exclude": True}, ] ).generate_model() +StableDiffusionTxt2ImgProcessingAPI = ReqTxt2Img class ResTxt2Img(BaseModel): images: List[str] = Field(default=None, title="Image", description="The generated image in base64 format.") @@ -239,6 +240,7 @@ ReqImg2Img = PydanticModelGenerator( {"key": "face_id", "type": Optional[ItemFaceID], "default": None, "exclude": True}, ] ).generate_model() +StableDiffusionImg2ImgProcessingAPI = ReqImg2Img class ResImg2Img(BaseModel): images: List[str] = Field(default=None, title="Image", description="The generated image in base64 format.") @@ -359,3 +361,13 @@ 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 ResNVML(BaseModel): # definition of http response + name: str = Field(title="Name") + version: dict = Field(title="Version") + pci: dict = Field(title="Version") + memory: dict = Field(title="Version") + clock: dict = Field(title="Version") + load: dict = Field(title="Version") + power: list = [] + state: str = Field(title="State") diff --git a/modules/api/nvml.py b/modules/api/nvml.py index b7bb3ddd5..d1116240f 100644 --- a/modules/api/nvml.py +++ b/modules/api/nvml.py @@ -1,8 +1,3 @@ -from typing import List -from fastapi import FastAPI -from pydantic import BaseModel, Field # pylint: disable=no-name-in-module - - try: import pynvml as nv nvml_ok = True @@ -12,17 +7,6 @@ except ImportError: nvml_initialized = False -class NVMLRes(BaseModel): # definition of http response - name: str = Field(title="Name") - version: dict = Field(title="Version") - pci: dict = Field(title="Version") - memory: dict = Field(title="Version") - clock: dict = Field(title="Version") - load: dict = Field(title="Version") - power: list = [] - state: str = Field(title="State") - - def get_reason(val): throttle = { 1: 'gpu idle', @@ -91,7 +75,3 @@ def get_nvml(): # log.debug(f'nvml failed: {e}') nvml_ok = False return [] - - -def nvml_api(app: FastAPI): - app.add_api_route("/sdapi/v1/nvml", get_nvml, methods=["GET"], response_model=List[NVMLRes]) diff --git a/modules/api/server.py b/modules/api/server.py index d8bfcbb3f..6c7bfed71 100644 --- a/modules/api/server.py +++ b/modules/api/server.py @@ -24,7 +24,7 @@ def get_motd(): motd += res.text return motd -def get_log_buffer(req: models.LogRequest = Depends()): +def get_log_buffer(req: models.ReqLog = Depends()): lines = shared.log.buffer[:req.lines] if req.lines > 0 else shared.log.buffer.copy() if req.clear: shared.log.buffer.clear() @@ -53,10 +53,10 @@ def set_config(req: Dict[str, Any]): def get_cmd_flags(): return vars(shared.cmd_opts) -def get_progress(req: models.ProgressRequest = Depends()): +def get_progress(req: models.ReqProgress = 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) + return models.ResProgress(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: @@ -70,7 +70,7 @@ def get_progress(req: models.ProgressRequest = Depends()): 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) + res = models.ResProgress(progress=progress, eta_relative=eta_relative, state=shared.state.dict(), current_image=current_image, textinfo=shared.state.textinfo) return res def post_interrupt(): @@ -113,4 +113,4 @@ def get_memory(): cuda = { 'error': 'unavailable' } except Exception as err: cuda = { 'error': f'{err}' } - return models.MemoryResponse(ram = ram, cuda = cuda) + return models.ResMemory(ram = ram, cuda = cuda) diff --git a/modules/scripts.py b/modules/scripts.py index 498fc5db6..a7951f484 100644 --- a/modules/scripts.py +++ b/modules/scripts.py @@ -428,7 +428,7 @@ class ScriptRunner: setattr(arg_info, field, v) api_args.append(arg_info) - script.api_info = api_models.ScriptInfo( + script.api_info = api_models.ItemScript( name=script.name, is_img2img=script.is_img2img, is_alwayson=script.alwayson, diff --git a/modules/ui_control.py b/modules/ui_control.py index b1f4b5fb2..0f346ab6f 100644 --- a/modules/ui_control.py +++ b/modules/ui_control.py @@ -403,7 +403,7 @@ def create_ui(_blocks: gr.Blocks=None): enabled_cb = gr.Checkbox(value= i==0, label="") process_id = gr.Dropdown(label="Processor", choices=processors.list_models(), value='None') model_id = gr.Dropdown(label="ControlNet", choices=controlnet.list_models(), value='None') - ui_common.create_refresh_button(model_id, controlnet.list_models, lambda: {"choices": controlnet.list_models(refresh=True)}, 'refresh_controlnet_models') + ui_common.create_refresh_button(model_id, controlnet.list_models, lambda: {"choices": controlnet.list_models(refresh=True)}, f'refresh_controlnet_models_{i}') model_strength = gr.Slider(label="Strength", minimum=0.01, maximum=2.0, step=0.01, value=1.0-i/10) control_start = gr.Slider(label="Start", minimum=0.0, maximum=1.0, step=0.05, value=0) control_end = gr.Slider(label="End", minimum=0.0, maximum=1.0, step=0.05, value=1.0) @@ -457,7 +457,7 @@ def create_ui(_blocks: gr.Blocks=None): enabled_cb = gr.Checkbox(value= i == 0, label="Enabled") process_id = gr.Dropdown(label="Processor", choices=processors.list_models(), value='None') model_id = gr.Dropdown(label="Adapter", choices=t2iadapter.list_models(), value='None') - ui_common.create_refresh_button(model_id, t2iadapter.list_models, lambda: {"choices": t2iadapter.list_models(refresh=True)}, 'refresh_adapter_models') + ui_common.create_refresh_button(model_id, t2iadapter.list_models, lambda: {"choices": t2iadapter.list_models(refresh=True)}, f'refresh_adapter_models_{i}') model_strength = gr.Slider(label="Strength", minimum=0.01, maximum=1.0, step=0.01, value=1.0-i/10) reset_btn = ui_components.ToolButton(value=ui_symbols.reset) image_upload = gr.UploadButton(label=ui_symbols.upload, file_types=['image'], elem_classes=['form', 'gradio-button', 'tool']) @@ -497,7 +497,7 @@ def create_ui(_blocks: gr.Blocks=None): enabled_cb = gr.Checkbox(value= i==0, label="") process_id = gr.Dropdown(label="Processor", choices=processors.list_models(), value='None') model_id = gr.Dropdown(label="ControlNet-XS", choices=xs.list_models(), value='None') - ui_common.create_refresh_button(model_id, xs.list_models, lambda: {"choices": xs.list_models(refresh=True)}, 'refresh_xs_models') + ui_common.create_refresh_button(model_id, xs.list_models, lambda: {"choices": xs.list_models(refresh=True)}, f'refresh_xs_models_{i}') model_strength = gr.Slider(label="Strength", minimum=0.01, maximum=1.0, step=0.01, value=1.0-i/10) control_start = gr.Slider(label="Start", minimum=0.0, maximum=1.0, step=0.05, value=0) control_end = gr.Slider(label="End", minimum=0.0, maximum=1.0, step=0.05, value=1.0) @@ -540,7 +540,7 @@ def create_ui(_blocks: gr.Blocks=None): enabled_cb = gr.Checkbox(value= i == 0, label="Enabled") process_id = gr.Dropdown(label="Processor", choices=processors.list_models(), value='None') model_id = gr.Dropdown(label="Model", choices=lite.list_models(), value='None') - ui_common.create_refresh_button(model_id, lite.list_models, lambda: {"choices": lite.list_models(refresh=True)}, 'refresh_lite_models') + ui_common.create_refresh_button(model_id, lite.list_models, lambda: {"choices": lite.list_models(refresh=True)}, f'refresh_lite_models_{i}') model_strength = gr.Slider(label="Strength", minimum=0.01, maximum=1.0, step=0.01, value=1.0-i/10) reset_btn = ui_components.ToolButton(value=ui_symbols.reset) image_upload = gr.UploadButton(label=ui_symbols.upload, file_types=['image'], elem_classes=['form', 'gradio-button', 'tool']) diff --git a/webui.py b/webui.py index f5e3ee086..f89e20560 100644 --- a/webui.py +++ b/webui.py @@ -174,8 +174,6 @@ def create_api(app): log.debug('Creating API') from modules.api.api import Api api = Api(app, queue_lock) - from modules.api.nvml import nvml_api - nvml_api(api) return api