CoolFace
Apppublic

huaweilin/VTBench

sourceHugging Faceupdated 1y agoView on Hugging Face
2likes
main.py100 linesDownload Raw Back to root
1import numpy as np2import os3import PIL4import pickle5import torch6import argparse7import json8from PIL import Image9import torch.nn as nn10import torch11from transformers import AutoProcessor, AutoModelForImageTextToText12from src.data_loader import DataCollatorForSupervisedDataset, get_dataset13from src.data_processing import tensor_to_pil14from src.model_processing import get_model15from PIL import Image16from accelerate import Accelerator17from torch.utils.data import DataLoader18from tqdm import tqdm19from concurrent.futures import ThreadPoolExecutor20 21parser = argparse.ArgumentParser()22parser.add_argument("--model_name", type=str, default="chameleon")23parser.add_argument("--model_path", type=str, default=None)24parser.add_argument("--dataset_name", type=str, default="task3-movie-posters")25parser.add_argument("--split_name", type=str, default="test")26parser.add_argument("--batch_size", default=8, type=int)27parser.add_argument("--output_dir", type=str, default=None)28parser.add_argument("--begin_id", default=0, type=int)29parser.add_argument("--n_take", default=-1, type=int)30args = parser.parse_args()31 32batch_size = args.batch_size33output_dir = args.output_dir34 35accelerator = Accelerator()36 37if accelerator.is_main_process and output_dir is not None:38    os.makedirs(output_dir, exist_ok=True)39    os.makedirs(f"{output_dir}/original_images", exist_ok=True)40    os.makedirs(f"{output_dir}/reconstructed_images", exist_ok=True)41    os.makedirs(f"{output_dir}/results", exist_ok=True)42 43model, data_params = get_model(args.model_path, args.model_name)44dataset = get_dataset(args.dataset_name, args.split_name, None if args.n_take <= 0 else args.n_take)45data_collator = DataCollatorForSupervisedDataset(args.dataset_name, **data_params)46dataloader = DataLoader(47    dataset, batch_size=batch_size, num_workers=0, collate_fn=data_collator48)49 50model, dataloader = accelerator.prepare(model, dataloader)51print("Model prepared...")52 53 54def save_results(55    pixel_values, reconstructed_image, idx, output_dir, data_params56):57    if reconstructed_image is None:58        return59 60    ori_img = tensor_to_pil(pixel_values, **data_params)61    rec_img = tensor_to_pil(reconstructed_image, **data_params)62 63    ori_img.save(f"{output_dir}/original_images/{idx:08d}.png")64    rec_img.save(f"{output_dir}/reconstructed_images/{idx:08d}.png")65 66    result = {67        "ori_img": ori_img,68        "rec_img": rec_img,69    }70 71    with open(f"{output_dir}/results/{idx:08d}.pickle", "wb") as fw:72        pickle.dump(result, fw)73 74 75executor = ThreadPoolExecutor(max_workers=16)76with torch.no_grad():77    print("Begin data loading...")78    for batch in tqdm(dataloader):79        pixel_values = batch["image"]80        reconstructed_images = model(pixel_values)81        if isinstance(reconstructed_images, tuple):82            reconstructed_images = reconstructed_images[0]83 84        if output_dir is not None:85            idx_list = batch["idx"]86            original_images = pixel_values.detach().cpu()87            if not isinstance(reconstructed_images, list):88                reconstructed_images = reconstructed_images.detach().cpu()89            for i in range(pixel_values.shape[0]):90                executor.submit(91                    save_results,92                    original_images[i],93                    reconstructed_images[i],94                    idx_list[i],95                    output_dir,96                    data_params,97                )98 99executor.shutdown(wait=True)100