CoolFace
Apppublic

acmyu/KeyframesAI

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