CoolFace
Apppublic

acmyu/KeyframesAI

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
test1.py185 linesDownload Raw Back to root
1import os2import torch3from PIL import Image4from src.models.stage1_prior_transformer import Stage1_PriorTransformer5from src.pipelines.stage1_prior_pipeline import Stage1_PriorPipeline6import torch.nn.functional as F7from transformers import (8    CLIPVisionModelWithProjection,9 10    CLIPImageProcessor,11)12import argparse13import numpy as np14 15import torch.multiprocessing as mp16import json17import time18 19 20# Read a text file and convert the coordinates into a tensor21def read_coordinates_file(file_path):22    coordinates_list = []23    with open(file_path, 'r') as file:24        for line in file:25            x, y = map(float, line.strip().split())26            coordinates_list.extend([x, y])27    coordinates_tensor = torch.tensor(coordinates_list, dtype=torch.float32).view(1, -1)28    return coordinates_tensor29 30def split_list_into_chunks(lst, n):31    chunk_size = len(lst) // n32    chunks = [lst[i:i + chunk_size] for i in range(0, len(lst), chunk_size)]33    if len(chunks) > n:34        last_chunk = chunks.pop()35        chunks[-1].extend(last_chunk)36    return chunks37 38def main(args):39 40    device = torch.device("cuda")41    generator = torch.Generator(device=device).manual_seed(args.seed_number)42 43    # save path44    save_dir = "{}/guidancescale{}_seed{}_numsteps{}/".format(args.save_path,  args.guidance_scale, args.seed_number, args.num_inference_steps)45    if not os.path.exists(save_dir):46        os.makedirs(save_dir, exist_ok=True)47 48    # prepare data aug49    clip_image_processor = CLIPImageProcessor()50 51    # prepare model52    model_ckpt = args.weights_name53 54 55    pipe = Stage1_PriorPipeline.from_pretrained(args.pretrained_model_name_or_path).to(device)56    pipe.prior= Stage1_PriorTransformer.from_pretrained(args.pretrained_model_name_or_path, subfolder="prior", num_embeddings=2,embedding_dim=1024, low_cpu_mem_usage=False, ignore_mismatched_sizes=True).to(device)57 58    prior_dict = torch.load(model_ckpt, map_location="cpu")["module"]59    pipe.prior.load_state_dict(prior_dict)60    pipe.enable_xformers_memory_efficient_attention()61 62    image_encoder = CLIPVisionModelWithProjection.from_pretrained(args.image_encoder_path).eval().to(device)63 64 65    print('====================== model load finish ===================')66 67 68    # start test69    start_time = time.time()70 71    #prepare data72    s_img_path = 'imgs/sm.png'73    t_img_path = 'imgs/target.png'74 75    s_pose_path = args.pose_path + select_test_data['source_image'].replace('.jpg', '.txt')76    t_pose_path = (args.pose_path + select_test_data["target_image"].replace(".jpg", ".txt"))77 78    # image_pair79    s_image = Image.open(s_img_path).convert("RGB").resize((args.img_width, args.img_height), Image.BICUBIC)80    #t_image = Image.open(t_img_path).convert("RGB").resize((args.img_width, args.img_height), Image.BICUBIC)81 82    s_pose = read_coordinates_file(s_pose_path).to(device).unsqueeze(1)83    t_pose = read_coordinates_file(t_pose_path).to(device).unsqueeze(1)84 85 86 87 88 89    clip_s_image = clip_image_processor(images=s_image, return_tensors="pt").pixel_values90    #clip_t_image = clip_image_processor(images=t_image, return_tensors="pt").pixel_values91 92 93    with torch.no_grad():94        s_img_embed = (image_encoder(clip_s_image.to(device)).image_embeds).unsqueeze(1)95        #target_embed = image_encoder(clip_t_image.to(device)).image_embeds96 97 98 99 100    output = pipe(101        s_embed = s_img_embed,102        s_pose = s_pose,103        t_pose = t_pose,104        num_images_per_prompt=1,105        num_inference_steps = args.num_inference_steps,106        generator = generator,107        guidance_scale = args.guidance_scale,108    )109 110    # save features111    feature = output[0].cpu().detach().numpy()112    np.save('embed.npy',  feature)113 114    # computer scores115    predict_embed = output[0]116 117    #cosine_similarities = F.cosine_similarity(predict_embed, target_embed)118    #sum_simm += cosine_similarities.item()119 120    end_time =time.time()121    print(end_time-start_time)122 123    """124    avg_simm = sum_simm/number125    with open (save_dir+'/a_results.txt', 'a') as ff:126        ff.write('number is {}, guidance_scale is {}, all averge simm is :{} \n'.format(number, args.guidance_scale, avg_simm))127    print('number is {}, guidance_scale is {}, all averge simm is :{}'.format(number, args.guidance_scale, avg_simm))128    """129 130 131if __name__ == "__main__":132    parser = argparse.ArgumentParser(description="Simple example of a prior model of stage1 script.")133    parser.add_argument("--pretrained_model_name_or_path",type=str,default="./kandinsky-2-2-prior",134        help="Path to pretrained model or model identifier from huggingface.co/models.",)135    parser.add_argument("--image_encoder_path",type=str,default="./OpenCLIP-ViT-H-14",136        help="Path to pretrained model or model identifier from huggingface.co/models.",)137    parser.add_argument("--img_path", type=str, default="./datasets/deepfashing/train_all_png/", help="image path", )138    parser.add_argument("--pose_path", type=str, default="./datasets/deepfashing/normalized_pose_txt/", help="pose path", )139    parser.add_argument("--json_path", type=str, default="./datasets/deepfashing/test_data.json", help="json path", )140    parser.add_argument("--save_path", type=str, default="./save_data/stage1", help="save path", )141    parser.add_argument("--guidance_scale",type=int,default=0,help="guidance_scale",)142    parser.add_argument("--seed_number",type=int,default=42,help="seed number",)143    parser.add_argument("--num_inference_steps",type=int,default=20,help="num_inference_steps",)144    parser.add_argument("--img_width",type=int,default=512,help="image width",)145    parser.add_argument("--img_height",type=int,default=512,help="image height",)146    parser.add_argument("--weights_name",type=str,default="s1_512.pt",help="weights number",)147 148 149    args = parser.parse_args()150    print(args)151    152    """153    # Set the number of GPUs.154    num_devices = torch.cuda.device_count()155 156    print("Using {} GPUs inference".format(num_devices))157 158    # load data159    test_data = json.load(open(args.json_path))160    select_test_datas = test_data161    print('The number of test data: {}'.format(len(select_test_datas)))162 163    # Create a process pool164    mp.set_start_method("spawn")165    data_list = split_list_into_chunks(select_test_datas, num_devices)166 167    processes = []168    for rank in range(num_devices):169        p = mp.Process(target=main, args=(args,rank, data_list[rank], ))170        processes.append(p)171        p.start()172 173 174    for rank, p in enumerate(processes):175        p.join()176    """177 178    main(args)179 180 181 182 183 184 185