update train script

This commit is contained in:
Vladimir Mandic
2023-10-01 13:59:13 -04:00
parent 6837bd9bc1
commit 1c28757061
11 changed files with 164 additions and 136 deletions
+4 -1
View File
@@ -1,6 +1,6 @@
# Change Log for SD.Next
## Update for 2023-09-30
## Update for 2023-10-01
**TBD**: Candidates before release:
- Integrate LoRA/Lyco for *backend:original*
@@ -8,6 +8,7 @@
for *backend:original* use extension: <https://github.com/ljleb/sd-webui-freeu>
- Switch Diffusers prompt parser
- Custom log destination
- Include chaiNNer extension as submodule
This is a big one, with some major changes and new functionality...
And probably the biggest release since introduction of **Diffusers**
@@ -21,6 +22,8 @@ Upgrades are still possible and supported, but above is recommended for best exp
- converted submenus from checkboxes to accordion elements
any ui state including state of open/closed menus can be saved as default!
see *System -> User interface -> Set menu states*
- new built-in theme **invokeai**
thanks @BinaryQuantumSoul
- small visual indicator bottom right of page showing internal server job state
- **Extra networks**:
- **Details**
+11 -3
View File
@@ -51,7 +51,7 @@ def get_latents(local_vae, images, weight_dtype):
img_tensors = img_tensors.to(device, weight_dtype)
with torch.no_grad():
latents = local_vae.encode(img_tensors).latent_dist.sample().float().to('cpu').numpy()
return latents
return latents, [images[0].shape[0], images[0].shape[1]]
def get_npz_filename_wo_ext(data_dir, image_key):
@@ -87,11 +87,19 @@ def create_vae_latents(local_params):
def process_batch(is_last):
for bucket in bucket_manager.buckets:
if (is_last and len(bucket) > 0) or len(bucket) >= args.batch:
latents = get_latents(vae, [img for _, img in bucket], weight_dtype)
latents, original_size = get_latents(vae, [img for _, img in bucket], weight_dtype)
assert latents.shape[2] == bucket[0][1].shape[0] // 8 and latents.shape[3] == bucket[0][1].shape[1] // 8, f'latent shape {latents.shape}, {bucket[0][1].shape}'
for (image_key, _), latent in zip(bucket, latents):
npz_file_name = get_npz_filename_wo_ext(args.input, image_key)
np.savez(npz_file_name, latent)
# np.savez(npz_file_name, latent)
kwargs = {}
np.savez(
npz_file_name,
latents=latent,
original_size=np.array(original_size),
crop_ltrb=np.array([0, 0]),
**kwargs,
)
bucket.clear()
data = [[(None, ip)] for ip in image_paths]
bucket_counts = {}
+18 -24
View File
@@ -80,6 +80,7 @@ def parse_args():
group_main.add_argument('--model', type=str, default='', required=False, help='base model to use for training, default: current loaded model')
group_main.add_argument('--name', type=str, default=None, required=True, help='output filename')
group_main.add_argument('--tag', type=str, default='person', required=False, help='primary tags, default: %(default)s')
group_main.add_argument('--comments', type=str, default='', required=False, help='comments to be added to trained model metadata, default: %(default)s')
group_data = parser.add_argument_group('Dataset')
group_data.add_argument('--input', type=str, default=None, required=True, help='input folder with training images')
@@ -98,6 +99,8 @@ def parse_args():
group_train.add_argument('--alpha', type=float, default=0, required=False, help='lora/lyco alpha for weights scaling, default: dim/2')
group_train.add_argument('--algo', type=str, default=None, choices=['locon', 'loha', 'lokr', 'ia3'], required=False, help='alternative lyco algoritm, default: %(default)s')
group_train.add_argument('--args', type=str, default=None, required=False, help='lora/lyco additional network arguments, default: %(default)s')
group_train.add_argument('--optimizer', type=str, default='AdamW', required=False, help='optimizer type, default: %(default)s')
# AdamW (default), AdamW8bit, PagedAdamW8bit, Lion8bit, PagedLion8bit, Lion, SGDNesterov, SGDNesterov8bit, DAdaptation(DAdaptAdamPreprint), DAdaptAdaGrad, DAdaptAdam, DAdaptAdan, DAdaptAdanIP, DAdaptLion, DAdaptSGD, AdaFactor
group_other = parser.add_argument_group('Other')
group_other.add_argument('--overwrite', default = False, action='store_true', help = "overwrite existing training, default: %(default)s")
@@ -140,15 +143,22 @@ def verify_args():
sdapi.postsync('/sdapi/v1/options', server_options.options)
else:
args.model = server_options.options.sd_model_checkpoint.split(' [')[0]
base_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
args.lora_dir = server_options.options.lora_dir
if not os.path.isabs(args.lora_dir):
args.lora_dir = os.path.join(base_dir, args.lora_dir)
args.lyco_dir = server_options.options.lyco_dir
if not os.path.isabs(args.lyco_dir):
args.lyco_dir = os.path.join(base_dir, args.lyco_dir)
args.ckpt_dir = server_options.options.ckpt_dir
if not os.path.isabs(args.ckpt_dir):
args.ckpt_dir = os.path.join(base_dir, args.ckpt_dir)
args.embeddings_dir = server_options.options.embeddings_dir
if not os.path.isfile(args.model):
attempt = os.path.abspath(os.path.join(args.ckpt_dir, args.model))
args.model = attempt if os.path.isfile(attempt) else args.model
if not os.path.isfile(args.model):
attempt = os.path.abspath(os.path.join(args.ckpt_dir, '..', args.model))
attempt = os.path.abspath(os.path.join(args.ckpt_dir, args.model + '.safetensors'))
args.model = attempt if os.path.isfile(attempt) else args.model
if not os.path.isfile(args.model):
log.error(f'cannot find loaded model: {args.model}')
@@ -242,7 +252,8 @@ def train_lora():
sys.path.append(lycoris_path)
log.debug('importing lora lib')
import train_network
train_network.train(options.lora)
trainer = train_network.NetworkTrainer()
trainer.train(options.lora)
if args.type == 'lyco':
log.debug('importing lycoris lib')
import importlib
@@ -267,12 +278,16 @@ def prepare_options():
options.lora.network_module = 'lycoris.kohya'
options.lora.in_json = os.path.join(args.process_dir, args.name + '.json')
# lora specific
options.lora.save_model_as = 'safetensors'
options.lora.pretrained_model_name_or_path = args.model
options.lora.output_name = args.name
options.lora.max_train_steps = args.steps
options.lora.network_dim = args.dim
options.lora.network_alpha = args.dim // 2 if args.alpha == 0 else args.alpha
options.lora.netwoork_args = []
options.lora.training_comment = args.comments
options.lora.sdpa = True
options.lora.optimizer_type = args.optimizer
if args.algo is not None:
options.lora.netwoork_args.append(f'algo={args.algo}')
if args.args is not None:
@@ -367,31 +382,10 @@ def process_inputs():
process.unload()
def check_versions():
if args.experimental:
log.info('experimental mode enabled')
return
log.info('checking accelerate')
error = False
import accelerate
if accelerate.__version__ != '0.19.0':
log.error(f'invalid accelerate version: accelerate=0.19.0 found={accelerate.__version__}')
error = True
log.info('checking diffusers')
import diffusers
if diffusers.__version__ != '0.10.2':
log.error(f'invalid diffusers version: diffusers=0.10.2 found={diffusers.__version__}')
error = True
if error:
log.info('> pip install accelerate==0.20.3 diffusers==0.10.2')
exit(1)
if __name__ == '__main__':
log.info('SD.Next train script')
parse_args()
setup_logging()
check_versions()
log.info('SD.Next Train')
sdapi.sd_url = args.server
if args.user is not None:
sdapi.sd_username = args.user
@@ -14,29 +14,35 @@ class ExtraNetworksPageLora(ui_extra_networks.ExtraNetworksPage):
def list_items(self):
for name, l in lora.available_loras.items():
path, _ext = os.path.splitext(l.filename)
possible_tags = l.metadata.get('ss_tag_frequency', {}) if l.metadata is not None else {}
if isinstance(possible_tags, str):
possible_tags = {}
tags = {}
for k, v in possible_tags.items():
words = k.split('_', 1) if '_' in k else [v, k]
tags[' '.join(words[1:])] = words[0]
name = os.path.splitext(os.path.relpath(l.filename, shared.cmd_opts.lora_dir))[0]
yield {
"type": 'Lora',
"name": name,
"filename": l.filename,
"hash": l.shorthash,
"search_term": self.search_terms_from_path(l.filename) + ' '.join(tags.keys()),
"preview": self.find_preview(path),
"description": self.find_description(path),
"info": self.find_info(path),
"prompt": json.dumps(f" <lora:{l.get_alias()}:{shared.opts.extra_networks_default_multiplier}>"),
"local_preview": f"{path}.{shared.opts.samples_format}",
"metadata": json.dumps(l.metadata, indent=4) if l.metadata else None,
"tags": tags,
}
try:
path, _ext = os.path.splitext(l.filename)
possible_tags = l.metadata.get('ss_tag_frequency', {}) if l.metadata is not None else {}
if isinstance(possible_tags, str):
possible_tags = {}
tags = {}
for k, v in possible_tags.items():
words = k.split('_', 1) if '_' in k else [v, k]
words = [str(w).replace('.json', '') for w in words]
if words[0] == '{}':
words[0] = 0
tags[' '.join(words[1:])] = words[0]
name = os.path.splitext(os.path.relpath(l.filename, shared.cmd_opts.lora_dir))[0]
yield {
"type": 'Lora',
"name": name,
"filename": l.filename,
"hash": l.shorthash,
"search_term": self.search_terms_from_path(l.filename) + ' '.join(tags.keys()),
"preview": self.find_preview(path),
"description": self.find_description(path),
"info": self.find_info(path),
"prompt": json.dumps(f" <lora:{l.get_alias()}:{shared.opts.extra_networks_default_multiplier}>"),
"local_preview": f"{path}.{shared.opts.samples_format}",
"metadata": json.dumps(l.metadata, indent=4) if l.metadata else None,
"tags": tags,
}
except Exception as e:
shared.log.debug(f"Extra networks error: type=lora file={name} {e}")
def allowed_directories_for_previews(self):
return [shared.cmd_opts.lora_dir]
+2 -1
View File
@@ -640,8 +640,9 @@ def create_ui(container, button_parent, tabname, skip_indexing = False):
pages = []
for page in get_pages():
if title is None or title == '' or title == page.title or len(page.html) == 0:
page.refresh()
page.page_time = 0
page.refresh_time = 0
page.refresh()
page.create_page(ui.tabname)
shared.log.debug(f"Refreshing Extra networks: page='{page.title}' items={len(page.items)} tab={ui.tabname}")
pages.append(page.html)
+19 -16
View File
@@ -14,22 +14,25 @@ class ExtraNetworksPageCheckpoints(ui_extra_networks.ExtraNetworksPage):
def list_items(self):
checkpoint: sd_models.CheckpointInfo
for name, checkpoint in sd_models.checkpoints_list.items():
fn = os.path.splitext(checkpoint.filename)[0]
record = {
"type": 'Model',
"name": checkpoint.name,
"title": checkpoint.title,
"filename": checkpoint.filename,
"hash": checkpoint.shorthash,
"search_term": self.search_terms_from_path(checkpoint.title),
"preview": self.find_preview(fn),
"local_preview": f"{fn}.{shared.opts.samples_format}",
"description": self.find_description(fn),
"info": self.find_info(fn),
"metadata": checkpoint.metadata,
"onclick": '"' + html.escape(f"""return selectCheckpoint({json.dumps(name)})""") + '"',
}
yield record
try:
fn = os.path.splitext(checkpoint.filename)[0]
record = {
"type": 'Model',
"name": checkpoint.name,
"title": checkpoint.title,
"filename": checkpoint.filename,
"hash": checkpoint.shorthash,
"search_term": self.search_terms_from_path(checkpoint.title),
"preview": self.find_preview(fn),
"local_preview": f"{fn}.{shared.opts.samples_format}",
"description": self.find_description(fn),
"info": self.find_info(fn),
"metadata": checkpoint.metadata,
"onclick": '"' + html.escape(f"""return selectCheckpoint({json.dumps(name)})""") + '"',
}
yield record
except Exception as e:
shared.log.debug(f"Extra networks error: type=model file={name} {e}")
def allowed_directories_for_previews(self):
return [v for v in [shared.opts.ckpt_dir, shared.opts.diffusers_dir, sd_models.model_path] if v is not None]
+16 -13
View File
@@ -12,19 +12,22 @@ class ExtraNetworksPageHypernetworks(ui_extra_networks.ExtraNetworksPage):
def list_items(self):
for name, path in shared.hypernetworks.items():
fn = os.path.splitext(path)[0]
name = os.path.relpath(fn, shared.opts.hypernetwork_dir)
yield {
"type": 'Hypernetwork',
"name": os.path.relpath(fn, shared.opts.hypernetwork_dir),
"filename": path,
"preview": self.find_preview(fn),
"description": self.find_description(fn),
"info": self.find_info(fn),
"search_term": self.search_terms_from_path(name),
"prompt": json.dumps(f"<hypernet:{name}:{shared.opts.extra_networks_default_multiplier}>"),
"local_preview": f"{fn}.{shared.opts.samples_format}",
}
try:
fn = os.path.splitext(path)[0]
name = os.path.relpath(fn, shared.opts.hypernetwork_dir)
yield {
"type": 'Hypernetwork',
"name": os.path.relpath(fn, shared.opts.hypernetwork_dir),
"filename": path,
"preview": self.find_preview(fn),
"description": self.find_description(fn),
"info": self.find_info(fn),
"search_term": self.search_terms_from_path(name),
"prompt": json.dumps(f"<hypernet:{name}:{shared.opts.extra_networks_default_multiplier}>"),
"local_preview": f"{fn}.{shared.opts.samples_format}",
}
except Exception as e:
shared.log.debug(f"Extra networks error: type=hypernetwork file={path} {e}")
def allowed_directories_for_previews(self):
return [shared.opts.hypernetwork_dir]
+25 -21
View File
@@ -97,27 +97,31 @@ class ExtraNetworksPageStyles(ui_extra_networks.ExtraNetworksPage):
def list_items(self):
for k, style in shared.prompt_styles.styles.items():
fn = os.path.splitext(getattr(style, 'filename', ''))[0]
name = getattr(style, 'name', '')
if name == '':
continue
txt = f'Prompt: {getattr(style, "prompt", "")}'
if len(getattr(style, 'negative_prompt', '')) > 0:
txt += f'\nNegative: {style.negative_prompt}'
yield {
"type": 'Style',
"name": name,
"title": k,
"filename": style.filename,
"search_term": f'{txt} {self.search_terms_from_path(name)}',
"preview": style.preview if getattr(style, 'preview', None) is not None and style.preview.startswith('data:') else self.find_preview(fn),
"description": style.description if getattr(style, 'description', None) is not None and len(style.description) > 0 else txt,
"prompt": getattr(style, 'prompt', ''),
"negative": getattr(style, 'negative_prompt', ''),
"extra": getattr(style, 'extra', ''),
"local_preview": f"{fn}.{shared.opts.samples_format}",
"onclick": '"' + html.escape(f"""return selectStyle({json.dumps(name)})""") + '"',
}
try:
fn = os.path.splitext(getattr(style, 'filename', ''))[0]
name = getattr(style, 'name', '')
if name == '':
continue
txt = f'Prompt: {getattr(style, "prompt", "")}'
if len(getattr(style, 'negative_prompt', '')) > 0:
txt += f'\nNegative: {style.negative_prompt}'
yield {
"type": 'Style',
"name": name,
"title": k,
"filename": style.filename,
"search_term": f'{txt} {self.search_terms_from_path(name)}',
"preview": style.preview if getattr(style, 'preview', None) is not None and style.preview.startswith('data:') else self.find_preview(fn),
"description": style.description if getattr(style, 'description', None) is not None and len(style.description) > 0 else txt,
"prompt": getattr(style, 'prompt', ''),
"negative": getattr(style, 'negative_prompt', ''),
"extra": getattr(style, 'extra', ''),
"local_preview": f"{fn}.{shared.opts.samples_format}",
"onclick": '"' + html.escape(f"""return selectStyle({json.dumps(name)})""") + '"',
}
except Exception as e:
shared.log.debug(f"Extra networks error: type=style file={k} {e}")
def allowed_directories_for_previews(self):
return [v for v in [shared.opts.styles_dir] if v is not None]
+20 -17
View File
@@ -39,23 +39,26 @@ class ExtraNetworksPageTextualInversion(ui_extra_networks.ExtraNetworksPage):
embeddings = []
embeddings = sorted(embeddings, key=lambda emb: emb.filename)
for embedding in embeddings:
path, _ext = os.path.splitext(embedding.filename)
tags = {}
if embedding.tag is not None:
tags[embedding.tag]=1
name = os.path.splitext(embedding.basename)[0]
yield {
"type": 'Embedding',
"name": name,
"filename": embedding.filename,
"preview": self.find_preview(path),
"description": self.find_description(path),
"info": self.find_info(path),
"search_term": self.search_terms_from_path(name),
"prompt": json.dumps(os.path.splitext(embedding.name)[0]),
"local_preview": f"{path}.{shared.opts.samples_format}",
"tags": tags,
}
try:
path, _ext = os.path.splitext(embedding.filename)
tags = {}
if embedding.tag is not None:
tags[embedding.tag]=1
name = os.path.splitext(embedding.basename)[0]
yield {
"type": 'Embedding',
"name": name,
"filename": embedding.filename,
"preview": self.find_preview(path),
"description": self.find_description(path),
"info": self.find_info(path),
"search_term": self.search_terms_from_path(name),
"prompt": json.dumps(os.path.splitext(embedding.name)[0]),
"local_preview": f"{path}.{shared.opts.samples_format}",
"tags": tags,
}
except Exception as e:
shared.log.debug(f"Extra networks error: type=embedding file={embedding.filename} {e}")
def allowed_directories_for_previews(self):
return list(sd_hijack.model_hijack.embedding_db.embedding_dirs)
+19 -16
View File
@@ -13,22 +13,25 @@ class ExtraNetworksPageVAEs(ui_extra_networks.ExtraNetworksPage):
def list_items(self):
for name, filename in sd_vae.vae_dict.items():
fn = os.path.splitext(filename)[0]
record = {
"type": 'VAE',
"name": name,
"title": name,
"filename": fn,
"hash": hashes.sha256_from_cache(filename, f"vae/{fn}"),
"search_term": self.search_terms_from_path(fn),
"preview": self.find_preview(fn),
"local_preview": f"{fn}.{shared.opts.samples_format}",
"description": self.find_description(fn),
"info": self.find_info(fn),
"metadata": {},
"onclick": '"' + html.escape(f"""return selectVAE({json.dumps(name)})""") + '"',
}
yield record
try:
fn = os.path.splitext(filename)[0]
record = {
"type": 'VAE',
"name": name,
"title": name,
"filename": fn,
"hash": hashes.sha256_from_cache(filename, f"vae/{fn}"),
"search_term": self.search_terms_from_path(fn),
"preview": self.find_preview(fn),
"local_preview": f"{fn}.{shared.opts.samples_format}",
"description": self.find_description(fn),
"info": self.find_info(fn),
"metadata": {},
"onclick": '"' + html.escape(f"""return selectVAE({json.dumps(name)})""") + '"',
}
yield record
except Exception as e:
shared.log.debug(f"Extra networks error: type=vae file={filename} {e}")
def allowed_directories_for_previews(self):
return [v for v in [shared.opts.vae_dir] if v is not None]
+1 -1
View File
@@ -59,7 +59,7 @@ numba==0.57.1
pandas==1.5.3
protobuf==3.20.3
pytorch_lightning==1.9.4
transformers==4.31.0
transformers==4.30.2
tomesd==0.1.3
urllib3==1.26.15
Pillow==9.5.0