Team Ai
Apppublic

ProgrammerParamesh/VirtualDress

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
inference_ootd_dc.py133 linesDownload Raw Back to ootd
1import pdb
2from pathlib import Path
3import sys
4PROJECT_ROOT = Path(__file__).absolute().parents[0].absolute()
5sys.path.insert(0, str(PROJECT_ROOT))
6import os
7import torch
8import numpy as np
9from PIL import Image
10import cv2
11
12import random
13import time
14import pdb
15
16from pipelines_ootd.pipeline_ootd import OotdPipeline
17from pipelines_ootd.unet_garm_2d_condition import UNetGarm2DConditionModel
18from pipelines_ootd.unet_vton_2d_condition import UNetVton2DConditionModel
19from diffusers import UniPCMultistepScheduler
20from diffusers import AutoencoderKL
21
22import torch.nn as nn
23import torch.nn.functional as F
24from transformers import AutoProcessor, CLIPVisionModelWithProjection
25from transformers import CLIPTextModel, CLIPTokenizer
26
27VIT_PATH = "../checkpoints/clip-vit-large-patch14"
28VAE_PATH = "../checkpoints/ootd"
29UNET_PATH = "../checkpoints/ootd/ootd_dc/checkpoint-36000"
30MODEL_PATH = "../checkpoints/ootd"
31
32class OOTDiffusionDC:
33
34    def __init__(self, gpu_id):
35        self.gpu_id = 'cuda:' + str(gpu_id)
36
37        vae = AutoencoderKL.from_pretrained(
38            VAE_PATH,
39            subfolder="vae",
40            torch_dtype=torch.float16,
41        )
42
43        unet_garm = UNetGarm2DConditionModel.from_pretrained(
44            UNET_PATH,
45            subfolder="unet_garm",
46            torch_dtype=torch.float16,
47            use_safetensors=True,
48        )
49        unet_vton = UNetVton2DConditionModel.from_pretrained(
50            UNET_PATH,
51            subfolder="unet_vton",
52            torch_dtype=torch.float16,
53            use_safetensors=True,
54        )
55
56        self.pipe = OotdPipeline.from_pretrained(
57            MODEL_PATH,
58            unet_garm=unet_garm,
59            unet_vton=unet_vton,
60            vae=vae,
61            torch_dtype=torch.float16,
62            variant="fp16",
63            use_safetensors=True,
64            safety_checker=None,
65            requires_safety_checker=False,
66        ).to(self.gpu_id)
67
68        self.pipe.scheduler = UniPCMultistepScheduler.from_config(self.pipe.scheduler.config)
69        
70        self.auto_processor = AutoProcessor.from_pretrained(VIT_PATH)
71        self.image_encoder = CLIPVisionModelWithProjection.from_pretrained(VIT_PATH).to(self.gpu_id)
72
73        self.tokenizer = CLIPTokenizer.from_pretrained(
74            MODEL_PATH,
75            subfolder="tokenizer",
76        )
77        self.text_encoder = CLIPTextModel.from_pretrained(
78            MODEL_PATH,
79            subfolder="text_encoder",
80        ).to(self.gpu_id)
81
82
83    def tokenize_captions(self, captions, max_length):
84        inputs = self.tokenizer(
85            captions, max_length=max_length, padding="max_length", truncation=True, return_tensors="pt"
86        )
87        return inputs.input_ids
88
89
90    def __call__(self,
91                model_type='hd',
92                category='upperbody',
93                image_garm=None,
94                image_vton=None,
95                mask=None,
96                image_ori=None,
97                num_samples=1,
98                num_steps=20,
99                image_scale=1.0,
100                seed=-1,
101    ):
102        if seed == -1:
103            random.seed(time.time())
104            seed = random.randint(0, 2147483647)
105        print('Initial seed: ' + str(seed))
106        generator = torch.manual_seed(seed)
107
108        with torch.no_grad():
109            prompt_image = self.auto_processor(images=image_garm, return_tensors="pt").to(self.gpu_id)
110            prompt_image = self.image_encoder(prompt_image.data['pixel_values']).image_embeds
111            prompt_image = prompt_image.unsqueeze(1)
112            if model_type == 'hd':
113                prompt_embeds = self.text_encoder(self.tokenize_captions([""], 2).to(self.gpu_id))[0]
114                prompt_embeds[:, 1:] = prompt_image[:]
115            elif model_type == 'dc':
116                prompt_embeds = self.text_encoder(self.tokenize_captions([category], 3).to(self.gpu_id))[0]
117                prompt_embeds = torch.cat([prompt_embeds, prompt_image], dim=1)
118            else:
119                raise ValueError("model_type must be \'hd\' or \'dc\'!")
120
121            images = self.pipe(prompt_embeds=prompt_embeds,
122                        image_garm=image_garm,
123                        image_vton=image_vton, 
124                        mask=mask,
125                        image_ori=image_ori,
126                        num_inference_steps=num_steps,
127                        image_guidance_scale=image_scale,
128                        num_images_per_prompt=num_samples,
129                        generator=generator,
130            ).images
131
132        return images
133