diff --git a/.eslintrc.json b/.eslintrc.json index 0fa3cf216..3d73788c3 100644 --- a/.eslintrc.json +++ b/.eslintrc.json @@ -22,6 +22,7 @@ "no-confusing-arrow":"off", "no-console":"off", "no-empty":"off", + "no-loop-func":"off", "no-mixed-operators":"off", "no-param-reassign":"off", "no-plusplus":"off", @@ -67,11 +68,11 @@ "switch_to_extras": "readonly", "get_tab_index": "readonly", "create_submit_args": "readonly", - "restart_reload": "readonly", + "restartReload": "readonly", "updateInput": "readonly", "toggleCompact": "readonly", // settings.js - "register_drag_drop": "readonly", + "registerDragDrop": "readonly", //extraNetworks.js "requestGet": "readonly", "getENActiveTab": "readonly", diff --git a/.gitignore b/.gitignore index ec0afa225..32c9bf88d 100644 --- a/.gitignore +++ b/.gitignore @@ -58,6 +58,8 @@ cache .idea/ /localizations -# unexcluded so folders get created +# force included !/models/VAE-approx !/models/VAE-approx/model.pt +!/models/Reference +!/models/Reference/**/* diff --git a/.gitmodules b/.gitmodules index 05e6e3c84..a6d3e9bf1 100644 --- a/.gitmodules +++ b/.gitmodules @@ -12,10 +12,6 @@ path = modules/lora url = https://github.com/kohya-ss/sd-scripts ignore = dirty -[submodule "extensions-builtin/clip-interrogator-ext"] - path = extensions-builtin/clip-interrogator-ext - url = https://github.com/Dahvikiin/clip-interrogator-ext.git - ignore = dirty [submodule "extensions-builtin/sd-webui-controlnet"] path = extensions-builtin/sd-webui-controlnet url = https://github.com/Mikubill/sd-webui-controlnet diff --git a/CHANGELOG.md b/CHANGELOG.md index 5995dcf3f..1b8dbdce2 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -11,6 +11,9 @@ Also, [Wiki](https://github.com/vladmandic/automatic/wiki) has been updated with Some highlights: [OpenVINO](https://github.com/vladmandic/automatic/wiki/OpenVINO), [IntelArc](https://github.com/vladmandic/automatic/wiki/Intel-ARC), [DirectML](https://github.com/vladmandic/automatic/wiki/DirectML), [ONNX/Olive>](https://github.com/vladmandic/automatic/wiki/ONNX-Runtime) - **Diffusers** + - since now **SD.Next** supports **12** different model types, we've added reference model for each type in + *Extra networks -> Reference* for easier select & auto-download + Models can still be downloaded manually, this is just a convenience feature & a showcase for supported models - new model type: [Segmind SSD-1B](https://huggingface.co/segmind/SSD-1B) its a *distilled* model, this time 50% smaller and faster version of SD-XL! (and quality does not suffer, its just more optimized) @@ -37,6 +40,9 @@ Some highlights: [OpenVINO](https://github.com/vladmandic/automatic/wiki/OpenVIN - extend support for [Free-U](https://github.com/ChenyangSi/FreeU) improve generations quality at no cost (other than finding params that work for you) - **General** + - attempt to auto-fix invalid samples which occure due to math errors in lower precision + example: `RuntimeWarning: invalid value encountered in cast: sample = sample.astype(np.uint8)` + begone **black images** *(note: if it proves as working, this solution will need to be expanded to cover all scenarios)* - add **Lora OFT** support, thanks @antis0007 and @ai-casanova - **Upscalers** - **compile** option, thanks @disty0 @@ -47,20 +53,27 @@ Some highlights: [OpenVINO](https://github.com/vladmandic/automatic/wiki/OpenVIN - new option: *settings -> system paths -> models* can be used to set custom base path for *all* models (previously only as cli option) - remove external clone of items in `/repositories` + - **Interrogator** module has been removed from `extensions-builtin` + and fully implemented (and improved) natively - **UI** - UI tweaks for default themes - UI switch core font in default theme to **noto-sans** previously default font was simply *system-ui*, but it lead to too much variations between browsers and platforms - updated **Context menu** - right-click on prompt or generate button + right-click on any button (e.g. generate button) - **Extra networks** - sort by name, size, date, etc. - switch between *gallery* and *list* views + - add tags from user metadata (in addition to tags in model metadata) for **lora** + - faster enumeration of all networks on server startup - **Packages** - updated `diffusers` to 0.22.0, `transformers` to 4.34.1 - update **openvino**, thanks @disty0 - - update **ipex**, thanks @disty0 - update **directml**, @lshqqytiger +- **Compute** + - **OpenVINO**: + - updated to mainstream `torch` *2.1.0* + - support for **ESRGAN** upscalers - **Fixes** - fix **freeu** for backend original and add it to xyz grid - fix loading diffuser models in huggingface format from non-standard location diff --git a/README.md b/README.md index 84e66df8d..e56eb9d23 100644 --- a/README.md +++ b/README.md @@ -124,10 +124,7 @@ SD.Next comes with several extensions pre-installed: - [ControlNet](https://github.com/Mikubill/sd-webui-controlnet) - [Agent Scheduler](https://github.com/ArtVentureX/sd-webui-agent-scheduler) -- [Multi-Diffusion Tiled Diffusion and VAE](https://github.com/pkuliyi2015/multidiffusion-upscaler-for-automatic1111) -- [LyCORIS](https://github.com/KohakuBlueleaf/a1111-sd-webui-lycoris) - [Image Browser](https://github.com/AlUlkesh/stable-diffusion-webui-images-browser) -- [CLiP Interrogator](https://github.com/pharmapsychotic/clip-interrogator-ext) - [Rembg Background Removal](https://github.com/AUTOMATIC1111/stable-diffusion-webui-rembg) ### **Collab** diff --git a/extensions-builtin/Lora/ui_extra_networks_lora.py b/extensions-builtin/Lora/ui_extra_networks_lora.py index d8049ce84..249f84890 100644 --- a/extensions-builtin/Lora/ui_extra_networks_lora.py +++ b/extensions-builtin/Lora/ui_extra_networks_lora.py @@ -14,10 +14,10 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): def create_item(self, name): l = networks.available_networks.get(name) - # alias = lora_on_disk.get_alias() try: path, _ext = os.path.splitext(l.filename) - possible_tags = l.metadata.get('ss_tag_frequency', {}) if l.metadata is not None else {} + name = os.path.splitext(os.path.relpath(l.filename, shared.cmd_opts.lora_dir))[0] + if shared.backend == shared.Backend.ORIGINAL: if l.sd_version == network.SdVersion.SDXL: return None @@ -30,6 +30,9 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): elif shared.sd_model_type == 'sd': if l.sd_version == network.SdVersion.SDXL: return None + + # tags from model metedata + possible_tags = l.metadata.get('ss_tag_frequency', {}) if l.metadata is not None else {} if isinstance(possible_tags, str): possible_tags = {} tags = {} @@ -39,7 +42,7 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): if words[0] == '{}': words[0] = 0 tags[' '.join(words[1:])] = words[0] - name = os.path.splitext(os.path.relpath(l.filename, shared.cmd_opts.lora_dir))[0] + item = { "type": 'Lora', "name": name, @@ -47,60 +50,29 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): "hash": l.shorthash, "search_term": self.search_terms_from_path(l.filename) + ' '.join(tags.keys()), "preview": self.find_preview(l.filename), - "description": self.find_description(l.filename), - "info": self.find_info(l.filename), "prompt": json.dumps(f" "), "local_preview": f"{path}.{shared.opts.samples_format}", "metadata": json.dumps(l.metadata, indent=4) if l.metadata else None, - "tags": tags, "mtime": os.path.getmtime(l.filename), "size": os.path.getsize(l.filename), } + info = self.find_info(l.filename) + item["info"] = info + item["description"] = self.find_description(l.filename, info) # use existing info instead of double-read + + # tags from user metadata + possible_tags = info.get('tags', []) + if not isinstance(possible_tags, list): + possible_tags = [v for v in possible_tags.values()] + for v in possible_tags: + tags[v] = 0 + item["tags"] = tags + return item except Exception as e: shared.log.debug(f"Extra networks error: type=lora file={name} {e}") return None - """ - item = { - "name": name, - "filename": lora_on_disk.filename, - "shorthash": lora_on_disk.shorthash, - "preview": self.find_preview(path), - "description": self.find_description(path), - "search_term": self.search_terms_from_path(lora_on_disk.filename) + " " + (lora_on_disk.hash or ""), - "local_preview": f"{path}.{shared.opts.samples_format}", - "metadata": lora_on_disk.metadata, - "sort_keys": {'default': index, **self.get_sort_keys(lora_on_disk.filename)}, - "sd_version": lora_on_disk.sd_version.name, - } - self.read_user_metadata(item) - activation_text = item["user_metadata"].get("activation text") - preferred_weight = item["user_metadata"].get("preferred weight", 0.0) - item["prompt"] = quote_js(f"") - if activation_text: - item["prompt"] += " + " + quote_js(" " + activation_text) - sd_version = item["user_metadata"].get("sd version") - if sd_version in network.SdVersion.__members__: - item["sd_version"] = sd_version - sd_version = network.SdVersion[sd_version] - else: - sd_version = lora_on_disk.sd_version - if shared.opts.lora_show_all or not enable_filter: - pass - elif sd_version == network.SdVersion.Unknown: - model_version = network.SdVersion.SDXL if shared.sd_model.is_sdxl else network.SdVersion.SD2 if shared.sd_model.is_sd2 else network.SdVersion.SD1 - if model_version.name in shared.opts.lora_hide_unknown_for_versions: - return None - elif shared.sd_model.is_sdxl and sd_version != network.SdVersion.SDXL: - return None - elif shared.sd_model.is_sd2 and sd_version != network.SdVersion.SD2: - return None - elif shared.sd_model.is_sd1 and sd_version != network.SdVersion.SD1: - return None - return item - """ - def list_items(self): for _index, name in enumerate(networks.available_networks): item = self.create_item(name) diff --git a/extensions-builtin/clip-interrogator-ext b/extensions-builtin/clip-interrogator-ext deleted file mode 160000 index 6e31272e1..000000000 --- a/extensions-builtin/clip-interrogator-ext +++ /dev/null @@ -1 +0,0 @@ -Subproject commit 6e31272e14308b4918f9785b1dda7cc1149e8838 diff --git a/extensions-builtin/sd-webui-controlnet b/extensions-builtin/sd-webui-controlnet index 3011ff6e7..fce6775a6 160000 --- a/extensions-builtin/sd-webui-controlnet +++ b/extensions-builtin/sd-webui-controlnet @@ -1 +1 @@ -Subproject commit 3011ff6e706d3fdd0cc7d2ac8ff0d59020b8f767 +Subproject commit fce6775a6dddef52ecd658259e909687d9dedf72 diff --git a/html/locale_en.json b/html/locale_en.json index e2480b253..f0c84af87 100644 --- a/html/locale_en.json +++ b/html/locale_en.json @@ -40,9 +40,9 @@ {"id":"","label":"disabled","localized":"","hint":""} ], "tabs": [ - {"id":"","label":"From Text","localized":"","hint":"Create image from text"}, - {"id":"","label":"From Image","localized":"","hint":"Create image from image"}, - {"id":"","label":"Process Image","localized":"","hint":"Process existing image"}, + {"id":"","label":"Text","localized":"","hint":"Create image from text"}, + {"id":"","label":"Image","localized":"","hint":"Create image from image"}, + {"id":"","label":"Process","localized":"","hint":"Process existing image"}, {"id":"","label":"Train","localized":"","hint":"Run training or model merging"}, {"id":"","label":"Models","localized":"","hint":"Convert or merge your models"}, {"id":"","label":"Interrogator","localized":"","hint":"Run interrogate to get description of your image"}, diff --git a/html/reference.json b/html/reference.json new file mode 100644 index 000000000..4a1ad9d34 --- /dev/null +++ b/html/reference.json @@ -0,0 +1,40 @@ +{ + "RunwayML SD 1.5": { + "path": "runwayml/stable-diffusion-v1-5" + }, + "StabilityAI SD 2.1": { + "path": "stabilityai/stable-diffusion-2-1-base" + }, + "StabilityAI SD-XL 1.0 Base": { + "path": "stabilityai/stable-diffusion-xl-base-1.0" + }, + "Segmind SSD-1B": { + "path": "segmind/SSD-1B" + }, + "Segmind Tiny": { + "path": "segmind/tiny-sd" + }, + "LCM Dreamshaper 7": { + "path": "SimianLuo/LCM_Dreamshaper_v7" + }, + "Warp Wuerstchen": { + "path": "warp-ai/wuerstchen" + }, + "Kandinsky 2.1": { + "path": "kandinsky-community/kandinsky-2-1" + }, + "Kandinsky 2.2": { + "path": "kandinsky-community/kandinsky-2-2-decoder" + }, + "DeepFloyd IF Medium": { + "path": "DeepFloyd/IF-I-M-v1.0" + }, + "Tsinghua UniDiffuser": { + "path": "thu-ml/unidiffuser-v1", + "desc": "UniDiffuser is a unified diffusion framework to fit all distributions relevant to a set of multi-modal data in one transformer. UniDiffuser is able to perform image, text, text-to-image, image-to-text, and image-text pair generation by setting proper timesteps without additional overhead.\nSpecifically, UniDiffuser employs a variation of transformer, called U-ViT, which parameterizes the joint noise prediction network. Other components perform as encoders and decoders of different modalities, including a pretrained image autoencoder from Stable Diffusion, a pretrained image ViT-B/32 CLIP encoder, a pretrained text ViT-L CLIP encoder, and a GPT-2 text decoder finetuned by ourselves.", + "preview": "unidiffuser-v1.jpg" + }, + "Sudo-AI Zero123": { + "path": "sudo-ai/zero123plus-v1.1" + } +} diff --git a/installer.py b/installer.py index c44e27a63..01a48ad8e 100644 --- a/installer.py +++ b/installer.py @@ -508,8 +508,6 @@ def check_torch(): import xformers if torch.__version__ != '2.0.1+cu118' and xformers.__version__ in ['0.0.22', '0.0.21', '0.0.20']: log.warning(f'Likely incompatible torch with: xformers=={xformers.__version__} installed: torch=={torch.__version__} required: torch==2.1.0+cu118 - build xformers manually or downgrade torch') - if 'cu118' not in torch.__version__: - log.warning(f'Likely incompatible Cuda with: xformers=={xformers.__version__} installed: torch=={torch.__version__} required: torch==2.1.0+cu118 - build xformers manually or downgrade torch') elif not args.experimental and not args.use_xformers: uninstall('xformers') except Exception as e: diff --git a/javascript/amethyst-nightfall.css b/javascript/amethyst-nightfall.css index 18d296730..2cbead34f 100644 --- a/javascript/amethyst-nightfall.css +++ b/javascript/amethyst-nightfall.css @@ -228,23 +228,11 @@ svg.feather.feather-image, .feather .feather-image { display: none } --neutral-800: #333333; --neutral-900: #111827; --neutral-950: #0b0f19; - --spacing-xxs: 1px; - --spacing-xs: 2px; - --spacing-sm: 4px; - --spacing-lg: 6px; - --spacing-xl: 8px; --radius-xxs: 0; --radius-xs: 0; --radius-md: 0; --radius-xl: 0; --radius-xxl: 0; - --text-xxs: 9px; - --text-xs: 10px; - --text-sm: 12px; - --text-md: 14px; - --text-lg: 16px; - --text-xl: 22px; - --text-xxl: 26px; --font: 'Source Sans Pro', 'ui-sans-serif', 'system-ui', sans-serif; --font-mono: 'IBM Plex Mono', 'ui-monospace', 'Consolas', monospace; --body-text-size: var(--text-md); diff --git a/javascript/black-orange.css b/javascript/black-orange.css index 4b23af2f5..dfc6f8753 100644 --- a/javascript/black-orange.css +++ b/javascript/black-orange.css @@ -38,9 +38,6 @@ --spacing-xxl: 6px; --line-sm: 1.2em; --line-md: 1.4em; - --text-sm: 12px; - --text-md: 14px; - --text-lg: 15px; } html { font-size: var(--font-size); } @@ -243,20 +240,11 @@ svg.feather.feather-image, .feather .feather-image { display: none } --neutral-800: #333333; --neutral-900: #111827; --neutral-950: #0b0f19; - --spacing-xxs: 1px; - --spacing-xs: 2px; - --spacing-sm: 4px; - --spacing-lg: 6px; - --spacing-xl: 8px; --radius-xxs: 0; --radius-xs: 0; --radius-md: 0; --radius-xl: 0; --radius-xxl: 0; - --text-xxs: 9px; - --text-xs: 10px; - --text-xl: 22px; - --text-xxl: 26px; --body-text-size: var(--text-md); --body-text-weight: 400; --embed-radius: var(--radius-lg); diff --git a/javascript/black-teal.css b/javascript/black-teal.css index a9094fe31..1918dfda5 100644 --- a/javascript/black-teal.css +++ b/javascript/black-teal.css @@ -33,9 +33,6 @@ --radius-lg: 4px; --line-sm: 1.2em; --line-md: 1.4em; - --text-sm: 12px; - --text-md: 14px; - --text-lg: 15px; } html { font-size: var(--font-size); font-family: var(--font); } @@ -235,21 +232,11 @@ textarea[rows="1"] { height: 33px !important; width: 99% !important; padding: 8p --neutral-800: #333333; --neutral-900: #111827; --neutral-950: #0b0f19; - --spacing-xxs: 1px; - --spacing-xs: 2px; - --spacing-sm: 3px; - --spacing-lg: 4px; - --spacing-xl: 5px; - --spacing-xxl: 6px; --radius-xxs: 0; --radius-xs: 0; --radius-md: 0; --radius-xl: 0; --radius-xxl: 0; - --text-xxs: 9px; - --text-xs: 10px; - --text-xl: 22px; - --text-xxl: 26px; --body-text-size: var(--text-md); --body-text-weight: 400; --embed-radius: var(--radius-lg); diff --git a/javascript/contextMenus.js b/javascript/contextMenus.js index 40e93f4e4..9cfe2dcb3 100644 --- a/javascript/contextMenus.js +++ b/javascript/contextMenus.js @@ -104,7 +104,7 @@ function initContextMenu() { }; for (const tab of ['txt2img', 'img2img']) { - for (const el of ['prompt > label > textarea', 'generate']) { + for (const el of ['generate', 'interrupt', 'skip', 'pause', 'paste', 'clear_prompt', 'extra_networks_btn']) { const id = `#${tab}_${el}`; appendContextMenuOption(id, 'Copy to clipboard', () => navigator.clipboard.writeText(document.querySelector(`#${tab}_prompt > label > textarea`).value)); appendContextMenuOption(id, 'Generate forever', () => generateForever(`#${tab}_generate`)); diff --git a/javascript/extensions.js b/javascript/extensions.js index 1e56a9a36..f94f9b57d 100644 --- a/javascript/extensions.js +++ b/javascript/extensions.js @@ -5,7 +5,7 @@ function extensions_apply(extensions_disabled_list, extensions_update_list, disa if (x.name.startsWith('enable_') && !x.checked) disable.push(x.name.substring(7)); if (x.name.startsWith('update_') && x.checked) update.push(x.name.substring(7)); }); - restart_reload(); + restartReload(); log('Extensions apply:', { disable, update }); return [JSON.stringify(disable), JSON.stringify(update), disable_all]; } diff --git a/javascript/extraNetworks.js b/javascript/extraNetworks.js index a0773514a..06c0cf819 100644 --- a/javascript/extraNetworks.js +++ b/javascript/extraNetworks.js @@ -1,4 +1,5 @@ const activePromptTextarea = {}; +let sortVal = 0; // helpers @@ -92,22 +93,28 @@ async function filterExtraNetworksForTab(tabname, searchTerm) { let found = 0; let items = 0; const t0 = performance.now(); - const cards = Array.from(gradioApp().querySelectorAll(`#${tabname}_extra_tabs div.card`)); - cards.forEach((elem) => { - items += 1; - if (searchTerm === '') { - elem.style.display = ''; - } else { - let text = `${elem.querySelector('.name').textContent.toLowerCase()} ${elem.querySelector('.search_term').textContent}`; - text = text.toLowerCase().replace('models--', 'Diffusers').replace('\\', '/'); - if (text.indexOf(searchTerm) === -1) { - elem.style.display = 'none'; - } else { + const pagename = getENActivePage(); + if (!pagename) return; + const allPages = Array.from(gradioApp().querySelectorAll('.extra-network-cards')); + const pages = allPages.filter((el) => el.id.includes(pagename.toLowerCase())); + for (const pg of pages) { + const cards = Array.from(pg.querySelectorAll('.card') || []); + cards.forEach((elem) => { + items += 1; + if (searchTerm === '') { elem.style.display = ''; - found += 1; + } else { + let text = elem.dataset.search.toLowerCase(); + text = text.toLowerCase().replace('models--', 'Diffusers').replace('\\', '/'); + if (text.indexOf(searchTerm) === -1) { + elem.style.display = 'none'; + } else { + elem.style.display = ''; + found += 1; + } } - } - }); + }); + } const t1 = performance.now(); if (found > 0) log(`filterExtraNetworks: text=${searchTerm} items=${items} match=${found} time=${Math.round(1000 * (t1 - t0)) / 1000000}`); else log(`filterExtraNetworks: text=all items=${items} time=${Math.round(1000 * (t1 - t0)) / 1000000}`); @@ -145,8 +152,6 @@ function tryToRemoveExtraNetworkFromPrompt(textarea, text) { return false; } -let sortVal = 0; - function sortExtraNetworks() { const sortDesc = ['Name [A-Z]', 'Name [Z-A]', 'Date [Newest]', 'Date [Oldest]', 'Size [Largest]', 'Size [Smallest]']; const pagename = getENActivePage(); @@ -160,8 +165,8 @@ function sortExtraNetworks() { if (num === 0) return 'sort: no cards'; cards.sort((a, b) => { // eslint-disable-line no-loop-func switch (sortVal) { - case 0: return a.dataset.name ? a.dataset.name.localeCompare(b.dataset.name) : 0; - case 1: return b.dataset.name ? b.dataset.name.localeCompare(a.dataset.name) : 0; + case 0: return a.dataset.name ? a.dataset.search.localeCompare(b.dataset.name) : 0; + case 1: return b.dataset.name ? b.dataset.search.localeCompare(a.dataset.name) : 0; case 2: return a.dataset.mtime && !isNaN(a.dataset.mtime) ? parseFloat(b.dataset.mtime) - parseFloat(a.dataset.mtime) : 0; case 3: return b.dataset.mtime && !isNaN(b.dataset.mtime) ? parseFloat(a.dataset.mtime) - parseFloat(b.dataset.mtime) : 0; case 4: return a.dataset.size && !isNaN(a.dataset.size) ? parseFloat(b.dataset.size) - parseFloat(a.dataset.size) : 0; diff --git a/javascript/invoked.css b/javascript/invoked.css index 684527817..0b540955b 100644 --- a/javascript/invoked.css +++ b/javascript/invoked.css @@ -224,23 +224,11 @@ button.selected {background: var(--button-primary-background-fill);} --neutral-800: #333333; --neutral-900: #111827; --neutral-950: #0b0f19; - --spacing-xxs: 1px; - --spacing-xs: 2px; - --spacing-sm: 4px; - --spacing-lg: 6px; - --spacing-xl: 8px; --radius-xxs: 0; --radius-xs: 0; --radius-md: 0; --radius-xl: 0; --radius-xxl: 0; - --text-xxs: 9px; - --text-xs: 10px; - --text-sm: 12px; - --text-md: 14px; - --text-lg: 16px; - --text-xl: 22px; - --text-xxl: 26px; --body-text-size: var(--text-md); --body-text-weight: 400; --embed-radius: var(--radius-lg); diff --git a/javascript/light-teal.css b/javascript/light-teal.css index 9b34212ec..52e540cfc 100644 --- a/javascript/light-teal.css +++ b/javascript/light-teal.css @@ -33,9 +33,6 @@ --radius-lg: 4px; --line-sm: 1.2em; --line-md: 1.4em; - --text-sm: 12px; - --text-md: 14px; - --text-lg: 15px; } html { font-size: var(--font-size); } @@ -297,20 +294,10 @@ svg.feather.feather-image, .feather .feather-image { display: none } --size-9: 64px; --slider_color: None; --slider-color: ; - --spacing-xxs: 1px; - --spacing-xs: 2px; - --spacing-sm: 3px; - --spacing-lg: 4px; - --spacing-xl: 5px; - --spacing-xxl: 6px; --stat-background-fill: linear-gradient(to right, var(--primary-400), var(--primary-600)); --table-border-color: var(--neutral-700); --table-even-background-fill: #222222; --table-odd-background-fill: #333333; --table-radius: var(--radius-lg); --table-row-focus: var(--color-accent-soft); - --text-lg: 16px; - --text-xs: 10px; - --text-xxl: 26px; - --text-xxs: 9px; } diff --git a/javascript/midnight-barbie.css b/javascript/midnight-barbie.css index 7049a96b9..32447c564 100644 --- a/javascript/midnight-barbie.css +++ b/javascript/midnight-barbie.css @@ -229,23 +229,11 @@ svg.feather.feather-image, .feather .feather-image { display: none } --neutral-800: #322c35; --neutral-900: #1b1127; --neutral-950: #140b19; - --spacing-xxs: 1px; - --spacing-xs: 2px; - --spacing-sm: 4px; - --spacing-lg: 6px; - --spacing-xl: 8px; --radius-xxs: 0; --radius-xs: 0; --radius-md: 0; --radius-xl: 0; --radius-xxl: 0; - --text-xxs: 9px; - --text-xs: 10px; - --text-sm: 12px; - --text-md: 14px; - --text-lg: 16px; - --text-xl: 22px; - --text-xxl: 26px; --body-text-size: var(--text-md); --body-text-weight: 400; --embed-radius: var(--radius-lg); diff --git a/javascript/progressBar.js b/javascript/progressBar.js index 5d8e65323..0133d3c1f 100644 --- a/javascript/progressBar.js +++ b/javascript/progressBar.js @@ -111,6 +111,7 @@ function requestProgress(id_task, progressEl, galleryEl, atEnd = null, onProgres }; const start = (id_task, id_live_preview) => { // eslint-disable-line no-shadow + if (!opts.live_previews_enable || opts.live_preview_refresh_period === 0 || opts.show_progress_every_n_steps === 0) return; request('./internal/progress', { id_task, id_live_preview }, (res) => { lastState = res; const elapsedFromStart = (new Date() - dateStart) / 1000; diff --git a/javascript/sdnext.css b/javascript/sdnext.css index 273224479..32e9edecd 100644 --- a/javascript/sdnext.css +++ b/javascript/sdnext.css @@ -25,7 +25,7 @@ textarea { overflow-y: auto !important; } .gradio-button.secondary-down { background: var(--button-secondary-background-fill); color: var(--button-secondary-text-color); } .gradio-button.secondary-down, .gradio-button.secondary-down:hover { box-shadow: 1px 1px 1px rgba(0,0,0,0.25) inset, 0px 0px 3px rgba(0,0,0,0.15) inset; } .gradio-button.secondary-down:hover { background: var(--button-secondary-background-fill-hover); color: var(--button-secondary-text-color-hover); } -.gradio-button.tool { max-width: min-content; min-width: min-content !important; align-self: end; font-size: 1.4em; color: var(--body-text-color) !important; margin-bottom: var(--spacing-md); align-self: center; } +.gradio-button.tool { max-width: min-content; min-width: min-content !important; align-self: end; font-size: 1.4em; color: var(--body-text-color) !important; margin-top: auto; margin-bottom: var(--spacing-md); align-self: center; } .gradio-checkbox { margin: 0.75em 1.5em 0 0; align-self: center; } .gradio-column { min-width: min(160px, 100%) !important; } .gradio-container { max-width: unset !important; padding: var(--block-label-padding) !important; } @@ -103,7 +103,7 @@ div#extras_scale_to_tab div.form{ flex-direction: row; } /* settings */ #si-sparkline-memo, #si-sparkline-load { background-color: #111; } #quicksettings { width: fit-content; } -#quicksettings > button { padding: 0 1em 0 0; align-self: end; margin-bottom: var(--text-lg); } +#quicksettings > button { padding: 0 1em 0 0; align-self: end; margin-bottom: var(--text-sm); } #settings { display: flex; gap: var(--layout-gap); } #settings div { border: none; gap: 0; margin: 0 0 var(--layout-gap) 0px; padding: 0; } #settings .gr-group { max-width: 70em; } @@ -185,7 +185,7 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt .extra-networks .tab-nav > button { margin-right: 0; height: 24px; padding: 2px 4px 2px 4px; } .extra-networks .buttons { position: absolute; right: 0; margin: -4px; background: var(--background-color); } .extra-networks .buttons > button { margin-left: -0.4em; height: 1.4em; color: var(--primary-300) !important; } -.extra-networks .custom-button { width: 120px; width: 100%; background: none; justify-content: left; text-align: left; padding: 2px 8px 2px 16px; text-indent: -8px; box-shadow: none; line-break: auto; } +.extra-networks .custom-button { width: 120px; width: 100%; background: none; justify-content: left; text-align: left; padding: 3px 3px 3px 12px; text-indent: -6px; box-shadow: none; line-break: auto; } .extra-networks .custom-button:hover { background: var(--button-primary-background-fill) } .extra-networks-tab { padding: 0 !important; } .extra-network-subdirs { background: var(--input-background-fill); overflow-x: hidden; overflow-y: auto; min-width: max(15%, 120px); padding-top: 0.5em; margin-top: -4px !important; } @@ -199,7 +199,7 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt .extra-network-cards .card:hover .preview { box-shadow: none; filter: grayscale(100%); } .extra-network-cards .card:hover .overlay { background: rgba(0, 0, 0, 0.40); } .extra-network-cards .card .overlay .tags { display: none; overflow-wrap: break-word; } -.extra-network-cards .card .overlay .tag { padding: 3px; background: rgba(70, 70, 70, 0.60); font-size: var(--text-lg); cursor: pointer; display: inline-block; margin-bottom: 4px; } +.extra-network-cards .card .overlay .tag { padding: 2px; margin: 2px; background: rgba(70, 70, 70, 0.60); font-size: var(--text-md); cursor: pointer; display: inline-block; } .extra-network-cards .card .actions > span { padding: 4px; } .extra-network-cards .card .actions > span:hover { color: var(--highlight-color); } .extra-network-cards .card:hover .actions { display: block; } @@ -263,3 +263,19 @@ table.settings-value-table td { padding: 0.4em; border: 1px solid #ccc; max-widt @keyframes move { from { background-position-x: 0, -40px; } to { background-position-x: 0, 40px; } } @keyframes spin { from { transform: rotate(0deg); } to { transform: rotate(360deg); } } @keyframes color { from { filter: hue-rotate(0deg) } to { filter: hue-rotate(360deg) } } + +:root, .light, .dark { + --text-xxs: 9px; + --text-xs: 10px; + --text-sm: 12px; + --text-md: 14px; + --text-lg: 15px; + --text-xl: 16px; + --text-xxl: 17px; + --spacing-xxs: 1px; + --spacing-xs: 2px; + --spacing-sm: 3px; + --spacing-lg: 4px; + --spacing-xl: 5px; + --spacing-xxl: 6px; +} diff --git a/javascript/settings.js b/javascript/settings.js index 3cd84b52a..dacc8bd87 100644 --- a/javascript/settings.js +++ b/javascript/settings.js @@ -89,7 +89,7 @@ onAfterUiUpdate(async () => { const jsdata = textarea.value; updateOpts(jsdata); executeCallbacks(optionsChangedCallbacks); - register_drag_drop(); + registerDragDrop(); Object.defineProperty(textarea, 'value', { set(newValue) { diff --git a/javascript/ui.js b/javascript/ui.js index 028cc9fd9..8581fa3db 100644 --- a/javascript/ui.js +++ b/javascript/ui.js @@ -211,7 +211,7 @@ function recalculate_prompts_inpaint(...args) { return Array.from(arguments); } -function register_drag_drop() { +function registerDragDrop() { const qs = gradioApp().getElementById('quicksettings'); if (!qs) return; qs.addEventListener('dragover', (evt) => { @@ -297,7 +297,7 @@ function getTranslation(...args) { return null; } -function monitor_server_status() { +function monitorServerStatus() { document.open(); document.write(` @@ -305,12 +305,12 @@ function monitor_server_status() {

