diff options
author | Philpax <me@philpax.me> | 2023-01-01 23:17:33 +0000 |
---|---|---|
committer | Philpax <me@philpax.me> | 2023-01-01 23:18:11 +0000 |
commit | b5819d9bf1794071139c640b5f1e72c84a0e051a (patch) | |
tree | 94d5333b81db9c2958648ab8912ea343fdf741f7 /modules/api | |
parent | 311354c0bb8930ea939d6aa6b3edd50c69301320 (diff) | |
download | stable-diffusion-webui-gfx803-b5819d9bf1794071139c640b5f1e72c84a0e051a.tar.gz stable-diffusion-webui-gfx803-b5819d9bf1794071139c640b5f1e72c84a0e051a.tar.bz2 stable-diffusion-webui-gfx803-b5819d9bf1794071139c640b5f1e72c84a0e051a.zip |
feat(api): add /sdapi/v1/embeddings
Diffstat (limited to 'modules/api')
-rw-r--r-- | modules/api/api.py | 8 | ||||
-rw-r--r-- | modules/api/models.py | 3 |
2 files changed, 11 insertions, 0 deletions
diff --git a/modules/api/api.py b/modules/api/api.py index 11daff0d..30bf3dac 100644 --- a/modules/api/api.py +++ b/modules/api/api.py @@ -100,6 +100,7 @@ class Api: self.add_api_route("/sdapi/v1/prompt-styles", self.get_prompt_styles, methods=["GET"], response_model=List[PromptStyleItem]) self.add_api_route("/sdapi/v1/artist-categories", self.get_artists_categories, methods=["GET"], response_model=List[str]) self.add_api_route("/sdapi/v1/artists", self.get_artists, methods=["GET"], response_model=List[ArtistItem]) + self.add_api_route("/sdapi/v1/embeddings", self.get_embeddings, methods=["GET"], response_model=EmbeddingsResponse) self.add_api_route("/sdapi/v1/refresh-checkpoints", self.refresh_checkpoints, methods=["POST"]) self.add_api_route("/sdapi/v1/create/embedding", self.create_embedding, methods=["POST"], response_model=CreateResponse) self.add_api_route("/sdapi/v1/create/hypernetwork", self.create_hypernetwork, methods=["POST"], response_model=CreateResponse) @@ -327,6 +328,13 @@ class Api: def get_artists(self): return [{"name":x[0], "score":x[1], "category":x[2]} for x in shared.artist_db.artists] + def get_embeddings(self): + db = sd_hijack.model_hijack.embedding_db + return { + "loaded": sorted(db.word_embeddings.keys()), + "skipped": sorted(db.skipped_embeddings), + } + def refresh_checkpoints(self): shared.refresh_checkpoints() diff --git a/modules/api/models.py b/modules/api/models.py index c446ce7a..a8472dc9 100644 --- a/modules/api/models.py +++ b/modules/api/models.py @@ -249,3 +249,6 @@ class ArtistItem(BaseModel): score: float = Field(title="Score") category: str = Field(title="Category") +class EmbeddingsResponse(BaseModel): + loaded: List[str] = Field(title="Loaded", description="Embeddings loaded for the current model") + skipped: List[str] = Field(title="Skipped", description="Embeddings skipped for the current model (likely due to architecture incompatibility)")
\ No newline at end of file |