update preview generation

This commit is contained in:
Vladimir Mandic
2023-02-07 09:53:19 -05:00
parent 5c6c0a11c5
commit f756da1596
4 changed files with 140 additions and 71 deletions
+53 -5
View File
@@ -5,7 +5,9 @@ import json
import time
import asyncio
import argparse
from pathlib import Path
sys.path.append(os.path.join(os.path.dirname(__file__), '..'))
sys.path.append(os.path.join(os.path.dirname(__file__), 'modules'))
from generate import sd, generate
from modules.util import Map, log
@@ -16,7 +18,7 @@ from modules.grid import grid
default = 'sd-v15-runwayml.ckpt [cc6cb27103]'
embeddings = ['blonde', 'bruntette', 'sexy', 'naked', 'mia', 'lin', 'kelly', 'hanna', 'rreid-random-v0']
exclude = ['sd-v20', 'sd-v21', 'inpainting', 'pix2pix']
prompt = "photo of beautiful woman <embedding>, photograph, posing, pose, high detailed, intricate, elegant, sharp focus, skin texture, looking forward, facing camera, 135mm, shot on dslr, canon 5d, 4k, modelshoot style, cinematic lighting"
prompt = "photo of <keyword> <embedding>, photograph, posing, pose, high detailed, intricate, elegant, sharp focus, skin texture, looking forward, facing camera, 135mm, shot on dslr, canon 5d, 4k, modelshoot style, cinematic lighting"
options = Map({
'generate': {
'restore_faces': True,
@@ -29,17 +31,20 @@ options = Map({
'sampler_name': 'DPM2 Karras',
'cfg_scale': 7,
'width': 512,
'height': 512
'height': 512,
},
'paths': {
"root": "/mnt/c/Users/mandi/OneDrive/Generative/Generate",
"generate": "image",
"upscale": "upscale",
"grid": "grid"
"grid": "grid",
},
'options': {
"sd_model_checkpoint": "sd-v15-runwayml",
"sd_vae": "vae-ft-mse-840000-ema-pruned.ckpt"
"sd_vae": "vae-ft-mse-840000-ema-pruned.ckpt",
},
'lora': {
'strength': 0.8,
}
})
@@ -97,6 +102,7 @@ async def models(params):
t0 = time.time()
for embedding in embeddings:
options.generate.prompt = prompt.replace('<embedding>', f'\"{embedding}\"')
options.generate.prompt = options.generate.prompt.replace('<keyword>', 'beautiful woman')
log.info({ 'model generating': model, 'embedding': embedding, 'prompt': options.generate.prompt })
data = await generate(options = options, quiet=True)
if 'image' in data:
@@ -118,6 +124,48 @@ async def models(params):
opt['sd_model_checkpoint'] = default
await post('/sdapi/v1/options', opt)
async def lora(params):
cmdflags = await get('/sdapi/v1/cmd-flags')
dir = cmdflags['lora_dir']
if not os.path.exists(dir):
log.error({ 'lora directory not found': dir })
return
models1 = [f for f in Path(dir).glob('*.safetensors')]
models2 = [f for f in Path(dir).glob('*.ckpt')]
models = [f.stem for f in models1 + models2]
log.info({ 'loras': len(models) })
for model in models:
fn = os.path.join(dir, model + '.png')
if os.path.exists(fn) and len(params.input) == 0: # if model preview exists and not manually included
log.info({ 'lora preview exists': model })
continue
images = []
labels = []
t0 = time.time()
keyword = model.replace('-', ' ')
options.generate.prompt = prompt.replace('<keyword>', f'\"{keyword}\"')
options.generate.prompt = options.generate.prompt.replace('<embedding>', '')
options.generate.prompt += f' <lora:{model}:{options.lora.strength}>'
log.info({ 'lora 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({ 'lora': model, 'embedding': 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({ 'lora 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 models(params)
await lora(params)
await close()
@@ -126,4 +174,4 @@ if __name__ == '__main__':
parser.add_argument('--output', type = str, default = '', required = False, help = 'output directory')
parser.add_argument('input', type = str, nargs = '*')
params = parser.parse_args()
asyncio.run(models(params))
asyncio.run(create_previews(params))