From 8b6ea6ef0b6138278602d23cdc40664f43318601 Mon Sep 17 00:00:00 2001 From: FxMorin <28154542+fxmorin@users.noreply.github.com> Date: Thu, 27 Feb 2025 00:21:25 -0500 Subject: [PATCH 1/2] Add API support --- scripts/daam/__init__.py | 3 ++- scripts/daam/api.py | 27 +++++++++++++++++++++++++++ scripts/daam_script.py | 9 ++++++--- 3 files changed, 35 insertions(+), 4 deletions(-) create mode 100644 scripts/daam/api.py diff --git a/scripts/daam/__init__.py b/scripts/daam/__init__.py index cc1779c..e169bc5 100644 --- a/scripts/daam/__init__.py +++ b/scripts/daam/__init__.py @@ -3,4 +3,5 @@ from .utils import * from .evaluate import * from .experiment import * -from .trace import * \ No newline at end of file +from .trace import * +from .api import * \ No newline at end of file diff --git a/scripts/daam/api.py b/scripts/daam/api.py new file mode 100644 index 0000000..e22e6d6 --- /dev/null +++ b/scripts/daam/api.py @@ -0,0 +1,27 @@ +from fastapi import FastAPI, Form +import gradio as gr + +api_attention_texts: str | None = None + + +def daam_api(_: gr.Blocks, app: FastAPI): + @app.post("/daam/v1/set-attention-text") + async def set_daam_attention_text( + texts: str = Form(description="Attention Texts for visualization") + ): + global api_attention_texts + api_attention_texts = None if len(texts) == 0 else texts + + @app.get("/daam/v1/get-attention-text") + async def return_daam_attention_text(): + return { + "texts": api_attention_texts + } + + +try: + from modules import script_callbacks + + script_callbacks.on_app_started(daam_api) +except: + print("[stable-diffusion-webui-daam] API failed to initialize") diff --git a/scripts/daam_script.py b/scripts/daam_script.py index be954a7..f435852 100644 --- a/scripts/daam_script.py +++ b/scripts/daam_script.py @@ -20,7 +20,7 @@ import modules.shared as shared from PIL import Image -from scripts.daam import trace, utils +from scripts.daam import trace, utils, api_attention_texts before_image_saved_handler = None @@ -70,8 +70,8 @@ def ui(self, is_img2img): self.tracers = None - return [attention_texts, hide_images, dont_save_images, hide_caption, use_grid, grid_layouyt, alpha, heatmap_image_scale, trace_each_layers, layers_as_row] - + return [attention_texts, hide_images, dont_save_images, hide_caption, use_grid, grid_layouyt, alpha, heatmap_image_scale, trace_each_layers, layers_as_row] + def process(self, p : StableDiffusionProcessing, attention_texts : str, @@ -88,6 +88,9 @@ def process(self, self.enabled = False # in case the assert fails assert opts.samples_save, "Cannot run Daam script. Enable 'Always save all generated images' setting." + if api_attention_texts: + attention_texts = api_attention_texts + self.images = [] self.hide_images = hide_images self.dont_save_images = dont_save_images From e528bb694b144da087feb9382cd5308ec5cc0889 Mon Sep 17 00:00:00 2001 From: FxMorin <28154542+fxmorin@users.noreply.github.com> Date: Thu, 27 Feb 2025 00:26:16 -0500 Subject: [PATCH 2/2] Add Override Settings for DAAM --- scripts/daam_script.py | 19 +++++++++++++++++-- 1 file changed, 17 insertions(+), 2 deletions(-) diff --git a/scripts/daam_script.py b/scripts/daam_script.py index f435852..b445b9c 100644 --- a/scripts/daam_script.py +++ b/scripts/daam_script.py @@ -23,6 +23,7 @@ from scripts.daam import trace, utils, api_attention_texts before_image_saved_handler = None +override_attention_texts: str | None = None class Script(scripts.Script): @@ -72,7 +73,17 @@ def ui(self, is_img2img): return [attention_texts, hide_images, dont_save_images, hide_caption, use_grid, grid_layouyt, alpha, heatmap_image_scale, trace_each_layers, layers_as_row] - def process(self, + def before_process(self, p, *args): + global override_attention_texts + if "daam-texts" in p.override_settings: + texts = p.override_settings["daam-texts"] + if texts is not None: + override_attention_texts = p.override_settings["daam-texts"] + del p.override_settings["daam-texts"] + return + override_attention_texts = None + + def process(self, p : StableDiffusionProcessing, attention_texts : str, hide_images : bool, @@ -88,7 +99,11 @@ def process(self, self.enabled = False # in case the assert fails assert opts.samples_save, "Cannot run Daam script. Enable 'Always save all generated images' setting." - if api_attention_texts: + global override_attention_texts + if override_attention_texts: + attention_texts = override_attention_texts + override_attention_texts = None + elif api_attention_texts: attention_texts = api_attention_texts self.images = []