CoolFace
Apppublic

acmyu/KeyframesAI

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
main.py1370 linesDownload Raw Back to root
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:

Showing the first 1,200 of 1370 lines. Download the file for the rest.