update extra networks

This commit is contained in:
Vladimir Mandic
2023-05-11 09:30:34 -04:00
parent 99b6133bc9
commit 05656a54fe
7 changed files with 47 additions and 38 deletions
+7 -1
View File
@@ -12,6 +12,7 @@ function setupExtraNetworksForTab(tabname) {
tabs.appendChild(descriptInput);
search.addEventListener('input', (evt) => {
searchTerm = search.value.toLowerCase();
console.log('HERE', searchTerm)
gradioApp().querySelectorAll(`#${tabname}_extra_tabs div.card`).forEach((elem) => {
text = `${elem.querySelector('.name').textContent.toLowerCase()} ${elem.querySelector('.search_term').textContent.toLowerCase()}`;
elem.style.display = text.indexOf(searchTerm) == -1 ? 'none' : '';
@@ -48,7 +49,7 @@ function tryToRemoveExtraNetworkFromPrompt(textarea, text) {
let replaced = false;
const newTextareaText = textarea.value.replaceAll(re_extranet_g, (found, index) => {
m = found.match(re_extranet);
if (m[1] == partToSearch) {
if (m[1] === partToSearch) {
replaced = true;
return '';
}
@@ -61,6 +62,11 @@ function tryToRemoveExtraNetworkFromPrompt(textarea, text) {
return false;
}
function refreshExtraNetworks(tabname) {
console.log('HERE2', tabname);
gradioApp().querySelector(`#${tabname}_extra_networks textarea`)?.dispatchEvent(new Event('input'));
}
function cardClicked(tabname, textToAdd, allowNegativePrompt) {
const textarea = allowNegativePrompt ? activePromptTextarea[tabname] : gradioApp().querySelector(`#${tabname}_prompt > label > textarea`);
if (!tryToRemoveExtraNetworkFromPrompt(textarea, textToAdd)) textarea.value = textarea.value + opts.extra_networks_add_text_separator + textToAdd;
+18 -20
View File
@@ -552,32 +552,36 @@ def save_image(image, path, basename, seed=None, prompt=None, extension='jpg', i
else:
exifinfo_data = params.pnginfo.get(pnginfo_section_name, '')
def atomically_save_image(image_to_save: Image, filename_without_extension: str, extension: str):
# save image with .tmp extension to avoid race condition when another process detects new image in the directory
fn = filename_without_extension + extension
def atomically_save_image(image: Image, basename: str, extension: str):
Image.MAX_IMAGE_PIXELS = None # disable check in Pillow and rely on check below to allow large custom image sizes
mp = round(image.width * image.height / 1000000)
if mp > shared.opts.img_max_size_mp:
shared.log.warning(f'Image size: {image.size} excedes {shared.opts.img_max_size_mp} MPixels')
fn = basename + extension
image_format = Image.registered_extensions()[extension]
log.debug(f'Saving image: {image_format} {fn}')
log.debug(f'Saving image: {image_format} {fn} {image.size}')
if image_format == 'PNG':
pnginfo_data = PngImagePlugin.PngInfo()
for k, v in params.pnginfo.items():
pnginfo_data.add_text(k, str(v))
image_to_save.save(fn, format=image_format, quality=opts.jpeg_quality, pnginfo=pnginfo_data)
image.save(fn, format=image_format, quality=opts.jpeg_quality, pnginfo=pnginfo_data)
elif image_format == 'JPEG':
if image_to_save.mode == 'RGBA':
if image.mode == 'RGBA':
shared.log.warning('Saving RGBA image as JPEG: Alpha channel will be lost')
image_to_save = image_to_save.convert("RGB")
elif image_to_save.mode == 'I;16':
image_to_save = image_to_save.point(lambda p: p * 0.0038910505836576).convert("L")
image = image.convert("RGB")
elif image.mode == 'I;16':
image = image.point(lambda p: p * 0.0038910505836576).convert("L")
exif_bytes = piexif.dump({ "Exif": { piexif.ExifIFD.UserComment: piexif.helper.UserComment.dump(exifinfo_data or "", encoding="unicode") } })
image_to_save.save(fn, format=image_format, quality=opts.jpeg_quality, exif=exif_bytes)
image.save(fn, format=image_format, quality=opts.jpeg_quality, exif=exif_bytes)
elif image_format == 'WEBP':
if image_to_save.mode == 'I;16':
image_to_save = image_to_save.point(lambda p: p * 0.0038910505836576).convert("RGB")
if image.mode == 'I;16':
image = image.point(lambda p: p * 0.0038910505836576).convert("RGB")
exif_bytes = piexif.dump({ "Exif": { piexif.ExifIFD.UserComment: piexif.helper.UserComment.dump(exifinfo_data or "", encoding="unicode") } })
image_to_save.save(fn, format=image_format, quality=opts.jpeg_quality, lossless=opts.webp_lossless, exif=exif_bytes)
image.save(fn, format=image_format, quality=opts.jpeg_quality, lossless=opts.webp_lossless, exif=exif_bytes)
else:
shared.log.warning(f'Unrecognized image format: {extension} attempting save as {image_format}')
image_to_save.save(fn, format=image_format, quality=opts.jpeg_quality)
image.save(fn, format=image_format, quality=opts.jpeg_quality)
filename, extension = os.path.splitext(params.filename)
if hasattr(os, 'statvfs'):
@@ -660,31 +664,25 @@ Steps: {json_info["steps"]}, Sampler: {sampler}, CFG scale: {json_info["scale"]}
def image_data(data):
import gradio as gr
try:
image = Image.open(io.BytesIO(data))
textinfo, _ = read_info_from_image(image)
return textinfo, None
except Exception:
pass
try:
text = data.decode('utf8')
assert len(text) < 10000
return text, None
except Exception:
pass
return gr.update(), None
def flatten(img, bgcolor):
"""replaces transparency with bgcolor (example: "#ffffff"), returning an RGB mode image with no transparency"""
if img.mode == "RGBA":
background = Image.new('RGBA', img.size, bgcolor)
background.paste(img, mask=img)
img = background
return img.convert('RGB')
+13 -13
View File
@@ -356,10 +356,7 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None)
current_checkpoint_info = None
if shared.sd_model:
current_checkpoint_info = shared.sd_model.sd_checkpoint_info
sd_hijack.model_hijack.undo_hijack(shared.sd_model)
shared.sd_model = None
devices.torch_gc()
shared.log.debug(f'Model unloaded: {memory_stats()}')
unload_model_weights()
do_inpainting_hijack()
devices.set_cuda_params()
if already_loaded_state_dict is not None:
@@ -374,18 +371,19 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None)
load_model(current_checkpoint_info, None)
return
shared.log.debug(f'Model dict loaded: {memory_stats()}')
clip_is_included_into_sd = sd1_clip_weight in state_dict or sd2_clip_weight in state_dict
sd_config = OmegaConf.load(checkpoint_config)
repair_config(sd_config)
timer.record("config")
shared.log.debug(f'Model config loaded: {memory_stats()}')
shared.log.info(f"Creating model from config: {checkpoint_config}")
sd_model = None
shared.log.debug(f'Model config: {sd_config.model.get("params", dict())}')
try:
clip_is_included_into_sd = sd1_clip_weight in state_dict or sd2_clip_weight in state_dict
with sd_disable_initialization.DisableInitialization(disable_clip=clip_is_included_into_sd):
sd_model = instantiate_from_config(sd_config.model)
except Exception:
sd_model = instantiate_from_config(sd_config.model)
shared.log.info(f"Model created from config: {checkpoint_config}")
sd_model.used_config = checkpoint_config
timer.record("create")
load_model_weights(sd_model, checkpoint_info, state_dict, timer)
@@ -409,6 +407,7 @@ def load_model(checkpoint_info=None, already_loaded_state_dict=None, timer=None)
script_callbacks.model_loaded_callback(sd_model)
timer.record("callbacks")
shared.log.info(f"Model loaded in {timer.summary()}")
current_checkpoint_info = None
devices.torch_gc()
shared.log.debug(f'Model load finished: {memory_stats()}')
@@ -424,10 +423,6 @@ def reload_model_weights(sd_model=None, info=None):
checkpoint_info = info or select_checkpoint()
if not sd_model:
sd_model = shared.sd_model
if shared.opts.model_reuse_dict and sd_model is not None:
shared.log.info('Reusing previous model dictionary')
else:
sd_model = None
if sd_model is None: # previous model load failed
current_checkpoint_info = None
else:
@@ -439,6 +434,12 @@ def reload_model_weights(sd_model=None, info=None):
else:
sd_model.to(devices.cpu)
sd_hijack.model_hijack.undo_hijack(sd_model)
if shared.opts.model_reuse_dict and sd_model is not None:
shared.log.info('Reusing previous model dictionary')
else:
unload_model_weights()
sd_model = None
shared.sd_model = None
timer = Timer()
state_dict = get_checkpoint_state_dict(checkpoint_info, timer)
checkpoint_config = sd_models_config.find_checkpoint_config(state_dict, checkpoint_info)
@@ -466,7 +467,6 @@ def reload_model_weights(sd_model=None, info=None):
def unload_model_weights(sd_model=None, _info=None):
from modules import sd_hijack
timer = Timer()
if shared.sd_model:
# shared.sd_model.cond_stage_model.to(devices.cpu)
# shared.sd_model.first_stage_model.to(devices.cpu)
@@ -474,8 +474,8 @@ def unload_model_weights(sd_model=None, _info=None):
sd_hijack.model_hijack.undo_hijack(shared.sd_model)
shared.sd_model = None
sd_model = None
devices.torch_gc()
shared.log.info(f"Unloaded weights {timer.summary()}")
devices.torch_gc()
shared.log.debug(f'Model weights unloaded: {memory_stats()}')
return sd_model
+4 -1
View File
@@ -189,6 +189,7 @@ class ExtraNetworksUi:
self.description_target_filename = None
self.description_input = None
self.tabname = None
self.search = None
def pages_in_preferred_order(pages):
@@ -212,8 +213,9 @@ def create_ui(container, button, tabname):
for page in ui.stored_extra_pages:
with gr.Tab(page.title, id=page.title.lower().replace(" ", "_")):
page_elem = gr.HTML(page.create_html(ui.tabname))
page_elem.change(fn=lambda: None, _js=f'() => refreshExtraNetworks("{tabname}")', inputs=[], outputs=[])
ui.pages.append(page_elem)
_filter = gr.Textbox('', show_label=False, elem_id=tabname+"_extra_search", placeholder="Search...", visible=False)
ui.search = gr.Textbox('', show_label=False, elem_id=tabname+"_extra_search", placeholder="Search...", visible=False)
ui.description_input = gr.TextArea('', show_label=False, elem_id=tabname+"_description_input", placeholder="Save/Replace Extra Network Description...", lines=2)
button_refresh = ToolButton(refresh_symbol, elem_id=tabname+"_extra_refresh")
button_close = ToolButton(close_symbol, elem_id=tabname+"_extra_close")
@@ -236,6 +238,7 @@ def create_ui(container, button, tabname):
for pg in ui.stored_extra_pages:
pg.refresh()
res.append(pg.create_html(ui.tabname))
ui.search.update(value = ui.search.value)
return res
button_refresh.click(fn=refresh, inputs=[], outputs=ui.pages)
+3 -1
View File
@@ -534,7 +534,9 @@ class Script(scripts.Script):
zs = process_axis(z_opt, z_values, z_values_dropdown)
Image.MAX_IMAGE_PIXELS = None # disable check in Pillow and rely on check below to allow large custom image sizes
grid_mp = round(len(xs) * len(ys) * len(zs) * p.width * p.height / 1000000)
assert grid_mp < shared.opts.img_max_size_mp, f'Error: Resulting grid would be too large ({grid_mp} MPixels) (max configured size is {shared.opts.img_max_size_mp} MPixels)'
if grid_mp > shared.opts.img_max_size_mp:
shared.log.warning(f'Grid size: {grid_mp} excedes {shared.opts.img_max_size_mp} MPixels')
return
def fix_axis_seeds(axis_opt, axis_list):
if axis_opt.label in ['Seed', 'Var. seed']: