diff --git a/cli/.pylintrc b/cli/.pylintrc
deleted file mode 100644
index 1231dff7a..000000000
--- a/cli/.pylintrc
+++ /dev/null
@@ -1,3 +0,0 @@
-# See https://pylint.pycqa.org/en/latest/user_guide/messages/message_control.html
-[MESSAGES CONTROL]
-enable=C,R,W,E,I
diff --git a/cli/README.md b/cli/README.md
index 1081e3215..70de255b1 100644
--- a/cli/README.md
+++ b/cli/README.md
@@ -1,9 +1,6 @@
# Stable-Diffusion Productivity Scripts
-*Notes*:
-- Offline scripts can be used with or without **Automatic WebUI**
-- Online scripts rely on **Automatic WebUI** API which should be started with `--api` parameter
-- All scripts have built-in `--help` parameter that can be used to get more information
+Note: All scripts have built-in `--help` parameter that can be used to get more information
@@ -18,32 +15,29 @@ Supports upsampling, face restoration and grid creation
By default uses parameters from `generate.json`
Parameters that are not specified will be randomized:
+
- Prompt will be dynamically created from template of random samples: `random.json`
- Sampler/Scheduler will be randomly picked from available ones
- CFG Scale set to 5-10
### Train
-Textual inversion embedding training
-> python train-ti.py
+Combined pipeline for **embeddings**, **lora**, **lycoris**, **dreambooth** and **hypernetwork**
+Optionally runs several image processing steps before training:
-Combined pipeline:
-1. Creates embedding
-2. Extracts images if input is movie
-3. Preprocesses images
-4. Runs training
+- keep original image
+- detect and extract face
+- detect and extract body
+- detect blur
+- detect dynamic range
+- attempt to upscale low resolution images
+- attempt to restore quality of low quality images
+- automatically generate captions using interrogate
+- resize image
+- square image
+- run image segmentation to remove background
-LoRA training
-> python train-lora.py
-
-Combined pipeline:
-1. Preprocesses images
-2. Runs training
-
-[Detailed documentation](https://github.com/vladmandic/automatic/wiki/Process.md)
-
-LoRA extract from model
-> python moidules/lora-extract.py
+> python train.py
@@ -51,107 +45,58 @@ LoRA extract from model
### Benchmark
-Benchmark your **Automatic WebUI**
-Note: Requires SD API
+> python run-benchmark.py
-> python modules/bench.py
+### Create Previews
-### Embedding Previews
+Create previews for **embeddings**, **lora**, **lycoris**, **dreambooth** and **hypernetwork**
-Create previews of embeddings using preview templates
-Note: Requires SD API
+> python create-previews.py
-> python modules/preview-embeddings.py
+## Image Grid
-## Grid
-
-Create flexible image grids from any number of images
-Note: Offline tool
-
-> python modiles/grid.py
+> python image-grid.py
### Image Watermark
Create invisible image watermark and remove existing EXIF tags
-Note: Offline tool
-> python modules/image-watermark.py
+> python image-watermark.py
-### Interrogate
+### Image Interrogate
Runs CLiP and Booru image interrogation
-Note: Requires SD API
-> python modules/interrogate.py
-
-### Interrogate-Offline
-
-Standalone implementation of GiT, CLiP and ViT image interrogation
-Note: Offline tool
-
-> python modules/interrogate-offline.py
-
-### Models Previews
-
-Create previews of models using built-in templates
-Note: Requires SD API
-
-> python modules/preview-models.py
+> python image-interrogate.py
### Palette Extract
Extract color palette from image(s)
-Note: Offline tool
-> python modules/palette-extract.py
-
-### Image Process
-
-Run image processing to extract face/body segments and run resolution/blur/dynamic-range checks
-Note: Offline except for interrogate to generate caption files which requires SD API
-
-> python modules/process.py
-
-[Detailed documentation](https://github.com/vladmandic/automatic/wiki/Process.md)
+> python image-palette.py
### Prompt Ideas
Generate complex prompt ideas
-Note: Offline tool
-> python modules/prompt-ideas.py
+> python prompt-ideas.py
### Prompt Promptist
Attempts to beautify the provided prompt
-Note: Offline tool
-> python modules/promptist.py
-
-### Training Loss-Chart
-
-Create loss-chart from training log
-Note: Offline tool, may require adjustment to train paths if used with other repos
-
-> python modules/train-losschart.py
-
-### Training Loss-Rate
-
-Create customizable loss rate to be used in training
-Note: Offline tool
-
-> python modules/train-lossrate.py
+> python prompt-promptist.py
### Video Extract
Extract frames from video files
-Note: Offline tool
-> python modules/video-extract.py
+> python video-extract.py
## Utility Scripts
+
### SDAPI
Utility module that handles async communication to Automatic API endpoints
diff --git a/cli/modules/preview-models.py b/cli/create-previews.py
similarity index 86%
rename from cli/modules/preview-models.py
rename to cli/create-previews.py
index 493d080c4..06c4eb097 100755
--- a/cli/modules/preview-models.py
+++ b/cli/create-previews.py
@@ -1,17 +1,15 @@
#!/usr/bin/env python
import os
-import sys
import json
import time
+import importlib
import asyncio
import argparse
from pathlib import Path
from util import Map, log
from sdapi import get, post, close
-from grid import grid
-
-sys.path.append(os.path.join(os.path.dirname(__file__), '..'))
from generate import generate # pylint: disable=import-error
+grid = importlib.import_module('image-grid').grid
default = 'sd-v15-runwayml.ckpt [cc6cb27103]'
@@ -63,7 +61,7 @@ options = Map({
"sd_vae": "vae-ft-mse-840000-ema-pruned.ckpt",
},
'lora': {
- 'strength': 0.9,
+ 'strength': 1.0,
},
'hypernetwork': {
'keyword': 'beautiful sexy woman',
@@ -115,6 +113,8 @@ async def preview_models(params):
log.info({ 'model load': model })
opt['sd_model_checkpoint'] = model
+ del opt['sd_lora']
+ del opt['sd_lyco']
await post('/sdapi/v1/options', opt)
opt = await get('/sdapi/v1/options')
images = []
@@ -142,6 +142,8 @@ async def preview_models(params):
if opt['sd_model_checkpoint'] != default and not params.fixed:
log.info({ 'model set default': default })
opt['sd_model_checkpoint'] = default
+ del opt['sd_lora']
+ del opt['sd_lyco']
await post('/sdapi/v1/options', opt)
@@ -263,11 +265,49 @@ async def hypernetwork(params):
log.info({ 'hypernetwork preview created': model, 'image': fn, 'images': len(images), 'grid': [image.width, image.height], 'time': round(t, 2), 'its': round(its, 2) })
+async def embedding(params):
+ opt = await get('/sdapi/v1/options')
+ folder = opt['embeddings_dir']
+ if not os.path.exists(folder):
+ log.error({ 'embeddings directory not found': folder })
+ return
+ models = [f.stem for f in Path(folder).glob('*.pt')]
+ log.info({ 'embeddings': len(models) })
+ for model in models:
+ fn = os.path.join(folder, model + '.preview' + options.format)
+ if os.path.exists(fn) and len(params.input) == 0: # if model preview exists and not manually included
+ log.info({ 'embedding preview exists': model })
+ continue
+ images = []
+ labels = []
+ t0 = time.time()
+ import re
+ keyword = '\"' + re.sub(r'\d', '', model) + '\"'
+ options.generate.batch_size = 4
+ options.generate.prompt = prompt.replace('', keyword)
+ options.generate.prompt = options.generate.prompt.replace('', '')
+ log.info({ 'embedding generating': model, 'keyword': keyword, 'prompt': options.generate.prompt })
+ data = await generate(options = options, quiet=True)
+ if 'image' in data:
+ for img in data['image']:
+ images.append(img)
+ labels.append(keyword)
+ else:
+ log.error({ 'lyco': model, 'keyword': keyword, 'error': data })
+ t1 = time.time()
+ image = grid(images = images, labels = labels, border = 8)
+ image.save(fn)
+ t = t1 - t0
+ its = 1.0 * options.generate.steps * len(images) / t
+ log.info({ 'embeding preview created': model, 'image': fn, 'images': len(images), 'grid': [image.width, image.height], 'time': round(t, 2), 'its': round(its, 2) })
+
+
async def create_previews(params):
await preview_models(params)
await lora(params)
await lyco(params)
await hypernetwork(params)
+ await embedding(params)
await close()
diff --git a/cli/generate.py b/cli/generate.py
index 5e617ace0..929a87634 100755
--- a/cli/generate.py
+++ b/cli/generate.py
@@ -32,9 +32,8 @@ from PIL import Image
from PIL.ExifTags import TAGS
from PIL.TiffImagePlugin import ImageFileDirectory_v2
-sys.path.append(os.path.join(os.path.dirname(__file__), 'modules'))
-from modules.sdapi import close, get, interrupt, post, session
-from modules.util import Map, log, safestring
+from sdapi import close, get, interrupt, post, session
+from util import Map, log, safestring
sd = {}
diff --git a/cli/modules/grid.py b/cli/image-grid.py
similarity index 100%
rename from cli/modules/grid.py
rename to cli/image-grid.py
diff --git a/cli/modules/interrogate.py b/cli/image-interrogate.py
similarity index 96%
rename from cli/modules/interrogate.py
rename to cli/image-interrogate.py
index 0442280dd..211980046 100755
--- a/cli/modules/interrogate.py
+++ b/cli/image-interrogate.py
@@ -11,7 +11,7 @@ import asyncio
import filetype
from PIL import Image
from util import log, Map
-import sdapi as sdapi
+import sdapi
stats = { 'captions': {}, 'keywords': {} }
@@ -96,7 +96,7 @@ async def main():
elif os.path.isdir(arg):
for root, _dirs, files in os.walk(arg):
for f in files:
- caption, keywords, _style = await interrogate(os.path.join(root, f))
+ _caption, _keywords, _style = await interrogate(os.path.join(root, f))
else:
log.error({ 'interrogate unknown file type': arg })
else:
diff --git a/cli/modules/palette-extract.py b/cli/image-palette.py
similarity index 82%
rename from cli/modules/palette-extract.py
rename to cli/image-palette.py
index 0472009c6..77eb5e130 100755
--- a/cli/modules/palette-extract.py
+++ b/cli/image-palette.py
@@ -5,23 +5,23 @@ import os
import io
import pathlib
import argparse
+import importlib
import pandas as pd
import numpy as np
import extcolors
import filetype
import matplotlib.pyplot as plt
import matplotlib.patches as patches
-import matplotlib.image as mpimg
from matplotlib.offsetbox import OffsetImage, AnnotationBbox
from colormap import rgb2hex
from PIL import Image
from util import log
-from grid import grid
+grid = importlib.import_module('image-grid').grid
-def color_to_df(input):
- colors_pre_list = str(input).replace('([(','').split(', (')[0:-1]
+def color_to_df(param):
+ colors_pre_list = str(param).replace('([(','').split(', (')[0:-1]
df_rgb = [i.split('), ')[0] + ')' for i in colors_pre_list]
- df_percent = [i.split('), ')[1].replace(')','') for i in colors_pre_list]
+ df_percent = [i.split('), ')[1].replace(')','') for i in colors_pre_list]
#convert RGB to HEX code
df_color_up = [rgb2hex(int(i.split(", ")[0].replace("(","")),
int(i.split(", ")[1]),
@@ -30,14 +30,14 @@ def color_to_df(input):
return df
-def palette(img, args, output):
+def palette(img, params, output):
size = 1024
img.thumbnail((size, size), Image.HAMMING)
-
+
#crate dataframe
- colors_x = extcolors.extract_from_image(img, tolerance = args.color, limit = 13)
+ colors_x = extcolors.extract_from_image(img, tolerance = params.color, limit = 13)
df_color = color_to_df(colors_x)
-
+
#annotate text
list_color = list(df_color['c_code'])
list_precent = [int(i) for i in list(df_color['occurence'])]
@@ -54,7 +54,7 @@ def palette(img, args, output):
imagebox = OffsetImage(data, zoom=2.5)
ab = AnnotationBbox(imagebox, (0, 0))
ax1.add_artist(ab)
-
+
#color palette
x_posi, y_posi, y_posi2 = 160, -260, -260
for c in list_color:
@@ -100,20 +100,20 @@ if __name__ == '__main__':
args = parser.parse_args()
log.info({ 'palette args': vars(args) })
if args.output != '':
- pathlib.Path(args.output).mkdir(parents = True, exist_ok = True)
+ pathlib.Path(args.output).mkdir(parents = True, exist_ok = True)
if not args.grid:
for arg in args.input:
if os.path.isfile(arg) and filetype.is_image(arg):
- img = Image.open(arg)
- output = os.path.join(args.output, pathlib.Path(arg).stem + '-' + args.suffix + '.jpg')
- palette(img, args, output)
+ image = Image.open(arg)
+ fn = os.path.join(args.output, pathlib.Path(arg).stem + '-' + args.suffix + '.jpg')
+ palette(image, args, fn)
elif os.path.isdir(arg):
for root, _dirs, files in os.walk(arg):
for f in files:
if filetype.is_image(os.path.join(root, f)):
- img = Image.open(os.path.join(root, f))
- output = os.path.join(args.output, pathlib.Path(f).stem + '-' + args.suffix + '.jpg')
- palette(img, args, output)
+ image = Image.open(os.path.join(root, f))
+ fn = os.path.join(args.output, pathlib.Path(f).stem + '-' + args.suffix + '.jpg')
+ palette(image, args, fn)
else:
images = []
for arg in args.input:
@@ -124,6 +124,6 @@ if __name__ == '__main__':
for f in files:
if filetype.is_image(os.path.join(root, f)):
images.append(Image.open(os.path.join(root, f)))
- img = grid(images)
- output = os.path.join(args.output, args.suffix + '.jpg')
- palette(img, args, output)
+ image = grid(images)
+ fn = os.path.join(args.output, args.suffix + '.jpg')
+ palette(image, args, fn)
diff --git a/cli/modules/image-watermark.py b/cli/image-watermark.py
similarity index 68%
rename from cli/modules/image-watermark.py
rename to cli/image-watermark.py
index 4b75fe481..16f3fb14a 100755
--- a/cli/modules/image-watermark.py
+++ b/cli/image-watermark.py
@@ -44,31 +44,31 @@ def set_exif(d: dict):
ifd[_TAGS[k]] = v
exif_stream = io.BytesIO()
ifd.save(exif_stream)
- bytes = b'Exif\x00\x00' + exif_stream.getvalue()
- return bytes
+ encoded = b'Exif\x00\x00' + exif_stream.getvalue()
+ return encoded
-def get_watermark(image, args):
+def get_watermark(image, params):
data = np.asarray(image)
- decoder = WatermarkDecoder(options.type, args.length)
- bytes = decoder.decode(data, options.method)
+ decoder = WatermarkDecoder(options.type, params.length)
+ decoded = decoder.decode(data, options.method)
try:
- watermark = str(bytes, 'UTF-8').replace('\x00', '')
+ s = str(decoded, 'UTF-8').replace('\x00', '')
except:
- watermark = ''
- return watermark
+ s = ''
+ return s
-def set_watermark(image, args):
+def set_watermark(image, params):
data = np.asarray(image)
encoder = WatermarkEncoder()
- encoder.set_watermark(options.type, args.wm.encode('utf-8'))
+ encoder.set_watermark(options.type, params.wm.encode('utf-8'))
encoded = encoder.encode(data, options.method)
image = Image.fromarray(encoded)
return image
-def watermark(args, file):
+def watermark(params, file):
if not os.path.exists(file):
log.error({ 'watermark': 'file not found' })
return
@@ -82,30 +82,30 @@ def watermark(args, file):
exif = get_exif(image)
- if args.command == 'read':
- watermark = get_watermark(image, args)
- log.info({ 'file': file, 'watermark': watermark, 'exif': exif, 'resolution': f'{image.width}x{image.height}' })
+ if params.command == 'read':
+ wm = get_watermark(image, params)
+ log.info({ 'file': file, 'watermark': wm, 'exif': exif, 'resolution': f'{image.width}x{image.height}' })
- elif args.command == 'write':
- metadata = b'' if args.strip else set_exif(exif)
- if args.output != '':
- pathlib.Path(args.output).mkdir(parents = True, exist_ok = True)
- image=set_watermark(image, args)
- fn = os.path.join(args.output, file)
+ elif params.command == 'write':
+ metadata = b'' if params.strip else set_exif(exif)
+ if params.output != '':
+ pathlib.Path(params.output).mkdir(parents = True, exist_ok = True)
+ image=set_watermark(image, params)
+ fn = os.path.join(params.output, file)
image.save(fn, exif=metadata)
- if args.verify:
+ if params.verify:
data = np.asarray(image)
- decoder = WatermarkDecoder(options.type, args.length)
- bytes = decoder.decode(data, options.method)
- if bytes.startswith(b'\xff'):
- watermark = ''
+ decoder = WatermarkDecoder(options.type, params.length)
+ decoded = decoder.decode(data, options.method)
+ if decoded.startswith(b'\xff'):
+ wm = ''
else:
- watermark = str(bytes, 'UTF-8').replace('\x00', '')
+ wm = str(decoded, 'UTF-8').replace('\x00', '')
else:
- watermark = args.wm
+ wm = params.wm
- log.info({ 'file': fn, 'watermark': watermark, 'exif': None if args.strip else exif, 'resolution': f'{image.width}x{image.height}' })
+ log.info({ 'file': fn, 'watermark': wm, 'exif': None if params.strip else exif, 'resolution': f'{image.width}x{image.height}' })
if __name__ == '__main__':
diff --git a/cli/modules/interrogate-offline.py b/cli/modules/interrogate-offline.py
deleted file mode 100755
index f3120ae63..000000000
--- a/cli/modules/interrogate-offline.py
+++ /dev/null
@@ -1,166 +0,0 @@
-#!/usr/bin/env python
-
-import os
-import gc
-import json
-import time
-import argparse
-import torch
-import filetype
-from PIL import Image
-import transformers
-from transformers import AutoProcessor, AutoModelForCausalLM
-from transformers import BlipProcessor, BlipForConditionalGeneration
-from transformers import VisionEncoderDecoderModel, ViTFeatureExtractor, AutoTokenizer
-from util import log, Map
-
-
-model = None
-processor = None
-extractor = None
-dtype = torch.float32
-device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
-
-options = Map({
- 'input': '',
- 'min': 8,
- 'max': 256,
- 'beams': 1,
- 'json': '',
- 'txt': False,
- 'tag': '',
- 'git': True,
- 'blip': True,
- 'precision': 'fp16',
- 'model': 'git',
-})
-
-
-def cleanup(s: str):
- s = s.split('"')[0].split('.')[0].split(' that')[0]
- s = s.split(' with a letter')[0].split(' with the number')[0].split(' with the word')[0]
- s = s.replace('arafed image of ', '')
- return s.replace('a ', '')
-
-
-def load_model(args):
- global model
- global processor
- global extractor
- transformers.logging.set_verbosity_error()
- if args.model == 'git':
- model_name = "microsoft/git-large-textcaps"
- if model is None:
- model = AutoModelForCausalLM.from_pretrained(model_name, torch_dtype=dtype)
- model.to(device)
- processor = AutoProcessor.from_pretrained(model_name, torch_dtype=dtype)
- log.info( { 'interrogate loaded model': model_name })
- elif args.model == 'blip':
- model_name = "Salesforce/blip-image-captioning-large"
- if model is None:
- model = BlipForConditionalGeneration.from_pretrained(model_name, torch_dtype=dtype)
- model.to(device)
- processor = BlipProcessor.from_pretrained(model_name, torch_dtype=dtype)
- log.info( { 'interrogate loaded model': model_name })
- elif args.model == 'vit':
- model_name = "nlpconnect/vit-gpt2-image-captioning"
- if model is None:
- model = VisionEncoderDecoderModel.from_pretrained(model_name, torch_dtype=dtype)
- model.to(device)
- extractor = ViTFeatureExtractor.from_pretrained(model_name, torch_dtype=dtype)
- processor = AutoTokenizer.from_pretrained(model_name, torch_dtype=dtype)
- log.info( { 'interrogate loaded model': model_name })
- else:
- log.info( { 'interrogate unknown model': args.model })
-
-
-def interrogate_files(params, files):
- args = Map({**options, **params})
- data = [f for f in files if filetype.is_image(f)]
- log.info({ 'interrogate files': len(files), 'images': len(data), 'args': args })
- load_model(args)
- metadata = {}
- for image_path in data:
- image = Image.open(image_path).convert('RGB')
- caption = ''
- if args.model == 'git':
- inputs = processor(images=[image], return_tensors="pt").to(device)
- ids = model.generate(pixel_values=inputs.pixel_values, num_beams=args.beams, min_length=args.min, max_length=args.max)
- caption = processor.batch_decode(ids, skip_special_tokens=True)[0]
- elif args.model == 'blip':
- inputs = processor(image, return_tensors="pt").to(device, dtype)
- ids = model.generate(**inputs, num_beams=args.beams, min_length=args.min, max_length=args.max)
- caption = processor.decode(ids[0], skip_special_tokens=True)
- elif args.model == 'vit':
- inputs = extractor(images=[image], return_tensors="pt").pixel_values.to(device)
- ids = model.generate(inputs, num_beams=args.beams, min_length=args.min, max_length=args.max)
- caption = processor.batch_decode(ids, skip_special_tokens=True)[0]
- else:
- log.error({ 'interrogate unknown model': args.model })
-
- caption = cleanup(caption)
- tags = ''
- if args.tag != '':
- tags += args.tag + ','
- tags += caption.split(' ')[0]
- if args.txt:
- with open(os.path.splitext(image_path)[0] + '.txt', "wt", encoding='utf-8') as f:
- f.write(caption + "\n")
- metadata[image_path] = { 'caption': caption, 'tags': tags }
- log.info({ 'interrogate image': image_path, 'moodel': args.model, 'caption': caption, 'tags': tags })
-
- if args.json != '':
- with open(args.json, "wt", encoding='utf-8') as f:
- f.write(json.dumps(metadata, indent=2) + "\n")
- return metadata
-
-
-def unload_model():
- global processor
- global model
- global extractor
- if model is not None:
- del model
- model = None
- if processor is not None:
- del processor
- processor = None
- if extractor is not None:
- del extractor
- extractor = None
- gc.collect()
- if torch.cuda.is_available():
- with torch.no_grad():
- torch.cuda.empty_cache()
- with torch.cuda.device('cuda'):
- torch.cuda.empty_cache()
- torch.cuda.ipc_collect()
-
-
-if __name__ == '__main__':
- parser = argparse.ArgumentParser(description = 'image interrogate')
- parser.add_argument('input', type=str, nargs='*', help='input file or directory')
- parser.add_argument('--model', default = 'git', choices = ['git', 'blip', 'vit'], help = "which model to use")
- parser.add_argument("--min", type=int, default=8, help="min length of caption")
- parser.add_argument("--max", type=int, default=256, help="max length of caption")
- parser.add_argument("--beams", type=int, default=1, help="number of beams to use")
- parser.add_argument("--json", type=str, default='', help="output json file")
- parser.add_argument("--tag", type=str, default='', help="append tag")
- parser.add_argument('--txt', default = False, action='store_true', help = "write captions to text files")
- params = parser.parse_args()
- log.info({ 'interrogate args': vars(params) })
- if len(params.input) == 0:
- parser.print_help()
- exit(1)
- files = []
- for loc in params.input:
- if os.path.isfile(loc):
- files.append(loc)
- elif os.path.isdir(loc):
- for root, _sub_dirs, dir in os.walk(loc):
- files = [os.path.join(root, f) for f in dir]
- t0 = time.time()
- metadata = interrogate_files(vars(params), files)
- t1 = time.time()
- log.info({ 'interrogate files': len(files), 'time': round(t1 - t0, 2) })
- unload_model()
diff --git a/cli/modules/lora-extract.py b/cli/modules/lora-extract.py
deleted file mode 100755
index 7f85d37f3..000000000
--- a/cli/modules/lora-extract.py
+++ /dev/null
@@ -1,144 +0,0 @@
-#!/usr/bin/env python
-
-"""
-Extract approximating LoRA by SVD from two SD models
-Based on:
-"""
-
-import os
-import sys
-import time
-import argparse
-import torch
-import transformers
-from tqdm import tqdm
-from util import log
-
-sys.path.append(os.path.join(os.path.dirname(__file__), '..', '..', 'modules', 'lora'))
-import library.model_util as model_util
-import networks.lora as lora
-
-
-def svd(args): # pylint: disable=redefined-outer-name
- device = 'cuda' if torch.cuda.is_available() and args.device == 'cuda' else 'cpu'
- transformers.logging.set_verbosity_error()
- CLAMP_QUANTILE = 0.99
- MIN_DIFF = 1e-6
- if args.precision == 'fp32':
- save_dtype = torch.float
- elif args.precision == 'fp16':
- save_dtype = torch.float16
- elif args.precision == 'bf16':
- save_dtype = torch.bfloat16
- else:
- save_dtype = None
- t0 = time.time()
- log.info({ 'loading model': args.original })
- text_encoder_o, _, unet_o = model_util.load_models_from_stable_diffusion_checkpoint(args.v2, args.original)
- log.info({ 'loading model': args.tuned })
- text_encoder_t, _, unet_t = model_util.load_models_from_stable_diffusion_checkpoint(args.v2, args.tuned)
- with torch.no_grad():
- torch.cuda.empty_cache()
- # create LoRA network to extract weights: Use dim (rank) as alpha
- lora_network_o = lora.create_network(1.0, args.dim, args.dim, None, text_encoder_o, unet_o)
- lora_network_t = lora.create_network(1.0, args.dim, args.dim, None, text_encoder_t, unet_t)
- assert len(lora_network_o.text_encoder_loras) == len(lora_network_t.text_encoder_loras), 'model version is different'
- # get diffs
- diffs = {}
- text_encoder_different = False
- for i, (lora_o, lora_t) in enumerate(zip(lora_network_o.text_encoder_loras, lora_network_t.text_encoder_loras)):
- lora_name = lora_o.lora_name
- module_o = lora_o.org_module
- module_t = lora_t.org_module
- diff = module_t.weight - module_o.weight
- # Text Encoder might be same
- if torch.max(torch.abs(diff)) > MIN_DIFF:
- text_encoder_different = True
- diff = diff.float()
- diffs[lora_name] = diff
-
- if not text_encoder_different:
- log.info({ 'lora': 'text encoder is same, extract U-Net only' })
- lora_network_o.text_encoder_loras = []
- diffs = {}
-
- for i, (lora_o, lora_t) in enumerate(zip(lora_network_o.unet_loras, lora_network_t.unet_loras)):
- lora_name = lora_o.lora_name
- module_o = lora_o.org_module
- module_t = lora_t.org_module
- diff = module_t.weight - module_o.weight
- diff = diff.float()
- diff = diff.to(device)
- diffs[lora_name] = diff
- t1 = time.time()
- log.info({ 'lora models': 'ready', 'time': round(t1 - t0, 2) })
-
- # make LoRA with svd
- log.info({ 'lora': 'calculating by svd' })
- rank = args.dim
- lora_weights = {}
- with torch.no_grad():
- for lora_name, mat in tqdm(list(diffs.items())):
- conv2d = len(mat.size()) == 4
- if conv2d:
- mat = mat.squeeze()
- U, S, Vh = torch.linalg.svd(mat)
- U = U[:, :rank]
- S = S[:rank]
- U = U @ torch.diag(S)
- Vh = Vh[:rank, :]
- dist = torch.cat([U.flatten(), Vh.flatten()])
- hi_val = torch.quantile(dist, CLAMP_QUANTILE)
- low_val = -hi_val
- U = U.clamp(low_val, hi_val)
- Vh = Vh.clamp(low_val, hi_val)
- lora_weights[lora_name] = (U, Vh)
- t2 = time.time()
-
- # make state dict for LoRA
- lora_network_o.apply_to(text_encoder_o, unet_o, text_encoder_different, True)
- lora_sd = lora_network_o.state_dict()
- log.info({ 'lora extracted weights': len(lora_sd), 'time': round(t2 - t1, 2) })
-
- for key in list(lora_sd.keys()):
- if 'alpha' in key:
- continue
- lora_name = key.split('.')[0]
- i = 0 if 'lora_up' in key else 1
- weights = lora_weights[lora_name][i]
- # print(key, i, weights.size(), lora_sd[key].size())
- if len(lora_sd[key].size()) == 4: # pylint: disable=unsubscriptable-object
- weights = weights.unsqueeze(2).unsqueeze(3)
- assert weights.size() == lora_sd[key].size(), f'size unmatch: {key}' # pylint: disable=unsubscriptable-object
- lora_sd[key] = weights # pylint: disable=unsupported-assignment-operation
-
- # load state dict to LoRA and save it
- info = lora_network_o.load_state_dict(lora_sd)
- log.info({ 'lora loading extracted weights': info })
-
- dir_name = os.path.dirname(args.save)
- if dir_name and not os.path.exists(dir_name):
- os.makedirs(dir_name, exist_ok=True)
-
- # minimum metadata
- metadata = {'ss_network_dim': str(args.dim), 'ss_network_alpha': str(args.dim)}
- lora_network_o.save_weights(args.save, save_dtype, metadata)
- t3 = time.time()
- log.info({ 'lora saved weights': args.save, 'time': round(t3 - t2, 2) })
-
-
-if __name__ == '__main__':
- parser = argparse.ArgumentParser(description = 'extract lora weights')
- parser.add_argument('--v2', action='store_true', help='load Stable Diffusion v2.x model / Stable Diffusion')
- parser.add_argument('--precision', type=str, default='fp16', choices=[None, 'fp32', 'fp16', 'bf16'], help='precision in saving, same to merging if omitted')
- parser.add_argument('--device', type=str, default='cuda', choices=['cpu', 'cuda'], help='use cpu or cuda if available')
- parser.add_argument('--original', type=str, default=None, required=True, help='Stable Diffusion original model: ckpt or safetensors file')
- parser.add_argument('--tuned', type=str, default=None, required=True, help='Stable Diffusion tuned model, LoRA is difference of `original to tuned`: ckpt or safetensors file')
- parser.add_argument('--save', type=str, default=None, required=True, help='destination file name: ckpt or safetensors file')
- parser.add_argument('--dim', type=int, default=4, help='dimension (rank) of LoRA')
- args = parser.parse_args()
- log.info({ 'extract lora args': vars(args) })
- if not os.path.exists(args.original) or not os.path.exists(args.tuned):
- log.error({ 'models not found': [args.original, args.tuned] })
- else:
- svd(args)
diff --git a/cli/modules/lora-latents.py b/cli/modules/lora-latents.py
deleted file mode 100755
index 4e1027f14..000000000
--- a/cli/modules/lora-latents.py
+++ /dev/null
@@ -1,160 +0,0 @@
-#!/usr/bin/env python
-
-import os
-import sys
-import json
-import pathlib
-import argparse
-import warnings
-
-import cv2
-import numpy as np
-import torch
-from PIL import Image
-from torchvision import transforms
-from tqdm import tqdm
-from util import log, Map
-
-sys.path.append(os.path.join(os.path.dirname(__file__), '..', '..', 'modules', 'lora'))
-import library.model_util as model_util
-import library.train_util as train_util
-
-warnings.filterwarnings('ignore')
-device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
-options = Map({
- 'batch': 1,
- 'input': '',
- 'json': '',
- 'max': 1024,
- 'min': 256,
- 'noupscale': False,
- 'precision': 'fp32',
- 'resolution': '512,512',
- 'steps': 64,
- 'vae': 'stabilityai/sd-vae-ft-mse'
-})
-vae = None
-
-
-def get_latents(vae, images, weight_dtype):
- image_transforms = transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.5], [0.5]) ])
- img_tensors = [image_transforms(image) for image in images]
- img_tensors = torch.stack(img_tensors)
- img_tensors = img_tensors.to(device, weight_dtype)
- with torch.no_grad():
- latents = vae.encode(img_tensors).latent_dist.sample().float().to('cpu').numpy()
- return latents
-
-
-def get_npz_filename_wo_ext(data_dir, image_key):
- return os.path.join(data_dir, os.path.splitext(os.path.basename(image_key))[0])
-
-
-def create_vae_latents(params):
- args = Map({**options, **params})
- log.info({ 'latents args': args })
- if args.steps % 8 > 0:
- log.warning({ 'latents': 'resolution is not multiple of 8' })
- image_paths = train_util.glob_images(args.input)
- if os.path.exists(args.json):
- log.info({ 'latents metadata': args.json, 'images': len(image_paths) })
- with open(args.json, 'rt', encoding='utf-8') as f:
- metadata = json.load(f)
- else:
- log.error({ 'latents metadata missing': args.json })
- return
- if args.precision == 'fp16':
- weight_dtype = torch.float16
- elif args.precision == 'bf16':
- weight_dtype = torch.bfloat16
- else:
- weight_dtype = torch.float32
- global vae
- if vae is None:
- vae = model_util.load_vae(args.vae, weight_dtype)
- vae.eval()
- vae.to(device, dtype=weight_dtype)
- max_reso = tuple([int(t) for t in args.resolution.split(',')])
- assert len(max_reso) == 2, f'illegal resolution: {args.resolution}'
- bucket_manager = train_util.BucketManager(args.noupscale, max_reso, args.min, args.max, args.steps)
- if not args.noupscale:
- bucket_manager.make_buckets()
- else:
- log.warning({ 'latents': 'min and max are ignored if noupscale is set' })
- img_ar_errors = []
- 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)
- 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)
- bucket.clear()
- data = [[(None, ip)] for ip in image_paths]
- bucket_counts = {}
- for data_entry in tqdm(data, smoothing=0.0):
- if data_entry[0] is None:
- continue
- img_tensor, image_path = data_entry[0]
- if img_tensor is not None:
- image = transforms.functional.to_pil_image(img_tensor)
- else:
- image = Image.open(image_path)
- image_key = os.path.basename(image_path)
- image_key = os.path.join(os.path.basename(pathlib.Path(image_path).parent), pathlib.Path(image_path).stem)
- if image_key not in metadata:
- metadata[image_key] = {}
- reso, resized_size, ar_error = bucket_manager.select_bucket(image.width, image.height)
- img_ar_errors.append(abs(ar_error))
- bucket_counts[reso] = bucket_counts.get(reso, 0) + 1
- metadata[image_key]['train_resolution'] = (reso[0] - reso[0] % 8, reso[1] - reso[1] % 8)
- if not args.noupscale:
- assert resized_size[0] == reso[0] or resized_size[1] == reso[1], f'internal error, resized size not match: {reso}, {resized_size}, {image.width}, {image.height}'
- assert resized_size[0] >= reso[0] and resized_size[1] >= reso[1], f'internal error, resized size too small: {reso}, {resized_size}, {image.width}, {image.height}'
- assert resized_size[0] >= reso[0] and resized_size[1] >= reso[1], f'internal error resized size is small: {resized_size}, {reso}'
- image = np.array(image)
- if resized_size[0] != image.shape[1] or resized_size[1] != image.shape[0]:
- image = cv2.resize(image, resized_size, interpolation=cv2.INTER_AREA)
- if resized_size[0] > reso[0]:
- trim_size = resized_size[0] - reso[0]
- image = image[:, trim_size//2:trim_size//2 + reso[0]]
- if resized_size[1] > reso[1]:
- trim_size = resized_size[1] - reso[1]
- image = image[trim_size//2:trim_size//2 + reso[1]]
- assert image.shape[0] == reso[1] and image.shape[1] == reso[0], f'internal error, illegal trimmed size: {image.shape}, {reso}'
- bucket_manager.add_image(reso, (image_key, image))
- process_batch(False)
-
- process_batch(True)
- vae.to('cpu')
-
- bucket_manager.sort()
- img_ar_errors = np.array(img_ar_errors)
- for i, reso in enumerate(bucket_manager.resos):
- count = bucket_counts.get(reso, 0)
- if count > 0:
- log.info({ 'latents bucket': i, 'resolution': reso, 'count': count, 'mean ar error': np.mean(img_ar_errors) })
- with open(args.json, 'wt', encoding='utf-8') as f:
- json.dump(metadata, f, indent=2)
-
-
-def unload_vae():
- global vae
- vae = None
-
-
-if __name__ == '__main__':
- parser = argparse.ArgumentParser()
- parser.add_argument('input', type=str, help='directory for train images')
- parser.add_argument('--json', type=str, required=True, help='metadata file to input')
- parser.add_argument('--vae', type=str, required=True, help='model name or path to encode latents')
- parser.add_argument('--batch', type=int, default=1, help='batch size in inference')
- parser.add_argument('--resolution', type=str, default='512,512', help='max resolution in fine tuning (width,height)')
- parser.add_argument('--min', type=int, default=256, help='minimum resolution for buckets')
- parser.add_argument('--max', type=int, default=1024, help='maximum resolution for buckets')
- parser.add_argument('--steps', type=int, default=64, help='steps of resolution for buckets, divisible by 8')
- parser.add_argument('--noupscale', action='store_true', help='make bucket for each image without upscaling')
- parser.add_argument('--precision', type=str, default='fp32', choices=['fp32', 'fp16', 'bf16'], help='use precision')
- params = parser.parse_args()
- create_vae_latents(vars(params))
diff --git a/cli/modules/models-diff.py b/cli/modules/models-diff.py
deleted file mode 100755
index fdeb235e6..000000000
--- a/cli/modules/models-diff.py
+++ /dev/null
@@ -1,74 +0,0 @@
-#!/usr/bin/env python
-# based on
-
-import safetensors
-import sys
-import torch
-from pathlib import Path
-import torch.nn as nn
-import torch.nn.functional as F
-import warnings
-from util import log
-
-warnings.filterwarnings("ignore", category=UserWarning)
-
-def cal_cross_attn(to_q, to_k, to_v, rand_input):
- hidden_dim, embed_dim = to_q.shape
- attn_to_q = nn.Linear(hidden_dim, embed_dim, bias=False)
- attn_to_k = nn.Linear(hidden_dim, embed_dim, bias=False)
- attn_to_v = nn.Linear(hidden_dim, embed_dim, bias=False)
- attn_to_q.load_state_dict({"weight": to_q})
- attn_to_k.load_state_dict({"weight": to_k})
- attn_to_v.load_state_dict({"weight": to_v})
-
- return torch.einsum(
- "ik, jk -> ik",
- F.softmax(torch.einsum("ij, kj -> ik", attn_to_q(rand_input), attn_to_k(rand_input)), dim=-1),
- attn_to_v(rand_input)
- )
-
-def load_model(path):
- if path.suffix == ".safetensors":
- return safetensors.torch.load_file(path, device="cpu")
- else:
- ckpt = torch.load(path, map_location="cpu")
- return ckpt["state_dict"] if "state_dict" in ckpt else ckpt
-
-def eval(model, n, input):
- qk = f"model.diffusion_model.output_blocks.{n}.1.transformer_blocks.0.attn1.to_q.weight"
- uk = f"model.diffusion_model.output_blocks.{n}.1.transformer_blocks.0.attn1.to_k.weight"
- vk = f"model.diffusion_model.output_blocks.{n}.1.transformer_blocks.0.attn1.to_v.weight"
- atoq, atok, atov = model[qk], model[uk], model[vk]
- attn = cal_cross_attn(atoq, atok, atov, input)
- return attn
-
-def main():
- file1 = Path(sys.argv[1])
- files = sys.argv[2:]
- seed = 114514
- torch.manual_seed(seed)
- model_a = load_model(file1)
- log.info(f"base: {file1.name}")
-
- map_attn_a = {}
- map_rand_input = {}
- for n in range(3, 11):
- hidden_dim, embed_dim = model_a[f"model.diffusion_model.output_blocks.{n}.1.transformer_blocks.0.attn1.to_q.weight"].shape
- rand_input = torch.randn([embed_dim, hidden_dim])
- map_attn_a[n] = eval(model_a, n, rand_input)
- map_rand_input[n] = rand_input
- del model_a
-
- for file2 in files:
- file2 = Path(file2)
- model_b = load_model(file2)
- sims = []
- for n in range(3, 11):
- attn_a = map_attn_a[n]
- attn_b = eval(model_b, n, map_rand_input[n])
- sim = torch.mean(torch.cosine_similarity(attn_a, attn_b))
- sims.append(sim)
- log.info(f"{file2}: {torch.mean(torch.stack(sims)) * 1e2:.2f}%")
-
-if __name__ == "__main__":
- main()
diff --git a/cli/modules/preview-embeddings.py b/cli/modules/preview-embeddings.py
deleted file mode 100755
index a64297925..000000000
--- a/cli/modules/preview-embeddings.py
+++ /dev/null
@@ -1,101 +0,0 @@
-#!/usr/bin/env python
-"""
-create preview images from embeddings
-"""
-import os
-import io
-import sys
-import json
-import base64
-import argparse
-from pathlib import Path
-from PIL import Image
-from inspect import getsourcefile
-from util import Map, log
-from sdapi import getsync, postsync
-from grid import grid
-
-template = 'photo of "{name}", {suffix}, high detailed, skin texture, looking forward, facing camera, 135mm, shot on dslr, 4k, modelshoot style'
-img2img_options = Map({
- 'prompt': None,
- 'negative_prompt': 'cartoon, drawing, cgi, sketch, comic, disfigured, deformed',
- 'init_images': [],
- 'sampler_name': 'DPM2 Karras',
- 'batch_size': 4,
- 'n_iter': 1,
- 'steps': 30,
- 'cfg_scale': 6,
- 'width': 512,
- 'height': 512,
- 'restore_faces': False
-})
-
-def encode(f):
- img = Image.open(f)
- with io.BytesIO() as stream:
- img.save(stream, 'JPEG')
- values = stream.getvalue()
- encoded = base64.b64encode(values).decode()
- return encoded
-
-def create_preview(name: str, suffix: str):
- options = getsync('/sdapi/v1/options')
- cmdflags = getsync('/sdapi/v1/cmd-flags')
- img2img_options['prompt'] = template.format(name = name, suffix = suffix)
- log.debug({ 'preview options': img2img_options })
- if len(img2img_options['init_images']) == 0:
- for i in range(img2img_options.batch_size):
- mask = os.path.join(os.path.dirname(getsourcefile(lambda:0)), 'preview-template'+ str(i+1) +'.jpg')
- if (not os.path.isfile(mask)):
- log.error({ 'preview': 'missing preview mask' })
- return
- img2img_options['init_images'].append(encode(mask))
- data = postsync('/sdapi/v1/img2img', img2img_options)
- if 'error' in data:
- log.error({ 'preview': data['error'], 'reason': data['reason'] })
- return
- info = Map(json.loads(data['info']))
- if not 'images' in data:
- log.error({ 'preview': 'no images' })
- return
- fn = os.path.join(cmdflags.embeddings_dir, name + '.preview.png')
- log.info({ 'preview': { 'name': fn, 'model': options.sd_model_checkpoint, 'seed': info.seed } })
- images = []
- for b64 in data['images']:
- images.append(Image.open(io.BytesIO(base64.b64decode(b64.split(",",1)[0]))))
- image = grid(images, None, square=True)
- image.save(fn)
-
-if __name__ == "__main__":
- log.info({ 'preview': 'start' })
- cmdflags = getsync('/sdapi/v1/cmd-flags')
-
- parser = argparse.ArgumentParser(description = 'generate embeddings previews')
- parser.add_argument('--overwrite', default = False, action='store_true', help = 'overwrite existing previews')
- parser.add_argument('input', type=str, nargs='*')
- params = parser.parse_args()
-
- if len(params.input) == 0:
- files = list(Path(cmdflags.embeddings_dir).glob('*.pt'))
- else:
- files = list(os.path.join(cmdflags.embeddings_dir, a + '.pt') for a in params.input if os.path.isfile(os.path.join(cmdflags.embeddings_dir, a + '.pt')))
- candidates = [str(f) for f in files]
- candidates.sort(key=os.path.getctime, reverse=True)
-
- files = []
- for f in candidates:
- fn = f.replace('.pt', '.preview.png')
- if os.path.isfile(f.replace('.pt', '.preview.png')):
- if params.overwrite:
- log.info({ 'preview add': fn })
- files.append(f)
- else:
- log.info({ 'preview skip': fn })
- else:
- log.info({ 'preview add': fn })
- files.append(f)
-
- log.info({ 'preview embeddings': len(files) })
- for f in files:
- name = Path(f).stem
- create_preview(name, 'person')
diff --git a/cli/modules/preview-template1.jpg b/cli/modules/preview-template1.jpg
deleted file mode 100644
index d6da3f5e9..000000000
Binary files a/cli/modules/preview-template1.jpg and /dev/null differ
diff --git a/cli/modules/preview-template2.jpg b/cli/modules/preview-template2.jpg
deleted file mode 100644
index adcdd3487..000000000
Binary files a/cli/modules/preview-template2.jpg and /dev/null differ
diff --git a/cli/modules/preview-template3.jpg b/cli/modules/preview-template3.jpg
deleted file mode 100644
index 4982d127a..000000000
Binary files a/cli/modules/preview-template3.jpg and /dev/null differ
diff --git a/cli/modules/preview-template4.jpg b/cli/modules/preview-template4.jpg
deleted file mode 100644
index 140dc9aaa..000000000
Binary files a/cli/modules/preview-template4.jpg and /dev/null differ
diff --git a/cli/modules/process.py b/cli/modules/process.py
deleted file mode 100755
index 975f9d1d6..000000000
--- a/cli/modules/process.py
+++ /dev/null
@@ -1,500 +0,0 @@
-#!/usr/bin/env python
-"""
-process people images
-- check image resolution
-- runs detection of face and body
-- extracts crop and performs checks:
- - visible: is face or body detected
- - in frame: for face based on box, for body based on number of visible keypoints
- - resolution: is cropped image still of sufficient resolution
- - optionaly upsample and restore face quality
- - blur: is image sharp enough
- - dynamic range: is image bright enough
- - similarity: compares image to all previously processed images to see if its unique enough
-- images are resized and optionally squared
-- face additionally runs through semantic segmentation to remove background
-- if image passes checks
- image padded and saved as extracted image
-- body requires that face is detected and in-frame,
- but does not have to pass all other checks as body performs its own checks
-- runs clip interrogation on extracted images to generate filewords
-"""
-
-import os
-import sys
-import io
-import math
-import base64
-import pathlib
-import argparse
-import logging
-import filetype
-import numpy as np
-import mediapipe as mp
-from PIL import Image, ImageOps
-from skimage.metrics import structural_similarity as ssim
-from scipy.stats import beta
-sys.path.append(os.path.join(os.path.dirname(__file__)))
-
-from util import log, Map
-from sdapi import postsync
-
-
-params = Map({
- # general settings, do not modify
- 'src': '', # source folder
- 'dst': '', # destination folder
- 'clear_dst': True, # remove all files from destination at the start
- 'format': '.jpg', # image format
- 'target_size': 512, # target resolution
- 'square_images': True, # should output images be squared
- 'segmentation_model': 0, # segmentation model 0/general 1/landscape
- 'segmentation_background': (192, 192, 192), # segmentation background color
- 'blur_samplesize': 60, # sample size to use for blur detection
- 'similarity_size': 64, # base similarity detection on reduced images
- # original image processing settings
- 'keep_original': False, # keep original image
- # face processing settings
- 'extract_face': False, # extract face from image
- 'face_score': 0.7, # min face detection score
- 'face_pad': 0.1, # pad face image percentage
- 'face_model': 1, # which face model to use 0/close-up 1/standard
- 'face_blur': False, # check for body blur
- 'face_blur_score': 1.5, # max score for face blur detection
- 'face_range': False, # check for body blur
- 'face_range_score': 0.15, # min score for face dynamic range detection
- 'face_restore': False, # attempt to restore face quality
- 'face_upscale': False, # attempt to scale small faces
- 'face_segmentation': False, # segmentation enabled
- # body processing settings
- 'extract_body': False, # extract body from image
- 'body_score': 0.9, # min body detection score
- 'body_visibility': 0.5, # min visibility score for each detected body part
- 'body_parts': 15, # min number of detected body parts with sufficient visibility
- 'body_pad': 0.2, # pad body image percentage
- 'body_model': 2, # body model to use 0/low 1/medium 2/high
- 'body_blur': False, # check for body blur
- 'body_blur_score': 1.8, # max score for body blur detection
- 'body_range': False, # check for body blur
- 'body_range_score': 0.15, # min score for body dynamic range detection
- 'body_segmentation': False, # segmentation enabled
- # similarity detection settings
- 'similarity_score': 0.8, # maximum similarity score before image is discarded
- # interrogate settings
- 'interrogate_model': ['clip', 'deepdanbooru'], # interrogate models
- 'interrogate_captions': True, # write captions to file
- 'tag_limit': 5, # number of tags to extract
-})
-face_model = None
-body_model = None
-segmentation_model = None
-
-
-def detect_blur(image):
- # based on
- bw = ImageOps.grayscale(image)
- cx, cy = image.size[0] // 2, image.size[1] // 2
- fft = np.fft.fft2(bw)
- fftShift = np.fft.fftshift(fft)
- fftShift[cy - params.blur_samplesize: cy + params.blur_samplesize, cx - params.blur_samplesize: cx + params.blur_samplesize] = 0
- fftShift = np.fft.ifftshift(fftShift)
- recon = np.fft.ifft2(fftShift)
- magnitude = np.log(np.abs(recon))
- mean = round(np.mean(magnitude), 2)
- return mean
-
-
-def detect_dynamicrange(image):
- # based on
- data = np.asarray(image)
- image = np.float32(data)
- RGB = [0.299, 0.587, 0.114]
- height, width = image.shape[:2]
- brightness_image = np.sqrt(image[..., 0] ** 2 * RGB[0] + image[..., 1] ** 2 * RGB[1] + image[..., 2] ** 2 * RGB[2])
- hist, _ = np.histogram(brightness_image, bins=256, range=(0, 255))
- img_brightness_pmf = hist / (height * width)
- dist = beta(2, 2)
- ys = dist.pdf(np.linspace(0, 1, 256))
- ref_pmf = ys / np.sum(ys)
- dot_product = np.dot(ref_pmf, img_brightness_pmf)
- squared_dist_a = np.sum(ref_pmf ** 2)
- squared_dist_b = np.sum(img_brightness_pmf ** 2)
- res = dot_product / math.sqrt(squared_dist_a * squared_dist_b)
- return round(res, 2)
-
-
-images = []
-def detect_simmilar(image):
- img = image.resize((params.similarity_size, params.similarity_size))
- img = ImageOps.grayscale(img)
- data = np.array(img)
- similarity = 0
- for i in images:
- val = ssim(data, i, data_range=255, channel_axis=None, gradient=False, full=False)
- if val > similarity:
- similarity = val
- images.append(data)
- return similarity
-
-
-def segmentation(image):
- global segmentation_model
- if segmentation_model is None:
- segmentation_model = mp.solutions.selfie_segmentation.SelfieSegmentation(model_selection=params.segmentation_model)
- data = np.array(image)
- results = segmentation_model.process(data)
- condition = np.stack((results.segmentation_mask,) * 3, axis=-1) > 0.1
- background = np.zeros(data.shape, dtype=np.uint8)
- background[:] = params.segmentation_background
- data = np.where(condition, data, background) # consider using a joint bilateral filter instead of pure combine
- segmented = Image.fromarray(data)
- return segmented
-
-
-def extract_face(img):
- if not params.extract_face:
- return None, True
- if img.mode == 'RGBA':
- img = img.convert('RGB')
- scale = max(img.size[0], img.size[1]) / params.target_size
- resized = img.copy()
- resized.thumbnail((params.target_size, params.target_size), Image.HAMMING)
-
- global face_model
- if face_model is None:
- face_model = mp.solutions.face_detection.FaceDetection(min_detection_confidence=params.face_score, model_selection=params.face_model)
- results = face_model.process(np.array(resized))
- if results.detections is None:
- return None, False
- box = results.detections[0].location_data.relative_bounding_box
- if box.xmin < 0 or box.ymin < 0 or (box.width - box.xmin) > 1 or (box.height - box.ymin) > 1:
- log.info({ 'process face skip': 'out of frame' })
- return None, False
- x = (box.xmin - params.face_pad / 2) * resized.width
- y = (box.ymin - params.face_pad / 2)* resized.height
- w = (box.width + params.face_pad) * resized.width
- h = (box.height + params.face_pad) * resized.height
- cx = x + w / 2
- cy = y + h / 2
- l = max(w, h) / 2
- square = [scale * (cx - l), scale * (cy - l), scale * (cx + l), scale * (cy + l)]
- square = [max(square[0], 0), max(square[1], 0), min(square[2], img.width), min(square[3], img.height)]
- cropped = img.crop(tuple(square))
-
- upscale = 1
- if params.face_restore or params.face_upscale:
- if (cropped.size[0] < params.target_size or cropped.size[1] < params.target_size) and params.face_upscale:
- upscale = 2
- kwargs = Map({
- 'image': encode(cropped),
- 'upscaler_1': 'SwinIR_4x' if params.face_upscale else None,
- 'codeformer_visibility': 1.0 if params.face_restore else 0.0,
- 'codeformer_weight': 0.15 if params.face_restore else 0.0,
- 'upscaling_resize': upscale,
- })
- original = [cropped.size[0], cropped.size[1]]
- res = postsync('/sdapi/v1/extra-single-image', kwargs)
- if 'image' not in res:
- log.error({ 'process face': 'upscale failed' })
- raise ValueError('upscale failed')
- cropped = Image.open(io.BytesIO(base64.b64decode(res['image'])))
- kwargs.image = [cropped.size[0], cropped.size[1]]
- upscaled = [cropped.size[0], cropped.size[1]]
- upscale = False if upscale == 1 else { 'original': original, 'upscaled': upscaled }
- log.info({ 'process face restore': params.face_restore, 'upscale': upscale })
-
- if cropped.size[0] < params.target_size and cropped.size[1] < params.target_size:
- log.info({ 'process face skip': 'low resolution', 'size': [cropped.size[0], cropped.size[1]] })
- return None, True
- cropped.thumbnail((params.target_size, params.target_size), Image.HAMMING)
-
- if params.square_images:
- squared = Image.new('RGB', (params.target_size, params.target_size))
- squared.paste(cropped, ((params.target_size - cropped.width) // 2, (params.target_size - cropped.height) // 2))
- if params.face_segmentation:
- squared = segmentation(squared)
- else:
- squared = cropped
-
- if params.face_blur:
- blur = detect_blur(squared)
- if blur > params.face_blur_score:
- log.info({ 'process face skip': 'blur check fail', 'blur': blur })
- return None, True
- else:
- log.debug({ 'process face blur': blur })
-
- if params.face_range:
- range = detect_dynamicrange(squared)
- if range < params.face_range_score:
- log.info({ 'process face skip': 'dynamic range check fail', 'range': range })
- return None, True
- else:
- log.debug({ 'process face dynamic range': range })
-
- similarity = detect_simmilar(squared)
- if similarity > params.similarity_score:
- log.info({ 'process face skip': 'similarity check fail', 'score': round(similarity, 2) })
- return None, True
-
- return squared, True
-
-
-def extract_body(img):
- if not params.extract_body:
- return None, True
- if img.mode == 'RGBA':
- img = img.convert('RGB')
- scale = max(img.size[0], img.size[1]) / params.target_size
- resized = img.copy()
- resized.thumbnail((params.target_size, params.target_size), Image.HAMMING)
-
- global body_model
- if body_model is None:
- body_model = mp.solutions.pose.Pose(static_image_mode=True, min_detection_confidence=params.body_score, model_complexity=params.body_model)
- results = body_model.process(np.array(resized))
- if results.pose_landmarks is None:
- return None, False
- x = [resized.width * (i.x - params.body_pad / 2) for i in results.pose_landmarks.landmark if i.visibility > params.body_visibility]
- y = [resized.height * (i.y - params.body_pad / 2) for i in results.pose_landmarks.landmark if i.visibility > params.body_visibility]
- if len(x) < params.body_parts:
- log.info({ 'process body skip': 'insufficient body parts', 'detected': len(x) })
- return None, True
- w = max(x) - min(x) + resized.width * params.body_pad
- h = max(y) - min(y) + resized.height * params.body_pad
- cx = min(x) + w / 2
- cy = min(y) + h / 2
- l = max(w, h) / 2
- square = [scale * (cx - l), scale * (cy - l), scale * (cx + l), scale * (cy + l)]
- square = [max(square[0], 0), max(square[1], 0), min(square[2], img.width), min(square[3], img.height)]
- cropped = img.crop(tuple(square))
- if cropped.size[0] < params.target_size and cropped.size[1] < params.target_size:
- log.info({ 'process body skip': 'low resolution', 'size': [cropped.size[0], cropped.size[1]] })
- return None, True
- cropped.thumbnail((params.target_size, params.target_size), Image.HAMMING)
-
- if params.square_images:
- squared = Image.new('RGB', (params.target_size, params.target_size))
- squared.paste(cropped, ((params.target_size - cropped.width) // 2, (params.target_size - cropped.height) // 2))
- if params.body_segmentation:
- squared = segmentation(squared)
- else:
- squared = cropped
-
- if params.body_blur:
- blur = detect_blur(squared)
- if blur > params.body_blur_score:
- log.info({ 'process body skip': 'blur check fail', 'blur': blur })
- return None, True
- else:
- log.debug({ 'process body blur': blur })
-
- if params.body_range:
- range = detect_dynamicrange(squared)
- if range < params.body_range_score:
- log.info({ 'process body skip': 'dynamic range check fail', 'range': range })
- return None, True
- else:
- log.debug({ 'process body dynamic range': range })
-
- similarity = detect_simmilar(squared)
- if similarity > params.similarity_score:
- log.info({ 'process body skip': 'similarity check fail', 'score': round(similarity, 2) })
- return None, True
-
- return squared, True
-
-
-def save_original(img):
- if img.mode == 'RGBA':
- img = img.convert('RGB')
- resized = img.copy()
- resized.thumbnail((params.target_size, params.target_size), Image.HAMMING)
- if params.square_images:
- squared = Image.new('RGB', (params.target_size, params.target_size))
- squared.paste(resized, ((params.target_size - resized.width) // 2, (params.target_size - resized.height) // 2))
- else:
- squared = resized
- return squared
-
-
-def encode(img):
- with io.BytesIO() as stream:
- img.save(stream, 'JPEG')
- values = stream.getvalue()
- encoded = base64.b64encode(values).decode()
- return encoded
-
-
-def interrogate(img, fn, intag = None):
- if len(params.interrogate_model) == 0:
- return
- caption = ''
- tags = []
- for model in params.interrogate_model:
- json = Map({ 'image': encode(img), 'model': model })
- res = postsync('/sdapi/v1/interrogate', json)
- if model == 'clip':
- caption = res.caption if 'caption' in res else ''
- caption = caption.split(',')[0].replace('a ', '')
- if intag is not None:
- caption = intag + ', ' + caption
- if model == 'deepdanbooru':
- tag = res.caption if 'caption' in res else ''
- tags = tag.split(',')
- tags = [t.replace('(', '').replace(')', '').replace('\\', '').split(':')[0].strip() for t in tags]
- if intag is not None:
- for t in intag.split(',')[::-1]:
- tags.insert(0, t.strip())
- if params.interrogate_captions:
- file = fn.replace(params.format, '.txt')
- f = open(file, 'w')
- f.write(caption)
- f.close()
- pos = 0 if len(tags) == 0 else 1
- tags.insert(pos, caption.split(' ')[1])
- if len(tags) > params.tag_limit:
- tags = tags[:params.tag_limit]
- log.info({ 'interrogate': caption, 'tags': tags })
- return caption, tags
-
-
-i = {}
-metadata = Map({})
-
-# entry point when used as module
-def process_file(f: str, dst: str = None, preview: bool = False, offline: bool = False, txt = None, tag = None, opts = []):
- def save(img, f, what):
- i[what] = i.get(what, 0) + 1
- if dst is None:
- dir = os.path.dirname(f)
- else:
- dir = dst
- base = os.path.basename(f).split('.')[0]
- parent = os.path.basename(pathlib.Path(dir))
- basename = str(i[what]).rjust(3, '0') + '-' + what + '-' + base
- fn = basename + params.format
- # log.debug({ 'save': fn })
- caption = ''
- tags = ''
- if not preview:
- img.save(os.path.join(dir, fn))
- if not offline:
- caption, tags = interrogate(img, os.path.join(dir, fn), tag)
- metadata[os.path.join(parent, basename)] = { 'caption': caption, 'tags': ','.join(tags) }
- return fn
-
- # overrides
- if len(opts) > 0:
- params.keep_original = True if 'original' in opts else False
- params.extract_face = True if 'face' in opts else False
- params.extract_body = True if 'body' in opts else False
- params.face_blur = True if 'blur' in opts else False
- params.body_blur = True if 'blur' in opts else False
- params.face_range = True if 'range' in opts else False
- params.body_range = True if 'range' in opts else False
- params.face_upscale = True if 'upscale' in opts else False
- params.face_restore = True if 'restore' in opts else False
-
- log.info({ 'processing': f })
- try:
- image = Image.open(f)
- except Exception as err:
- log.error({ 'image': f, 'error': err })
- return 0, {}
-
- image = ImageOps.exif_transpose(image) # rotate image according to EXIF orientation
- if txt is not None:
- params.interrogate_captions = txt
-
- if image.width < 512 or image.height < 512:
- log.info({ 'process skip': 'low resolution', 'resolution': [image.width, image.height] })
- return 0, {}
- log.debug({ 'resolution': [image.width, image.height], 'mp': round((image.width * image.height) / 1024 / 1024, 1) })
-
- face, ok = extract_face(image)
- if face is not None:
- fn = save(face, f, 'face')
- log.info({ 'extract face': fn })
- else:
- log.debug({ 'no face': f })
-
- if not ok:
- return 0, {}
-
- body, ok = extract_body(image)
- if body is not None:
- fn = save(body, f, 'body')
- log.info({ 'extract body': fn })
- else:
- log.debug({ 'no body': f })
-
- if params.keep_original:
- resized = save_original(image)
- fn = save(resized, f, 'original')
- log.info({ 'original': fn })
-
- image.close()
- return i, metadata
-
-def process_images(src: str, dst: str, args = None):
- params.src = src
- params.dst = dst
- if args is not None:
- params.update(args)
- log.info({ 'processing': params })
- if not os.path.isdir(src):
- log.error({ 'process': 'not a folder', 'src': src })
- else:
- if os.path.isdir(dst) and params.clear_dst:
- log.info({ 'clear dst': dst })
- i = [os.path.join(dst, f) for f in os.listdir(dst) if os.path.isfile(os.path.join(dst, f)) and filetype.is_image(os.path.join(dst, f))]
- for f in i:
- os.remove(f)
- pathlib.Path(dst).mkdir(parents=True, exist_ok=True)
- for root, _sub_dirs, files in os.walk(src):
- for f in files:
- i, _metadata = process_file(os.path.join(root, f), dst)
- return i
-
-
-def unload_models():
- global face_model
- if face_model is not None:
- face_model = None
- global body_model
- if body_model is not None:
- body_model = None
- global segmentation_model
- if segmentation_model is not None:
- segmentation_model = None
-
-
-if __name__ == '__main__':
- # log.setLevel(logging.DEBUG)
- parser = argparse.ArgumentParser(description = 'dataset processor')
- parser.add_argument('--output', type=str, required=True, help='folder to store images')
- parser.add_argument('--preview', default=False, action='store_true', help = "run processing but do not store results")
- parser.add_argument('--offline', default=False, action='store_true', help = "run only processing steps that do not require running server")
- parser.add_argument('--debug', default=False, action='store_true', help = "enable debug logging")
- parser.add_argument('input', type=str, nargs='*')
- args = parser.parse_args()
- params.dst = args.output
- if args.debug:
- log.setLevel(logging.DEBUG)
- log.debug({ 'debug': True })
- log.info({ 'processing': params })
- if not os.path.exists(params.dst) and not args.preview:
- pathlib.Path(params.dst).mkdir(parents=True, exist_ok=True)
- files = []
- for loc in args.input:
- if os.path.isfile(loc):
- files.append(loc)
- elif os.path.isdir(loc):
- for root, _sub_dirs, dir in os.walk(loc):
- for f in dir:
- files.append(os.path.join(root, f))
- for f in files:
- process_file(f, params.dst, args.preview, args.offline)
- log.info({ 'processed': i, 'inputs': len(files) })
- # print(json.dumps(metadata, indent=2))
diff --git a/cli/modules/train-losschart.py b/cli/modules/train-losschart.py
deleted file mode 100755
index 74523f637..000000000
--- a/cli/modules/train-losschart.py
+++ /dev/null
@@ -1,191 +0,0 @@
-#!/usr/bin/env python
-
-import io
-import os
-import sys
-import json
-import pathlib
-import logging
-import torch
-import numpy as np
-from PIL import Image, ImageFont, ImageDraw
-from matplotlib import pyplot as plt
-from util import log, Map
-
-
-def settings(logdir: str, name: str):
- filename = os.path.join(logdir, name, 'settings.json')
- with open(filename, 'r', encoding='utf-8') as f:
- data = json.load(f)
- data = Map(data)
- log.debug({ 'settings': data })
- return data
-
-
-def plot(logdir: str, name: str):
- f = os.path.join(logdir, name, 'train.csv')
- if not os.path.isfile(f):
- log.debug({ 'train log missing': f })
- return
- name = pathlib.Path(f).parent.name
- img = os.path.join(logdir, f"{name}.train.png")
-
- step, loss, rate = plt.np.loadtxt(f, delimiter = ',', skiprows = 1, usecols = [0, 3, 4], unpack = True)
- d = settings(logdir, name)
- # window = d.get('gradient_step', 1) * d.get('batch_size', 1)
- window = d.get('save_embedding_every', 1)
- try:
- log.debug({ 'loss plot': name, 'output': img, 'data': f, 'records': len(step) })
- except:
- return # no data
- if len(step) < 5:
- return
-
- plt.rcParams.update({'font.variant':'small-caps'})
- plt.rc('axes', edgecolor='gray')
- plt.rc('font', size=10)
- plt.rc('font', variant='small-caps')
- plt.grid(color='gray', linewidth=1, axis='both', alpha=0.5)
- plt.rcParams['figure.figsize'] = [14, 14]
- plt.rcParams['figure.facecolor'] = 'black'
- figure, axis = plt.subplots(2, 1)
-
- # create top graph
- ax0 = axis[0]
- ax0.set_facecolor(color = (0.1, 0.1, 0.1, 0.5))
- ax0.tick_params(axis='x', labelcolor='white')
- ax0.set_xlabel('step'.upper(), color='white')
- ax0.set_xlim(0, d.steps)
- ax0.set_ylim(0, 0.5)
- ax0.set_axisbelow(True)
- ax0.xaxis.grid(color='gray', linestyle='dashed')
- ax0.yaxis.grid(color='gray', linestyle='dashed')
-
- # loss values
- ax0.plot(step, loss, color='gray', label='loss value')
- ax0.set_ylabel('loss value'.upper(), color='#CE6400')
- ax0.tick_params(axis='y', labelcolor='#CE6400')
-
- # trendline
- z = np.polyfit(step, loss, 1)
- p = np.poly1d(z)
- ax0.plot(step, p(loss), color='#5020F0', linewidth=3, label='loss trendline')
-
- # moving average
- if len(loss) > window:
- maval = []
- minval = []
- maxval = []
- for ind in range(window - 1):
- maval.insert(0, np.nan)
- minval.insert(0, np.nan)
- maxval.insert(0, np.nan)
- for ind in range(len(loss) - window + 1):
- maval.append(np.mean(loss[ind:ind+window]))
- minval.append(np.min(loss[ind:ind+window]))
- maxval.append(np.max(loss[ind:ind+window]))
- ax0.plot(step, maval, color='#CE6400', linewidth=5, label='average loss value')
- ax0.plot(step, minval, color='#500010', linewidth=5, label='min loss per epoch')
- ax0.plot(step, maxval, color='#005010', linewidth=5, label='max loss per epoch')
-
- # learning rate
- ax0_right = ax0.twinx()
- ax0_right.set_ylabel('learn rate'.upper(), color='cyan')
- ax0_right.plot(step, rate, color='cyan', linewidth=3, linestyle='dashed', label='learn rate')
- ax0_right.tick_params(axis='y', labelcolor='cyan')
-
- # axis legend
- handles0, labels0 = ax0.get_legend_handles_labels() # because ax2 is twin, both are included
- ax0.legend(handles0, labels0, loc="best")
-
- # embeddings
- ax1 = axis[1]
- ax1.set_facecolor(color = (0.1, 0.1, 0.1, 0.5))
- ax1.set_ylabel('vector average'.upper(), color=(1, 0.2, 0.5))
- ax1.tick_params(axis='y', labelcolor=(1, 0.2, 0.5))
- ax1.set_xlim(0, d.steps)
- ax1_right = ax1.twinx()
- ax1_right.set_ylabel('vector norm'.upper(), color=(0.2, 1.0, 0.5))
- ax1_right.tick_params(axis='y', labelcolor=(0.2, 1.0, 0.5))
- ax1.xaxis.grid(color='gray', linestyle='dashed')
- ax1.yaxis.grid(color='gray', linestyle='dashed')
- x = []
- avg = []
- norm = []
- avg_v = [[] for y in range(d.num_vectors_per_token)]
- norm_v = [[] for y in range(d.num_vectors_per_token)]
- embedding_files = sorted(pathlib.Path(os.path.join(logdir, name, 'embeddings')).glob('*.pt'), key=os.path.getmtime)
- for f in embedding_files:
- embed = torch.load(f, map_location=torch.device("cpu")) # pylint: disable=no-member
- x.append(embed["step"] + 1)
- token = list(embed["string_to_token"].keys())[0]
- tensors = embed["string_to_param"][token]
- val = tensors.detach().numpy()
- data = val.flatten()
- avg.append(np.average(np.abs(data)))
- norm.append(np.linalg.norm(data))
- for i in range(val.shape[0]):
- avg_v[i].append(np.average(np.abs(val[i])))
- norm_v[i].append(np.linalg.norm(val[i]))
- ax1.plot(x, avg, color=(1, 0.2, 0.5), linewidth=3, label='all vectors average value')
- ax1_right.plot(x, norm, color=(0.2, 1, 0.5), linewidth=3, label='all vectors norm value', linestyle='dashed')
- for i in range(d.num_vectors_per_token):
- ax1.plot(x, avg_v[i], color= (i / (d.num_vectors_per_token + 1), 0.2, 0.5), linewidth=1, label=f'vector={i} average value')
- ax1_right.plot(x, norm_v[i], color= (0.2, i / (d.num_vectors_per_token + 1), 0.5), linewidth=1, label=f'vector={i} norm value', linestyle='dashed')
-
- # axis legend
- handles1, labels1 = ax1.get_legend_handles_labels() # because ax2 is twin, both are included
- ax1.legend(handles1, labels1, loc="upper left")
- ax1_right.legend(loc="upper right")
-
- # create chart and convert to pil
- figure.tight_layout()
- buf = io.BytesIO()
- plt.savefig(buf, format='png')
- pltimg = Image.open(buf)
- size = (pltimg.size[0], pltimg.size[1] + 240)
- image = Image.new('RGB', size = size, color = (206, 100, 0))
- font = ImageFont.truetype('DejaVuSansMono', 18)
- image.paste(pltimg, box=(0, 240))
- buf.close()
- plt.close()
-
- # text
- textl = f"""NAME: {d.embedding_name.upper()}
-IMAGES: {d.num_of_dataset_images}
-VECTORS: {d.num_vectors_per_token}
-STEPS: {d.steps}
-BATCH-SIZE: {d.batch_size}
-GRADIENT-STEP: {d.gradient_step}
-SAMPLING-METHOD: {d.latent_sampling_method}
-MODEL: {d.model_name.upper()}
-LEARN-RATE: {d.learn_rate.replace(' ', '')}
-"""
-
- minval = f"{round(np.min(loss), 4)} @ {round(step[np.argmin(loss)])}"
- maxval = f"{round(np.max(loss), 4)} @ {round(step[np.argmax(loss)])}"
- textr = f"""{d.datetime}
-LOSS: {round(loss[-1], 4)}
-MIN: {minval}
-MAX: {maxval}
-TREND: {z[0]:.5f}
-"""
- if len(avg) > 0:
- textr += f"VECTOR AVG: {avg[-1]:.3f}\n"
- if len(norm) > 0:
- textr += f"VECTOR NORM: {norm[-1]:.3f}\n"
- ctx = ImageDraw.Draw(image)
- ctx.text((8, 8), textl, font = font, fill = (255, 255, 255), spacing = 8)
- ctx.text((image.size[0] - 220, 8), textr, font = font, fill = (255, 255, 255))
-
- image.save(img)
-
-
-if __name__ == "__main__":
- log.setLevel(logging.DEBUG)
- if len(sys.argv) == 2:
- arg = sys.argv[1]
- log.debug({ 'args': arg })
- plot(os.path.dirname(arg), os.path.basename(arg))
- else:
- log.error({ 'loss chart': 'specify embedding name'})
diff --git a/cli/modules/train-lossrate.py b/cli/modules/train-lossrate.py
deleted file mode 100755
index dd6b394dd..000000000
--- a/cli/modules/train-lossrate.py
+++ /dev/null
@@ -1,132 +0,0 @@
-#!/usr/bin/env python
-"""
-auto-generate learn-rate
-"""
-import io
-import math
-import logging
-import numpy as np
-from PIL import Image, ImageFont, ImageDraw
-from matplotlib import pyplot as plt
-from util import log, Map
-
-
-loss_types = ['linear', 'log', 'linalg', 'power']
-
-
-def gen_steps(steps, step):
- return [x for x in range(1, steps + step) if x % step == 0]
-
-
-def gen_loss_rate(steps: int, step: int, loss_start: float, loss_end: float, loss_type: loss_types, power: int = 3):
- def norm(val):
- return ((loss_start - loss_end) * val) / val.max() + loss_end
-
- steps_val = gen_steps(steps, step)
-
- if loss_type == 'linear':
- loss_val = np.interp(steps_val, [steps_val[0], steps_val[-1]], [loss_start, loss_end])
-
- elif loss_type == 'log':
- loss_val = np.logspace(loss_start, 0, num=len(steps_val), base=math.e)
- loss_val = norm(loss_val - loss_val.min())
-
- elif loss_type == 'linalg':
- loss_val = np.array(steps_val[::-1], dtype='float')
- loss_val = norm(loss_val / np.linalg.norm(loss_val))
-
- elif loss_type == 'power':
- loss_val = np.array([math.pow(x, power) for x in range(len(steps_val))][::-1])
- loss_val = norm(loss_val)
-
- else:
- return []
-
- return loss_val
-
-
-def gen_loss_rate_str(steps: int, step: int, loss_start: float, loss_end: float, loss_type: loss_types, power: int = 3):
- steps_val = gen_steps(steps, step)
- loss_val = gen_loss_rate(steps, step, loss_start, loss_end, loss_type, power)
- loss_rate = [f"{loss_val[i]:.4f}:{steps_val[i]}" for i in range(len(steps_val))]
- loss_rate = ', '.join(loss_rate)
- log.debug({ 'loss_rate': loss_rate, 'function': loss_type, 'power': power })
- return loss_rate
-
-
-def example_plot(steps: int, step: int, loss_start: float, loss_end: float):
- plt.rcParams.update({'font.variant':'small-caps'})
- plt.rc('axes', edgecolor='gray')
- plt.rc('font', size=10)
- plt.rc('font', variant='small-caps')
- plt.grid(color='gray', linewidth=1, axis='both', alpha=0.5)
- plt.rcParams['figure.figsize'] = [14, 6]
- plt.figure(facecolor='black')
-
- loss_rates = []
-
- ax1 = plt.subplot(1, 2, 1)
- ax1.set_facecolor('grey')
- ax1.set_xlabel('step'.upper(), color='white')
- ax1.set_ylabel('loss value'.upper(), color='white')
- ax1.xaxis.grid(color='gray', linestyle='dashed')
- ax1.tick_params(axis='x', labelcolor='white')
- ax1.tick_params(axis='y', labelcolor='white')
- ax1.legend(loc="best")
- for loss_type in [x for x in loss_types if x != 'power']:
- col = (np.random.random(), np.random.random(), np.random.random())
- x = gen_steps(steps, step)
- y = gen_loss_rate(loss_type = loss_type, steps = steps, step = step, loss_start = loss_start, loss_end = loss_end)
- ax1.plot(x, y, label=loss_type, color = col)
- loss = gen_loss_rate_str(loss_type = loss_type, steps = steps, step = step, loss_start = loss_start, loss_end = loss_end)
- loss_rates.append(f"LOSS {loss} TYPE {loss_type}")
- handles, labels = ax1.get_legend_handles_labels()
- ax1.legend(handles, labels)
-
- ax2 = plt.subplot(1, 2, 2)
- ax2.set_facecolor('grey')
- ax2.set_xlabel('step'.upper(), color='white')
- ax2.set_ylabel('loss value'.upper(), color='white')
- ax2.xaxis.grid(color='gray', linestyle='dashed')
- ax2.tick_params(axis='x', labelcolor='white')
- ax2.tick_params(axis='y', labelcolor='white')
- ax2.legend(loc="best")
- for power in range(1, 11):
- col = (np.random.random(), np.random.random(), np.random.random())
- x = gen_steps(steps, step)
- y = gen_loss_rate(loss_type = 'power', power = power, steps = steps, step = step, loss_start = loss_start, loss_end = loss_end)
- ax2.plot(x, y, label=f"power={pow}", color = col)
- loss = gen_loss_rate_str(loss_type = 'power', power = power, steps = steps, step = step, loss_start = loss_start, loss_end = loss_end)
- loss_rates.append(f"LOSS {loss} TYPE power={power}")
- handles, labels = ax2.get_legend_handles_labels()
- ax2.legend(handles, labels)
-
- plt.tight_layout()
- buf = io.BytesIO()
- plt.savefig(buf, format='png')
- pltimg = Image.open(buf)
- size = (pltimg.size[0], pltimg.size[1] + 300)
- image = Image.new('RGB', size = size, color = (206, 100, 0))
- font = ImageFont.truetype('DejaVuSansMono', 14)
- image.paste(pltimg, box=(0, 300))
- buf.close()
-
- # text
- rates = "\n".join(loss_rates)
- text = f"STEPS {steps} STEP {step} LOSS-START {loss_start} LOSS-END {loss_end}\n" + rates
-
- ctx = ImageDraw.Draw(image)
- ctx.text((8, 8), text, font = font, fill = (255, 255, 255), spacing = 8)
- image.save('lossrate.jpg')
-
-
-if __name__ == "__main__":
- log.setLevel(logging.DEBUG)
- arg = Map({
- "steps": 500,
- "step": 50,
- "loss_start": 0.01,
- "loss_end": 0.001
- })
- log.debug({ 'options': arg })
- example_plot(**arg)
diff --git a/cli/modules/prompt-ideas.py b/cli/prompt-ideas.py
similarity index 88%
rename from cli/modules/prompt-ideas.py
rename to cli/prompt-ideas.py
index 18efd8d57..55b9a2d7b 100755
--- a/cli/modules/prompt-ideas.py
+++ b/cli/prompt-ideas.py
@@ -10,18 +10,13 @@ from transformers import GPT2Tokenizer, GPT2LMHeadModel
from util import log
-tokenizer = None
-model = None
-
-
def prompt(text: str, temp: float = 0.9, top: int = 8, penalty: float = 1.2, alpha: float = 0.6, num: int = 5, length: int = 80):
- global tokenizer, model # pylint: disable=global-statement
- if tokenizer is None:
- tokenizer = GPT2Tokenizer.from_pretrained('distilgpt2')
- tokenizer.add_special_tokens({'pad_token': '[PAD]'})
- if model is None:
- model = GPT2LMHeadModel.from_pretrained('FredZhang7/distilgpt2-stable-diffusion-v2')
+ log.info({ 'loading': 'tokenizer' })
+ tokenizer = GPT2Tokenizer.from_pretrained('distilgpt2')
+ tokenizer.add_special_tokens({'pad_token': '[PAD]'})
input_ids = tokenizer(text, return_tensors='pt').input_ids
+ log.info({ 'loading': 'model' })
+ model = GPT2LMHeadModel.from_pretrained('FredZhang7/distilgpt2-stable-diffusion-v2')
output = model.generate(input_ids,
do_sample = True,
temperature = temp,
diff --git a/cli/modules/prompt-promptist.py b/cli/prompt-promptist.py
similarity index 80%
rename from cli/modules/prompt-promptist.py
rename to cli/prompt-promptist.py
index 59fea45d5..23e670e59 100755
--- a/cli/modules/prompt-promptist.py
+++ b/cli/prompt-promptist.py
@@ -5,24 +5,27 @@ use microsoft promptist to beautify prompt
"""
import sys
-from transformers import AutoModelForCausalLM, AutoTokenizer
from util import log
-
-def load_prompter():
+def load_model():
+ log.info({ 'loading': 'model' })
+ from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained("microsoft/Promptist") # pylint: disable=redefined-outer-name
+ return model
+
+def load_tokenizer():
+ log.info({ 'loading': 'tokenizer' })
+ from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained("gpt2") # pylint: disable=redefined-outer-name
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "left"
- return model, tokenizer
-
-
-model, tokenizer = load_prompter()
-
+ return tokenizer
def beautify(plain_text):
+ tokenizer = load_tokenizer()
input_ids = tokenizer(plain_text.strip() + " Rephrase:", return_tensors = "pt").input_ids
eos_id = tokenizer.eos_token_id
+ model = load_model()
outputs = model.generate(input_ids, do_sample = False, max_new_tokens = 75, num_beams = 8, num_return_sequences = 8, eos_token_id = eos_id, pad_token_id = eos_id, length_penalty = -1.0)
output_texts = tokenizer.batch_decode(outputs, skip_special_tokens = True)
texts = []
diff --git a/cli/random/detectmodel.py b/cli/random/detectmodel.py
deleted file mode 100755
index 0d7096123..000000000
--- a/cli/random/detectmodel.py
+++ /dev/null
@@ -1,43 +0,0 @@
-#!/usr/bin/env python
-"""
-Detect model type
-
-Works for v1 and v2-base (EPS models), both standard inference and inpainting
-
-But looking at model dumps between EPS and V type models, its only about parametrization, there are no differences in actual model (its just weighted differently without any structural difference)
-So i don't see easy way to auto-detect if model should be run in `EPS` or `V` mode
-Only difference are some calculations in `ldm/models/diffusion/ddpm.py` and by then we already need to know which code-path to trigger
-(maaaybe there could be a way by looking at some cherry-picked base tensors min/max range, but I dont see that as reliable)
-"""
-
-import os
-import sys
-
-import torch
-
-def signature(model):
- if model is None:
- return None
- try:
- size = model['state_dict']['model.diffusion_model.input_blocks.1.1.transformer_blocks.0.attn2.to_k.weight'].shape[1]
- unet = model['state_dict']['model.diffusion_model.input_blocks.0.0.weight'].shape[1]
- except:
- return 'unknown'
- guess = 'v1' if size == 768 else 'v2' # 768 for v1 and 1024 for v2
- guess += '-inference' if unet == 4 else '-inpainting' # inference models have shorter inputs, 4 for inference 9 for inpainting
- return guess
-
-def load(file: str):
- try:
- model = torch.load(file, map_location='cpu')
- return model
- except Exception as err:
- print(f"Error loading {f}: {err}")
-
-if __name__ == "__main__":
- sys.argv.pop(0)
- for f in sys.argv:
- if os.path.isfile(f):
- print(f"Model {f} is of type {signature(load(f))}")
- else:
- print(f"{f} is not a file")
diff --git a/cli/random/versions.py b/cli/random/versions.py
deleted file mode 100755
index 136f7214b..000000000
--- a/cli/random/versions.py
+++ /dev/null
@@ -1,33 +0,0 @@
-#!/usr/bin/env python
-"""
-print module versions
-"""
-
-import importlib
-import pkg_resources
-
-modules = [
- 'diffusers', 'xformers', 'tokenizers', 'accelerate', 'safetensors'
-]
-
-def get_torch():
- try:
- torch = importlib.import_module('torch')
- print('torch:', { 'version': torch.__version__ })
- print('cuda:', { 'available': torch.cuda.is_available(), 'version': torch.version.cuda, 'arch': torch.cuda.get_arch_list() })
- print('device:', { 'name': torch.cuda.get_device_name(torch.cuda.current_device()) })
- except Exception as err:
- print('torch:', { 'error': err })
-
-
-def version(name: str):
- try:
- ver = pkg_resources.get_distribution(name).version
- print(f"{name}: {ver}")
- except Exception as err:
- print(f"{name} error: {err}")
-
-if __name__ == "__main__": # create & train test embedding when used from cli
- get_torch()
- for module in modules:
- version(module)
diff --git a/cli/requirements.txt b/cli/requirements.txt
index b8f807b41..50732dc8b 100644
--- a/cli/requirements.txt
+++ b/cli/requirements.txt
@@ -2,3 +2,4 @@ mediapipe
colormap
invisible-watermark
filetype
+albumentations
diff --git a/cli/modules/bench.py b/cli/run-benchmark.py
similarity index 100%
rename from cli/modules/bench.py
rename to cli/run-benchmark.py
diff --git a/cli/modules/sdapi.py b/cli/sdapi.py
similarity index 97%
rename from cli/modules/sdapi.py
rename to cli/sdapi.py
index c2201ef9c..2a258955e 100755
--- a/cli/modules/sdapi.py
+++ b/cli/sdapi.py
@@ -1,14 +1,14 @@
#!/usr/bin/env python
+#pylint: disable=redefined-outer-name
"""
helper methods that creates HTTP session with managed connection pool
provides async HTTP get/post methods and several helper methods
"""
import sys
-import json
-import aiohttp
import asyncio
import logging
+import aiohttp
import requests
from util import Map, log
@@ -76,7 +76,6 @@ def getsync(endpoint: str, json: dict = None):
except Exception as err:
log.error({ 'session': err })
return {}
-
async def post(endpoint: str, json: dict = None):
@@ -136,9 +135,9 @@ def progresssync():
def options():
- options = getsync('/sdapi/v1/options')
+ opts = getsync('/sdapi/v1/options')
flags = getsync('/sdapi/v1/cmd-flags')
- return { 'options': options, 'flags': flags }
+ return { 'options': opts, 'flags': flags }
def shutdown():
@@ -184,6 +183,7 @@ if __name__ == "__main__":
if 'options' in sys.argv:
opt = options()
log.debug({ 'options' })
+ import json
print(json.dumps(opt['options'], indent = 2))
log.debug({ 'cmd-flags' })
print(json.dumps(opt['flags'], indent = 2))
diff --git a/cli/random/dynamotest.py b/cli/torch-compile.py
similarity index 100%
rename from cli/random/dynamotest.py
rename to cli/torch-compile.py
diff --git a/cli/train-lora.py b/cli/train-lora.py
deleted file mode 100755
index 0feb3e8fa..000000000
--- a/cli/train-lora.py
+++ /dev/null
@@ -1,274 +0,0 @@
-#!/usr/bin/env python
-
-"""
-Extract approximating LoRA by SVD from two SD models
-Based on:
-
-Train LoRA with custom preprocessing, tagging and bucketing
-
-Disabled/broken:
-- `accelerate` with *dynamo* enabled
-- `xformers` due to *faketensors* requirement
-- `mem_eff_attn` due to *forwardfunc* mismatch
-- 'use_8bit_adam` due to *bitsandbyttes* CUDA errors
-"""
-
-import os
-import re
-import gc
-import sys
-import json
-import time
-import shutil
-import argparse
-import tempfile
-import torch
-import logging
-import importlib
-import transformers
-from pathlib import Path
-from modules.util import log, Map, get_memory
-import modules.process
-import modules.sdapi
-
-latents = importlib.import_module('modules.lora-latents')
-
-lora_path = os.path.abspath(os.path.join(os.path.dirname(__file__), os.pardir, 'modules', 'lora'))
-sys.path.append(lora_path)
-lycoris_path = os.path.abspath(os.path.join(os.path.dirname(__file__), os.pardir, 'modules', 'lycoris'))
-sys.path.append(lycoris_path)
-from train_network import train
-
-options = Map({
- "bucket_no_upscale": False,
- "bucket_reso_steps": 64,
- "cache_latents": True,
- "caption_dropout_every_n_epochs": None,
- "caption_dropout_rate": 0.0,
- "caption_extension": ".txt",
- "caption_extention": ".txt",
- "caption_tag_dropout_rate": 0.0,
- "clip_skip": None,
- "color_aug": False,
- "dataset_repeats": 1,
- "debug_dataset": False,
- "enable_bucket": False,
- "face_crop_aug_range": None,
- "flip_aug": False,
- "full_fp16": False,
- "gradient_accumulation_steps": 1,
- "gradient_checkpointing": False,
- "in_json": "",
- "keep_tokens": None,
- "learning_rate": 5e-05,
- "log_prefix": None,
- "logging_dir": None,
- "lr_scheduler_num_cycles": 1,
- "lr_scheduler_power": 1,
- "lr_scheduler": "cosine",
- "lr_warmup_steps": 0,
- "max_bucket_reso": 1024,
- "max_data_loader_n_workers": 8,
- "max_grad_norm": 0.0,
- "max_token_length": None,
- "max_train_epochs": None,
- "max_train_steps": 5000,
- "mem_eff_attn": False,
- "min_bucket_reso": 256,
- "mixed_precision": "fp16",
- "network_alpha": 1.0,
- "network_args": None,
- "network_dim": 16,
- "network_module": "networks.lora",
- "network_train_text_encoder_only": False,
- "network_train_unet_only": False,
- "network_weights": None,
- "no_metadata": False,
- "output_dir": "",
- "output_name": "",
- "persistent_data_loader_workers": False,
- "pretrained_model_name_or_path": "",
- "prior_loss_weight": 1.0,
- "random_crop": False,
- "reg_data_dir": None,
- "resolution": "512,512",
- "resume": None,
- "save_every_n_epochs": None,
- "save_last_n_epochs_state": None,
- "save_last_n_epochs": None,
- "save_model_as": "ckpt",
- "save_n_epoch_ratio": None,
- "save_precision": "fp16",
- "save_state": False,
- "seed": 42,
- "shuffle_caption": False,
- "text_encoder_lr": 5e-05,
- "train_batch_size": 1,
- "train_data_dir": "",
- "training_comment": "mood-magic",
- "unet_lr": 0.001,
- "use_8bit_adam": False,
- "v_parameterization": False,
- "v2": False,
- "vae": None,
- "xformers": False,
-})
-
-
-def mem_stats():
- gc.collect()
- if torch.cuda.is_available():
- with torch.no_grad():
- torch.cuda.empty_cache()
- with torch.cuda.device('cuda'):
- torch.cuda.empty_cache()
- torch.cuda.ipc_collect()
- mem = get_memory()
- log.info({ 'memory': { 'ram': mem.ram, 'gpu': mem.gpu } })
-
-
-if __name__ == '__main__':
- parser = argparse.ArgumentParser(description = 'train lora')
- parser.add_argument('--model', type=str, default=None, required=False, help='original model to use a base for training, default: active model')
- parser.add_argument('--input', '--dataset', type=str, default=None, required=True, help='input folder with training images')
- parser.add_argument('--output', '--lora', type=str, default=None, required=True, help='lora name')
- parser.add_argument('--tag', type=str, default=None, required=False, help='primary tag')
- parser.add_argument('--dir', type=str, default=None, required=False, help='folder containing lora checkpoints')
- parser.add_argument('--interim', type=int, default=0, help = 'save interim checkpoints after n epoch')
- parser.add_argument('--process', type=str, default='original', required=False, help='list of processing steps: original,face,body,blur,range,upscale,restore')
- parser.add_argument('--noprocess', default = False, action='store_true', help = 'skip processing and use existing input data')
- parser.add_argument('--notrain', default = False, action='store_true', help = 'just run processing and skip training')
- parser.add_argument('--nocaptions', default = False, action='store_true', help = 'skip creating captions and tags')
- parser.add_argument('--nolatents', default = False, action='store_true', help = 'skip generating vae latents')
- parser.add_argument('--offline', default = False, action='store_true', help = 'do not use webui server for processing')
- parser.add_argument('--shutdown', default = False, action='store_true', help = 'shutdown webui server')
- parser.add_argument('--gradient', type=int, default=1, required=False, help='gradient accumulation steps, default: %(default)s')
- parser.add_argument('--steps', type=int, default=4000, required=False, help='training steps, default: %(default)s')
- parser.add_argument('--dim', type=int, default=40, required=False, help='network dimension, default: %(default)s')
- parser.add_argument('--repeats', type=int, default=10, required=False, help='number of repeats per image, default: %(default)s')
- parser.add_argument('--alpha', type=float, default=0, required=False, help='alpha for weights scaling, default: half of dim')
- parser.add_argument('--batch', type=int, default=1, required=False, help='batch size, default: %(default)s')
- parser.add_argument('--lr', type=float, default=1e-04, required=False, help='model learning rate, default: %(default)s')
- parser.add_argument('--unetlr', type=float, default=1e-04, required=False, help='unet learning rate, default: %(default)s')
- parser.add_argument('--textlr', type=float, default=5e-05, required=False, help='text encoder learning rate, default: %(default)s')
- parser.add_argument('--dreambooth', default=False, action='store_true', help = "use dreambooth style training")
- parser.add_argument('--lycoris', default=False, action='store_true', help = "use lycoris style training")
- parser.add_argument('--debug', default=False, action='store_true', help = "enable debug logging")
- args = parser.parse_args()
- defaults = Map({ 'options': {}, 'flags': {} }) if args.offline else Map(modules.sdapi.options())
-
- if args.debug:
- log.setLevel(logging.DEBUG)
- log.debug({ 'debug': True })
- if args.model is None:
- args.model = defaults.options.get('sd_model_checkpoint', None)
- args.model = args.model.split(' [')[0] if args.model is not None else None
- if args.dir is None:
- args.dir = defaults.flags.get('lora_dir', None)
- if not os.path.isabs(args.model) and args.dir is not None and not os.path.exists(args.model):
- args.model = os.path.abspath(os.path.join(args.dir, os.pardir, 'Stable-diffusion', args.model))
- if args.dir is None:
- args.dir = os.path.join(args.input, 'lora')
- if not os.path.exists(args.model) or not os.path.isfile(args.model):
- log.error({ 'lora cannot find model': args.model })
- exit(1)
- if not os.path.exists(args.input) or not os.path.isdir(args.input):
- log.error({ 'lora cannot find training dir': args.input })
- exit(1)
- if not os.path.exists(args.dir) or not os.path.isdir(args.dir):
- log.error({ 'lora cannot find training dir': args.dir })
- exit(1)
- options.pretrained_model_name_or_path = args.model
- options.output_dir = args.dir
- options.output_name = args.output
- options.max_train_steps = args.steps
- options.network_dim = args.dim
- options.network_alpha = args.dim // 2 if args.alpha == 0 else args.alpha
- options.gradient_accumulation_steps = args.gradient
- options.save_every_n_epochs = args.interim if args.interim > 0 else None
- options.learning_rate = args.lr
- options.unet_lr = args.unetlr
- options.text_encoder_lr = args.textlr
- options.train_batch_size = args.batch
- log.info({ 'train lora args': vars(options) })
- transformers.logging.set_verbosity_error()
- mem_stats()
-
- json_file = os.path.join(tempfile.gettempdir(), args.output, args.output + '.json')
- base = os.path.join(tempfile.gettempdir(), args.output)
- options.train_data_dir = base
- res = None
-
- if args.dreambooth:
- log.info({ 'using dreambooth style training': True })
- options.in_json = None
- else:
- options.in_json = json_file
-
- for root, _sub_dirs, folder in os.walk(args.input):
- files = [os.path.join(root, f) for f in folder]
-
- if not args.noprocess:
- # preprocess
- processing_options = args.process.split(',')
- processing_options = [opt.strip() for opt in re.split(',| ', args.process)]
- log.info({ 'processing steps': processing_options })
-
- if os.path.exists(json_file):
- os.remove(json_file)
-
- steps = [step for step in processing_options if step in ['face', 'body', 'original']]
- for step in steps:
- # processing_options = [step for step in processing_options if step not in ['face', 'body', 'original']].append(step)
- if step == 'face':
- opts = [step for step in processing_options if step not in ['body', 'original']]
- if step == 'body':
- opts = [step for step in processing_options if step not in ['face', 'original', 'upscale', 'restore']]
- if step == 'original':
- opts = [step for step in processing_options if step not in ['face', 'body', 'upscale', 'restore', 'blur', 'range']]
- log.info({ 'processing step': opts })
- concept = step
- if concept == 'original' and args.tag is not None:
- concept = args.tag.split(',')[0].strip()
- dir = os.path.join(base, str(args.repeats) + '_' + concept)
- if os.path.exists(dir):
- shutil.rmtree(dir, ignore_errors=True)
- Path(dir).mkdir(parents=True, exist_ok=True)
-
- for f in files:
- try:
- res, metadata = modules.process.process_file(f = f, dst = dir, preview = False, offline = args.offline, txt = args.dreambooth, tag = args.tag, opts = opts)
- if not args.dreambooth:
- with open(json_file, "w") as outfile:
- outfile.write(json.dumps(metadata, indent=2))
- except ValueError as e:
- exit(1)
- log.info({ 'processed step': step, 'outputs': res, 'inputs': len(files), 'metadata': json_file, 'path': dir })
-
- modules.process.unload_models()
- mem_stats()
-
-
- dirs = [os.path.join(base, dir) for dir in os.listdir(base) if os.path.isdir(os.path.join(base, dir))]
- log.info({ 'input datasets': dirs, 'metadata': json_file })
-
- if not args.nolatents and not args.dreambooth:
- # create latents
- for dir in dirs:
- latents.create_vae_latents(Map({ 'input': dir, 'json': json_file }))
- latents.unload_vae()
- mem_stats()
- else:
- log.info({ 'skip processing': len(files), 'metadata': json_file, 'path': dir })
-
- if args.shutdown:
- log.info({ 'server shutdown required': True })
- modules.sdapi.shutdown()
- time.sleep(1)
-
- if args.lycoris:
- log.info({ 'using lycoris network': True })
- options.network_module = 'lycoris.kohya'
- if not args.notrain:
- train(options)
- mem_stats()
diff --git a/cli/train-ti.py b/cli/train-ti.py
deleted file mode 100755
index a2848c74f..000000000
--- a/cli/train-ti.py
+++ /dev/null
@@ -1,591 +0,0 @@
-#!/usr/bin/env python
-# pylint: disable=no-member
-"""
-simple implementation of training api: `/sdapi/v1/train`
-- supports: create embedding, image preprocess, train embedding (with all known parameters)
-- does not (yet) support: create hyper-network, train hyper-network
-- compatible with progress api: `/sdapi/v1/progress`
-- if interrupted, auto-continues from last known step
-- create and preprocess executed as sync jobs
-- train is executed as async job with progress monitoring
-"""
-
-import argparse
-import asyncio
-import logging
-import math
-import os
-import sys
-import time
-import importlib
-from pathlib import Path, PurePath
-
-import filetype
-from PIL import Image
-
-sys.path.append(os.path.join(os.path.dirname(__file__), 'modules'))
-from modules.util import Map, log, set_logfile
-from modules.sdapi import close, get, interrupt, post, progress, session
-from modules.process import process_images
-from modules.grid import grid
-create_preview = importlib.import_module('modules.preview-embeddings').create_preview
-plot = importlib.import_module('modules.train-losschart').plot
-extract = importlib.import_module('modules.video-extract').extract
-gen_loss_rate_str = importlib.import_module('modules.train-lossrate').gen_loss_rate_str
-
-images = []
-args = {}
-options = None
-cmdflags = None
-args = Map({
- "training_model": "sd-v15-runwayml.ckpt",
- "extract_video": {
- "rate": 0,
- "fps": 5,
- "vstart": 0,
- "vend": 0
- },
- "create_embedding": {
- "name": "test",
- "num_vectors_per_token": 1,
- "overwrite_old": False,
- "init_text": "*"
- },
- "preprocess": {
- "id_task": 0,
- "process_src": "",
- "process_dst": "",
- "process_width": 512,
- "process_height": 512,
- "process_flip": False,
- "process_split": False,
- "process_caption": True,
- "process_caption_deepbooru": False,
- "preprocess_txt_action": "ignore",
- "process_focal_crop": True,
- "process_focal_crop_face_weight": 0.9,
- "process_focal_crop_entropy_weight": 0.3,
- "process_focal_crop_edges_weight": 0.5,
- "process_focal_crop_debug": False,
- "split_threshold": 0.5,
- "overlap_ratio": 0.2,
- "process_multicrop": None,
- "process_multicrop_mindim": None,
- "process_multicrop_maxdim": None,
- "process_multicrop_minarea": None,
- "process_multicrop_maxarea": None,
- "process_multicrop_objective": None,
- "process_multicrop_threshold": None,
- },
- "train_embedding": {
- "id_task": 0,
- "embedding_name": "",
- "learn_rate": -1,
- "batch_size": 1,
- "steps": 500,
- "data_root": "",
- "log_directory": "train/log",
- "template_filename": "subject_filewords.txt",
- "gradient_step": 20,
- "training_width": 512,
- "training_height": 512,
- "shuffle_tags": False,
- "tag_drop_out": 0,
- "clip_grad_mode": "disabled",
- "clip_grad_value": "0.1",
- "latent_sampling_method": "once",
- "create_image_every": -1,
- "save_embedding_every": -1,
- "save_image_with_stored_embedding": False,
- "preview_from_txt2img": False,
- "preview_prompt": "",
- "preview_negative_prompt": "blurry, duplicate, ugly, deformed, low res, watermark, text",
- "preview_steps": 20,
- "preview_sampler_index": 0,
- "preview_cfg_scale": 6,
- "preview_seed": -1,
- "preview_width": 512,
- "preview_height": 512,
- "varsize": False,
- "use_weight": False,
- },
-})
-
-
-async def plotloss(params):
- logdir = os.path.abspath(os.path.join(cmdflags['embeddings_dir'], '../train/log'))
- try:
- plot(logdir, params.name)
- except Exception as err:
- log.warning({ 'loss chart error': err })
-
-
-async def captions(docs: list):
- exclude = ['a', 'in', 'on', 'out', 'at', 'the', 'and', 'with', 'next', 'to', 'it', 'for', 'of', 'into', 'that']
- d = dict()
- for f in docs:
- text = open(f, 'r', encoding='utf-8')
- for line in text:
- line = line.strip()
- line = line.lower()
- words = line.split(" ")
- for word in words:
- if word in exclude:
- continue
- d[word] = d[word] + 1 if word in d else 1
- pairs = ((value, key) for (key,value) in d.items())
- sort = sorted(pairs, reverse = True)
- if len(sort) > 10:
- del sort[10:]
- d = {k: v for v, k in sort}
- log.info({ 'top captions': d })
-
-
-async def preprocess_cleanup(params):
- log.info({ 'preprocess cleanup': params.dst })
- for f in Path(params.dst).glob('*.png'):
- f.unlink()
- for f in Path(params.dst).glob('*.jpg'):
- f.unlink()
- for f in Path(params.dst).glob('*.txt'):
- f.unlink()
- try:
- if os.path.isdir(params.dst):
- Path(params.dst).rmdir()
- except Exception as err:
- log.warning({ 'preprocess cleanup': params.dst, 'error': err })
-
-
-async def preprocess_builtin(params):
- global images # pylint: disable=global-statement
- log.debug({ 'preprocess start' })
- files = [os.path.join(params.src, f) for f in os.listdir(params.src) if os.path.isfile(os.path.join(params.src, f))]
- candidates = [f for f in files if filetype.is_image(f)]
- not_images = [f for f in files if (not filetype.is_image(f) and not f.endswith('.txt'))]
- images = []
- low_res = []
- for f in candidates:
- img = Image.open(f)
- mp = (img.size[0] * img.size[1]) / 1024 / 1024
- if mp < 1 or img.size[0] < 512 or img.size[1] < 512:
- low_res.append(f)
- os.rename(f, f + '.skip')
- else:
- images.append(f)
- log.debug({ 'preprocess skipping': not_images })
- log.debug({ 'preprocess low res': low_res })
- args.preprocess.process_src = params.src
- args.preprocess.process_dst = params.dst
- log.debug({ 'preprocess args': args.preprocess })
- _res = await post('/sdapi/v1/preprocess', json = args.preprocess)
- processed = [os.path.join(params.dst, f) for f in os.listdir(params.dst) if os.path.isfile(os.path.join(params.dst, f))]
- processed_imgs = [f for f in processed if f.endswith('.png')]
- processed_docs = [f for f in processed if f.endswith('.txt')]
- log.info({ 'preprocess': {
- 'source': params.src,
- 'destination': params.dst,
- 'files': len(files),
- 'images': len(images),
- 'processed': len(processed_imgs),
- 'captions': len(processed_docs),
- 'skipped': len(not_images),
- 'low-res': len(low_res) }
- })
- if len(processed_docs) > 0:
- await captions(processed_docs)
- return len(processed_imgs)
-
-
-async def preprocess(params):
- global images # pylint: disable=global-statement
- res = 0
- if os.path.isfile(params.src):
- if not filetype.is_video(params.src):
- kind = filetype.guess(params.src)
- log.error({ 'preprocess error': { 'not a valid movie file': params.src, 'guess': kind } })
- else:
- extract_dst = os.path.join(params.dst, 'extract')
- log.debug({ 'preprocess args': args.extract_video })
- images = extract(params.src, extract_dst, rate = args.extract_video.rate, fps = args.extract_video.fps, start = args.extract_video.vstart, end = args.extract_video.vend) # extract keyframes from movie
- if images > 0:
- params.src = extract_dst
- res = await preprocess(params) # call again but now with keyframes
- else:
- log.error({ 'preprocess video extract': 'no images' })
- elif os.path.isdir(params.src):
- if params.overwrite:
- await preprocess_cleanup(params)
- elif os.path.isdir(params.dst):
- log.error({ 'preprocess output folder already exists': params.dst })
- return 0
-
- if params.preprocess == 'builtin':
- res = await preprocess_builtin(params)
- i = [os.path.join(params.dst, f) for f in os.listdir(params.dst) if os.path.isfile(os.path.join(params.dst, f)) and filetype.is_image(os.path.join(params.dst, f))]
- images = [Image.open(img) for img in i]
- res = len(images)
-
- elif params.preprocess == 'custom':
- t0 = time.perf_counter()
- args.preprocess.process_src = params.src
- args.preprocess.process_dst = params.dst
- process_images(src = params.src, dst = params.dst)
- i = [os.path.join(params.dst, f) for f in os.listdir(params.dst) if os.path.isfile(os.path.join(params.dst, f)) and filetype.is_image(os.path.join(params.dst, f))]
- images = [Image.open(img) for img in i]
- t1 = time.perf_counter()
- log.info({ 'preprocess': { 'source': params.src, 'destination': params.dst, 'images': len(images), 'time': round(t1 - t0, 2) } })
- res = len(images)
-
- else:
- args.preprocess.process_dst = params.src
- i = [os.path.join(params.src, f) for f in os.listdir(params.src) if os.path.isfile(os.path.join(params.src, f)) and filetype.is_image(os.path.join(params.src, f))]
- images = [Image.open(img) for img in i]
- res = len(images)
-
- else:
- log.error({ 'preprocess error': { 'not a valid input': params.src } })
- if len(images) > 0:
- img = grid(images, labels = None, width = 2048, height = 2048, border = 8, square = True)
- logdir = os.path.abspath(os.path.join(cmdflags['embeddings_dir'], '../train/log'))
- Path(logdir).mkdir(parents = True, exist_ok = True)
- fn = os.path.join(logdir, params.name + '.inputs.jpg')
- img.save(fn)
- log.info({ 'preprocess input grid': fn })
- return res
-
-
-async def check(params):
- global options # pylint: disable=global-statement
- options = await get('/sdapi/v1/options')
- global cmdflags # pylint: disable=global-statement
- cmdflags = await get('/sdapi/v1/cmd-flags')
-
- logdir = os.path.abspath(os.path.join(cmdflags['embeddings_dir'], '../train/log', params.name))
- logfile = os.path.abspath(os.path.join(cmdflags['embeddings_dir'], '../train/log', params.name + '.train.log'))
- set_logfile(logfile)
-
- log.info({ 'checking server options' })
-
- options['training_image_repeats_per_epoch'] = 1
-
- if params.skipmodel:
- log.info({ 'using model': options['sd_model_checkpoint'] })
- else:
- log.debug({ 'check model': args.training_model })
- if len(args.training_model) > 0 and not options['sd_model_checkpoint'].startswith(args.training_model):
- models = await get('/sdapi/v1/sd-models')
- models = [obj["title"] for obj in models]
- found = [i for i in models if i.startswith(args.training_model)]
- if len(found) == 0:
- log.error({ 'model not found': args.training_model, 'available': models })
- exit()
- else:
- log.warning({ 'switching model': found[0] })
- options['sd_model_checkpoint'] = found[0]
-
- log.debug({ 'check embedding': params.name })
-
- lst = os.path.join(cmdflags['embeddings_dir'])
- log.debug({ 'embeddings folder': lst })
- path = Path(cmdflags['embeddings_dir']).glob(f'{params.name}.pt*')
- matches = [f for f in path]
- for match in matches:
- if params.overwrite:
- log.info({ 'delete embedding': match.name })
- os.remove(os.path.join(cmdflags['embeddings_dir'], match.name))
- else:
- log.error({ 'embedding exists': match.name })
- await close()
- exit()
- f = os.path.join(logdir, 'train.csv')
- if os.path.isfile(f):
- if params.overwrite:
- log.info({ 'delete training log': f })
- os.remove(os.path.join(logdir, 'train.csv'))
- else:
- log.warning({ 'training log exists': f })
- f = os.path.join(logdir, '..', params.name, '.png')
- if os.path.isfile(f):
- if params.overwrite:
- log.info({ 'delete training graph': f })
- os.remove(f)
-
- log.debug({ 'options': 'update' })
- await post('/sdapi/v1/options', options)
-
- return
-
-
-async def create(params):
- log.debug({ 'create start' })
- if not os.path.isdir(args.preprocess.process_dst):
- log.error({ 'train source not found': args.preprocess.process_dst })
- exit()
- if params.vectors == -1: # dynamically determine number of vectors depending on number of input images
- if len(images) <= 20:
- vectors = 2
- elif len(images) <= 100:
- vectors = 4
- else:
- vectors = 6
- else:
- vectors = params.vectors
- if os.path.exists(params.name) and os.path.isfile(params.name):
- log.info({ 'deleting existing embedding': { 'name': params.name } })
- os.remove(params.name)
- args.create_embedding.name = params.name
- words = params.init.split(',')
- if len(words) > vectors:
- params.init = ','.join(words[:vectors])
- log.warning({ 'create embedding init words cut': params.init })
- args.create_embedding.init_text = params.init
- args.create_embedding.num_vectors_per_token = vectors
- log.debug({ 'create args': args.create_embedding })
- res = await post('/sdapi/v1/create/embedding', args.create_embedding)
- if 'info' in res:
- log.info({ 'create embedding': { 'name': params.name, 'init': params.init, 'vectors': vectors, 'message': res.info } })
- else:
- log.error({ 'create failed:', res })
- return None
- log.debug({ 'create end' })
- return params.name
-
-
-async def train(params):
- log.debug({ 'train start' })
- args.train_embedding.embedding_name = params.name
-
- imgs = [f for f in os.listdir(args.preprocess.process_dst) if os.path.isfile(os.path.join(args.preprocess.process_dst, f)) and filetype.is_image(os.path.join(args.preprocess.process_dst, f))]
- args.train_embedding.data_root = args.preprocess.process_dst
- if len(imgs) == 0:
- log.error({ 'train no input images in folder': args.preprocess.process_dst })
- return
-
- if params.grad == -1:
- args.train_embedding.gradient_step = len(imgs) // args.train_embedding.batch_size
- divisor = args.train_embedding.gradient_step // 60
- args.train_embedding.gradient_step = args.train_embedding.gradient_step // (1 + divisor)
- log.info({ 'dynamic gradient step': args.train_embedding.gradient_step })
- if params.steps == -1:
- args.train_embedding.steps = params.maxsteps // args.train_embedding.gradient_step
- log.info({ 'dynamic steps': args.train_embedding.steps, 'estimated total steps': args.train_embedding.steps * args.train_embedding.gradient_step * args.train_embedding.batch_size })
-
- epoch_size = args.train_embedding.batch_size * args.train_embedding.gradient_step
- if args.train_embedding.create_image_every == -1:
- args.train_embedding.create_image_every = args.train_embedding.steps // 10
- if args.train_embedding.save_embedding_every == -1:
- args.train_embedding.save_embedding_every = args.train_embedding.steps // 10
- if args.train_embedding.learn_rate == -1:
- loss_args = {
- "steps": args.train_embedding.steps,
- "step": args.train_embedding.create_image_every,
- "loss_start": params.rstart,
- "loss_end": params.rend,
- "loss_type": 'power',
- "power": params.rdescend
- }
- args.train_embedding.learn_rate = gen_loss_rate_str(**loss_args)
- log.info({ 'dynamic learn-rate': loss_args })
- log.debug({ 'learn rate': args.train_embedding.learn_rate, 'params': loss_args })
-
- log.info({ 'train embedding': {
- 'name': params.name,
- 'source': args.preprocess.process_dst,
- 'images': len(imgs),
- 'steps': args.train_embedding.steps,
- 'batch': args.train_embedding.batch_size,
- 'gradient-step': args.train_embedding.gradient_step,
- 'sampling': args.train_embedding.latent_sampling_method,
- 'epoch-size': epoch_size }
- })
- log.info({ 'learn rate': args.train_embedding.learn_rate })
- log.debug({ 'train args': args.train_embedding })
- t0 = time.time()
- res = await post('/sdapi/v1/train/embedding', args.train_embedding)
- log.info({ 'train result': res })
- t1 = time.time()
- log.info({ 'train embedding finished': { 'name': params.name, 'time': round(t1 - t0) } })
- log.debug({ 'train end': res.info if 'info' in res else res })
- return
-
-
-async def pipeline(params):
- log.debug({ 'pipeline start' })
-
- # interrupt
- await interrupt()
-
- # preprocess
- num = await preprocess(params)
- if num == 0:
- log.warning({ 'preprocess': 'no resulting images'})
- return
-
- # create embedding
- name = await create(params)
- if not params.name in name:
- log.error({ 'create embedding failed': name })
- return
-
- # train embedding
- await train(params)
-
- await plotloss(params)
-
- # create_preview(params.name, params.init)
-
- log.debug({ 'pipeline end' })
- return
-
-
-async def monitor(params):
- step = 0
- t0 = time.perf_counter()
- t1 = time.perf_counter()
- log.info({' starting monitor': t0 })
- finished = 0
- while True:
- await asyncio.sleep(params.monitor)
- res = await progress()
- if not 'state' in res:
- log.info({ 'monitor disconnected': res })
- break
- if (res.state.job_count == params.steps and res.state.job_no >= res.state.job_count) or (res.eta_relative < 0) or (res.interrupted) or (res.state.job_count == 0): # need exit case if interrupted or failed
- if res.interrupted:
- log.info({ 'monitor interrupted': { 'embedding': params.name } })
- break # exit for monitor job
- else:
- finished += 1
- if finished >= 2: # do it more than once since preprocessing job can finish just in time for monitor to finish
- log.info({ 'monitor finished': { 'embedding': params.name } })
- break
- else:
- if res.state.job_no == 0:
- step = 0
- t0 = time.perf_counter()
- t1 = time.perf_counter()
- try:
- if 'Loss:' in res.textinfo:
- text = res.textinfo.split('
')[0].split()
- loss = float(text[-1])
- else:
- loss = -1
- except:
- loss = -1
- if math.isnan(loss):
- log.error({ 'monitor': { 'progress': round(100 * res.progress), 'embedding': params.name, 'eta': round(res.eta_relative), 'step': res.state.job_no, 'steps': res.state.job_count, 'loss': 'nan' } })
- await interrupt()
- else:
- elapsed = t1 - t0
- log.info({ 'monitor': {
- 'job': res.state.job,
- 'progress': round(100 * res.progress),
- 'embedding': params.name,
- 'epoch': (1 + res.state.job_no // len(images)) if len(images) > 0 else 'n/a',
- 'step': res.state.job_no,
- 'steps': res.state.job_count,
- 'loss': loss if loss > -1 else 'n/a',
- 'total': round(1.0 * elapsed * res.state.job_count / res.state.job_no) if res.state.job_no > 0 and t1 != t0 else 'n/a',
- 'elapsed': round(elapsed),
- 'remaining': round(res.eta_relative),
- 'it/s': round((res.state.job_no - step) / (time.perf_counter() - t1), 2) }
- })
- if step % 10 == 0:
- await plotloss(params)
- step = res.state.job_no
- t1 = time.perf_counter()
- return
-
-
-async def main():
- parser = argparse.ArgumentParser(description="sd train ti pipeline")
- parser.add_argument("--name", type = str, required = True, help = "embedding name, set to auto to use src folder name")
- parser.add_argument("--src", type = str, required = True, help = "source image folder or movie file")
- parser.add_argument("--init", type = str, default = "person", required = False, help = "initialization class, default: %(default)s")
- parser.add_argument("--dst", type = str, default = "/tmp", required = False, help = "destination image folder for processed images, default: %(default)s")
- parser.add_argument("--steps", type = int, default = -1, required = False, help = "training steps, default: %(default)s")
- parser.add_argument("--maxsteps", type = int, default = 5000, required = False, help = "max training steps used when dynamic gradient is active, default: %(default)s")
- parser.add_argument("--vectors", type = int, default = -1, required = False, help = "number of vectors per token, default: dynamic based on number of input images")
- parser.add_argument("--batch", type = int, default = 1, required = False, help = "batch size, default: %(default)s")
- parser.add_argument("--rate", type = str, default = "", required = False, help = "learn rate, default: dynamic")
- parser.add_argument("--rstart", type = float, default = 0.02, required = False, help = "starting learn rate if using dynamic rate, default: %(default)s")
- parser.add_argument("--rend", type = float, default = 0.0005, required = False, help = "ending learn rate if using dynamic rate, default: %(default)s")
- parser.add_argument("--rdescend", type = float, default = 2, required = False, help = "learn rate descend power when using dynamic rate, default: %(default)s")
- parser.add_argument("--grad", type = int, default = -1, required = False, help = "accumulate gradient over n images, default: : %(default)s")
- parser.add_argument("--type", type = str, default = 'subject', required = False, help = "training type: subject/style/unknown, default: %(default)s")
- parser.add_argument('--overwrite', default = False, action='store_true', help = "overwrite existing embedding, default: %(default)s")
- parser.add_argument("--vstart", type = float, default = 0, required = False, help = "if processing video skip first n seconds, default: %(default)s")
- parser.add_argument("--vend", type = float, default = 0, required = False, help = "if processing video skip last n seconds, default: %(default)s")
- parser.add_argument('--skipcaption', default = False, action='store_true', help = "do not auto-generate captions, default: %(default)s")
- parser.add_argument('--skipmodel', default = False, action='store_true', help = "skip model validation and switch, default: %(default)s")
- parser.add_argument('--preprocess', type = str, choices=['builtin', 'custom', 'none'], default = 'custom', help = "preprocessing type, default: %(default)s")
- parser.add_argument('--nocleanup', default = False, action='store_true', help = "skip cleanup after completion, default: %(default)s")
- parser.add_argument("--monitor", type = int, default = 30, required = False, help = "progress monitor frequency, default: : %(default)s")
- parser.add_argument('--debug', default = False, action='store_true', help = "print extra debug information, default: %(default)s")
- params = parser.parse_args()
- if params.debug:
- log.setLevel(logging.DEBUG)
- log.debug({ 'debug': True })
- log.debug({ 'args': params.__dict__ })
- home = Path(sys.argv[0]).parent
- global args # pylint: disable=global-statement
- if params.vstart > 0:
- args.extract_video.vstart = params.vstart
- if params.vend > 0:
- args.extract_video.vend = params.vend
- if params.steps > -1:
- args.train_embedding.steps = params.steps
- if params.batch > -1:
- args.train_embedding.batch_size = params.batch
- if params.rate != '':
- args.train_embedding.learn_rate = params.rate
- if params.grad > -1:
- args.train_embedding.gradient_step = params.grad
- if params.type == 'subject':
- if params.skipcaption:
- args.train_embedding.template_filename = 'subject.txt'
- args.preprocess.process_caption = False
- else:
- args.train_embedding.template_filename = 'subject_filewords.txt'
- elif params.type == 'style':
- if params.skipcaption:
- args.train_embedding.template_filename = 'style.txt'
- args.preprocess.process_caption = False
- else:
- args.train_embedding.template_filename = 'style_filewords.txt'
- else:
- if params.skipcaption:
- args.train_embedding.template_filename = 'unknown.txt'
- args.preprocess.process_caption = False
- else:
- args.train_embedding.template_filename = 'unknown_filewords.txt'
- if params.name == 'auto':
- params.name = PurePath(params.src).name
- log.info({ 'training name': params.name })
- if params.dst == "/tmp":
- params.dst = os.path.join("/tmp/train", params.name)
- log.debug({ 'args': params.__dict__ })
- params.src = os.path.abspath(params.src)
- params.dst = os.path.abspath(params.dst)
-
- try:
- await session()
- await check(params)
- a = asyncio.create_task(pipeline(params))
- b = asyncio.create_task(monitor(params))
- await asyncio.gather(a, b) # wait for both pipeline and monitor to finish
- except Exception as e:
- log.error({ 'exception': e })
- finally:
- if not params.nocleanup:
- await preprocess_cleanup(params)
- await close()
- return
-
-if __name__ == "__main__":
- log.info({ 'train textual inversion' })
- try:
- asyncio.run(main())
- except KeyboardInterrupt:
- log.warning({ 'interrupted': 'keyboard request' })
- # asyncio.run(interrupt())
diff --git a/cli/train/console.py b/cli/train/console.py
deleted file mode 100644
index e69de29bb..000000000
diff --git a/cli/train/latents.py b/cli/train/latents.py
index b6a5787aa..b0509f2cd 100755
--- a/cli/train/latents.py
+++ b/cli/train/latents.py
@@ -44,13 +44,13 @@ options = Map({
vae = None
-def get_latents(vae, images, weight_dtype):
+def get_latents(local_vae, images, weight_dtype):
image_transforms = transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.5], [0.5]) ])
img_tensors = [image_transforms(image) for image in images]
img_tensors = torch.stack(img_tensors)
img_tensors = img_tensors.to(device, weight_dtype)
with torch.no_grad():
- latents = vae.encode(img_tensors).latent_dist.sample().float().to('cpu').numpy()
+ latents = local_vae.encode(img_tensors).latent_dist.sample().float().to('cpu').numpy()
return latents
@@ -58,8 +58,8 @@ def get_npz_filename_wo_ext(data_dir, image_key):
return os.path.join(data_dir, os.path.splitext(os.path.basename(image_key))[0])
-def create_vae_latents(params):
- args = Map({**options, **params})
+def create_vae_latents(local_params):
+ args = Map({**options, **local_params})
console.log(f'create vae latents args: {args}')
image_paths = train_util.glob_images(args.input)
if os.path.exists(args.json):
@@ -73,7 +73,7 @@ def create_vae_latents(params):
weight_dtype = torch.bfloat16
else:
weight_dtype = torch.float32
- global vae
+ global vae # pylint: disable=global-statement
if vae is None:
vae = model_util.load_vae(args.vae, weight_dtype)
vae.eval()
@@ -142,7 +142,7 @@ def create_vae_latents(params):
def unload_vae():
- global vae
+ global vae # pylint: disable=global-statement
vae = None
diff --git a/cli/train/process.py b/cli/train/process.py
index 8aedcd872..a2104bf0c 100644
--- a/cli/train/process.py
+++ b/cli/train/process.py
@@ -1,15 +1,13 @@
+ # pylint: disable=global-statement
import os
-import sys
import io
import math
import base64
-import pathlib
import numpy as np
import mediapipe as mp
from PIL import Image, ImageOps
from skimage.metrics import structural_similarity as ssim
from scipy.stats import beta
-sys.path.append(os.path.join(os.path.dirname(__file__)))
import util
import sdapi
@@ -23,9 +21,9 @@ all_images_by_type = {}
class Result(object):
- def __init__(self, type: str, input: str, tag: str = None, requested: list = []):
- self.type = type
- self.input = input
+ def __init__(self, typ: str, fn: str, tag: str = None, requested: list = []):
+ self.type = typ
+ self.input = fn
self.output = ''
self.basename = ''
self.message = ''
@@ -56,8 +54,8 @@ def detect_dynamicrange(image: Image):
data = np.asarray(image)
image = np.float32(data)
RGB = [0.299, 0.587, 0.114]
- height, width = image.shape[:2]
- brightness_image = np.sqrt(image[..., 0] ** 2 * RGB[0] + image[..., 1] ** 2 * RGB[1] + image[..., 2] ** 2 * RGB[2])
+ height, width = image.shape[:2] # pylint: disable=unsubscriptable-object
+ brightness_image = np.sqrt(image[..., 0] ** 2 * RGB[0] + image[..., 1] ** 2 * RGB[1] + image[..., 2] ** 2 * RGB[2]) # pylint: disable=unsubscriptable-object
hist, _ = np.histogram(brightness_image, bins=256, range=(0, 255))
img_brightness_pmf = hist / (height * width)
dist = beta(2, 2)
@@ -264,7 +262,7 @@ def save_image(res: Result, folder: str):
def file(filename: str, folder: str, tag = None, requested = []):
# initialize result dict
- res = Result(input = filename, type='unknown', tag=tag, requested = requested)
+ res = Result(fn = filename, typ='unknown', tag=tag, requested = requested)
# open image
try:
res.image = Image.open(filename)
diff --git a/cli/train/sdapi.py b/cli/train/sdapi.py
index e85148e59..f642e1bad 100644
--- a/cli/train/sdapi.py
+++ b/cli/train/sdapi.py
@@ -1,7 +1,5 @@
-import sys
-import json
-import aiohttp
import asyncio
+import aiohttp
import requests
from util import Map
@@ -89,9 +87,9 @@ def progress():
def options():
- options = getsync('/sdapi/v1/options')
+ opt = getsync('/sdapi/v1/options')
flags = getsync('/sdapi/v1/cmd-flags')
- return { 'options': options, 'flags': flags }
+ return { 'options': opt, 'flags': flags }
def shutdown():
diff --git a/cli/train/train.py b/cli/train/train.py
index 29eebf0ad..ea521adaf 100755
--- a/cli/train/train.py
+++ b/cli/train/train.py
@@ -46,8 +46,6 @@ lycoris_path = os.path.abspath(os.path.join(os.path.dirname(__file__), os.pardir
sys.path.append(lycoris_path)
import train_network
-print('HERE6')
-
# globals
args = None
valid_steps = ['original', 'face', 'body', 'blur', 'range', 'upscale', 'restore', 'interrogate', 'resize', 'square', 'segment']
@@ -69,26 +67,29 @@ def mem_stats():
def parse_args():
global args # pylint: disable=global-statement
- parser = argparse.ArgumentParser(description = 'train')
- # basic section
- parser.add_argument('--type', type=str, choices=['embedding', 'lora', 'lycoris', 'dreambooth'], default=None, required=True, help='training type')
- parser.add_argument('--name', type=str, default=None, required=True, help='output filename')
- parser.add_argument('--overwrite', default = False, action='store_true', help = "overwrite existing training, default: %(default)s")
- parser.add_argument('--tag', type=str, default='person', required=False, help='primary tags, default: %(default)s')
- parser.add_argument('--input', type=str, default=None, required=True, help='input folder with training images')
- parser.add_argument('--output', type=str, default='', required=False, help='where to store processed images, default is system temp/train')
- parser.add_argument('--process', type=str, default='original,interrogate,resize,square', required=False, help=f'list of possible processing steps: {valid_steps}, default: %(default)s')
+ parser = argparse.ArgumentParser(description = 'Train')
- # global params
- parser.add_argument('--gradient', type=int, default=1, required=False, help='gradient accumulation steps, default: %(default)s')
- parser.add_argument('--steps', type=int, default=2500, required=False, help='training steps, default: %(default)s')
- parser.add_argument('--batch', type=int, default=1, required=False, help='batch size, default: %(default)s')
- parser.add_argument('--lr', type=float, default=1e-04, required=False, help='model learning rate, default: %(default)s')
- parser.add_argument('--dim', type=int, default=40, required=False, help='network dimension or number of vectors, default: %(default)s')
+ group_main = parser.add_argument_group('Main')
+ group_main.add_argument('--type', type=str, choices=['embedding', 'lora', 'lycoris', 'dreambooth'], default=None, required=True, help='training type')
+ group_main.add_argument('--name', type=str, default=None, required=True, help='output filename')
+ group_main.add_argument('--overwrite', default = False, action='store_true', help = "overwrite existing training, default: %(default)s")
+ group_main.add_argument('--tag', type=str, default='person', required=False, help='primary tags, 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')
+ group_data.add_argument('--output', type=str, default='', required=False, help='where to store processed images, default is system temp/train')
+ group_data.add_argument('--process', type=str, default='original,interrogate,resize,square', required=False, help=f'list of possible processing steps: {valid_steps}, default: %(default)s')
+
+ group_train = parser.add_argument_group('Train')
+ group_train.add_argument('--gradient', type=int, default=1, required=False, help='gradient accumulation steps, default: %(default)s')
+ group_train.add_argument('--steps', type=int, default=2500, required=False, help='training steps, default: %(default)s')
+ group_train.add_argument('--batch', type=int, default=1, required=False, help='batch size, default: %(default)s')
+ group_train.add_argument('--lr', type=float, default=1e-04, required=False, help='model learning rate, default: %(default)s')
+ group_train.add_argument('--dim', type=int, default=40, required=False, help='network dimension or number of vectors, default: %(default)s')
# lora params
- parser.add_argument('--repeats', type=int, default=10, required=False, help='number of repeats per image, default: %(default)s')
- parser.add_argument('--alpha', type=float, default=0, required=False, help='alpha for weights scaling, default: dim/2')
+ group_train.add_argument('--repeats', type=int, default=10, required=False, help='number of repeats per image, default: %(default)s')
+ group_train.add_argument('--alpha', type=float, default=0, required=False, help='alpha for weights scaling, default: dim/2')
args = parser.parse_args()
diff --git a/cli/train/util.py b/cli/train/util.py
index 944cca106..e67b4f403 100755
--- a/cli/train/util.py
+++ b/cli/train/util.py
@@ -42,8 +42,8 @@ def get_memory():
return Map(mem)
-class Map(dict):
- __slots__ = ('__dict__')
+class Map(dict): # pylint: disable=C0205
+ __slots__ = ('__dict__') # pylint: disable=C0325
def __init__(self, *args, **kwargs):
super(Map, self).__init__(*args, **kwargs)
for arg in args:
diff --git a/cli/modules/util.py b/cli/util.py
similarity index 97%
rename from cli/modules/util.py
rename to cli/util.py
index ef185bfa4..c1ca29c94 100755
--- a/cli/modules/util.py
+++ b/cli/util.py
@@ -66,8 +66,8 @@ def get_memory():
return Map(mem)
-class Map(dict):
- __slots__ = ('__dict__')
+class Map(dict): # pylint: disable=C0205
+ __slots__ = ('__dict__') # pylint: disable=C0325
def __init__(self, *args, **kwargs):
super(Map, self).__init__(*args, **kwargs)
for arg in args:
diff --git a/cli/modules/video-extract.py b/cli/video-extract.py
similarity index 97%
rename from cli/modules/video-extract.py
rename to cli/video-extract.py
index edd0caf31..9bc4544e6 100755
--- a/cli/modules/video-extract.py
+++ b/cli/video-extract.py
@@ -16,8 +16,8 @@ def probe(src: str):
result = subprocess.run(cmd, shell = True, capture_output = True, text = True, check = True)
data = json.loads(result.stdout)
stream = [x for x in data['streams'] if x["codec_type"] == "video"][0]
- format = data['format'] if 'format' in data else {}
- res = {**stream, **format}
+ fmt = data['format'] if 'format' in data else {}
+ res = {**stream, **fmt}
video = Map({
'codec': res.get('codec_name', 'unknown') + '/' + res.get('codec_tag_string', ''),
'resolution': [int(res.get('width', 0)), int(res.get('height', 0))],
diff --git a/cli/xformers.sh b/cli/xformers.sh
deleted file mode 100755
index 5bd352db9..000000000
--- a/cli/xformers.sh
+++ /dev/null
@@ -1,12 +0,0 @@
-#!/usr/bin/env bash
-echo "Installing xformers"
-
-NVCC_FLAGS="--use_fast_math"
-FORCE_CUDA="1"
-TORCH_CUDA_ARCH_LIST="8.6"
-pip install ninja -q
-pip uninstall xformers -y 2>/dev/null
-pip install -v -U git+https://github.com/facebookresearch/xformers.git@main#egg=xformers
-pip show torch
-pip show xformers
-python -m xformers.info
diff --git a/extensions-builtin/sd-webui-controlnet b/extensions-builtin/sd-webui-controlnet
index eaf993908..721159f7b 160000
--- a/extensions-builtin/sd-webui-controlnet
+++ b/extensions-builtin/sd-webui-controlnet
@@ -1 +1 @@
-Subproject commit eaf993908310ac0324f3a024438b8c47196d7056
+Subproject commit 721159f7b7f224db06b5f712167bbc315d5cd835
diff --git a/modules/api/api.py b/modules/api/api.py
index 1405c7a63..485a217c7 100644
--- a/modules/api/api.py
+++ b/modules/api/api.py
@@ -398,13 +398,15 @@ class Api:
options.update({k: shared.opts.data.get(k, shared.opts.data_labels.get(k).default)})
else:
options.update({k: shared.opts.data.get(k, None)})
-
+ if 'sd_lyco' in options:
+ del options['sd_lyco']
+ if 'sd_lora' in options:
+ del options['sd_lora']
return options
def set_config(self, req: Dict[str, Any]):
for k, v in req.items():
shared.opts.set(k, v)
-
shared.opts.save(shared.config_filename)
return
@@ -573,7 +575,7 @@ class Api:
ram = { 'error': f'{err}' }
try:
import torch
- if shared.cmd_opts.use_ipex():
+ if shared.cmd_opts.use_ipex:
system = { 'free': (torch.xpu.get_device_properties("xpu").total_memory - torch.xpu.memory_allocated()), 'used': torch.xpu.memory_allocated(), 'total': torch.xpu.get_device_properties("xpu").total_memory }
s = dict(torch.xpu.memory_stats("xpu"))
allocated = { 'current': s['allocated_bytes.all.current'], 'peak': s['allocated_bytes.all.peak'] }
diff --git a/modules/shared.py b/modules/shared.py
index 6c35c9252..ed0a1b4b3 100644
--- a/modules/shared.py
+++ b/modules/shared.py
@@ -712,7 +712,7 @@ def restore_defaults(restart=True):
if os.path.exists(cmd_opts.ui_config):
log.info('Restoring UI defaults')
os.remove(cmd_opts.ui_config)
- restart_server(True)
+ restart_server(restart)
def listfiles(dirname):