| """Set Transformer using only sequence embeddings and relative position.""" |
|
|
| from __future__ import annotations |
|
|
| from dataclasses import asdict, dataclass |
|
|
| import torch |
| from torch import nn |
| from torch.nn import functional as F |
|
|
|
|
| @dataclass(frozen=True) |
| class ModelConfig: |
| architecture: str = "setnet" |
| esm_dimension: int = 1280 |
| pfam_dimension: int = 64 |
| pfam_vocab_size: int = 0 |
| hidden_dimension: int = 256 |
| output_dimension: int = 256 |
| attention_heads: int = 8 |
| inducing_points: int = 32 |
| feedforward_dimension: int = 512 |
| attention_blocks: int = 2 |
| position_bins: int = 64 |
| dropout: float = 0.1 |
|
|
| @classmethod |
| def from_dict(cls, values: dict[str, object]) -> "ModelConfig": |
| return cls(**{key: values[key] for key in asdict(cls()) if key in values}) |
|
|
|
|
| class GeneProjection(nn.Module): |
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| self.projection = nn.Linear(config.esm_dimension, config.hidden_dimension) |
| self.normalization = nn.LayerNorm(config.hidden_dimension) |
| self.position = nn.Embedding(config.position_bins, config.hidden_dimension) |
| self.dropout = nn.Dropout(config.dropout) |
| self.position_bins = config.position_bins |
|
|
| def forward(self, embeddings: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: |
| bins = (positions * (self.position_bins - 1)).long().clamp(0, self.position_bins - 1) |
| projected = F.gelu(self.normalization(self.projection(embeddings))) |
| return self.dropout(projected + self.position(bins)) |
|
|
|
|
| class PfamGeneProjection(nn.Module): |
| """Project ESM genes together with a BGC-level Pfam inventory embedding.""" |
|
|
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| if config.pfam_vocab_size < 3: |
| raise ValueError("Pfam vocabulary must contain padding, unknown, and one domain token") |
| self.pfam_embedding = nn.Embedding( |
| config.pfam_vocab_size, config.pfam_dimension, padding_idx=0 |
| ) |
| self.projection = nn.Linear( |
| config.esm_dimension + config.pfam_dimension, config.hidden_dimension |
| ) |
| self.normalization = nn.LayerNorm(config.hidden_dimension) |
| self.position = nn.Embedding(config.position_bins, config.hidden_dimension) |
| self.dropout = nn.Dropout(config.dropout) |
| self.position_bins = config.position_bins |
|
|
| def forward( |
| self, |
| embeddings: torch.Tensor, |
| positions: torch.Tensor, |
| pfam_tokens: torch.Tensor, |
| ) -> torch.Tensor: |
| bins = (positions * (self.position_bins - 1)).long().clamp(0, self.position_bins - 1) |
| token_mask = pfam_tokens.ne(0).unsqueeze(-1) |
| token_values = self.pfam_embedding(pfam_tokens).masked_fill(~token_mask, 0.0) |
| counts = token_mask.sum(dim=1).clamp_min(1) |
| pfam_summary = token_values.sum(dim=1) / counts |
| pfam_summary = pfam_summary.unsqueeze(1).expand(-1, embeddings.shape[1], -1) |
| merged = torch.cat((embeddings, pfam_summary), dim=-1) |
| projected = F.gelu(self.normalization(self.projection(merged))) |
| return self.dropout(projected + self.position(bins)) |
|
|
|
|
| class InducedSelfAttentionBlock(nn.Module): |
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| dimension = config.hidden_dimension |
| self.inducing = nn.Parameter(torch.empty(config.inducing_points, dimension)) |
| nn.init.xavier_uniform_(self.inducing) |
| self.inducing_norm = nn.LayerNorm(dimension) |
| self.input_norm = nn.LayerNorm(dimension) |
| self.summary_attention = nn.MultiheadAttention( |
| dimension, config.attention_heads, config.dropout, batch_first=True |
| ) |
| self.output_attention = nn.MultiheadAttention( |
| dimension, config.attention_heads, config.dropout, batch_first=True |
| ) |
| self.output_norm = nn.LayerNorm(dimension) |
| self.feedforward = nn.Sequential( |
| nn.Linear(dimension, config.feedforward_dimension), |
| nn.GELU(), |
| nn.Dropout(config.dropout), |
| nn.Linear(config.feedforward_dimension, dimension), |
| nn.Dropout(config.dropout), |
| ) |
|
|
| def forward(self, values: torch.Tensor, padding_mask: torch.Tensor | None) -> torch.Tensor: |
| batch_size = values.shape[0] |
| inducing = self.inducing.unsqueeze(0).expand(batch_size, -1, -1) |
| normalized_values = self.input_norm(values) |
| summary, _ = self.summary_attention( |
| inducing, normalized_values, normalized_values, key_padding_mask=padding_mask |
| ) |
| normalized_summary = self.inducing_norm(summary) |
| update, _ = self.output_attention( |
| normalized_values, normalized_summary, normalized_summary |
| ) |
| output = values + update |
| return output + self.feedforward(self.output_norm(output)) |
|
|
|
|
| class PoolingByAttention(nn.Module): |
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| self.seed = nn.Parameter(torch.empty(1, config.hidden_dimension)) |
| nn.init.xavier_uniform_(self.seed) |
| self.normalization = nn.LayerNorm(config.hidden_dimension) |
| self.attention = nn.MultiheadAttention( |
| config.hidden_dimension, config.attention_heads, config.dropout, batch_first=True |
| ) |
|
|
| def forward(self, values: torch.Tensor, padding_mask: torch.Tensor | None) -> torch.Tensor: |
| seed = self.seed.unsqueeze(0).expand(values.shape[0], -1, -1) |
| output, _ = self.attention( |
| seed, |
| self.normalization(values), |
| self.normalization(values), |
| key_padding_mask=padding_mask, |
| ) |
| return output.squeeze(1) |
|
|
|
|
| class LeakageFreeBGCSetNet(nn.Module): |
| """Map a variable-length BGC gene set to a normalized embedding.""" |
|
|
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") |
|
|
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| self.config = config |
| self.gene_projection = GeneProjection(config) |
| self.blocks = nn.ModuleList( |
| InducedSelfAttentionBlock(config) for _ in range(config.attention_blocks) |
| ) |
| self.pooling = PoolingByAttention(config) |
| self.output = nn.Linear(config.hidden_dimension, config.output_dimension) |
|
|
| def encode_genes( |
| self, |
| gene_embeddings: torch.Tensor, |
| relative_positions: torch.Tensor, |
| padding_mask: torch.Tensor | None = None, |
| ) -> torch.Tensor: |
| values = self.gene_projection(gene_embeddings, relative_positions) |
| for block in self.blocks: |
| values = block(values, padding_mask) |
| return values |
|
|
| def forward( |
| self, |
| gene_embeddings: torch.Tensor, |
| relative_positions: torch.Tensor, |
| padding_mask: torch.Tensor | None = None, |
| ) -> torch.Tensor: |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) |
| pooled = self.pooling(genes, padding_mask) |
| return F.normalize(self.output(pooled), p=2, dim=-1) |
|
|
|
|
| class MaskedGenePredictionHead(nn.Module): |
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| self.layers = nn.Sequential( |
| nn.Linear(config.hidden_dimension, config.feedforward_dimension), |
| nn.GELU(), |
| nn.Linear(config.feedforward_dimension, config.esm_dimension), |
| ) |
|
|
| def forward(self, contextual_embeddings: torch.Tensor) -> torch.Tensor: |
| return self.layers(contextual_embeddings) |
|
|
| class GatedDeepSets(nn.Module): |
| """Permutation-invariant learned pooling without gene-gene interactions.""" |
|
|
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") |
|
|
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| self.config = config |
| self.gene_projection = GeneProjection(config) |
| self.gate = nn.Sequential( |
| nn.Linear(config.hidden_dimension, config.hidden_dimension // 2), |
| nn.GELU(), |
| nn.Linear(config.hidden_dimension // 2, 1), |
| ) |
| self.output = nn.Linear(config.hidden_dimension, config.output_dimension) |
|
|
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): |
| return self.gene_projection(gene_embeddings, relative_positions) |
|
|
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) |
| logits = self.gate(genes).squeeze(-1) |
| if padding_mask is not None: |
| logits = logits.masked_fill(padding_mask, -torch.finfo(logits.dtype).max) |
| weights = torch.softmax(logits, dim=-1) |
| if padding_mask is not None: |
| weights = weights.masked_fill(padding_mask, 0.0) |
| pooled = (genes * weights.unsqueeze(-1)).sum(dim=1) |
| return F.normalize(self.output(pooled), p=2, dim=-1) |
|
|
|
|
| class RotarySelfAttentionBlock(nn.Module): |
| """Bidirectional self-attention with rotary position phases.""" |
|
|
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| d = config.hidden_dimension |
| self.heads = config.attention_heads |
| self.head_dim = d // self.heads |
| if self.head_dim % 2: |
| raise ValueError("Rotary attention requires an even per-head dimension") |
| self.norm = nn.LayerNorm(d) |
| self.qkv = nn.Linear(d, 3 * d) |
| self.output = nn.Linear(d, d) |
| self.dropout = nn.Dropout(config.dropout) |
| self.ffn_norm = nn.LayerNorm(d) |
| self.ffn = nn.Sequential( |
| nn.Linear(d, config.feedforward_dimension), nn.GELU(), |
| nn.Dropout(config.dropout), nn.Linear(config.feedforward_dimension, d), |
| nn.Dropout(config.dropout), |
| ) |
|
|
| def _rotate(self, values: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: |
| half = self.head_dim // 2 |
| inv = 1.0 / (10000.0 ** ( |
| torch.arange(half, device=values.device, dtype=values.dtype) / half |
| )) |
| angles = positions[:, None, :, None] * 128.0 * inv[None, None, None, :] |
| cos, sin = angles.cos(), angles.sin() |
| first, second = values[..., :half], values[..., half:] |
| return torch.cat((first * cos - second * sin, first * sin + second * cos), dim=-1) |
|
|
| def forward(self, values: torch.Tensor, positions: torch.Tensor, padding_mask=None) -> torch.Tensor: |
| normalized = self.norm(values) |
| batch, length, dimension = normalized.shape |
| qkv = self.qkv(normalized).view(batch, length, 3, self.heads, self.head_dim) |
| q, k, v = qkv.unbind(dim=2) |
| q, k = q.transpose(1, 2), k.transpose(1, 2) |
| v = v.transpose(1, 2) |
| q, k = self._rotate(q, positions), self._rotate(k, positions) |
| scores = torch.matmul(q, k.transpose(-2, -1)) / (self.head_dim ** 0.5) |
| if padding_mask is not None: |
| scores = scores.masked_fill( |
| padding_mask[:, None, None, :], -torch.finfo(scores.dtype).max |
| ) |
| attention = self.dropout(torch.softmax(scores, dim=-1)) |
| contextual = torch.matmul(attention, v).transpose(1, 2).reshape(batch, length, dimension) |
| output = values + self.output(contextual) |
| output = output + self.ffn(self.ffn_norm(output)) |
| if padding_mask is not None: |
| output = output.masked_fill(padding_mask.unsqueeze(-1), 0.0) |
| return output |
|
|
|
|
| class RoPETransformer(nn.Module): |
| """BGC-MAP-inspired bidirectional RoPE encoder with attentive pooling.""" |
|
|
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") |
|
|
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| self.config = config |
| self.gene_projection = GeneProjection(config) |
| self.blocks = nn.ModuleList( |
| RotarySelfAttentionBlock(config) for _ in range(config.attention_blocks) |
| ) |
| self.pooling = PoolingByAttention(config) |
| self.output = nn.Linear(config.hidden_dimension, config.output_dimension) |
|
|
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): |
| values = self.gene_projection(gene_embeddings, relative_positions) |
| for block in self.blocks: |
| values = block(values, relative_positions, padding_mask) |
| return values |
|
|
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) |
| return F.normalize(self.output(self.pooling(genes, padding_mask)), p=2, dim=-1) |
|
|
|
|
| class CrossAttentionPool(nn.Module): |
| """Learned context queries cross-attend to a self-attended BGC sequence.""" |
|
|
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") |
|
|
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| self.config = config |
| d = config.hidden_dimension |
| self.gene_projection = GeneProjection(config) |
| layer = nn.TransformerEncoderLayer( |
| d, config.attention_heads, config.feedforward_dimension, config.dropout, |
| batch_first=True, norm_first=True, activation="gelu", |
| ) |
| self.encoder = nn.TransformerEncoder(layer, config.attention_blocks) |
| self.queries = nn.Parameter(torch.empty(2, d)) |
| nn.init.xavier_uniform_(self.queries) |
| self.cross = nn.MultiheadAttention(d, config.attention_heads, config.dropout, batch_first=True) |
| self.output = nn.Linear(d, config.output_dimension) |
|
|
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): |
| values = self.gene_projection(gene_embeddings, relative_positions) |
| return self.encoder(values, src_key_padding_mask=padding_mask) |
|
|
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) |
| queries = self.queries.unsqueeze(0).expand(genes.shape[0], -1, -1) |
| attended, _ = self.cross(queries, genes, genes, key_padding_mask=padding_mask) |
| return F.normalize(self.output(attended.mean(dim=1)), p=2, dim=-1) |
|
|
|
|
| class LocalGlobalEncoder(nn.Module): |
| """PST-inspired local adjacency message passing followed by global pooling.""" |
|
|
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") |
|
|
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| self.config = config |
| d = config.hidden_dimension |
| self.gene_projection = GeneProjection(config) |
| self.local = nn.Conv1d(d, d, kernel_size=3, padding=1, groups=1) |
| self.local_gate = nn.Sequential(nn.Linear(d, d), nn.Sigmoid()) |
| self.global_block = InducedSelfAttentionBlock(config) |
| self.pooling = PoolingByAttention(config) |
| self.output = nn.Linear(d, config.output_dimension) |
|
|
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): |
| values = self.gene_projection(gene_embeddings, relative_positions) |
| local = self.local(values.transpose(1, 2)).transpose(1, 2) |
| values = values + local * self.local_gate(values) |
| return self.global_block(values, padding_mask) |
|
|
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) |
| return F.normalize(self.output(self.pooling(genes, padding_mask)), p=2, dim=-1) |
|
|
|
|
| class DilatedCNN(nn.Module): |
| """BiGCARP-inspired residual dilated convolutional encoder.""" |
|
|
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") |
|
|
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| self.config = config |
| d = config.hidden_dimension |
| self.gene_projection = GeneProjection(config) |
| self.blocks = nn.ModuleList( |
| nn.Sequential( |
| nn.Conv1d(d, d, 3, padding=dilation, dilation=dilation), |
| nn.GELU(), nn.Dropout(config.dropout), nn.Conv1d(d, d, 1), |
| ) for dilation in (1, 2, 4, 8, 16) |
| ) |
| self.pooling = PoolingByAttention(config) |
| self.output = nn.Linear(d, config.output_dimension) |
|
|
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): |
| values = self.gene_projection(gene_embeddings, relative_positions) |
| for block in self.blocks: |
| updated = block(values.transpose(1, 2)).transpose(1, 2) |
| values = values + updated |
| if padding_mask is not None: |
| values = values.masked_fill(padding_mask.unsqueeze(-1), 0.0) |
| return values |
|
|
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) |
| return F.normalize(self.output(self.pooling(genes, padding_mask)), p=2, dim=-1) |
|
|
|
|
| class HierarchicalLocalGlobal(nn.Module): |
| """Multi-scale local convolutions followed by a small global Transformer.""" |
|
|
| input_names = ("gene_embeddings", "relative_positions", "padding_mask") |
|
|
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| self.config = config |
| d = config.hidden_dimension |
| self.gene_projection = GeneProjection(config) |
| self.local3 = nn.Conv1d(d, d // 2, 3, padding=1) |
| self.local5 = nn.Conv1d(d, d // 2, 5, padding=2) |
| self.merge = nn.Linear(d, d) |
| layer = nn.TransformerEncoderLayer( |
| d, config.attention_heads, config.feedforward_dimension, config.dropout, |
| batch_first=True, norm_first=True, activation="gelu", |
| ) |
| self.global_encoder = nn.TransformerEncoder( |
| layer, max(1, config.attention_blocks // 2) |
| ) |
| self.pooling = PoolingByAttention(config) |
| self.output = nn.Linear(d, config.output_dimension) |
|
|
| def encode_genes(self, gene_embeddings, relative_positions, padding_mask=None): |
| values = self.gene_projection(gene_embeddings, relative_positions) |
| transposed = values.transpose(1, 2) |
| local = torch.cat((self.local3(transposed), self.local5(transposed)), dim=1).transpose(1, 2) |
| values = values + self.merge(local) |
| return self.global_encoder(values, src_key_padding_mask=padding_mask) |
|
|
| def forward(self, gene_embeddings, relative_positions, padding_mask=None): |
| genes = self.encode_genes(gene_embeddings, relative_positions, padding_mask) |
| return F.normalize(self.output(self.pooling(genes, padding_mask)), p=2, dim=-1) |
|
|
|
|
| class PfamAugmentedSetNet(nn.Module): |
| """SetNet whose gene projection is conditioned on the BGC Pfam inventory.""" |
|
|
| input_names = ("gene_embeddings", "relative_positions", "padding_mask", "pfam_tokens") |
| uses_pfam = True |
|
|
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| self.config = config |
| self.gene_projection = PfamGeneProjection(config) |
| self.blocks = nn.ModuleList( |
| InducedSelfAttentionBlock(config) for _ in range(config.attention_blocks) |
| ) |
| self.pooling = PoolingByAttention(config) |
| self.output = nn.Linear(config.hidden_dimension, config.output_dimension) |
|
|
| def encode_genes( |
| self, |
| gene_embeddings: torch.Tensor, |
| relative_positions: torch.Tensor, |
| padding_mask: torch.Tensor | None = None, |
| pfam_tokens: torch.Tensor | None = None, |
| ) -> torch.Tensor: |
| if pfam_tokens is None: |
| raise ValueError("Pfam-augmented SetNet requires pfam_tokens") |
| values = self.gene_projection(gene_embeddings, relative_positions, pfam_tokens) |
| for block in self.blocks: |
| values = block(values, padding_mask) |
| return values |
|
|
| def forward( |
| self, |
| gene_embeddings: torch.Tensor, |
| relative_positions: torch.Tensor, |
| padding_mask: torch.Tensor | None = None, |
| pfam_tokens: torch.Tensor | None = None, |
| ) -> torch.Tensor: |
| genes = self.encode_genes( |
| gene_embeddings, relative_positions, padding_mask, pfam_tokens |
| ) |
| pooled = self.pooling(genes, padding_mask) |
| return F.normalize(self.output(pooled), p=2, dim=-1) |
|
|
|
|
| class WeightedPfamJaccard(nn.Module): |
| """Learn nonnegative Pfam importance weights for differentiable set Jaccard.""" |
|
|
| input_names = ("pfam_tokens",) |
| uses_pfam = True |
|
|
| def __init__(self, config: ModelConfig) -> None: |
| super().__init__() |
| if config.pfam_vocab_size < 3: |
| raise ValueError("Pfam vocabulary must contain padding, unknown, and one domain token") |
| initial = float(torch.log(torch.expm1(torch.tensor(1.0)))) |
| self.raw_weights = nn.Parameter( |
| torch.full((config.pfam_vocab_size,), initial, dtype=torch.float32) |
| ) |
| with torch.no_grad(): |
| self.raw_weights[0] = -20.0 |
|
|
| def domain_weights(self) -> torch.Tensor: |
| positive = F.softplus(self.raw_weights) |
| return torch.cat((positive[:1] * 0.0, positive[1:])) |
|
|
| def pairwise_jaccard(self, pfam_tokens: torch.Tensor) -> torch.Tensor: |
| if pfam_tokens.ndim != 2: |
| raise ValueError("Pfam tokens must have shape [batch, domains]") |
| batch_size = pfam_tokens.shape[0] |
| vocabulary = self.raw_weights.shape[0] |
| presence = torch.zeros( |
| batch_size, vocabulary, device=pfam_tokens.device, dtype=torch.float32 |
| ) |
| presence.scatter_(1, pfam_tokens.clamp_min(0), 1.0) |
| presence[:, 0] = 0.0 |
| weighted = presence * self.domain_weights().to(pfam_tokens.device) |
| totals = weighted.sum(dim=1) |
| intersection = weighted @ presence.T |
| union = totals[:, None] + totals[None, :] - intersection |
| return intersection / union.clamp_min(1e-8) |
|
|
| def forward(self, pfam_tokens: torch.Tensor) -> torch.Tensor: |
| return self.domain_weights() |
|
|
|
|
| ARCHITECTURES = { |
| "setnet": LeakageFreeBGCSetNet, |
| "weighted_pfam_jaccard": WeightedPfamJaccard, |
| "pfam_setnet": PfamAugmentedSetNet, |
| "gated_deepsets": GatedDeepSets, |
| "rope_transformer": RoPETransformer, |
| "cross_attention": CrossAttentionPool, |
| "local_global": LocalGlobalEncoder, |
| "dilated_cnn": DilatedCNN, |
| "hierarchical": HierarchicalLocalGlobal, |
| } |
|
|
|
|
| def build_model(config: ModelConfig) -> nn.Module: |
| try: |
| model_class = ARCHITECTURES[config.architecture] |
| except KeyError as exc: |
| raise ValueError(f"Unknown architecture: {config.architecture}") from exc |
| return model_class(config) |
|
|