Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion scripts/daam/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,5 @@
from .utils import *
from .evaluate import *
from .experiment import *
from .trace import *
from .trace import *
from .api import *
27 changes: 27 additions & 0 deletions scripts/daam/api.py
Original file line number Diff line number Diff line change
@@ -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")
26 changes: 22 additions & 4 deletions scripts/daam_script.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,9 +20,10 @@
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
override_attention_texts: str | None = None

class Script(scripts.Script):

Expand Down Expand Up @@ -70,9 +71,19 @@ 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]

def process(self,
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 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,
Expand All @@ -88,6 +99,13 @@ 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."

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 = []
self.hide_images = hide_images
self.dont_save_images = dont_save_images
Expand Down