CoolFace
Apppublic

acmyu/KeyframesAI

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
train.py267 linesDownload Raw Back to root
1import glob
2import os
3import torch
4from torch import nn, optim
5import torch.nn.functional as F
6import torchvision.transforms.functional as FF
7from PIL import Image
8import numpy as np
9from diffusers import UniPCMultistepScheduler
10from src.models.stage2_inpaint_unet_2d_condition import Stage2_InapintUNet2DConditionModel
11from accelerate import Accelerator
12
13from torchvision import transforms
14from diffusers.models.controlnet import ControlNetConditioningEmbedding
15from transformers import CLIPImageProcessor
16from transformers import Dinov2Model
17from diffusers import AutoencoderKL, DDPMScheduler, UNet2DConditionModel,ControlNetModel,DDIMScheduler
18from src.pipelines.PCDMs_pipeline import PCDMsPipeline
19from single_extract_pose import inference_pose
20
21
22device = "cuda"
23pretrained_model_name_or_path ="stabilityai/stable-diffusion-2-1-base"
24image_encoder_path = "facebook/dinov2-giant"
25model_ckpt_path = "./pcdms_ckpt.pt"   # ckpt path
26
27num_samples = 1
28image_size = (512, 512)
29s_img_path = 'imgs/sm.png' # input image 1
30target_pose_img = 'imgs/pose.png' # input image 2
31
32
33def image_grid(imgs, rows, cols):
34    assert len(imgs) == rows * cols
35    w, h = imgs[0].size
36    print(w, h)
37    grid = Image.new("RGB", size=(cols * w, rows * h))
38    grid_w, grid_h = grid.size
39
40    for i, img in enumerate(imgs):
41        grid.paste(img, box=(i % cols * w, i // cols * h))
42    return grid
43
44def load_mydict(model_ckpt_path):
45    model_sd = torch.load(model_ckpt_path, map_location="cpu")["module"]
46
47    image_proj_model_dict = {}
48    pose_proj_dict = {}
49    unet_dict = {}
50    for k in model_sd.keys():
51        if k.startswith("pose_proj"):
52            pose_proj_dict[k.replace("pose_proj.", "")] = model_sd[k]
53
54        elif k.startswith("image_proj_model"):
55            image_proj_model_dict[k.replace("image_proj_model.", "")] = model_sd[k]
56
57
58        elif k.startswith("unet"):
59            unet_dict[k.replace("unet.", "")] = model_sd[k]
60        else:
61            print(k)
62    return image_proj_model_dict, pose_proj_dict, unet_dict
63
64class ImageProjModel(torch.nn.Module):
65    """SD model with image prompt"""
66    def __init__(self, in_dim, hidden_dim, out_dim, dropout = 0.):
67        super().__init__()
68
69        self.net = nn.Sequential(
70            nn.Linear(in_dim, hidden_dim),
71            nn.GELU(),
72            nn.Dropout(dropout),
73            nn.LayerNorm(hidden_dim),
74            nn.Linear(hidden_dim, out_dim),
75            nn.Dropout(dropout)
76        )
77
78    def forward(self, x):  
79        return self.net(x)
80
81
82
83clip_image_processor = CLIPImageProcessor()
84img_transform = transforms.Compose([
85    transforms.ToTensor(),
86    transforms.Normalize([0.5], [0.5]),
87])
88
89generator = torch.Generator(device=device).manual_seed(42)
90unet = 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)
91vae = AutoencoderKL.from_pretrained(pretrained_model_name_or_path,subfolder="vae").to(device, dtype=torch.float16)
92image_encoder = Dinov2Model.from_pretrained(image_encoder_path).to(device, dtype=torch.float16)
93noise_scheduler = DDIMScheduler(
94    num_train_timesteps=1000,
95    beta_start=0.00085,
96    beta_end=0.012,
97    beta_schedule="scaled_linear",
98    clip_sample=False,
99    set_alpha_to_one=False,
100    steps_offset=1,
101)
102
103#noise_scheduler = DDPMScheduler.from_pretrained(pretrained_model_name_or_path, subfolder="scheduler")
104
105print('====================== model load finish ===================')
106
107
108
109class SDModel(torch.nn.Module):
110    """SD model with image prompt"""
111    def __init__(self, unet) -> None:
112        super().__init__()
113        self.unet = unet
114        
115        self.image_proj_model = ImageProjModel(in_dim=1536, hidden_dim=768, out_dim=1024).to(device).to(dtype=torch.float16)
116        self.pose_proj = ControlNetConditioningEmbedding(
117            conditioning_embedding_channels=320,
118            block_out_channels=(16, 32, 96, 256),
119            conditioning_channels=3).to(device).to(dtype=torch.float16)
120        
121        # load weight
122        image_proj_model_dict, pose_proj_dict, unet_dict = load_mydict(model_ckpt_path)
123        self.image_proj_model.load_state_dict(image_proj_model_dict)
124        self.pose_proj.load_state_dict(pose_proj_dict)
125        unet.load_state_dict(unet_dict)
126
127
128    def forward(self, s_img_path, t_pose_path, t_img_path, epoch):
129
130        pipe = PCDMsPipeline.from_pretrained(pretrained_model_name_or_path, unet=self.unet,  torch_dtype=torch.float16, scheduler=noise_scheduler,feature_extractor=None,safety_checker=None).to(device)
131        
132        t_pose = inference_pose(t_img_path, image_size=(image_size[1], image_size[0])).convert("RGB").resize(image_size, Image.BICUBIC)
133        target_img = Image.open(t_img_path).convert("RGB").resize(image_size, Image.BICUBIC)
134        
135        
136        s_img = Image.open(s_img_path).convert("RGB").resize(image_size, Image.BICUBIC)
137        black_image = Image.new("RGB", s_img.size, (0, 0, 0)).resize(image_size, Image.BICUBIC)
138
139        s_img_t_mask = Image.new("RGB", (s_img.width * 2, s_img.height))
140        s_img_t_mask.paste(s_img, (0, 0))
141        s_img_t_mask.paste(black_image, (s_img.width, 0))
142
143        s_pose = inference_pose(s_img_path, image_size=(image_size[1], image_size[0])).resize(image_size, Image.BICUBIC)
144        print('source image width: {}, height: {}'.format(s_pose.width, s_pose.height))
145        #t_pose = Image.open(t_pose_path).convert("RGB").resize((image_size), Image.BICUBIC)
146
147        st_pose = Image.new("RGB", (s_pose.width * 2, s_pose.height))
148        st_pose.paste(s_pose, (0, 0))
149        st_pose.paste(t_pose, (s_pose.width, 0))
150
151
152        clip_s_img = clip_image_processor(images=s_img, return_tensors="pt").pixel_values
153        vae_image = torch.unsqueeze(img_transform(s_img_t_mask), 0)
154        cond_st_pose = torch.unsqueeze(img_transform(st_pose), 0)
155
156        mask1 = torch.ones((1, 1, int(image_size[0] / 8), int(image_size[1] / 8))).to(device, dtype=torch.float16)
157        mask0 = torch.zeros((1, 1, int(image_size[0] / 8), int(image_size[1] / 8))).to(device, dtype=torch.float16)
158        mask = torch.cat([mask1, mask0], dim=3)
159
160        st_img = (Image.new("RGB", (image_size[0] * 2, image_size[1])))
161        st_img.paste(s_img, (0, 0))
162        st_img.paste(target_img, (image_size[0], 0))
163        st_img.save('tar.png')
164        st_img = torch.unsqueeze(img_transform(st_img), 0)
165        
166        
167
168        with torch.inference_mode():
169            cond_pose = self.pose_proj(cond_st_pose.to(dtype=torch.float16, device=device))
170            simg_mask_latents = pipe.vae.encode(vae_image.to(device, dtype=torch.float16)).latent_dist.sample()
171            simg_mask_latents = simg_mask_latents * 0.18215
172
173            images_embeds = image_encoder(clip_s_img.to(device, dtype=torch.float16)).last_hidden_state
174            image_prompt_embeds = self.image_proj_model(images_embeds)
175            uncond_image_prompt_embeds = self.image_proj_model(torch.zeros_like(images_embeds))
176            
177            latents = pipe.vae.encode(st_img.to(device, dtype=torch.float16)).latent_dist.sample()
178            latents = latents * pipe.vae.config.scaling_factor
179            noise = torch.randn_like(latents)
180            timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (1,),device=latents.device, )
181            timesteps = timesteps.long()
182            target = noise_scheduler.get_velocity(latents, noise, timesteps)
183
184        bs_embed, seq_len, _ = image_prompt_embeds.shape
185        image_prompt_embeds = image_prompt_embeds.repeat(1, num_samples, 1)
186        image_prompt_embeds = image_prompt_embeds.view(bs_embed * num_samples, seq_len, -1)
187        uncond_image_prompt_embeds = uncond_image_prompt_embeds.repeat(1, num_samples, 1)
188        uncond_image_prompt_embeds = uncond_image_prompt_embeds.view(bs_embed * num_samples, seq_len, -1)
189        
190
191        output, model_pred = pipe(
192            simg_mask_latents= simg_mask_latents,
193            mask = mask,
194            cond_pose = cond_pose,
195            prompt_embeds=image_prompt_embeds,
196            negative_prompt_embeds=uncond_image_prompt_embeds,
197            height=image_size[1],
198            width=image_size[0]*2,
199            num_images_per_prompt=num_samples,
200            guidance_scale=2.0,
201            generator=generator,
202            num_inference_steps=50,
203        )
204        output = output.images[-1]
205        output.save('out'+str(epoch)+'.png')
206        
207        """
208        with torch.inference_mode():
209            output = torch.unsqueeze(img_transform(output), 0)
210            latents = pipe.vae.encode(output.to(device, dtype=torch.float16)).latent_dist.sample()
211            latents = latents * pipe.vae.config.scaling_factor
212            noise = torch.randn_like(latents)
213            timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (1,),device=latents.device, )
214            timesteps = timesteps.long()
215            model_pred = noise_scheduler.get_velocity(latents, noise, timesteps)
216        """
217        
218        return model_pred, target
219
220
221
222# Training setup
223sd_model = SDModel(unet)
224sd_model.train()
225optimizer = optim.AdamW(sd_model.parameters(), lr=1e-5)
226loss_fn = nn.MSELoss()
227
228
229accelerator = Accelerator()
230sd_model, optimizer = accelerator.prepare(sd_model, optimizer)
231
232
233prev = sd_model.unet.state_dict()
234
235# Fine-tuning loop
236num_epochs = 5
237for epoch in range(num_epochs):
238    for s_img_path, t_pose_path, t_img_path in zip(['imgs/sm.png'], ['imgs/p1.png'], ['imgs/target.png']):
239        with accelerator.accumulate(sd_model):
240            optimizer.zero_grad()
241            
242            model_pred, target = sd_model(s_img_path, t_pose_path, t_img_path, epoch)
243            
244            #loss = loss_fn(torch.unsqueeze(img_transform(output), 0), torch.unsqueeze(img_transform(target_img),0))
245            loss = F.mse_loss(model_pred.float(), target.float(), reduction="mean")
246            loss.requires_grad = True
247            
248            accelerator.backward(loss)
249            optimizer.step()
250            
251            set1 = set(prev.items())
252            set2 = set(sd_model.unet.state_dict().items())
253            dif = set1 ^ set2
254            print(len(dif))
255            prev = sd_model.unet.state_dict()
256    
257    print(f"Epoch [{epoch+1}/{num_epochs}], Loss: {loss.item()}")
258        
259
260
261
262# Save fine-tuned model
263torch.save(sd_model, "fine_tuned_pcdms.pt")
264#sd_model.save_checkpoint("outputs", "0", {})
265print("Fine-tuning completed. Model saved.")
266
267