cleanup api model defs

This commit is contained in:
Vladimir Mandic
2024-01-28 10:25:26 -05:00
parent 1ed97c7d89
commit b79d8ea6c3
9 changed files with 38 additions and 44 deletions
+4 -2
View File
@@ -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
+3 -1
View File
@@ -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)
+9 -9
View File
@@ -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)
+12
View File
@@ -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")
-20
View File
@@ -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])
+5 -5
View File
@@ -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)
+1 -1
View File
@@ -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,
+4 -4
View File
@@ -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'])
-2
View File
@@ -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