HuangLab/CELL-E_2-Image_Prediction
4
1import torch2from torchvision import transforms3from math import pi4import torchvision.transforms.functional as TF5 6 7# Define helper functions8def exists(val):9 """Check if a variable exists"""10 return val is not None11 12 13def uniq(arr):14 return {el: True for el in arr}.keys()15 16 17def default(val, d):18 """If a value exists, return it; otherwise, return a default value"""19 return val if exists(val) else d20 21 22def max_neg_value(t):23 return -torch.finfo(t.dtype).max24 25 26def cast_tuple(val, depth=1):27 if isinstance(val, list):28 val = tuple(val)29 return val if isinstance(val, tuple) else (val,) * depth30 31 32def is_empty(t):33 """Check if a tensor is empty"""34 # Return True if the number of elements in the tensor is zero, else False35 return t.nelement() == 036 37 38def masked_mean(t, mask, dim=1):39 """40 Compute the mean of a tensor, masked by a given mask41 42 Args:43 t (torch.Tensor): input tensor of shape (batch_size, seq_len, hidden_dim)44 mask (torch.Tensor): mask tensor of shape (batch_size, seq_len)45 dim (int): dimension along which to compute the mean (default=1)46 47 Returns:48 torch.Tensor: masked mean tensor of shape (batch_size, hidden_dim)49 """50 t = t.masked_fill(~mask[:, :, None], 0.0)51 return t.sum(dim=1) / mask.sum(dim=1)[..., None]52 53 54def set_requires_grad(model, value):55 """56 Set whether or not the model's parameters require gradients57 58 Args:59 model (torch.nn.Module): the PyTorch model to modify60 value (bool): whether or not to require gradients61 """62 for param in model.parameters():63 param.requires_grad = value64 65 66def eval_decorator(fn):67 """68 Decorator function to evaluate a given function69 70 Args:71 fn (callable): function to evaluate72 73 Returns:74 callable: the decorated function75 """76 77 def inner(model, *args, **kwargs):78 was_training = model.training79 model.eval()80 out = fn(model, *args, **kwargs)81 model.train(was_training)82 return out83 84 return inner85 86 87def log(t, eps=1e-20):88 """89 Compute the natural logarithm of a tensor90 91 Args:92 t (torch.Tensor): input tensor93 eps (float): small value to add to prevent taking the log of 0 (default=1e-20)94 95 Returns:96 torch.Tensor: the natural logarithm of the input tensor97 """98 return torch.log(t + eps)99 100 101def gumbel_noise(t):102 """103 Generate Gumbel noise104 105 Args:106 t (torch.Tensor): input tensor107 108 Returns:109 torch.Tensor: a tensor of Gumbel noise with the same shape as the input tensor110 """111 noise = torch.zeros_like(t).uniform_(0, 1)112 return -log(-log(noise))113 114 115def gumbel_sample(t, temperature=0.9, dim=-1):116 """117 Sample from a Gumbel-softmax distribution118 119 Args:120 t (torch.Tensor): input tensor of shape (batch_size, num_classes)121 temperature (float): temperature for the Gumbel-softmax distribution (default=0.9)122 dim (int): dimension along which to sample (default=-1)123 124 Returns:125 torch.Tensor: a tensor of samples from the Gumbel-softmax distribution with the same shape as the input tensor126 """127 return (t / max(temperature, 1e-10)) + gumbel_noise(t)128 129 130def top_k(logits, thres=0.5):131 """132 Return a tensor where all but the top k values are set to negative infinity133 134 Args:135 logits (torch.Tensor): input tensor of shape (batch_size, num_classes)136 thres (float): threshold for the top k values (default=0.5)137 138 Returns:139 torch.Tensor: a tensor with the same shape as the input tensor, where all but the top k values are set to negative infinity140 """141 num_logits = logits.shape[-1]142 k = max(int((1 - thres) * num_logits), 1)143 val, ind = torch.topk(logits, k)144 probs = torch.full_like(logits, float("-inf"))145 probs.scatter_(-1, ind, val)146 return probs147 148 149def gamma_func(mode="cosine", scale=0.15):150 """Return a function that takes a single input r and returns a value based on the selected mode"""151 152 # Define a different function based on the selected mode153 if mode == "linear":154 return lambda r: 1 - r155 elif mode == "cosine":156 return lambda r: torch.cos(r * pi / 2)157 elif mode == "square":158 return lambda r: 1 - r**2159 elif mode == "cubic":160 return lambda r: 1 - r**3161 elif mode == "scaled-cosine":162 return lambda r: scale * (torch.cos(r * pi / 2))163 else:164 # Raise an error if the selected mode is not implemented165 raise NotImplementedError166 167 168class always:169 """Helper class to always return a given value"""170 171 def __init__(self, val):172 self.val = val173 174 def __call__(self, x, *args, **kwargs):175 return self.val176 177 178class DivideMax(torch.nn.Module):179 def __init__(self, dim):180 super().__init__()181 self.dim = dim182 183 def forward(self, x):184 maxes = x.amax(dim=self.dim, keepdim=True).detach()185 return x / maxes186 187def replace_outliers(image, percentile=0.0001):188 189 lower_bound, upper_bound = torch.quantile(image, percentile), torch.quantile(190 image, 1 - percentile191 )192 mask = (image <= upper_bound) & (image >= lower_bound)193 194 valid_pixels = image[mask]195 196 image[~mask] = torch.clip(image[~mask], min(valid_pixels), max(valid_pixels))197 198 return image199 200 201def process_image(image, dataset=None, image_type=None):202 image /= image.max()203 204 if dataset == "HPA":205 if image_type == 'nucleus':206 normalize = (0.0655, 0.0650)207 208 elif image_type == 'protein':209 normalize = (0.1732, 0.1208)210 211 elif dataset == "OpenCell":212 213 if image_type == 'nucleus':214 normalize = (0.0272, 0.0244)215 216 elif image_type == 'protein':217 normalize = (0.0486, 0.0671)218 219 t_forms = []220 221 t_forms.append(transforms.RandomCrop(256))222 223 # t_forms.append(transforms.Normalize(normalize[0],normalize[1]))224 225 226 image = transforms.Compose(t_forms)(image)227 228 return image229 