Team Ai
Modelpublic

recursionpharma/OpenPhenom

sourceHugging Faceupdated 7mo agoView on Hugging Face
22likes1.2kdownloads
mae_utils.py71 linesDownload Raw Back to root
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