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