CoolFace
Apppublic

reimash/ai-toolkit

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
test_vae.py131 linesDownload Raw Back to testing
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