mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
Merge remote-tracking branch 'upstream/master' into style
This commit is contained in:
@@ -1,5 +1,12 @@
|
||||
# Change Log for SD.Next
|
||||
|
||||
## Update for 05/28/2023
|
||||
|
||||
- settings search option
|
||||
- fully common save/zip/delete (new) options in all tabs
|
||||
- system info live gpu memory and load graphs for nvidia gpus
|
||||
- minor style changes
|
||||
|
||||
## Update for 05/26/2023
|
||||
|
||||
Some quality-of-life improvements...
|
||||
|
||||
@@ -48,3 +48,4 @@ Tech that can be integrated as part of the core workflow...
|
||||
## Random
|
||||
|
||||
- Bunch of stuff: <https://pharmapsychotic.com/tools.html>
|
||||
- <https://github.com/AUTOMATIC1111/stable-diffusion-webui/compare/89f9faa...20ae71f>
|
||||
|
||||
Submodule extensions-builtin/sd-extension-system-info updated: 4915b98576...8046b15445
Submodule extensions-builtin/sd-webui-controlnet updated: cdba83b6e1...09cb9a32d1
Submodule extensions-builtin/stable-diffusion-webui-images-browser updated: c61fae964a...75af6d0c32
+6
-4
@@ -258,7 +258,7 @@ def check_torch():
|
||||
elif allow_rocm and (shutil.which('rocminfo') is not None or os.path.exists('/opt/rocm/bin/rocminfo') or os.path.exists('/dev/kfd')):
|
||||
log.info('AMD ROCm toolkit detected')
|
||||
os.environ.setdefault('HSA_OVERRIDE_GFX_VERSION', '10.3.0')
|
||||
os.environ.setdefault('PYTORCH_HIP_ALLOC_CONF', 'garbage_collection_threshold:0.9,max_split_size_mb:512')
|
||||
os.environ.setdefault('PYTORCH_HIP_ALLOC_CONF', 'garbage_collection_threshold:0.8,max_split_size_mb:512')
|
||||
torch_command = os.environ.get('TORCH_COMMAND', 'torch==2.0.0 torchvision==0.15.1 --index-url https://download.pytorch.org/whl/rocm5.4.2')
|
||||
xformers_package = os.environ.get('XFORMERS_PACKAGE', 'none')
|
||||
elif allow_ipex and args.use_ipex and shutil.which('sycl-ls') is not None:
|
||||
@@ -289,9 +289,11 @@ def check_torch():
|
||||
log.info(f'Torch {torch.__version__}')
|
||||
if args.use_ipex and allow_ipex:
|
||||
import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import
|
||||
log.info(f'Torch backend: Intel OneAPI {torch.__version__}')
|
||||
log.info(f'Torch backend: Intel IPEX {ipex.__version__}')
|
||||
log.info(f'{os.popen("icpx --version").read().rstrip()}')
|
||||
log.info(f'Torch detected GPU: {torch.xpu.get_device_name("xpu")} VRAM {round(torch.xpu.get_device_properties("xpu").total_memory / 1024 / 1024)}')
|
||||
elif torch.cuda.is_available() and (allow_cuda or allow_rocm):
|
||||
# log.debug(f'Torch allocator: {torch.cuda.get_allocator_backend()}')
|
||||
if torch.version.cuda and allow_cuda:
|
||||
log.info(f'Torch backend: nVidia CUDA {torch.version.cuda} cuDNN {torch.backends.cudnn.version() if torch.backends.cudnn.is_available() else "N/A"}')
|
||||
elif torch.version.hip and allow_rocm:
|
||||
@@ -521,7 +523,7 @@ def set_environment():
|
||||
os.environ.setdefault('ACCELERATE', 'True')
|
||||
os.environ.setdefault('FORCE_CUDA', '1')
|
||||
os.environ.setdefault('ATTN_PRECISION', 'fp16')
|
||||
os.environ.setdefault('PYTORCH_CUDA_ALLOC_CONF', 'garbage_collection_threshold:0.9,max_split_size_mb:512')
|
||||
os.environ.setdefault('PYTORCH_CUDA_ALLOC_CONF', 'garbage_collection_threshold:0.8,max_split_size_mb:512')
|
||||
os.environ.setdefault('CUDA_LAUNCH_BLOCKING', '0')
|
||||
os.environ.setdefault('CUDA_CACHE_DISABLE', '0')
|
||||
os.environ.setdefault('CUDA_AUTO_BOOST', '1')
|
||||
@@ -703,7 +705,7 @@ def extensions_preload(force = False):
|
||||
from modules.paths_internal import extensions_builtin_dir, extensions_dir
|
||||
extension_folders = [extensions_builtin_dir] if args.safe else [extensions_builtin_dir, extensions_dir]
|
||||
for ext_dir in extension_folders:
|
||||
preload_extensions(ext_dir, parser, args.debug)
|
||||
preload_extensions(ext_dir, parser)
|
||||
except:
|
||||
log.error('Error running extension preloading')
|
||||
if args.profile:
|
||||
|
||||
@@ -81,12 +81,12 @@ svg.feather.feather-image, .feather .feather-image { display: none }
|
||||
#quicksettings > div, #quicksettings > fieldset { min-width: 24em; max-width: 26em; line-height: 2em; }
|
||||
#refresh_sd_model_checkpoint { height: 48px; margin-left: -14px; background: #333333; box-shadow: none; }
|
||||
#open_folder_extras, #footer, #style_pos_col, #style_neg_col, #roll_col, #extras_upscaler_2, #extras_upscaler_2_visibility, #txt2img_seed_resize_from_w, #txt2img_seed_resize_from_h { display: none; }
|
||||
#refresh_txt2img_styles, #refresh_img2img_styles { height: 40px; }
|
||||
#refresh_txt2img_styles, #refresh_img2img_styles { height: 45px; margin-left: -8px; }
|
||||
#save-animation { border-radius: 0 !important; margin-bottom: 16px; background-color: #111111; }
|
||||
#script_list { padding: 4px; margin-top: 20px; margin-bottom: 20px; }
|
||||
#settings > div.flex-wrap { width: 15em; }
|
||||
#tab_extensions table { background-color: #222222; }
|
||||
#txt2img_actions_column, #img2img_actions_column { min-width: 260px !important; max-width: 260px !important; }
|
||||
#txt2img_actions_column, #img2img_actions_column { min-width: 280px !important; max-width: 280px !important; gap: 0.6em }
|
||||
#txt2img_cfg_scale { min-width: 200px; }
|
||||
#txt2img_checkboxes, #img2img_checkboxes { background-color: transparent; }
|
||||
#txt2img_checkboxes, #img2img_checkboxes { margin-bottom: 0.2em; }
|
||||
@@ -102,7 +102,7 @@ svg.feather.feather-image, .feather .feather-image { display: none }
|
||||
#txt2img_subseed_row { padding: 0; margin-top: 16px; }
|
||||
#txt2img_subseed_show, #img2img_subseed_show { display: None }
|
||||
#txt2img_subseed_strength { margin-top: 0; }
|
||||
#txt2img_tools, #img2img_tools { margin-top: 54px; scale: 120%; margin-left: 26px; }
|
||||
#txt2img_tools, #img2img_tools { margin-top: 54px; scale: 120%; margin-left: 40px; filter: hue-rotate(180deg) saturate(0.5); }
|
||||
#txtimg_hr_finalres { max-width: 200px; }
|
||||
#pnginfo_html2_info { margin-top: -18px; background-color: var(--input-background-fill); padding: var(--input-padding) }
|
||||
#txt2img_extra_refresh, #txt2img_extra_close { height: 1.7em; }
|
||||
|
||||
@@ -92,6 +92,7 @@ button.custom-button{
|
||||
|
||||
.performance p{
|
||||
display: inline-block;
|
||||
color: var(--primary-100) !important
|
||||
}
|
||||
|
||||
.performance .time {
|
||||
@@ -115,6 +116,8 @@ a{
|
||||
|
||||
#txt2img_gallery img, #img2img_gallery img, #extras_gallery img{
|
||||
object-fit: scale-down;
|
||||
width: -webkit-fill-available !important;
|
||||
height: inherit !important;
|
||||
}
|
||||
#txt2img_actions_column, #img2img_actions_column {
|
||||
gap: 0.5em;
|
||||
@@ -725,3 +728,7 @@ div.controlnet_main_options {
|
||||
}
|
||||
|
||||
#modelmerger_interp_description { margin-top: 1em; margin-bottom: 1em; }
|
||||
|
||||
.thumbnail-item > img {
|
||||
|
||||
}
|
||||
@@ -260,6 +260,22 @@ onUiUpdate(() => {
|
||||
});
|
||||
};
|
||||
}
|
||||
const settings_search = gradioApp().querySelectorAll('#settings_search > label > textarea')[0];
|
||||
settings_search.oninput = (e) => {
|
||||
setTimeout(() => {
|
||||
gradioApp().querySelectorAll('#settings > div').forEach((elem) => {
|
||||
if (elem.id === 'settings_tab_licenses') return;
|
||||
elem.style.display = 'block';
|
||||
});
|
||||
gradioApp().querySelectorAll('#tab_settings .tabitem').forEach((section) => {
|
||||
section.querySelectorAll('.block').forEach((setting) => {
|
||||
const visible = setting.innerText.toLowerCase().includes(e.target.value.toLowerCase()) || setting.id.toLowerCase().includes(e.target.value.toLowerCase());
|
||||
if (setting.parentElement.classList.contains('form')) setting.parentElement.style.display = visible ? 'flex' : 'none';
|
||||
else setting.style.display = visible ? 'block' : 'none';
|
||||
});
|
||||
});
|
||||
}, 50);
|
||||
};
|
||||
});
|
||||
|
||||
onOptionsChanged(() => {
|
||||
|
||||
@@ -44,6 +44,8 @@ group.add_argument('--use-directml', default = False, action='store_true', help
|
||||
group.add_argument("--use-cuda", default=False, action='store_true', help="Force use nVidia CUDA backend, default: %(default)s")
|
||||
group.add_argument("--use-rocm", default=False, action='store_true', help="Force use AMD ROCm backend, default: %(default)s")
|
||||
group.add_argument('--subpath', type=str, help='Customize the URL subpath for usage with reverse proxy')
|
||||
group.add_argument('--backend', type=str, choices=[None, 'original', 'diffusers'], default=None, required=False, help='force backend type')
|
||||
|
||||
|
||||
# removed args are added here as hidden in fixed format for compatbility reasons
|
||||
group.add_argument("-f", action='store_true', help=argparse.SUPPRESS) # allows running as root; implemented outside of webui
|
||||
|
||||
@@ -671,7 +671,7 @@ class DeepDanbooruModel(nn.Module):
|
||||
t_771 = torch.sigmoid(t_770)
|
||||
return t_771
|
||||
|
||||
def load_state_dict(self, state_dict, **kwargs):
|
||||
def load_state_dict(self, state_dict, **kwargs): # pylint: disable=arguments-differ,unused-argument
|
||||
self.tags = state_dict.get('tags', [])
|
||||
|
||||
super(DeepDanbooruModel, self).load_state_dict({k: v for k, v in state_dict.items() if k != 'tags'})
|
||||
|
||||
+3
-1
@@ -81,6 +81,8 @@ def torch_gc(force=False):
|
||||
|
||||
|
||||
def test_fp16():
|
||||
if shared.cmd_opts.experimental:
|
||||
return True
|
||||
try:
|
||||
x = torch.tensor([[1.5,.0,.0,.0]]).to(device).half()
|
||||
layerNorm = torch.nn.LayerNorm(4, eps=0.00001, elementwise_affine=True, dtype=torch.float16, device=device)
|
||||
@@ -114,7 +116,7 @@ def set_cuda_params():
|
||||
pass
|
||||
global dtype, dtype_vae, dtype_unet, unet_needs_upcast # pylint: disable=global-statement
|
||||
ok = test_fp16()
|
||||
if shared.cmd_opts.use_directml: # TODO DirectML does not have full autocast capabilities
|
||||
if shared.cmd_opts.use_directml and not shared.cmd_opts.experimental: # TODO DirectML does not have full autocast capabilities
|
||||
shared.opts.no_half = True
|
||||
shared.opts.no_half_vae = True
|
||||
if ok and shared.opts.cuda_dtype == 'FP32':
|
||||
|
||||
@@ -17,7 +17,7 @@ def mod2normal(state_dict):
|
||||
if 'conv_first.weight' in state_dict:
|
||||
crt_net = {}
|
||||
items = []
|
||||
for k, v in state_dict.items():
|
||||
for k, _v in state_dict.items():
|
||||
items.append(k)
|
||||
|
||||
crt_net['model.0.weight'] = state_dict['conv_first.weight']
|
||||
@@ -53,7 +53,7 @@ def resrgan2normal(state_dict, nb=23):
|
||||
re8x = 0
|
||||
crt_net = {}
|
||||
items = []
|
||||
for k, v in state_dict.items():
|
||||
for k, _v in state_dict.items():
|
||||
items.append(k)
|
||||
|
||||
crt_net['model.0.weight'] = state_dict['conv_first.weight']
|
||||
@@ -186,7 +186,7 @@ class UpscalerESRGAN(Upscaler):
|
||||
elif "conv_first.weight" in state_dict:
|
||||
state_dict = mod2normal(state_dict)
|
||||
elif "model.0.weight" not in state_dict:
|
||||
raise Exception("The file is not a recognized ESRGAN model.")
|
||||
raise TypeError("The file is not a recognized ESRGAN model.")
|
||||
|
||||
in_nc, out_nc, nf, nb, plus, mscale = infer_params(state_dict)
|
||||
|
||||
|
||||
@@ -24,9 +24,6 @@ def run_pnginfo(image):
|
||||
for key, text in items.items():
|
||||
if key != 'UserComment':
|
||||
info += f"<div><b>{html.escape(str(key))}</b>: {html.escape(str(text))}</div>"
|
||||
if len(info) == 0:
|
||||
message = "Nothing found in the image."
|
||||
info = f"<div><p>{message}<p></div>"
|
||||
return '', geninfo, info
|
||||
|
||||
|
||||
|
||||
@@ -350,7 +350,9 @@ def create_override_settings_dict(text_pairs):
|
||||
|
||||
|
||||
def connect_paste(button, local_paste_fields, input_comp, override_settings_component, tabname):
|
||||
|
||||
def paste_func(prompt):
|
||||
shared.log.debug(f'paste prompt: {prompt}')
|
||||
if prompt is not None and 'Negative prompt' not in prompt and 'Steps' not in prompt:
|
||||
prompt = None
|
||||
if not prompt and not shared.cmd_opts.hide_ui_dir_config:
|
||||
|
||||
@@ -13,9 +13,8 @@ loaded_gfpgan_model = None
|
||||
|
||||
def gfpgann():
|
||||
import facexlib
|
||||
import gfpgan
|
||||
global loaded_gfpgan_model
|
||||
global model_path
|
||||
import gfpgan # pylint: disable=unused-import
|
||||
global loaded_gfpgan_model # pylint: disable=global-statement
|
||||
if loaded_gfpgan_model is not None:
|
||||
loaded_gfpgan_model.gfpgan.to(devices.device_gfpgan)
|
||||
return loaded_gfpgan_model
|
||||
@@ -54,7 +53,7 @@ def gfpgan_fix_faces(np_image):
|
||||
send_model_to(model, devices.device_gfpgan)
|
||||
|
||||
np_image_bgr = np_image[:, :, ::-1]
|
||||
cropped_faces, restored_faces, gfpgan_output_bgr = model.enhance(np_image_bgr, has_aligned=False, only_center_face=False, paste_back=True)
|
||||
_cropped_faces, _restored_faces, gfpgan_output_bgr = model.enhance(np_image_bgr, has_aligned=False, only_center_face=False, paste_back=True)
|
||||
np_image = gfpgan_output_bgr[:, :, ::-1]
|
||||
|
||||
model.face_helper.clean_all()
|
||||
@@ -69,7 +68,6 @@ gfpgan_constructor = None
|
||||
|
||||
|
||||
def setup_model(dirname):
|
||||
global model_path
|
||||
if not os.path.exists(model_path):
|
||||
os.makedirs(model_path)
|
||||
|
||||
@@ -77,9 +75,9 @@ def setup_model(dirname):
|
||||
import gfpgan
|
||||
import facexlib
|
||||
|
||||
global user_path
|
||||
global have_gfpgan
|
||||
global gfpgan_constructor
|
||||
global user_path # pylint: disable=global-statement
|
||||
global have_gfpgan # pylint: disable=global-statement
|
||||
global gfpgan_constructor # pylint: disable=global-statement
|
||||
|
||||
load_file_from_url_orig = gfpgan.utils.load_file_from_url
|
||||
facex_load_file_from_url_orig = facexlib.detection.load_file_from_url
|
||||
|
||||
+2
-2
@@ -5,7 +5,7 @@ import modules.scripts
|
||||
from modules import sd_samplers, shared
|
||||
from modules.generation_parameters_copypaste import create_override_settings_dict
|
||||
from modules.processing import Processed, StableDiffusionProcessingImg2Img, process_images
|
||||
from modules.ui import plaintext_to_html
|
||||
from modules.ui import plaintext_to_html, infotext_to_html
|
||||
import modules.processing as processing
|
||||
from modules.memstats import memory_stats
|
||||
|
||||
@@ -165,4 +165,4 @@ def img2img(id_task: str, mode: int, prompt: str, negative_prompt: str, prompt_s
|
||||
p.close()
|
||||
generation_info_js = processed.js()
|
||||
shared.log.debug(f'Processed: {len(processed.images)} Memory: {memory_stats()} img')
|
||||
return processed.images, generation_info_js, plaintext_to_html(processed.info), plaintext_to_html(processed.comments)
|
||||
return processed.images, generation_info_js, infotext_to_html(processed.info), plaintext_to_html(processed.comments)
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import modules.errors as errors
|
||||
|
||||
@@ -7,9 +6,8 @@ import modules.errors as errors
|
||||
localizations = {}
|
||||
|
||||
|
||||
def list_localizations(dirname):
|
||||
def list_localizations(dirname): # pylint: disable=unused-argument
|
||||
localizations.clear()
|
||||
return localizations
|
||||
"""
|
||||
for file in os.listdir(dirname):
|
||||
fn, ext = os.path.splitext(file)
|
||||
@@ -23,6 +21,8 @@ def list_localizations(dirname):
|
||||
fn, ext = os.path.splitext(file.filename)
|
||||
localizations[fn] = file.path
|
||||
"""
|
||||
return localizations
|
||||
|
||||
|
||||
def localization_js(current_localization_name):
|
||||
fn = localizations.get(current_localization_name, None)
|
||||
|
||||
+2
-2
@@ -6,7 +6,7 @@ cpu = torch.device("cpu")
|
||||
|
||||
|
||||
def send_everything_to_cpu():
|
||||
global module_in_gpu
|
||||
global module_in_gpu # pylint: disable=global-statement
|
||||
|
||||
if module_in_gpu is not None:
|
||||
module_in_gpu.to(cpu)
|
||||
@@ -22,7 +22,7 @@ def setup_for_low_vram(sd_model, use_medvram):
|
||||
we add this as forward_pre_hook to a lot of modules and this way all but one of them will
|
||||
be in CPU
|
||||
"""
|
||||
global module_in_gpu
|
||||
global module_in_gpu # pylint: disable=global-statement
|
||||
|
||||
module = parents.get(module, module)
|
||||
|
||||
|
||||
+1
-2
@@ -4,7 +4,7 @@ from PIL import Image, ImageFilter, ImageOps
|
||||
def get_crop_region(mask, pad=0):
|
||||
"""finds a rectangular region that contains all masked ares in an image. Returns (x1, y1, x2, y2) coordinates of the rectangle.
|
||||
For example, if a user has painted the top-right part of a 512x512 image", the result may be (256, 0, 512, 256)"""
|
||||
|
||||
|
||||
h, w = mask.shape
|
||||
|
||||
crop_left = 0
|
||||
@@ -96,4 +96,3 @@ def fill(image, mask):
|
||||
image_mod.alpha_composite(blurred)
|
||||
|
||||
return image_mod.convert("RGB")
|
||||
|
||||
|
||||
+3
-7
@@ -24,17 +24,13 @@ class MemUsageMonitor(threading.Thread):
|
||||
#torch.cuda.is_available() reports False when using IPEX.
|
||||
if shared.cmd_opts.use_ipex:
|
||||
self.cuda_mem_get_info()
|
||||
torch.cuda.memory_stats("xpu")
|
||||
torch.xpu.memory_stats("xpu")
|
||||
else:
|
||||
self.disabled = True
|
||||
else:
|
||||
try:
|
||||
if shared.cmd_opts.use_ipex:
|
||||
self.cuda_mem_get_info()
|
||||
torch.cuda.memory_stats("xpu")
|
||||
else:
|
||||
self.cuda_mem_get_info()
|
||||
torch.cuda.memory_stats(self.device)
|
||||
self.cuda_mem_get_info()
|
||||
torch.cuda.memory_stats(self.device)
|
||||
except Exception:
|
||||
self.disabled = True
|
||||
|
||||
|
||||
@@ -48,8 +48,8 @@ def run_postprocessing(extras_mode, image, image_folder: List[tempfile.NamedTemp
|
||||
outpath = output_dir
|
||||
else:
|
||||
outpath = opts.outdir_samples or opts.outdir_extras_samples
|
||||
infotext = ''
|
||||
for image, name, ext in zip(image_data, image_names, image_ext):
|
||||
infotext = ''
|
||||
if shared.state.interrupted:
|
||||
shared.log.debug('Postprocess interrupted')
|
||||
break
|
||||
@@ -62,10 +62,12 @@ def run_postprocessing(extras_mode, image, image_folder: List[tempfile.NamedTemp
|
||||
basename = os.path.splitext(os.path.basename(name))[0]
|
||||
else:
|
||||
basename = ''
|
||||
infotext = ", ".join([k if k == v else f'{k}: {generation_parameters_copypaste.quote(v)}' for k, v in pp.info.items() if v is not None])
|
||||
_geninfo, items = images.read_info_from_image(image)
|
||||
for k, v in items.items():
|
||||
pp.image.info[k] = v
|
||||
if 'parameters' in items:
|
||||
infotext = items['parameters'] + ', '
|
||||
infotext = infotext + ", ".join([k if k == v else f'{k}: {generation_parameters_copypaste.quote(v)}' for k, v in pp.info.items() if v is not None])
|
||||
pp.image.info["postprocessing"] = infotext
|
||||
if save_output:
|
||||
images.save_image(pp.image, path=outpath, basename=basename, seed=None, prompt=None, extension=ext or opts.samples_format, info=infotext, short_filename=True, no_prompt=True, grid=False, pnginfo_section_name="extras", existing_info=pp.image.info, forced_filename=None)
|
||||
@@ -73,7 +75,7 @@ def run_postprocessing(extras_mode, image, image_folder: List[tempfile.NamedTemp
|
||||
outputs.append(pp.image)
|
||||
|
||||
devices.torch_gc()
|
||||
return outputs, ui_common.plaintext_to_html(infotext), ''
|
||||
return outputs, ui_common.infotext_to_html(infotext), pp.image.info
|
||||
|
||||
|
||||
def run_extras(extras_mode, resize_mode, image, image_folder, input_dir, output_dir, show_extras_results, gfpgan_visibility, codeformer_visibility, codeformer_weight, upscaling_resize, upscaling_resize_w, upscaling_resize_h, upscaling_crop, extras_upscaler_1, extras_upscaler_2, extras_upscaler_2_visibility, upscale_first: bool, save_output: bool = True): #pylint: disable=unused-argument
|
||||
|
||||
@@ -19,7 +19,7 @@ from installer import git_commit
|
||||
import modules.sd_hijack
|
||||
from modules import devices, prompt_parser, masking, sd_samplers, lowvram, generation_parameters_copypaste, script_callbacks, extra_networks, sd_vae_approx, scripts # pylint: disable=unused-import
|
||||
from modules.sd_hijack import model_hijack
|
||||
from modules.shared import opts, cmd_opts, state, log
|
||||
from modules.shared import opts, cmd_opts, state, log, backend, Backend
|
||||
import modules.shared as shared
|
||||
import modules.paths as paths
|
||||
import modules.face_restoration
|
||||
@@ -220,7 +220,7 @@ class StableDiffusionProcessing:
|
||||
source_image = devices.cond_cast_float(source_image)
|
||||
# HACK: Using introspection as the Depth2Image model doesn't appear to uniquely
|
||||
# identify itself with a field common to all models. The conditioning_key is also hybrid.
|
||||
if opts.sd_backend == 'Diffusers': # TODO: img2img_image_conditioning
|
||||
if backend == Backend.DIFFUSERS: # TODO: Diffusers img2img_image_conditioning
|
||||
return latent_image.new_zeros(latent_image.shape[0], 5, 1, 1)
|
||||
if isinstance(self.sd_model, LatentDepth2ImageDiffusion):
|
||||
return self.depth2img_image_conditioning(source_image)
|
||||
@@ -522,7 +522,7 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed:
|
||||
assert p.prompt is not None
|
||||
seed = get_fixed_seed(p.seed)
|
||||
subseed = get_fixed_seed(p.subseed)
|
||||
if opts.sd_backend == 'Original':
|
||||
if backend == Backend.ORIGINAL:
|
||||
modules.sd_hijack.model_hijack.apply_circular(p.tiling)
|
||||
modules.sd_hijack.model_hijack.clear_comments()
|
||||
comments = {}
|
||||
@@ -573,11 +573,11 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed:
|
||||
cache[0] = (required_prompts, steps)
|
||||
return cache[1]
|
||||
|
||||
ema_scope_context = p.sd_model.ema_scope if opts.sd_backend == 'Original' else nullcontext
|
||||
ema_scope_context = p.sd_model.ema_scope if backend == Backend.ORIGINAL else nullcontext
|
||||
with torch.no_grad(), ema_scope_context():
|
||||
with devices.autocast():
|
||||
p.init(p.all_prompts, p.all_seeds, p.all_subseeds)
|
||||
if shared.opts.live_previews_enable and opts.show_progress_type == "Approx NN" and opts.sd_backend == 'Original':
|
||||
if shared.opts.live_previews_enable and opts.show_progress_type == "Approx NN" and backend == Backend.ORIGINAL:
|
||||
sd_vae_approx.model()
|
||||
if state.job_count == -1:
|
||||
state.job_count = p.n_iter
|
||||
@@ -618,7 +618,7 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed:
|
||||
if p.n_iter > 1:
|
||||
shared.state.job = f"Batch {n+1} out of {p.n_iter}"
|
||||
|
||||
if opts.sd_backend == 'Original':
|
||||
if backend == Backend.ORIGINAL:
|
||||
uc = get_conds_with_caching(prompt_parser.get_learned_conditioning, negative_prompts, p.steps * step_multiplier, cached_uc)
|
||||
c = get_conds_with_caching(prompt_parser.get_multicond_learned_conditioning, prompts, p.steps * step_multiplier, cached_c)
|
||||
if len(model_hijack.comments) > 0:
|
||||
@@ -649,7 +649,7 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed:
|
||||
devices.torch_gc()
|
||||
if p.scripts is not None:
|
||||
p.scripts.postprocess_batch(p, x_samples_ddim, batch_number=n)
|
||||
else: # TODO Diffusers
|
||||
else: # TODO Diffusers main processing
|
||||
generator = [torch.Generator(device="cpu").manual_seed(s) for s in seeds]
|
||||
if shared.sd_model.scheduler.name != p.sampler_name:
|
||||
sampler = sd_samplers.all_samplers_map.get(p.sampler_name, None)
|
||||
@@ -671,7 +671,7 @@ def process_images_inner(p: StableDiffusionProcessing) -> Processed:
|
||||
|
||||
for i, x_sample in enumerate(x_samples_ddim):
|
||||
p.batch_index = i
|
||||
if opts.sd_backend == 'Original':
|
||||
if backend == Backend.ORIGINAL:
|
||||
x_sample = 255. * np.moveaxis(x_sample.cpu().numpy(), 0, 2)
|
||||
x_sample = x_sample.astype(np.uint8)
|
||||
else:
|
||||
|
||||
@@ -13,7 +13,7 @@ from typing import List
|
||||
import lark
|
||||
import torch
|
||||
from compel import Compel
|
||||
from modules.shared import log, opts
|
||||
from modules.shared import opts
|
||||
|
||||
# a prompt like this: "fantasy landscape with a [mountain:lake:0.25] and [an oak:a christmas tree:0.75][ in foreground::0.6][ in background:0.25] [shoddy:masterful:0.5]"
|
||||
# will be represented with prompt_schedule like this (assuming steps=100):
|
||||
|
||||
@@ -6,7 +6,7 @@ import modules.errors as errors
|
||||
preloaded = []
|
||||
|
||||
|
||||
def load_module(path, detailed=False):
|
||||
def load_module(path):
|
||||
module_spec = importlib.util.spec_from_file_location(os.path.basename(path), path)
|
||||
module = importlib.util.module_from_spec(module_spec)
|
||||
try:
|
||||
@@ -17,7 +17,7 @@ def load_module(path, detailed=False):
|
||||
|
||||
|
||||
|
||||
def preload_extensions(extensions_dir, parser, detailed=False):
|
||||
def preload_extensions(extensions_dir, parser):
|
||||
if not os.path.isdir(extensions_dir):
|
||||
return
|
||||
for dirname in sorted(os.listdir(extensions_dir)):
|
||||
@@ -28,7 +28,7 @@ def preload_extensions(extensions_dir, parser, detailed=False):
|
||||
if not os.path.isfile(preload_script):
|
||||
continue
|
||||
try:
|
||||
module = load_module(preload_script, detailed)
|
||||
module = load_module(preload_script)
|
||||
if hasattr(module, 'preload'):
|
||||
module.preload(parser)
|
||||
except Exception as e:
|
||||
|
||||
@@ -31,23 +31,19 @@ class ScriptPostprocessing:
|
||||
The return value should be a dictionary that maps parameter names to components used in processing.
|
||||
Values of those components will be passed to process() function.
|
||||
"""
|
||||
|
||||
pass
|
||||
pass # pylint: disable=unnecessary-pass
|
||||
|
||||
def process(self, pp: PostprocessedImage, **args):
|
||||
"""
|
||||
This function is called to postprocess the image.
|
||||
args contains a dictionary with all values returned by components from ui()
|
||||
"""
|
||||
|
||||
pass
|
||||
pass # pylint: disable=unnecessary-pass
|
||||
|
||||
def image_changed(self):
|
||||
pass
|
||||
|
||||
|
||||
|
||||
|
||||
def wrap_call(func, filename, funcname, *args, default=None, **kwargs):
|
||||
try:
|
||||
res = func(*args, **kwargs)
|
||||
@@ -66,7 +62,7 @@ class ScriptPostprocessingRunner:
|
||||
def initialize_scripts(self, scripts_data):
|
||||
self.scripts = []
|
||||
|
||||
for script_class, path, basedir, script_module in scripts_data:
|
||||
for script_class, path, _basedir, _script_module in scripts_data:
|
||||
script: ScriptPostprocessing = script_class()
|
||||
script.filename = path
|
||||
|
||||
@@ -124,7 +120,7 @@ class ScriptPostprocessingRunner:
|
||||
script_args = args[script.args_from:script.args_to]
|
||||
|
||||
process_args = {}
|
||||
for (name, component), value in zip(script.controls.items(), script_args):
|
||||
for (name, _component), value in zip(script.controls.items(), script_args):
|
||||
process_args[name] = value
|
||||
|
||||
script.process(pp, **process_args)
|
||||
|
||||
@@ -35,10 +35,10 @@ class DisableInitialization:
|
||||
return original
|
||||
|
||||
def __enter__(self):
|
||||
def do_nothing(*args, **kwargs):
|
||||
def do_nothing(*args, **kwargs): # pylint: disable=unused-argument
|
||||
pass
|
||||
|
||||
def create_model_and_transforms_without_pretrained(*args, pretrained=None, **kwargs):
|
||||
def create_model_and_transforms_without_pretrained(*args, pretrained=None, **kwargs): # pylint: disable=unused-argument
|
||||
return self.create_model_and_transforms(*args, pretrained=None, **kwargs)
|
||||
|
||||
def CLIPTextModel_from_pretrained(pretrained_model_name_or_path, *model_args, **kwargs):
|
||||
@@ -61,16 +61,16 @@ class DisableInitialization:
|
||||
if res is None:
|
||||
res = original(url, *args, local_files_only=False, **kwargs)
|
||||
return res
|
||||
except Exception as e:
|
||||
except Exception:
|
||||
return original(url, *args, local_files_only=False, **kwargs)
|
||||
|
||||
def transformers_utils_hub_get_from_cache(url, *args, local_files_only=False, **kwargs):
|
||||
def transformers_utils_hub_get_from_cache(url, *args, local_files_only=False, **kwargs): # pylint: disable=unused-argument
|
||||
return transformers_utils_hub_get_file_from_cache(self.transformers_utils_hub_get_from_cache, url, *args, **kwargs)
|
||||
|
||||
def transformers_tokenization_utils_base_cached_file(url, *args, local_files_only=False, **kwargs):
|
||||
def transformers_tokenization_utils_base_cached_file(url, *args, local_files_only=False, **kwargs): # pylint: disable=unused-argument
|
||||
return transformers_utils_hub_get_file_from_cache(self.transformers_tokenization_utils_base_cached_file, url, *args, **kwargs)
|
||||
|
||||
def transformers_configuration_utils_cached_file(url, *args, local_files_only=False, **kwargs):
|
||||
def transformers_configuration_utils_cached_file(url, *args, local_files_only=False, **kwargs): # pylint: disable=unused-argument
|
||||
return transformers_utils_hub_get_file_from_cache(self.transformers_configuration_utils_cached_file, url, *args, **kwargs)
|
||||
|
||||
self.replace(torch.nn.init, 'kaiming_uniform_', do_nothing)
|
||||
@@ -78,16 +78,15 @@ class DisableInitialization:
|
||||
self.replace(torch.nn.init, '_no_grad_uniform_', do_nothing)
|
||||
|
||||
if self.disable_clip:
|
||||
self.create_model_and_transforms = self.replace(open_clip, 'create_model_and_transforms', create_model_and_transforms_without_pretrained)
|
||||
self.CLIPTextModel_from_pretrained = self.replace(ldm.modules.encoders.modules.CLIPTextModel, 'from_pretrained', CLIPTextModel_from_pretrained)
|
||||
self.transformers_modeling_utils_load_pretrained_model = self.replace(transformers.modeling_utils.PreTrainedModel, '_load_pretrained_model', transformers_modeling_utils_load_pretrained_model)
|
||||
self.transformers_tokenization_utils_base_cached_file = self.replace(transformers.tokenization_utils_base, 'cached_file', transformers_tokenization_utils_base_cached_file)
|
||||
self.transformers_configuration_utils_cached_file = self.replace(transformers.configuration_utils, 'cached_file', transformers_configuration_utils_cached_file)
|
||||
self.transformers_utils_hub_get_from_cache = self.replace(transformers.utils.hub, 'get_from_cache', transformers_utils_hub_get_from_cache)
|
||||
self.create_model_and_transforms = self.replace(open_clip, 'create_model_and_transforms', create_model_and_transforms_without_pretrained) # pylint: disable=attribute-defined-outside-init
|
||||
self.CLIPTextModel_from_pretrained = self.replace(ldm.modules.encoders.modules.CLIPTextModel, 'from_pretrained', CLIPTextModel_from_pretrained) # pylint: disable=attribute-defined-outside-init
|
||||
self.transformers_modeling_utils_load_pretrained_model = self.replace(transformers.modeling_utils.PreTrainedModel, '_load_pretrained_model', transformers_modeling_utils_load_pretrained_model) # pylint: disable=attribute-defined-outside-init
|
||||
self.transformers_tokenization_utils_base_cached_file = self.replace(transformers.tokenization_utils_base, 'cached_file', transformers_tokenization_utils_base_cached_file) # pylint: disable=attribute-defined-outside-init
|
||||
self.transformers_configuration_utils_cached_file = self.replace(transformers.configuration_utils, 'cached_file', transformers_configuration_utils_cached_file) # pylint: disable=attribute-defined-outside-init
|
||||
self.transformers_utils_hub_get_from_cache = self.replace(transformers.utils.hub, 'get_from_cache', transformers_utils_hub_get_from_cache) # pylint: disable=attribute-defined-outside-init
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
for obj, field, original in self.replaced:
|
||||
setattr(obj, field, original)
|
||||
|
||||
self.replaced.clear()
|
||||
|
||||
|
||||
@@ -5,15 +5,15 @@ import ldm.modules.diffusionmodules.openaimodel
|
||||
|
||||
|
||||
def BasicTransformerBlock_forward(self, x, context=None):
|
||||
return checkpoint(self._forward, x, context)
|
||||
return checkpoint(self._forward, x, context) # pylint: disable=protected-access
|
||||
|
||||
|
||||
def AttentionBlock_forward(self, x):
|
||||
return checkpoint(self._forward, x)
|
||||
return checkpoint(self._forward, x) # pylint: disable=protected-access
|
||||
|
||||
|
||||
def ResBlock_forward(self, x, emb):
|
||||
return checkpoint(self._forward, x, emb)
|
||||
return checkpoint(self._forward, x, emb) # pylint: disable=protected-access
|
||||
|
||||
|
||||
stored = []
|
||||
@@ -43,4 +43,3 @@ def remove():
|
||||
ldm.modules.diffusionmodules.openaimodel.AttentionBlock.forward = stored[2]
|
||||
|
||||
stored.clear()
|
||||
|
||||
|
||||
@@ -70,7 +70,7 @@ def process_text_old(self: sd_hijack_clip.FrozenCLIPEmbedderWithCustomWordsBase,
|
||||
|
||||
|
||||
def forward_old(self: sd_hijack_clip.FrozenCLIPEmbedderWithCustomWordsBase, texts):
|
||||
batch_multipliers, remade_batch_tokens, used_custom_terms, hijack_comments, hijack_fixes, token_count = process_text_old(self, texts)
|
||||
batch_multipliers, remade_batch_tokens, used_custom_terms, hijack_comments, hijack_fixes, _token_count = process_text_old(self, texts)
|
||||
|
||||
self.hijack.comments += hijack_comments
|
||||
|
||||
|
||||
@@ -4,12 +4,11 @@ import ldm.models.diffusion.ddpm
|
||||
import ldm.models.diffusion.ddim
|
||||
import ldm.models.diffusion.plms
|
||||
|
||||
from ldm.models.diffusion.ddpm import LatentDiffusion
|
||||
from ldm.models.diffusion.plms import PLMSSampler
|
||||
from ldm.models.diffusion.ddim import DDIMSampler, noise_like
|
||||
from ldm.models.diffusion.ddpm import LatentDiffusion # pylint: disable=unused-import
|
||||
from ldm.models.diffusion.plms import PLMSSampler # pylint: disable=unused-import
|
||||
from ldm.models.diffusion.ddim import DDIMSampler, noise_like # pylint: disable=unused-import
|
||||
from ldm.models.diffusion.sampling_util import norm_thresholding
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def p_sample_plms(self, x, c, t, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False,
|
||||
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
|
||||
@@ -63,7 +62,6 @@ def p_sample_plms(self, x, c, t, index, repeat_noise=False, use_original_steps=F
|
||||
if quantize_denoised:
|
||||
pred_x0, _, *_ = self.model.first_stage_model.quantize(pred_x0)
|
||||
if dynamic_threshold is not None:
|
||||
from ldm.models.diffusion.sampling_util import norm_thresholding
|
||||
pred_x0 = norm_thresholding(pred_x0, dynamic_threshold)
|
||||
# direction pointing to x_t
|
||||
dir_xt = (1. - a_prev - sigma_t**2).sqrt() * e_t
|
||||
|
||||
@@ -51,7 +51,7 @@ def get_available_vram():
|
||||
|
||||
|
||||
# see https://github.com/basujindal/stable-diffusion/pull/117 for discussion
|
||||
def split_cross_attention_forward_v1(self, x, context=None, mask=None):
|
||||
def split_cross_attention_forward_v1(self, x, context=None, mask=None): # pylint: disable=unused-argument
|
||||
h = self.heads
|
||||
|
||||
q_in = self.to_q(x)
|
||||
@@ -90,7 +90,7 @@ def split_cross_attention_forward_v1(self, x, context=None, mask=None):
|
||||
|
||||
|
||||
# taken from https://github.com/Doggettx/stable-diffusion and modified
|
||||
def split_cross_attention_forward(self, x, context=None, mask=None):
|
||||
def split_cross_attention_forward(self, x, context=None, mask=None): # pylint: disable=unused-argument
|
||||
h = self.heads
|
||||
q_in = self.to_q(x)
|
||||
context = default(context, x)
|
||||
@@ -231,7 +231,7 @@ def einsum_op(q, k, v):
|
||||
# Tested on i7 with 8MB L3 cache.
|
||||
return einsum_op_tensor_mem(q, k, v, 32)
|
||||
|
||||
def split_cross_attention_forward_invokeAI(self, x, context=None, mask=None):
|
||||
def split_cross_attention_forward_invokeAI(self, x, context=None, mask=None): # pylint: disable=unused-argument
|
||||
h = self.heads
|
||||
|
||||
q = self.to_q(x)
|
||||
@@ -315,7 +315,7 @@ def sub_quad_attention(q, k, v, q_chunk_size=1024, kv_chunk_size=None, kv_chunk_
|
||||
if chunk_threshold_bytes is not None and qk_matmul_size_bytes <= chunk_threshold_bytes:
|
||||
# the big matmul fits into our memory limit; do everything in 1 chunk,
|
||||
# i.e. send it down the unchunked fast-path
|
||||
query_chunk_size = q_tokens
|
||||
query_chunk_size = q_tokens # pylint: disable=unused-variable
|
||||
kv_chunk_size = k_tokens
|
||||
|
||||
with devices.without_autocast(disable=q.dtype == v.dtype):
|
||||
@@ -336,7 +336,7 @@ def get_xformers_flash_attention_op(q, k, v):
|
||||
|
||||
try:
|
||||
flash_attention_op = xformers.ops.MemoryEfficientAttentionFlashAttentionOp
|
||||
fw, bw = flash_attention_op
|
||||
fw, _bw = flash_attention_op
|
||||
if fw.supports(xformers.ops.fmha.Inputs(query=q, key=k, value=v, attn_bias=None)):
|
||||
return flash_attention_op
|
||||
except Exception as e:
|
||||
@@ -345,7 +345,7 @@ def get_xformers_flash_attention_op(q, k, v):
|
||||
return None
|
||||
|
||||
|
||||
def xformers_attention_forward(self, x, context=None, mask=None):
|
||||
def xformers_attention_forward(self, x, context=None, mask=None): # pylint: disable=unused-argument
|
||||
h = self.heads
|
||||
q_in = self.to_q(x)
|
||||
context = default(context, x)
|
||||
@@ -481,7 +481,7 @@ def xformers_attnblock_forward(self, x):
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
b, c, h, w = q.shape
|
||||
b, c, h, w = q.shape # pylint: disable=unused-variable
|
||||
q, k, v = map(lambda t: rearrange(t, 'b c h w -> b (h w) c'), (q, k, v))
|
||||
dtype = q.dtype
|
||||
if shared.opts.upcast_attn:
|
||||
@@ -503,7 +503,7 @@ def sdp_attnblock_forward(self, x):
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
b, c, h, w = q.shape
|
||||
b, c, h, w = q.shape # pylint: disable=unused-variable
|
||||
q, k, v = map(lambda t: rearrange(t, 'b c h w -> b (h w) c'), (q, k, v))
|
||||
dtype = q.dtype
|
||||
if shared.opts.upcast_attn:
|
||||
@@ -531,7 +531,7 @@ def sub_quad_attnblock_forward(self, x):
|
||||
q = self.q(h_)
|
||||
k = self.k(h_)
|
||||
v = self.v(h_)
|
||||
b, c, h, w = q.shape
|
||||
b, c, h, w = q.shape # pylint: disable=unused-variable
|
||||
q, k, v = map(lambda t: rearrange(t, 'b c h w -> b (h w) c'), (q, k, v))
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
|
||||
@@ -47,25 +47,25 @@ def apply_model(orig_func, self, x_noisy, t, cond, **kwargs):
|
||||
|
||||
|
||||
class GELUHijack(torch.nn.GELU, torch.nn.Module):
|
||||
def __init__(self, *args, **kwargs):
|
||||
def __init__(self, *args, **kwargs): # pylint: disable=super-init-not-called
|
||||
torch.nn.GELU.__init__(self, *args, **kwargs)
|
||||
def forward(self, x):
|
||||
def forward(self, input): # pylint: disable=redefined-builtin
|
||||
if devices.unet_needs_upcast:
|
||||
return torch.nn.GELU.forward(self.float(), x.float()).to(devices.dtype_unet)
|
||||
return torch.nn.GELU.forward(self.float(), input.float()).to(devices.dtype_unet)
|
||||
else:
|
||||
return torch.nn.GELU.forward(self, x)
|
||||
return torch.nn.GELU.forward(self, input)
|
||||
|
||||
|
||||
ddpm_edit_hijack = None
|
||||
def hijack_ddpm_edit():
|
||||
global ddpm_edit_hijack
|
||||
global ddpm_edit_hijack # pylint: disable=global-statement
|
||||
if not ddpm_edit_hijack:
|
||||
CondFunc('modules.models.diffusion.ddpm_edit.LatentDiffusion.decode_first_stage', first_stage_sub, first_stage_cond)
|
||||
CondFunc('modules.models.diffusion.ddpm_edit.LatentDiffusion.encode_first_stage', first_stage_sub, first_stage_cond)
|
||||
ddpm_edit_hijack = CondFunc('modules.models.diffusion.ddpm_edit.LatentDiffusion.apply_model', apply_model, unet_needs_upcast)
|
||||
|
||||
|
||||
unet_needs_upcast = lambda *args, **kwargs: devices.unet_needs_upcast
|
||||
unet_needs_upcast = lambda *args, **kwargs: devices.unet_needs_upcast # pylint: disable=unnecessary-lambda-assignment
|
||||
CondFunc('ldm.models.diffusion.ddpm.LatentDiffusion.apply_model', apply_model, unet_needs_upcast)
|
||||
CondFunc('ldm.modules.diffusionmodules.openaimodel.timestep_embedding', lambda orig_func, timesteps, *args, **kwargs: orig_func(timesteps, *args, **kwargs).to(torch.float32 if timesteps.dtype == torch.int64 else devices.dtype_unet), unet_needs_upcast)
|
||||
if version.parse(torch.__version__) <= version.parse("1.13.2") or torch.cuda.is_available() or shared.cmd_opts.use_ipex:
|
||||
@@ -73,8 +73,8 @@ if version.parse(torch.__version__) <= version.parse("1.13.2") or torch.cuda.is_
|
||||
CondFunc('ldm.modules.attention.GEGLU.forward', lambda orig_func, self, x: orig_func(self.float(), x.float()).to(devices.dtype_unet), unet_needs_upcast)
|
||||
CondFunc('open_clip.transformer.ResidualAttentionBlock.__init__', lambda orig_func, *args, **kwargs: kwargs.update({'act_layer': GELUHijack}) and False or orig_func(*args, **kwargs), lambda _, *args, **kwargs: kwargs.get('act_layer') is None or kwargs['act_layer'] == torch.nn.GELU)
|
||||
|
||||
first_stage_cond = lambda _, self, *args, **kwargs: devices.unet_needs_upcast and self.model.diffusion_model.dtype == torch.float16
|
||||
first_stage_sub = lambda orig_func, self, x, **kwargs: orig_func(self, x.to(devices.dtype_vae), **kwargs)
|
||||
first_stage_cond = lambda _, self, *args, **kwargs: devices.unet_needs_upcast and self.model.diffusion_model.dtype == torch.float16 # pylint: disable=unnecessary-lambda-assignment
|
||||
first_stage_sub = lambda orig_func, self, x, **kwargs: orig_func(self, x.to(devices.dtype_vae), **kwargs) # pylint: disable=unnecessary-lambda-assignment
|
||||
CondFunc('ldm.models.diffusion.ddpm.LatentDiffusion.decode_first_stage', first_stage_sub, first_stage_cond)
|
||||
CondFunc('ldm.models.diffusion.ddpm.LatentDiffusion.encode_first_stage', first_stage_sub, first_stage_cond)
|
||||
CondFunc('ldm.models.diffusion.ddpm.LatentDiffusion.get_first_stage_encoding', lambda orig_func, *args, **kwargs: orig_func(*args, **kwargs).float(), first_stage_cond)
|
||||
|
||||
+16
-11
@@ -28,14 +28,14 @@ checkpoints_loaded = collections.OrderedDict()
|
||||
skip_next_load = False
|
||||
|
||||
|
||||
class CheckpointInfo: # TODO Diffusers
|
||||
class CheckpointInfo:
|
||||
def __init__(self, filename):
|
||||
name = ''
|
||||
self.name = None
|
||||
self.hash = None
|
||||
self.filename = filename
|
||||
abspath = os.path.abspath(filename)
|
||||
if shared.opts.sd_backend == 'Original':
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
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):
|
||||
@@ -48,7 +48,6 @@ class CheckpointInfo: # TODO Diffusers
|
||||
self.hash = model_hash(self.filename)
|
||||
self.sha256 = hashes.sha256_from_cache(self.filename, f"checkpoint/{name}")
|
||||
else: # TODO Diffusers
|
||||
# sd_model.unet.config._name_or_path.split("/")[-2]
|
||||
repo = [r for r in modelloader.diffuser_repos if filename == r['filename']]
|
||||
if len(repo) == 0:
|
||||
shared.log.error(f'Cannot find diffuser model: {filename}')
|
||||
@@ -105,7 +104,7 @@ def checkpoint_tiles():
|
||||
def list_models():
|
||||
checkpoints_list.clear()
|
||||
checkpoint_aliases.clear()
|
||||
if shared.opts.sd_backend == 'Original':
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
model_list = modelloader.load_models(model_path=os.path.join(models_path, 'Stable-diffusion'), model_url=None, command_path=shared.opts.ckpt_dir, ext_filter=[".ckpt", ".safetensors"], download_name=None, ext_blacklist=[".vae.ckpt", ".vae.safetensors"])
|
||||
else:
|
||||
model_list = modelloader.load_diffusers(model_path=os.path.join(models_path, 'Diffusers'), command_path=shared.opts.diffusers_dir)
|
||||
@@ -114,7 +113,7 @@ def list_models():
|
||||
if checkpoint_info.name is not None:
|
||||
checkpoint_info.register()
|
||||
if shared.cmd_opts.ckpt is not None:
|
||||
if not os.path.exists(shared.cmd_opts.ckpt) and shared.opts.sd_backend == 'Original':
|
||||
if not os.path.exists(shared.cmd_opts.ckpt) and shared.backend == shared.Backend.ORIGINAL:
|
||||
if shared.cmd_opts.ckpt.lower() != "none":
|
||||
shared.log.warning(f"Requested checkpoint not found: {shared.cmd_opts.ckpt}")
|
||||
else:
|
||||
@@ -228,7 +227,7 @@ def read_metadata_from_safetensors(filename):
|
||||
|
||||
|
||||
def read_state_dict(checkpoint_file, map_location=None): # pylint: disable=unused-argument
|
||||
if shared.opts.sd_backend == 'Diffusers':
|
||||
if shared.backend == shared.Backend.DIFFUSERS:
|
||||
return None
|
||||
try:
|
||||
pl_sd = None
|
||||
@@ -377,9 +376,9 @@ class SdModelData:
|
||||
if self.sd_model is None:
|
||||
with self.lock:
|
||||
try:
|
||||
if shared.opts.sd_backend == 'Original':
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
load_model()
|
||||
elif shared.opts.sd_backend == 'Diffusers':
|
||||
elif shared.backend == shared.Backend.DIFFUSERS:
|
||||
load_diffuser()
|
||||
else:
|
||||
shared.log.error(f"Unknown Stable Diffusion backend: {shared.opts.sd_backend}")
|
||||
@@ -430,6 +429,12 @@ def load_diffuser(checkpoint_info=None, already_loaded_state_dict=None, timer=No
|
||||
shared.log.info(f'Loading diffuser model: {checkpoint_info.filename}')
|
||||
scheduler = diffusers.UniPCMultistepScheduler.from_pretrained(checkpoint_info.filename, subfolder="scheduler")
|
||||
sd_model = diffusers.DiffusionPipeline.from_pretrained(checkpoint_info.filename, scheduler=scheduler, **diffusor_config)
|
||||
if shared.cmd_opts.medvram:
|
||||
sd_model.enable_model_cpu_offload()
|
||||
if shared.cmd_opts.lowvram:
|
||||
sd_model.enable_sequential_cpu_offload()
|
||||
if shared.opts.cross_attention_optimization == "xFormers":
|
||||
sd_model.enable_xformers_memory_efficient_attention()
|
||||
sd_model.sd_checkpoint_info = checkpoint_info
|
||||
sd_model.sd_model_checkpoint = checkpoint_info.filename
|
||||
sd_model.sd_model_hash = checkpoint_info.hash
|
||||
@@ -540,7 +545,7 @@ def reload_model_weights(sd_model=None, info=None):
|
||||
sd_model.to(devices.cpu)
|
||||
if shared.opts.model_reuse_dict and sd_model is not None:
|
||||
shared.log.info('Reusing previous model dictionary')
|
||||
sd_hijack.model_hijack.undo_hijack(sd_model) # TODO double undo hijack
|
||||
sd_hijack.model_hijack.undo_hijack(sd_model)
|
||||
else:
|
||||
unload_model_weights()
|
||||
sd_model = None
|
||||
@@ -551,7 +556,7 @@ def reload_model_weights(sd_model=None, info=None):
|
||||
if sd_model is None or checkpoint_config != sd_model.used_config:
|
||||
del sd_model
|
||||
checkpoints_loaded.clear()
|
||||
if shared.opts.sd_backend == 'Original':
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
load_model(checkpoint_info, already_loaded_state_dict=state_dict, timer=timer)
|
||||
else:
|
||||
load_diffuser(checkpoint_info, already_loaded_state_dict=state_dict, timer=timer)
|
||||
@@ -576,7 +581,7 @@ def unload_model_weights(sd_model=None, _info=None):
|
||||
from modules import sd_hijack
|
||||
if model_data.sd_model:
|
||||
model_data.sd_model.to(devices.cpu)
|
||||
if shared.opts.sd_backend == 'Original':
|
||||
if shared.backend == shared.Backend.ORIGINAL:
|
||||
sd_hijack.model_hijack.undo_hijack(model_data.sd_model)
|
||||
model_data.sd_model = None
|
||||
sd_model = None
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
from modules import sd_samplers_compvis, sd_samplers_kdiffusion, sd_samplers_diffusors, shared
|
||||
from modules.sd_samplers_common import samples_to_image_grid, sample_to_image # pylint: disable=unused-import
|
||||
from modules.shared import opts
|
||||
from modules.shared import backend, Backend
|
||||
|
||||
if opts.sd_backend == 'Original':
|
||||
if backend == Backend.ORIGINAL:
|
||||
all_samplers = [
|
||||
*sd_samplers_kdiffusion.samplers_data_k_diffusion,
|
||||
*sd_samplers_compvis.samplers_data_compvis,
|
||||
@@ -23,7 +23,7 @@ def create_sampler(name, model):
|
||||
else:
|
||||
config = all_samplers[0]
|
||||
assert config is not None, f'bad sampler name: {name}'
|
||||
if opts.sd_backend == 'Original':
|
||||
if backend == Backend.ORIGINAL:
|
||||
sampler = config.constructor(model)
|
||||
sampler.config = config
|
||||
return sampler
|
||||
|
||||
@@ -40,7 +40,7 @@ class VanillaStableDiffusionSampler:
|
||||
|
||||
self.conditioning_key = sd_model.model.conditioning_key
|
||||
|
||||
def number_of_needed_noises(self, p):
|
||||
def number_of_needed_noises(self, p): # pylint: disable=unused-argument
|
||||
return 0
|
||||
|
||||
def launch_sampling(self, steps, func):
|
||||
@@ -97,7 +97,10 @@ class VanillaStableDiffusionSampler:
|
||||
unconditional_conditioning = unconditional_conditioning[:, :cond.shape[1]]
|
||||
|
||||
if self.mask is not None:
|
||||
img_orig = self.sampler.model.q_sample(self.init_latent, ts)
|
||||
if shared.cmd_opts.use_ipex:
|
||||
img_orig = self.sampler.model.q_sample(self.init_latent, ts.type(torch.int64))
|
||||
else:
|
||||
img_orig = self.sampler.model.q_sample(self.init_latent, ts)
|
||||
x = img_orig * self.mask + self.nmask * x
|
||||
|
||||
# Wrap the image conditioning back up since the DDIM code can accept the dict directly.
|
||||
@@ -125,7 +128,7 @@ class VanillaStableDiffusionSampler:
|
||||
self.update_step(res[1])
|
||||
return x, ts, cond, uncond, res
|
||||
|
||||
def unipc_after_update(self, x, model_x):
|
||||
def unipc_after_update(self, x, model_x): # pylint: disable=unused-argument
|
||||
self.update_step(x)
|
||||
|
||||
def initialize(self, p):
|
||||
|
||||
@@ -97,10 +97,10 @@ class CFGDenoiser(torch.nn.Module):
|
||||
|
||||
if shared.sd_model.model.conditioning_key == "crossattn-adm":
|
||||
image_uncond = torch.zeros_like(image_cond)
|
||||
make_condition_dict = lambda c_crossattn, c_adm: {"c_crossattn": c_crossattn, "c_adm": c_adm}
|
||||
make_condition_dict = lambda c_crossattn, c_adm: {"c_crossattn": c_crossattn, "c_adm": c_adm} # pylint: disable=unnecessary-lambda-assignment
|
||||
else:
|
||||
image_uncond = image_cond
|
||||
make_condition_dict = lambda c_crossattn, c_concat: {"c_crossattn": c_crossattn, "c_concat": [c_concat]}
|
||||
make_condition_dict = lambda c_crossattn, c_concat: {"c_crossattn": c_crossattn, "c_concat": [c_concat]} # pylint: disable=unnecessary-lambda-assignment
|
||||
|
||||
if not is_edit_model:
|
||||
x_in = torch.cat([torch.stack([x[i] for _ in range(n)]) for i, n in enumerate(repeats)] + [x])
|
||||
|
||||
@@ -32,7 +32,7 @@ class VAEApprox(nn.Module):
|
||||
|
||||
|
||||
def model():
|
||||
global sd_vae_approx_model
|
||||
global sd_vae_approx_model # pylint: disable=global-statement
|
||||
|
||||
if sd_vae_approx_model is None:
|
||||
model_path = os.path.join(paths.models_path, "VAE-approx", "model.pt")
|
||||
|
||||
+14
-2
@@ -4,10 +4,10 @@ import time
|
||||
import json
|
||||
import datetime
|
||||
import urllib.request
|
||||
from enum import Enum
|
||||
import gradio as gr
|
||||
import tqdm
|
||||
import requests
|
||||
# from ldm.models.diffusion.ddpm import LatentDiffusion
|
||||
from modules import errors, ui_components, shared_items, cmd_args
|
||||
from modules.paths_internal import models_path, script_path, data_path, sd_configs_path, sd_default_config, sd_model_file, default_sd_model_file, extensions_dir, extensions_builtin_dir # pylint: disable=W0611
|
||||
import modules.interrogate
|
||||
@@ -72,6 +72,11 @@ ui_reorder_categories = [
|
||||
]
|
||||
|
||||
|
||||
class Backend(Enum):
|
||||
ORIGINAL = 1
|
||||
DIFFUSERS = 2
|
||||
|
||||
|
||||
def reload_hypernetworks():
|
||||
from modules.hypernetworks import hypernetwork
|
||||
global hypernetworks # pylint: disable=W0603
|
||||
@@ -444,7 +449,7 @@ options_templates.update(options_section(('ui', "Live previews"), {
|
||||
"live_previews_enable": OptionInfo(True, "Show live previews of the created image"),
|
||||
"show_progress_grid": OptionInfo(True, "Show previews of all images generated in a batch as a grid"),
|
||||
"notification_audio_enable": OptionInfo(False, "Play a sound when images are finished generating"),
|
||||
"notification_audio_path": OptionInfo("html/notification.mp3","Path to notification sound",component_args=hide_dirs),
|
||||
"notification_audio_path": OptionInfo("html/notification.mp3","Path to notification sound", component_args=hide_dirs),
|
||||
"show_progress_every_n_steps": OptionInfo(1, "Show new live preview image every N sampling steps. Set to -1 to show after completion of batch.", gr.Slider, {"minimum": -1, "maximum": 32, "step": 1}),
|
||||
"show_progress_type": OptionInfo("Approx NN", "Image creation progress preview mode", gr.Radio, {"choices": ["Full", "Approx NN", "Approx cheap"]}),
|
||||
"live_preview_content": OptionInfo("Combined", "Live preview subject", gr.Radio, {"choices": ["Combined", "Prompt", "Negative prompt"]}),
|
||||
@@ -634,6 +639,13 @@ opts = Options()
|
||||
config_filename = cmd_opts.config
|
||||
opts.load(config_filename)
|
||||
cmd_opts = cmd_args.compatibility_args(opts, cmd_opts)
|
||||
if cmd_opts.backend == 'diffusers':
|
||||
log.info('Overriding backend to Diffusers')
|
||||
opts.data['sd_backend'] = 'Diffusers'
|
||||
if cmd_opts.backend == 'original':
|
||||
log.info('Overriding backend to Diffusers')
|
||||
opts.data['sd_backend'] = 'Original'
|
||||
backend = Backend.DIFFUSERS if opts.sd_backend == 'Diffusers' else Backend.ORIGINAL
|
||||
|
||||
prompt_styles = modules.styles.StyleDatabase(opts.styles_dir)
|
||||
cmd_opts.disable_extension_access = (cmd_opts.share or cmd_opts.listen or (cmd_opts.server_name or False)) and not cmd_opts.insecure
|
||||
|
||||
@@ -19,7 +19,7 @@ from torch.utils.checkpoint import checkpoint
|
||||
|
||||
|
||||
def narrow_trunc(
|
||||
input: Tensor,
|
||||
input: Tensor, # pylint: disable=redefined-builtin
|
||||
dim: int,
|
||||
start: int,
|
||||
length: int
|
||||
@@ -79,8 +79,8 @@ def _query_chunk_attention(
|
||||
summarize_chunk: SummarizeChunk,
|
||||
kv_chunk_size: int,
|
||||
) -> Tensor:
|
||||
batch_x_heads, k_tokens, k_channels_per_head = key.shape
|
||||
_, _, v_channels_per_head = value.shape
|
||||
_batch_x_heads, k_tokens, _k_channels_per_head = key.shape
|
||||
_, _, _v_channels_per_head = value.shape
|
||||
|
||||
def chunk_scanner(chunk_idx: int) -> AttnChunk:
|
||||
key_chunk = narrow_trunc(
|
||||
@@ -113,7 +113,6 @@ def _query_chunk_attention(
|
||||
return all_values / all_weights
|
||||
|
||||
|
||||
# TODO: refactor CrossAttention#get_attention_scores to share code with this
|
||||
def _get_attention_scores_no_kv_chunking(
|
||||
query: Tensor,
|
||||
key: Tensor,
|
||||
@@ -164,7 +163,7 @@ def efficient_dot_product_attention(
|
||||
Returns:
|
||||
Output of shape `[batch * num_heads, query_tokens, channels_per_head]`.
|
||||
"""
|
||||
batch_x_heads, q_tokens, q_channels_per_head = query.shape
|
||||
_batch_x_heads, q_tokens, q_channels_per_head = query.shape
|
||||
_, k_tokens, _ = key.shape
|
||||
scale = q_channels_per_head ** -0.5
|
||||
|
||||
@@ -202,7 +201,7 @@ def efficient_dot_product_attention(
|
||||
value=value,
|
||||
)
|
||||
|
||||
# TODO: maybe we should use torch.empty_like(query) to allocate storage in-advance,
|
||||
# maybe we should use torch.empty_like(query) to allocate storage in-advance,
|
||||
# and pass slices to be mutated, instead of torch.cat()ing the returned slices
|
||||
res = torch.cat([
|
||||
compute_query_chunk_attn(
|
||||
|
||||
@@ -207,7 +207,7 @@ class EmbeddingDatabase:
|
||||
continue
|
||||
|
||||
def load_textual_inversion_embeddings(self, force_reload=False):
|
||||
if shared.opts.sd_backend == 'Diffusers': # TODO Diffusers
|
||||
if shared.backend == shared.Backend.DIFFUSERS: # TODO Diffusers
|
||||
return
|
||||
if not force_reload:
|
||||
need_reload = False
|
||||
|
||||
+2
-2
@@ -3,7 +3,7 @@ from modules import sd_samplers, shared
|
||||
from modules.generation_parameters_copypaste import create_override_settings_dict
|
||||
from modules.processing import StableDiffusionProcessingTxt2Img, process_images
|
||||
# from modules.shared import opts, sd_model, debug
|
||||
from modules.ui import plaintext_to_html
|
||||
from modules.ui import plaintext_to_html, infotext_to_html
|
||||
from modules.memstats import memory_stats
|
||||
|
||||
|
||||
@@ -58,4 +58,4 @@ def txt2img(id_task: str, prompt: str, negative_prompt: str, prompt_styles, step
|
||||
p.close()
|
||||
generation_info_js = processed.js()
|
||||
shared.log.debug(f'Processed: {len(processed.images)} Memory: {memory_stats()} txt')
|
||||
return processed.images, generation_info_js, plaintext_to_html(processed.info), plaintext_to_html(processed.comments)
|
||||
return processed.images, generation_info_js, infotext_to_html(processed.info), plaintext_to_html(processed.comments)
|
||||
|
||||
+14
-9
@@ -13,7 +13,7 @@ from modules.call_queue import wrap_gradio_gpu_call, wrap_queued_call, wrap_grad
|
||||
from modules import sd_hijack, sd_models, script_callbacks, ui_extensions, deepbooru, sd_vae, extra_networks, ui_common, ui_postprocessing
|
||||
from modules.ui_components import FormRow, FormColumn, FormGroup, ToolButton, FormHTML # pylint: disable=unused-import
|
||||
from modules.paths import script_path, data_path
|
||||
from modules.shared import opts, cmd_opts
|
||||
from modules.shared import opts, cmd_opts, backend, Backend
|
||||
from modules import prompt_parser
|
||||
import modules.codeformer_model
|
||||
import modules.generation_parameters_copypaste as parameters_copypaste
|
||||
@@ -50,12 +50,12 @@ sample_img2img = sample_img2img if os.path.exists(sample_img2img) else None
|
||||
# Important that they exactly match script.js for tooltip to work.
|
||||
random_symbol = '\U0001f3b2\ufe0f' # 🎲️
|
||||
reuse_symbol = '\u267b\ufe0f' # ♻️
|
||||
paste_symbol = '\u2199\ufe0f' # ↙
|
||||
refresh_symbol = '\U0001f504' # 🔄
|
||||
save_style_symbol = '\U0001f4be' # 💾
|
||||
apply_style_symbol = '\U0001f4cb' # 📋
|
||||
clear_prompt_symbol = '\U0001f5d1\ufe0f' # 🗑️
|
||||
extra_networks_symbol = '\U0001F3B4' # 🎴
|
||||
paste_symbol = '\U0001F4D8' # '\u2199\ufe0f' # ↙
|
||||
refresh_symbol = '\U0001F504' # 🔄
|
||||
save_style_symbol = '\U0001F6C5' # '\U0001f4be' # 💾
|
||||
apply_style_symbol = '\U0001F9F3' # '\U0001f4cb' # 📋
|
||||
clear_prompt_symbol = '\U0001F6AE' # '\U0001f5d1\ufe0f' # 🗑️
|
||||
extra_networks_symbol = '\U0001F310' # '\U0001F3B4' # 🎴
|
||||
switch_values_symbol = '\U000021C5' # ⇅
|
||||
|
||||
|
||||
@@ -63,6 +63,10 @@ def plaintext_to_html(text):
|
||||
return ui_common.plaintext_to_html(text)
|
||||
|
||||
|
||||
def infotext_to_html(text):
|
||||
return ui_common.infotext_to_html(text)
|
||||
|
||||
|
||||
def send_gradio_gallery_to_image(x):
|
||||
if len(x) == 0:
|
||||
return None
|
||||
@@ -204,7 +208,7 @@ def update_token_counter(text, steps):
|
||||
prompt_schedules = [[[steps, text]]]
|
||||
flat_prompts = reduce(lambda list1, list2: list1+list2, prompt_schedules)
|
||||
prompts = [prompt_text for step, prompt_text in flat_prompts]
|
||||
if opts.sd_backend == 'Original':
|
||||
if backend == Backend.ORIGINAL:
|
||||
token_count, max_length = max([sd_hijack.model_hijack.get_prompt_lengths(prompt) for prompt in prompts], key=lambda args: args[0])
|
||||
else:
|
||||
tokenizer = modules.shared.sd_model.tokenizer
|
||||
@@ -1323,9 +1327,10 @@ def create_ui():
|
||||
unload_sd_model = gr.Button(value='Unload checkpoint', variant='primary', elem_id="sett_unload_sd_model")
|
||||
reload_sd_model = gr.Button(value='Reload checkpoint', variant='primary', elem_id="sett_reload_sd_model")
|
||||
# reload_script_bodies = gr.Button(value='Reload scripts', variant='primary', elem_id="settings_reload_script_bodies")
|
||||
with gr.Row():
|
||||
_settings_search = gr.Text(label="Search", elem_id="settings_search")
|
||||
|
||||
result = gr.HTML(elem_id="settings_result")
|
||||
|
||||
quicksettings_names = opts.quicksettings_list
|
||||
quicksettings_names = {x: i for i, x in enumerate(quicksettings_names) if x != 'quicksettings'}
|
||||
quicksettings_list = []
|
||||
|
||||
+102
-102
@@ -1,6 +1,7 @@
|
||||
import json
|
||||
import html
|
||||
import os
|
||||
import shutil
|
||||
import platform
|
||||
import subprocess
|
||||
import gradio as gr
|
||||
@@ -8,6 +9,7 @@ from modules import call_queue, shared
|
||||
from modules.generation_parameters_copypaste import image_from_url_text
|
||||
import modules.images
|
||||
|
||||
|
||||
folder_symbol = '\U0001f4c2' # 📂
|
||||
|
||||
|
||||
@@ -16,60 +18,99 @@ def update_generation_info(generation_info, html_info, img_index):
|
||||
generation_info = json.loads(generation_info)
|
||||
if img_index < 0 or img_index >= len(generation_info["infotexts"]):
|
||||
return html_info, gr.update()
|
||||
return plaintext_to_html(generation_info["infotexts"][img_index]), gr.update()
|
||||
html_text = infotext_to_html(generation_info["infotexts"][img_index])
|
||||
return html_text, gr.update()
|
||||
except Exception:
|
||||
pass
|
||||
# if the json parse or anything else fails, just return the old html_info
|
||||
return html_info, gr.update()
|
||||
|
||||
|
||||
def plaintext_to_html(text):
|
||||
text = "<p>" + "<br>\n".join([f"{html.escape(x)}" for x in text.split('\n')]) + "</p>"
|
||||
return text
|
||||
res = '<p class="plaintext">' + "<br>\n".join([f"{html.escape(x)}" for x in text.split('\n')]) + '</p>'
|
||||
return res
|
||||
|
||||
|
||||
def infotext_to_html(text):
|
||||
res = '<p class="html_info">Prompt: ' + html.escape(text or '').replace('\n', '<br>') + '</p>'
|
||||
sections = res.split('Steps:') # before and after prompt+negprompt'
|
||||
if len(sections) > 1:
|
||||
res = sections[0] + '<br>Steps: ' + sections[1].strip().replace(', ', ' | ')
|
||||
res = res.replace('<br><br>', '<br>')
|
||||
return res
|
||||
|
||||
|
||||
def delete_files(js_data, images, _do_make_zip, index):
|
||||
try:
|
||||
data = json.loads(js_data)
|
||||
except Exception:
|
||||
data = { 'index_of_first_image': 0 }
|
||||
start_index = 0
|
||||
if index > -1 and shared.opts.save_selected_only and (index >= data['index_of_first_image']):
|
||||
images = [images[index]]
|
||||
start_index = index
|
||||
filenames = []
|
||||
filenames = []
|
||||
fullfns = []
|
||||
for _image_index, filedata in enumerate(images, start_index):
|
||||
if 'name' in filedata and os.path.isfile(filedata['name']):
|
||||
fullfn = filedata['name']
|
||||
filenames.append(os.path.basename(fullfn))
|
||||
try:
|
||||
os.remove(fullfn)
|
||||
fullfns.append(fullfn)
|
||||
shared.log.info(f"Deleting image: {fullfn}")
|
||||
except Exception as e:
|
||||
shared.log.error(f'Error deleting file: {fullfn} {e}')
|
||||
images = [image for image in images if image['name'] not in fullfns]
|
||||
return images, plaintext_to_html(f"Deleted: {filenames[0] if len(filenames) > 0 else 'none'}")
|
||||
|
||||
|
||||
def save_files(js_data, images, do_make_zip, index):
|
||||
if js_data is None or len(js_data) == 0:
|
||||
return
|
||||
filenames = []
|
||||
fullfns = []
|
||||
os.makedirs(shared.opts.outdir_save, exist_ok=True)
|
||||
|
||||
#quick dictionary to class object conversion. Its necessary due apply_filename_pattern requiring it
|
||||
class MyObject:
|
||||
class MyObject: #quick dictionary to class object conversion. Its necessary due apply_filename_pattern requiring it
|
||||
def __init__(self, d=None):
|
||||
if d is not None:
|
||||
for key, value in d.items():
|
||||
setattr(self, key, value)
|
||||
|
||||
data = json.loads(js_data)
|
||||
try:
|
||||
data = json.loads(js_data)
|
||||
except Exception:
|
||||
data = { 'index_of_first_image': 0 }
|
||||
p = MyObject(data)
|
||||
path = shared.opts.outdir_save
|
||||
save_to_dirs = shared.opts.use_save_to_dirs_for_ui
|
||||
extension: str = shared.opts.samples_format
|
||||
start_index = 0
|
||||
if index > -1 and shared.opts.save_selected_only and (index >= data["index_of_first_image"]): # ensures we are looking at a specific non-grid picture, and we have save_selected_only
|
||||
if index > -1 and shared.opts.save_selected_only and (index >= data['index_of_first_image']): # ensures we are looking at a specific non-grid picture, and we have save_selected_only # pylint: disable=no-member
|
||||
images = [images[index]]
|
||||
start_index = index
|
||||
os.makedirs(shared.opts.outdir_save, exist_ok=True)
|
||||
filenames = []
|
||||
fullfns = []
|
||||
for image_index, filedata in enumerate(images, start_index):
|
||||
image = image_from_url_text(filedata)
|
||||
is_grid = image_index < p.index_of_first_image # pylint: disable=no-member
|
||||
i = 0 if is_grid else (image_index - p.index_of_first_image) # pylint: disable=no-member
|
||||
if len(p.all_seeds) <= i: # pylint: disable=no-member
|
||||
p.all_seeds.append(p.seed) # pylint: disable=no-member
|
||||
if len(p.all_prompts) <= i: # pylint: disable=no-member
|
||||
p.all_prompts.append(p.prompt) # pylint: disable=no-member
|
||||
fullfn, txt_fullfn = modules.images.save_image(image, path, "", seed=p.all_seeds[i], prompt=p.all_prompts[i], extension=extension, info=p.infotexts[image_index], grid=is_grid, p=p, save_to_dirs=save_to_dirs) # pylint: disable=no-member
|
||||
if fullfn is None:
|
||||
continue
|
||||
filename = os.path.relpath(fullfn, path)
|
||||
filenames.append(filename)
|
||||
fullfns.append(fullfn)
|
||||
if txt_fullfn:
|
||||
filenames.append(os.path.basename(txt_fullfn))
|
||||
fullfns.append(txt_fullfn)
|
||||
if 'name' in filedata and ('tmp' not in filedata['name']) and os.path.isfile(filedata['name']):
|
||||
fullfn = filedata['name']
|
||||
filenames.append(os.path.basename(fullfn))
|
||||
fullfns.append(fullfn)
|
||||
shutil.copy(fullfn, shared.opts.outdir_save)
|
||||
shared.log.info(f"Copying image: {fullfn} -> {shared.opts.outdir_save}")
|
||||
else:
|
||||
image = image_from_url_text(filedata)
|
||||
is_grid = image_index < p.index_of_first_image # pylint: disable=no-member
|
||||
i = 0 if is_grid else (image_index - p.index_of_first_image) # pylint: disable=no-member
|
||||
if len(p.all_seeds) <= i: # pylint: disable=no-member
|
||||
p.all_seeds.append(p.seed) # pylint: disable=no-member
|
||||
if len(p.all_prompts) <= i: # pylint: disable=no-member
|
||||
p.all_prompts.append(p.prompt) # pylint: disable=no-member
|
||||
fullfn, txt_fullfn = modules.images.save_image(image, shared.opts.outdir_save, "", seed=p.all_seeds[i], prompt=p.all_prompts[i], extension=shared.opts.samples_format, info=p.infotexts[image_index], grid=is_grid, p=p, save_to_dirs=shared.opts.use_save_to_dirs_for_ui) # pylint: disable=no-member
|
||||
if fullfn is None:
|
||||
continue
|
||||
filename = os.path.relpath(fullfn, shared.opts.outdir_save)
|
||||
filenames.append(filename)
|
||||
fullfns.append(fullfn)
|
||||
if txt_fullfn:
|
||||
filenames.append(os.path.basename(txt_fullfn))
|
||||
fullfns.append(txt_fullfn)
|
||||
if do_make_zip:
|
||||
zip_filepath = os.path.join(path, "images.zip")
|
||||
zip_filepath = os.path.join(shared.opts.outdir_save, "images.zip")
|
||||
from zipfile import ZipFile
|
||||
with ZipFile(zip_filepath, "w") as zip_file:
|
||||
for i in range(len(fullfns)):
|
||||
@@ -105,87 +146,46 @@ def create_output_panel(tabname, outdir):
|
||||
with gr.Group(elem_id=f"{tabname}_gallery_container"):
|
||||
result_gallery = gr.Gallery(value=['html/logo.png'], label='Output', show_label=False, elem_id=f"{tabname}_gallery").style(preview=False, container=False, columns=[1,2,3,4,5,6]) # <576px, <768px, <992px, <1200px, <1400px, >1400px
|
||||
|
||||
generation_info = None
|
||||
with gr.Column():
|
||||
with gr.Row(elem_id=f"image_buttons_{tabname}", elem_classes="image-buttons"):
|
||||
open_folder_button = gr.Button('show', visible=not shared.cmd_opts.hide_ui_dir_config)
|
||||
|
||||
if tabname != "extras":
|
||||
save = gr.Button('save', elem_id=f'save_{tabname}')
|
||||
save_zip = gr.Button('zip', elem_id=f'save_zip_{tabname}')
|
||||
|
||||
save = gr.Button('save', elem_id=f'save_{tabname}')
|
||||
save_zip = gr.Button('zip', elem_id=f'save_zip_{tabname}')
|
||||
delete = gr.Button('delete', elem_id=f'delete_{tabname}')
|
||||
buttons = parameters_copypaste.create_buttons(["img2img", "inpaint", "extras"])
|
||||
|
||||
open_folder_button.click(
|
||||
fn=lambda: open_folder(shared.opts.outdir_samples or outdir),
|
||||
inputs=[],
|
||||
outputs=[],
|
||||
)
|
||||
|
||||
if tabname != "extras":
|
||||
download_files = gr.File(None, file_count="multiple", interactive=False, show_label=False, visible=False, elem_id=f'download_files_{tabname}')
|
||||
|
||||
with gr.Group():
|
||||
html_info = gr.HTML(elem_id=f'html_info_{tabname}', elem_classes="infotext")
|
||||
html_log = gr.HTML(elem_id=f'html_log_{tabname}')
|
||||
|
||||
generation_info = gr.Textbox(visible=False, elem_id=f'generation_info_{tabname}')
|
||||
if tabname == 'txt2img' or tabname == 'img2img':
|
||||
generation_info_button = gr.Button(visible=False, elem_id=f"{tabname}_generation_info_button")
|
||||
generation_info_button.click(
|
||||
fn=update_generation_info,
|
||||
_js="function(x, y, z){ return [x, y, selected_gallery_index()] }",
|
||||
inputs=[generation_info, html_info, html_info],
|
||||
outputs=[html_info, html_info],
|
||||
show_progress=False,
|
||||
)
|
||||
|
||||
save.click(
|
||||
fn=call_queue.wrap_gradio_call(save_files),
|
||||
_js="(x, y, z, w) => [x, y, false, selected_gallery_index()]",
|
||||
inputs=[
|
||||
generation_info,
|
||||
result_gallery,
|
||||
html_info,
|
||||
html_info,
|
||||
],
|
||||
outputs=[
|
||||
download_files,
|
||||
html_log,
|
||||
],
|
||||
show_progress=False,
|
||||
)
|
||||
|
||||
save_zip.click(
|
||||
fn=call_queue.wrap_gradio_call(save_files),
|
||||
_js="(x, y, z, w) => [x, y, true, selected_gallery_index()]",
|
||||
inputs=[
|
||||
generation_info,
|
||||
result_gallery,
|
||||
html_info,
|
||||
html_info,
|
||||
],
|
||||
outputs=[
|
||||
download_files,
|
||||
html_log,
|
||||
]
|
||||
)
|
||||
|
||||
else:
|
||||
html_info_x = gr.HTML(elem_id=f'html_info_x_{tabname}')
|
||||
open_folder_button.click(fn=lambda: open_folder(shared.opts.outdir_samples or outdir), inputs=[], outputs=[])
|
||||
download_files = gr.File(None, file_count="multiple", interactive=False, show_label=False, visible=False, elem_id=f'download_files_{tabname}')
|
||||
with gr.Group():
|
||||
html_info = gr.HTML(elem_id=f'html_info_{tabname}', elem_classes="infotext")
|
||||
html_log = gr.HTML(elem_id=f'html_log_{tabname}')
|
||||
generation_info = gr.Textbox(visible=False, elem_id=f'generation_info_{tabname}')
|
||||
generation_info_button = gr.Button(visible=False, elem_id=f"{tabname}_generation_info_button")
|
||||
generation_info_button.click(fn=update_generation_info, _js="function(x, y, z){ return [x, y, selected_gallery_index()] }", show_progress=False,
|
||||
inputs=[generation_info, html_info, html_info],
|
||||
outputs=[html_info, html_info],
|
||||
)
|
||||
save.click(fn=call_queue.wrap_gradio_call(save_files), _js="(x, y, z, w) => [x, y, false, selected_gallery_index()]", show_progress=False,
|
||||
inputs=[generation_info, result_gallery, html_info, html_info],
|
||||
outputs=[download_files, html_log],
|
||||
)
|
||||
save_zip.click(fn=call_queue.wrap_gradio_call(save_files), _js="(x, y, z, w) => [x, y, true, selected_gallery_index()]",
|
||||
inputs=[generation_info, result_gallery, html_info, html_info],
|
||||
outputs=[download_files, html_log],
|
||||
)
|
||||
delete.click(fn=call_queue.wrap_gradio_call(delete_files), _js="(x, y, z, w) => [x, y, true, selected_gallery_index()]",
|
||||
inputs=[generation_info, result_gallery, html_info, html_info],
|
||||
outputs=[result_gallery, html_log],
|
||||
)
|
||||
|
||||
paste_field_names = []
|
||||
if tabname == "txt2img":
|
||||
paste_field_names = modules.scripts.scripts_txt2img.paste_field_names
|
||||
elif tabname == "img2img":
|
||||
paste_field_names = modules.scripts.scripts_img2img.paste_field_names
|
||||
|
||||
else:
|
||||
paste_field_names = []
|
||||
for paste_tabname, paste_button in buttons.items():
|
||||
parameters_copypaste.register_paste_params_button(parameters_copypaste.ParamBinding(
|
||||
paste_button=paste_button, tabname=paste_tabname, source_tabname="txt2img" if tabname == "txt2img" else None, source_image_component=result_gallery,
|
||||
paste_field_names=paste_field_names
|
||||
paste_button=paste_button, tabname=paste_tabname, source_tabname=("txt2img" if tabname == "txt2img" else None), source_image_component=result_gallery, paste_field_names=paste_field_names
|
||||
))
|
||||
|
||||
return result_gallery, generation_info if tabname != "extras" else html_info_x, html_info, html_log
|
||||
return result_gallery, generation_info, html_info, html_log
|
||||
|
||||
+445
-445
@@ -1,445 +1,445 @@
|
||||
import json
|
||||
import os.path
|
||||
import shutil
|
||||
import errno
|
||||
import html
|
||||
from datetime import datetime
|
||||
import git
|
||||
import gradio as gr
|
||||
from modules import extensions, shared, paths, errors
|
||||
from modules.call_queue import wrap_gradio_gpu_call
|
||||
|
||||
|
||||
extensions_index = "https://vladmandic.github.io/sd-data/pages/extensions.json"
|
||||
hide_tags = ["localization"]
|
||||
extensions_list = []
|
||||
sort_ordering = {
|
||||
"default": (True, lambda x: x.get('sort_default', '')),
|
||||
"user extensions": (True, lambda x: x.get('sort_user', '')),
|
||||
"update avilable": (True, lambda x: x.get('sort_update', '')),
|
||||
"updated date": (True, lambda x: x.get('updated', '2000-01-01T00:00')),
|
||||
"created date": (False, lambda x: x.get('created', '2000-01-01T00:00')),
|
||||
"name": (False, lambda x: x.get('name', '').lower()),
|
||||
"enabled": (False, lambda x: x.get('sort_enabled', '').lower()),
|
||||
"size": (True, lambda x: x.get('size', 0)),
|
||||
"stars": (True, lambda x: x.get('stars', 0)),
|
||||
"commits": (True, lambda x: x.get('commits', 0)),
|
||||
"issues": (True, lambda x: x.get('issues', 0)),
|
||||
}
|
||||
|
||||
|
||||
def update_extension_list():
|
||||
global extensions_list # pylint: disable=global-statement
|
||||
try:
|
||||
with open(os.path.join(paths.script_path, "html", "extensions.json"), "r", encoding="utf-8") as f:
|
||||
extensions_list = json.loads(f.read())
|
||||
shared.log.debug(f'Extensions list loaded: {os.path.join(paths.script_path, "html", "extensions.json")}')
|
||||
except:
|
||||
shared.log.debug(f'Extensions list failed to load: {os.path.join(paths.script_path, "html", "extensions.json")}')
|
||||
found = []
|
||||
for ext in extensions.extensions:
|
||||
ext.read_info_from_repo()
|
||||
for ext in extensions_list:
|
||||
installed = [extension for extension in extensions.extensions
|
||||
if extension.git_name == ext['name']
|
||||
or extension.name == ext['name']
|
||||
or (extension.remote or '').startswith(ext['url'].replace('.git', ''))]
|
||||
if len(installed) > 0:
|
||||
found.append(installed[0])
|
||||
not_matched = [extension for extension in extensions.extensions if extension not in found]
|
||||
for ext in not_matched:
|
||||
entry = {
|
||||
"name": ext.name or "",
|
||||
"description": ext.description or "",
|
||||
"url": ext.remote or "",
|
||||
"tags": [],
|
||||
"stars": 0,
|
||||
"issues": 0,
|
||||
"commits": 0,
|
||||
"size": 0,
|
||||
"long": ext.git_name or ext.name or "",
|
||||
"added": ext.ctime,
|
||||
"created": ext.ctime,
|
||||
"updated": ext.mtime,
|
||||
}
|
||||
extensions_list.append(entry)
|
||||
|
||||
|
||||
def check_access():
|
||||
assert not shared.cmd_opts.disable_extension_access, "extension access disabled because of command line flags"
|
||||
|
||||
|
||||
def apply_and_restart(disable_list, update_list, disable_all):
|
||||
check_access()
|
||||
shared.log.debug(f'Extensions apply: disable={disable_list} update={update_list}')
|
||||
disabled = json.loads(disable_list)
|
||||
assert type(disabled) == list, f"wrong disable_list data for apply_and_restart: {disable_list}"
|
||||
update = json.loads(update_list)
|
||||
assert type(update) == list, f"wrong update_list data for apply_and_restart: {update_list}"
|
||||
update = set(update)
|
||||
for ext in extensions.extensions:
|
||||
if ext.name not in update:
|
||||
continue
|
||||
try:
|
||||
ext.fetch_and_reset_hard()
|
||||
except Exception as e:
|
||||
errors.display(e, f'extensions apply update: {ext.name}')
|
||||
shared.opts.disabled_extensions = disabled
|
||||
shared.opts.disable_all_extensions = disable_all
|
||||
shared.opts.save(shared.config_filename)
|
||||
shared.restart_server(restart=True)
|
||||
|
||||
|
||||
def check_updates(_id_task, disable_list, search_text, sort_column):
|
||||
check_access()
|
||||
disabled = json.loads(disable_list)
|
||||
assert type(disabled) == list, f"wrong disable_list data for apply_and_restart: {disable_list}"
|
||||
exts = [ext for ext in extensions.extensions if ext.remote is not None and ext.name not in disabled]
|
||||
shared.log.info(f'Extensions update check: update={len(exts)} disabled={len(disable_list)}')
|
||||
shared.state.job_count = len(exts)
|
||||
for ext in exts:
|
||||
shared.state.textinfo = ext.name
|
||||
try:
|
||||
ext.check_updates()
|
||||
if ext.can_update:
|
||||
ext.fetch_and_reset_hard()
|
||||
ext.read_info_from_repo()
|
||||
commit_date = ext.commit_date or 1577836800
|
||||
shared.log.info(f'Extensions updated: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}')
|
||||
else:
|
||||
commit_date = ext.commit_date or 1577836800
|
||||
shared.log.debug(f'Extensions no update available: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}')
|
||||
except FileNotFoundError as e:
|
||||
if 'FETCH_HEAD' not in str(e):
|
||||
raise
|
||||
except Exception:
|
||||
errors.display(e, f'extensions check update: {ext.name}')
|
||||
shared.state.nextjob()
|
||||
return refresh_extensions_list_from_data(search_text, sort_column), "Extension update complete | Restart required"
|
||||
|
||||
|
||||
def make_commit_link(commit_hash, remote, text=None):
|
||||
if text is None:
|
||||
text = commit_hash[:8]
|
||||
if remote.startswith("https://github.com/"):
|
||||
href = os.path.join(remote, "commit", commit_hash)
|
||||
return f'<a href="{href}" target="_blank">{text}</a>'
|
||||
else:
|
||||
return text
|
||||
|
||||
|
||||
def normalize_git_url(url):
|
||||
if url is None:
|
||||
return ""
|
||||
url = url.replace(".git", "")
|
||||
return url
|
||||
|
||||
|
||||
def install_extension_from_url(dirname, url, branch_name, search_text, sort_column):
|
||||
check_access()
|
||||
assert url, 'No URL specified'
|
||||
if dirname is None or dirname == "":
|
||||
*parts, last_part = url.split('/') # pylint: disable=unused-variable
|
||||
last_part = normalize_git_url(last_part)
|
||||
dirname = last_part
|
||||
target_dir = os.path.join(extensions.extensions_dir, dirname)
|
||||
shared.log.info(f'Installing extension: {url} into {target_dir}')
|
||||
assert not os.path.exists(target_dir), f'Extension directory already exists: {target_dir}'
|
||||
normalized_url = normalize_git_url(url)
|
||||
assert len([x for x in extensions.extensions if normalize_git_url(x.remote) == normalized_url]) == 0, 'Extension with this URL is already installed'
|
||||
tmpdir = os.path.join(paths.data_path, "tmp", dirname)
|
||||
try:
|
||||
shutil.rmtree(tmpdir, True)
|
||||
if not branch_name:
|
||||
# if no branch is specified, use the default branch
|
||||
with git.Repo.clone_from(url, tmpdir) as repo:
|
||||
repo.remote().fetch()
|
||||
for submodule in repo.submodules:
|
||||
submodule.update()
|
||||
else:
|
||||
with git.Repo.clone_from(url, tmpdir, branch=branch_name) as repo:
|
||||
repo.remote().fetch()
|
||||
for submodule in repo.submodules:
|
||||
submodule.update()
|
||||
try:
|
||||
os.rename(tmpdir, target_dir)
|
||||
except OSError as err:
|
||||
if err.errno == errno.EXDEV:
|
||||
shutil.move(tmpdir, target_dir)
|
||||
else:
|
||||
raise err
|
||||
from launch import run_extension_installer
|
||||
run_extension_installer(target_dir)
|
||||
extensions.list_extensions()
|
||||
return [refresh_extensions_list_from_data(search_text, sort_column), html.escape(f"Extension installed: {target_dir} | Restart required")]
|
||||
finally:
|
||||
shutil.rmtree(tmpdir, True)
|
||||
|
||||
|
||||
def install_extension(extension_to_install, search_text, sort_column):
|
||||
shared.log.info(f'Extension install: {extension_to_install}')
|
||||
code, message = install_extension_from_url(None, extension_to_install, None, search_text, sort_column)
|
||||
return code, message
|
||||
|
||||
|
||||
def uninstall_extension(extension_path, search_text, sort_column):
|
||||
def errorRemoveReadonly(func, path, exc):
|
||||
import stat
|
||||
excvalue = exc[1]
|
||||
shared.log.debug(f'Exception during cleanup: {func} {path} {excvalue.strerror}')
|
||||
if func in (os.rmdir, os.remove, os.unlink) and excvalue.errno == errno.EACCES:
|
||||
shared.log.debug(f'Retrying cleanup: {path}')
|
||||
os.chmod(path, stat.S_IRWXU | stat.S_IRWXG | stat.S_IRWXO)
|
||||
func(path)
|
||||
|
||||
ext = [extension for extension in extensions.extensions if os.path.abspath(extension.path) == os.path.abspath(extension_path)]
|
||||
if len(ext) > 0 and os.path.isdir(extension_path):
|
||||
found = ext[0]
|
||||
try:
|
||||
shutil.rmtree(found.path, ignore_errors=False, onerror=errorRemoveReadonly)
|
||||
except Exception as e:
|
||||
shared.log.warning(f'Extension uninstall failed: {found.path} {e}')
|
||||
extensions.extensions = [extension for extension in extensions.extensions if os.path.abspath(found.path) != os.path.abspath(extension_path)]
|
||||
update_extension_list()
|
||||
code = refresh_extensions_list_from_data(search_text, sort_column)
|
||||
shared.log.info(f'Extension uninstalled: {found.path}')
|
||||
return code, f"Extension uninstalled: {found.path} | Restart required"
|
||||
else:
|
||||
shared.log.warning(f'Extension uninstall cannot find extension: {extension_path}')
|
||||
code = refresh_extensions_list_from_data(search_text, sort_column)
|
||||
return code, f"Extension uninstalled failed: {extension_path}"
|
||||
|
||||
|
||||
def update_extension(extension_path, search_text, sort_column):
|
||||
exts = [extension for extension in extensions.extensions if os.path.abspath(extension.path) == os.path.abspath(extension_path)]
|
||||
shared.state.job_count = len(exts)
|
||||
for ext in exts:
|
||||
shared.log.debug(f'Extensions update start: {ext.name} {ext.commit_hash} {ext.commit_date}')
|
||||
shared.state.textinfo = ext.name
|
||||
try:
|
||||
ext.check_updates()
|
||||
if ext.can_update:
|
||||
ext.fetch_and_reset_hard()
|
||||
ext.read_info_from_repo()
|
||||
commit_date = ext.commit_date or 1577836800
|
||||
shared.log.info(f'Extensions updated: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}')
|
||||
else:
|
||||
commit_date = ext.commit_date or 1577836800
|
||||
shared.log.info(f'Extensions no update available: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}')
|
||||
except FileNotFoundError as e:
|
||||
if 'FETCH_HEAD' not in str(e):
|
||||
raise
|
||||
except Exception as e:
|
||||
shared.log.error(f'Extensions update failed: {ext.name}')
|
||||
errors.display(e, f'extensions check update: {ext.name}')
|
||||
shared.log.debug(f'Extensions update finish: {ext.name} {ext.commit_hash} {ext.commit_date}')
|
||||
shared.state.nextjob()
|
||||
return refresh_extensions_list_from_data(search_text, sort_column), f"Extension updated | {extension_path} | Restart required"
|
||||
|
||||
|
||||
def refresh_extensions_list(search_text, sort_column):
|
||||
global extensions_list # pylint: disable=global-statement
|
||||
import urllib.request
|
||||
try:
|
||||
with urllib.request.urlopen(extensions_index) as response:
|
||||
text = response.read()
|
||||
extensions_list = json.loads(text)
|
||||
with open(os.path.join(paths.script_path, "html", "extensions.json"), "w", encoding="utf-8") as outfile:
|
||||
json_object = json.dumps(extensions_list, indent=2)
|
||||
outfile.write(json_object)
|
||||
shared.log.debug(f'Updated extensions list: {len(extensions_list)} {extensions_index} {outfile}')
|
||||
except Exception as e:
|
||||
shared.log.warning(f'Updated extensions list failed: {extensions_index} {e}')
|
||||
update_extension_list()
|
||||
code = refresh_extensions_list_from_data(search_text, sort_column)
|
||||
return code, f'Extensions | {len(extensions.extensions)} registered | {len(extensions_list)} available'
|
||||
|
||||
|
||||
def search_extensions(search_text, sort_column):
|
||||
code = refresh_extensions_list_from_data(search_text, sort_column)
|
||||
return code, f'Search | {search_text} | {sort_column}'
|
||||
|
||||
|
||||
def refresh_extensions_list_from_data(search_text, sort_column):
|
||||
shared.log.debug(f'Extensions manager: refresh list search="{search_text}" sort="{sort_column}"')
|
||||
code = """
|
||||
<table id="extensions">
|
||||
<colgroup>
|
||||
<col style="width: 1%; background: var(--table-border-color)">
|
||||
<col style="width: 20%; background: var(--table-border-color)">
|
||||
<col style="width: 59%;">
|
||||
<col style="width: 5%; background: var(--panel-background-fill)">
|
||||
<col style="width: 10%; background: var(--panel-background-fill)">
|
||||
<col style="width: 5%; background: var(--table-border-color)">
|
||||
</colgroup>
|
||||
<thead style="font-size: 110%; border-style: solid; border-bottom: 1px var(--button-primary-border-color) solid">
|
||||
<tr>
|
||||
<th>Enabled</th>
|
||||
<th>Extension</th>
|
||||
<th>Description</th>
|
||||
<th>Type</th>
|
||||
<th>Current version</th>
|
||||
<th></th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>"""
|
||||
for ext in extensions_list:
|
||||
extension = [extension for extension in extensions.extensions if extension.git_name == ext['name'] or extension.name == ext['name']]
|
||||
if len(extension) > 0:
|
||||
extension[0].read_info_from_repo()
|
||||
ext['installed'] = len(extension) > 0
|
||||
ext['commit_date'] = extension[0].commit_date if len(extension) > 0 else 1577836800
|
||||
ext['is_builtin'] = extension[0].is_builtin if len(extension) > 0 else False
|
||||
ext['version'] = extension[0].version if len(extension) > 0 else ''
|
||||
ext['enabled'] = extension[0].enabled if len(extension) > 0 else ''
|
||||
ext['remote'] = extension[0].remote if len(extension) > 0 else None
|
||||
ext['path'] = extension[0].path if len(extension) > 0 else ''
|
||||
ext['sort_default'] = f"{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}"
|
||||
sort_reverse, sort_function = sort_ordering[sort_column]
|
||||
|
||||
def dt(x: str):
|
||||
val = ext.get(x, None)
|
||||
if val is not None:
|
||||
return datetime.fromisoformat(val[:-1]).strftime('%a %b%d %Y %H:%M')
|
||||
else:
|
||||
return "N/A"
|
||||
|
||||
for ext in sorted(extensions_list, key=sort_function, reverse=sort_reverse):
|
||||
name = ext.get("name", "unknown")
|
||||
added = dt('added')
|
||||
created = dt('created')
|
||||
pushed = dt('pushed')
|
||||
updated = dt('updated')
|
||||
url = ext.get('url', None)
|
||||
size = ext.get('size', 0)
|
||||
stars = ext.get('stars', 0)
|
||||
issues = ext.get('issues', 0)
|
||||
commits = ext.get('commits', 0)
|
||||
description = ext.get("description", "")
|
||||
installed = ext.get("installed", False)
|
||||
enabled = ext.get("enabled", False)
|
||||
path = ext.get("path", "")
|
||||
remote = ext.get("remote", None)
|
||||
commit_date = ext.get("commit_date", 1577836800) or 1577836800
|
||||
update_available = (remote is not None) & (installed) & (datetime.utcfromtimestamp(commit_date + 60 * 60) < datetime.fromisoformat(ext.get('updated', '2000-01-01T00:00:00.000Z')[:-1]))
|
||||
ext['sort_user'] = f"{'0' if ext['is_builtin'] else '1'}{'1' if ext['installed'] else '0'}{ext.get('name', '')}"
|
||||
ext['sort_enabled'] = f"{'0' if ext['enabled'] else '1'}{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}"
|
||||
ext['sort_update'] = f"{'1' if update_available else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}"
|
||||
tags = ext.get("tags", [])
|
||||
tags_string = ' '.join(tags)
|
||||
tags = tags + ["installed"] if installed else tags
|
||||
if len([x for x in tags if x in hide_tags]) > 0:
|
||||
continue
|
||||
if search_text and search_text.strip():
|
||||
if search_text.lower() not in html.escape(name).lower() and search_text.lower() not in html.escape(description).lower() and search_text.lower() not in html.escape(tags_string).lower():
|
||||
continue
|
||||
version_code = ''
|
||||
type_code = ''
|
||||
install_code = ''
|
||||
enabled_code = ''
|
||||
if installed:
|
||||
type_code = f"""<div class="type">{"SYSTEM" if ext['is_builtin'] else 'USER'}</div>"""
|
||||
version_code = f"""<div class="version" style="background: {"--input-border-color-focus" if update_available else "inherit"}">{ext['version']}</div>"""
|
||||
enabled_code = f"""<input class="gr-check-radio gr-checkbox" name="enable_{html.escape(name)}" type="checkbox" {'checked="checked"' if enabled else ''}>"""
|
||||
masked_path = html.escape(path.replace('\\', '/'))
|
||||
if not ext['is_builtin']:
|
||||
install_code = f"""<button onclick="uninstall_extension(this, '{masked_path}')" class="lg secondary gradio-button custom-button extension-button">uninstall</button>"""
|
||||
if update_available:
|
||||
install_code += f"""<button onclick="update_extension(this, '{masked_path}')" class="lg secondary gradio-button custom-button extension-button">update</button>"""
|
||||
else:
|
||||
install_code = f"""<button onclick="install_extension(this, '{html.escape(url)}')" class="lg secondary gradio-button custom-button extension-button">install</button>"""
|
||||
tags_text = ", ".join([f"<span class='extension-tag'>{x}</span>" for x in tags])
|
||||
code += f"""
|
||||
<tr>
|
||||
<td{' class="extension_status"' if ext['installed'] else ''}>{enabled_code}</td>
|
||||
<td><a href="{html.escape(url)}" target="_blank" class="name">{html.escape(name)}</a><br>{tags_text}</td>
|
||||
<td>{html.escape(description)}
|
||||
<p class="info"><span class="date">Created {html.escape(created)} | Added {html.escape(added)} | Pushed {html.escape(pushed)} | Updated {html.escape(updated)}</span></p>
|
||||
<p class="info"><span class="date">Stars {html.escape(str(stars))} | Size {html.escape(str(size))} | Commits {html.escape(str(commits))} | Issues {html.escape(str(issues))}</span></p>
|
||||
</td>
|
||||
<td>{type_code}</td>
|
||||
<td>{version_code}</td>
|
||||
<td>{install_code}</td>
|
||||
</tr>"""
|
||||
code += "</tbody></table>"
|
||||
return code
|
||||
|
||||
|
||||
def create_ui():
|
||||
import modules.ui
|
||||
with gr.Blocks(analytics_enabled=False) as ui:
|
||||
extensions_disable_all = gr.Radio(label="Disable all extensions", choices=["none", "user", "all"], value=shared.opts.disable_all_extensions, elem_id="extensions_disable_all", visible=False)
|
||||
extensions_disabled_list = gr.Text(elem_id="extensions_disabled_list", visible=False).style(container=False)
|
||||
extensions_update_list = gr.Text(elem_id="extensions_update_list", visible=False).style(container=False)
|
||||
with gr.Tabs(elem_id="tabs_extensions"):
|
||||
with gr.TabItem("Manage Extensions", id="manage"):
|
||||
with gr.Row(elem_id="extensions_installed_top"):
|
||||
extension_to_install = gr.Text(elem_id="extension_to_install", visible=False)
|
||||
install_extension_button = gr.Button(elem_id="install_extension_button", visible=False)
|
||||
uninstall_extension_button = gr.Button(elem_id="uninstall_extension_button", visible=False)
|
||||
update_extension_button = gr.Button(elem_id="update_extension_button", visible=False)
|
||||
with gr.Column(scale=4):
|
||||
search_text = gr.Text(label="Search")
|
||||
info = gr.HTML('Note: After any operation such as install/uninstall or enable/disable, please restart the server')
|
||||
with gr.Column(scale=1):
|
||||
sort_column = gr.Dropdown(value="default", label="Sort by", choices=list(sort_ordering.keys()), multiselect=False)
|
||||
with gr.Column(scale=1):
|
||||
refresh_extensions_button = gr.Button(value="Refresh extension list", variant="primary")
|
||||
check = gr.Button(value="Update installed extensions", variant="primary")
|
||||
apply = gr.Button(value="Apply changes & restart server", variant="primary")
|
||||
update_extension_list()
|
||||
extensions_table = gr.HTML(refresh_extensions_list_from_data(search_text.value, sort_column.value))
|
||||
check.click(
|
||||
fn=wrap_gradio_gpu_call(check_updates, extra_outputs=[gr.update()]),
|
||||
_js="extensions_check",
|
||||
inputs=[info, extensions_disabled_list, search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
apply.click(
|
||||
fn=apply_and_restart,
|
||||
_js="extensions_apply",
|
||||
inputs=[extensions_disabled_list, extensions_update_list, extensions_disable_all],
|
||||
outputs=[],
|
||||
)
|
||||
refresh_extensions_button.click(
|
||||
fn=modules.ui.wrap_gradio_call(refresh_extensions_list, extra_outputs=[gr.update(), gr.update()]),
|
||||
inputs=[search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
install_extension_button.click(
|
||||
fn=modules.ui.wrap_gradio_call(install_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]),
|
||||
inputs=[extension_to_install, search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
uninstall_extension_button.click(
|
||||
fn=modules.ui.wrap_gradio_call(uninstall_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]),
|
||||
inputs=[extension_to_install, search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
update_extension_button.click(
|
||||
fn=modules.ui.wrap_gradio_call(update_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]),
|
||||
inputs=[extension_to_install, search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
search_text.change(
|
||||
fn=modules.ui.wrap_gradio_call(search_extensions, extra_outputs=[gr.update(), gr.update()]),
|
||||
inputs=[search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
sort_column.change(
|
||||
fn=modules.ui.wrap_gradio_call(search_extensions, extra_outputs=[gr.update(), gr.update()]),
|
||||
inputs=[search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
with gr.TabItem("Manual install", id="install_from_url"):
|
||||
install_url = gr.Text(label="URL for extension's git repository")
|
||||
install_branch = gr.Text(label="Specific branch name", placeholder="Leave empty for default main branch")
|
||||
install_dirname = gr.Text(label="Local directory name", placeholder="Leave empty for auto")
|
||||
install_button = gr.Button(value="Install", variant="primary")
|
||||
info = gr.HTML(elem_id="extension_info")
|
||||
install_button.click(
|
||||
fn=modules.ui.wrap_gradio_call(install_extension_from_url, extra_outputs=[gr.update()]),
|
||||
inputs=[install_dirname, install_url, install_branch, search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
return ui
|
||||
import json
|
||||
import os.path
|
||||
import shutil
|
||||
import errno
|
||||
import html
|
||||
from datetime import datetime
|
||||
import git
|
||||
import gradio as gr
|
||||
from modules import extensions, shared, paths, errors
|
||||
from modules.call_queue import wrap_gradio_gpu_call
|
||||
|
||||
|
||||
extensions_index = "https://vladmandic.github.io/sd-data/pages/extensions.json"
|
||||
hide_tags = ["localization"]
|
||||
extensions_list = []
|
||||
sort_ordering = {
|
||||
"default": (True, lambda x: x.get('sort_default', '')),
|
||||
"user extensions": (True, lambda x: x.get('sort_user', '')),
|
||||
"update avilable": (True, lambda x: x.get('sort_update', '')),
|
||||
"updated date": (True, lambda x: x.get('updated', '2000-01-01T00:00')),
|
||||
"created date": (True, lambda x: x.get('created', '2000-01-01T00:00')),
|
||||
"name": (False, lambda x: x.get('name', '').lower()),
|
||||
"enabled": (False, lambda x: x.get('sort_enabled', '').lower()),
|
||||
"size": (True, lambda x: x.get('size', 0)),
|
||||
"stars": (True, lambda x: x.get('stars', 0)),
|
||||
"commits": (True, lambda x: x.get('commits', 0)),
|
||||
"issues": (True, lambda x: x.get('issues', 0)),
|
||||
}
|
||||
|
||||
|
||||
def update_extension_list():
|
||||
global extensions_list # pylint: disable=global-statement
|
||||
try:
|
||||
with open(os.path.join(paths.script_path, "html", "extensions.json"), "r", encoding="utf-8") as f:
|
||||
extensions_list = json.loads(f.read())
|
||||
shared.log.debug(f'Extensions list loaded: {os.path.join(paths.script_path, "html", "extensions.json")}')
|
||||
except:
|
||||
shared.log.debug(f'Extensions list failed to load: {os.path.join(paths.script_path, "html", "extensions.json")}')
|
||||
found = []
|
||||
for ext in extensions.extensions:
|
||||
ext.read_info_from_repo()
|
||||
for ext in extensions_list:
|
||||
installed = [extension for extension in extensions.extensions
|
||||
if extension.git_name == ext['name']
|
||||
or extension.name == ext['name']
|
||||
or (extension.remote or '').startswith(ext['url'].replace('.git', ''))]
|
||||
if len(installed) > 0:
|
||||
found.append(installed[0])
|
||||
not_matched = [extension for extension in extensions.extensions if extension not in found]
|
||||
for ext in not_matched:
|
||||
entry = {
|
||||
"name": ext.name or "",
|
||||
"description": ext.description or "",
|
||||
"url": ext.remote or "",
|
||||
"tags": [],
|
||||
"stars": 0,
|
||||
"issues": 0,
|
||||
"commits": 0,
|
||||
"size": 0,
|
||||
"long": ext.git_name or ext.name or "",
|
||||
"added": ext.ctime,
|
||||
"created": ext.ctime,
|
||||
"updated": ext.mtime,
|
||||
}
|
||||
extensions_list.append(entry)
|
||||
|
||||
|
||||
def check_access():
|
||||
assert not shared.cmd_opts.disable_extension_access, "extension access disabled because of command line flags"
|
||||
|
||||
|
||||
def apply_and_restart(disable_list, update_list, disable_all):
|
||||
check_access()
|
||||
shared.log.debug(f'Extensions apply: disable={disable_list} update={update_list}')
|
||||
disabled = json.loads(disable_list)
|
||||
assert type(disabled) == list, f"wrong disable_list data for apply_and_restart: {disable_list}"
|
||||
update = json.loads(update_list)
|
||||
assert type(update) == list, f"wrong update_list data for apply_and_restart: {update_list}"
|
||||
update = set(update)
|
||||
for ext in extensions.extensions:
|
||||
if ext.name not in update:
|
||||
continue
|
||||
try:
|
||||
ext.fetch_and_reset_hard()
|
||||
except Exception as e:
|
||||
errors.display(e, f'extensions apply update: {ext.name}')
|
||||
shared.opts.disabled_extensions = disabled
|
||||
shared.opts.disable_all_extensions = disable_all
|
||||
shared.opts.save(shared.config_filename)
|
||||
shared.restart_server(restart=True)
|
||||
|
||||
|
||||
def check_updates(_id_task, disable_list, search_text, sort_column):
|
||||
check_access()
|
||||
disabled = json.loads(disable_list)
|
||||
assert type(disabled) == list, f"wrong disable_list data for apply_and_restart: {disable_list}"
|
||||
exts = [ext for ext in extensions.extensions if ext.remote is not None and ext.name not in disabled]
|
||||
shared.log.info(f'Extensions update check: update={len(exts)} disabled={len(disable_list)}')
|
||||
shared.state.job_count = len(exts)
|
||||
for ext in exts:
|
||||
shared.state.textinfo = ext.name
|
||||
try:
|
||||
ext.check_updates()
|
||||
if ext.can_update:
|
||||
ext.fetch_and_reset_hard()
|
||||
ext.read_info_from_repo()
|
||||
commit_date = ext.commit_date or 1577836800
|
||||
shared.log.info(f'Extensions updated: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}')
|
||||
else:
|
||||
commit_date = ext.commit_date or 1577836800
|
||||
shared.log.debug(f'Extensions no update available: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}')
|
||||
except FileNotFoundError as e:
|
||||
if 'FETCH_HEAD' not in str(e):
|
||||
raise
|
||||
except Exception:
|
||||
errors.display(e, f'extensions check update: {ext.name}')
|
||||
shared.state.nextjob()
|
||||
return refresh_extensions_list_from_data(search_text, sort_column), "Extension update complete | Restart required"
|
||||
|
||||
|
||||
def make_commit_link(commit_hash, remote, text=None):
|
||||
if text is None:
|
||||
text = commit_hash[:8]
|
||||
if remote.startswith("https://github.com/"):
|
||||
href = os.path.join(remote, "commit", commit_hash)
|
||||
return f'<a href="{href}" target="_blank">{text}</a>'
|
||||
else:
|
||||
return text
|
||||
|
||||
|
||||
def normalize_git_url(url):
|
||||
if url is None:
|
||||
return ""
|
||||
url = url.replace(".git", "")
|
||||
return url
|
||||
|
||||
|
||||
def install_extension_from_url(dirname, url, branch_name, search_text, sort_column):
|
||||
check_access()
|
||||
assert url, 'No URL specified'
|
||||
if dirname is None or dirname == "":
|
||||
*parts, last_part = url.split('/') # pylint: disable=unused-variable
|
||||
last_part = normalize_git_url(last_part)
|
||||
dirname = last_part
|
||||
target_dir = os.path.join(extensions.extensions_dir, dirname)
|
||||
shared.log.info(f'Installing extension: {url} into {target_dir}')
|
||||
assert not os.path.exists(target_dir), f'Extension directory already exists: {target_dir}'
|
||||
normalized_url = normalize_git_url(url)
|
||||
assert len([x for x in extensions.extensions if normalize_git_url(x.remote) == normalized_url]) == 0, 'Extension with this URL is already installed'
|
||||
tmpdir = os.path.join(paths.data_path, "tmp", dirname)
|
||||
try:
|
||||
shutil.rmtree(tmpdir, True)
|
||||
if not branch_name:
|
||||
# if no branch is specified, use the default branch
|
||||
with git.Repo.clone_from(url, tmpdir) as repo:
|
||||
repo.remote().fetch()
|
||||
for submodule in repo.submodules:
|
||||
submodule.update()
|
||||
else:
|
||||
with git.Repo.clone_from(url, tmpdir, branch=branch_name) as repo:
|
||||
repo.remote().fetch()
|
||||
for submodule in repo.submodules:
|
||||
submodule.update()
|
||||
try:
|
||||
os.rename(tmpdir, target_dir)
|
||||
except OSError as err:
|
||||
if err.errno == errno.EXDEV:
|
||||
shutil.move(tmpdir, target_dir)
|
||||
else:
|
||||
raise err
|
||||
from launch import run_extension_installer
|
||||
run_extension_installer(target_dir)
|
||||
extensions.list_extensions()
|
||||
return [refresh_extensions_list_from_data(search_text, sort_column), html.escape(f"Extension installed: {target_dir} | Restart required")]
|
||||
finally:
|
||||
shutil.rmtree(tmpdir, True)
|
||||
|
||||
|
||||
def install_extension(extension_to_install, search_text, sort_column):
|
||||
shared.log.info(f'Extension install: {extension_to_install}')
|
||||
code, message = install_extension_from_url(None, extension_to_install, None, search_text, sort_column)
|
||||
return code, message
|
||||
|
||||
|
||||
def uninstall_extension(extension_path, search_text, sort_column):
|
||||
def errorRemoveReadonly(func, path, exc):
|
||||
import stat
|
||||
excvalue = exc[1]
|
||||
shared.log.debug(f'Exception during cleanup: {func} {path} {excvalue.strerror}')
|
||||
if func in (os.rmdir, os.remove, os.unlink) and excvalue.errno == errno.EACCES:
|
||||
shared.log.debug(f'Retrying cleanup: {path}')
|
||||
os.chmod(path, stat.S_IRWXU | stat.S_IRWXG | stat.S_IRWXO)
|
||||
func(path)
|
||||
|
||||
ext = [extension for extension in extensions.extensions if os.path.abspath(extension.path) == os.path.abspath(extension_path)]
|
||||
if len(ext) > 0 and os.path.isdir(extension_path):
|
||||
found = ext[0]
|
||||
try:
|
||||
shutil.rmtree(found.path, ignore_errors=False, onerror=errorRemoveReadonly)
|
||||
except Exception as e:
|
||||
shared.log.warning(f'Extension uninstall failed: {found.path} {e}')
|
||||
extensions.extensions = [extension for extension in extensions.extensions if os.path.abspath(found.path) != os.path.abspath(extension_path)]
|
||||
update_extension_list()
|
||||
code = refresh_extensions_list_from_data(search_text, sort_column)
|
||||
shared.log.info(f'Extension uninstalled: {found.path}')
|
||||
return code, f"Extension uninstalled: {found.path} | Restart required"
|
||||
else:
|
||||
shared.log.warning(f'Extension uninstall cannot find extension: {extension_path}')
|
||||
code = refresh_extensions_list_from_data(search_text, sort_column)
|
||||
return code, f"Extension uninstalled failed: {extension_path}"
|
||||
|
||||
|
||||
def update_extension(extension_path, search_text, sort_column):
|
||||
exts = [extension for extension in extensions.extensions if os.path.abspath(extension.path) == os.path.abspath(extension_path)]
|
||||
shared.state.job_count = len(exts)
|
||||
for ext in exts:
|
||||
shared.log.debug(f'Extensions update start: {ext.name} {ext.commit_hash} {ext.commit_date}')
|
||||
shared.state.textinfo = ext.name
|
||||
try:
|
||||
ext.check_updates()
|
||||
if ext.can_update:
|
||||
ext.fetch_and_reset_hard()
|
||||
ext.read_info_from_repo()
|
||||
commit_date = ext.commit_date or 1577836800
|
||||
shared.log.info(f'Extensions updated: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}')
|
||||
else:
|
||||
commit_date = ext.commit_date or 1577836800
|
||||
shared.log.info(f'Extensions no update available: {ext.name} {ext.commit_hash[:8]} {datetime.utcfromtimestamp(commit_date)}')
|
||||
except FileNotFoundError as e:
|
||||
if 'FETCH_HEAD' not in str(e):
|
||||
raise
|
||||
except Exception as e:
|
||||
shared.log.error(f'Extensions update failed: {ext.name}')
|
||||
errors.display(e, f'extensions check update: {ext.name}')
|
||||
shared.log.debug(f'Extensions update finish: {ext.name} {ext.commit_hash} {ext.commit_date}')
|
||||
shared.state.nextjob()
|
||||
return refresh_extensions_list_from_data(search_text, sort_column), f"Extension updated | {extension_path} | Restart required"
|
||||
|
||||
|
||||
def refresh_extensions_list(search_text, sort_column):
|
||||
global extensions_list # pylint: disable=global-statement
|
||||
import urllib.request
|
||||
try:
|
||||
with urllib.request.urlopen(extensions_index) as response:
|
||||
text = response.read()
|
||||
extensions_list = json.loads(text)
|
||||
with open(os.path.join(paths.script_path, "html", "extensions.json"), "w", encoding="utf-8") as outfile:
|
||||
json_object = json.dumps(extensions_list, indent=2)
|
||||
outfile.write(json_object)
|
||||
shared.log.debug(f'Updated extensions list: {len(extensions_list)} {extensions_index} {outfile}')
|
||||
except Exception as e:
|
||||
shared.log.warning(f'Updated extensions list failed: {extensions_index} {e}')
|
||||
update_extension_list()
|
||||
code = refresh_extensions_list_from_data(search_text, sort_column)
|
||||
return code, f'Extensions | {len(extensions.extensions)} registered | {len(extensions_list)} available'
|
||||
|
||||
|
||||
def search_extensions(search_text, sort_column):
|
||||
code = refresh_extensions_list_from_data(search_text, sort_column)
|
||||
return code, f'Search | {search_text} | {sort_column}'
|
||||
|
||||
|
||||
def refresh_extensions_list_from_data(search_text, sort_column):
|
||||
shared.log.debug(f'Extensions manager: refresh list search="{search_text}" sort="{sort_column}"')
|
||||
code = """
|
||||
<table id="extensions">
|
||||
<colgroup>
|
||||
<col style="width: 1%; background: var(--table-border-color)">
|
||||
<col style="width: 20%; background: var(--table-border-color)">
|
||||
<col style="width: 59%;">
|
||||
<col style="width: 5%; background: var(--panel-background-fill)">
|
||||
<col style="width: 10%; background: var(--panel-background-fill)">
|
||||
<col style="width: 5%; background: var(--table-border-color)">
|
||||
</colgroup>
|
||||
<thead style="font-size: 110%; border-style: solid; border-bottom: 1px var(--button-primary-border-color) solid">
|
||||
<tr>
|
||||
<th>Enabled</th>
|
||||
<th>Extension</th>
|
||||
<th>Description</th>
|
||||
<th>Type</th>
|
||||
<th>Current version</th>
|
||||
<th></th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>"""
|
||||
for ext in extensions_list:
|
||||
extension = [extension for extension in extensions.extensions if extension.git_name == ext['name'] or extension.name == ext['name']]
|
||||
if len(extension) > 0:
|
||||
extension[0].read_info_from_repo()
|
||||
ext['installed'] = len(extension) > 0
|
||||
ext['commit_date'] = extension[0].commit_date if len(extension) > 0 else 1577836800
|
||||
ext['is_builtin'] = extension[0].is_builtin if len(extension) > 0 else False
|
||||
ext['version'] = extension[0].version if len(extension) > 0 else ''
|
||||
ext['enabled'] = extension[0].enabled if len(extension) > 0 else ''
|
||||
ext['remote'] = extension[0].remote if len(extension) > 0 else None
|
||||
ext['path'] = extension[0].path if len(extension) > 0 else ''
|
||||
ext['sort_default'] = f"{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}"
|
||||
sort_reverse, sort_function = sort_ordering[sort_column]
|
||||
|
||||
def dt(x: str):
|
||||
val = ext.get(x, None)
|
||||
if val is not None:
|
||||
return datetime.fromisoformat(val[:-1]).strftime('%a %b%d %Y %H:%M')
|
||||
else:
|
||||
return "N/A"
|
||||
|
||||
for ext in sorted(extensions_list, key=sort_function, reverse=sort_reverse):
|
||||
name = ext.get("name", "unknown")
|
||||
added = dt('added')
|
||||
created = dt('created')
|
||||
pushed = dt('pushed')
|
||||
updated = dt('updated')
|
||||
url = ext.get('url', None)
|
||||
size = ext.get('size', 0)
|
||||
stars = ext.get('stars', 0)
|
||||
issues = ext.get('issues', 0)
|
||||
commits = ext.get('commits', 0)
|
||||
description = ext.get("description", "")
|
||||
installed = ext.get("installed", False)
|
||||
enabled = ext.get("enabled", False)
|
||||
path = ext.get("path", "")
|
||||
remote = ext.get("remote", None)
|
||||
commit_date = ext.get("commit_date", 1577836800) or 1577836800
|
||||
update_available = (remote is not None) & (installed) & (datetime.utcfromtimestamp(commit_date + 60 * 60) < datetime.fromisoformat(ext.get('updated', '2000-01-01T00:00:00.000Z')[:-1]))
|
||||
ext['sort_user'] = f"{'0' if ext['is_builtin'] else '1'}{'1' if ext['installed'] else '0'}{ext.get('name', '')}"
|
||||
ext['sort_enabled'] = f"{'0' if ext['enabled'] else '1'}{'1' if ext['is_builtin'] else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}"
|
||||
ext['sort_update'] = f"{'1' if update_available else '0'}{'1' if ext['installed'] else '0'}{ext.get('updated', '2000-01-01T00:00')}"
|
||||
tags = ext.get("tags", [])
|
||||
tags_string = ' '.join(tags)
|
||||
tags = tags + ["installed"] if installed else tags
|
||||
if len([x for x in tags if x in hide_tags]) > 0:
|
||||
continue
|
||||
if search_text and search_text.strip():
|
||||
if search_text.lower() not in html.escape(name).lower() and search_text.lower() not in html.escape(description).lower() and search_text.lower() not in html.escape(tags_string).lower():
|
||||
continue
|
||||
version_code = ''
|
||||
type_code = ''
|
||||
install_code = ''
|
||||
enabled_code = ''
|
||||
if installed:
|
||||
type_code = f"""<div class="type">{"SYSTEM" if ext['is_builtin'] else 'USER'}</div>"""
|
||||
version_code = f"""<div class="version" style="background: {"--input-border-color-focus" if update_available else "inherit"}">{ext['version']}</div>"""
|
||||
enabled_code = f"""<input class="gr-check-radio gr-checkbox" name="enable_{html.escape(name)}" type="checkbox" {'checked="checked"' if enabled else ''}>"""
|
||||
masked_path = html.escape(path.replace('\\', '/'))
|
||||
if not ext['is_builtin']:
|
||||
install_code = f"""<button onclick="uninstall_extension(this, '{masked_path}')" class="lg secondary gradio-button custom-button extension-button">uninstall</button>"""
|
||||
if update_available:
|
||||
install_code += f"""<button onclick="update_extension(this, '{masked_path}')" class="lg secondary gradio-button custom-button extension-button">update</button>"""
|
||||
else:
|
||||
install_code = f"""<button onclick="install_extension(this, '{html.escape(url)}')" class="lg secondary gradio-button custom-button extension-button">install</button>"""
|
||||
tags_text = ", ".join([f"<span class='extension-tag'>{x}</span>" for x in tags])
|
||||
code += f"""
|
||||
<tr>
|
||||
<td{' class="extension_status"' if ext['installed'] else ''}>{enabled_code}</td>
|
||||
<td><a href="{html.escape(url)}" target="_blank" class="name">{html.escape(name)}</a><br>{tags_text}</td>
|
||||
<td>{html.escape(description)}
|
||||
<p class="info"><span class="date">Created {html.escape(created)} | Added {html.escape(added)} | Pushed {html.escape(pushed)} | Updated {html.escape(updated)}</span></p>
|
||||
<p class="info"><span class="date">Stars {html.escape(str(stars))} | Size {html.escape(str(size))} | Commits {html.escape(str(commits))} | Issues {html.escape(str(issues))}</span></p>
|
||||
</td>
|
||||
<td>{type_code}</td>
|
||||
<td>{version_code}</td>
|
||||
<td>{install_code}</td>
|
||||
</tr>"""
|
||||
code += "</tbody></table>"
|
||||
return code
|
||||
|
||||
|
||||
def create_ui():
|
||||
import modules.ui
|
||||
with gr.Blocks(analytics_enabled=False) as ui:
|
||||
extensions_disable_all = gr.Radio(label="Disable all extensions", choices=["none", "user", "all"], value=shared.opts.disable_all_extensions, elem_id="extensions_disable_all", visible=False)
|
||||
extensions_disabled_list = gr.Text(elem_id="extensions_disabled_list", visible=False).style(container=False)
|
||||
extensions_update_list = gr.Text(elem_id="extensions_update_list", visible=False).style(container=False)
|
||||
with gr.Tabs(elem_id="tabs_extensions"):
|
||||
with gr.TabItem("Manage Extensions", id="manage"):
|
||||
with gr.Row(elem_id="extensions_installed_top"):
|
||||
extension_to_install = gr.Text(elem_id="extension_to_install", visible=False)
|
||||
install_extension_button = gr.Button(elem_id="install_extension_button", visible=False)
|
||||
uninstall_extension_button = gr.Button(elem_id="uninstall_extension_button", visible=False)
|
||||
update_extension_button = gr.Button(elem_id="update_extension_button", visible=False)
|
||||
with gr.Column(scale=4):
|
||||
search_text = gr.Text(label="Search")
|
||||
info = gr.HTML('Note: After any operation such as install/uninstall or enable/disable, please restart the server')
|
||||
with gr.Column(scale=1):
|
||||
sort_column = gr.Dropdown(value="default", label="Sort by", choices=list(sort_ordering.keys()), multiselect=False)
|
||||
with gr.Column(scale=1):
|
||||
refresh_extensions_button = gr.Button(value="Refresh extension list", variant="primary")
|
||||
check = gr.Button(value="Update installed extensions", variant="primary")
|
||||
apply = gr.Button(value="Apply changes & restart server", variant="primary")
|
||||
update_extension_list()
|
||||
extensions_table = gr.HTML(refresh_extensions_list_from_data(search_text.value, sort_column.value))
|
||||
check.click(
|
||||
fn=wrap_gradio_gpu_call(check_updates, extra_outputs=[gr.update()]),
|
||||
_js="extensions_check",
|
||||
inputs=[info, extensions_disabled_list, search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
apply.click(
|
||||
fn=apply_and_restart,
|
||||
_js="extensions_apply",
|
||||
inputs=[extensions_disabled_list, extensions_update_list, extensions_disable_all],
|
||||
outputs=[],
|
||||
)
|
||||
refresh_extensions_button.click(
|
||||
fn=modules.ui.wrap_gradio_call(refresh_extensions_list, extra_outputs=[gr.update(), gr.update()]),
|
||||
inputs=[search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
install_extension_button.click(
|
||||
fn=modules.ui.wrap_gradio_call(install_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]),
|
||||
inputs=[extension_to_install, search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
uninstall_extension_button.click(
|
||||
fn=modules.ui.wrap_gradio_call(uninstall_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]),
|
||||
inputs=[extension_to_install, search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
update_extension_button.click(
|
||||
fn=modules.ui.wrap_gradio_call(update_extension, extra_outputs=[gr.update(), gr.update(), gr.update()]),
|
||||
inputs=[extension_to_install, search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
search_text.change(
|
||||
fn=modules.ui.wrap_gradio_call(search_extensions, extra_outputs=[gr.update(), gr.update()]),
|
||||
inputs=[search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
sort_column.change(
|
||||
fn=modules.ui.wrap_gradio_call(search_extensions, extra_outputs=[gr.update(), gr.update()]),
|
||||
inputs=[search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
with gr.TabItem("Manual install", id="install_from_url"):
|
||||
install_url = gr.Text(label="URL for extension's git repository")
|
||||
install_branch = gr.Text(label="Specific branch name", placeholder="Leave empty for default main branch")
|
||||
install_dirname = gr.Text(label="Local directory name", placeholder="Leave empty for auto")
|
||||
install_button = gr.Button(value="Install", variant="primary")
|
||||
info = gr.HTML(elem_id="extension_info")
|
||||
install_button.click(
|
||||
fn=modules.ui.wrap_gradio_call(install_extension_from_url, extra_outputs=[gr.update()]),
|
||||
inputs=[install_dirname, install_url, install_branch, search_text, sort_column],
|
||||
outputs=[extensions_table, info],
|
||||
)
|
||||
return ui
|
||||
|
||||
@@ -12,7 +12,7 @@ class ExtraNetworksPageHypernetworks(ui_extra_networks.ExtraNetworksPage):
|
||||
|
||||
def list_items(self):
|
||||
for name, path in shared.hypernetworks.items():
|
||||
path, ext = os.path.splitext(path)
|
||||
path, _ext = os.path.splitext(path)
|
||||
yield {
|
||||
"name": name,
|
||||
"filename": path,
|
||||
|
||||
@@ -3,13 +3,18 @@ from modules import scripts_postprocessing, scripts, shared, gfpgan_model, codef
|
||||
import modules.generation_parameters_copypaste as parameters_copypaste
|
||||
from modules.call_queue import wrap_gradio_gpu_call, wrap_queued_call, wrap_gradio_call # pylint: disable=unused-import
|
||||
from modules.extras import run_pnginfo
|
||||
from modules.ui_common import infotext_to_html
|
||||
|
||||
|
||||
def wrap_pnginfo(image):
|
||||
_, geninfo, info = run_pnginfo(image)
|
||||
return '', infotext_to_html(geninfo), info, geninfo
|
||||
|
||||
|
||||
def submit_click(tab_index, extras_image, image_batch, extras_batch_input_dir, extras_batch_output_dir, show_extras_results, *script_inputs):
|
||||
result_images, html_info_x, html_info = postprocessing.run_postprocessing(tab_index, extras_image, image_batch, extras_batch_input_dir, extras_batch_output_dir, show_extras_results, *script_inputs)
|
||||
if result_images is not None and len(result_images) > 0:
|
||||
_html_info, _generation_info, html_info_x = run_pnginfo(result_images[0])
|
||||
return result_images, html_info_x, html_info
|
||||
|
||||
result_images, geninfo, _js_info = postprocessing.run_postprocessing(tab_index, extras_image, image_batch, extras_batch_input_dir, extras_batch_output_dir, show_extras_results, *script_inputs)
|
||||
return result_images, geninfo, '{}', ''
|
||||
|
||||
|
||||
def create_ui():
|
||||
@@ -37,32 +42,21 @@ def create_ui():
|
||||
skip = gr.Button('Skip', elem_id=f"{id_part}_skip", variant='secondary')
|
||||
skip.click(fn=lambda: shared.state.skip(), inputs=[], outputs=[])
|
||||
interrupt.click(fn=lambda: shared.state.interrupt(), inputs=[], outputs=[])
|
||||
result_images, html_info_x, html_info, _html_log = ui_common.create_output_panel("extras", shared.opts.outdir_extras_samples)
|
||||
html_info = gr.HTML(elem_id="pnginfo_html_info")
|
||||
generation_info = gr.Textbox(elem_id="pnginfo_generation_info", label="Parameters", visible=False)
|
||||
generation_info_pretty = gr.Textbox(elem_id="pnginfo_generation_info_pretty", label="Parameters")
|
||||
gr.HTML('Full metadata')
|
||||
html2_info = gr.HTML(elem_id="pnginfo_html2_info")
|
||||
result_images, generation_info, html_info, html_log = ui_common.create_output_panel("extras", shared.opts.outdir_extras_samples)
|
||||
gr.HTML('File metadata')
|
||||
exif_info = gr.HTML(elem_id="pnginfo_html_info")
|
||||
gen_info = gr.Text(elem_id="pnginfo_gen_info", visible=False)
|
||||
for tabname, button in buttons.items():
|
||||
parameters_copypaste.register_paste_params_button(parameters_copypaste.ParamBinding(paste_button=button, tabname=tabname, source_text_component=generation_info, source_image_component=extras_image))
|
||||
|
||||
def pretty_geninfo(generation_info: str):
|
||||
if generation_info is None:
|
||||
return ''
|
||||
sections = generation_info.split('Steps:')
|
||||
if len(sections) > 1:
|
||||
param = sections[0].strip() + '\nSteps:' + sections[1].strip().replace(', ', '\n')
|
||||
return param
|
||||
return generation_info
|
||||
parameters_copypaste.register_paste_params_button(parameters_copypaste.ParamBinding(paste_button=button, tabname=tabname, source_text_component=gen_info, source_image_component=extras_image))
|
||||
|
||||
tab_single.select(fn=lambda: 0, inputs=[], outputs=[tab_index])
|
||||
tab_batch.select(fn=lambda: 1, inputs=[], outputs=[tab_index])
|
||||
tab_batch_dir.select(fn=lambda: 2, inputs=[], outputs=[tab_index])
|
||||
generation_info.change(fn=pretty_geninfo, inputs=[generation_info], outputs=[generation_info_pretty])
|
||||
_dummy = gr.HTML(visible=False)
|
||||
extras_image.change(
|
||||
fn=wrap_gradio_call(run_pnginfo),
|
||||
fn=wrap_gradio_call(wrap_pnginfo),
|
||||
inputs=[extras_image],
|
||||
outputs=[html_info, generation_info, html2_info],
|
||||
outputs=[_dummy, html_info, exif_info, gen_info],
|
||||
)
|
||||
submit.click(
|
||||
fn=call_queue.wrap_gradio_gpu_call(submit_click, extra_outputs=[None, '']),
|
||||
@@ -73,12 +67,13 @@ def create_ui():
|
||||
extras_batch_input_dir,
|
||||
extras_batch_output_dir,
|
||||
show_extras_results,
|
||||
*script_inputs
|
||||
*script_inputs,
|
||||
],
|
||||
outputs=[
|
||||
result_images,
|
||||
html_info_x,
|
||||
html_info,
|
||||
generation_info,
|
||||
html_log,
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user