whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
3.09 kB
"""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])