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