recursionpharma/OpenPhenom
221.2k
1# © Recursion Pharmaceuticals 20242import math3 4import torch5 6 7def flatten_images(8 img: torch.Tensor, patch_size: int, channel_agnostic: bool = False9) -> torch.Tensor:10 """11 Flattens 2D images into tokens with the same pixel values12 13 Parameters14 ----------15 img : input image tensor (N, C, H, W)16 17 Returns18 -------19 flattened_img: flattened image tensor (N, L, patch_size**2 * C)20 """21 22 if (img.shape[2] != img.shape[3]) or (img.shape[2] % patch_size != 0):23 raise ValueError("image H must equal image W and be divisible by patch_size")24 in_chans = img.shape[1]25 26 h = w = int(img.shape[2] // patch_size)27 x = img.reshape(shape=(img.shape[0], in_chans, h, patch_size, w, patch_size))28 29 if channel_agnostic:30 x = torch.permute(x, (0, 1, 2, 4, 3, 5)) # NCHPWQ -> NCHWPQ31 x = x.reshape(shape=(img.shape[0], in_chans * h * w, int(patch_size**2)))32 else:33 x = torch.permute(x, (0, 2, 4, 3, 5, 1)) # NCHPWQ -> NHWPQC34 x = x.reshape(shape=(img.shape[0], h * w, int(patch_size**2 * in_chans)))35 return x36 37 38def unflatten_tokens(39 tokens: torch.Tensor,40 patch_size: int,41 num_modalities: int = 1,42 channel_agnostic: bool = False,43) -> torch.Tensor:44 """45 Unflattens tokens (N,L,patch_size**2 * C) into image tensor (N,C,H,W) with the pixel values46 47 Parameters48 ----------49 tokens : input token tensor (N,L,patch_size**2 * C)50 51 Returns52 -------53 img: image tensor (N,C,H,W)54 """55 if num_modalities > 1 and not channel_agnostic:56 raise ValueError("Multiple modalities requires channel agnostic unflattening.")57 58 h = w = int(math.sqrt(tokens.shape[1] // num_modalities))59 if h * w != (tokens.shape[1] // num_modalities):60 raise ValueError("sqrt of number of tokens not integer")61 62 if channel_agnostic:63 x = tokens.reshape(shape=(tokens.shape[0], -1, h, w, patch_size, patch_size))64 x = torch.permute(x, (0, 1, 2, 4, 3, 5)) # NCHWPQ -> NCHPWQ65 else:66 x = tokens.reshape(shape=(tokens.shape[0], h, w, patch_size, patch_size, -1))67 x = torch.permute(x, (0, 5, 1, 3, 2, 4)) # NHWPQC -> NCHPWQ68 img = x.reshape(shape=(x.shape[0], -1, h * patch_size, h * patch_size))69 70 return img71 