Sumit-Jethani/Self_Supervised_Learning_using_Masked_AutoEncoders
0
1import numpy as np2import torch3from model import patchify, unpatchify4from config import PATCH_SIZE, IMG_SIZE5 6MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)7STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)8 9 10def denormalise(t):11 img = (t.cpu().float() * STD + MEAN).clamp(0, 1)12 return (img.permute(1, 2, 0).numpy() * 255).astype(np.uint8)13 14 15def make_masked_image(image, mask, patch_size=PATCH_SIZE, img_size=IMG_SIZE):16 w = img_size // patch_size17 out = denormalise(image).copy().astype(float)18 for idx in range(mask.size(0)):19 if mask[idx] == 1:20 r = (idx // w) * patch_size21 c = (idx % w) * patch_size22 out[r:r+patch_size, c:c+patch_size] = 12723 return out.astype(np.uint8)24 25 26def reconstruct_image(pred_patches, original, mask):27 orig_p = patchify(original.unsqueeze(0)).squeeze(0)28 mean = orig_p.mean(dim=-1, keepdim=True)29 var = orig_p.var(dim=-1, keepdim=True)30 pred_d = pred_patches * (var + 1e-6).sqrt() + mean31 comp = torch.where(mask.unsqueeze(-1).bool(), pred_d, orig_p)32 return denormalise(unpatchify(comp.unsqueeze(0)).squeeze(0))