diff --git a/CHANGELOG.md b/CHANGELOG.md index 614a2a0ab..1c47398af 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,8 +8,10 @@ Mostly a service release - default is standard `torch.no_grad` new option is `torch.inference_only` which is slightly faster and uses less vram, but only works on some gpus - updated gradio +- styles support for subfolders - clean-up logging - capture system info in startup log + - better diagnostic output - capture extension output - capture ldm output - cleaner server restart diff --git a/extensions-builtin/sd-webui-controlnet b/extensions-builtin/sd-webui-controlnet index 2b12b2760..b15636ed3 160000 --- a/extensions-builtin/sd-webui-controlnet +++ b/extensions-builtin/sd-webui-controlnet @@ -1 +1 @@ -Subproject commit 2b12b2760dcde0e29c0d33640739ba3fe2c0cd32 +Subproject commit b15636ed35eff934af69985bcdfbc407cfedfe7d diff --git a/installer.py b/installer.py index 9111d0dcc..f4c4e0a3b 100644 --- a/installer.py +++ b/installer.py @@ -760,6 +760,28 @@ def check_extensions(): return round(newest_all) +def get_version(): + version = None + if version is None: + try: + res = subprocess.run('git log --pretty=format:"%h %ad" -1 --date=short', stdout = subprocess.PIPE, stderr = subprocess.PIPE, shell=True, check=True) + ver = res.stdout.decode(encoding = 'utf8', errors='ignore') if len(res.stdout) > 0 else ' ' + githash, updated = ver.split(' ') + res = subprocess.run('git remote get-url origin', stdout = subprocess.PIPE, stderr = subprocess.PIPE, shell=True, check=True) + origin = res.stdout.decode(encoding = 'utf8', errors='ignore') if len(res.stdout) > 0 else '' + res = subprocess.run('git rev-parse --abbrev-ref HEAD', stdout = subprocess.PIPE, stderr = subprocess.PIPE, shell=True, check=True) + branch_name = res.stdout.decode(encoding = 'utf8', errors='ignore') if len(res.stdout) > 0 else '' + version = { + 'app': 'sd.next', + 'updated': updated, + 'hash': githash, + 'url': origin.replace('\n', '') + '/tree/' + branch_name.replace('\n', '') + } + except Exception: + version = { 'app': 'sd.next', 'version': 'unknown' } + return version + + # check version of the main repo and optionally upgrade it def check_version(offline=False, reset=True): # pylint: disable=unused-argument if args.skip_all: @@ -768,8 +790,7 @@ def check_version(offline=False, reset=True): # pylint: disable=unused-argument log.error('Not a git repository') if not args.ignore: sys.exit(1) - ver = git('log -1 --pretty=format:"%h %ad"') - log.info(f'Version: {ver}') + log.info(f'Version: {print_dict(get_version())}') if args.version or args.skip_git: return commit = git('rev-parse HEAD') diff --git a/javascript/black-teal.css b/javascript/black-teal.css index 207e5f2f8..aa79023e6 100644 --- a/javascript/black-teal.css +++ b/javascript/black-teal.css @@ -77,7 +77,7 @@ svg.feather.feather-image, .feather .feather-image { display: none } .py-6 { padding-bottom: 0; } .tabs { background-color: var(--background-color); } .block.token-counter span { background-color: var(--input-background-fill) !important; box-shadow: 2px 2px 2px #111; border: none !important; font-size: 0.8rem; } -.tab-nav { zoom: 120%; margin-bottom: 10px; border-bottom: 2px solid var(--highlight-color) !important; padding-bottom: 2px; } +.tab-nav { zoom: 120%; margin-top: 10px; margin-bottom: 10px; border-bottom: 2px solid var(--highlight-color) !important; padding-bottom: 2px; } .label-wrap { margin: 16px 0px 8px 0px; } .gradio-slider input[type="number"] { width: 4.5em; font-size: 0.8rem; height: 20px; } .gradio-button.tool { border: none; background: none; box-shadow: none; filter: hue-rotate(340deg) saturate(0.5); } diff --git a/javascript/extraNetworks.js b/javascript/extraNetworks.js index 9f9d6af88..65931e286 100644 --- a/javascript/extraNetworks.js +++ b/javascript/extraNetworks.js @@ -284,7 +284,7 @@ function extraNetworksSearchButton(event) { updateInput(searchTextarea); } -function extraNetworksRefreshButton() { +function getENActivePage() { const tabname = getENActiveTab(); const page = gradioApp().querySelector(`#${tabname}_extra_networks > .tabs > .tab-nav > .selected`); return page ? page.innerText : ''; diff --git a/javascript/light-teal.css b/javascript/light-teal.css index a7f285951..accdeea16 100644 --- a/javascript/light-teal.css +++ b/javascript/light-teal.css @@ -77,7 +77,7 @@ svg.feather.feather-image, .feather .feather-image { display: none } .py-6 { padding-bottom: 0; } .tabs { background-color: var(--background-color); } .block.token-counter span { background-color: var(--input-background-fill) !important; box-shadow: 2px 2px 2px #111; border: none !important; font-size: 0.8rem; } -.tab-nav { zoom: 120%; margin-bottom: 10px; border-bottom: 2px solid var(--highlight-color) !important; padding-bottom: 2px; } +.tab-nav { zoom: 120%; margin-top: 10px; margin-bottom: 10px; border-bottom: 2px solid var(--highlight-color) !important; padding-bottom: 2px; } .label-wrap { margin: 16px 0px 8px 0px; } .gradio-slider input[type="number"] { width: 4.5em; font-size: 0.8rem; height: 20px; } .gradio-button.tool { border: none; background: none; box-shadow: none; filter: hue-rotate(340deg) saturate(0.5); } diff --git a/javascript/promptBracketChecker.js b/javascript/promptChecker.js similarity index 84% rename from javascript/promptBracketChecker.js rename to javascript/promptChecker.js index f5aa1f79a..66fa4442d 100644 --- a/javascript/promptBracketChecker.js +++ b/javascript/promptChecker.js @@ -3,15 +3,17 @@ // Counts open and closed brackets (round, square, curly) in the prompt and negative prompt text boxes in the txt2img and img2img tabs. // If there's a mismatch, the keyword counter turns red and if you hover on it, a tooltip tells you what's wrong. +let promptCheckerInitialized = false; + function checkBrackets(textArea, counterElt) { const counts = {}; - (textArea.value.match(/[(){}[\]]/g) || []).forEach((bracket) => { counts[bracket] = (counts[bracket] || 0) + 1; }); const errors = []; function checkPair(open, close, kind) { if (counts[open] !== counts[close]) errors.push(`${open}...${close} - Detected ${counts[open] || 0} opening and ${counts[close] || 0} closing ${kind}.`); } + (textArea.value.match(/[(){}[\]]/g) || []).forEach((bracket) => { counts[bracket] = (counts[bracket] || 0) + 1; }); checkPair('(', ')', 'round brackets'); checkPair('[', ']', 'square brackets'); checkPair('{', '}', 'curly brackets'); @@ -22,10 +24,14 @@ function checkBrackets(textArea, counterElt) { function setupBracketChecking(idPrompt, idCounter) { const textarea = gradioApp().querySelector(`#${idPrompt} > label > textarea`); const counter = gradioApp().getElementById(idCounter); - if (textarea && counter) textarea.addEventListener('input', () => checkBrackets(textarea, counter)); + if (!textarea || !counter) return; + if (!promptCheckerInitialized) log('promptChecker'); + promptCheckerInitialized = true; + textarea.addEventListener('input', () => checkBrackets(textarea, counter)); } onAfterUiUpdate(() => { + if (promptCheckerInitialized) return; setupBracketChecking('txt2img_prompt', 'txt2img_token_counter'); setupBracketChecking('txt2img_neg_prompt', 'txt2img_negative_token_counter'); setupBracketChecking('img2img_prompt', 'img2img_token_counter'); diff --git a/javascript/style.css b/javascript/style.css index 381598f15..b1255e14d 100644 --- a/javascript/style.css +++ b/javascript/style.css @@ -19,7 +19,7 @@ div.gradio-html.min{ min-height: 0; } .gradio-dropdown label span:not(.has-info), .gradio-textbox label span:not(.has-info), .gradio-number label span:not(.has-info) { margin-bottom: 0; } .gradio-dropdown ul.options li.item { padding: 0.05em 0; } .gradio-dropdown ul.options li.item:not(:has(.hide)) { background-color: var(--neutral-100); } -.gradio-dropdown ul.options{ z-index: 3000; min-width: fit-content; max-width: inherit; white-space: nowrap; } +.gradio-dropdown ul.options { z-index: 3000; min-width: fit-content; max-width: inherit; max-height: 25vh !important; white-space: nowrap; } .gradio-dropdown:not(.multiselect) .wrap-inner.wrap-inner.wrap-inner{ flex-wrap: unset; } .gradio-dropdown.multiselect .token-remove.remove-all.remove-all{ display: flex; } .gradio-dropdown.multiselect div.wrap-inner { overflow-x: hidden; overflow-y: auto; max-height: 50vh; overflow-wrap: anywhere; } diff --git a/launch.py b/launch.py index 488ff8366..1b2106992 100644 --- a/launch.py +++ b/launch.py @@ -40,7 +40,7 @@ def get_custom_args(): current = getattr(args, arg) if current != default: custom[arg] = getattr(args, arg) - installer.log.info(f'Command line args: {installer.print_dict(custom)}') + installer.log.info(f'Command line args: {sys.argv[1:]} {installer.print_dict(custom)}') @lru_cache() @@ -121,7 +121,7 @@ def get_memory_stats(): process = psutil.Process(os.getpid()) res = process.memory_info() ram_total = 100 * res.rss / process.memory_percent() - return f'used: {gb(res.rss)} total: {gb(ram_total)}' + return f'used={gb(res.rss)} total={gb(ram_total)}' def start_server(immediate=True, server=None): @@ -143,7 +143,6 @@ def start_server(immediate=True, server=None): # installer.log.debug(f'Loading module: {module_spec}') server = importlib.util.module_from_spec(module_spec) installer.log.debug(f'Starting module: {server}') - installer.log.info(f"Server arguments: {sys.argv[1:]}") get_custom_args() module_spec.loader.exec_module(server) uvicorn = None @@ -218,9 +217,10 @@ if __name__ == "__main__": except Exception: alive = False requests = 0 - if round(time.time()) % 120 == 0: - state = f'job="{instance.state.job}" {instance.state.job_no}/{instance.state.job_count}' - installer.log.debug(f'Server alive={alive} requests={requests} memory {get_memory_stats()} {state}') + if round(time.time()) % 10 == 0: + state = f'job="{instance.state.job}" {instance.state.job_no}/{instance.state.job_count}' if instance.state.job != '' or instance.state.job_no != 0 or instance.state.job_count != 0 else 'idle' + uptime = round(time.time() - instance.state.server_start) + installer.log.debug(f'Server alive={alive} jobs={instance.state.total_jobs} requests={requests} uptime={uptime}s memory {get_memory_stats()} {state}') if not alive: if uv is not None and uv.wants_restart: installer.log.info('Server restarting...') diff --git a/modules/api/api.py b/modules/api/api.py index 39436d14a..72fdf0829 100644 --- a/modules/api/api.py +++ b/modules/api/api.py @@ -131,7 +131,7 @@ class Api: self.add_api_route("/sdapi/v1/hypernetworks", self.get_hypernetworks, methods=["GET"], response_model=List[models.HypernetworkItem]) self.add_api_route("/sdapi/v1/face-restorers", self.get_face_restorers, methods=["GET"], response_model=List[models.FaceRestorerItem]) self.add_api_route("/sdapi/v1/realesrgan-models", self.get_realesrgan_models, methods=["GET"], response_model=List[models.RealesrganItem]) - 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/prompt-styles", self.get_prompt_styles, methods=["GET"], response_model=List[models.StyleItem]) 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]) @@ -479,10 +479,8 @@ class Api: def get_prompt_styles(self): styleList = [] - for k in shared.prompt_styles.styles: - style = shared.prompt_styles.styles[k] - styleList.append({"name":style[0], "prompt": style[1], "negative_prompt": style[2]}) - + for _k, v in shared.prompt_styles.styles.items(): + styleList.append(v) return styleList def get_embeddings(self): diff --git a/modules/api/models.py b/modules/api/models.py index 142d4db55..da4158dcd 100644 --- a/modules/api/models.py +++ b/modules/api/models.py @@ -264,10 +264,13 @@ class RealesrganItem(BaseModel): path: Optional[str] = Field(title="Path") scale: Optional[int] = Field(title="Scale") -class PromptStyleItem(BaseModel): +class StyleItem(BaseModel): name: str = Field(title="Name") prompt: Optional[str] = Field(title="Prompt") negative_prompt: Optional[str] = Field(title="Negative Prompt") + extra: Optional[str] = Field(title="Extra") + filename: Optional[str] = Field(title="Filename") + preview: Optional[str] = Field(title="Preview") class ArtistItem(BaseModel): name: str = Field(title="Name") diff --git a/modules/modelloader.py b/modules/modelloader.py index 152465941..afd2bcac3 100644 --- a/modules/modelloader.py +++ b/modules/modelloader.py @@ -434,6 +434,6 @@ def load_upscalers(): datas += scaler.scalers shared.sd_upscalers = sorted( datas, - # Special case for UpscalerNone keeps it at the beginning of the list. - key=lambda x: x.name.lower() if not isinstance(x.scaler, (UpscalerNone, UpscalerLanczos, UpscalerNearest)) else "" + key=lambda x: x.name.lower() if not isinstance(x.scaler, (UpscalerNone, UpscalerLanczos, UpscalerNearest)) else "" # Special case for UpscalerNone keeps it at the beginning of the list. ) + shared.log.debug(f"Loaded upscalers: items={len(shared.sd_upscalers)}") diff --git a/modules/paths_internal.py b/modules/paths_internal.py index 44afc4ad3..83c097e46 100644 --- a/modules/paths_internal.py +++ b/modules/paths_internal.py @@ -18,4 +18,4 @@ cmd_opts_pre = parser_pre.parse_known_args()[0] data_path = cmd_opts_pre.data_dir models_path = cmd_opts_pre.models_dir if os.path.isabs(cmd_opts_pre.models_dir) else os.path.join(data_path, cmd_opts_pre.models_dir) extensions_dir = os.path.join(data_path, "extensions") -extensions_builtin_dir = os.path.join(script_path, "extensions-builtin") +extensions_builtin_dir = "extensions-builtin" diff --git a/modules/script_loading.py b/modules/script_loading.py index 49489c4ed..64f16e681 100644 --- a/modules/script_loading.py +++ b/modules/script_loading.py @@ -19,7 +19,7 @@ def load_module(path): setup_logging() # reset since scripts can hijaack logging for line in stdout.getvalue().splitlines(): if len(line) > 0: - errors.log.info(f'Extension: script={os.path.relpath(path)} {line.strip()}') + errors.log.info(f"Extension: script='{os.path.relpath(path)}' {line.strip()}") except Exception as e: errors.display(e, f'Module load: {path}') return module diff --git a/modules/sd_samplers.py b/modules/sd_samplers.py index e0f6493c8..61a56c82a 100644 --- a/modules/sd_samplers.py +++ b/modules/sd_samplers.py @@ -23,7 +23,7 @@ def list_samplers(backend_name = shared.backend): samplers = all_samplers samplers_for_img2img = all_samplers samplers_map = {} - shared.log.debug(f'Available samplers: {[x.name for x in all_samplers]}') + # shared.log.debug(f'Available samplers: {[x.name for x in all_samplers]}') def find_sampler_config(name): diff --git a/modules/shared.py b/modules/shared.py index bf5f46933..292c60b4d 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -94,6 +94,7 @@ class State: job = "" job_no = 0 job_count = 0 + total_jobs = 0 processing_has_refined_job_count = False job_timestamp = '0' sampling_step = 0 @@ -141,6 +142,7 @@ class State: return obj def begin(self, title=""): + self.total_jobs += 1 self.current_image = None self.current_image_sampling_step = 0 self.current_latent = None diff --git a/modules/styles.py b/modules/styles.py index a959fac89..82ad05593 100644 --- a/modules/styles.py +++ b/modules/styles.py @@ -3,22 +3,18 @@ from __future__ import annotations import csv import os import json -import shutil -import typing from installer import log from modules import paths -if typing.TYPE_CHECKING: - # Only import this when code is being type-checked, it doesn't have any effect at runtime - from .processing import StableDiffusionProcessing - - -class PromptStyle(typing.NamedTuple): - name: str - prompt: str - negative_prompt: str - extra: str = "" +class Style(): + def __init__(self, name: str, prompt: str = "", negative_prompt: str = "", extra: str = "", filename: str = "", preview: str = ""): + self.name = name + self.prompt = prompt + self.negative_prompt = negative_prompt + self.extra = extra + self.filename = filename + self.preview = preview def merge_prompts(style_prompt: str, prompt: str) -> str: @@ -43,7 +39,7 @@ def apply_styles_to_prompt(prompt, styles): class StyleDatabase: def __init__(self, opts): - self.no_style = PromptStyle("None", "", "") + self.no_style = Style("None") self.styles = {} self.path = opts.styles_dir if os.path.isfile(opts.styles_dir): @@ -54,8 +50,8 @@ class StyleDatabase: self.mkdir() self.save_styles(opts.styles_dir, verbose=True) log.debug(f'Migrated styles: file={legacy_file} folder={self.path}') + self.reload() self.mkdir() - self.reload() def mkdir(self): if not os.path.isdir(self.path): @@ -64,15 +60,21 @@ class StyleDatabase: def reload(self): self.styles.clear() - for fn in os.listdir(self.path): - if not fn.lower().endswith(".json"): - continue - with open(os.path.join(self.path, fn), 'r', encoding='utf-8') as f: - try: - style = json.load(f) - self.styles[style["name"]] = PromptStyle(style["name"], style["prompt"], style["negative"], style["extra"]) - except Exception as e: - log.error(f'Failed to load style: file={fn} error={e}') + def list_folder(folder): + for filename in os.listdir(folder): + fn = os.path.join(folder, filename) + if os.path.isfile(fn) and fn.lower().endswith(".json"): + with open(fn, 'r', encoding='utf-8') as f: + try: + style = json.load(f) + fn = os.path.splitext(os.path.relpath(fn, self.path))[0] + self.styles[style["name"]] = Style(style["name"], style.get("prompt", ""), style.get("negative", ""), style.get("extra", ""), fn, style.get("preview", "")) + except Exception as e: + log.error(f'Failed to load style: file={fn} error={e}') + elif os.path.isdir(fn): + list_folder(fn) + + list_folder(self.path) log.debug(f'Loaded styles: folder={self.path} items={len(self.styles.keys())}') def get_style_prompts(self, styles): @@ -94,6 +96,7 @@ class StyleDatabase: "prompt": self.styles[name].prompt, "negative": self.styles[name].negative_prompt, "extra": "", + "preview": "", } fn = os.path.join(path, name + ".json") try: @@ -110,13 +113,12 @@ class StyleDatabase: reader = csv.DictReader(file, skipinitialspace=True) for row in reader: try: - prompt = row["prompt"] if "prompt" in row else row["text"] - negative_prompt = row.get("negative_prompt", "") - self.styles[row["name"]] = PromptStyle(row["name"], prompt, negative_prompt) + self.styles[row["name"]] = Style(row["name"], row["prompt"] if "prompt" in row else row["text"], row.get("negative_prompt", "")) except Exception: log.error(f'Styles error: file={legacy_file} row={row}') log.debug(f'Loaded legacy styles: file={legacy_file} items={len(self.styles.keys())}') + """ def save_csv(self, path: str) -> None: import tempfile basedir = os.path.dirname(path) @@ -124,8 +126,9 @@ class StyleDatabase: os.makedirs(basedir, exist_ok=True) fd, temp_path = tempfile.mkstemp(".csv") with os.fdopen(fd, "w", encoding="utf-8-sig", newline='') as file: - writer = csv.DictWriter(file, fieldnames=PromptStyle._fields) + writer = csv.DictWriter(file, fieldnames=Style._fields) writer.writeheader() writer.writerows(style._asdict() for k, style in self.styles.items()) log.debug(f'Saved legacy styles: {path} {len(self.styles.keys())}') shutil.move(temp_path, path) + """ diff --git a/modules/ui.py b/modules/ui.py index cd5c3b3eb..ad3754416 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -73,7 +73,7 @@ def send_gradio_gallery_to_image(x): def add_style(name: str, prompt: str, negative_prompt: str): if name is None: return [gr_show() for x in range(4)] - style = modules.styles.PromptStyle(name, prompt, negative_prompt) + style = modules.styles.Style(name, prompt, negative_prompt) modules.shared.prompt_styles.styles[style.name] = style modules.shared.prompt_styles.save_styles(modules.shared.opts.styles_dir) return [gr.Dropdown.update(visible=True, choices=list(modules.shared.prompt_styles.styles)) for _ in range(2)] @@ -235,11 +235,11 @@ def create_toprow(is_img2img): with gr.Row(): with gr.Column(scale=80): with gr.Row(): - prompt = gr.Textbox(label="Prompt", elem_id=f"{id_part}_prompt", show_label=False, lines=3, placeholder="Prompt (press Ctrl+Enter or Alt+Enter to generate)", elem_classes=["prompt"]) + prompt = gr.Textbox(elem_id=f"{id_part}_prompt", show_label=False, lines=3, placeholder="Prompt", elem_classes=["prompt"]) with gr.Row(): with gr.Column(scale=80): with gr.Row(): - negative_prompt = gr.Textbox(label="Negative prompt", elem_id=f"{id_part}_neg_prompt", show_label=False, lines=3, placeholder="Negative prompt (press Ctrl+Enter or Alt+Enter to generate)", elem_classes=["prompt"]) + negative_prompt = gr.Textbox(elem_id=f"{id_part}_neg_prompt", show_label=False, lines=3, placeholder="Negative prompt", elem_classes=["prompt"]) button_interrogate = None button_deepbooru = None if is_img2img: @@ -270,7 +270,7 @@ def create_toprow(is_img2img): negative_token_button = gr.Button(visible=False, elem_id=f"{id_part}_negative_token_button") with gr.Row(elem_id=f"{id_part}_styles_row"): prompt_styles = gr.Dropdown(label="Styles", elem_id=f"{id_part}_styles", choices=[k for k, v in modules.shared.prompt_styles.styles.items()], value=[], multiselect=True) - create_refresh_button(prompt_styles, modules.shared.prompt_styles.reload, lambda: {"choices": [k for k, v in modules.shared.prompt_styles.styles.items()]}, f"refresh_{id_part}_styles") + # create_refresh_button(prompt_styles, modules.shared.prompt_styles.reload, lambda: {"choices": [k for k, v in modules.shared.prompt_styles.styles.items()]}, f"refresh_{id_part}_styles") prompt_styles_btn = gr.Button('Apply', elem_id=f"{id_part}_styles_select", visible=False) prompt_styles_btn.click(_js="applyStyles", fn=parse_style, inputs=[prompt_styles], outputs=[prompt_styles]) return prompt, prompt_styles, negative_prompt, submit, button_interrogate, button_deepbooru, prompt_style_apply, save_style, paste, extra_networks_button, token_counter, token_button, negative_token_counter, negative_token_button diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index b12b21859..bc33ba7a4 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -64,7 +64,7 @@ def get_metadata(page: str = "", item: str = ""): metadata = page.metadata.get(item, 'none') if metadata is None: metadata = '' - shared.log.debug(f'Extra networks metadata: page={page} item={item} len={len(metadata)}') + shared.log.debug(f"Extra networks metadata: page='{page}' item={item} len={len(metadata)}") return JSONResponse({"metadata": metadata}) @@ -75,7 +75,7 @@ def get_info(page: str = "", item: str = ""): info = page.info.get(item, 'none') if info is None: info = '' - shared.log.debug(f'Extra networks info: page={page} item={item} len={len(info)}') + shared.log.debug(f"Extra networks info: page='{page}' item={item} len={len(info)}") return JSONResponse({"info": info}) @@ -150,7 +150,7 @@ class ExtraNetworksPage: def is_empty(self, folder): for f in listdir(folder): _fn, ext = os.path.splitext(f) - if ext.lower() in ['.ckpt', '.safetensors', '.pt'] or os.path.isdir(os.path.join(folder, f)): + if ext.lower() in ['.ckpt', '.safetensors', '.pt', '.json'] or os.path.isdir(os.path.join(folder, f)): return False return True @@ -164,21 +164,20 @@ class ExtraNetworksPage: continue try: img = Image.open(f) - if img.width > 1024 or img.height > 1024 or os.path.getsize(f) > 70000: + if img.width > 1024 or img.height > 1024 or os.path.getsize(f) > 65536: img = img.convert('RGB') img.thumbnail((512, 512), Image.HAMMING) - img.save(fn) + img.save(fn, quality=50) img.close() created += 1 except Exception as e: shared.log.error(f'Extra network error creating thumbnail: {f} {e}') if created > 0: - shared.log.info(f"Extra network created thumbnails: {self.name} {created}") + shared.log.info(f"Extra network thumbnails: {self.name} created={created}") self.missing_thumbs.clear() def create_page(self, tabname, skip = False): - if self.refresh_time is not None and self.refresh_time > refresh_time: - # shared.log.debug(f'Extra networks: {self.name} items={len(self.items)} tab={tabname} cached') + if self.refresh_time is not None and self.refresh_time > refresh_time: # cached page return self.html t0 = time.time() self_name_id = self.name.replace(" ", "_") @@ -219,7 +218,7 @@ class ExtraNetworksPage: else: return '' t1 = time.time() - shared.log.debug(f'Extra networks: page={self.name} items={len(self.items)} subdirs={len(subdirs)} tab={tabname} dirs={self.allowed_directories_for_previews()} time={round(t1-t0, 2)}') + shared.log.debug(f"Extra networks: page='{self.name}' items={len(self.items)} subdirs={len(subdirs)} tab={tabname} dirs={self.allowed_directories_for_previews()} time={round(t1-t0, 2)}") threading.Thread(target=self.create_thumb).start() def list_items(self): @@ -302,12 +301,48 @@ class ExtraNetworksPage: pass return '' + def save_preview(self, index, images, filename): + try: + image = image_from_url_text(images[int(index)]) + except Exception as e: + shared.log.error(f'Extra network save preview: {filename} {e}') + return + is_allowed = False + for page in extra_pages: + if any(path_is_parent(x, filename) for x in page.allowed_directories_for_previews()): + is_allowed = True + break + if not is_allowed: + shared.log.error(f'Extra network save preview: {filename} not allowed') + return + if image.width > 512 or image.height > 512: + image = image.convert('RGB') + image.thumbnail((512, 512), Image.HAMMING) + image.save(filename, quality=50) + fn, _ext = os.path.splitext(filename) + thumb = fn + '.thumb.jpg' + if os.path.exists(thumb): + shared.log.debug(f'Extra network delete thumbnail: {thumb}') + os.remove(thumb) + shared.log.info(f'Extra network save preview: {filename}') + + def save_description(self, filename, desc): + lastDotIndex = filename.rindex('.') + filename = filename[0:lastDotIndex]+".txt" + if desc != "": + try: + with open(filename, 'w', encoding='utf-8') as f: + f.write(desc) + shared.log.info(f'Extra network save description: {filename} {desc}') + except Exception as e: + shared.log.error(f'Extra network save description: {filename} {e}') + def initialize(): extra_pages.clear() -def register_default_pages(): +def register_pages(): from modules.ui_extra_networks_textual_inversion import ExtraNetworksPageTextualInversion from modules.ui_extra_networks_hypernets import ExtraNetworksPageHypernetworks from modules.ui_extra_networks_checkpoints import ExtraNetworksPageCheckpoints @@ -360,10 +395,6 @@ def create_ui(container, button, tabname, skip_indexing = False): is_visible = not is_visible return is_visible, gr.update(visible=is_visible), gr.update(variant=("secondary-down" if is_visible else "secondary")) - state_visible = gr.State(value=False) # pylint: disable=abstract-class-instantiated - button.click(fn=toggle_visibility, inputs=[state_visible], outputs=[state_visible, container, button]) - button_close.click(fn=toggle_visibility, inputs=[state_visible], outputs=[state_visible, container]) - def en_refresh(title): res = [] for page in extra_pages: @@ -371,12 +402,15 @@ def create_ui(container, button, tabname, skip_indexing = False): page.refresh() page.refresh_time = None page.create_page(ui.tabname) - shared.log.debug(f"Refreshing Extra networks: page={page.title} items={len(page.items)} tab={ui.tabname}") + shared.log.debug(f"Refreshing Extra networks: page='{page.title}' items={len(page.items)} tab={ui.tabname}") res.append(page.html) ui.search.update(value = ui.search.value) return res - button_refresh.click(_js='extraNetworksRefreshButton', fn=en_refresh, inputs=[ui.search], outputs=ui.pages) + state_visible = gr.State(value=False) # pylint: disable=abstract-class-instantiated + button.click(fn=toggle_visibility, inputs=[state_visible], outputs=[state_visible, container, button]) + button_close.click(fn=toggle_visibility, inputs=[state_visible], outputs=[state_visible, container]) + button_refresh.click(_js='getENActivePage', fn=en_refresh, inputs=[ui.search], outputs=ui.pages) return ui @@ -388,54 +422,37 @@ def path_is_parent(parent_path, child_path): def setup_ui(ui, gallery): - def save_preview(index, images, filename): - if len(images) == 0: - for page in extra_pages: - page.create_page(ui.tabname) - return [page.html for page in extra_pages] - index = int(index) - index = 0 if index < 0 else index - index = len(images) - 1 if index >= len(images) else index - img_info = images[index if index >= 0 else 0] - image = image_from_url_text(img_info) - is_allowed = False - for extra_page in extra_pages: - if any(path_is_parent(x, filename) for x in extra_page.allowed_directories_for_previews()): - is_allowed = True - break - assert is_allowed, f'writing to {filename} is not allowed' - image.save(filename) - fn, _ext = os.path.splitext(filename) - thumb = fn + '.thumb.jpg' - if os.path.exists(thumb): - shared.log.debug(f'Extra network delete thumbnail: {thumb}') - os.remove(thumb) - shared.log.info(f'Extra network save preview: {filename}') - return [page.create_page(ui.tabname) for page in extra_pages] + def save_preview(pagename, index, images, filename): + res = [] + for page in extra_pages: + if pagename is None or pagename == '' or pagename == page.title or len(page.html) == 0: + page.save_preview(index, images, filename) + res.append(page.create_page(ui.tabname)) + else: + res.append(page.html) + return res + ui.button_save_preview.click( fn=save_preview, - _js="function(x, y, z) {return [selected_gallery_index(), y, z]}", - inputs=[ui.preview_target_filename, gallery, ui.preview_target_filename], - outputs=[*ui.pages] + _js="function(t, i, y, z) {return [getENActivePage(), selected_gallery_index(), y, z]}", + inputs=[ui.search, ui.preview_target_filename, gallery, ui.preview_target_filename], + outputs=ui.pages ) - # write description to a file - def save_description(filename, desc): - lastDotIndex = filename.rindex('.') - filename = filename[0:lastDotIndex]+".txt" - if desc != "": - try: - with open(filename,'w', encoding='utf-8') as f: - f.write(desc) - shared.log.info(f'Extra network save description: {filename} {desc}') - except Exception as e: - shared.log.error(f'Extra network save description: {filename} {e}') - return [page.create_page(ui.tabname) for page in extra_pages] + def save_description(pagename, filename, desc): + res = [] + for page in extra_pages: + if pagename is None or pagename == '' or pagename == page.title or len(page.html) == 0: + page.save_description(filename, desc) + res.append(page.create_page(ui.tabname)) + else: + res.append(page.html) + return res ui.button_save_description.click( fn=save_description, - _js="function(x, y) { return [x, y] }", - inputs=[ui.description_target_filename, ui.description], - outputs=[*ui.pages] + _js="function(t, x, y) { return [getENActivePage(), x, y] }", + inputs=[ui.search, ui.description_target_filename, ui.description], + outputs=ui.pages ) diff --git a/modules/ui_extra_networks_styles.py b/modules/ui_extra_networks_styles.py index 974ea0614..c3e402d70 100644 --- a/modules/ui_extra_networks_styles.py +++ b/modules/ui_extra_networks_styles.py @@ -1,7 +1,6 @@ import os import html import json - from modules import shared, ui_extra_networks @@ -12,21 +11,53 @@ class ExtraNetworksPageStyles(ui_extra_networks.ExtraNetworksPage): def refresh(self): shared.prompt_styles.reload() - def list_items(self): + """ + import io + import base64 + from PIL import Image + + def image2str(image): + buff = io.BytesIO() + image.save(buff, format="JPEG", quality=80) + encoded = base64.b64encode(buff.getvalue()) + return encoded + + def str2image(data): + buff = io.BytesIO(base64.b64decode(data)) + return Image.open(buff) + + def save_preview(self, index, images, filename): + from modules.generation_parameters_copypaste import image_from_url_text + try: + image = image_from_url_text(images[int(index)]) + except Exception: + shared.log.error(f'Extra network save preview: {filename} no image') + return + if image.width > 512 or image.height > 512: + image = image.convert('RGB').thumbnail((512, 512), Image.HAMMING) for k in shared.prompt_styles.styles.keys(): - path = os.path.join(shared.opts.styles_dir, k) - txt = f'Prompt: {shared.prompt_styles.styles[k].prompt}' - negative = shared.prompt_styles.styles[k].negative_prompt - if negative is not None and len(negative) > 0: - txt += f'\nNegative: {negative}' + if k == filename: + shared.prompt_styles.styles[k].preview = image2str(image) + break + + def save_description(self, filename, desc): + pass + """ + + def list_items(self): + for k, v in shared.prompt_styles.styles.items(): + fn = os.path.join(shared.opts.styles_dir, v.filename) + txt = f'Prompt: {v.prompt}' + if len(v.negative_prompt) > 0: + txt += f'\nNegative: {v.negative_prompt}' yield { - "name": k, - "search_term": path, - "filename": path, - "preview": self.find_preview(path), + "name": v.name, + "search_term": f'{txt} /{v.filename}', + "filename": v.filename, + "preview": self.find_preview(fn), "description": txt, "onclick": '"' + html.escape(f"""return selectStyle({json.dumps(k)})""") + '"', - "local_preview": f"{path}.{shared.opts.samples_format}", + "local_preview": f"{fn}.{shared.opts.samples_format}", } def allowed_directories_for_previews(self): diff --git a/webui.py b/webui.py index 98856ee60..f4496d169 100644 --- a/webui.py +++ b/webui.py @@ -13,7 +13,7 @@ import torch # pylint: disable=wrong-import-order from modules import timer, errors, paths # pylint: disable=unused-import local_url = None -from installer import log, git_commit, print_dict +from installer import log, git_commit import ldm.modules.encoders.modules # pylint: disable=W0611,C0411,E0401 from modules import shared, extensions, extra_networks, ui_tempdir, ui_extra_networks, modelloader # pylint: disable=ungrouped-imports from modules.paths import create_paths @@ -116,9 +116,10 @@ def initialize(): modules.textual_inversion.textual_inversion.list_textual_inversion_templates() shared.reload_hypernetworks() + shared.prompt_styles.reload() ui_extra_networks.initialize() - ui_extra_networks.register_default_pages() + ui_extra_networks.register_pages() extra_networks.initialize() extra_networks.register_default_extra_networks() timer.startup.record("extra-networks") @@ -196,15 +197,13 @@ def async_policy(): super().__init__() self.loop = self.get_event_loop() self.loop.set_exception_handler(self.handle_exception) - log.debug(f"Event loop: {self.loop}") + # log.debug(f"Event loop: {self.loop}") asyncio.set_event_loop_policy(AnyThreadEventLoopPolicy()) def start_common(): log.debug('Entering start sequence') - if cmd_opts.debug and hasattr(shared, 'get_version'): - log.debug(f'Version: {print_dict(shared.get_version())}') logging.disable(logging.NOTSET if cmd_opts.debug else logging.DEBUG) if shared.cmd_opts.data_dir is not None and len(shared.cmd_opts.data_dir) > 0: log.info(f'Using data path: {shared.cmd_opts.data_dir}') @@ -265,7 +264,7 @@ def start_ui(): ui_tempdir.register_tmp_file(shared.demo, os.path.join(cmd_opts.data_dir, 'x')) shared.log.info(f'Local URL: {local_url}') if cmd_opts.docs: - shared.log.info(f'API Docs: {local_url[:-1]}/docs') # {local_url[:-1]}?view=api + shared.log.info(f'API Docs: {local_url[:-1]}/docs') # pylint: disable=unsubscriptable-object if share_url is not None: shared.log.info(f'Share URL: {share_url}') shared.log.debug(f'Gradio registered functions: {len(shared.demo.fns)}')