reimash/ai-toolkit
0
1import argparse2import os3from PIL import Image4import torch5from torchvision.transforms import Resize, ToTensor6from diffusers import AutoencoderKL7from pytorch_fid import fid_score8from skimage.metrics import peak_signal_noise_ratio as psnr9import lpips10from tqdm import tqdm11from torchvision import transforms12 13device = torch.device("cuda" if torch.cuda.is_available() else "cpu")14 15def load_images(folder_path):16 images = []17 for filename in os.listdir(folder_path):18 if filename.lower().endswith(('.png', '.jpg', '.jpeg')):19 img_path = os.path.join(folder_path, filename)20 images.append(img_path)21 return images22 23 24def paramiter_count(model):25 state_dict = model.state_dict()26 paramiter_count = 027 for key in state_dict:28 paramiter_count += torch.numel(state_dict[key])29 return int(paramiter_count)30 31 32def calculate_metrics(vae, images, max_imgs=-1, save_output=False):33 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")34 vae = vae.to(device)35 lpips_model = lpips.LPIPS(net='alex').to(device)36 37 rfid_scores = []38 psnr_scores = []39 lpips_scores = []40 41 # transform = transforms.Compose([42 # transforms.Resize(256, antialias=True),43 # transforms.CenterCrop(256)44 # ])45 # needs values between -1 and 146 to_tensor = ToTensor()47 48 # remove _reconstructed.png files49 images = [img for img in images if not img.endswith("_reconstructed.png")]50 51 if max_imgs > 0 and len(images) > max_imgs:52 images = images[:max_imgs]53 54 for img_path in tqdm(images):55 try:56 img = Image.open(img_path).convert('RGB')57 # img_tensor = to_tensor(transform(img)).unsqueeze(0).to(device)58 img_tensor = to_tensor(img).unsqueeze(0).to(device)59 img_tensor = 2 * img_tensor - 160 # if width or height is not divisible by 8, crop it61 if img_tensor.shape[2] % 8 != 0 or img_tensor.shape[3] % 8 != 0:62 img_tensor = img_tensor[:, :, :img_tensor.shape[2] // 8 * 8, :img_tensor.shape[3] // 8 * 8]63 64 except Exception as e:65 print(f"Error processing {img_path}: {e}")66 continue67 68 69 with torch.no_grad():70 reconstructed = vae.decode(vae.encode(img_tensor).latent_dist.sample()).sample71 72 # Calculate rFID73 # rfid = fid_score.calculate_frechet_distance(vae, img_tensor, reconstructed)74 # rfid_scores.append(rfid)75 76 # Calculate PSNR77 psnr_val = psnr(img_tensor.cpu().numpy(), reconstructed.cpu().numpy())78 psnr_scores.append(psnr_val)79 80 # Calculate LPIPS81 lpips_val = lpips_model(img_tensor, reconstructed).item()82 lpips_scores.append(lpips_val)83 84 # avg_rfid = sum(rfid_scores) / len(rfid_scores)85 avg_rfid = 086 avg_psnr = sum(psnr_scores) / len(psnr_scores)87 avg_lpips = sum(lpips_scores) / len(lpips_scores)88 89 if save_output:90 filename_no_ext = os.path.splitext(os.path.basename(img_path))[0]91 folder = os.path.dirname(img_path)92 save_path = os.path.join(folder, filename_no_ext + "_reconstructed.png")93 reconstructed = (reconstructed + 1) / 294 reconstructed = reconstructed.clamp(0, 1)95 reconstructed = transforms.ToPILImage()(reconstructed[0].cpu())96 reconstructed.save(save_path)97 98 return avg_rfid, avg_psnr, avg_lpips99 100 101def main():102 parser = argparse.ArgumentParser(description="Calculate average rFID, PSNR, and LPIPS for VAE reconstructions")103 parser.add_argument("--vae_path", type=str, required=True, help="Path to the VAE model")104 parser.add_argument("--image_folder", type=str, required=True, help="Path to the folder containing images")105 parser.add_argument("--max_imgs", type=int, default=-1, help="Max num of images. Default is -1 for all images.")106 # boolean store true107 parser.add_argument("--save_output", action="store_true", help="Save the output images")108 args = parser.parse_args()109 110 if os.path.isfile(args.vae_path):111 vae = AutoencoderKL.from_single_file(args.vae_path)112 else:113 try:114 vae = AutoencoderKL.from_pretrained(args.vae_path)115 except:116 vae = AutoencoderKL.from_pretrained(args.vae_path, subfolder="vae")117 vae.eval()118 vae = vae.to(device)119 print(f"Model has {paramiter_count(vae)} parameters")120 images = load_images(args.image_folder)121 122 avg_rfid, avg_psnr, avg_lpips = calculate_metrics(vae, images, args.max_imgs, args.save_output)123 124 # print(f"Average rFID: {avg_rfid}")125 print(f"Average PSNR: {avg_psnr}")126 print(f"Average LPIPS: {avg_lpips}")127 128 129if __name__ == "__main__":130 main()131 