HuangLab/CELL-E_2-Image_Prediction
4
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 