huaweilin/VTBench
2
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 