From 36324361bd5df4efc92c3d14e5e11454c2548087 Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Thu, 21 Sep 2023 11:58:24 -0400 Subject: [PATCH] refactor en folder handling --- CHANGELOG.md | 8 +- extensions-builtin/Lora/lora.py | 16 +--- .../Lora/ui_extra_networks_lora.py | 16 ++-- javascript/extraNetworks.js | 2 +- modules/api/api.py | 7 +- modules/api/models.py | 7 +- modules/hashes.py | 6 +- modules/hijack/ddpm_edit.py | 6 +- modules/hypernetworks/hypernetwork.py | 2 +- modules/postprocess/swinir_model_arch.py | 5 +- modules/postprocess/swinir_model_arch_v2.py | 10 +-- modules/sd_models.py | 80 ++++++++----------- modules/shared.py | 1 - modules/shared_items.py | 5 -- modules/styles.py | 23 +++--- .../textual_inversion/textual_inversion.py | 13 ++- modules/ui.py | 2 +- modules/ui_extra_networks.py | 24 +++--- modules/ui_extra_networks_checkpoints.py | 22 ++--- modules/ui_extra_networks_hypernets.py | 15 ++-- modules/ui_extra_networks_styles.py | 21 ++--- .../ui_extra_networks_textual_inversion.py | 9 ++- modules/uni_pc/uni_pc.py | 4 +- wiki | 2 +- 24 files changed, 135 insertions(+), 171 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index e74e340cf..1bd947618 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,9 +2,11 @@ ## Update for 2023-09-20 -- Added **change log** to UI, see *System -> Changelog* -- **Extra networks**: faster search, ability to show/hide/sort networks -- **Upscalers**: complete refactor... +- Added **change log** to UI, see *System -> Changelog* +- **Extra networks**: + - faster search, ability to show/hide/sort networks + - refactored subfolder handling +- **Upscalers**: complete refactor... - more high quality upscalers available by default - unified init/download/execute/progress code - easier installation diff --git a/extensions-builtin/Lora/lora.py b/extensions-builtin/Lora/lora.py index f45595cf6..31b0ff8f2 100644 --- a/extensions-builtin/Lora/lora.py +++ b/extensions-builtin/Lora/lora.py @@ -224,21 +224,15 @@ def load_lora(name, lora_on_disk): def load_loras(names, multipliers=None): already_loaded = {} - for lora in loaded_loras: if lora.name in names: already_loaded[lora.name] = lora - loaded_loras.clear() - loras_on_disk = [available_lora_aliases.get(name, None) for name in names] if any(x is None for x in loras_on_disk): list_available_loras() - loras_on_disk = [available_lora_aliases.get(name, None) for name in names] - failed_to_load_loras = [] - recompile_model = False if shared.opts.cuda_compile and shared.opts.cuda_compile_backend == "openvino_fx": if len(names) == len(shared.compiled_model_state.lora_model): @@ -270,12 +264,10 @@ def load_loras(names, multipliers=None): continue lora.mentioned_name = name lora_on_disk.read_hash() - if lora is None: failed_to_load_loras.append(name) print(f"Couldn't find Lora with name {name}") continue - lora.multiplier = multipliers[i] if multipliers else 1.0 loaded_loras.append(lora) @@ -291,25 +283,20 @@ def lora_calc_updown(lora, module, target): with torch.no_grad(): up = module.up.weight.to(target.device, dtype=target.dtype) down = module.down.weight.to(target.device, dtype=target.dtype) - if up.shape[2:] == (1, 1) and down.shape[2:] == (1, 1): updown = (up.squeeze(2).squeeze(2) @ down.squeeze(2).squeeze(2)).unsqueeze(2).unsqueeze(3) elif up.shape[2:] == (3, 3) or down.shape[2:] == (3, 3): updown = torch.nn.functional.conv2d(down.permute(1, 0, 2, 3), up).permute(1, 0, 2, 3) else: updown = up @ down - updown = updown * lora.multiplier * (module.alpha / module.up.weight.shape[1] if module.alpha else 1.0) - return updown def lora_restore_weights_from_backup(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.MultiheadAttention]): weights_backup = getattr(self, "lora_weights_backup", None) - if weights_backup is None: return - if isinstance(self, torch.nn.MultiheadAttention): self.in_proj_weight.copy_(weights_backup[0]) self.out_proj.weight.copy_(weights_backup[1]) @@ -450,6 +437,7 @@ def lora_MultiheadAttention_load_state_dict(self, *args, **kwargs): def list_available_loras(): + from modules.paths_internal import script_path available_loras.clear() available_lora_aliases.clear() forbidden_lora_aliases.clear() @@ -457,6 +445,8 @@ def list_available_loras(): forbidden_lora_aliases.update({"none": 1, "Addams": 1}) os.makedirs(shared.cmd_opts.lora_dir, exist_ok=True) for filename in sorted([*filter(extension_filter(['.PT', '.CKPT', '.SAFETENSORS']), directory_files(shared.cmd_opts.lora_dir))], key=str.lower): + if filename.startswith(script_path): + filename = os.path.relpath(filename, script_path) name = os.path.splitext(os.path.basename(filename))[0] entry = LoraOnDisk(name, filename) available_loras[name] = entry diff --git a/extensions-builtin/Lora/ui_extra_networks_lora.py b/extensions-builtin/Lora/ui_extra_networks_lora.py index e99105142..e416f959d 100644 --- a/extensions-builtin/Lora/ui_extra_networks_lora.py +++ b/extensions-builtin/Lora/ui_extra_networks_lora.py @@ -1,5 +1,5 @@ -import json import os +import json import lora from modules import shared, ui_extra_networks @@ -15,10 +15,6 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): def list_items(self): for name, l in lora.available_loras.items(): path, _ext = os.path.splitext(l.filename) - alias = l.get_alias() - prompt = f" " - prompt = json.dumps(prompt) - metadata = json.dumps(l.metadata, indent=4) if l.metadata else None possible_tags = l.metadata.get('ss_tag_frequency', {}) if l.metadata is not None else {} if isinstance(possible_tags, str): possible_tags = {} @@ -26,18 +22,18 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage): for k, v in possible_tags.items(): words = k.split('_', 1) if '_' in k else [v, k] tags[' '.join(words[1:])] = words[0] + name = os.path.splitext(os.path.relpath(l.filename, shared.cmd_opts.lora_dir))[0] yield { "name": name, - "filename": path, - "fullname": l.filename, + "filename": l.filename, "hash": l.shorthash, + "search_term": self.search_terms_from_path(l.filename) + ' '.join(tags.keys()), "preview": self.find_preview(path), "description": self.find_description(path), "info": self.find_info(path), - "search_term": self.search_terms_from_path(l.filename) + ' '.join(tags.keys()), - "prompt": prompt, + "prompt": json.dumps(f" "), "local_preview": f"{path}.{shared.opts.samples_format}", - "metadata": metadata, + "metadata": json.dumps(l.metadata, indent=4) if l.metadata else None, "tags": tags, } diff --git a/javascript/extraNetworks.js b/javascript/extraNetworks.js index 43ef498a8..1b6de2316 100644 --- a/javascript/extraNetworks.js +++ b/javascript/extraNetworks.js @@ -308,7 +308,7 @@ function extraNetworksSearchButton(event) { const tabname = getENActiveTab(); const searchTextarea = gradioApp().querySelector(`#${tabname}_extra_tabs > div > div > textarea`); const button = event.target; - const text = button.classList.contains('search-all') ? '' : `/${button.textContent.trim()}/`; + const text = button.classList.contains('search-all') ? '' : `${button.textContent.trim()}/`; searchTextarea.value = text; updateInput(searchTextarea); } diff --git a/modules/api/api.py b/modules/api/api.py index 6a053cb03..6b436d7f3 100644 --- a/modules/api/api.py +++ b/modules/api/api.py @@ -129,7 +129,6 @@ class Api: self.add_api_route("/sdapi/v1/sd-models", self.get_sd_models, methods=["GET"], response_model=List[models.SDModelItem]) 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.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"]) @@ -468,7 +467,7 @@ class Api: ] def get_sd_models(self): - return [{"title": x.title, "model_name": x.model_name, "hash": x.shorthash, "sha256": x.sha256, "filename": x.filename, "config": find_checkpoint_config_near_filename(x)} for x in checkpoints_list.values()] + return [{"title": x.title, "name": x.name, "filename": x.filename, "type": x.type, "hash": x.shorthash, "sha256": x.sha256, "config": find_checkpoint_config_near_filename(x)} for x in checkpoints_list.values()] def get_hypernetworks(self): return [{"name": name, "path": shared.hypernetworks[name]} for name in shared.hypernetworks] @@ -476,10 +475,6 @@ class Api: def get_face_restorers(self): return [{"name":x.name(), "cmd_dir": getattr(x, "cmd_dir", None)} for x in shared.face_restorers] - def get_realesrgan_models(self): - from modules.postprocess.realesrgan_model import get_realesrgan_models - return [{"name":x.name,"path":x.data_path, "scale":x.scale} for x in get_realesrgan_models(None)] - def get_prompt_styles(self): return [{ 'name': v.name, 'prompt': v.prompt, 'negative_prompt': v.negative_prompt, 'extra': v.extra, 'filename': v.filename, 'preview': v.preview} for v in shared.prompt_styles.styles.values()] diff --git a/modules/api/models.py b/modules/api/models.py index 9241592eb..3ed73a42e 100644 --- a/modules/api/models.py +++ b/modules/api/models.py @@ -245,10 +245,11 @@ class UpscalerItem(BaseModel): class SDModelItem(BaseModel): title: str = Field(title="Title") - model_name: str = Field(title="Model Name") - hash: Optional[str] = Field(title="Short hash") - sha256: Optional[str] = Field(title="sha256 hash") + name: str = Field(title="Model Name") filename: str = Field(title="Filename") + type: str = Field(title="Model type") + sha256: Optional[str] = Field(title="SHA256 hash") + hash: Optional[str] = Field(title="Short hash") config: Optional[str] = Field(title="Config file") class HypernetworkItem(BaseModel): diff --git a/modules/hashes.py b/modules/hashes.py index 288a175d5..4bbc80850 100644 --- a/modules/hashes.py +++ b/modules/hashes.py @@ -25,7 +25,7 @@ def calculate_sha256(filename, quiet=False): hash_sha256 = hashlib.sha256() blksize = 1024 * 1024 if not quiet: - with progress.open(filename, 'rb', description=f'Calculating model hash: [cyan]{filename}', auto_refresh=True, console=shared.console) as f: + with progress.open(filename, 'rb', description=f'Calculating hash: [cyan]{filename}', auto_refresh=True, console=shared.console) as f: for chunk in iter(lambda: f.read(blksize), b""): hash_sha256.update(chunk) else: @@ -56,8 +56,9 @@ def sha256(filename, title, use_addnet_hash=False): return None if not os.path.isfile(filename): return None + shared.state.begin("hashing") if use_addnet_hash: - with progress.open(filename, 'rb', description=f'Calculating model hash: [cyan]{filename}', auto_refresh=True, console=shared.console) as f: + with progress.open(filename, 'rb', description=f'Calculating hash: [cyan]{filename}', auto_refresh=True, console=shared.console) as f: sha256_value = addnet_hash_safetensors(f) else: sha256_value = calculate_sha256(filename) @@ -65,6 +66,7 @@ def sha256(filename, title, use_addnet_hash=False): "mtime": os.path.getmtime(filename), "sha256": sha256_value } + shared.state.end() dump_cache() return sha256_value diff --git a/modules/hijack/ddpm_edit.py b/modules/hijack/ddpm_edit.py index ad067dd8b..ea25d60fb 100644 --- a/modules/hijack/ddpm_edit.py +++ b/modules/hijack/ddpm_edit.py @@ -1024,7 +1024,7 @@ class LatentDiffusion(DDPM): elif self.parameterization == "eps": target = noise else: - raise NotImplementedError() + raise NotImplementedError loss_simple = self.get_loss(model_output, target, mean=False).mean([1, 2, 3]) loss_dict.update({f'{prefix}/loss_simple': loss_simple.mean()}) @@ -1063,7 +1063,7 @@ class LatentDiffusion(DDPM): elif self.parameterization == "x0": x_recon = model_out else: - raise NotImplementedError() + raise NotImplementedError if clip_denoised: x_recon.clamp_(-1., 1.) @@ -1423,7 +1423,7 @@ class DiffusionWrapper(pl.LightningModule): cc = c_crossattn[0] out = self.diffusion_model(x, t, y=cc) else: - raise NotImplementedError() + raise NotImplementedError return out diff --git a/modules/hypernetworks/hypernetwork.py b/modules/hypernetworks/hypernetwork.py index 7b319b1b7..d36edb0d1 100644 --- a/modules/hypernetworks/hypernetwork.py +++ b/modules/hypernetworks/hypernetwork.py @@ -288,7 +288,7 @@ def list_hypernetworks(path): fn = os.path.join(folder, filename) if os.path.isfile(fn) and fn.lower().endswith(".pt"): name = os.path.splitext(os.path.basename(fn))[0] - res[name] = filename + res[name] = fn elif os.path.isdir(fn) and not fn.startswith('.'): list_folder(fn) diff --git a/modules/postprocess/swinir_model_arch.py b/modules/postprocess/swinir_model_arch.py index 4f6696a46..73898b3a5 100644 --- a/modules/postprocess/swinir_model_arch.py +++ b/modules/postprocess/swinir_model_arch.py @@ -279,8 +279,7 @@ class SwinTransformerBlock(nn.Module): return x def extra_repr(self) -> str: - return f"dim={self.dim}, input_resolution={self.input_resolution}, num_heads={self.num_heads}, " \ - f"window_size={self.window_size}, shift_size={self.shift_size}, mlp_ratio={self.mlp_ratio}" + return f"dim={self.dim}, input_resolution={self.input_resolution}, num_heads={self.num_heads}, window_size={self.window_size}, shift_size={self.shift_size}, mlp_ratio={self.mlp_ratio}" def flops(self): flops = 0 @@ -587,7 +586,7 @@ class Upsample(nn.Sequential): m.append(nn.Conv2d(num_feat, 9 * num_feat, 3, 1, 1)) m.append(nn.PixelShuffle(3)) else: - raise ValueError(f'scale {scale} is not supported. ' 'Supported scales: 2^n and 3.') + raise ValueError(f'scale {scale} is not supported. Supported scales: 2^n and 3.') super(Upsample, self).__init__(*m) diff --git a/modules/postprocess/swinir_model_arch_v2.py b/modules/postprocess/swinir_model_arch_v2.py index 991e3212d..19bcab441 100644 --- a/modules/postprocess/swinir_model_arch_v2.py +++ b/modules/postprocess/swinir_model_arch_v2.py @@ -174,8 +174,7 @@ class WindowAttention(nn.Module): return x def extra_repr(self) -> str: - return f'dim={self.dim}, window_size={self.window_size}, ' \ - f'pretrained_window_size={self.pretrained_window_size}, num_heads={self.num_heads}' + return f'dim={self.dim}, window_size={self.window_size}, pretrained_window_size={self.pretrained_window_size}, num_heads={self.num_heads}' def flops(self, N): # calculate flops for 1 window with token length of N @@ -307,8 +306,7 @@ class SwinTransformerBlock(nn.Module): return x def extra_repr(self) -> str: - return f"dim={self.dim}, input_resolution={self.input_resolution}, num_heads={self.num_heads}, " \ - f"window_size={self.window_size}, shift_size={self.shift_size}, mlp_ratio={self.mlp_ratio}" + return f"dim={self.dim}, input_resolution={self.input_resolution}, num_heads={self.num_heads}, window_size={self.window_size}, shift_size={self.shift_size}, mlp_ratio={self.mlp_ratio}" def flops(self): flops = 0 @@ -620,7 +618,7 @@ class Upsample(nn.Sequential): m.append(nn.Conv2d(num_feat, 9 * num_feat, 3, 1, 1)) m.append(nn.PixelShuffle(3)) else: - raise ValueError(f'scale {scale} is not supported. ' 'Supported scales: 2^n and 3.') + raise ValueError(f'scale {scale} is not supported. Supported scales: 2^n and 3.') super(Upsample, self).__init__(*m) class Upsample_hf(nn.Sequential): @@ -641,7 +639,7 @@ class Upsample_hf(nn.Sequential): m.append(nn.Conv2d(num_feat, 9 * num_feat, 3, 1, 1)) m.append(nn.PixelShuffle(3)) else: - raise ValueError(f'scale {scale} is not supported. ' 'Supported scales: 2^n and 3.') + raise ValueError(f'scale {scale} is not supported. Supported scales: 2^n and 3.') super(Upsample_hf, self).__init__(*m) diff --git a/modules/sd_models.py b/modules/sd_models.py index 4ace52368..69106650c 100644 --- a/modules/sd_models.py +++ b/modules/sd_models.py @@ -23,7 +23,7 @@ from modules import paths, shared, modelloader, devices, script_callbacks, sd_va from modules.sd_hijack_inpainting import do_inpainting_hijack from modules.timer import Timer from modules.memstats import memory_stats -from modules.paths_internal import models_path +from modules.paths_internal import models_path, script_path try: import diffusers @@ -44,71 +44,55 @@ sd_metadata_timer = 0 class CheckpointInfo: def __init__(self, filename): - name = '' self.name = None self.hash = None self.filename = filename self.type = '' - abspath = os.path.abspath(filename) + filename = os.path.abspath(filename) + if filename.startswith(script_path): + filename = os.path.relpath(filename, script_path) + relname = os.path.relpath(filename, model_path) + relname = os.path.relpath(filename, shared.cmd_opts.ckpt_dir) + relname, ext = os.path.splitext(relname) + ext = ext.lower()[1:] - if os.path.isfile(abspath): # ckpt or safetensor - if shared.opts.ckpt_dir is not None and abspath.startswith(shared.opts.ckpt_dir): - name = abspath.replace(shared.opts.ckpt_dir, '') - elif abspath.startswith(model_path): - name = abspath.replace(model_path, '') - else: - name = os.path.basename(filename) - if name.startswith("\\") or name.startswith("/"): - name = name[1:] - self.name = name - self.sha256 = hashes.sha256_from_cache(self.filename, f"checkpoint/{name}") - self.hash = self.sha256[0:8] if self.sha256 is not None else None - self.path = abspath - self.type = abspath.split('.')[-1].lower() - self.name_for_extra = os.path.splitext(os.path.basename(filename))[0] - self.model_name = os.path.splitext(name.replace("/", "_").replace("\\", "_"))[0] + if os.path.isfile(filename): # ckpt or safetensor + self.name = relname + self.filename = filename + self.sha256 = hashes.sha256_from_cache(self.filename, f"checkpoint/{os.path.basename(relname)}.{ext}") + self.type = ext + # self.model_name = os.path.splitext(name.replace("/", "_").replace("\\", "_"))[0] else: # maybe a diffuser repo = [r for r in modelloader.diffuser_repos if filename == r['filename']] if len(repo) == 0: - if filename.lower() != 'none': - shared.log.error(f'Cannot find diffuser model: {filename}') - else: - shared.log.info(f'Skipping model load: {filename}') - return - self.name = repo[0]['name'] - self.hash = repo[0]['hash'][:8] - self.sha256 = repo[0]['hash'] - self.path = repo[0]['path'] - self.type = 'diffusers' - self.name_for_extra = repo[0]['name'] - self.model_name = repo[0]['name'] - if os.path.isfile(repo[0]['model_info']): - file_path = repo[0]['model_info'] - self.model_info = shared.readfile(file_path, silent=True) + self.name = relname + self.filename = filename + self.sha256 = None + self.type = 'unknown' + else: + self.name = repo[0]['name'] + self.filename = repo[0]['path'] + self.sha256 = repo[0]['hash'] + self.type = 'diffusers' self.shorthash = self.sha256[0:10] if self.sha256 else None self.title = self.name if self.shorthash is None else f'{self.name} [{self.shorthash}]' - self.ids = [self.hash, self.model_name, self.title, self.name, f'{self.name} [{self.hash}]'] + ([self.shorthash, self.sha256, f'{self.name} [{self.shorthash}]'] if self.shorthash else []) - self.metadata = {} - _, ext = os.path.splitext(self.filename) - if ext.lower() == ".safetensors": - try: - self.metadata = read_metadata_from_safetensors(filename) - except Exception as e: - errors.display(e, f"reading checkpoint metadata: {filename}") + self.path = self.filename + self.model_name = os.path.basename(self.name) + # shared.log.debug(f'Checkpoint: type={self.type} name={self.name} filename={self.filename} hash={self.shorthash} title={self.title}') + self.metadata = read_metadata_from_safetensors(filename) def register(self): checkpoints_list[self.title] = self - for i in self.ids: - checkpoint_aliases[i] = self + for i in [self.name, self.filename, self.shorthash, self.title]: + if i is not None: + checkpoint_aliases[i] = self def calculate_shorthash(self): self.sha256 = hashes.sha256(self.filename, f"checkpoint/{self.name}") if self.sha256 is None: return self.shorthash = self.sha256[0:10] - if self.shorthash not in self.ids: - self.ids += [self.shorthash, self.sha256, f'{self.name} [{self.shorthash}]'] checkpoints_list.pop(self.title) self.title = f'{self.name} [{self.shorthash}]' self.register() @@ -340,9 +324,11 @@ def read_metadata_from_safetensors(filename): res = sd_metadata.get(filename, None) if res is not None: return res - res = {} + if not filename.endswith(".safetensors"): + return {} if shared.cmd_opts.no_metadata: return {} + res = {} try: t0 = time.time() with open(filename, mode="rb") as file: diff --git a/modules/shared.py b/modules/shared.py index 988bd3051..8de276203 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -656,7 +656,6 @@ options_templates.update(options_section(('extra_networks', "Extra Networks"), { "extra_networks_card_cover": OptionInfo("sidebar", "UI position", gr.Radio, lambda: {"choices": ["cover", "inline", "sidebar"]}), "extra_networks_height": OptionInfo(53, "UI height (%)", gr.Slider, {"minimum": 10, "maximum": 100, "step": 1}), "extra_networks_sidebar_width": OptionInfo(35, "UI sidebar width (%)", gr.Slider, {"minimum": 10, "maximum": 80, "step": 1}), - "extra_networks_card_lazy": OptionInfo(True, "UI card preview lazy loading", gr.Checkbox, { "visible": False }), "extra_networks_card_size": OptionInfo(160, "UI card size (px)", gr.Slider, {"minimum": 20, "maximum": 2000, "step": 1}), "extra_networks_card_square": OptionInfo(True, "UI disable variable aspect ratio"), "extra_networks_card_fit": OptionInfo("cover", "UI image contain method", gr.Radio, lambda: {"choices": ["contain", "cover", "fill"], "visible": False}), diff --git a/modules/shared_items.py b/modules/shared_items.py index 061fe4929..c095f7ac9 100644 --- a/modules/shared_items.py +++ b/modules/shared_items.py @@ -1,8 +1,3 @@ -def realesrgan_models_names(): - import modules.realesrgan_model - return [x.name for x in modules.realesrgan_model.get_realesrgan_models(None)] - - def postprocessing_scripts(): import modules.scripts return modules.scripts.scripts_postproc.scripts diff --git a/modules/styles.py b/modules/styles.py index 8491a43a6..f23c99c01 100644 --- a/modules/styles.py +++ b/modules/styles.py @@ -10,14 +10,13 @@ from modules import paths class Style(): def __init__(self, name: str, prompt: str = "", negative_prompt: str = "", extra: str = "", filename: str = "", preview: str = ""): - self.name = re.sub(r'[\t\r\n]', '', name).strip() + 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: if "{prompt}" in style_prompt: res = style_prompt.replace("{prompt}", prompt) @@ -61,13 +60,15 @@ class StyleDatabase: self.styles.clear() def list_folder(folder): for filename in os.listdir(folder): - fn = os.path.join(folder, filename) + fn = os.path.abspath(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", "")) + basename = os.path.splitext(os.path.basename(fn))[0] + name = re.sub(r'[\t\r\n]', '', style.get("name", basename)).strip() + name = os.path.join(os.path.dirname(os.path.relpath(fn, self.path)), name) + self.styles[style["name"]] = Style(name=name, prompt=style.get("prompt", ""), negative_prompt=style.get("negative", ""), extra=style.get("extra", ""), filename=fn, preview=style.get("preview", "")) except Exception as e: log.error(f'Failed to load style: file={fn} error={e}') elif os.path.isdir(fn) and not fn.startswith('.'): @@ -77,17 +78,21 @@ class StyleDatabase: self.styles = dict(sorted(self.styles.items(), key=lambda style: style[1].filename)) log.debug(f'Loaded styles: folder={self.path} items={len(self.styles.keys())}') + def find_style(self, name): + found = [style for style in self.styles.values() if style.name == name] + return found[0] if len(found) > 0 else self.no_style + def get_style_prompts(self, styles): - return [self.styles.get(x, self.no_style).prompt for x in styles] + return [self.find_style(x).prompt for x in styles] def get_negative_style_prompts(self, styles): - return [self.styles.get(x, self.no_style).negative_prompt for x in styles] + return [self.find_style(x).negative_prompt for x in styles] def apply_styles_to_prompt(self, prompt, styles): - return apply_styles_to_prompt(prompt, [self.styles.get(x, self.no_style).prompt for x in styles]) + return apply_styles_to_prompt(prompt, [self.find_style(x).prompt for x in styles]) def apply_negative_styles_to_prompt(self, prompt, styles): - return apply_styles_to_prompt(prompt, [self.styles.get(x, self.no_style).negative_prompt for x in styles]) + return apply_styles_to_prompt(prompt, [self.find_style(x).negative_prompt for x in styles]) def save_styles(self, path, verbose=False): for name in list(self.styles): diff --git a/modules/textual_inversion/textual_inversion.py b/modules/textual_inversion/textual_inversion.py index b959034ef..69126a337 100644 --- a/modules/textual_inversion/textual_inversion.py +++ b/modules/textual_inversion/textual_inversion.py @@ -29,18 +29,19 @@ def list_textual_inversion_templates(): class Embedding: - def __init__(self, vec, name, step=None): + def __init__(self, vec, name, filename=None, step=None): self.vec = vec self.name = name self.tag = name self.step = step + self.filename = filename + self.basename = os.path.relpath(filename, shared.opts.embeddings_dir) if filename is not None else None self.shape = None self.vectors = 0 self.cached_checksum = None self.sd_checkpoint = None self.sd_checkpoint_name = None self.optimizer_state_dict = None - self.filename = None def save(self, filename): embedding_data = { @@ -131,8 +132,7 @@ class EmbeddingDatabase: pipe.text_encoder.resize_token_embeddings(len(pipe.tokenizer)) return name = os.path.basename(fn) - embedding = Embedding(vec=None, name=name) - embedding.filename = path + embedding = Embedding(vec=None, name=name, filename=path) try: if hasattr(pipe,"load_textual_inversion"): pipe.load_textual_inversion(path, cache_dir=shared.opts.diffusers_dir, local_files_only=True) @@ -203,14 +203,13 @@ class EmbeddingDatabase: vec = emb.detach().to(devices.device, dtype=torch.float32) # name = data.get('name', name) - embedding = Embedding(vec, name) + embedding = Embedding(vec=vec, name=name, filename=path) embedding.tag = data.get('name', None) embedding.step = data.get('step', None) embedding.sd_checkpoint = data.get('sd_checkpoint', None) embedding.sd_checkpoint_name = data.get('sd_checkpoint_name', None) embedding.vectors = vec.shape[0] embedding.shape = vec.shape[-1] - embedding.filename = path if self.expected_shape == -1 or self.expected_shape == embedding.shape: self.register_embedding(embedding, shared.sd_model) else: @@ -291,7 +290,7 @@ def create_embedding(name, num_vectors_per_token, overwrite_old, init_text='*'): if not overwrite_old and os.path.exists(fn): shared.log.warning(f"Embedding already exists: {fn}") else: - embedding = Embedding(vec, name) + embedding = Embedding(vec=vec, name=name, filename=fn) embedding.step = 0 embedding.save(fn) shared.log.info(f'Created embedding: {fn} vectors {num_vectors_per_token} init {init_text}') diff --git a/modules/ui.py b/modules/ui.py index 292935024..5a7070ccd 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -267,7 +267,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=[k for k, v in modules.shared.prompt_styles.styles.items()], 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) # 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]) diff --git a/modules/ui_extra_networks.py b/modules/ui_extra_networks.py index bf70cf4d7..625e7dd1e 100644 --- a/modules/ui_extra_networks.py +++ b/modules/ui_extra_networks.py @@ -110,7 +110,7 @@ class ExtraNetworksPage: {card_extra} - + ''' @@ -139,13 +139,8 @@ class ExtraNetworksPage: preview = f"./sd_extra_networks/thumb?filename={quoted_filename}&mtime={mtime}" return preview - def search_terms_from_path(self, filename, possible_directories=None): - abspath = os.path.abspath(filename) - for parentdir in (possible_directories if possible_directories is not None else self.allowed_directories_for_previews()): - parentdir = os.path.abspath(parentdir) - if abspath.startswith(parentdir): - return abspath[len(parentdir):].replace('\\', '/') - return "" + def search_terms_from_path(self, filename): + return filename.replace('\\', '/') def is_empty(self, folder): for f in listdir(folder): @@ -228,21 +223,20 @@ class ExtraNetworksPage: return [] def create_html(self, item, tabname): - try: + # try: args = { "tabname": json.dumps(tabname), "name": item["name"].replace('_', ' '), "title": item["name"], "tags": '|'.join([item.get("tags")] if isinstance(item.get("tags", {}), str) else list(item.get("tags", {}).keys())), - "preview": html.escape(item.get("preview", None)), + "preview": html.escape(item.get("preview", self.link_preview('html/card-no-preview.png'))), "width": shared.opts.extra_networks_card_size, "height": shared.opts.extra_networks_card_size if shared.opts.extra_networks_card_square else 'auto', "fit": shared.opts.extra_networks_card_fit, - "loading": "lazy" if shared.opts.extra_networks_card_lazy else "eager", "prompt": item.get("prompt", None), "search_term": item.get("search_term", ""), "description": item.get("description") or "", - "local_preview": item["local_preview"], + "local_preview": item.get("local_preview"), "card_click": item.get("onclick", '"' + html.escape(f'return cardClicked({item.get("prompt", None)}, {"true" if self.allow_negative_prompt else "false"})') + '"'), "card_save_preview": '"' + html.escape('return saveCardPreview(event)') + '"', "card_save_desc": '"' + html.escape('return saveCardDescription(event)') + '"', @@ -260,9 +254,9 @@ class ExtraNetworksPage: if alias is not None: args['title'] += f'\nAlias: {alias}' return self.card.format(**args) - except Exception as e: - shared.log.error(f'Extra networks item error: page={tabname} item={item["name"]} {e}') - return "" + # except Exception as e: + # shared.log.error(f'Extra networks item error: page={tabname} item={item["name"]} {e}') + # return "" def find_preview(self, path): preview_extensions = ["jpg", "jpeg", "png", "webp", "tiff", "jp2"] diff --git a/modules/ui_extra_networks_checkpoints.py b/modules/ui_extra_networks_checkpoints.py index 39ce706be..4a2eb33e0 100644 --- a/modules/ui_extra_networks_checkpoints.py +++ b/modules/ui_extra_networks_checkpoints.py @@ -15,21 +15,21 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage): def list_items(self): checkpoint: sd_models.CheckpointInfo for name, checkpoint in sd_models.checkpoints_list.items(): - path, _ext = os.path.splitext(checkpoint.filename) - yield { - "name": checkpoint.name_for_extra, + fn = os.path.splitext(checkpoint.filename)[0] + record = { + "name": checkpoint.name, "title": checkpoint.title, - "filename": path, - "fullname": checkpoint.filename, + "filename": checkpoint.filename, "hash": checkpoint.shorthash, - "preview": self.find_preview(path), - "description": self.find_description(path), - "info": self.find_info(path), - "search_term": f'{self.search_terms_from_path(checkpoint.filename)} {(checkpoint.sha256 or "")} /{checkpoint.type}/', - "onclick": '"' + html.escape(f"""return selectCheckpoint({json.dumps(name)})""") + '"', - "local_preview": f"{path}.{shared.opts.samples_format}", + "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), "metadata": checkpoint.metadata, + "onclick": '"' + html.escape(f"""return selectCheckpoint({json.dumps(name)})""") + '"', } + 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] diff --git a/modules/ui_extra_networks_hypernets.py b/modules/ui_extra_networks_hypernets.py index a61f5ffb6..51fafb56d 100644 --- a/modules/ui_extra_networks_hypernets.py +++ b/modules/ui_extra_networks_hypernets.py @@ -12,16 +12,17 @@ class ExtraNetworksPageHypernetworks(ui_extra_networks.ExtraNetworksPage): def list_items(self): for name, path in shared.hypernetworks.items(): - path, _ext = os.path.splitext(path) + fn = os.path.splitext(path)[0] + name = os.path.relpath(fn, shared.opts.hypernetwork_dir) yield { - "name": name, + "name": os.path.relpath(fn, shared.opts.hypernetwork_dir), "filename": path, - "preview": self.find_preview(path), - "description": self.find_description(path), - "info": self.find_info(path), - "search_term": self.search_terms_from_path(path), + "preview": self.find_preview(fn), + "description": self.find_description(fn), + "info": self.find_info(fn), + "search_term": self.search_terms_from_path(name), "prompt": json.dumps(f""), - "local_preview": f"{path}.preview.{shared.opts.samples_format}", + "local_preview": f"{fn}.{shared.opts.samples_format}", } def allowed_directories_for_previews(self): diff --git a/modules/ui_extra_networks_styles.py b/modules/ui_extra_networks_styles.py index db772efd1..ba56f8f9d 100644 --- a/modules/ui_extra_networks_styles.py +++ b/modules/ui_extra_networks_styles.py @@ -45,19 +45,20 @@ class ExtraNetworksPageStyles(ui_extra_networks.ExtraNetworksPage): """ 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}' + for k, style in shared.prompt_styles.styles.items(): + fn = os.path.splitext(style.filename)[0] + txt = f'Prompt: {style.prompt}' + if len(style.negative_prompt) > 0: + txt += f'\nNegative: {style.negative_prompt}' yield { - "name": v.name, - "search_term": f'{txt} /{v.filename}', - "filename": v.filename, + "name": style.name, + "title": k, + "filename": style.filename, + "search_term": f'{txt} {self.search_terms_from_path(style.name)}', "preview": self.find_preview(fn), - "description": txt, - "onclick": '"' + html.escape(f"""return selectStyle({json.dumps(k)})""") + '"', "local_preview": f"{fn}.{shared.opts.samples_format}", + "description": txt, + "onclick": '"' + html.escape(f"""return selectStyle({json.dumps(style.name)})""") + '"', } def allowed_directories_for_previews(self): diff --git a/modules/ui_extra_networks_textual_inversion.py b/modules/ui_extra_networks_textual_inversion.py index 78e900858..3be7309b3 100644 --- a/modules/ui_extra_networks_textual_inversion.py +++ b/modules/ui_extra_networks_textual_inversion.py @@ -26,7 +26,7 @@ class ExtraNetworksPageTextualInversion(ui_extra_networks.ExtraNetworksPage): for filename in os.listdir(folder): fn = os.path.join(folder, filename) if os.path.isfile(fn) and (fn.lower().endswith(".pt") or fn.lower().endswith(".safetensors")): - embedding = Embedding(0, os.path.basename(fn)) + embedding = Embedding(vec=0, name=os.path.basename(fn), filename=fn) embedding.filename = fn embeddings.append(embedding) elif os.path.isdir(fn) and not fn.startswith('.'): @@ -45,15 +45,16 @@ class ExtraNetworksPageTextualInversion(ui_extra_networks.ExtraNetworksPage): tags = {} if embedding.tag is not None: tags[embedding.tag]=1 + name = os.path.splitext(embedding.basename)[0] yield { - "name": os.path.splitext(embedding.name)[0], + "name": name, "filename": embedding.filename, "preview": self.find_preview(path), "description": self.find_description(path), "info": self.find_info(path), - "search_term": self.search_terms_from_path(embedding.filename), + "search_term": self.search_terms_from_path(name), "prompt": json.dumps(os.path.splitext(embedding.name)[0]), - "local_preview": f"{path}.preview.{shared.opts.samples_format}", + "local_preview": f"{path}.{shared.opts.samples_format}", "tags": tags, } diff --git a/modules/uni_pc/uni_pc.py b/modules/uni_pc/uni_pc.py index fa16a0b48..7a79239ac 100644 --- a/modules/uni_pc/uni_pc.py +++ b/modules/uni_pc/uni_pc.py @@ -667,7 +667,7 @@ class UniPC: elif self.variant == 'bh2': B_h = torch.expm1(hh) else: - raise NotImplementedError() + raise NotImplementedError for i in range(1, order + 1): R.append(torch.pow(rks, i - 1)) @@ -802,7 +802,7 @@ class UniPC: model_prev_list[-1] = model_x progress.update(task, advance=1, description=f"Progress {round(len(vec_t) * step / (time.time() - t), 2)}it/s") else: - raise NotImplementedError() + raise NotImplementedError if denoise_to_zero: x = self.denoise_to_zero_fn(x, torch.ones((x.shape[0],)).to(device) * t_0) return x diff --git a/wiki b/wiki index d43376f66..8a182523a 160000 --- a/wiki +++ b/wiki @@ -1 +1 @@ -Subproject commit d43376f66fe454d2911a3b284077910df2b16b23 +Subproject commit 8a182523ac30a69bd707fb47a7bc828b63942ed6