mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
update cli
This commit is contained in:
@@ -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
|
||||
+30
-85
@@ -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
|
||||
|
||||
<br>
|
||||
|
||||
@@ -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
|
||||
|
||||
<br>
|
||||
|
||||
@@ -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
|
||||
|
||||
<br>
|
||||
|
||||
## Utility Scripts
|
||||
|
||||
### SDAPI
|
||||
|
||||
Utility module that handles async communication to Automatic API endpoints
|
||||
|
||||
@@ -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>', keyword)
|
||||
options.generate.prompt = options.generate.prompt.replace('<embedding>', '')
|
||||
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()
|
||||
|
||||
|
||||
+2
-3
@@ -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 = {}
|
||||
|
||||
@@ -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:
|
||||
@@ -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)
|
||||
@@ -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__':
|
||||
@@ -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()
|
||||
@@ -1,144 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
"""
|
||||
Extract approximating LoRA by SVD from two SD models
|
||||
Based on: <https://github.com/kohya-ss/sd-scripts/blob/main/networks/extract_lora_from_models.py>
|
||||
"""
|
||||
|
||||
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)
|
||||
@@ -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))
|
||||
@@ -1,74 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
# based on <https://huggingface.co/JosephusCheung/ASimilarityCalculatior>
|
||||
|
||||
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()
|
||||
@@ -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')
|
||||
Binary file not shown.
|
Before Width: | Height: | Size: 7.1 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 8.0 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 7.6 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 9.1 KiB |
@@ -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 <https://github.com/karthik9319/Blur-Detection/>
|
||||
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 <https://towardsdatascience.com/measuring-enhancing-image-quality-attributes-234b0f250e10>
|
||||
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))
|
||||
@@ -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'})
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
@@ -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 = []
|
||||
@@ -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")
|
||||
@@ -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)
|
||||
@@ -2,3 +2,4 @@ mediapipe
|
||||
colormap
|
||||
invisible-watermark
|
||||
filetype
|
||||
albumentations
|
||||
|
||||
@@ -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))
|
||||
@@ -1,274 +0,0 @@
|
||||
#!/usr/bin/env python
|
||||
|
||||
"""
|
||||
Extract approximating LoRA by SVD from two SD models
|
||||
Based on: <https://github.com/kohya-ss/sd-scripts/blob/main/networks/train_network.py>
|
||||
|
||||
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()
|
||||
-591
@@ -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('<br/>')[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())
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
+3
-5
@@ -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():
|
||||
|
||||
+20
-19
@@ -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()
|
||||
|
||||
|
||||
+2
-2
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
@@ -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))],
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user