Team Ai
Apppublic

HuangLab/CELL-E_2-Image_Prediction

sourceHugging Facemitupdated 2y agoView on Hugging Face
4likes
dataloader.py309 linesDownload Raw Back to root
1import os2import numpy as np3from PIL import Image, ImageSequence4import json5import pandas as pd6 7import torch8from torch.utils.data import Dataset9from torchvision import transforms10import torchvision.transforms.functional as TF11 12from celle.utils import replace_outliers13 14def simple_conversion(seq):15    """Create 26-dim embedding"""16    chars = [17        "-",18        "M",19        "R",20        "H",21        "K",22        "D",23        "E",24        "S",25        "T",26        "N",27        "Q",28        "C",29        "U",30        "G",31        "P",32        "A",33        "V",34        "I",35        "F",36        "Y",37        "W",38        "L",39        "O",40        "X",41        "Z",42        "B",43        "J",44    ]45 46    nums = range(len(chars))47 48    seqs_x = np.zeros(len(seq))49 50    for idx, char in enumerate(seq):51 52        lui = chars.index(char)53 54        seqs_x[idx] = nums[lui]55 56    return torch.tensor([seqs_x]).long()57 58 59class CellLoader(Dataset):60    """imports mined opencell images with protein sequence"""61 62    def __init__(63        self,64        data_csv=None,65        dataset=None,66        split_key=None,67        resize=600,68        crop_size=600,69        crop_method="random",70        sequence_mode="simple",71        vocab="bert",72        threshold="median",73        text_seq_len=0,74        pad_mode="random",75    ):76        self.data_csv = data_csv77        self.dataset = dataset78        self.image_folders = []79        self.crop_method = crop_method80        self.resize = resize81        self.crop_size = crop_size82        self.sequence_mode = sequence_mode83        self.threshold = threshold84        self.text_seq_len = int(text_seq_len)85        self.vocab = vocab86        self.pad_mode = pad_mode87 88        if self.sequence_mode == "embedding" or self.sequence_mode == "onehot":89 90 91            if self.vocab == "esm1b" or self.vocab == "esm2":92                from esm import Alphabet93 94                self.tokenizer = Alphabet.from_architecture(95                    "ESM-1b"96                ).get_batch_converter()97                self.text_seq_len += 298 99        if data_csv:100 101            data = pd.read_csv(data_csv)102 103            self.parent_path = os.path.dirname(data_csv).split(data_csv)[0]104 105            if split_key == "train":106                self.data = data[data["split"] == "train"]107            elif split_key == "val":108                self.data = data[data["split"] == "val"]109            else:110                self.data = data111 112            self.data = self.data.reset_index(drop=True)113            114                115 116    def __len__(self):117        return len(self.data)118 119    def __getitem__(120        self,121        idx,122        get_sequence=True,123        get_images=True,124    ):125        if get_sequence and self.text_seq_len > 0:126 127            protein_vector = self.get_protein_vector(idx)128 129        else:130            protein_vector = torch.zeros((1, 1))131 132        if get_images:133 134            nucleus, target, threshold = self.get_images(idx, self.dataset)135        else:136            nucleus, target, threshold = torch.zeros((3, 1))137 138        data_dict = {139            "nucleus": nucleus.float(),140            "target": target.float(),141            "threshold": threshold.float(),142            "sequence": protein_vector.long(),143        }144 145        return data_dict146 147    def get_protein_vector(self, idx):148 149        if "protein_sequence" not in self.data.columns:150 151            metadata = self.retrieve_metadata(idx)152            protein_sequence = metadata["sequence"]153        else:154            protein_sequence = self.data.iloc[idx]["protein_sequence"]155 156        protein_vector = self.tokenize_sequence(protein_sequence)157 158        return protein_vector159 160    def get_images(self, idx, dataset):161 162        if dataset == "HPA":163 164            nucleus = Image.open(165                os.path.join(166                    self.parent_path, self.data.iloc[idx]["nucleus_image_path"]167                )168            )169 170            target = Image.open(171                os.path.join(self.parent_path, self.data.iloc[idx]["target_image_path"])172            )173 174            nucleus = TF.to_tensor(nucleus)[0]175            target = TF.to_tensor(target)[0]176 177            image = torch.stack([nucleus, target], axis=0)178 179            normalize = (0.0655, 0.0650), (0.1732, 0.1208)180 181        elif dataset == "OpenCell":182            image = Image.open(183                os.path.join(self.parent_path, self.data.iloc[idx]["image_path"])184            )185            nucleus, target = [page.copy() for page in ImageSequence.Iterator(image)]186 187            nucleus = replace_outliers(torch.divide(TF.to_tensor(nucleus), 65536))[0]188            target = replace_outliers(torch.divide(TF.to_tensor(target), 65536))[0]189 190            image = torch.stack([nucleus, target], axis=0)191 192            normalize = (193                (0.0272, 0.0244),194                (0.0486, 0.0671),195            )196 197        # # from https://discuss.pytorch.org/t/how-to-apply-same-transform-on-a-pair-of-picture/14914198 199        t_forms = [transforms.Resize(self.resize, antialias=None)]200 201        if self.crop_method == "random":202 203            t_forms.append(transforms.RandomCrop(self.crop_size))204            t_forms.append(transforms.RandomHorizontalFlip(p=0.5))205            t_forms.append(transforms.RandomVerticalFlip(p=0.5))206 207        elif self.crop_method == "center":208 209            t_forms.append(transforms.CenterCrop(self.crop_size))210 211        t_forms.append(transforms.Normalize(normalize[0], normalize[1]))212 213        image = transforms.Compose(t_forms)(image)214 215        nucleus, target = image216 217        nucleus /= torch.abs(nucleus).max()218        target -= target.min()219        target /= target.max()220 221        nucleus = nucleus.unsqueeze(0)222        target = target.unsqueeze(0)223 224        threshold = target225 226        if self.threshold == "mean":227 228            threshold = 1.0 * (threshold > (torch.mean(threshold)))229 230        elif self.threshold == "median":231 232            threshold = 1.0 * (threshold > (torch.median(threshold)))233 234        elif self.threshold == "1090_IQR":235 236            p10 = torch.quantile(threshold, 0.1, None)237            p90 = torch.quantile(threshold, 0.9, None)238            threshold = torch.clip(threshold, p10, p90)239 240        nucleus = torch.nan_to_num(nucleus, 0.0, 1.0, 0.0)241        target = torch.nan_to_num(target, 0.0, 1.0, 0.0)242        threshold = torch.nan_to_num(threshold, 0.0, 1.0, 0.0)243 244        return nucleus, target, threshold245 246    def retrieve_metadata(self, idx):247        with open(248            os.path.join(self.parent_path, self.data.iloc[idx]["metadata_path"])249        ) as f:250            metadata = json.load(f)251 252        return metadata253 254    def tokenize_sequence(self, protein_sequence):255 256        pad_token = 0257 258        if self.sequence_mode == "simple":259            protein_vector = simple_conversion(protein_sequence)260 261        elif self.sequence_mode == "center":262            protein_sequence = protein_sequence.center(self.text_seq_length, "-")263            protein_vector = simple_conversion(protein_sequence)264 265        elif self.sequence_mode == "alternating":266            protein_sequence = protein_sequence.center(self.text_seq_length, "-")267            protein_sequence = protein_sequence[::18]268            protein_sequence = protein_sequence.center(269                int(self.text_seq_length / 18) + 1, "-"270            )271            protein_vector = simple_conversion(protein_sequence)272 273 274        elif self.sequence_mode == "embedding":275 276            if self.vocab == "esm1b" or self.vocab == "esm2":277                pad_token = 1278                protein_vector = self.tokenizer([("", protein_sequence)])[-1]279 280        if protein_vector.shape[-1] < self.text_seq_len:281 282            diff = self.text_seq_len - protein_vector.shape[-1]283 284            if self.pad_mode == "end":285                protein_vector = torch.nn.functional.pad(286                    protein_vector, (0, diff), "constant", pad_token287                )288            elif self.pad_mode == "random":289                split = diff - np.random.randint(0, diff + 1)290 291                protein_vector = torch.cat(292                    [torch.ones(1, split) * 0, protein_vector], dim=1293                )294 295                protein_vector = torch.nn.functional.pad(296                    protein_vector, (0, diff - split), "constant", pad_token297                )298 299        elif protein_vector.shape[-1] > self.text_seq_len:300            start_int = np.random.randint(301                0, protein_vector.shape[-1] - self.text_seq_len302            )303 304            protein_vector = protein_vector[305                :, start_int : start_int + self.text_seq_len306            ]307 308        return protein_vector.long()309