From 11fa3aff6d671012a9127c9e9afb3808d1aaf89a Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Tue, 25 Apr 2023 09:21:38 -0400 Subject: [PATCH] ti fixes --- .gitignore | 2 + javascript/progressbar.js | 2 +- javascript/textualInversion.js | 1 + .../textual_inversion/textual_inversion.py | 52 ++++++++----------- modules/textual_inversion/ui.py | 5 -- modules/ui.py | 33 +----------- modules/ui_extra_networks_hypernets.py | 2 - .../ui_extra_networks_textual_inversion.py | 4 +- setup.py | 10 +++- 9 files changed, 39 insertions(+), 72 deletions(-) diff --git a/.gitignore b/.gitignore index fa8769047..8a9c92f78 100644 --- a/.gitignore +++ b/.gitignore @@ -32,6 +32,8 @@ venv /models/**/* /interrogate/**/* /train/log/**/* +/textual_inversion/**/* +/detected_maps/**/* /tmp /log /cert diff --git a/javascript/progressbar.js b/javascript/progressbar.js index e9b6bf5d5..7d218ff05 100644 --- a/javascript/progressbar.js +++ b/javascript/progressbar.js @@ -77,7 +77,7 @@ function requestProgress(id_task, progressbarContainer, gallery, atEnd = null, o setTitle("") if (divProgress) parentProgressbar.removeChild(divProgress) if (parentGallery) parentGallery.removeChild(livePreview) - atEnd() + if (atEnd) atEnd() } var fun = function(id_task, id_live_preview){ diff --git a/javascript/textualInversion.js b/javascript/textualInversion.js index 883092547..8a50ec601 100644 --- a/javascript/textualInversion.js +++ b/javascript/textualInversion.js @@ -5,6 +5,7 @@ function start_training_textual_inversion(){ gradioApp().querySelector('#ti_error').innerHTML='' var id = randomId() const onProgress = (progress) => gradioApp().getElementById('ti_progress').innerHTML = progress.textinfo; + // requestProgress(id_task, progressbarContainer, gallery, atEnd = null, onProgress = null, once = false) { requestProgress(id, gradioApp().getElementById('ti_output'), gradioApp().getElementById('ti_gallery'), null, onProgress, false) var res = args_to_array(arguments) res[0] = id diff --git a/modules/textual_inversion/textual_inversion.py b/modules/textual_inversion/textual_inversion.py index 9bcf205ab..24ccf743f 100644 --- a/modules/textual_inversion/textual_inversion.py +++ b/modules/textual_inversion/textual_inversion.py @@ -1,23 +1,19 @@ import os -from collections import namedtuple - -import torch -import tqdm import html import csv +from collections import namedtuple +import torch +import tqdm import safetensors.torch - +from rich import print # pylint: disable=redefined-builtin import numpy as np from PIL import Image, PngImagePlugin from torch.utils.tensorboard import SummaryWriter - from modules import shared, devices, sd_hijack, processing, sd_models, images, sd_samplers, sd_hijack_checkpoint, errors import modules.textual_inversion.dataset from modules.textual_inversion.learn_schedule import LearnRateScheduler - from modules.textual_inversion.image_embedding import embedding_to_b64, embedding_from_b64, insert_image_data_embed, extract_image_data_embed, caption_image_overlay from modules.textual_inversion.logging import save_settings_to_file -from rich import print # pylint: disable=redefined-builtin TextualInversionTemplate = namedtuple("TextualInversionTemplate", ["name", "path"]) textual_inversion_templates = {} @@ -26,7 +22,7 @@ textual_inversion_templates = {} def list_textual_inversion_templates(): textual_inversion_templates.clear() - for root, dirs, fns in os.walk(shared.opts.embeddings_templates_dir): + for root, _dirs, fns in os.walk(shared.opts.embeddings_templates_dir): for fn in fns: path = os.path.join(root, fn) @@ -198,7 +194,7 @@ class EmbeddingDatabase: if not os.path.isdir(embdir.path): return - for root, dirs, fns in os.walk(embdir.path, followlinks=True): + for root, _dirs, fns in os.walk(embdir.path, followlinks=True): for fn in fns: try: fullfn = os.path.join(root, fn) @@ -214,7 +210,7 @@ class EmbeddingDatabase: def load_textual_inversion_embeddings(self, force_reload=False): if not force_reload: need_reload = False - for path, embdir in self.embedding_dirs.items(): + for _path, embdir in self.embedding_dirs.items(): if embdir.has_changed(): need_reload = True break @@ -227,7 +223,7 @@ class EmbeddingDatabase: self.skipped_embeddings.clear() self.expected_shape = self.get_expected_shape() - for path, embdir in self.embedding_dirs.items(): + for _path, embdir in self.embedding_dirs.items(): self.load_from_dir(embdir) embdir.update() @@ -270,13 +266,13 @@ def create_embedding(name, num_vectors_per_token, overwrite_old, init_text='*'): # Remove illegal characters from name. name = "".join( x for x in name if (x.isalnum() or x in "._- ")) fn = os.path.join(shared.opts.embeddings_dir, f"{name}.pt") - if not overwrite_old: - assert not os.path.exists(fn), f"file {fn} already exists" - - embedding = Embedding(vec, name) - embedding.step = 0 - embedding.save(fn) - + if not overwrite_old and os.path.exists(fn): + print(f"Embedding already exists: {fn}") + else: + embedding = Embedding(vec, name) + embedding.step = 0 + embedding.save(fn) + print(f'Created embedding: {fn} vectors {num_vectors_per_token} init {init_text}') return fn @@ -288,7 +284,7 @@ def write_loss(log_directory, filename, step, epoch_len, values): return write_csv_header = False if os.path.exists(os.path.join(log_directory, filename)) else True - with open(os.path.join(log_directory, filename), "a+", newline='') as fout: + with open(os.path.join(log_directory, filename), "a+", newline='', encoding='utf-8') as fout: csv_writer = csv.DictWriter(fout, fieldnames=["step", "epoch", "epoch_step", *(values.keys())]) if write_csv_header: @@ -460,7 +456,7 @@ def train_embedding(id_task, embedding_name, learn_rate, batch_size, gradient_st try: sd_hijack_checkpoint.add() - for i in range((steps-initial_step) * gradient_step): + for _i in range((steps-initial_step) * gradient_step): if scheduler.finished: break if shared.state.interrupted: @@ -574,7 +570,7 @@ def train_embedding(id_task, embedding_name, learn_rate, batch_size, gradient_st if image is not None: shared.state.assign_current_image(image) - last_saved_image, last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False) + last_saved_image, _last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False) last_saved_image += f", prompt: {preview_text}" if shared.opts.training_enable_tensorboard and shared.opts.training_tensorboard_save_images: @@ -587,18 +583,17 @@ def train_embedding(id_task, embedding_name, learn_rate, batch_size, gradient_st info = PngImagePlugin.PngInfo() data = torch.load(last_saved_file) info.add_text("sd-ti-embedding", embedding_to_b64(data)) - - title = "<{}>".format(data.get('name', '???')) + title = f"<{data.get('name', '???')}>" try: vectorSize = list(data['string_to_param'].values())[0].shape[0] - except Exception as e: + except Exception: vectorSize = '?' checkpoint = sd_models.select_checkpoint() footer_left = checkpoint.model_name - footer_mid = '[{}]'.format(checkpoint.shorthash) - footer_right = '{}v {}s'.format(vectorSize, steps_done) + footer_mid = f'[{checkpoint.shorthash}]' + footer_right = f'{vectorSize}v {steps_done}s' captioned_image = caption_image_overlay(image, title, footer_left, footer_mid, footer_right) captioned_image = insert_image_data_embed(captioned_image, data) @@ -606,7 +601,7 @@ def train_embedding(id_task, embedding_name, learn_rate, batch_size, gradient_st captioned_image.save(last_saved_image_chunks, "PNG", pnginfo=info) embedding_yet_to_be_embedded = False - last_saved_image, last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False) + last_saved_image, _last_text_info = images.save_image(image, images_dir, "", p.seed, p.prompt, shared.opts.samples_format, processed.infotexts[0], p=p, forced_filename=forced_filename, save_to_dirs=False) last_saved_image += f", prompt: {preview_text}" shared.state.job_no = embedding.step @@ -624,7 +619,6 @@ Last saved image: {html.escape(last_saved_image)}
save_embedding(embedding, optimizer, checkpoint, embedding_name, filename, remove_cached_checksum=True) except Exception as e: errors.display(e, 'embedding train') - pass finally: pbar.leave = False pbar.close() diff --git a/modules/textual_inversion/ui.py b/modules/textual_inversion/ui.py index af7ddb1e7..8afd8ef04 100644 --- a/modules/textual_inversion/ui.py +++ b/modules/textual_inversion/ui.py @@ -1,7 +1,5 @@ import html - import gradio as gr - import modules.textual_inversion.textual_inversion import modules.textual_inversion.preprocess from modules import sd_hijack, shared @@ -9,9 +7,7 @@ from modules import sd_hijack, shared def create_embedding(name, initialization_text, nvpt, overwrite_old): filename = modules.textual_inversion.textual_inversion.create_embedding(name, nvpt, overwrite_old, init_text=initialization_text) - sd_hijack.model_hijack.embedding_db.load_textual_inversion_embeddings() - return gr.Dropdown.update(choices=sorted(sd_hijack.model_hijack.embedding_db.word_embeddings.keys())), f"Created: {filename}", "" @@ -42,4 +38,3 @@ Embedding saved to {html.escape(filename)} finally: if not apply_optimizations: sd_hijack.apply_optimizations() - diff --git a/modules/ui.py b/modules/ui.py index 0077102c3..83249d39c 100644 --- a/modules/ui.py +++ b/modules/ui.py @@ -1127,7 +1127,7 @@ def create_ui(): gradient_step = gr.Number(label='Gradient accumulation steps', value=1, precision=0, elem_id="train_gradient_step") dataset_directory = gr.Textbox(label='Dataset directory', placeholder="Path to directory with input images", elem_id="train_dataset_directory") - log_directory = gr.Textbox(label='Log directory', placeholder="Path to directory where to write outputs", value="textual_inversion", elem_id="train_log_directory") + log_directory = gr.Textbox(label='Log directory', placeholder="Path to directory where to write outputs", value="train/log/embeddings", elem_id="train_log_directory") with FormRow(): template_file = gr.Dropdown(label='Prompt template', value="style_filewords.txt", elem_id="train_template_file", choices=get_textual_inversion_template_names()) @@ -1136,7 +1136,7 @@ def create_ui(): training_width = gr.Slider(minimum=64, maximum=2048, step=8, label="Width", value=512, elem_id="train_training_width") training_height = gr.Slider(minimum=64, maximum=2048, step=8, label="Height", value=512, elem_id="train_training_height") varsize = gr.Checkbox(label="Do not resize images", value=False, elem_id="train_varsize") - steps = gr.Number(label='Max steps', value=100000, precision=0, elem_id="train_steps") + steps = gr.Number(label='Max steps', value=1000, precision=0, elem_id="train_steps") with FormRow(): create_image_every = gr.Number(label='Save an image to log directory every N steps, 0 to disable', value=500, precision=0, elem_id="train_create_image_every") @@ -1740,32 +1740,3 @@ def reload_javascript(): if not hasattr(shared, 'GradioTemplateResponseOriginal'): shared.GradioTemplateResponseOriginal = gradio.routes.templates.TemplateResponse - - -def versions_html(): - import torch - import launch - - python_version = ".".join([str(x) for x in sys.version_info[0:3]]) - commit = launch.commit_hash() - short_commit = commit[0:8] - - if shared.xformers_available: - import xformers - xformers_version = xformers.__version__ - else: - xformers_version = "N/A" - - return f""" -python: {python_version} - •  -torch: {getattr(torch, '__long_version__',torch.__version__)} - •  -xformers: {xformers_version} - •  -gradio: {gr.__version__} - •  -commit: {short_commit} - •  -checkpoint: N/A -""" diff --git a/modules/ui_extra_networks_hypernets.py b/modules/ui_extra_networks_hypernets.py index 3deb9cbac..545898486 100644 --- a/modules/ui_extra_networks_hypernets.py +++ b/modules/ui_extra_networks_hypernets.py @@ -1,6 +1,5 @@ import json import os - from modules import shared, ui_extra_networks @@ -14,7 +13,6 @@ class ExtraNetworksPageHypernetworks(ui_extra_networks.ExtraNetworksPage): def list_items(self): for name, path in shared.hypernetworks.items(): path, ext = os.path.splitext(path) - yield { "name": name, "filename": path, diff --git a/modules/ui_extra_networks_textual_inversion.py b/modules/ui_extra_networks_textual_inversion.py index b8dd408d8..1abf39675 100644 --- a/modules/ui_extra_networks_textual_inversion.py +++ b/modules/ui_extra_networks_textual_inversion.py @@ -25,12 +25,12 @@ class ExtraNetworksPageTextualInversion(ui_extra_networks.ExtraNetworksPage): for embedding in embeddings: path, _ext = os.path.splitext(embedding.filename) yield { - "name": embedding.name, + "name": os.path.splitext(embedding.name)[0], "filename": embedding.filename, "preview": self.find_preview(path), "description": self.find_description(path), "search_term": self.search_terms_from_path(embedding.filename), - "prompt": json.dumps(embedding.name), + "prompt": json.dumps(os.path.splitext(embedding.name)[0]), "local_preview": f"{path}.preview.{shared.opts.samples_format}", } diff --git a/setup.py b/setup.py index 2e97322f2..a9499dc6b 100644 --- a/setup.py +++ b/setup.py @@ -136,12 +136,18 @@ def update(folder): branch = git('branch', folder) if 'main' in branch: git('checkout main', folder) + branch = 'main' elif 'master' in branch: git('checkout master', folder) + branch = 'master' else: log.warning(f'Unknown branch for: {folder}') - git('pull --autostash --rebase', folder) - branch = git('branch', folder) + branch = None + if branch is None: + git('pull --autostash --rebase', folder) + else: + git(f'pull origin {branch} --autostash --rebase', folder) + # branch = git('branch', folder) # clone git repository