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