Waiting for server...

@@ -318,12 +318,12 @@ function monitor_server_status() { document.close(); } -function restart_reload() { +function restartReload() { document.body.style = 'background: #222222; font-size: 1rem; font-family:monospace; margin-top:20%; color:lightgray; text-align:center'; document.body.innerHTML = '

Server shutdown in progress...

'; fetch('/sdapi/v1/progress') - .then((res) => setTimeout(restart_reload, 1000)) - .catch((e) => setTimeout(monitor_server_status, 500)); + .then((res) => setTimeout(restartReload, 1000)) + .catch((e) => setTimeout(monitorServerStatus, 500)); return []; } @@ -351,6 +351,12 @@ function selectVAE(name) { log(`Change VAE: ${desiredVAEName}`); } +function selectReference(name) { + console.log('HERE', name); + desiredCheckpointName = name; + gradioApp().getElementById('change_reference').click(); +} + function currentImg2imgSourceResolution(_a, _b, scaleBy) { const img = gradioApp().querySelector('#mode_img2img > div[style="display: block;"] img'); return img ? [img.naturalWidth, img.naturalHeight, scaleBy] : [0, 0, scaleBy]; @@ -361,7 +367,7 @@ function updateImg2imgResizeToTextAfterChangingImage() { return []; } -function create_theme_element() { +function createThemeElement() { const el = document.createElement('img'); el.id = 'theme-preview'; el.className = 'theme-preview'; @@ -393,7 +399,7 @@ function previewTheme() { if (theme) { window.open(theme.subdomain, '_blank'); } else { - const el = document.getElementById('theme-preview') || create_theme_element(); + const el = document.getElementById('theme-preview') || createThemeElement(); el.style.display = el.style.display === 'block' ? 'none' : 'block'; name = name.replace('/', '-'); el.src = `/file=html/${name}.jpg`; diff --git a/models/Reference/unidiffuser-v1.jpg b/models/Reference/unidiffuser-v1.jpg new file mode 100644 index 000000000..72f104c95 Binary files /dev/null and b/models/Reference/unidiffuser-v1.jpg differ diff --git a/modules/images.py b/modules/images.py index e061040f1..2c4ecefe4 100644 --- a/modules/images.py +++ b/modules/images.py @@ -413,7 +413,7 @@ class FilenameGenerator: [part := part.replace(word, '_') for word in invalid_files] # pylint: disable=expression-not-assigned newparts.append(part) fn = Path(*newparts) - max_length = os.statvfs(__file__).f_namemax - 32 if hasattr(os, 'statvfs') else 230 + max_length = max(230, os.statvfs(__file__).f_namemax - 32 if hasattr(os, 'statvfs') else 230) fn = str(fn)[:max_length-max(4, len(ext))].rstrip(invalid_suffix) + ext debug(f'Filename sanitize: input="{filename}" parts={parts} output="{fn}" ext={ext} max={max_length} len={len(fn)}') return fn diff --git a/modules/loader.py b/modules/loader.py index c725d6d24..3dbb0c4a8 100644 --- a/modules/loader.py +++ b/modules/loader.py @@ -41,4 +41,4 @@ errors.install([gradio]) import diffusers # pylint: disable=W0611,C0411 timer.startup.record("diffusers") -errors.log.debug(f'Load packages: torch={getattr(torch, "__long_version__", torch.__version__)} diffusers={diffusers.__version__} gradio={gradio.__version__}') +errors.log.info(f'Load packages: torch={getattr(torch, "__long_version__", torch.__version__)} diffusers={diffusers.__version__} gradio={gradio.__version__}') diff --git a/modules/lora b/modules/lora index 2a23713f7..95ae56bd2 160000 --- a/modules/lora +++ b/modules/lora @@ -1 +1 @@ -Subproject commit 2a23713f71628b2d1b88a51035b3e4ee2b5dbe46 +Subproject commit 95ae56bd22c285ccb2fe5fca96d92f39842bb99b diff --git a/modules/modelloader.py b/modules/modelloader.py index 86cfdf3c7..4d4858cec 100644 --- a/modules/modelloader.py +++ b/modules/modelloader.py @@ -213,21 +213,24 @@ def download_diffusers_model(hub_id: str, cache_dir: str = None, download_config pipeline_dir = None ok = True + err = None try: pipeline_dir = DiffusionPipeline.download(hub_id, **download_config) except Exception as e: + err = e ok = False - shared.log.warning(f"Diffusers download error: {hub_id} {e}") - if not ok: + # shared.log.warning(f"Diffusers download error: {hub_id} {e}") + if not ok and 'Repository Not Found' not in str(err): try: download_config.pop('load_connected_pipeline') download_config.pop('variant') pipeline_dir = hf.snapshot_download(hub_id, **download_config) - except Exception as e: - shared.log.warning(f"Diffusers hub download error: {hub_id} {e}") + except Exception: + # shared.log.warning(f"Diffusers download error: {hub_id} {e}") + pass if pipeline_dir is None: - shared.log.error(f"Diffusers no pipeline folder: {hub_id}") + shared.log.error(f"Diffusers download error: {hub_id} {err}") return None try: # TODO diffusers is this real error? @@ -314,6 +317,23 @@ def find_diffuser(name: str): return None +def load_reference(name: str): + found = [r for r in diffuser_repos if name == r['name'] or name == r['friendly'] or name == r['path']] + if len(found) > 0: # already downloaded + shared.log.debug(f'Reference model: {found[0]}') + return True + shared.log.debug(f'Reference download: {name}') + model_dir = download_diffusers_model(name, shared.opts.diffusers_dir) + if model_dir is None: + shared.log.debug(f'Reference download failed: {name}') + return False + else: + shared.log.debug(f'Reference download complete: {name}') + from modules import sd_models + sd_models.list_models() + return True + + modelloader_directories = {} cache_last = 0 cache_time = 1 diff --git a/modules/processing.py b/modules/processing.py index 3a5a4b435..b0213db25 100644 --- a/modules/processing.py +++ b/modules/processing.py @@ -4,6 +4,7 @@ import math import time import hashlib import random +import warnings from contextlib import nullcontext from typing import Any, Dict, List import torch @@ -726,23 +727,16 @@ def process_images(p: StableDiffusionProcessing) -> Processed: def validate_sample(sample): - ok = True - try: - sample = sample.astype(np.uint8) - return sample - except (Exception, Warning, RuntimeWarning) as e: - shared.log.error(f'Failed to validate sample values: {e}') - ok = False - if not ok: - try: - sample = np.nan_to_num(sample, nan=0, posinf=255, neginf=0) - sample = sample.astype(np.uint8) - shared.log.debug('Corrected sample values') - except (Exception, Warning, RuntimeWarning) as e: - shared.log.error(f'Failed to correct sample values: {e}') - sample = np.zeros_like(sample) - sample = sample.astype(np.uint8) - return sample + with warnings.catch_warnings(record=True) as w: + cast = sample.astype(np.uint8) + if len(w) > 0: + nans = np.isnan(sample).sum() + shared.log.error(f'Failed to validate samples: sample={sample.shape} invalid={nans}') + cast = np.nan_to_num(sample) + minimum, maximum, mean = np.min(cast), np.max(cast), np.mean(cast) + cast = cast.astype(np.uint8) + shared.log.warning(f'Attempted to correct samples: min={minimum:.2f} max={maximum:.2f} mean={mean:.2f}') + return cast def process_images_inner(p: StableDiffusionProcessing) -> Processed: diff --git a/modules/processing_diffusers.py b/modules/processing_diffusers.py index 3207ee147..a12bd6d86 100644 --- a/modules/processing_diffusers.py +++ b/modules/processing_diffusers.py @@ -143,6 +143,7 @@ def process_diffusers(p: StableDiffusionProcessing, seeds, prompts, negative_pro decoded = full_vae_decode(latents=latents, model=shared.sd_model) else: decoded = taesd_vae_decode(latents=latents) + # decoded = validate_sample(decoded) # TODO validate sample imgs = model.image_processor.postprocess(decoded, output_type=output_type) shared.state.job = prev_job return imgs diff --git a/modules/sd_models.py b/modules/sd_models.py index f68d82f1e..dc33538a2 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -203,13 +203,13 @@ def list_models(): def update_model_hashes(): txt = [] lst = [ckpt for ckpt in checkpoints_list.values() if ckpt.hash is None] - shared.log.info(f'Models list: short hash missing for {len(lst)} out of {len(checkpoints_list)} models') + # shared.log.info(f'Models list: short hash missing for {len(lst)} out of {len(checkpoints_list)} models') for ckpt in lst: ckpt.hash = model_hash(ckpt.filename) - txt.append(f'Calculated short hash: {ckpt.title} {ckpt.hash}') - txt.append(f'Updated short hashes for {len(lst)} out of {len(checkpoints_list)} models') + # txt.append(f'Calculated short hash: {ckpt.title} {ckpt.hash}') + # txt.append(f'Updated short hashes for {len(lst)} out of {len(checkpoints_list)} models') lst = [ckpt for ckpt in checkpoints_list.values() if ckpt.sha256 is None or ckpt.shorthash is None] - shared.log.info(f'Models list: full hash missing for {len(lst)} out of {len(checkpoints_list)} models') + shared.log.info(f'Models list: hash missing={len(lst)} total={len(checkpoints_list)}') for ckpt in lst: ckpt.sha256 = hashes.sha256(ckpt.filename, f"checkpoint/{ckpt.name}") ckpt.shorthash = ckpt.sha256[0:10] if ckpt.sha256 is not None else None diff --git a/modules/ui.py b/modules/ui.py index f80ad290c..f18fc1456 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -10,7 +10,7 @@ import numpy as np from PIL import Image from modules.call_queue import wrap_gradio_gpu_call, wrap_queued_call, wrap_gradio_call -from modules import sd_hijack, sd_models, script_callbacks, ui_extensions, deepbooru, extra_networks, ui_common, ui_postprocessing, ui_loadsave, ui_train, ui_models +from modules import sd_hijack, sd_models, script_callbacks, ui_extensions, deepbooru, extra_networks, ui_common, ui_postprocessing, ui_loadsave, ui_train, ui_models, ui_interrogate from modules.ui_components import FormRow, FormGroup, ToolButton, FormHTML from modules.paths import script_path, data_path from modules.shared import opts, cmd_opts @@ -263,7 +263,7 @@ def create_toprow(is_img2img): pause = gr.Button('Pause', elem_id=f"{id_part}_pause") pause.click(fn=lambda: modules.shared.state.pause(), _js='checkPaused', inputs=[], outputs=[]) with gr.Row(elem_id=f"{id_part}_tools"): - button_paste = gr.Button(value='Restore', variant='secondary', elem_id="paste") # symbols.paste + button_paste = gr.Button(value='Restore', variant='secondary', elem_id=f"{id_part}_paste") # symbols.paste button_clear = gr.Button(value='Clear', variant='secondary', elem_id=f"{id_part}_clear_prompt_btn") # symbols.clear button_extra = gr.Button(value='Networks', variant='secondary', elem_id=f"{id_part}_extra_networks_btn") # symbols.networks button_clear.click(fn=lambda *x: ['', ''], inputs=[prompt, negative_prompt], outputs=[prompt, negative_prompt], show_progress=False) @@ -273,8 +273,7 @@ def create_toprow(is_img2img): negative_token_counter = gr.HTML(value="0/75", elem_id=f"{id_part}_negative_token_counter", elem_classes=["token-counter"]) 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=[style.name for style in modules.shared.prompt_styles.styles.values()], value=[], multiselect=True) - prompt_styles = gr.Dropdown(label="Styles", elem_id=f"{id_part}_styles", choices=['aaa'], value=[], multiselect=True) + prompt_styles = gr.Dropdown(label="Styles", elem_id=f"{id_part}_styles", choices=[style.name for style in modules.shared.prompt_styles.styles.values()], value=[], multiselect=True) prompt_styles_btn_refresh = ToolButton(symbols.refresh, elem_id=f"{id_part}_styles_refresh", visible=True) prompt_styles_btn_refresh.click(fn=lambda: gr.update(choices=[style.name for style in modules.shared.prompt_styles.styles.values()]), inputs=[], outputs=[prompt_styles]) prompt_styles_btn_select = gr.Button('Select', elem_id=f"{id_part}_styles_select", visible=False) @@ -636,7 +635,6 @@ def create_ui(startup_timer = None): img2img_batch_inpaint_mask_dir = gr.Textbox(label="Inpaint batch mask directory", **modules.shared.hide_dirs, elem_id="img2img_batch_inpaint_mask_dir") img2img_tabs = [tab_img2img, tab_sketch, tab_inpaint, tab_inpaint_color, tab_inpaint_upload, tab_batch] - for i, tab in enumerate(img2img_tabs): tab.select(fn=lambda tabnum=i: tabnum, inputs=[], outputs=[img2img_selected_tab]) @@ -904,6 +902,11 @@ def create_ui(startup_timer = None): ui_models.create_ui() timer.startup.record("ui-models") + with gr.Blocks(analytics_enabled=False) as interrogate_interface: + ui_interrogate.create_ui() + timer.startup.record("ui-interrogate") + + def create_setting_component(key, is_quicksettings=False): def fun(): return opts.data[key] if key in opts.data else opts.data_labels[key].default @@ -1103,11 +1106,12 @@ def create_ui(startup_timer = None): timer.startup.record("ui-settings") interfaces = [ - (txt2img_interface, "From Text", "txt2img"), - (img2img_interface, "From Image", "img2img"), - (extras_interface, "Process Image", "process"), + (txt2img_interface, "Text", "txt2img"), + (img2img_interface, "Image", "img2img"), + (extras_interface, "Process", "process"), (train_interface, "Train", "train"), (models_interface, "Models", "models"), + (interrogate_interface, "Interrogate", "interrogate"), ] interfaces += script_callbacks.ui_tabs_callback() interfaces += [(settings_interface, "System", "system")] @@ -1153,9 +1157,9 @@ def create_ui(startup_timer = None): inputs=components, outputs=[text_settings, result], ) - defaults_submit.click(fn=lambda: modules.shared.restore_defaults(restart=True), _js="restart_reload") - restart_submit.click(fn=lambda: modules.shared.restart_server(restart=True), _js="restart_reload") - shutdown_submit.click(fn=lambda: modules.shared.restart_server(restart=False), _js="restart_reload") + defaults_submit.click(fn=lambda: modules.shared.restore_defaults(restart=True), _js="restartReload") + restart_submit.click(fn=lambda: modules.shared.restart_server(restart=True), _js="restartReload") + shutdown_submit.click(fn=lambda: modules.shared.restart_server(restart=False), _js="restartReload") for _i, k, _item in quicksettings_list: component = component_dict[k] @@ -1190,6 +1194,21 @@ def create_ui(startup_timer = None): outputs=[component_dict['sd_vae'], text_settings], ) + def reference_submit(model): + from modules import modelloader + loaded = modelloader.load_reference(model) + if loaded: + return model if loaded else opts.sd_model_checkpoint + print('HERE', model, loaded) + return loaded + + button_set_reference = gr.Button('Change reference', elem_id='change_reference', visible=False) + button_set_reference.click( + fn=reference_submit, + _js="function(v){ return desiredCheckpointName; }", + inputs=[component_dict['sd_model_checkpoint']], + outputs=[component_dict['sd_model_checkpoint']], + ) component_keys = [k for k in opts.data_labels.keys() if k in component_dict] def get_settings_values(): diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index 9cc5c24cc..1804a1199 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -15,19 +15,19 @@ from collections import OrderedDict import gradio as gr from PIL import Image from starlette.responses import FileResponse, JSONResponse -from modules import shared, scripts, modelloader +from modules import paths, shared, scripts, modelloader from modules.ui_components import ToolButton import modules.ui_symbols as symbols + allowed_dirs = [] dir_cache = {} # key=path, value=(mtime, listdir(path)) refresh_time = 0 extra_pages = shared.extra_networks debug = shared.log.info if os.environ.get('SD_EN_DEBUG', None) is not None else lambda *args, **kwargs: None card_full = ''' -
+
-
{title}
@@ -39,22 +39,21 @@ card_full = '''
''' card_list = ''' -
+
🛈 
{title}
 
-
''' def listdir(path): - debug(f'EN list-dir: {path}') if not os.path.exists(path): return [] if path in dir_cache and os.path.getmtime(path) == dir_cache[path][0]: return dir_cache[path][1] else: + # debug(f'EN list-dir list: {path}') dir_cache[path] = (os.path.getmtime(path), [os.path.join(path, f) for f in os.listdir(path)]) return dir_cache[path][1] @@ -138,6 +137,9 @@ class ExtraNetworksPage: self.refresh_time = 0 self.page_time = 0 self.list_time = 0 + self.info_time = 0 + self.desc_time = 0 + self.dirs = {} self.view = shared.opts.extra_networks_view self.card = card_full if shared.opts.extra_networks_view == 'gallery' else card_list @@ -210,7 +212,6 @@ class ExtraNetworksPage: self.missing_thumbs.clear() def create_items(self, tabname): - debug(f'EN create-items: {self.name}') if self.refresh_time is not None and self.refresh_time > refresh_time: # cached results return t0 = time.time() @@ -223,7 +224,8 @@ class ExtraNetworksPage: for item in self.items: self.metadata[item["name"]] = item.get("metadata", {}) t1 = time.time() - self.list_time = round(t1-t0, 2) + debug(f'EN create-items: page={self.name} items={len(self.items)} time={t1-t0:.2f}') + self.list_time += t1-t0 def create_page(self, tabname, skip = False): @@ -237,8 +239,11 @@ class ExtraNetworksPage: allowed_folders = [os.path.abspath(x) for x in self.allowed_directories_for_previews()] for parentdir, dirs in {d: modelloader.directory_directories(d) for d in allowed_folders}.items(): for tgt in dirs.keys(): - if shared.opts.diffusers_dir in tgt: - subdirs[os.path.basename(shared.opts.diffusers_dir)] = 1 + if shared.backend == shared.Backend.DIFFUSERS: + if os.path.join(paths.models_path, 'Reference') in tgt: + subdirs['Reference'] = 1 + if shared.opts.diffusers_dir in tgt: + subdirs[os.path.basename(shared.opts.diffusers_dir)] = 1 if 'models--' in tgt: continue subdir = tgt[len(parentdir):].replace("\\", "/") @@ -255,6 +260,7 @@ class ExtraNetworksPage: self.create_items(tabname) self.create_xyz_grid() htmls = [] + self.items.sort(key=lambda x: x["mtime"], reverse=True) for item in self.items: htmls.append(self.create_html(item, tabname)) self.html += ''.join(htmls) @@ -263,7 +269,7 @@ class ExtraNetworksPage: self.html = f"
{subdirs_html}
{self.html}
" else: return '' - 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={self.list_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()} list={self.list_time:.2f} desc={self.desc_time:.2f} info={self.info_time:.2f}") if len(self.missing_thumbs) > 0: threading.Thread(target=self.create_thumb).start() return self.html @@ -280,7 +286,7 @@ class ExtraNetworksPage: "tabname": tabname, "page": self.name, "name": item["name"], - "title": item["name"].replace('_', ' '), + "title": os.path.basename(item["name"].replace('_', ' ')), "filename": item["filename"], "tags": '|'.join([item.get("tags")] if isinstance(item.get("tags", {}), str) else list(item.get("tags", {}).keys())), "preview": html.escape(item.get("preview", self.link_preview('html/card-no-preview.png'))), @@ -288,7 +294,7 @@ class ExtraNetworksPage: "height": shared.opts.extra_networks_card_size if shared.opts.extra_networks_card_square else 'auto', "fit": shared.opts.extra_networks_card_fit, "prompt": item.get("prompt", None), - "search_term": item.get("search_term", ""), + "search": item.get("search_term", ""), "description": item.get("description") or "", "card_click": item.get("onclick", '"' + html.escape(f'return cardClicked({item.get("prompt", None)}, {"true" if self.allow_negative_prompt else "false"})') + '"'), "mtime": item.get("mtime", 0), @@ -305,8 +311,9 @@ class ExtraNetworksPage: def find_preview_file(self, path): fn = os.path.splitext(path)[0] preview_extensions = ["jpg", "jpeg", "png", "webp", "tiff", "jp2"] + files = listdir(os.path.dirname(path)) for file in [f'{fn}{mid}{ext}' for ext in preview_extensions for mid in ['.thumb.', '.preview.', '.']]: - if os.path.exists(file): + if file in files: return file return 'html/card-no-preview.png' @@ -315,14 +322,16 @@ class ExtraNetworksPage: return self.link_preview('html/card-no-preview.png') fn = os.path.splitext(path)[0] preview_extensions = ["jpg", "jpeg", "png", "webp", "tiff", "jp2"] + files = listdir(os.path.dirname(path)) for file in [f'{fn}{mid}{ext}' for ext in preview_extensions for mid in ['.thumb.', '.', '.preview.']]: - if os.path.exists(file): + if file in files: if '.thumb.' not in file: self.missing_thumbs.append(file) return self.link_preview(file) return self.link_preview('html/card-no-preview.png') - def find_description(self, path): + def find_description(self, path, info=None): + t0 = time.time() class HTMLFilter(HTMLParser): text = "" def handle_data(self, data): @@ -332,7 +341,8 @@ class ExtraNetworksPage: self.text += '\n' fn = os.path.splitext(path)[0] + '.txt' - if os.path.exists(fn): + # if os.path.exists(fn): + if fn in listdir(os.path.dirname(path)): try: with open(fn, "r", encoding="utf-8", errors="replace") as f: txt = f.read() @@ -340,20 +350,27 @@ class ExtraNetworksPage: return txt except OSError: pass - info = self.find_info(path) + if info is None: + info = self.find_info(path) desc = info.get('description', '') or '' f = HTMLFilter() f.feed(desc) + t1 = time.time() + self.desc_time += t1-t0 return f.text def find_info(self, path): + t0 = time.time() fn = os.path.splitext(path)[0] + '.json' - if os.path.exists(fn): + # if os.path.exists(fn): + data = {} + if fn in listdir(os.path.dirname(path)): data = shared.readfile(fn, silent=True) if type(data) is list: data = data[0] - return data - return {} + t1 = time.time() + self.info_time += t1-t0 + return data def initialize(): @@ -524,8 +541,8 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): for page in get_pages(): page.create_page(ui.tabname, skip_indexing) with gr.Tab(page.title, id=page.title.lower().replace(" ", "_"), elem_classes="extra-networks-tab") as tab: - hmtl = gr.HTML(page.html, elem_id=f'{tabname}{page.name}_extra_page', elem_classes="extra-networks-page") - ui.pages.append(hmtl) + page_html = gr.HTML(page.html, elem_id=f'{tabname}{page.name}_extra_page', elem_classes="extra-networks-page") + ui.pages.append(page_html) tab.select(ui_tab_change, _js="getENActivePage", inputs=[ui.button_details], outputs=[ui.button_scan, ui.button_save, ui.button_model]) # ui.tabs.change(fn=ui_tab_change, inputs=[], outputs=[ui.button_scan, ui.button_save]) @@ -724,7 +741,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): return ui_refresh_click(title) def ui_save_click(): - from modules import paths, generation_parameters_copypaste + from modules import generation_parameters_copypaste filename = os.path.join(paths.data_path, "params.txt") if os.path.exists(filename): with open(filename, "r", encoding="utf8") as file: @@ -736,7 +753,7 @@ def create_ui(container, button_parent, tabname, skip_indexing = False): return res def ui_quicksave_click(name): - from modules import paths, generation_parameters_copypaste + from modules import generation_parameters_copypaste fn = os.path.join(paths.data_path, "params.txt") if os.path.exists(fn): with open(fn, "r", encoding="utf8") as file: diff --git a/modules/ui_extra_networks_checkpoints.py b/modules/ui_extra_networks_checkpoints.py index b2ffbd80f..30849c088 100644 --- a/modules/ui_extra_networks_checkpoints.py +++ b/modules/ui_extra_networks_checkpoints.py @@ -1,8 +1,9 @@ import html import json import os -from modules import shared, ui_extra_networks, sd_models +from modules import shared, ui_extra_networks, sd_models, paths +reference_dir = os.path.join(paths.models_path, 'Reference') class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): def __init__(self): @@ -11,12 +12,35 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): def refresh(self): shared.refresh_checkpoints() + def list_reference(self): + if shared.backend != shared.Backend.DIFFUSERS: + return [] + reference_models = shared.readfile(os.path.join('html', 'reference.json')) + for k, v in reference_models.items(): + name = os.path.join(reference_dir, k) + yield { + "type": 'Model', + "name": name, + "title": name, + "filename": v['path'], + "search_term": self.search_terms_from_path(name), + "preview": self.find_preview(os.path.join(reference_dir, os.path.basename(v['path']))), + "local_preview": f"{os.path.splitext(name)[0]}.{shared.opts.samples_format}", + "onclick": '"' + html.escape(f"""return selectReference({json.dumps(v['path'])})""") + '"', + "hash": None, + "mtime": 0, + "size": 0, + "info": {}, + "metadata": {}, + "description": v.get('desc', ''), + } + def list_items(self): checkpoint: sd_models.CheckpointInfo checkpoints = sd_models.checkpoints_list.copy() for name, checkpoint in checkpoints.items(): try: - fn = os.path.splitext(checkpoint.filename)[0] + exists = os.path.exists(checkpoint.filename) record = { "type": 'Model', "name": checkpoint.name, @@ -24,18 +48,20 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): "filename": checkpoint.filename, "hash": checkpoint.shorthash, "search_term": self.search_terms_from_path(checkpoint.title), - "preview": self.find_preview(fn), - "local_preview": f"{fn}.{shared.opts.samples_format}", - "description": self.find_description(fn), - "info": self.find_info(fn), + "preview": self.find_preview(checkpoint.filename), + "local_preview": f"{os.path.splitext(checkpoint.filename)[0]}.{shared.opts.samples_format}", "metadata": checkpoint.metadata, "onclick": '"' + html.escape(f"""return selectCheckpoint({json.dumps(name)})""") + '"', - "mtime": os.path.getmtime(checkpoint.filename), - "size": os.path.getsize(checkpoint.filename), + "mtime": os.path.getmtime(checkpoint.filename) if exists else 0, + "size": os.path.getsize(checkpoint.filename) if exists else 0, } + record["info"] = self.find_info(checkpoint.filename) + record["description"] = self.find_description(checkpoint.filename, record["info"]) yield record except Exception as e: shared.log.debug(f"Extra networks error: type=model file={name} {e}") + for record in self.list_reference(): + yield record def allowed_directories_for_previews(self): - return [v for v in [shared.opts.ckpt_dir, shared.opts.diffusers_dir, sd_models.model_path] if v is not None] + return [v for v in [shared.opts.ckpt_dir, shared.opts.diffusers_dir, reference_dir, sd_models.model_path] if v is not None] diff --git a/modules/ui_extra_networks_hypernets.py b/modules/ui_extra_networks_hypernets.py index cf6a02ae7..2fde91634 100644 --- a/modules/ui_extra_networks_hypernets.py +++ b/modules/ui_extra_networks_hypernets.py @@ -14,17 +14,17 @@ class ExtraNetworksPageHypernetworks(ui_extra_networks.ExtraNetworksPage): for name, path in shared.hypernetworks.items(): try: fn = os.path.splitext(path)[0] - name = os.path.relpath(fn, shared.opts.hypernetwork_dir) + name = os.path.relpath(os.path.splitext(path)[0], shared.opts.hypernetwork_dir) yield { "type": 'Hypernetwork', - "name": os.path.relpath(fn, shared.opts.hypernetwork_dir), + "name": name, "filename": path, - "preview": self.find_preview(fn), - "description": self.find_description(fn), - "info": self.find_info(fn), + "preview": self.find_preview(path), + "description": self.find_description(path), + "info": self.find_info(path), "search_term": self.search_terms_from_path(name), - "prompt": json.dumps(f""), - "local_preview": f"{fn}.{shared.opts.samples_format}", + "prompt": json.dumps(f""), + "local_preview": f"{os.path.splitext(path)[0]}.{shared.opts.samples_format}", "mtime": os.path.getmtime(path), "size": os.path.getsize(path), } diff --git a/modules/ui_extra_networks_textual_inversion.py b/modules/ui_extra_networks_textual_inversion.py index 5c750b813..004e4cecc 100644 --- a/modules/ui_extra_networks_textual_inversion.py +++ b/modules/ui_extra_networks_textual_inversion.py @@ -45,20 +45,21 @@ class ExtraNetworksPageTextualInversion(ui_extra_networks.ExtraNetworksPage): if embedding.tag is not None: tags[embedding.tag]=1 name = os.path.splitext(embedding.basename)[0] - yield { + record = { "type": 'Embedding', "name": name, "filename": embedding.filename, - "preview": self.find_preview(path), - "description": self.find_description(path), - "info": self.find_info(path), + "preview": self.find_preview(embedding.filename), "search_term": self.search_terms_from_path(name), - "prompt": json.dumps(os.path.splitext(embedding.name)[0]), + "prompt": json.dumps(f" {os.path.splitext(embedding.name)[0]}"), "local_preview": f"{path}.{shared.opts.samples_format}", "tags": tags, "mtime": os.path.getmtime(embedding.filename), "size": os.path.getsize(embedding.filename), } + record["info"] = self.find_info(embedding.filename) + record["description"] = self.find_description(embedding.filename, record["info"]) + yield record except Exception as e: shared.log.debug(f"Extra networks error: type=embedding file={embedding.filename} {e}") diff --git a/modules/ui_extra_networks_vae.py b/modules/ui_extra_networks_vae.py index 09e7565e9..e05920a6b 100644 --- a/modules/ui_extra_networks_vae.py +++ b/modules/ui_extra_networks_vae.py @@ -14,23 +14,22 @@ class ExtraNetworksPageVAEs(ui_extra_networks.ExtraNetworksPage): def list_items(self): for name, filename in sd_vae.vae_dict.items(): try: - fn = os.path.splitext(filename)[0] record = { "type": 'VAE', "name": name, "title": name, - "filename": fn, - "hash": hashes.sha256_from_cache(filename, f"vae/{fn}"), - "search_term": self.search_terms_from_path(fn), - "preview": self.find_preview(fn), - "local_preview": f"{fn}.{shared.opts.samples_format}", - "description": self.find_description(fn), - "info": self.find_info(fn), + "filename": filename, + "hash": hashes.sha256_from_cache(filename, f"vae/{filename}"), + "search_term": self.search_terms_from_path(filename), + "preview": self.find_preview(filename), + "local_preview": f"{os.path.splitext(filename)[0]}.{shared.opts.samples_format}", "metadata": {}, "onclick": '"' + html.escape(f"""return selectVAE({json.dumps(name)})""") + '"', "mtime": os.path.getmtime(filename), "size": os.path.getsize(filename), } + record["info"] = self.find_info(filename) + record["description"] = self.find_description(filename, record["info"]) yield record except Exception as e: shared.log.debug(f"Extra networks error: type=vae file={filename} {e}") diff --git a/modules/ui_interrogate.py b/modules/ui_interrogate.py new file mode 100644 index 000000000..3adc5b72c --- /dev/null +++ b/modules/ui_interrogate.py @@ -0,0 +1,267 @@ +import os +import base64 +from io import BytesIO +import gradio as gr +import open_clip +import torch +from PIL import Image +from pydantic import BaseModel, Field # pylint: disable=no-name-in-module +from fastapi import FastAPI +from fastapi.exceptions import HTTPException +from clip_interrogator import Config, Interrogator +import modules.generation_parameters_copypaste as parameters_copypaste +from modules import devices, lowvram, shared, paths + + +ci = None +low_vram = False + + +class BatchWriter: + def __init__(self, folder): + self.folder = folder + self.csv, self.file = None, None + + def add(self, file, prompt): + txt_file = os.path.splitext(file)[0] + ".txt" + with open(os.path.join(self.folder, txt_file), 'w', encoding='utf-8') as f: + f.write(prompt) + + def close(self): + if self.file is not None: + self.file.close() + + +def load(clip_model_name): + global ci # pylint: disable=global-statement + if ci is None: + config = Config(device=devices.get_optimal_device(), cache_path=os.path.join(paths.models_path, 'clip-interrogator'), clip_model_name=clip_model_name, quiet=True) + if low_vram: + config.apply_low_vram_defaults() + shared.log.info(f'Interrogate load: config={config}') + ci = Interrogator(config) + elif clip_model_name != ci.config.clip_model_name: + ci.config.clip_model_name = clip_model_name + shared.log.info(f'Interrogate load: config={ci.config}') + ci.load_clip_model() + + +def unload(): + if ci is not None: + shared.log.debug('Interrogate offload') + ci.caption_model = ci.caption_model.to(devices.cpu) + ci.clip_model = ci.clip_model.to(devices.cpu) + ci.caption_offloaded = True + ci.clip_offloaded = True + devices.torch_gc() + + +def image_analysis(image, clip_model_name): + load(clip_model_name) + image = image.convert('RGB') + image_features = ci.image_to_features(image) + top_mediums = ci.mediums.rank(image_features, 5) + top_artists = ci.artists.rank(image_features, 5) + top_movements = ci.movements.rank(image_features, 5) + top_trendings = ci.trendings.rank(image_features, 5) + top_flavors = ci.flavors.rank(image_features, 5) + medium_ranks = dict(zip(top_mediums, ci.similarities(image_features, top_mediums))) + artist_ranks = dict(zip(top_artists, ci.similarities(image_features, top_artists))) + movement_ranks = dict(zip(top_movements, ci.similarities(image_features, top_movements))) + trending_ranks = dict(zip(top_trendings, ci.similarities(image_features, top_trendings))) + flavor_ranks = dict(zip(top_flavors, ci.similarities(image_features, top_flavors))) + return medium_ranks, artist_ranks, movement_ranks, trending_ranks, flavor_ranks + + +def interrogate(image, mode, caption=None): + shared.log.info(f'Interrogate: image={image} mode={mode} config={ci.config}') + if mode == 'best': + prompt = ci.interrogate(image, caption=caption) + elif mode == 'caption': + prompt = ci.generate_caption(image) if caption is None else caption + elif mode == 'classic': + prompt = ci.interrogate_classic(image, caption=caption) + elif mode == 'fast': + prompt = ci.interrogate_fast(image, caption=caption) + elif mode == 'negative': + prompt = ci.interrogate_negative(image) + else: + raise RuntimeError(f"Unknown mode {mode}") + return prompt + + +def image_to_prompt(image, mode, clip_model_name): + shared.state.begin() + shared.state.job = 'interrogate' + try: + if shared.cmd_opts.lowvram or shared.cmd_opts.medvram: + lowvram.send_everything_to_cpu() + devices.torch_gc() + load(clip_model_name) + image = image.convert('RGB') + shared.log.info(f'Interrogate: image={image} mode={mode} config={ci.config}') + prompt = interrogate(image, mode) + except Exception as e: + prompt = f"Exception {type(e)}" + shared.log.error(f'Interrogate: {e}') + shared.state.end() + return prompt + + +def get_models(): + return ['/'.join(x) for x in open_clip.list_pretrained()] + + +def batch_process(batch_files, batch_folder, batch_str, mode, clip_model, write): + files = [] + if batch_files is not None: + files += [f.name for f in batch_files] + if batch_folder is not None: + files += [f.name for f in batch_folder] + if batch_str is not None and len(batch_str) > 0 and os.path.exists(batch_str) and os.path.isdir(batch_str): + files += [os.path.join(batch_str, f) for f in os.listdir(batch_str) if f.lower().endswith(('.png', '.jpg', '.jpeg', '.webp'))] + if len(files) == 0: + shared.log.error('Interrogate batch no images') + return '' + shared.log.info(f'Interrogate batch: images={len(files)} mode={mode} config={ci.config}') + shared.state.begin() + shared.state.job = 'batch interrogate' + prompts = [] + try: + if shared.cmd_opts.lowvram or shared.cmd_opts.medvram: + lowvram.send_everything_to_cpu() + devices.torch_gc() + load(clip_model) + captions = [] + # first pass: generate captions + for file in files: + caption = "" + try: + if shared.state.interrupted: + break + image = Image.open(file).convert('RGB') + caption = ci.generate_caption(image) + except Exception as e: + shared.log.error(f'Interrogate caption: {e}') + finally: + captions.append(caption) + # second pass: interrogate + if write: + writer = BatchWriter(os.path.dirname(files[0])) + for idx, file in enumerate(files): + try: + if shared.state.interrupted: + break + image = Image.open(file).convert('RGB') + prompt = interrogate(image, mode, caption=captions[idx]) + prompts.append(prompt) + if write: + writer.add(file, prompt) + except OSError as e: + shared.log.error(f'Interrogate batch: {e}') + if write: + writer.close() + ci.config.quiet = False + unload() + except Exception as e: + shared.log.error(f'Interrogate batch: {e}') + shared.state.end() + return '\n\n'.join(prompts) + + +def create_ui(): + global low_vram # pylint: disable=global-statement + low_vram = shared.cmd_opts.lowvram or shared.cmd_opts.medvram + if not low_vram and torch.cuda.is_available(): + device = devices.get_optimal_device() + vram_total = torch.cuda.get_device_properties(device).total_memory + if vram_total <= 12*1024*1024*1024: + low_vram = True + with gr.Row(elem_id="interrogate_tab"): + with gr.Column(): + with gr.Tab("Image"): + with gr.Row(): + image = gr.Image(type='pil', label="Image") + with gr.Row(): + prompt = gr.Textbox(label="Prompt", lines=3) + with gr.Row(): + medium = gr.Label(label="Medium", num_top_classes=5) + artist = gr.Label(label="Artist", num_top_classes=5) + movement = gr.Label(label="Movement", num_top_classes=5) + trending = gr.Label(label="Trending", num_top_classes=5) + flavor = gr.Label(label="Flavor", num_top_classes=5) + with gr.Row(): + interrogate_btn = gr.Button("Interrogate", variant='primary') + analyze_btn = gr.Button("Analyze", variant='primary') + unload_btn = gr.Button("Unload") + with gr.Row(): + buttons = parameters_copypaste.create_buttons(["txt2img", "img2img", "extras"]) + for tabname, button in buttons.items(): + parameters_copypaste.register_paste_params_button(parameters_copypaste.ParamBinding(paste_button=button, tabname=tabname, source_text_component=prompt, source_image_component=image,)) + with gr.Tab("Batch"): + with gr.Row(): + batch_files = gr.File(label="Files", show_label=True, file_count='multiple', file_types=['image'], type='file', interactive=True, height=100) + with gr.Row(): + batch_folder = gr.File(label="Folder", show_label=True, file_count='directory', file_types=['image'], type='file', interactive=True, height=100) + with gr.Row(): + batch_str = gr.Text(label="Folder", value="", interactive=True) + with gr.Row(): + batch = gr.Text(label="Prompts", lines=10) + with gr.Row(): + write = gr.Checkbox(label='Write prompts to files', value=False) + with gr.Row(): + batch_btn = gr.Button("Interrogate", variant='primary') + with gr.Column(): + with gr.Row(): + clip_model = gr.Dropdown(get_models(), value='ViT-L-14/openai', label='CLIP Model') + with gr.Row(): + mode = gr.Radio(['best', 'fast', 'classic', 'caption', 'negative'], label='Mode', value='best') + interrogate_btn.click(image_to_prompt, inputs=[image, mode, clip_model], outputs=prompt) + analyze_btn.click(image_analysis, inputs=[image, clip_model], outputs=[medium, artist, movement, trending, flavor]) + unload_btn.click(unload) + batch_btn.click(batch_process, inputs=[batch_files, batch_folder, batch_str, mode, clip_model, write], outputs=[batch]) + + +def decode_base64_to_image(encoding): + if encoding.startswith("data:image/"): + encoding = encoding.split(";")[1].split(",")[1] + try: + image = Image.open(BytesIO(base64.b64decode(encoding))) + return image + except Exception as e: + raise HTTPException(status_code=500, detail="Invalid encoded image") from e + + +def mount_interrogator_api(_: gr.Blocks, app: FastAPI): # TODO redesign interrogator api + + class InterrogatorAnalyzeRequest(BaseModel): + image: str = Field(default="", title="Image", description="Image to work on, must be a Base64 string containing the image's data.") + clip_model_name: str = Field(default="ViT-L-14/openai", title="Model", description="The interrogate model used. See the models endpoint for a list of available models.") + + class InterrogatorPromptRequest(InterrogatorAnalyzeRequest): + mode: str = Field(default="fast", title="Mode", description="The mode used to generate the prompt. Can be one of: best, fast, classic, negative.") + + @app.get("/interrogator/models") + async def api_get_models(): + return ["/".join(x) for x in open_clip.list_pretrained()] + + @app.post("/interrogator/prompt") + async def api_get_prompt(analyzereq: InterrogatorPromptRequest): + image_b64 = analyzereq.image + if image_b64 is None: + raise HTTPException(status_code=404, detail="Image not found") + img = decode_base64_to_image(image_b64) + prompt = image_to_prompt(img, analyzereq.mode, analyzereq.clip_model_name) + return {"prompt": prompt} + + @app.post("/interrogator/analyze") + async def api_analyze(analyzereq: InterrogatorAnalyzeRequest): + image_b64 = analyzereq.image + if image_b64 is None: + raise HTTPException(status_code=404, detail="Image not found") + + img = decode_base64_to_image(image_b64) + (medium_ranks, artist_ranks, movement_ranks, trending_ranks, flavor_ranks) = image_analysis(img, analyzereq.clip_model_name) + return {"medium": medium_ranks, "artist": artist_ranks, "movement": movement_ranks, "trending": trending_ranks, "flavor": flavor_ranks} + +# script_callbacks.on_app_started(mount_interrogator_api) diff --git a/modules/ui_postprocessing.py b/modules/ui_postprocessing.py index b28597653..334afcc9a 100644 --- a/modules/ui_postprocessing.py +++ b/modules/ui_postprocessing.py @@ -25,7 +25,7 @@ def create_ui(): with gr.TabItem('Single Image', id="single_image", elem_id="extras_single_tab") as tab_single: extras_image = gr.Image(label="Source", source="upload", interactive=True, type="pil", elem_id="extras_image") with gr.TabItem('Process Batch', id="batch_process", elem_id="extras_batch_process_tab") as tab_batch: - image_batch = gr.Files(label="Batch Process", interactive=True, elem_id="extras_image_batch") + image_batch = gr.Files(label="Batch process", interactive=True, elem_id="extras_image_batch") with gr.TabItem('Process Folder', id="batch_from_directory", elem_id="extras_batch_directory_tab") as tab_batch_dir: extras_batch_input_dir = gr.Textbox(label="Input directory", **shared.hide_dirs, placeholder="A directory on the same machine where the server is running.", elem_id="extras_batch_input_dir") extras_batch_output_dir = gr.Textbox(label="Output directory", **shared.hide_dirs, placeholder="Leave blank to save images to the default path.", elem_id="extras_batch_output_dir") diff --git a/modules/ui_tempdir.py b/modules/ui_tempdir.py index 72fe53b7d..f9b38de14 100644 --- a/modules/ui_tempdir.py +++ b/modules/ui_tempdir.py @@ -3,7 +3,7 @@ import tempfile from collections import namedtuple from pathlib import Path import gradio as gr -from PIL import PngImagePlugin +from PIL import Image, PngImagePlugin from modules import shared, errors @@ -36,7 +36,7 @@ def check_tmp_file(gradio, filename): return ok -def pil_to_temp_file(self, img, dir: str, format="png") -> str: # pylint: disable=redefined-builtin,unused-argument +def pil_to_temp_file(self, img: Image, dir: str, format="png") -> str: # pylint: disable=redefined-builtin,unused-argument """ # original gradio implementation bytes_data = gr.processing_utils.encode_pil_to_bytes(img, format) @@ -62,9 +62,12 @@ def pil_to_temp_file(self, img, dir: str, format="png") -> str: # pylint: disabl if isinstance(key, str) and isinstance(value, str): metadata.add_text(key, value) use_metadata = True + if not os.path.exists(dir): + os.makedirs(dir, exist_ok=True) + shared.log.debug(f'Created temp folder: path="{dir}"') with tempfile.NamedTemporaryFile(delete=False, suffix=".png", dir=dir) as tmp: - img.save(tmp, pnginfo=(metadata if use_metadata else None)) name = tmp.name + img.save(name, pnginfo=(metadata if use_metadata else None)) shared.log.debug(f'Saving temp: image="{name}"') return name diff --git a/requirements.txt b/requirements.txt index 68b614704..f5fb52c44 100644 --- a/requirements.txt +++ b/requirements.txt @@ -45,6 +45,7 @@ dctorch httpx==0.24.1 compel==2.0.2 torchsde==0.2.6 +clip-interrogator==0.6.0 antlr4-python3-runtime==4.9.3 requests==2.31.0 tqdm==4.66.1 diff --git a/webui.bat b/webui.bat index 2d12762c6..11aaf6b3a 100755 --- a/webui.bat +++ b/webui.bat @@ -1,3 +1,7 @@ +:: -------------------------------------------------------------------------------------------------------------- +:: Do not make any changes to this file, change the variables in webui-user.bat instead and call this file +:: -------------------------------------------------------------------------------------------------------------- + @echo off if not defined PYTHON (set PYTHON=python) diff --git a/webui.ps1 b/webui.ps1 index eac963fea..a607ac6f1 100755 --- a/webui.ps1 +++ b/webui.ps1 @@ -1,3 +1,7 @@ +# -------------------------------------------------------------------------------------------------------------- +# Do not make any changes to this file, change the variables in webui-user.ps1 instead and call this file +# -------------------------------------------------------------------------------------------------------------- + function ShowStdOutStdErr { Write-Output "exit code: $LASTEXITCODE" diff --git a/webui.sh b/webui.sh index b05921579..692cc36c1 100755 --- a/webui.sh +++ b/webui.sh @@ -1,8 +1,7 @@ #!/usr/bin/env bash -################################################# -# Please do not make any changes to this file, # -# change the variables in webui-user.sh instead # -################################################# +# ------------------------------------------------------------------------------------------------------------- +# Do not make any changes to this file, change the variables in webui-user.sh instead and call this file +# ------------------------------------------------------------------------------------------------------------- # change to local directory cd -- "$(dirname -- "$0")" @@ -18,9 +17,9 @@ then fi # python3 executable -if [[ -z "${python_cmd}" ]] +if [[ -z "${PYTHON}" ]] then - python_cmd="python3" + PYTHON="python3" fi # git executable @@ -51,7 +50,7 @@ then exit 1 fi -for preq in "${GIT}" "${python_cmd}" +for preq in "${GIT}" "${PYTHON}" do if ! hash "${preq}" &>/dev/null then @@ -60,7 +59,7 @@ do fi done -if ! "${python_cmd}" -c "import venv" &>/dev/null +if ! "${PYTHON}" -c "import venv" &>/dev/null then echo "Error: python3-venv is not installed" exit 1 @@ -69,7 +68,7 @@ fi echo "Create and activate python venv" if [[ ! -d "${venv_dir}" ]] then - "${python_cmd}" -m venv "${venv_dir}" + "${PYTHON}" -m venv "${venv_dir}" first_launch=1 fi @@ -102,5 +101,5 @@ then exec ipexrun --multi-task-manager 'taskset' --memory-allocator 'jemalloc' launch.py "$@" else echo "Launching launch.py..." - exec "${python_cmd}" launch.py "$@" + exec "${PYTHON}" launch.py "$@" fi diff --git a/wiki b/wiki index 3d5e2a213..e999774e3 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit 3d5e2a21302375a297ac11ba41df7131361b9fdd +Subproject commit e999774e3096ceb89a264548fdfaaa76d891c0df