Views
No views yet
1import torch
2import numpy as np
3from PIL import Image
4from einops import repeat
5from datasets import load_dataset, concatenate_datasets
6from IPython.display import display, HTML
7from torchvision.transforms import ToPILImage, PILToTensor, Compose
8from torchvision.transforms import Resize, RandomCrop, CenterCrop, RandomHorizontalFlip, RandomVerticalFlip, RandomRotation
9from vit_pytorch.mae import MAE
10from vit_pytorch.simple_vit_with_register_tokens import SimpleViT
11from einops.layers.torch import Rearrange
12class Args: pass1device = "cpu"
2checkpoint = torch.load("v0.0.1.pt",map_location="cpu")
3args = checkpoint['args']
4args.crops_per_sample = 1
5
6encoder = SimpleViT(
7 image_size = args.img_dim[1],
8 channels = args.img_dim[0],
9 patch_size = args.patch_size,
10 num_classes = args.num_classes,
11 dim = args.embed_dim,
12 depth = args.depth,
13 heads = args.heads,
14 mlp_dim = args.mlp_dim,
15 dim_head = args.embed_dim//args.heads,
16).to(device)
17
18model = MAE(
19 encoder=encoder,
20 decoder_dim=args.embed_dim,
21 masking_ratio=args.masking_ratio,
22 decoder_depth=args.decoder_depth,
23 decoder_heads=args.heads,
24 decoder_dim_head=args.embed_dim//args.heads,
25).to(device)
26
27model.load_state_dict(checkpoint['model_state_dict'])<All keys matched successfully>dataset = load_dataset("danjacobellis/cell_synthetic_labels")1transforms = Compose([
2 RandomCrop(896),
3 RandomRotation(22.5),
4 CenterCrop(672),
5 Resize(224, interpolation=Image.Resampling.LANCZOS),
6 RandomVerticalFlip(0.5),
7 RandomHorizontalFlip(0.5),
8 PILToTensor(),
9])
10
11def collate_fn(batch):
12 batch_size = len(batch)*args.crops_per_sample
13 inputs = torch.zeros(
14 (batch_size, args.img_dim[0], args.img_dim[1], args.img_dim[2]),
15 dtype=torch.uint8
16 )
17 for i_sample, sample in enumerate(batch):
18 img = sample['image']
19 for i_crop in range(args.crops_per_sample):
20 ind = i_sample*args.crops_per_sample + i_crop
21 inputs[ind,:,:,:] = transforms(img)
22
23 return inputs1data_loader_valid = torch.utils.data.DataLoader(
2 dataset['validation'],
3 batch_size=8,
4 shuffle=False,
5 num_workers=args.num_workers,
6 drop_last=False,
7 pin_memory=True,
8 collate_fn=collate_fn
9)1with torch.no_grad():
2 x = next(iter(data_loader_valid))
3 x = x.to(torch.float)
4 x = x / 255
5 x = x.to(device)
6
7 patches = model.to_patch(x)
8 batch, num_patches, *_ = patches.shape
9
10 tokens = model.patch_to_emb(patches)
11 tokens += model.encoder.pos_embedding.to(device, dtype=tokens.dtype)
12
13 num_masked = int(model.masking_ratio * num_patches)
14 rand_indices = torch.rand(batch, num_patches, device = device).argsort(dim = -1)
15 masked_indices, unmasked_indices = rand_indices[:, :num_masked], rand_indices[:, num_masked:]
16
17 batch_range = torch.arange(batch, device = device)[:, None]
18 tokens = tokens[batch_range, unmasked_indices]
19
20 masked_patches = patches[batch_range, masked_indices]
21 encoded_tokens = model.encoder.transformer(tokens)
22 decoder_tokens = model.enc_to_dec(encoded_tokens)
23 unmasked_decoder_tokens = decoder_tokens + model.decoder_pos_emb(unmasked_indices)
24
25 mask_tokens = repeat(model.mask_token, 'd -> b n d', b = batch, n = num_masked)
26 mask_tokens = mask_tokens + model.decoder_pos_emb(masked_indices)
27
28 decoder_tokens = torch.zeros(batch, num_patches, model.decoder_dim, device=device)
29 decoder_tokens[batch_range, unmasked_indices] = unmasked_decoder_tokens
30 decoder_tokens[batch_range, masked_indices] = mask_tokens
31 decoded_tokens = model.decoder(decoder_tokens)
32
33 mask_tokens = decoded_tokens[batch_range, masked_indices]
34 pred_pixel_values = model.to_pixels(mask_tokens)
35
36 recon_loss = torch.nn.functional.mse_loss(pred_pixel_values, masked_patches)1def reconstruct_image(self, patches, model_input, masked_indices=None, pred_pixel_values=None, patch_size=8):
2 patches = patches.cpu()
3 masked_indices_in = masked_indices is not None
4 predicted_pixels_in = pred_pixel_values is not None
5 if masked_indices_in:
6 masked_indices = masked_indices.cpu()
7 if predicted_pixels_in:
8 pred_pixel_values = pred_pixel_values.cpu()
9 patch_width = patch_height = patch_size
10 reconstructed_image = patches.clone()
11 if masked_indices_in or predicted_pixels_in:
12 for i in range(reconstructed_image.shape[0]):
13 if masked_indices_in and predicted_pixels_in:
14 reconstructed_image[i, masked_indices[i].cpu()] = pred_pixel_values[i, :].cpu().float()
15 elif masked_indices_in:
16 reconstructed_image[i, masked_indices[i].cpu()] = 0
17 invert_patch = Rearrange('b (h w) (p1 p2 c) -> b c (h p1) (w p2)', w=int(model_input.shape[3] / patch_width),
18 h=int(model_input.shape[2] / patch_height), c=model_input.shape[1],
19 p1=patch_height, p2=patch_width)
20 reconstructed_image = invert_patch(reconstructed_image)
21 reconstructed_image = reconstructed_image.numpy().transpose(0, 2, 3, 1)
22 return reconstructed_image.transpose(0, 3, 1, 2)1with torch.no_grad():
2 reconstructed_images1 = reconstruct_image(
3 model,
4 patches,
5 x,
6 masked_indices=masked_indices,
7 pred_pixel_values=pred_pixel_values,
8 patch_size=16
9 )
10 reconstructed_images2 = reconstruct_image(
11 model,
12 patches,
13 x,
14 masked_indices=masked_indices,
15 patch_size=16
16 )1for i_img, img in enumerate(x):
2 rec1 = reconstructed_images1[i_img]
3 rec2 = reconstructed_images2[i_img]
4 display(ToPILImage()(img[0]))
5 display(ToPILImage()(rec2[0]))
6 display(ToPILImage()(rec1[0]))























!jupyter nbconvert --to markdown README.ipynb[NbConvertApp] Converting notebook README.ipynb to markdown
[NbConvertApp] Support files will be in README_files/
[NbConvertApp] Writing 7517 bytes to README.md