Team Ai
Apppublic

HuangLab/CELL-E_2-Image_Prediction

sourceHugging Facemitupdated 2y agoView on Hugging Face
4likes
utils.py229 linesDownload Raw Back to celle
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