acmyu/KeyframesAI
0
1import logging
2import math
3import os
4from typing import Any, Dict, List, Optional, Tuple, Union
5#from diffusers.models.controlnet import ControlNetConditioningEmbedding
6from diffusers.models.controlnets.controlnet import ControlNetConditioningEmbedding
7import torch
8from torch import nn
9import torch.nn.functional as F
10import torch.utils.checkpoint
11import transformers
12from accelerate import Accelerator
13from accelerate.logging import get_logger
14from accelerate.utils import ProjectConfiguration, set_seed
15
16from tqdm.auto import tqdm
17from src.configs.stage2_config import args
18
19import diffusers
20from diffusers import (
21 AutoencoderKL,
22 DDPMScheduler,
23)
24from diffusers.optimization import get_scheduler
25from diffusers.utils import check_min_version, is_wandb_available
26from src.dataset.stage2_dataset import InpaintDataset, InpaintCollate_fn
27from transformers import CLIPVisionModelWithProjection
28from transformers import Dinov2Model
29from src.models.stage2_inpaint_unet_2d_condition import Stage2_InapintUNet2DConditionModel
30
31
32
33import glob
34import os
35import torch
36from torch import nn
37from PIL import Image, ImageOps
38import numpy as np
39from diffusers import UniPCMultistepScheduler
40from src.models.stage2_inpaint_unet_2d_condition import Stage2_InapintUNet2DConditionModel
41
42from torchvision import transforms
43#from diffusers.models.controlnet import ControlNetConditioningEmbedding
44from transformers import CLIPImageProcessor
45from transformers import Dinov2Model
46from diffusers import AutoencoderKL, DDPMScheduler, UNet2DConditionModel,ControlNetModel,DDIMScheduler
47from src.pipelines.PCDMs_pipeline import PCDMsPipeline
48#from single_extract_pose import inference_pose
49
50
51import spaces
52from libs.easy_dwpose import DWposeDetector
53from libs.easy_dwpose.draw import draw_openpose
54from libs.film import Predictor
55from PIL import Image
56import cv2
57import os
58import gradio as gr
59import rembg
60import uuid
61import gc
62from numba import cuda
63import requests
64import json
65
66from huggingface_hub import hf_hub_download, HfApi
67from numba import cuda
68from multiprocessing import Pool, Process, Queue
69import torch.multiprocessing as mp
70
71# Inputs ===================================================================================================
72
73input_img = "sm.png"
74train_imgs = ["target.png"]
75in_vid = "walk.mp4"
76out_vid = 'out.mp4'
77
78"""
79train_steps = 100
80inference_steps = 10
81fps = 12
82"""
83
84debug = False
85save_model = True
86should_gen_vid = False
87max_batch_size = 8
88max_frame_count = 200
89no_bg_final = True
90
91def save_temp_imgs(imgs):
92 os.makedirs('temp', exist_ok=True)
93 results = []
94
95 api = HfApi()
96
97
98 for i, img in enumerate(imgs):
99
100 #img_name = 'temp/'+str(uuid.uuid4())+'.png'
101 img_name = 'temp/'+str(i)+'.png'
102 img.save(img_name)
103
104 """
105 url = 'https://tmpfiles.org/api/v1/upload'
106
107 try:
108 response = requests.post(url, files={'file': open(img_name, 'rb')})
109
110 # Check for successful response (status code 200)
111 response.raise_for_status()
112
113 # Print the server's response
114 print("Status Code:", response.status_code)
115
116 data = response.json()
117 print("Response JSON:", data)
118 results.append(data['data']['url'])
119
120 except requests.exceptions.RequestException as e:
121 print(f"An error occurred: {e}")
122 """
123
124 results.append('https://huggingface.co/datasets/acmyu/KeyframesAIFiles/resolve/main/'+img_name)
125
126 api.upload_file(
127 path_or_fileobj='temp',
128 path_in_repo='temp',
129 repo_id="acmyu/KeyframesAIFiles",
130 repo_type="dataset",
131 )
132
133 return results
134
135
136def getThumbnails(imgs):
137 thumbs = []
138 thumb_size = (512, 512)
139 for img in imgs:
140 th = img.copy()
141 th.thumbnail(thumb_size)
142 thumbs.append(th)
143 return thumbs
144
145
146# Pose detection ==============================================================================================
147
148def load_models():
149 dwpose = DWposeDetector(device="cuda")
150 rembg_session = rembg.new_session("u2netp")
151
152 pcdms_model = hf_hub_download(repo_id="acmyu/PCDMs", filename="pcdms_ckpt.pt")
153
154 # Load scheduler
155 noise_scheduler = DDPMScheduler.from_pretrained("stabilityai/stable-diffusion-2-1-base", subfolder="scheduler")
156
157 # Load model
158 image_encoder_p = Dinov2Model.from_pretrained('facebook/dinov2-giant')
159 image_encoder_g = CLIPVisionModelWithProjection.from_pretrained('laion/CLIP-ViT-H-14-laion2B-s32B-b79K')#("openai/clip-vit-base-patch32")
160
161 vae = AutoencoderKL.from_pretrained("stabilityai/stable-diffusion-2-1-base", subfolder="vae")
162 unet = Stage2_InapintUNet2DConditionModel.from_pretrained(
163 "stabilityai/stable-diffusion-2-1-base",
164 torch_dtype=torch.float16,
165 subfolder="unet",
166 in_channels=9,
167 low_cpu_mem_usage=False,
168 ignore_mismatched_sizes=True)
169
170
171 return dwpose, rembg_session, pcdms_model, noise_scheduler, image_encoder_p, image_encoder_g, vae, unet
172
173
174#load_models()
175
176def img_pad(img, tw, th, transparent=False):
177 #print('pad', tw, th)
178 img.thumbnail((tw, th))
179 if transparent:
180 new_img = Image.new('RGBA', (tw, th), (0, 0, 0, 0))
181 else:
182 new_img = Image.new("RGB", (tw, th), (0, 0, 0))
183 left = (tw - img.width) // 2
184 top = (th - img.height) // 2
185 #print(left, top)
186 new_img.paste(img, (left, top))
187 return new_img
188
189
190def resize_pad(img, tw, th, transparent):
191 w, h = img.size
192 orig_tw = tw
193 orig_th = th
194
195 if tw/th > w/h:
196 tw = int(th * w/h)
197 elif tw/th < w/h:
198 th = int(tw * h/w)
199
200 img = img.resize((tw, th), Image.BICUBIC)
201
202 return img_pad(img, orig_tw, orig_th, True)
203
204
205def resize_and_pad(img, target_img):
206 tw, th = target_img.size
207 return resize_pad(img, tw, th, False)
208
209
210def remove_zero_pad(image):
211 image = np.array(image)
212 dummy = np.argwhere(image != 0) # assume blackground is zero
213 max_y = dummy[:, 0].max()
214 min_y = dummy[:, 0].min()
215 min_x = dummy[:, 1].min()
216 max_x = dummy[:, 1].max()
217 crop_image = image[min_y:max_y, min_x:max_x]
218
219 return Image.fromarray(crop_image)
220
221
222def get_pose(img, dwpose, outfile, crop=False):
223 #pil_image = Image.open("imgs/"+img).convert("RGB")
224 #skeleton = dwpose(pil_image, output_type="np", include_hands=True, include_face=False)
225
226 img.thumbnail((512,512))
227 out_img, pose = dwpose(img, include_hands=True, include_face=False)
228
229 #print(pose['bodies'])
230
231 if crop:
232 bbox = out_img.getbbox()
233 out_img = out_img.crop(bbox)
234 out_img = ImageOps.expand(out_img, border=int(out_img.width*0.2), fill=(0,0,0))
235
236 return out_img, pose
237
238
239def extract_frames(video_path, fps):
240 video_capture = cv2.VideoCapture(video_path)
241 frame_count = 0
242 frames = []
243
244 fps_in = video_capture.get(cv2.CAP_PROP_FPS)
245 fps_out = fps
246
247 index_in = -1
248 index_out = -1
249
250 while True:
251 success = video_capture.grab()
252 if not success: break
253 index_in += 1
254
255 if frame_count > max_frame_count:
256 break
257
258 out_due = int(index_in / fps_in * fps_out)
259 if out_due > index_out:
260 success, frame = video_capture.retrieve()
261 if not success:
262 break
263 index_out += 1
264
265 frame_count += 1
266 frames.append(Image.fromarray(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)))
267
268 video_capture.release()
269 print(f"Extracted {frame_count} frames")
270 return frames
271
272
273def removebg(img, rembg_session, transparent=False):
274
275 if transparent:
276 result = Image.new('RGBA', img.size, (0, 0, 0, 0))
277 else:
278 result = Image.new("RGB", img.size, "#ffffff")
279 out = rembg.remove(img, session=rembg_session)
280 result.paste(out, mask=out)
281 return result
282
283
284def prepare_inputs_train(images, bg_remove, dwpose, rembg_session):
285 print("remove background", bg_remove)
286 if bg_remove:
287 images = [removebg(img, rembg_session) for img in images]
288
289 in_img = images[0]
290 in_pose, _ = get_pose(in_img, dwpose, "in_pose.png")
291 train_poses = []
292 train_imgs = [resize_and_pad(img, in_img) for img in images[1:]]
293
294 for i, img in enumerate(train_imgs):
295 train_pose, _ = get_pose(img, dwpose, "tr_pose"+str(i)+".png")
296 train_poses.append(train_pose)
297
298 return in_img, in_pose, train_imgs, train_poses
299
300
301def prepare_inputs_inference(in_img, in_vid, frames, fps, dwpose, rembg_session, bg_remove, resize_inputs, is_app=False, target_poses=None):
302 progress=gr.Progress(track_tqdm=True)
303
304 print("prepare_inputs_inference")
305
306 in_pose, _ = get_pose(in_img, dwpose, "in_pose.png")
307
308 print(in_vid)
309 print(frames)
310 if in_vid:
311 frames = extract_frames(in_vid, fps)
312 for f in frames:
313 f.thumbnail((512,512))
314
315 print("remove background", bg_remove)
316 if bg_remove:
317 in_img = removebg(in_img, rembg_session)
318 #frames = [removebg(img, rembg_session) for img in frames]
319 if debug:
320 for i, frame in enumerate(frames):
321 frame.save("out/frame_"+str(i)+".png")
322
323 print("vid: ", in_vid, fps)
324
325 progress_bar = tqdm(range(len(frames)), initial=0, desc="Frames")
326 if not target_poses:
327 target_poses = []
328 target_poses_coords = []
329 max_left = max_top = 999999
330 max_right = max_bottom = 0
331 it = frames
332 if is_app:
333 it = progress.tqdm(frames, desc="Pose Detection")
334 for f in it:
335 tpose, tpose_coords = get_pose(f, dwpose, "tar_pose"+str(len(target_poses))+".png")
336 #print(tpose_coords)
337 coords = {}
338 for k in tpose_coords:
339 if k == 'bodies_multi':
340 coords['bodies'] = tpose_coords[k].tolist()
341 elif k in ['hands']:
342 coords[k] = tpose_coords[k].tolist()
343 elif k in ['num_candidates']:
344 coords[k] = tpose_coords[k]
345 #print(coords)
346 target_poses.append(tpose)
347 target_poses_coords.append(json.dumps(coords))
348 progress_bar.update(1)
349
350
351 target_poses_cropped = []
352 for tpose in target_poses:
353 if resize_inputs:
354 bbox = tpose.getbbox()
355 left, top, right, bottom = bbox
356 max_left = min(max_left, left)
357 max_top = min(max_top, top)
358 max_right = max(max_right, right)
359 max_bottom = max(max_bottom, bottom)
360
361 tpose = tpose.crop((max_left, max_top, max_right, max_bottom))
362 tpose = ImageOps.expand(tpose, border=int(tpose.width*0.2), fill=(0,0,0))
363
364 tpose = resize_and_pad(tpose, in_img)
365
366
367 if debug:
368 tpose.save("out/"+"tar_pose"+str(len(target_poses_cropped))+".png")
369 target_poses_cropped.append(tpose)
370
371 #target_poses_cropped[0].save("pose.png")
372 return in_img, target_poses_cropped, in_pose, target_poses_coords, frames
373
374
375def prepare_inputs(images, in_vid, fps, bg_remove, dwpose, rembg_session, resize_inputs, is_app=False):
376
377 in_img, in_pose, train_imgs, train_poses = prepare_inputs_train(images, bg_remove, dwpose, rembg_session)
378
379 in_img, target_poses_cropped, _, _, _ = prepare_inputs_inference(in_img, in_vid, [], fps, dwpose, rembg_session, bg_remove, resize_inputs, is_app)
380
381
382 return in_img, in_pose, train_imgs, train_poses, target_poses_cropped
383
384
385# Training ===================================================================================================
386
387# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
388check_min_version("0.18.0.dev0")
389
390logger = get_logger(__name__)
391
392
393class ImageProjModel_p(torch.nn.Module):
394 """SD model with image prompt"""
395
396 def __init__(self, in_dim, hidden_dim, out_dim, dropout = 0.):
397 super().__init__()
398
399 self.net = nn.Sequential(
400 nn.Linear(in_dim, hidden_dim),
401 nn.GELU(),
402 nn.Dropout(dropout),
403 nn.LayerNorm(hidden_dim),
404 nn.Linear(hidden_dim, out_dim),
405 nn.Dropout(dropout)
406 )
407
408 def forward(self, x):
409 return self.net(x)
410
411class ImageProjModel_g(torch.nn.Module):
412 """SD model with image prompt"""
413
414 def __init__(self, in_dim, hidden_dim, out_dim, dropout = 0.):
415 super().__init__()
416
417 self.net = nn.Sequential(
418 nn.Linear(in_dim, hidden_dim),
419 nn.GELU(),
420 nn.Dropout(dropout),
421 nn.LayerNorm(hidden_dim),
422 nn.Linear(hidden_dim, out_dim),
423 nn.Dropout(dropout)
424 )
425
426 def forward(self, x): # b, 257,1280
427 return self.net(x)
428
429
430class SDModel(torch.nn.Module):
431 """SD model with image prompt"""
432 def __init__(self, unet) -> None:
433 super().__init__()
434 self.image_proj_model_p = ImageProjModel_p(in_dim=1536, hidden_dim=768, out_dim=1024)
435
436 self.unet = unet
437 self.pose_proj = ControlNetConditioningEmbedding(
438 conditioning_embedding_channels=320,
439 block_out_channels=(16, 32, 96, 256),
440 conditioning_channels=3)
441
442
443 def forward(self, noisy_latents, timesteps, simg_f_p, timg_f_g, pose_f):
444
445 extra_image_embeddings_p = self.image_proj_model_p(simg_f_p)
446 extra_image_embeddings_g = timg_f_g
447
448 print(extra_image_embeddings_p.size())
449 print(extra_image_embeddings_g.size())
450
451 encoder_image_hidden_states = torch.cat([extra_image_embeddings_p ,extra_image_embeddings_g], dim=1)
452 pose_cond = self.pose_proj(pose_f)
453
454 pred_noise = self.unet(noisy_latents, timesteps, class_labels=timg_f_g, encoder_hidden_states=encoder_image_hidden_states,my_pose_cond=pose_cond).sample
455 return pred_noise
456
457def load_training_checkpoint(model, pcdms_model, tag=None, **kwargs):
458 #model_sd = torch.load(load_dir, map_location="cpu")["module"]
459 model_sd = torch.load(
460 pcdms_model,
461 map_location="cpu"
462 )["module"]
463
464
465 image_proj_model_dict = {}
466 pose_proj_dict = {}
467 unet_dict = {}
468 for k in model_sd.keys():
469 if k.startswith("pose_proj"):
470 pose_proj_dict[k.replace("pose_proj.", "")] = model_sd[k]
471
472 elif k.startswith("image_proj_model_p"):
473 image_proj_model_dict[k.replace("image_proj_model_p.", "")] = model_sd[k]
474
475 elif k.startswith("image_proj_model."):
476 image_proj_model_dict[k.replace("image_proj_model.", "")] = model_sd[k]
477
478
479 elif k.startswith("unet"):
480 unet_dict[k.replace("unet.", "")] = model_sd[k]
481 else:
482 print(k)
483
484 model.pose_proj.load_state_dict(pose_proj_dict)
485 model.image_proj_model_p.load_state_dict(image_proj_model_dict)
486 model.unet.load_state_dict(unet_dict)
487
488 return model, 0, 0
489
490
491def checkpoint_model(checkpoint_folder, ckpt_id, model, epoch, last_global_step, **kwargs):
492 """Utility function for checkpointing model + optimizer dictionaries
493 The main purpose for this is to be able to resume training from that instant again
494 """
495 checkpoint_state_dict = {
496 "epoch": epoch,
497 "last_global_step": last_global_step,
498 }
499 # Add extra kwargs too
500 checkpoint_state_dict.update(kwargs)
501
502 success = model.save_checkpoint(checkpoint_folder, ckpt_id, checkpoint_state_dict)
503 status_msg = f"checkpointing: checkpoint_folder={checkpoint_folder}, ckpt_id={ckpt_id}"
504 if success:
505 logging.info(f"Success {status_msg}")
506 else:
507 logging.warning(f"Failure {status_msg}")
508 return
509
510
511@spaces.GPU(duration=600)
512def train(modelId, in_image, in_pose, train_images, train_poses, train_steps, pcdms_model, noise_scheduler, image_encoder_p, image_encoder_g, vae, unet, finetune=True, is_app=False):
513 logging_dir = 'outputs/logging'
514 print('start train')
515
516 progress=gr.Progress(track_tqdm=True)
517
518 accelerator = Accelerator(
519 log_with=args.report_to,
520 project_dir=logging_dir,
521 mixed_precision=args.mixed_precision,
522 gradient_accumulation_steps=args.gradient_accumulation_steps
523 )
524
525 # Make one log on every process with the configuration for debugging.
526 #logging.basicConfig(
527 # format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
528 # datefmt="%m/%d/%Y %H:%M:%S",
529 # level=logging.INFO, )
530
531 print(accelerator.state)
532 if accelerator.is_local_main_process:
533 transformers.utils.logging.set_verbosity_warning()
534 diffusers.utils.logging.set_verbosity_info()
535 else:
536 transformers.utils.logging.set_verbosity_error()
537 diffusers.utils.logging.set_verbosity_error()
538
539 # If passed along, set the training seed now.
540 set_seed(42)
541
542 # Handle the repository creation
543 if accelerator.is_main_process:
544 os.makedirs('outputs', exist_ok=True)
545
546
547 """
548 unet = Stage2_InapintUNet2DConditionModel.from_pretrained("stabilityai/stable-diffusion-2-1-base", subfolder="unet",
549 in_channels=9, class_embed_type="projection" ,projection_class_embeddings_input_dim=1024,
550 low_cpu_mem_usage=False, ignore_mismatched_sizes=True)
551 """
552 image_encoder_p.requires_grad_(False)
553 image_encoder_g.requires_grad_(False)
554 vae.requires_grad_(False)
555
556 sd_model = SDModel(unet=unet)
557 sd_model.train()
558
559
560 if args.gradient_checkpointing:
561 sd_model.enable_gradient_checkpointing()
562
563
564 # Enable TF32 for faster training on Ampere GPUs,
565 # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices
566 if args.allow_tf32:
567 torch.backends.cuda.matmul.allow_tf32 = True
568
569 learning_rate = 1e-4
570 train_batch_size = min(len(train_images), max_batch_size) #len(train_images) % 16
571
572
573 # Optimizer creation
574 params_to_optimize = sd_model.parameters()
575 optimizer = torch.optim.AdamW(
576 params_to_optimize,
577 lr=learning_rate,
578 betas=(args.adam_beta1, args.adam_beta2),
579 weight_decay=args.adam_weight_decay,
580 eps=args.adam_epsilon,
581 )
582
583 inputs = [{
584 "source_image": in_image,
585 "source_pose": in_pose,
586 "target_image": timg,
587 "target_pose": tpose,
588 } for timg, tpose in zip(train_images, train_poses)]
589
590 """
591 inputs = {[
592 "source_image": Image.open('imgs/sm.png'),
593 "source_pose": Image.open('imgs/sm_pose.jpg'),
594 "target_image": Image.open('imgs/target.png'),
595 "target_pose": Image.open('imgs/target_pose.jpg'),
596 ]}
597 """
598
599 #print(inputs)
600
601 dataset = InpaintDataset(
602 inputs,
603 'imgs/',
604 size=(args.img_width, args.img_height), # w h
605 imgp_drop_rate=0.1,
606 imgg_drop_rate=0.1,
607 )
608
609 """
610 dataset = InpaintDataset(
611 args.json_path,
612 args.image_root_path,
613 size=(args.img_width, args.img_height), # w h
614 imgp_drop_rate=0.1,
615 imgg_drop_rate=0.1,
616 )
617 """
618
619 train_sampler = torch.utils.data.distributed.DistributedSampler(
620 dataset, num_replicas=accelerator.num_processes, rank=accelerator.process_index, shuffle=True)
621
622 train_dataloader = torch.utils.data.DataLoader(
623 dataset,
624 sampler=train_sampler,
625 collate_fn=InpaintCollate_fn,
626 batch_size=train_batch_size,
627 num_workers=0,)
628
629
630 # Scheduler and math around the number of training steps.
631 overrode_max_train_steps = False
632 num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
633 if args.max_train_steps is None:
634 args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
635 overrode_max_train_steps = True
636 args.max_train_steps = train_steps
637
638 lr_scheduler = get_scheduler(
639 args.lr_scheduler,
640 optimizer=optimizer,
641 num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes,
642 num_training_steps=args.max_train_steps * accelerator.num_processes,
643 num_cycles=args.lr_num_cycles,
644 power=args.lr_power,
645 )
646
647 # Prepare everything with our `accelerator`.
648 sd_model, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(sd_model, optimizer, train_dataloader, lr_scheduler)
649
650 # For mixed precision training we cast the text_encoder and vae weights to half-precision
651 # as these models are only used for inference, keeping weights in full precision is not required.
652 weight_dtype = torch.float32
653 """
654 if accelerator.mixed_precision == "fp16":
655 weight_dtype = torch.float16
656 elif accelerator.mixed_precision == "bf16":
657 weight_dtype = torch.bfloat16
658 """
659
660 # Move vae, unet and text_encoder to device and cast to weight_dtype
661 vae.to(accelerator.device, dtype=weight_dtype)
662 sd_model.unet.to(accelerator.device, dtype=weight_dtype)
663 image_encoder_p.to(accelerator.device, dtype=weight_dtype)
664 image_encoder_g.to(accelerator.device, dtype=weight_dtype)
665
666 # We need to recalculate our total training steps as the size of the training dataloader may have changed.
667 num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)
668 if overrode_max_train_steps:
669 args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch
670 # Afterwards we recalculate our number of training epochs
671 args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
672
673
674 args.num_train_epochs = train_steps
675
676
677 # Train!
678 total_batch_size = (
679 train_batch_size
680 * accelerator.num_processes
681 * args.gradient_accumulation_steps
682 )
683
684 print("***** Running training *****")
685 print(f" Num batches each epoch = {len(train_dataloader)}")
686 print(f" Num Epochs = {args.num_train_epochs}")
687 print(f" Instantaneous batch size per device = {train_batch_size}")
688 print(
689 f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}"
690 )
691 print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
692 print(f" Total optimization steps = {args.max_train_steps}")
693
694
695 if args.resume_from_checkpoint:
696 # New Code #
697 # Loads the DeepSpeed checkpoint from the specified path
698 prior_model, last_epoch, last_global_step = load_training_checkpoint(
699 sd_model,
700 pcdms_model,
701 **{"load_optimizer_states": True, "load_lr_scheduler_states": True},
702 )
703 print(f"Resumed from checkpoint: {args.resume_from_checkpoint}, global step: {last_global_step}")
704 starting_epoch = last_epoch
705 global_steps = last_global_step
706 sd_model = sd_model
707 else:
708 global_steps = 0
709 starting_epoch = 0
710 sd_model = sd_model
711
712 progress_bar = tqdm(range(global_steps, args.max_train_steps), initial=global_steps, desc="Steps",
713 # Only show the progress bar once on each machine.
714 disable=not accelerator.is_local_main_process, )
715
716 bsz = train_batch_size
717
718 if not finetune or train_steps == 0:
719 accelerator.wait_for_everyone()
720 accelerator.end_training()
721
722 checkpoint_state_dict = {
723 "epoch": 0,
724 "module": {k: v.cpu() for k, v in sd_model.state_dict().items()}, #sd_model.state_dict(),
725 }
726 torch.save(checkpoint_state_dict, modelId+".pt")
727
728 del sd_model
729 gc.collect()
730 torch.cuda.empty_cache()
731 return
732 #return {k: v.cpu() for k, v in sd_model.state_dict().items()}
733
734
735 it = range(starting_epoch, args.num_train_epochs)
736 if is_app:
737 it = progress.tqdm(it, desc="Fine-tuning")
738 for epoch in it:
739 for step, batch in enumerate(train_dataloader):
740 with accelerator.accumulate(sd_model):
741 with torch.no_grad():
742 # Convert images to latent space
743 latents = vae.encode(batch["source_target_image"].to(dtype=weight_dtype)).latent_dist.sample()
744 latents = latents * vae.config.scaling_factor
745
746 # Get the masked image latents
747 masked_latents = vae.encode(batch["vae_source_mask_image"].to(dtype=weight_dtype)).latent_dist.sample()
748 masked_latents = masked_latents * vae.config.scaling_factor
749
750 bsz = batch["target_image"].size(dim=0)
751
752 # mask
753 mask1 = torch.ones((bsz, 1, int(args.img_height / 8), int(args.img_width / 8))).to(accelerator.device, dtype=weight_dtype)
754 mask0 = torch.zeros((bsz, 1, int(args.img_height / 8), int(args.img_width / 8))).to(accelerator.device, dtype=weight_dtype)
755 mask = torch.cat([mask1, mask0], dim=3)
756 # Get the image embedding for conditioning
757 cond_image_feature_p = image_encoder_p(batch["source_image"].to(accelerator.device, dtype=weight_dtype))
758 cond_image_feature_p = (cond_image_feature_p.last_hidden_state)
759
760
761 cond_image_feature_g = image_encoder_g(batch["target_image"].to(accelerator.device, dtype=weight_dtype), ).image_embeds
762 cond_image_feature_g =cond_image_feature_g.unsqueeze(1)
763
764 # Sample noise that we'll add to the latents
765 noise = torch.randn_like(latents)
766 if args.noise_offset:
767 # https://www.crosslabs.org//blog/diffusion-with-offset-noise
768 noise += args.noise_offset * torch.randn(
769 (latents.shape[0], latents.shape[1], 1, 1), device=latents.device
770 )
771
772 # Sample a random timestep for each image
773 #timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (train_batch_size,),device=latents.device, )
774 timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (bsz,),device=latents.device, )
775 timesteps = timesteps.long()
776
777
778
779 # Add noise to the latents according to the noise magnitude at each timestep (this is the forward diffusion process)
780 noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps)
781
782 #print(noisy_latents.size(), mask.size(), masked_latents.size())
783
784 noisy_latents = torch.cat([noisy_latents, mask, masked_latents], dim=1)
785 # Get the text embedding for conditioning
786
787
788 cond_pose = batch["source_target_pose"].to(dtype=weight_dtype)
789
790 #print(noisy_latents.size())
791 #print(cond_image_feature_p.size())
792 #print(cond_image_feature_g.size())
793 #print(cond_pose.size())
794
795 # Predict the noise residual
796 model_pred = sd_model(noisy_latents, timesteps, cond_image_feature_p,cond_image_feature_g, cond_pose, )
797
798 # Get the target for loss depending on the prediction type
799 if noise_scheduler.config.prediction_type == "epsilon":
800 target = noise
801 elif noise_scheduler.config.prediction_type == "v_prediction":
802 target = noise_scheduler.get_velocity(latents, noise, timesteps)
803 else:
804 raise ValueError(
805 f"Unknown prediction type {noise_scheduler.config.prediction_type}"
806 )
807
808 loss = F.mse_loss(model_pred.float(), target.float(), reduction="mean")
809
810 accelerator.backward(loss)
811 if accelerator.sync_gradients:
812 params_to_clip = sd_model.parameters()
813 accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm)
814 optimizer.step()
815 lr_scheduler.step()
816 optimizer.zero_grad(set_to_none=args.set_grads_to_none)
817
818 # Checks if the accelerator has performed an optimization step behind the scenes
819 if accelerator.sync_gradients:
820 global_steps += 1
821
822 if global_steps >= args.max_train_steps:
823 break
824
825
826 logs = {"loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]}
827 print(logs)
828 progress_bar.set_postfix(**logs)
829
830 progress_bar.update(1)
831
832 # Create the pipeline using the trained modules and save it.
833 accelerator.wait_for_everyone()
834 accelerator.end_training()
835
836 sd_model.unet.cpu()
837 sd_model.cpu()
838 del vae
839 del image_encoder_p
840 del image_encoder_g
841
842 if save_model: #if global_steps % args.checkpointing_steps == 0 or global_steps == args.max_train_steps:
843 print('saving', modelId)
844
845 checkpoint_state_dict = {
846 "epoch": 0,
847 "module": {k: v.cpu() for k, v in sd_model.state_dict().items()}, #sd_model.state_dict(),
848 }
849 print(list(sd_model.state_dict().keys())[:20])
850 torch.save(checkpoint_state_dict, modelId+".pt")
851
852 del sd_model
853 gc.collect()
854 torch.cuda.empty_cache()
855 print('done train')
856 print(torch.cuda.memory_allocated()/1024**2)
857 return
858
859 del sd_model
860 gc.collect()
861 torch.cuda.empty_cache()
862 return {k: v.cpu() for k, v in sd_model.state_dict().items()}
863
864
865
866
867# Pose-transfer ===================================================================================================
868
869
870device = "cuda"
871
872class ImageProjModel(torch.nn.Module):
873 """SD model with image prompt"""
874 def __init__(self, in_dim, hidden_dim, out_dim, dropout = 0.):
875 super().__init__()
876
877 self.net = nn.Sequential(
878 nn.Linear(in_dim, hidden_dim),
879 nn.GELU(),
880 nn.Dropout(dropout),
881 nn.LayerNorm(hidden_dim),
882 nn.Linear(hidden_dim, out_dim),
883 nn.Dropout(dropout)
884 )
885
886 def forward(self, x):
887 return self.net(x)
888
889def image_grid(imgs, rows, cols):
890 assert len(imgs) == rows * cols
891 w, h = imgs[0].size
892 print(w, h)
893 grid = Image.new("RGB", size=(cols * w, rows * h))
894 grid_w, grid_h = grid.size
895
896 for i, img in enumerate(imgs):
897 grid.paste(img, box=(i % cols * w, i // cols * h))
898 return grid
899
900def load_mydict(modelId, finetuned_model):
901 if save_model:
902 model_ckpt_path = modelId+'.pt'
903 model_sd = torch.load(model_ckpt_path, map_location="cpu")["module"]
904 else:
905 model_sd = finetuned_model #torch.load(model_ckpt_path, map_location="cpu")["module"]
906
907 image_proj_model_dict = {}
908 pose_proj_dict = {}
909 unet_dict = {}
910 for k in model_sd.keys():
911 if k.startswith("pose_proj"):
912 pose_proj_dict[k.replace("pose_proj.", "")] = model_sd[k]
913
914 elif k.startswith("image_proj_model_p"):
915 image_proj_model_dict[k.replace("image_proj_model_p.", "")] = model_sd[k]
916 elif k.startswith("image_proj_model"):
917 image_proj_model_dict[k.replace("image_proj_model.", "")] = model_sd[k]
918
919
920 elif k.startswith("unet"):
921 unet_dict[k.replace("unet.", "")] = model_sd[k]
922 else:
923 print(k)
924 return image_proj_model_dict, pose_proj_dict, unet_dict
925
926
927
928@spaces.GPU(duration=600)
929def inference(modelId, in_image, in_pose, target_poses, inference_steps, finetuned_model, vae, unet, image_encoder, is_app=False):
930 print('start inference')
931 progress=gr.Progress(track_tqdm=True)
932
933 if not save_model:
934 finetuned_model = {k: v.cuda() for k, v in finetuned_model.items()}
935
936 device = "cuda"
937 pretrained_model_name_or_path ="stabilityai/stable-diffusion-2-1-base"
938 image_encoder_path = "facebook/dinov2-giant"
939 #model_ckpt_path = "./pcdms_ckpt.pt" # ckpt path
940 model_ckpt_path = modelId+'.pt'
941
942
943 clip_image_processor = CLIPImageProcessor()
944 img_transform = transforms.Compose([
945 transforms.ToTensor(),
946 transforms.Normalize([0.5], [0.5]),
947 ])
948
949 generator = torch.Generator(device=device).manual_seed(42)
950
951 """
952 unet = Stage2_InapintUNet2DConditionModel.from_pretrained(pretrained_model_name_or_path, torch_dtype=torch.float16,subfolder="unet",in_channels=9, low_cpu_mem_usage=False, ignore_mismatched_sizes=True).to(device)
953 vae = AutoencoderKL.from_pretrained(pretrained_model_name_or_path,subfolder="vae").to(device, dtype=torch.float16)
954 image_encoder = Dinov2Model.from_pretrained(image_encoder_path).to(device, dtype=torch.float16)
955 """
956 noise_scheduler = DDIMScheduler(
957 num_train_timesteps=1000,
958 beta_start=0.00085,
959 beta_end=0.012,
960 beta_schedule="scaled_linear",
961 clip_sample=False,
962 set_alpha_to_one=False,
963 steps_offset=1,
964 )
965
966 unet = unet.to(device, dtype=torch.float16)
967 vae = vae.to(device, dtype=torch.float16)
968 image_encoder = image_encoder.to(device, dtype=torch.float16)
969
970
971 image_proj_model = ImageProjModel(in_dim=1536, hidden_dim=768, out_dim=1024).to(device).to(dtype=torch.float16)
972 pose_proj_model = ControlNetConditioningEmbedding(
973 conditioning_embedding_channels=320,
974 block_out_channels=(16, 32, 96, 256),
975 conditioning_channels=3).to(device).to(dtype=torch.float16)
976
977
978 # load weight
979 print('loading', modelId)
980 image_proj_model_dict, pose_proj_dict, unet_dict = load_mydict(modelId, finetuned_model)
981 print('loaded', modelId)
982 image_proj_model.load_state_dict(image_proj_model_dict)
983 pose_proj_model.load_state_dict(pose_proj_dict)
984 unet.load_state_dict(unet_dict)
985
986
987 pipe = PCDMsPipeline.from_pretrained("stabilityai/stable-diffusion-2-1-base", unet=unet, torch_dtype=torch.float16, scheduler=noise_scheduler,feature_extractor=None,safety_checker=None).to(device)
988
989 print('====================== model load finish ===================')
990
991 results = []
992 progress_bar = tqdm(range(len(target_poses)), initial=0, desc="Frames")
993
994
995 it = target_poses
996 if is_app:
997 it = progress.tqdm(it, desc="Pose Transfer")
998 for pose in it:
999
1000 num_samples = 1
1001 image_size = (512, 512)
1002 s_img_path = 'imgs/'+input_img # input image 1
1003 #target_pose_img = 'imgs/pose_'+str(n)+'.png' # input image 2
1004
1005 #t_pose = inference_pose(target_pose_img, image_size=(image_size[1], image_size[0])).resize(image_size, Image.BICUBIC)
1006 #t_pose = Image.open(target_pose_img).convert("RGB").resize((image_size), Image.BICUBIC)
1007 t_pose = pose.convert("RGB").resize((image_size), Image.BICUBIC)
1008 #t_pose = resize_and_pad(pose.convert("RGB"))
1009
1010
1011 #s_img = Image.open(s_img_path)
1012 width_orig, height_orig = in_image.size
1013 s_img = in_image.convert("RGB").resize(image_size, Image.BICUBIC)
1014 #s_img = resize_and_pad(in_image.convert("RGB"))
1015 black_image = Image.new("RGB", s_img.size, (0, 0, 0)).resize(image_size, Image.BICUBIC)
1016
1017 s_img_t_mask = Image.new("RGB", (s_img.width * 2, s_img.height))
1018 s_img_t_mask.paste(s_img, (0, 0))
1019 s_img_t_mask.paste(black_image, (s_img.width, 0))
1020
1021 #s_pose = inference_pose(s_img_path, image_size=(image_size[1], image_size[0])).resize(image_size, Image.BICUBIC)
1022 #s_pose = Image.open('imgs/sm_pose.jpg').convert("RGB").resize(image_size, Image.BICUBIC)
1023 s_pose = in_pose.convert("RGB").resize(image_size, Image.BICUBIC)
1024 #s_pose = resize_and_pad(in_pose.convert("RGB"))
1025 print('source image width: {}, height: {}'.format(s_pose.width, s_pose.height))
1026 #t_pose = Image.open(target_pose_img).convert("RGB").resize((image_size), Image.BICUBIC)
1027
1028 st_pose = Image.new("RGB", (s_pose.width * 2, s_pose.height))
1029 st_pose.paste(s_pose, (0, 0))
1030 st_pose.paste(t_pose, (s_pose.width, 0))
1031
1032
1033 clip_s_img = clip_image_processor(images=s_img, return_tensors="pt").pixel_values
1034 vae_image = torch.unsqueeze(img_transform(s_img_t_mask), 0)
1035 cond_st_pose = torch.unsqueeze(img_transform(st_pose), 0)
1036
1037 mask1 = torch.ones((1, 1, int(image_size[0] / 8), int(image_size[1] / 8))).to(device, dtype=torch.float16)
1038 mask0 = torch.zeros((1, 1, int(image_size[0] / 8), int(image_size[1] / 8))).to(device, dtype=torch.float16)
1039 mask = torch.cat([mask1, mask0], dim=3)
1040
1041
1042 with torch.inference_mode():
1043 cond_pose = pose_proj_model(cond_st_pose.to(dtype=torch.float16, device=device))
1044 simg_mask_latents = pipe.vae.encode(vae_image.to(device, dtype=torch.float16)).latent_dist.sample()
1045 simg_mask_latents = simg_mask_latents * 0.18215
1046
1047 images_embeds = image_encoder(clip_s_img.to(device, dtype=torch.float16)).last_hidden_state
1048 image_prompt_embeds = image_proj_model(images_embeds)
1049 uncond_image_prompt_embeds = image_proj_model(torch.zeros_like(images_embeds))
1050
1051 bs_embed, seq_len, _ = image_prompt_embeds.shape
1052 image_prompt_embeds = image_prompt_embeds.repeat(1, num_samples, 1)
1053 image_prompt_embeds = image_prompt_embeds.view(bs_embed * num_samples, seq_len, -1)
1054 uncond_image_prompt_embeds = uncond_image_prompt_embeds.repeat(1, num_samples, 1)
1055 uncond_image_prompt_embeds = uncond_image_prompt_embeds.view(bs_embed * num_samples, seq_len, -1)
1056
1057 output, _ = pipe(
1058 simg_mask_latents= simg_mask_latents,
1059 mask = mask,
1060 cond_pose = cond_pose,
1061 prompt_embeds=image_prompt_embeds,
1062 negative_prompt_embeds=uncond_image_prompt_embeds,
1063 height=image_size[1],
1064 width=image_size[0]*2,
1065 num_images_per_prompt=num_samples,
1066 guidance_scale=2.0,
1067 generator=generator,
1068 num_inference_steps=inference_steps,
1069 )
1070
1071 output = output.images[-1]
1072
1073 result = output.crop((image_size[0], 0, image_size[0] * 2, image_size[1]))
1074 result = result.resize((width_orig, height_orig), Image.BICUBIC)
1075 #result = remove_zero_pad(result)
1076
1077 if debug:
1078 result.save('out/'+str(len(results))+'.png')
1079 results.append(result)
1080 progress_bar.update(1)
1081
1082 del unet
1083 del vae
1084 del image_encoder
1085 del image_proj_model
1086 del pose_proj_model
1087
1088 if not save_model:
1089 del finetuned_model
1090
1091 gc.collect()
1092 torch.cuda.empty_cache()
1093 print(torch.cuda.memory_allocated()/1024**2)
1094
1095 return results
1096
1097
1098def gen_vid(frames, video_name, fps, codec):
1099 progress=gr.Progress(track_tqdm=True)
1100
1101 frame = cv2.cvtColor(np.array(frames[0]), cv2.COLOR_RGB2BGR)
1102 height, width, layers = frame.shape
1103
1104 #video = cv2.VideoWriter(video_name, 0, 1, (width,height))
1105 if codec == 'mp4':
1106 video = cv2.VideoWriter(video_name, cv2.VideoWriter_fourcc(*'mp4v'), fps, (width, height))
1107 else:
1108 video = cv2.VideoWriter(video_name, cv2.VideoWriter_fourcc(*'VP90'), fps, (width, height))
1109
1110 for r in progress.tqdm(frames, desc="Creating video"):
1111 image = cv2.cvtColor(np.array(r), cv2.COLOR_RGB2BGR)
1112 video.write(image)
1113
1114 #cv2.destroyAllWindows()
1115 #video.release()
1116
1117
1118
1119def run(images, video_path, train_steps=100, inference_steps=10, fps=12, bg_remove=False, resize_inputs=True, finetune=True, is_app=False):
1120 print("==== Load Models ====")
1121 dwpose, rembg_session, pcdms_model, noise_scheduler, image_encoder_p, image_encoder_g, vae, unet = load_models()
1122
1123 print("==== Pose Detection ====")
1124 in_img, in_pose, train_imgs, train_poses, target_poses = prepare_inputs(images, video_path, fps, bg_remove, dwpose, rembg_session, resize_inputs, is_app=is_app)
1125
1126 if save_model:
1127 train("fine_tuned_pcdms", in_img, in_pose, train_imgs, train_poses, train_steps, pcdms_model, noise_scheduler, image_encoder_p, image_encoder_g, vae, unet, finetune, is_app)
1128 print('next')
1129 results = inference("fine_tuned_pcdms", in_img, in_pose, target_poses, inference_steps, None, vae, unet, image_encoder_p, is_app)
1130
1131 else:
1132 print("==== Finetuning ====")
1133 finetuned_model = train("fine_tuned_pcdms", in_img, in_pose, train_imgs, train_poses, train_steps, pcdms_model, noise_scheduler, image_encoder_p, image_encoder_g, vae, unet, finetune, is_app)
1134
1135 print("==== Pose Transfer ====")
1136 results = inference("fine_tuned_pcdms", in_img, in_pose, target_poses, inference_steps, finetuned_model, vae, unet, image_encoder_p, is_app)
1137
1138 return results
1139
1140
1141def run_train_impl(images, train_steps=100, modelId="fine_tuned_pcdms", bg_remove=True, resize_inputs=True, finetune=True):
1142 finetune=True
1143 is_app=True
1144 images = [img[0] for img in images]
1145
1146 dwpose, rembg_session, pcdms_model, noise_scheduler, image_encoder_p, image_encoder_g, vae, unet = load_models()
1147
1148 if resize_inputs:
1149 resize = 'target'
1150 else:
1151 resize = 'none'
1152
1153 in_img, in_pose, train_imgs, train_poses = prepare_inputs_train(images, bg_remove, dwpose, rembg_session)
1154
1155 train(modelId, in_img, in_pose, train_imgs, train_poses, train_steps, pcdms_model, noise_scheduler, image_encoder_p, image_encoder_g, vae, unet, finetune, is_app)
1156
1157 gc.collect()
1158 torch.cuda.empty_cache()
1159
1160def run_train(images, train_steps=100, modelId="fine_tuned_pcdms", bg_remove=True, resize_inputs=True):
1161 run_train_impl(images, train_steps, modelId, bg_remove, resize_inputs)
1162
1163 """
1164 mp.set_start_method('spawn', force=True)
1165 p = mp.Process(target=run_train_impl, args=(images, train_steps, modelId, bg_remove, resize_inputs))
1166 p.start()
1167 p.join()
1168 """
1169
1170
1171def run_inference_impl(images, video_path, frames, train_steps=100, inference_steps=10, fps=12, modelId="fine_tuned_pcdms", img_width=1920, img_height=1080, bg_remove=True, resize_inputs=True):
1172 finetune=True
1173 is_app=True
1174
1175
1176 dwpose, rembg_session, pcdms_model, noise_scheduler, image_encoder_p, image_encoder_g, vae, unet = load_models()
1177
1178 if not os.path.exists(modelId+".pt"):
1179 run_train(images, train_steps, modelId, bg_remove, resize_inputs)
1180
1181 images = [img[0] for img in images]
1182 in_img = images[0]
1183 if frames:
1184 frames = [img[0] for img in frames]
1185
1186 in_img, target_poses, in_pose, target_poses_coords, orig_frames = prepare_inputs_inference(in_img, video_path, frames, fps, dwpose, rembg_session, bg_remove, resize_inputs, is_app)
1187 #target_poses[0].save('inf_pose.png')
1188
1189 results = inference(modelId, in_img, in_pose, target_poses, inference_steps, None, vae, unet, image_encoder_p, is_app)
1190 #urls = save_temp_imgs(results)
1191
1192 if should_gen_vid:
1193 if debug:
1194 gen_vid(results, out_vid+'.mp4', fps, 'mp4')
1195 else:
1196 gen_vid(results, out_vid+'.webm', fps, 'webm')
1197
1198
1199 # postprocessing
1200 if no_bg_final:
