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