aaronfeller's picture
Publish minimal StabilityArc inference model
4691fb1 verified
Raw History Blame Contribute Delete
7.6 kB
"""Shared RoPE stability decoder and symmetric pair residual."""
from __future__ import annotations
import math
from dataclasses import dataclass
from pathlib import Path
from typing import Sequence
import torch
from torch import nn
from torch.nn import functional as F
AMINO_ACIDS = "ACDEFGHIKLMNPQRSTVWY"
AMINO_ACID_TO_INDEX = {residue: index for index, residue in enumerate(AMINO_ACIDS)}
ESMC_EMBEDDING_DIM = 1152
@dataclass
class MutationTask:
wildtype_embedding: torch.Tensor
positions: torch.Tensor
target_residues: torch.Tensor
class RotarySelfAttention(nn.Module):
def __init__(self, model_dim: int, num_heads: int, dropout: float) -> None:
super().__init__()
self.num_heads = num_heads
self.head_dim = model_dim // num_heads
self.dropout = dropout
self.query_key_value = nn.Linear(model_dim, model_dim * 3)
self.output = nn.Linear(model_dim, model_dim)
def _rotate(self, values: torch.Tensor) -> torch.Tensor:
first, second = values.chunk(2, dim=-1)
return torch.cat([-second, first], dim=-1)
def _apply_rope(self, values: torch.Tensor) -> torch.Tensor:
length = values.size(-2)
positions = torch.arange(length, device=values.device, dtype=torch.float32)
frequencies = torch.exp(-math.log(10000) * torch.arange(0, self.head_dim, 2, device=values.device, dtype=torch.float32) / self.head_dim)
angles = positions[:, None] * frequencies[None, :]
cosine = torch.cat([angles.cos(), angles.cos()], dim=-1).to(values.dtype)[None, None, :, :]
sine = torch.cat([angles.sin(), angles.sin()], dim=-1).to(values.dtype)[None, None, :, :]
return values * cosine + self._rotate(values) * sine
def forward(self, tokens: torch.Tensor) -> torch.Tensor:
batch, length, _ = tokens.shape
query, key, value = self.query_key_value(tokens).chunk(3, dim=-1)
shape = (batch, length, self.num_heads, self.head_dim)
query = self._apply_rope(query.reshape(shape).transpose(1, 2))
key = self._apply_rope(key.reshape(shape).transpose(1, 2))
value = value.reshape(shape).transpose(1, 2)
attended = F.scaled_dot_product_attention(query, key, value, dropout_p=self.dropout if self.training else 0.0)
return self.output(attended.transpose(1, 2).reshape(batch, length, -1))
class RotaryTransformerEncoderLayer(nn.Module):
def __init__(self, model_dim: int, num_heads: int, dropout: float) -> None:
super().__init__()
self.attention_norm = nn.LayerNorm(model_dim)
self.attention = RotarySelfAttention(model_dim, num_heads, dropout)
self.feedforward_norm = nn.LayerNorm(model_dim)
self.feedforward = nn.Sequential(
nn.Linear(model_dim, model_dim * 4),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(model_dim * 4, model_dim),
)
self.dropout = nn.Dropout(dropout)
def forward(self, tokens: torch.Tensor) -> torch.Tensor:
tokens = tokens + self.dropout(self.attention(self.attention_norm(tokens)))
return tokens + self.dropout(self.feedforward(self.feedforward_norm(tokens)))
class StabilityBackbone(nn.Module):
def __init__(self, model_dim: int, depth: int, num_heads: int) -> None:
super().__init__()
self.input_proj = nn.Linear(ESMC_EMBEDDING_DIM, model_dim)
self.encoder = nn.ModuleList(RotaryTransformerEncoderLayer(model_dim, num_heads, 0.1) for _ in range(depth))
self.norm = nn.LayerNorm(model_dim)
self.score_head = nn.Linear(model_dim, len(AMINO_ACIDS))
def forward(self, embedding: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
hidden = self.input_proj(embedding)
for layer in self.encoder:
hidden = layer(hidden)
hidden = self.norm(hidden)
return hidden, self.score_head(hidden)
class StabilityArc(nn.Module):
def __init__(self, model_dim: int = 256, depth: int = 4, num_heads: int = 8) -> None:
super().__init__()
self.backbone = StabilityBackbone(model_dim, depth, num_heads)
self.site_projection = nn.Sequential(nn.Linear(model_dim, model_dim), nn.GELU(), nn.Linear(model_dim, model_dim))
self.pair_residue = nn.Embedding(len(AMINO_ACIDS), model_dim)
self.pair_head = nn.Sequential(nn.Linear(model_dim * 3 + 2, model_dim), nn.GELU(), nn.Linear(model_dim, 1))
def score_task_components(self, task: MutationTask, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
hidden, scores = self.backbone(task.wildtype_embedding.to(device).unsqueeze(0))
site = self.site_projection(hidden.squeeze(0))
return scores.squeeze(0), site
def score_edit_sets(self, scores: torch.Tensor, site: torch.Tensor, contact_map: torch.Tensor, edit_sets: Sequence[Sequence[tuple[int, int]]], device: torch.device) -> torch.Tensor:
direct = []
pair_positions_left: list[int] = []
pair_positions_right: list[int] = []
pair_residues_left: list[int] = []
pair_residues_right: list[int] = []
pair_groups: list[int] = []
for group, edits in enumerate(edit_sets):
direct.append(scores[[position for position, _ in edits], [residue for _, residue in edits]].sum())
for left, (left_position, left_residue) in enumerate(edits):
for right_position, right_residue in edits[left + 1:]:
pair_positions_left.append(left_position)
pair_positions_right.append(right_position)
pair_residues_left.append(left_residue)
pair_residues_right.append(right_residue)
pair_groups.append(group)
predictions = torch.stack(direct)
if not pair_groups:
return predictions
left = site[torch.tensor(pair_positions_left, dtype=torch.long, device=device)] + self.pair_residue(torch.tensor(pair_residues_left, dtype=torch.long, device=device))
right = site[torch.tensor(pair_positions_right, dtype=torch.long, device=device)] + self.pair_residue(torch.tensor(pair_residues_right, dtype=torch.long, device=device))
distance = (torch.tensor(pair_positions_left, dtype=torch.float32, device=device) - torch.tensor(pair_positions_right, dtype=torch.float32, device=device)).abs().unsqueeze(-1) / max(site.size(0) - 1, 1)
contact = contact_map.to(device)[
torch.tensor(pair_positions_left, dtype=torch.long, device=device),
torch.tensor(pair_positions_right, dtype=torch.long, device=device),
].unsqueeze(-1)
features = torch.cat([left + right, (left - right).abs(), left * right, distance, contact], dim=-1)
residuals = self.pair_head(features).squeeze(-1)
return predictions + predictions.new_zeros(len(edit_sets)).scatter_add_(0, torch.tensor(pair_groups, dtype=torch.long, device=device), residuals)
def load_stabilityarc(checkpoint: Path, device: str = "cpu") -> tuple[StabilityArc, list[str]]:
payload = torch.load(checkpoint, map_location="cpu", weights_only=True)
config = payload["config"]
if config["position_encoding"] != "rope" or config["pair_residual"] != "contact-aware":
raise ValueError("Checkpoint must be a direct RoPE stability scorer with a contact-aware pair residual")
model = StabilityArc(config["model_dim"], config["depth"], config["num_heads"])
model.load_state_dict(payload["state_dict"], strict=True)
return model.to(device).eval(), list(payload["training_proteins"])