update lora script

This commit is contained in:
Vladimir Mandic
2023-02-26 11:05:53 -05:00
parent 1f714beccb
commit 99142d6477
2 changed files with 26 additions and 6 deletions
+18 -4
View File
@@ -206,7 +206,7 @@ def extract_face(img):
if params.square_images:
squared = Image.new('RGB', (params.target_size, params.target_size))
squared.paste(cropped, (0, 0))
squared.paste(cropped, ((params.target_size - cropped.width) // 2, (params.target_size - cropped.height) // 2))
if params.face_segmentation:
squared = segmentation(squared)
else:
@@ -269,7 +269,7 @@ def extract_body(img):
if params.square_images:
squared = Image.new('RGB', (params.target_size, params.target_size))
squared.paste(cropped, (0, 0))
squared.paste(cropped, ((params.target_size - cropped.width) // 2, (params.target_size - cropped.height) // 2))
if params.body_segmentation:
squared = segmentation(squared)
else:
@@ -297,6 +297,19 @@ def extract_body(img):
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')
@@ -391,8 +404,9 @@ def process_file(f: str, dst: str = None, preview: bool = False, offline: bool =
log.debug({ 'no body': f })
if params.keep_original:
fn = save(image, f, 'original')
log.info({ 'original': fn })
resized = save_original(image)
fn = save(resized, f, 'original')
log.info({ 'keep original': fn })
image.close()
return i, metadata
+8 -2
View File
@@ -68,6 +68,7 @@ options = Map({
"save_last_n_epochs_state": None,
"save_state": False,
"resume": None,
"max_grad_norm": 0.0,
"train_batch_size": 1,
"max_token_length": None,
"use_8bit_adam": False,
@@ -144,6 +145,7 @@ if __name__ == '__main__':
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('--debug', default=False, action='store_true', help = "enable debug logging")
args = parser.parse_args()
if args.debug:
@@ -180,7 +182,11 @@ if __name__ == '__main__':
Path(dir).mkdir(parents=True, exist_ok=True)
json_file = os.path.join(tempfile.gettempdir(), args.output, args.output + '.json')
options.train_data_dir = os.path.join(tempfile.gettempdir(), args.output)
options.in_json = json_file
if args.dreambooth:
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]
@@ -203,7 +209,7 @@ if __name__ == '__main__':
with open(json_file, "w") as outfile:
outfile.write(json.dumps(metadata, indent=2))
if not args.nolatents:
if not args.nolatents and not args.dreambooth:
# create latents
latents.create_vae_latents(Map({ 'input': dir, 'json': json_file }))
latents.unload_vae()