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])
|