File size: 3,089 Bytes
c87881a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
"""Losses and leakage-free gene-set augmentation."""

from __future__ import annotations

from collections.abc import Sequence

import torch
from torch.nn import functional as F


def supervised_contrastive_loss(
    embeddings: torch.Tensor,
    group_ids: Sequence[str],
    temperature: float = 0.07,
) -> torch.Tensor:
    if embeddings.ndim != 2 or embeddings.shape[0] != len(group_ids):
        raise ValueError("Embeddings and group IDs have inconsistent batch dimensions")
    if temperature <= 0:
        raise ValueError("Temperature must be positive")
    normalized = F.normalize(embeddings, dim=1)
    logits = normalized @ normalized.T / temperature
    identity = torch.eye(len(group_ids), dtype=torch.bool, device=embeddings.device)
    groups = torch.tensor(
        [[left == right for right in group_ids] for left in group_ids],
        dtype=torch.bool,
        device=embeddings.device,
    )
    positives = groups & ~identity
    positive_counts = positives.sum(dim=1)
    if torch.any(positive_counts == 0):
        raise ValueError("Every item must have at least one positive in its batch")
    logits = logits.masked_fill(identity, float("-inf"))
    log_probabilities = logits - torch.logsumexp(logits, dim=1, keepdim=True)
    positive_log_probability = log_probabilities.masked_fill(~positives, 0.0).sum(dim=1)
    positive_log_probability = positive_log_probability / positive_counts
    return -positive_log_probability.mean()


def augment_gene_sets(
    embeddings: torch.Tensor,
    positions: torch.Tensor,
    padding_mask: torch.Tensor,
    gene_dropout: float,
    position_jitter: float,
    generator: torch.Generator | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
    augmented_embeddings = embeddings.clone()
    augmented_positions = positions.clone()
    augmented_mask = padding_mask.clone()
    for row in range(embeddings.shape[0]):
        valid_indices = torch.nonzero(~padding_mask[row], as_tuple=False).flatten()
        maximum_drop = max(0, len(valid_indices) - 2)
        drop_count = min(maximum_drop, int(len(valid_indices) * gene_dropout))
        if drop_count:
            order = torch.randperm(len(valid_indices), generator=generator, device="cpu")
            dropped = valid_indices.cpu()[order[:drop_count]].to(embeddings.device)
            augmented_mask[row, dropped] = True
            augmented_embeddings[row, dropped] = 0
        noise = torch.empty(positions.shape[1], device=positions.device)
        noise.uniform_(-position_jitter, position_jitter, generator=generator)
        augmented_positions[row] = (augmented_positions[row] + noise).clamp(0.0, 1.0)
        augmented_positions[row, augmented_mask[row]] = 0.0
    return augmented_embeddings, augmented_positions, augmented_mask


def masked_gene_loss(
    predicted: torch.Tensor,
    target: torch.Tensor,
    masked_positions: torch.Tensor,
) -> torch.Tensor:
    if not masked_positions.any():
        raise ValueError("At least one gene must be masked")
    return F.mse_loss(predicted[masked_positions], target[masked_positions])