acmyu/KeyframesAI
0
1import os
2import json
3import cv2
4import torch
5from torch import nn
6from PIL import Image
7import numpy as np
8from diffusers import UniPCMultistepScheduler
9import torch.nn.functional as F
10from torchvision import transforms
11from diffusers import AutoencoderKL, DDPMScheduler, UNet2DConditionModel
12from transformers import CLIPImageProcessor
13from src.pipelines.stage3_refined_pipeline import Stage3_RefinedPipeline
14import argparse
15from transformers import Dinov2Model
16from typing import Any, Dict, List, Optional, Tuple, Union
17from skimage.metrics import structural_similarity as compare_ssim
18import torch
19import torch.nn as nn
20import torch.multiprocessing as mp
21import json
22import time
23def split_list_into_chunks(lst, n):
24 chunk_size = len(lst) // n
25 chunks = [lst[i:i + chunk_size] for i in range(0, len(lst), chunk_size)]
26 if len(chunks) > n:
27 last_chunk = chunks.pop()
28 chunks[-1].extend(last_chunk)
29 return chunks
30
31def image_grid(imgs, rows, cols):
32 assert len(imgs) == rows * cols
33
34 w, h = imgs[0].size
35 grid = Image.new("RGB", size=(cols * w, rows * h))
36 grid_w, grid_h = grid.size
37
38 for i, img in enumerate(imgs):
39 grid.paste(img, box=(i % cols * w, i // cols * h))
40 return grid
41
42def zero_module(module):
43 for p in module.parameters():
44 nn.init.zeros_(p)
45 return module
46
47
48
49
50
51class ImageProjModel_p(torch.nn.Module):
52 """SD model with image prompt"""
53
54 def __init__(self, in_dim, hidden_dim, out_dim, dropout = 0.):
55 super().__init__()
56
57 self.net = nn.Sequential(
58 nn.Linear(in_dim, hidden_dim),
59 nn.GELU(),
60 nn.Dropout(dropout),
61 nn.LayerNorm(hidden_dim),
62 nn.Linear(hidden_dim, out_dim),
63 nn.Dropout(dropout)
64 )
65
66 def forward(self, x): # b, 257,1280
67 return self.net(x)
68
69
70
71
72def inference():
73
74 device = "cuda"
75 generator = torch.Generator(device=device).manual_seed(42)
76
77 clip_image_processor = CLIPImageProcessor()
78
79 img_transform = transforms.Compose([
80 transforms.ToTensor(),
81 transforms.Normalize([0.5], [0.5]),
82 ])
83
84
85 # model define
86 image_proj_model_p_dict = {}
87 unet_dict = {}
88
89 image_encoder_p = Dinov2Model.from_pretrained('facebook/dinov2-giant').to(device).eval()
90
91 image_proj_model_p = ImageProjModel_p(in_dim=1536, hidden_dim=768, out_dim=1024).to(device).eval()
92
93 #model_ckpt = "{}/mp_rank_00_model_states.pt".format('{save_ckpt}')
94 model_ckpt = "s3_512.pt"
95 with torch.no_grad():
96 model_sd = torch.load(model_ckpt)["module"]
97
98 for k in model_sd.keys():
99 if k.startswith("image_proj_model_p"):
100 image_proj_model_p_dict[k.replace("image_proj_model_p.", "")] = model_sd[k]
101
102 elif k.startswith("unet"):
103 unet_dict[k.replace("unet.", "")] = model_sd[k]
104
105 else:
106 print(k)
107
108
109 image_proj_model_p.load_state_dict(image_proj_model_p_dict)
110
111 pipe = Stage3_RefinedPipeline.from_pretrained("stabilityai/stable-diffusion-2-1-base",torch_dtype=torch.float16).to(device)
112
113 pipe.unet= UNet2DConditionModel.from_pretrained("stabilityai/stable-diffusion-2-1-base", subfolder="unet",
114 in_channels=8, low_cpu_mem_usage=False, ignore_mismatched_sizes=True).to(device)
115
116 pipe.unet.load_state_dict(unet_dict)
117
118 pipe.scheduler = UniPCMultistepScheduler.from_config(pipe.scheduler.config)
119 pipe.enable_xformers_memory_efficient_attention()
120
121
122
123
124
125 all_ssim = []
126
127 s_img_path = 'imgs/sm.png'
128 #t_img_path = 'imgs/expected.png'
129 gen_t_img_path = 'imgs/coarse.png'
130
131 s_img = Image.open(s_img_path).convert("RGB").resize((512,512), Image.BICUBIC)
132 #t_img = Image.open(t_img_path).convert("RGB").resize((512,512), Image.BICUBIC)
133 gen_t_img = Image.open(gen_t_img_path).convert("RGB").resize((512,512), Image.BICUBIC)
134
135
136
137 clip_processor_s_img = clip_image_processor(images=s_img, return_tensors="pt").pixel_values
138 s_img_f = image_encoder_p(clip_processor_s_img.to(device)).last_hidden_state
139 s_img_proj_f = image_proj_model_p(s_img_f) # s_img
140
141
142 vae_gen_t_image = torch.unsqueeze(img_transform(gen_t_img), 0)
143
144
145
146
147
148 output = pipe(
149 height=512,
150 width=512,
151 guidance_rescale=2.0,
152 vae_gen_t_image=vae_gen_t_image,
153 s_img_proj_f=s_img_proj_f,
154 num_images_per_prompt=4,
155 guidance_scale=1.0,
156 generator=generator,
157 num_inference_steps=20,
158 )
159
160 for i, r in enumerate(output.images):
161 r.save('out'+str(i)+'.png')
162
163 save_output = []
164 result = output.images[0].crop((512, 0, 512 * 2, 512))
165 save_output.append(result.resize((352, 512), Image.BICUBIC))
166 save_output.insert(0, gen_t_img.resize((352, 512), Image.BICUBIC))
167 save_output.insert(0, s_img.resize((352, 512), Image.BICUBIC))
168 grid = image_grid(save_output, 1, 3)
169 grid.save("out.png")
170
171
172
173
174if __name__ == "__main__":
175
176 inference()
177
178
179
180 