diff --git a/modules/api/api.py b/modules/api/api.py index c4bc4d49b..0a35a8fcf 100644 --- a/modules/api/api.py +++ b/modules/api/api.py @@ -9,10 +9,12 @@ from fastapi import FastAPI, APIRouter, Depends from fastapi.security import HTTPBasic, HTTPBasicCredentials from fastapi.exceptions import HTTPException from PIL import PngImagePlugin,Image + import piexif import piexif.helper import gradio as gr from modules import errors, shared, sd_samplers, deepbooru, sd_hijack, images, scripts, ui, postprocessing +from modules.sd_vae import vae_dict from modules.api import models from modules.processing import StableDiffusionProcessingTxt2Img, StableDiffusionProcessingImg2Img, process_images from modules.textual_inversion.textual_inversion import create_embedding, train_embedding @@ -132,6 +134,7 @@ class Api: self.add_api_route("/sdapi/v1/prompt-styles", self.get_prompt_styles, methods=["GET"], response_model=List[models.PromptStyleItem]) 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-vaes", 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) @@ -445,6 +448,10 @@ class Api: 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 [ diff --git a/modules/api/models.py b/modules/api/models.py index 704ebfe92..a14d6d299 100644 --- a/modules/api/models.py +++ b/modules/api/models.py @@ -232,6 +232,10 @@ class SamplerItem(BaseModel): aliases: List[str] = Field(title="Aliases") options: Dict[str, str] = Field(title="Options") +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")