From 94019274e1c0ff2d24af35fdeda9265b0996d04a Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Thu, 26 Feb 2026 16:27:30 -0600 Subject: [PATCH 01/19] creation of the ESM encoder wrapper, accounting for cached edge attributes, using constants, fixing tests --- src/__init__.py | 8 +- src/encoder_base.py | 20 +++-- src/esm_encoder.py | 84 +++++++++++++++++++ src/gvp_encoder.py | 191 +++++++++++++++++++++++++++++++++--------- src/slae_encoder.py | 10 +-- tests/test_encoder.py | 181 +++++++++++++++++++++++++++++---------- 6 files changed, 393 insertions(+), 101 deletions(-) create mode 100644 src/esm_encoder.py diff --git a/src/__init__.py b/src/__init__.py index 43d4a48..d7a8a58 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -6,5 +6,9 @@ Importing this module triggers encoder registration. """ -from src import gvp_encoder, slae -from src.encoder_base import BaseProteinEncoder, build_encoder, register_encoder +from src import esm_encoder, gvp_encoder, slae_encoder +from src.encoder_base import ( + BaseProteinEncoder as BaseProteinEncoder, + build_encoder as build_encoder, + register_encoder as register_encoder, +) diff --git a/src/encoder_base.py b/src/encoder_base.py index 115272c..3496584 100644 --- a/src/encoder_base.py +++ b/src/encoder_base.py @@ -5,6 +5,7 @@ - BaseProteinEncoder: Abstract base class that all encoders must implement - Registry pattern: Decorator-based registration and build_encoder() function """ + from __future__ import annotations from abc import ABC, abstractmethod @@ -16,7 +17,6 @@ if TYPE_CHECKING: from torch_geometric.data import HeteroData - class BaseProteinEncoder(ABC, nn.Module): """ Abstract base class for protein encoders. @@ -38,7 +38,9 @@ def encoder_type(self) -> str: pass @abstractmethod - def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: + def forward( + self, data: HeteroData + ) -> tuple[torch.Tensor, torch.Tensor, tuple | None]: """ Encode protein data. @@ -46,9 +48,11 @@ def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: data: HeteroData with protein nodes Returns: - tuple of (s, V) where: + tuple of (s, V, pp_edge_attr) where: s: (N, scalar_dim) scalar features V: (N, vector_dim, 3) vector features + pp_edge_attr: tuple (s_edge, V_edge) for PP edges, or None if encoder + doesn't process edges (e.g., SLAE, ESM) """ pass @@ -67,11 +71,9 @@ def from_config(cls, config: dict, device: torch.device) -> BaseProteinEncoder: """ pass - # global encoder registry _ENCODER_REGISTRY: dict[str, BaseProteinEncoder] = {} - def register_encoder(name: str): """ Decorator to register an encoder class. @@ -81,13 +83,14 @@ def register_encoder(name: str): class MyEncoder(BaseProteinEncoder): ... """ + def decorator(cls: BaseProteinEncoder) -> BaseProteinEncoder: if name in _ENCODER_REGISTRY: raise ValueError(f"Encoder '{name}' is already registered") _ENCODER_REGISTRY[name] = cls return cls - return decorator + return decorator def get_encoder_class(name: str) -> BaseProteinEncoder: """ @@ -107,7 +110,6 @@ def get_encoder_class(name: str) -> BaseProteinEncoder: raise KeyError(f"Unknown encoder type '{name}'. Available: {available}") return _ENCODER_REGISTRY[name] - def build_encoder(config: dict, device: torch.device) -> BaseProteinEncoder: """ Build encoder from configuration dict. @@ -121,8 +123,8 @@ def build_encoder(config: dict, device: torch.device) -> BaseProteinEncoder: Returns: Instantiated encoder implementing BaseProteinEncoder """ - if 'encoder_type' not in config: + if "encoder_type" not in config: raise ValueError("'encoder_type' must be specified in config") - encoder_type = config['encoder_type'] + encoder_type = config["encoder_type"] encoder_cls = get_encoder_class(encoder_type) return encoder_cls.from_config(config, device) diff --git a/src/esm_encoder.py b/src/esm_encoder.py new file mode 100644 index 0000000..0707aa4 --- /dev/null +++ b/src/esm_encoder.py @@ -0,0 +1,84 @@ +# esm_encoder.py +""" +ESM embeddings wrapper. + +This encoder reads pre-computed ESM3 embeddings from data and returns +them directly as scalar features with zero vector channels. Downstream +GVP message-passing layers (including protein-protein edges) provide all +geometric processing. +""" + +from __future__ import annotations + +import torch +from torch_geometric.data import HeteroData + +from src.encoder_base import BaseProteinEncoder, register_encoder + +@register_encoder("esm") +class ESMEncoder(BaseProteinEncoder): + """ + ESM encoder that reads cached embeddings from data. + + Returns (esm_embedding, empty_vectors) with output_dims = (esm_dim, 0). + No learnable parameters — all geometric processing happens in the + downstream ProteinWaterUpdate layers (pp, wp, pw, ww edges). + """ + + def __init__(self, esm_dim: int = 1536): + """ + Initialize ESMEncoder. + + Args: + esm_dim: Dimension of ESM embeddings (default: 1536 for ESM3-open) + """ + super().__init__() + self._esm_dim = esm_dim + + @property + def output_dims(self) -> tuple[int, int]: + """Return (esm_dim, 0) — scalars only.""" + return self._esm_dim, 0 + + @property + def encoder_type(self) -> str: + return "esm" + + def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, None]: + """ + Read cached ESM embeddings and return (s, V, None). + + Args: + data: HeteroData with data['protein'].esm_embedding + + Returns: + s: (N, esm_dim) — raw ESM embeddings + V: (N, 0, 3) — empty vector features + pp_edge_attr: None — ESM doesn't process edges + """ + if "esm_embedding" not in data["protein"]: + raise NotImplementedError( + "ESM encoder requires cached embeddings. " + "Please provide pre-computed esm_embedding in data['protein']. " + "Run scripts/generate_esm_embeddings.py first." + ) + + embeddings = data["protein"].esm_embedding + V = embeddings.new_empty(embeddings.size(0), 0, 3) + return embeddings, V, None + + @classmethod + def from_config(cls, config: dict, device: torch.device) -> ESMEncoder: + """ + Construct ESMEncoder from config dict. + + Args: + config: Configuration dictionary with: + - esm_dim: ESM embedding dimension (default: 1536) + device: Device to place the encoder on + + Returns: + Instantiated ESMEncoder + """ + esm_dim = config.get("esm_dim", 1536) + return cls(esm_dim=esm_dim).to(device) diff --git a/src/gvp_encoder.py b/src/gvp_encoder.py index a477947..0c11974 100644 --- a/src/gvp_encoder.py +++ b/src/gvp_encoder.py @@ -17,14 +17,24 @@ from torch_geometric.data import Batch, Data, HeteroData from torch_scatter import scatter_add, scatter_max, scatter_mean -from src.constants import EDGE_PP +from src.constants import EDGE_PP, NODE_FEATURE_DIM, NUM_RBF, RBF_CUTOFF from src.encoder_base import BaseProteinEncoder, register_encoder from src.gvp import GVP, EdgeUpdate, GVPConvLayer -from src.utils import rbf as _rbf +from src.utils import rbf def _edge_vectors(pos: torch.Tensor, edge_index: torch.Tensor): - """Compute edge vectors and distances.""" + """ + Compute edge vectors and distances from node positions. + + Args: + pos: (N, 3) node position coordinates + edge_index: (2, E) edge indices with source in row 0, destination in row 1 + + Returns: + rij: (E,) edge distances clamped to minimum 1e-4 + r_hat: (E, 3) unit vectors pointing from source to destination + """ src, dst = edge_index[0], edge_index[1] vec = pos[dst] - pos[src] rij = torch.linalg.norm(vec, dim=-1).clamp(min=1e-4) @@ -37,13 +47,13 @@ def make_encoder_data(data: HeteroData) -> Data: Build a homogeneous Data with protein nodes for GVP encoder. Extracts protein subgraph from HeteroData for use with GVP encoder. - Edge features are computed by the encoder itself. + If PP edge features are cached in the HeteroData, they are copied through. Args: data: HeteroData with protein nodes Returns: - enc_data: Data with x, pos, edge_index + enc_data: Data with x, pos, edge_index, and optionally cached edge features """ device = data['protein'].pos.device prot = data['protein'] @@ -51,7 +61,7 @@ def make_encoder_data(data: HeteroData) -> Data: x = prot.x pos = prot.pos - # protein-protein edges (topology only - features computed by encoder) + # protein-protein edges if EDGE_PP in data.edge_types: edge_index = data[EDGE_PP].edge_index else: @@ -63,6 +73,14 @@ def make_encoder_data(data: HeteroData) -> Data: edge_index=edge_index, ) + # Copy cached edge features if available + if EDGE_PP in data.edge_types: + pp_edge = data[EDGE_PP] + if hasattr(pp_edge, 'edge_rbf'): + enc_data.edge_rbf = pp_edge.edge_rbf + if hasattr(pp_edge, 'edge_unit'): + enc_data.edge_unit = pp_edge.edge_unit + # batch for multi-complex batches if hasattr(prot, "batch"): enc_data.batch = prot.batch @@ -80,12 +98,12 @@ class ProteinGVPEncoder(nn.Module): def __init__( self, - node_scalar_in: int = 17, + node_scalar_in: int = NODE_FEATURE_DIM, node_vec_in: int = 1, hidden_dims: tuple[int, int] = (256, 32), - edge_scalar_in: int = 16, + edge_scalar_in: int = NUM_RBF, edge_vec_in: int = 1, - edge_scalar_out: int = 16, + edge_scalar_out: int = NUM_RBF, n_layers: int = 3, n_message: int = 2, n_feedforward: int = 2, @@ -99,9 +117,37 @@ def __init__( pool_aggr: Literal["mean", "sum", "max"] = "mean", update_w_distance: bool = True, distance_dim: int | None = None, - radius: float = 8.0, - num_edge_rbf: int = 16, + radius: float = RBF_CUTOFF, + num_edge_rbf: int = NUM_RBF, + use_edge_update: bool = True, ): + """ + Initialize GVP encoder for protein structure processing. + + Args: + node_scalar_in: Input scalar feature dimension (e.g., element one-hot) + node_vec_in: Input vector feature channels (typically 1 for orientation) + hidden_dims: (scalar_dim, vector_dim) hidden layer dimensions + edge_scalar_in: Input edge scalar dimension (RBF features) + edge_vec_in: Input edge vector channels (unit vectors) + edge_scalar_out: Output edge scalar dimension + n_layers: Number of GVP convolution layers + n_message: Number of GVPs in message function + n_feedforward: Number of GVPs in feedforward function + drop_rate: Dropout rate for regularization + vector_gate: Whether to use vector gating in GVP layers + scalar_activation: Activation function for scalar channels + vector_activation: Activation function for vector gating + init_vec_zero: If True, initialize input vectors as zeros + pooled_dim: Output dimension when pooling by residue + pool_residue: If True, pool atom features to residue level + pool_aggr: Aggregation method for residue pooling ('mean', 'sum', 'max') + update_w_distance: Include distance features in edge updates + distance_dim: Dimension for distance conditioning, defaults to edge_scalar_in + radius: Distance cutoff in Angstroms for RBF encoding + num_edge_rbf: Number of RBF basis functions + use_edge_update: Whether to update edge features between layers + """ super().__init__() self.node_scalar_in = node_scalar_in self.node_vec_in = node_vec_in @@ -119,6 +165,7 @@ def __init__( self.radius = radius self.num_edge_rbf = num_edge_rbf self.pooled_dim = pooled_dim + self.use_edge_update = use_edge_update distance_dim = distance_dim or edge_scalar_in self.distance_dim = distance_dim @@ -161,12 +208,15 @@ def __init__( for _ in range(n_layers) ]) - self.edge_update = EdgeUpdate( - n_node_scalars=S_hid, - s_edge_width=self.s_edge_width, - update_w_distance=update_w_distance, - distance_dim=distance_dim, - ) + if use_edge_update: + self.edge_update = EdgeUpdate( + n_node_scalars=S_hid, + s_edge_width=self.s_edge_width, + update_w_distance=update_w_distance, + distance_dim=distance_dim, + ) + else: + self.edge_update = None self.atom_readout = nn.Sequential( nn.Linear(S_hid + V_hid, pooled_dim), @@ -178,6 +228,15 @@ def __init__( @staticmethod def _tuple_to_scalar_dense(x_tuple: tuple) -> torch.Tensor: + """ + Convert GVP tuple to dense scalar representation. + + Args: + x_tuple: (s, V) where s is (N, scalar_dim) and V is (N, vector_dim, 3) + + Returns: + (N, scalar_dim + vector_dim) concatenation of scalars and vector norms + """ s, V = x_tuple vnorm = torch.linalg.norm(V, dim=-1) return torch.cat([s, vnorm], dim=-1) @@ -185,29 +244,73 @@ def _tuple_to_scalar_dense(x_tuple: tuple) -> torch.Tensor: def _pool_by_residue( self, atom_embed: torch.Tensor, residue_index: torch.Tensor, num_residues: int ) -> torch.Tensor: + """ + Pool atom-level embeddings to residue level. + + Args: + atom_embed: (N_atoms, embed_dim) atom embeddings + residue_index: (N_atoms,) residue index per atom + num_residues: Total number of residues + + Returns: + (num_residues, embed_dim) pooled residue embeddings + """ aggr = self.pool_aggr if aggr == "mean": return scatter_mean(atom_embed, residue_index, dim=0, dim_size=num_residues) - if aggr == "sum": + elif aggr == "sum": return scatter_add(atom_embed, residue_index, dim=0, dim_size=num_residues) - if aggr == "max": + elif aggr == "max": out, _ = scatter_max(atom_embed, residue_index, dim=0, dim_size=num_residues) return out - raise ValueError(f"Unknown pool_aggr={aggr!r}") + else: + raise ValueError(f"Unknown pool_aggr={aggr!r}") @staticmethod - def _initial_node_tuple(x_scalar: torch.Tensor, device=None) -> tuple: + def _initial_node_tuple( + x_scalar: torch.Tensor, device: torch.device | None = None + ) -> tuple[torch.Tensor, torch.Tensor]: zeros = torch.zeros(x_scalar.size(0), 1, 3, device=x_scalar.device if device is None else device) return (x_scalar, zeros) - def _compute_edge_attr(self, pos: torch.Tensor, edge_index: torch.Tensor): - d, u = _edge_vectors(pos, edge_index) - s_edge_raw = _rbf(d, num_gaussians=self.num_edge_rbf, cutoff=self.radius) + def _compute_edge_attr(self, data: Batch): + """ + Build edge attributes from positions or cached features. + + If cached edge features (edge_rbf, edge_unit) are available in data, + use them directly. Otherwise, compute from positions. + + Args: + data: Batch with pos, edge_index, and optionally edge_rbf, edge_unit + + Returns: + (s_edge, V_edge): Tuple of scalar and vector edge features + s_edge_raw: Raw RBF features (for distance conditioning) + """ + # Use cached features if available + if hasattr(data, 'edge_rbf') and hasattr(data, 'edge_unit'): + s_edge_raw = data.edge_rbf + u = data.edge_unit + else: + # Fallback: compute from positions + d, u = _edge_vectors(data.pos, data.edge_index) + s_edge_raw = rbf(d, num_gaussians=self.num_edge_rbf, cutoff=self.radius) + s_edge = self.edge_in_proj(s_edge_raw) V_edge = u.unsqueeze(1) return (s_edge, V_edge), s_edge_raw - def forward(self, data: Batch): + def forward(self, data: Batch) -> tuple[tuple, tuple | None]: + """ + Forward pass through the GVP encoder. + + Args: + data: Batch with node features, positions, and edge indices + + Returns: + x: tuple (s, V) of node scalar and vector features + edge_attr: tuple (s_edge, V_edge) of edge features, or None if use_edge_update=False + """ x_scalar = self.input_scalar_encoder(data.x) if self.init_vec_zero or not hasattr(data, "node_v"): @@ -216,26 +319,29 @@ def forward(self, data: Batch): node_features = (x_scalar, data.node_v.unsqueeze(1)) x = self.input_gvp(node_features) - edge_attr, dist_feat = self._compute_edge_attr(data.pos, data.edge_index) + edge_attr, dist_feat = self._compute_edge_attr(data) for layer in self.layers: x = layer(x, data.edge_index, edge_attr) - edge_attr = self.edge_update( - node_tuple=x, - edge_index=data.edge_index, - edge_attr=edge_attr, - distance_feat=(dist_feat if self.update_w_distance else None), - ) + if self.edge_update is not None: + edge_attr = self.edge_update( + node_tuple=x, + edge_index=data.edge_index, + edge_attr=edge_attr, + distance_feat=(dist_feat if self.update_w_distance else None), + ) if self.pool_residue: - assert hasattr(data, "residue_index") and hasattr(data, "num_residues"), \ - "Pooling requires data.residue_index (N,) and data.num_residues (int)." + if not (hasattr(data, "residue_index") and hasattr(data, "num_residues")): + raise ValueError("Pooling requires data.residue_index and data.num_residues") atom_dense = self._tuple_to_scalar_dense(x) atom_embed = self.atom_readout(atom_dense) res_embed = self._pool_by_residue(atom_embed, data.residue_index, int(data.num_residues)) - return res_embed + return res_embed, None # No edge features when pooling - return x + # Return edge_attr only if edge_update was used + final_edge_attr = edge_attr if self.edge_update is not None else None + return x, final_edge_attr def load_encoder_from_checkpoint( @@ -360,7 +466,7 @@ def encoder_type(self) -> str: """Return encoder type identifier.""" return 'gvp' - def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: + def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, tuple | None]: """ Encode protein data. @@ -368,15 +474,17 @@ def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: data: HeteroData with protein nodes Returns: - Tuple of (s, V) features + s: (N, scalar_dim) scalar features + V: (N, vector_dim, 3) vector features + pp_edge_attr: tuple (s_edge, V_edge) for PP edges, or None if edge updates disabled """ # Convert HeteroData to homogeneous Data for GVP encoder enc_data = make_encoder_data(data) with torch.set_grad_enabled(not self._freeze): - s, V = self.encoder(enc_data) + (s, V), edge_attr = self.encoder(enc_data) - return s, V + return s, V, edge_attr @classmethod def from_config(cls, config: dict, device: torch.device) -> GVPEncoder: @@ -389,6 +497,7 @@ def from_config(cls, config: dict, device: torch.device) -> GVPEncoder: - node_scalar_in: Input feature dimension (default: 16) - hidden_s, hidden_v: Hidden dimensions - freeze_encoder: Whether to freeze encoder + - use_edge_update: Whether to use edge updates (default: True) device: Device to place the encoder on Returns: @@ -399,6 +508,7 @@ def from_config(cls, config: dict, device: torch.device) -> GVPEncoder: hidden_s = config.get('hidden_s', 256) hidden_v = config.get('hidden_v', 32) freeze = config.get('freeze_encoder', False) + use_edge_update = config.get('use_edge_update', True) if encoder_ckpt: encoder, _ = load_encoder_from_checkpoint( @@ -411,6 +521,7 @@ def from_config(cls, config: dict, device: torch.device) -> GVPEncoder: node_scalar_in=node_scalar_in, hidden_dims=(hidden_s, hidden_v), edge_scalar_in=16, + use_edge_update=use_edge_update, ).to(device) return cls(encoder=encoder, freeze=freeze) diff --git a/src/slae_encoder.py b/src/slae_encoder.py index 33ee84f..6b6b9bb 100644 --- a/src/slae_encoder.py +++ b/src/slae_encoder.py @@ -13,7 +13,6 @@ from src.encoder_base import BaseProteinEncoder, register_encoder - @register_encoder('slae') class SLAEEncoder(BaseProteinEncoder): """ @@ -37,9 +36,9 @@ def output_dims(self) -> tuple[int, int]: def encoder_type(self) -> str: return 'slae' - def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: + def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, None]: """ - Read cached SLAE embeddings and return (s, V). + Read cached SLAE embeddings and return (s, V, None). Args: data: HeteroData with data['protein'].slae_embedding @@ -47,17 +46,18 @@ def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: Returns: s: (N, slae_dim) — raw SLAE embeddings V: (N, 0, 3) — empty vector features + pp_edge_attr: None — SLAE doesn't process edges """ if 'slae_embedding' not in data['protein']: raise NotImplementedError( "SLAE encoder requires cached embeddings. " "Please provide pre-computed slae_embedding in data['protein']. " - "Run scripts/precompute_slae_embeddings.py first." + "Run scripts/generate_slae_embeddings.py first." ) embeddings = data['protein'].slae_embedding V = embeddings.new_empty(embeddings.size(0), 0, 3) - return embeddings, V + return embeddings, V, None @classmethod def from_config(cls, config: dict, device: torch.device) -> SLAEEncoder: diff --git a/tests/test_encoder.py b/tests/test_encoder.py index f48531b..185bf8a 100644 --- a/tests/test_encoder.py +++ b/tests/test_encoder.py @@ -15,8 +15,8 @@ from torch_geometric.data import Data, HeteroData from src.encoder_base import build_encoder, get_encoder_class -from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder, make_encoder_data -from src.slae import SLAEEncoder +from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder +from src.slae_encoder import SLAEEncoder # ============== Fixtures ============== @@ -113,8 +113,8 @@ def test_slae_registered(self): assert cls is SLAEEncoder def test_unknown_encoder_raises(self): - """Unknown encoder type should raise ValueError.""" - with pytest.raises(ValueError, match="Unknown encoder type"): + """Unknown encoder type should raise KeyError.""" + with pytest.raises(KeyError, match="Unknown encoder type"): get_encoder_class('nonexistent_encoder') def test_build_encoder_gvp(self, device): @@ -143,12 +143,6 @@ def test_build_encoder_slae(self, device): assert encoder.encoder_type == 'slae' assert encoder.output_dims == (128, 0) - def test_build_encoder_default_gvp(self, device): - """build_encoder should default to GVP if encoder_type not specified.""" - config = {'node_scalar_in': 16} - encoder = build_encoder(config, device) - assert isinstance(encoder, GVPEncoder) - # ============== Base Interface Tests ============== @@ -172,11 +166,13 @@ def test_gvp_implements_interface(self, device, sample_hetero_data): assert len(encoder.output_dims) == 2 assert isinstance(encoder.encoder_type, str) - # Check forward returns (s, V) tuple - s, V = encoder(sample_hetero_data) + # Check forward returns (s, V, pp_edge_attr) tuple + s, V, pp_edge_attr = encoder(sample_hetero_data) assert s.shape[0] == sample_hetero_data['protein'].num_nodes assert V.shape[0] == sample_hetero_data['protein'].num_nodes assert V.shape[2] == 3 + # GVP encoder should return edge features + assert pp_edge_attr is not None or encoder.encoder.edge_update is None def test_slae_implements_interface(self, device, sample_hetero_data_with_slae): """SLAEEncoder should implement all required interface methods.""" @@ -188,11 +184,13 @@ def test_slae_implements_interface(self, device, sample_hetero_data_with_slae): assert encoder.output_dims == (128, 0) assert isinstance(encoder.encoder_type, str) - # Check forward returns (s, V) tuple - s, V = encoder(sample_hetero_data_with_slae) + # Check forward returns (s, V, pp_edge_attr) tuple + s, V, pp_edge_attr = encoder(sample_hetero_data_with_slae) assert s.shape[0] == sample_hetero_data_with_slae['protein'].num_nodes assert s.shape[1] == 128 assert V.shape == (sample_hetero_data_with_slae['protein'].num_nodes, 0, 3) + # SLAE encoder should return None for edge features + assert pp_edge_attr is None def test_from_config_class_method(self, device): """Both encoders should have from_config class method.""" @@ -232,11 +230,13 @@ def test_encoder_initialization(self, simple_encoder): def test_encoder_forward_with_pooling(self, simple_encoder, sample_homogeneous_data): """Test forward pass with residue pooling.""" - output = simple_encoder(sample_homogeneous_data) + output, edge_attr = simple_encoder(sample_homogeneous_data) assert output.shape == (sample_homogeneous_data.num_residues, simple_encoder.pooled_dim) + # Pooling mode returns None for edge features + assert edge_attr is None def test_encoder_forward_no_pooling(self, sample_homogeneous_data): - """Test encoder without residue pooling returns (s, V) tuple.""" + """Test encoder without residue pooling returns ((s, V), edge_attr) tuple.""" encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 16), @@ -246,12 +246,15 @@ def test_encoder_forward_no_pooling(self, sample_homogeneous_data): num_edge_rbf=16, ) - output = encoder(sample_homogeneous_data) + (s, v), edge_attr = encoder(sample_homogeneous_data) - assert isinstance(output, tuple) - s, v = output assert s.shape == (sample_homogeneous_data.num_nodes, 64) assert v.shape == (sample_homogeneous_data.num_nodes, 16, 3) + # edge_attr should be a tuple (s_edge, V_edge) + assert edge_attr is not None + s_edge, V_edge = edge_attr + assert s_edge.dim() == 2 + assert V_edge.dim() == 3 def test_encoder_empty_graph(self, simple_encoder): """Test encoder handles empty graphs gracefully.""" @@ -263,8 +266,10 @@ def test_encoder_empty_graph(self, simple_encoder): num_residues=0, ) - output = simple_encoder(data) + output, edge_attr = simple_encoder(data) assert output.shape == (0, simple_encoder.pooled_dim) + # Pooling mode returns None for edge features + assert edge_attr is None class TestGVPEncoderWrapper: @@ -297,7 +302,7 @@ def test_wrapper_encoder_type(self, device): assert encoder.encoder_type == 'gvp' def test_wrapper_forward(self, device, sample_hetero_data): - """Wrapper forward should return (s, V) from HeteroData.""" + """Wrapper forward should return (s, V, edge_attr) from HeteroData.""" base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 16), @@ -307,10 +312,15 @@ def test_wrapper_forward(self, device, sample_hetero_data): encoder = GVPEncoder(encoder=base_encoder, freeze=False) - s, V = encoder(sample_hetero_data) + s, V, edge_attr = encoder(sample_hetero_data) assert s.shape == (sample_hetero_data['protein'].num_nodes, 64) assert V.shape == (sample_hetero_data['protein'].num_nodes, 16, 3) + # edge_attr should be a tuple (s_edge, V_edge) when edge_update is enabled + assert edge_attr is not None + s_edge, V_edge = edge_attr + assert s_edge.dim() == 2 # (E, scalar_dim) + assert V_edge.dim() == 3 # (E, 1, 3) def test_wrapper_freeze(self, device): """Freeze parameter should disable gradients.""" @@ -327,26 +337,6 @@ def test_wrapper_freeze(self, device): assert not p.requires_grad -class TestMakeEncoderData: - """Tests for make_encoder_data helper function.""" - - def test_converts_hetero_to_homo(self, sample_hetero_data): - """Should convert HeteroData to homogeneous Data.""" - enc_data = make_encoder_data(sample_hetero_data) - - assert isinstance(enc_data, Data) - assert hasattr(enc_data, 'x') - assert hasattr(enc_data, 'pos') - assert hasattr(enc_data, 'edge_index') - - def test_preserves_batch(self, sample_hetero_data): - """Should preserve batch attribute.""" - enc_data = make_encoder_data(sample_hetero_data) - - assert hasattr(enc_data, 'batch') - assert enc_data.batch.shape[0] == sample_hetero_data['protein'].num_nodes - - # ============== SLAE Encoder Tests ============== class TestSLAEEncoder: @@ -363,16 +353,18 @@ def test_encoder_type(self, device): assert encoder.encoder_type == 'slae' def test_encoder_forward(self, device, sample_hetero_data_with_slae): - """Forward pass should return (s, V) tuple with raw embeddings.""" + """Forward pass should return (s, V, None) tuple with raw embeddings.""" encoder = SLAEEncoder(slae_dim=128).to(device) - s, V = encoder(sample_hetero_data_with_slae) + s, V, pp_edge_attr = encoder(sample_hetero_data_with_slae) n_atoms = sample_hetero_data_with_slae['protein'].num_nodes assert s.shape == (n_atoms, 128) assert V.shape == (n_atoms, 0, 3) # Raw embeddings should be identical to input assert torch.allclose(s, sample_hetero_data_with_slae['protein'].slae_embedding) + # SLAE encoder doesn't return edge features + assert pp_edge_attr is None def test_encoder_missing_embeddings_error(self, device, sample_hetero_data): """Should raise NotImplementedError when embeddings are missing.""" @@ -386,7 +378,7 @@ def test_encoder_no_nans(self, device, sample_hetero_data_with_slae): """Output should not contain NaNs or Infs.""" encoder = SLAEEncoder(slae_dim=128).to(device) - s, V = encoder(sample_hetero_data_with_slae) + s, V, _ = encoder(sample_hetero_data_with_slae) assert not torch.isnan(s).any(), "Scalar output contains NaNs" assert not torch.isinf(s).any(), "Scalar output contains Infs" @@ -397,6 +389,105 @@ def test_encoder_no_learnable_params(self, device): assert sum(p.numel() for p in encoder.parameters()) == 0 +# ============== ESM Encoder Tests ============== + +class TestESMEncoder: + """Tests for ESMEncoder (BaseProteinEncoder implementation).""" + + def test_esm_registered(self): + """ESM encoder should be registered.""" + from src.esm_encoder import ESMEncoder + cls = get_encoder_class('esm') + assert cls is ESMEncoder + + def test_build_encoder_esm(self, device): + """Should build ESM encoder from config.""" + from src import build_encoder + encoder = build_encoder({'encoder_type': 'esm'}, device) + assert encoder.encoder_type == 'esm' + + def test_encoder_output_dims(self, device): + """Encoder should expose correct output_dims.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + assert encoder.output_dims == (1536, 0) + + def test_encoder_type(self, device): + """Encoder should return 'esm' as encoder_type.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + assert encoder.encoder_type == 'esm' + + def test_encoder_forward(self, device, sample_hetero_data): + """Forward pass should return (s, V, None) tuple with raw embeddings.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + + # Add mock ESM embeddings + n_atoms = sample_hetero_data['protein'].num_nodes + sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 1536, device=device) + + s, V, pp_edge_attr = encoder(sample_hetero_data) + + assert s.shape == (n_atoms, 1536) + assert V.shape == (n_atoms, 0, 3) + # Raw embeddings should be identical to input + assert torch.allclose(s, sample_hetero_data['protein'].esm_embedding) + # ESM encoder doesn't return edge features + assert pp_edge_attr is None + + def test_encoder_missing_embeddings_error(self, device, sample_hetero_data): + """Should raise NotImplementedError when embeddings are missing.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + + # sample_hetero_data does NOT have esm_embedding + with pytest.raises(NotImplementedError, match="requires cached embeddings"): + encoder(sample_hetero_data) + + def test_encoder_no_learnable_params(self, device): + """ESM encoder should have no learnable parameters.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + assert sum(p.numel() for p in encoder.parameters()) == 0 + + def test_from_config(self, device): + """Should construct from config dict.""" + from src.esm_encoder import ESMEncoder + config = {'esm_dim': 2048} + encoder = ESMEncoder.from_config(config, device) + assert encoder.output_dims == (2048, 0) + + def test_esm_encoder_no_nans(self, device, sample_hetero_data): + """Output should not contain NaNs or Infs.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + + # Add mock ESM embeddings + n_atoms = sample_hetero_data['protein'].num_nodes + sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 1536, device=device) + + s, V, _ = encoder(sample_hetero_data) + + assert not torch.isnan(s).any(), "Scalar output contains NaNs" + assert not torch.isinf(s).any(), "Scalar output contains Infs" + + def test_esm_encoder_device_placement(self, device, sample_hetero_data): + """Verify tensors are on the correct device.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + + # Add mock ESM embeddings on correct device + n_atoms = sample_hetero_data['protein'].num_nodes + sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 1536, device=device) + + s, V, _ = encoder(sample_hetero_data) + + # Compare device types (handles cuda vs cuda:0) + assert s.device.type == device.type, f"Expected device type {device.type}, got {s.device.type}" + assert V.device.type == device.type, f"Expected device type {device.type}, got {V.device.type}" + + # ============== Encoder Interoperability Tests ============== class TestEncoderInteroperability: From ed9ce25bcc4c43d9350515a2282d606c1ff37dc2 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Tue, 3 Mar 2026 13:41:52 -0600 Subject: [PATCH 02/19] addressing changes in encoder classes --- src/__init__.py | 8 +- src/constants.py | 7 ++ src/encoder_base.py | 76 ++++++++++-- src/esm_encoder.py | 55 +++++++++ src/flow.py | 2 +- src/gvp_encoder.py | 279 +++++++++++++++++++++++++++++------------- src/slae_encoder.py | 39 ++---- tests/test_encoder.py | 181 ++++++++++++++++++++------- tests/test_flow.py | 10 +- tests/test_forward.py | 6 +- 10 files changed, 480 insertions(+), 183 deletions(-) create mode 100644 src/esm_encoder.py diff --git a/src/__init__.py b/src/__init__.py index 43d4a48..d7a8a58 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -6,5 +6,9 @@ Importing this module triggers encoder registration. """ -from src import gvp_encoder, slae -from src.encoder_base import BaseProteinEncoder, build_encoder, register_encoder +from src import esm_encoder, gvp_encoder, slae_encoder +from src.encoder_base import ( + BaseProteinEncoder as BaseProteinEncoder, + build_encoder as build_encoder, + register_encoder as register_encoder, +) diff --git a/src/constants.py b/src/constants.py index dbb7481..fdd24b6 100644 --- a/src/constants.py +++ b/src/constants.py @@ -6,6 +6,13 @@ IDE autocompletion and refactoring support. """ +# Node feature dimensions +NODE_FEATURE_DIM = 16 # Default node scalar feature dimension + +# RBF (Radial Basis Function) parameters +NUM_RBF = 16 # Number of RBF basis functions +RBF_CUTOFF = 8.0 # Distance cutoff in Angstroms for RBF encoding + # Edge type tuples: (src_node_type, edge_name, dst_node_type) EDGE_PP = ('protein', 'pp', 'protein') # protein -> protein EDGE_WW = ('water', 'ww', 'water') # water -> water diff --git a/src/encoder_base.py b/src/encoder_base.py index 115272c..0a1c525 100644 --- a/src/encoder_base.py +++ b/src/encoder_base.py @@ -5,6 +5,7 @@ - BaseProteinEncoder: Abstract base class that all encoders must implement - Registry pattern: Decorator-based registration and build_encoder() function """ + from __future__ import annotations from abc import ABC, abstractmethod @@ -16,7 +17,6 @@ if TYPE_CHECKING: from torch_geometric.data import HeteroData - class BaseProteinEncoder(ABC, nn.Module): """ Abstract base class for protein encoders. @@ -29,16 +29,18 @@ class BaseProteinEncoder(ABC, nn.Module): @abstractmethod def output_dims(self) -> tuple[int, int]: """Return (scalar_dim, vector_dim) output dimensions.""" - pass + raise NotImplementedError("Subclasses must implement output_dims") @property @abstractmethod def encoder_type(self) -> str: """Return encoder type identifier ('gvp', 'slae', 'esm', etc.).""" - pass + raise NotImplementedError("Subclasses must implement encoder_type") @abstractmethod - def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: + def forward( + self, data: HeteroData + ) -> tuple[torch.Tensor, torch.Tensor, tuple | None]: """ Encode protein data. @@ -46,11 +48,13 @@ def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: data: HeteroData with protein nodes Returns: - tuple of (s, V) where: + tuple of (s, V, pp_edge_attr) where: s: (N, scalar_dim) scalar features V: (N, vector_dim, 3) vector features + pp_edge_attr: tuple (s_edge, V_edge) for PP edges, or None if encoder + doesn't process edges (e.g., SLAE, ESM) """ - pass + raise NotImplementedError("Subclasses must implement forward") @classmethod @abstractmethod @@ -65,13 +69,61 @@ def from_config(cls, config: dict, device: torch.device) -> BaseProteinEncoder: Returns: Instantiated encoder """ - pass + raise NotImplementedError("Subclasses must implement from_config") + + +class CachedEmbeddingEncoder(BaseProteinEncoder): + """ + Base class for encoders that read pre-computed embeddings from data. + + Subclasses only need to define: + - encoder_type property (return string like 'slae', 'esm') + - from_config class method + """ + + def __init__(self, embedding_dim: int, embedding_key: str): + """ + Initialize CachedEmbeddingEncoder. + + Args: + embedding_dim: Dimension of the cached embeddings + embedding_key: Key to look up embeddings in data['protein'] + """ + super().__init__() + self._embedding_dim = embedding_dim + self._embedding_key = embedding_key + + @property + def output_dims(self) -> tuple[int, int]: + """Return (embedding_dim, 0) — scalars only.""" + return self._embedding_dim, 0 + + def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, None]: + """ + Read cached embeddings and return (s, V, None). + + Args: + data: HeteroData with cached embeddings in data['protein'] + + Returns: + s: (N, embedding_dim) — raw embeddings + V: (N, 0, 3) — empty vector features + pp_edge_attr: None — cached embedding encoders don't process edges + """ + if self._embedding_key not in data['protein']: + raise NotImplementedError( + f"{self.encoder_type.upper()} encoder requires cached embeddings. " + f"Please provide pre-computed {self._embedding_key} in data['protein']." + ) + + embeddings = data['protein'][self._embedding_key] + V = embeddings.new_empty(embeddings.size(0), 0, 3) + return embeddings, V, None # global encoder registry _ENCODER_REGISTRY: dict[str, BaseProteinEncoder] = {} - def register_encoder(name: str): """ Decorator to register an encoder class. @@ -81,13 +133,14 @@ def register_encoder(name: str): class MyEncoder(BaseProteinEncoder): ... """ + def decorator(cls: BaseProteinEncoder) -> BaseProteinEncoder: if name in _ENCODER_REGISTRY: raise ValueError(f"Encoder '{name}' is already registered") _ENCODER_REGISTRY[name] = cls return cls - return decorator + return decorator def get_encoder_class(name: str) -> BaseProteinEncoder: """ @@ -107,7 +160,6 @@ def get_encoder_class(name: str) -> BaseProteinEncoder: raise KeyError(f"Unknown encoder type '{name}'. Available: {available}") return _ENCODER_REGISTRY[name] - def build_encoder(config: dict, device: torch.device) -> BaseProteinEncoder: """ Build encoder from configuration dict. @@ -121,8 +173,8 @@ def build_encoder(config: dict, device: torch.device) -> BaseProteinEncoder: Returns: Instantiated encoder implementing BaseProteinEncoder """ - if 'encoder_type' not in config: + if "encoder_type" not in config: raise ValueError("'encoder_type' must be specified in config") - encoder_type = config['encoder_type'] + encoder_type = config["encoder_type"] encoder_cls = get_encoder_class(encoder_type) return encoder_cls.from_config(config, device) diff --git a/src/esm_encoder.py b/src/esm_encoder.py new file mode 100644 index 0000000..5c3bcb8 --- /dev/null +++ b/src/esm_encoder.py @@ -0,0 +1,55 @@ +# esm_encoder.py +""" +ESM embeddings wrapper. + +This encoder reads pre-computed ESM3 embeddings from data and returns +them directly as scalar features with zero vector channels. Downstream +GVP message-passing layers (including protein-protein edges) provide all +geometric processing. +""" + +from __future__ import annotations + +import torch + +from src.encoder_base import CachedEmbeddingEncoder, register_encoder + + +@register_encoder("esm") +class ESMEncoder(CachedEmbeddingEncoder): + """ + ESM encoder that reads cached embeddings from data. + + Returns (esm_embedding, empty_vectors) with output_dims = (esm_dim, 0). + No learnable parameters — all geometric processing happens in the + downstream ProteinWaterUpdate layers (pp, wp, pw, ww edges). + """ + + def __init__(self, esm_dim: int = 1536): + """ + Initialize ESMEncoder. + + Args: + esm_dim: Dimension of ESM embeddings (default: 1536 for ESM3-open) + """ + super().__init__(embedding_dim=esm_dim, embedding_key="esm_embedding") + + @property + def encoder_type(self) -> str: + return "esm" + + @classmethod + def from_config(cls, config: dict, device: torch.device) -> ESMEncoder: + """ + Construct ESMEncoder from config dict. + + Args: + config: Configuration dictionary with: + - esm_dim: ESM embedding dimension (default: 1536) + device: Device to place the encoder on + + Returns: + Instantiated ESMEncoder + """ + esm_dim = config.get("esm_dim", 1536) + return cls(esm_dim=esm_dim).to(device) diff --git a/src/flow.py b/src/flow.py index ed0025a..501623e 100644 --- a/src/flow.py +++ b/src/flow.py @@ -286,7 +286,7 @@ def forward(self, device = data['protein'].pos.device # Single unified encoder call - works for ANY encoder type - s_all, v_all = self.encoder(data) + s_all, v_all, _edge_attr = self.encoder(data) # Bridge encoder dims -> flow dims # Pass tuple when encoder has vector outputs, tensor when scalar-only diff --git a/src/gvp_encoder.py b/src/gvp_encoder.py index a477947..caac54b 100644 --- a/src/gvp_encoder.py +++ b/src/gvp_encoder.py @@ -17,57 +17,10 @@ from torch_geometric.data import Batch, Data, HeteroData from torch_scatter import scatter_add, scatter_max, scatter_mean -from src.constants import EDGE_PP +from src.constants import EDGE_PP, NODE_FEATURE_DIM, NUM_RBF, RBF_CUTOFF from src.encoder_base import BaseProteinEncoder, register_encoder from src.gvp import GVP, EdgeUpdate, GVPConvLayer -from src.utils import rbf as _rbf - - -def _edge_vectors(pos: torch.Tensor, edge_index: torch.Tensor): - """Compute edge vectors and distances.""" - src, dst = edge_index[0], edge_index[1] - vec = pos[dst] - pos[src] - rij = torch.linalg.norm(vec, dim=-1).clamp(min=1e-4) - r_hat = vec / rij[:, None] - return rij, r_hat - - -def make_encoder_data(data: HeteroData) -> Data: - """ - Build a homogeneous Data with protein nodes for GVP encoder. - - Extracts protein subgraph from HeteroData for use with GVP encoder. - Edge features are computed by the encoder itself. - - Args: - data: HeteroData with protein nodes - - Returns: - enc_data: Data with x, pos, edge_index - """ - device = data['protein'].pos.device - prot = data['protein'] - - x = prot.x - pos = prot.pos - - # protein-protein edges (topology only - features computed by encoder) - if EDGE_PP in data.edge_types: - edge_index = data[EDGE_PP].edge_index - else: - edge_index = torch.empty(2, 0, dtype=torch.long, device=device) - - enc_data = Data( - x=x, - pos=pos, - edge_index=edge_index, - ) - - # batch for multi-complex batches - if hasattr(prot, "batch"): - enc_data.batch = prot.batch - - return enc_data +from src.utils import rbf class ProteinGVPEncoder(nn.Module): @@ -80,12 +33,12 @@ class ProteinGVPEncoder(nn.Module): def __init__( self, - node_scalar_in: int = 17, + node_scalar_in: int = NODE_FEATURE_DIM, node_vec_in: int = 1, hidden_dims: tuple[int, int] = (256, 32), - edge_scalar_in: int = 16, + edge_scalar_in: int = NUM_RBF, edge_vec_in: int = 1, - edge_scalar_out: int = 16, + edge_scalar_out: int = NUM_RBF, n_layers: int = 3, n_message: int = 2, n_feedforward: int = 2, @@ -99,9 +52,37 @@ def __init__( pool_aggr: Literal["mean", "sum", "max"] = "mean", update_w_distance: bool = True, distance_dim: int | None = None, - radius: float = 8.0, - num_edge_rbf: int = 16, + radius: float = RBF_CUTOFF, + num_edge_rbf: int = NUM_RBF, + use_edge_update: bool = True, ): + """ + Initialize GVP encoder for protein structure processing. + + Args: + node_scalar_in: Input scalar feature dimension (e.g., element one-hot) + node_vec_in: Input vector feature channels (typically 1 for orientation) + hidden_dims: (scalar_dim, vector_dim) hidden layer dimensions + edge_scalar_in: Input edge scalar dimension (RBF features) + edge_vec_in: Input edge vector channels (unit vectors) + edge_scalar_out: Output edge scalar dimension + n_layers: Number of GVP convolution layers + n_message: Number of GVPs in message function + n_feedforward: Number of GVPs in feedforward function + drop_rate: Dropout rate for regularization + vector_gate: Whether to use vector gating in GVP layers + scalar_activation: Activation function for scalar channels + vector_activation: Activation function for vector gating + init_vec_zero: If True, initialize input vectors as zeros + pooled_dim: Output dimension when pooling by residue + pool_residue: If True, pool atom features to residue level + pool_aggr: Aggregation method for residue pooling ('mean', 'sum', 'max') + update_w_distance: Include distance features in edge updates + distance_dim: Dimension for distance conditioning, defaults to edge_scalar_in + radius: Distance cutoff in Angstroms for RBF encoding + num_edge_rbf: Number of RBF basis functions + use_edge_update: Whether to update edge features between layers + """ super().__init__() self.node_scalar_in = node_scalar_in self.node_vec_in = node_vec_in @@ -119,6 +100,7 @@ def __init__( self.radius = radius self.num_edge_rbf = num_edge_rbf self.pooled_dim = pooled_dim + self.use_edge_update = use_edge_update distance_dim = distance_dim or edge_scalar_in self.distance_dim = distance_dim @@ -161,12 +143,15 @@ def __init__( for _ in range(n_layers) ]) - self.edge_update = EdgeUpdate( - n_node_scalars=S_hid, - s_edge_width=self.s_edge_width, - update_w_distance=update_w_distance, - distance_dim=distance_dim, - ) + if use_edge_update: + self.edge_update = EdgeUpdate( + n_node_scalars=S_hid, + s_edge_width=self.s_edge_width, + update_w_distance=update_w_distance, + distance_dim=distance_dim, + ) + else: + self.edge_update = None self.atom_readout = nn.Sequential( nn.Linear(S_hid + V_hid, pooled_dim), @@ -178,6 +163,15 @@ def __init__( @staticmethod def _tuple_to_scalar_dense(x_tuple: tuple) -> torch.Tensor: + """ + Convert GVP tuple to dense scalar representation. + + Args: + x_tuple: (s, V) where s is (N, scalar_dim) and V is (N, vector_dim, 3) + + Returns: + (N, scalar_dim + vector_dim) concatenation of scalars and vector norms + """ s, V = x_tuple vnorm = torch.linalg.norm(V, dim=-1) return torch.cat([s, vnorm], dim=-1) @@ -185,29 +179,92 @@ def _tuple_to_scalar_dense(x_tuple: tuple) -> torch.Tensor: def _pool_by_residue( self, atom_embed: torch.Tensor, residue_index: torch.Tensor, num_residues: int ) -> torch.Tensor: + """ + Pool atom-level embeddings to residue level. + + Args: + atom_embed: (N_atoms, embed_dim) atom embeddings + residue_index: (N_atoms,) residue index per atom + num_residues: Total number of residues + + Returns: + (num_residues, embed_dim) pooled residue embeddings + """ aggr = self.pool_aggr if aggr == "mean": return scatter_mean(atom_embed, residue_index, dim=0, dim_size=num_residues) - if aggr == "sum": + elif aggr == "sum": return scatter_add(atom_embed, residue_index, dim=0, dim_size=num_residues) - if aggr == "max": + elif aggr == "max": out, _ = scatter_max(atom_embed, residue_index, dim=0, dim_size=num_residues) return out - raise ValueError(f"Unknown pool_aggr={aggr!r}") + else: + raise ValueError(f"Unknown pool_aggr={aggr!r}") @staticmethod - def _initial_node_tuple(x_scalar: torch.Tensor, device=None) -> tuple: + def _initial_node_tuple( + x_scalar: torch.Tensor, device: torch.device | None = None + ) -> tuple[torch.Tensor, torch.Tensor]: zeros = torch.zeros(x_scalar.size(0), 1, 3, device=x_scalar.device if device is None else device) return (x_scalar, zeros) - def _compute_edge_attr(self, pos: torch.Tensor, edge_index: torch.Tensor): - d, u = _edge_vectors(pos, edge_index) - s_edge_raw = _rbf(d, num_gaussians=self.num_edge_rbf, cutoff=self.radius) + @staticmethod + def _edge_vectors(pos: torch.Tensor, edge_index: torch.Tensor): + """ + Compute edge vectors and distances from node positions. + + Args: + pos: (N, 3) node position coordinates + edge_index: (2, E) edge indices with source in row 0, destination in row 1 + + Returns: + rij: (E,) edge distances clamped to minimum 1e-4 + r_hat: (E, 3) unit vectors pointing from source to destination + """ + src, dst = edge_index[0], edge_index[1] + vec = pos[dst] - pos[src] + rij = torch.linalg.norm(vec, dim=-1).clamp(min=1e-4) + r_hat = vec / rij[:, None] + return rij, r_hat + + def _compute_edge_attr(self, data: Batch): + """ + Build edge attributes from positions or cached features. + + If cached edge features (edge_rbf, edge_unit) are available in data, + use them directly. Otherwise, compute from positions. + + Args: + data: Batch with pos, edge_index, and optionally edge_rbf, edge_unit + + Returns: + (s_edge, V_edge): Tuple of scalar and vector edge features + s_edge_raw: Raw RBF features (for distance conditioning) + """ + # Use cached features if available + if hasattr(data, 'edge_rbf') and hasattr(data, 'edge_unit'): + s_edge_raw = data.edge_rbf + u = data.edge_unit + else: + # Fallback: compute from positions + d, u = self._edge_vectors(data.pos, data.edge_index) + s_edge_raw = rbf(d, num_gaussians=self.num_edge_rbf, cutoff=self.radius) + s_edge = self.edge_in_proj(s_edge_raw) V_edge = u.unsqueeze(1) return (s_edge, V_edge), s_edge_raw - def forward(self, data: Batch): + def forward(self, data: Batch) -> tuple[tuple, tuple | None]: + """ + Forward pass through the GVP encoder. + + Args: + data: Batch with node features, positions, and edge indices + + Returns: + x: tuple (s, V) of node scalar and vector features + edge_attr: tuple (s_edge, V_edge) of edge features, or None if use_edge_update=False + """ x_scalar = self.input_scalar_encoder(data.x) if self.init_vec_zero or not hasattr(data, "node_v"): @@ -216,26 +273,29 @@ def forward(self, data: Batch): node_features = (x_scalar, data.node_v.unsqueeze(1)) x = self.input_gvp(node_features) - edge_attr, dist_feat = self._compute_edge_attr(data.pos, data.edge_index) + edge_attr, dist_feat = self._compute_edge_attr(data) for layer in self.layers: x = layer(x, data.edge_index, edge_attr) - edge_attr = self.edge_update( - node_tuple=x, - edge_index=data.edge_index, - edge_attr=edge_attr, - distance_feat=(dist_feat if self.update_w_distance else None), - ) + if self.edge_update is not None: + edge_attr = self.edge_update( + node_tuple=x, + edge_index=data.edge_index, + edge_attr=edge_attr, + distance_feat=(dist_feat if self.update_w_distance else None), + ) if self.pool_residue: - assert hasattr(data, "residue_index") and hasattr(data, "num_residues"), \ - "Pooling requires data.residue_index (N,) and data.num_residues (int)." + if not (hasattr(data, "residue_index") and hasattr(data, "num_residues")): + raise ValueError("Pooling requires data.residue_index and data.num_residues") atom_dense = self._tuple_to_scalar_dense(x) atom_embed = self.atom_readout(atom_dense) res_embed = self._pool_by_residue(atom_embed, data.residue_index, int(data.num_residues)) - return res_embed + return res_embed, None # No edge features when pooling - return x + # Return edge_attr only if edge_update was used + final_edge_attr = edge_attr if self.edge_update is not None else None + return x, final_edge_attr def load_encoder_from_checkpoint( @@ -360,7 +420,53 @@ def encoder_type(self) -> str: """Return encoder type identifier.""" return 'gvp' - def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: + @staticmethod + def make_encoder_data(data: HeteroData) -> Data: + """ + Build a homogeneous Data with protein nodes for GVP encoder. + + Extracts protein subgraph from HeteroData for use with GVP encoder. + If PP edge features are cached in the HeteroData, they are copied through. + + Args: + data: HeteroData with protein nodes + + Returns: + enc_data: Data with x, pos, edge_index, and optionally cached edge features + """ + device = data['protein'].pos.device + prot = data['protein'] + + x = prot.x + pos = prot.pos + + # protein-protein edges + if EDGE_PP in data.edge_types: + edge_index = data[EDGE_PP].edge_index + else: + edge_index = torch.empty(2, 0, dtype=torch.long, device=device) + + enc_data = Data( + x=x, + pos=pos, + edge_index=edge_index, + ) + + # Copy cached edge features if available + if EDGE_PP in data.edge_types: + pp_edge = data[EDGE_PP] + if hasattr(pp_edge, 'edge_rbf'): + enc_data.edge_rbf = pp_edge.edge_rbf + if hasattr(pp_edge, 'edge_unit'): + enc_data.edge_unit = pp_edge.edge_unit + + # batch for multi-complex batches + if hasattr(prot, "batch"): + enc_data.batch = prot.batch + + return enc_data + + def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, tuple | None]: """ Encode protein data. @@ -368,15 +474,17 @@ def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: data: HeteroData with protein nodes Returns: - Tuple of (s, V) features + s: (N, scalar_dim) scalar features + V: (N, vector_dim, 3) vector features + pp_edge_attr: tuple (s_edge, V_edge) for PP edges, or None if edge updates disabled """ # Convert HeteroData to homogeneous Data for GVP encoder - enc_data = make_encoder_data(data) + enc_data = self.make_encoder_data(data) with torch.set_grad_enabled(not self._freeze): - s, V = self.encoder(enc_data) + (s, V), edge_attr = self.encoder(enc_data) - return s, V + return s, V, edge_attr @classmethod def from_config(cls, config: dict, device: torch.device) -> GVPEncoder: @@ -389,6 +497,7 @@ def from_config(cls, config: dict, device: torch.device) -> GVPEncoder: - node_scalar_in: Input feature dimension (default: 16) - hidden_s, hidden_v: Hidden dimensions - freeze_encoder: Whether to freeze encoder + - use_edge_update: Whether to use edge updates (default: True) device: Device to place the encoder on Returns: @@ -399,6 +508,7 @@ def from_config(cls, config: dict, device: torch.device) -> GVPEncoder: hidden_s = config.get('hidden_s', 256) hidden_v = config.get('hidden_v', 32) freeze = config.get('freeze_encoder', False) + use_edge_update = config.get('use_edge_update', True) if encoder_ckpt: encoder, _ = load_encoder_from_checkpoint( @@ -411,6 +521,7 @@ def from_config(cls, config: dict, device: torch.device) -> GVPEncoder: node_scalar_in=node_scalar_in, hidden_dims=(hidden_s, hidden_v), edge_scalar_in=16, + use_edge_update=use_edge_update, ).to(device) return cls(encoder=encoder, freeze=freeze) diff --git a/src/slae_encoder.py b/src/slae_encoder.py index 33ee84f..bf9331f 100644 --- a/src/slae_encoder.py +++ b/src/slae_encoder.py @@ -9,13 +9,12 @@ from __future__ import annotations import torch -from torch_geometric.data import HeteroData -from src.encoder_base import BaseProteinEncoder, register_encoder +from src.encoder_base import CachedEmbeddingEncoder, register_encoder @register_encoder('slae') -class SLAEEncoder(BaseProteinEncoder): +class SLAEEncoder(CachedEmbeddingEncoder): """ SLAE encoder that reads cached embeddings from data. @@ -25,39 +24,17 @@ class SLAEEncoder(BaseProteinEncoder): """ def __init__(self, slae_dim: int = 128): - super().__init__() - self._slae_dim = slae_dim - - @property - def output_dims(self) -> tuple[int, int]: - """Return (slae_dim, 0) — scalars only.""" - return self._slae_dim, 0 - - @property - def encoder_type(self) -> str: - return 'slae' - - def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: """ - Read cached SLAE embeddings and return (s, V). + Initialize SLAEEncoder. Args: - data: HeteroData with data['protein'].slae_embedding - - Returns: - s: (N, slae_dim) — raw SLAE embeddings - V: (N, 0, 3) — empty vector features + slae_dim: Dimension of SLAE embeddings (default: 128) """ - if 'slae_embedding' not in data['protein']: - raise NotImplementedError( - "SLAE encoder requires cached embeddings. " - "Please provide pre-computed slae_embedding in data['protein']. " - "Run scripts/precompute_slae_embeddings.py first." - ) + super().__init__(embedding_dim=slae_dim, embedding_key='slae_embedding') - embeddings = data['protein'].slae_embedding - V = embeddings.new_empty(embeddings.size(0), 0, 3) - return embeddings, V + @property + def encoder_type(self) -> str: + return 'slae' @classmethod def from_config(cls, config: dict, device: torch.device) -> SLAEEncoder: diff --git a/tests/test_encoder.py b/tests/test_encoder.py index f48531b..185bf8a 100644 --- a/tests/test_encoder.py +++ b/tests/test_encoder.py @@ -15,8 +15,8 @@ from torch_geometric.data import Data, HeteroData from src.encoder_base import build_encoder, get_encoder_class -from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder, make_encoder_data -from src.slae import SLAEEncoder +from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder +from src.slae_encoder import SLAEEncoder # ============== Fixtures ============== @@ -113,8 +113,8 @@ def test_slae_registered(self): assert cls is SLAEEncoder def test_unknown_encoder_raises(self): - """Unknown encoder type should raise ValueError.""" - with pytest.raises(ValueError, match="Unknown encoder type"): + """Unknown encoder type should raise KeyError.""" + with pytest.raises(KeyError, match="Unknown encoder type"): get_encoder_class('nonexistent_encoder') def test_build_encoder_gvp(self, device): @@ -143,12 +143,6 @@ def test_build_encoder_slae(self, device): assert encoder.encoder_type == 'slae' assert encoder.output_dims == (128, 0) - def test_build_encoder_default_gvp(self, device): - """build_encoder should default to GVP if encoder_type not specified.""" - config = {'node_scalar_in': 16} - encoder = build_encoder(config, device) - assert isinstance(encoder, GVPEncoder) - # ============== Base Interface Tests ============== @@ -172,11 +166,13 @@ def test_gvp_implements_interface(self, device, sample_hetero_data): assert len(encoder.output_dims) == 2 assert isinstance(encoder.encoder_type, str) - # Check forward returns (s, V) tuple - s, V = encoder(sample_hetero_data) + # Check forward returns (s, V, pp_edge_attr) tuple + s, V, pp_edge_attr = encoder(sample_hetero_data) assert s.shape[0] == sample_hetero_data['protein'].num_nodes assert V.shape[0] == sample_hetero_data['protein'].num_nodes assert V.shape[2] == 3 + # GVP encoder should return edge features + assert pp_edge_attr is not None or encoder.encoder.edge_update is None def test_slae_implements_interface(self, device, sample_hetero_data_with_slae): """SLAEEncoder should implement all required interface methods.""" @@ -188,11 +184,13 @@ def test_slae_implements_interface(self, device, sample_hetero_data_with_slae): assert encoder.output_dims == (128, 0) assert isinstance(encoder.encoder_type, str) - # Check forward returns (s, V) tuple - s, V = encoder(sample_hetero_data_with_slae) + # Check forward returns (s, V, pp_edge_attr) tuple + s, V, pp_edge_attr = encoder(sample_hetero_data_with_slae) assert s.shape[0] == sample_hetero_data_with_slae['protein'].num_nodes assert s.shape[1] == 128 assert V.shape == (sample_hetero_data_with_slae['protein'].num_nodes, 0, 3) + # SLAE encoder should return None for edge features + assert pp_edge_attr is None def test_from_config_class_method(self, device): """Both encoders should have from_config class method.""" @@ -232,11 +230,13 @@ def test_encoder_initialization(self, simple_encoder): def test_encoder_forward_with_pooling(self, simple_encoder, sample_homogeneous_data): """Test forward pass with residue pooling.""" - output = simple_encoder(sample_homogeneous_data) + output, edge_attr = simple_encoder(sample_homogeneous_data) assert output.shape == (sample_homogeneous_data.num_residues, simple_encoder.pooled_dim) + # Pooling mode returns None for edge features + assert edge_attr is None def test_encoder_forward_no_pooling(self, sample_homogeneous_data): - """Test encoder without residue pooling returns (s, V) tuple.""" + """Test encoder without residue pooling returns ((s, V), edge_attr) tuple.""" encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 16), @@ -246,12 +246,15 @@ def test_encoder_forward_no_pooling(self, sample_homogeneous_data): num_edge_rbf=16, ) - output = encoder(sample_homogeneous_data) + (s, v), edge_attr = encoder(sample_homogeneous_data) - assert isinstance(output, tuple) - s, v = output assert s.shape == (sample_homogeneous_data.num_nodes, 64) assert v.shape == (sample_homogeneous_data.num_nodes, 16, 3) + # edge_attr should be a tuple (s_edge, V_edge) + assert edge_attr is not None + s_edge, V_edge = edge_attr + assert s_edge.dim() == 2 + assert V_edge.dim() == 3 def test_encoder_empty_graph(self, simple_encoder): """Test encoder handles empty graphs gracefully.""" @@ -263,8 +266,10 @@ def test_encoder_empty_graph(self, simple_encoder): num_residues=0, ) - output = simple_encoder(data) + output, edge_attr = simple_encoder(data) assert output.shape == (0, simple_encoder.pooled_dim) + # Pooling mode returns None for edge features + assert edge_attr is None class TestGVPEncoderWrapper: @@ -297,7 +302,7 @@ def test_wrapper_encoder_type(self, device): assert encoder.encoder_type == 'gvp' def test_wrapper_forward(self, device, sample_hetero_data): - """Wrapper forward should return (s, V) from HeteroData.""" + """Wrapper forward should return (s, V, edge_attr) from HeteroData.""" base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 16), @@ -307,10 +312,15 @@ def test_wrapper_forward(self, device, sample_hetero_data): encoder = GVPEncoder(encoder=base_encoder, freeze=False) - s, V = encoder(sample_hetero_data) + s, V, edge_attr = encoder(sample_hetero_data) assert s.shape == (sample_hetero_data['protein'].num_nodes, 64) assert V.shape == (sample_hetero_data['protein'].num_nodes, 16, 3) + # edge_attr should be a tuple (s_edge, V_edge) when edge_update is enabled + assert edge_attr is not None + s_edge, V_edge = edge_attr + assert s_edge.dim() == 2 # (E, scalar_dim) + assert V_edge.dim() == 3 # (E, 1, 3) def test_wrapper_freeze(self, device): """Freeze parameter should disable gradients.""" @@ -327,26 +337,6 @@ def test_wrapper_freeze(self, device): assert not p.requires_grad -class TestMakeEncoderData: - """Tests for make_encoder_data helper function.""" - - def test_converts_hetero_to_homo(self, sample_hetero_data): - """Should convert HeteroData to homogeneous Data.""" - enc_data = make_encoder_data(sample_hetero_data) - - assert isinstance(enc_data, Data) - assert hasattr(enc_data, 'x') - assert hasattr(enc_data, 'pos') - assert hasattr(enc_data, 'edge_index') - - def test_preserves_batch(self, sample_hetero_data): - """Should preserve batch attribute.""" - enc_data = make_encoder_data(sample_hetero_data) - - assert hasattr(enc_data, 'batch') - assert enc_data.batch.shape[0] == sample_hetero_data['protein'].num_nodes - - # ============== SLAE Encoder Tests ============== class TestSLAEEncoder: @@ -363,16 +353,18 @@ def test_encoder_type(self, device): assert encoder.encoder_type == 'slae' def test_encoder_forward(self, device, sample_hetero_data_with_slae): - """Forward pass should return (s, V) tuple with raw embeddings.""" + """Forward pass should return (s, V, None) tuple with raw embeddings.""" encoder = SLAEEncoder(slae_dim=128).to(device) - s, V = encoder(sample_hetero_data_with_slae) + s, V, pp_edge_attr = encoder(sample_hetero_data_with_slae) n_atoms = sample_hetero_data_with_slae['protein'].num_nodes assert s.shape == (n_atoms, 128) assert V.shape == (n_atoms, 0, 3) # Raw embeddings should be identical to input assert torch.allclose(s, sample_hetero_data_with_slae['protein'].slae_embedding) + # SLAE encoder doesn't return edge features + assert pp_edge_attr is None def test_encoder_missing_embeddings_error(self, device, sample_hetero_data): """Should raise NotImplementedError when embeddings are missing.""" @@ -386,7 +378,7 @@ def test_encoder_no_nans(self, device, sample_hetero_data_with_slae): """Output should not contain NaNs or Infs.""" encoder = SLAEEncoder(slae_dim=128).to(device) - s, V = encoder(sample_hetero_data_with_slae) + s, V, _ = encoder(sample_hetero_data_with_slae) assert not torch.isnan(s).any(), "Scalar output contains NaNs" assert not torch.isinf(s).any(), "Scalar output contains Infs" @@ -397,6 +389,105 @@ def test_encoder_no_learnable_params(self, device): assert sum(p.numel() for p in encoder.parameters()) == 0 +# ============== ESM Encoder Tests ============== + +class TestESMEncoder: + """Tests for ESMEncoder (BaseProteinEncoder implementation).""" + + def test_esm_registered(self): + """ESM encoder should be registered.""" + from src.esm_encoder import ESMEncoder + cls = get_encoder_class('esm') + assert cls is ESMEncoder + + def test_build_encoder_esm(self, device): + """Should build ESM encoder from config.""" + from src import build_encoder + encoder = build_encoder({'encoder_type': 'esm'}, device) + assert encoder.encoder_type == 'esm' + + def test_encoder_output_dims(self, device): + """Encoder should expose correct output_dims.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + assert encoder.output_dims == (1536, 0) + + def test_encoder_type(self, device): + """Encoder should return 'esm' as encoder_type.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + assert encoder.encoder_type == 'esm' + + def test_encoder_forward(self, device, sample_hetero_data): + """Forward pass should return (s, V, None) tuple with raw embeddings.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + + # Add mock ESM embeddings + n_atoms = sample_hetero_data['protein'].num_nodes + sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 1536, device=device) + + s, V, pp_edge_attr = encoder(sample_hetero_data) + + assert s.shape == (n_atoms, 1536) + assert V.shape == (n_atoms, 0, 3) + # Raw embeddings should be identical to input + assert torch.allclose(s, sample_hetero_data['protein'].esm_embedding) + # ESM encoder doesn't return edge features + assert pp_edge_attr is None + + def test_encoder_missing_embeddings_error(self, device, sample_hetero_data): + """Should raise NotImplementedError when embeddings are missing.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + + # sample_hetero_data does NOT have esm_embedding + with pytest.raises(NotImplementedError, match="requires cached embeddings"): + encoder(sample_hetero_data) + + def test_encoder_no_learnable_params(self, device): + """ESM encoder should have no learnable parameters.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + assert sum(p.numel() for p in encoder.parameters()) == 0 + + def test_from_config(self, device): + """Should construct from config dict.""" + from src.esm_encoder import ESMEncoder + config = {'esm_dim': 2048} + encoder = ESMEncoder.from_config(config, device) + assert encoder.output_dims == (2048, 0) + + def test_esm_encoder_no_nans(self, device, sample_hetero_data): + """Output should not contain NaNs or Infs.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + + # Add mock ESM embeddings + n_atoms = sample_hetero_data['protein'].num_nodes + sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 1536, device=device) + + s, V, _ = encoder(sample_hetero_data) + + assert not torch.isnan(s).any(), "Scalar output contains NaNs" + assert not torch.isinf(s).any(), "Scalar output contains Infs" + + def test_esm_encoder_device_placement(self, device, sample_hetero_data): + """Verify tensors are on the correct device.""" + from src.esm_encoder import ESMEncoder + encoder = ESMEncoder(esm_dim=1536).to(device) + + # Add mock ESM embeddings on correct device + n_atoms = sample_hetero_data['protein'].num_nodes + sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 1536, device=device) + + s, V, _ = encoder(sample_hetero_data) + + # Compare device types (handles cuda vs cuda:0) + assert s.device.type == device.type, f"Expected device type {device.type}, got {s.device.type}" + assert V.device.type == device.type, f"Expected device type {device.type}, got {V.device.type}" + + # ============== Encoder Interoperability Tests ============== class TestEncoderInteroperability: diff --git a/tests/test_flow.py b/tests/test_flow.py index 7990f6a..d754bd7 100644 --- a/tests/test_flow.py +++ b/tests/test_flow.py @@ -16,7 +16,7 @@ ProteinWaterUpdate, build_knn_edges, ) -from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder, make_encoder_data +from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder @pytest.fixture @@ -148,7 +148,7 @@ def test_with_batch(self, device): class TestMakeEncoderData: def test_basic_output(self, simple_hetero_data): - enc_data = make_encoder_data(simple_hetero_data) + enc_data = GVPEncoder.make_encoder_data(simple_hetero_data) assert isinstance(enc_data, Data) assert hasattr(enc_data, 'x') @@ -156,7 +156,7 @@ def test_basic_output(self, simple_hetero_data): assert hasattr(enc_data, 'edge_index') def test_shapes(self, simple_hetero_data): - enc_data = make_encoder_data(simple_hetero_data) + enc_data = GVPEncoder.make_encoder_data(simple_hetero_data) n_nodes = simple_hetero_data['protein'].pos.size(0) n_edges = simple_hetero_data['protein', 'pp', 'protein'].edge_index.size(1) @@ -166,7 +166,7 @@ def test_shapes(self, simple_hetero_data): assert enc_data.edge_index.shape == (2, n_edges) def test_batch_preserved(self, batched_hetero_data): - enc_data = make_encoder_data(batched_hetero_data) + enc_data = GVPEncoder.make_encoder_data(batched_hetero_data) assert hasattr(enc_data, 'batch') assert enc_data.batch.shape[0] == batched_hetero_data['protein'].pos.size(0) @@ -177,7 +177,7 @@ def test_no_edges(self, device): data['protein'].x = torch.randn(10, 16, device=device) # No edges defined - enc_data = make_encoder_data(data) + enc_data = GVPEncoder.make_encoder_data(data) assert enc_data.edge_index.shape == (2, 0) diff --git a/tests/test_forward.py b/tests/test_forward.py index d15da6d..5458900 100644 --- a/tests/test_forward.py +++ b/tests/test_forward.py @@ -11,7 +11,7 @@ from torch_geometric.data import HeteroData from src.flow import FlowMatcher, FlowWaterGVP -from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder, make_encoder_data +from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder def _iter_tensors(obj): @@ -187,7 +187,7 @@ def test_forward_pass_no_nan_with_module_hooks(device): ).to(device) # Quick pre-check: protein encoder input features created from pp edges - enc_data = make_encoder_data(data) + enc_data = GVPEncoder.make_encoder_data(data) assert_edge_index_in_range(enc_data.edge_index, enc_data.x.size(0), enc_data.x.size(0), "pp edge_index") # Also validate knn edges are sane (catches orientation / k issues) @@ -361,7 +361,7 @@ def test_forward_with_duplicate_protein_coords_localizes_nan(device): ).to(device) # ---- Pre-check: encoder input edges ---- - enc_data = make_encoder_data(data) + enc_data = GVPEncoder.make_encoder_data(data) for tensor in _iter_tensors(enc_data): assert torch.isfinite(tensor).all() From adf7c993e0178315761f0347a88f861c0652fbc0 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Tue, 3 Mar 2026 13:52:56 -0600 Subject: [PATCH 03/19] using constants for edge dims checking --- src/gvp_encoder.py | 17 ++++++++--------- 1 file changed, 8 insertions(+), 9 deletions(-) diff --git a/src/gvp_encoder.py b/src/gvp_encoder.py index 260d65d..66d7293 100644 --- a/src/gvp_encoder.py +++ b/src/gvp_encoder.py @@ -17,7 +17,6 @@ from torch_geometric.data import Batch, Data, HeteroData from torch_scatter import scatter_add, scatter_max, scatter_mean -from src.constants import EDGE_PP, NODE_FEATURE_DIM, NUM_RBF, RBF_CUTOFF from src.constants import EDGE_PP, NODE_FEATURE_DIM, NUM_RBF, RBF_CUTOFF from src.encoder_base import BaseProteinEncoder, register_encoder from src.gvp import GVP, EdgeUpdate, GVPConvLayer @@ -305,8 +304,8 @@ def load_encoder_from_checkpoint( device: str = "cuda", default_hidden_dims: tuple[int, int] = (256, 64), default_pooled_dim: int = 128, - default_num_edge_rbf: int = 16, - default_radius: float = 8.0, + default_num_edge_rbf: int = NUM_RBF, + default_radius: float = RBF_CUTOFF, ) -> tuple[ProteinGVPEncoder, dict[str, Any]]: """ Load pretrained ProteinGVPEncoder from checkpoint. @@ -366,7 +365,7 @@ def load_encoder_from_checkpoint( hidden_dims=hidden_dims, edge_scalar_in=args.get("num_edge_rbf", default_num_edge_rbf), edge_vec_in=1, - edge_scalar_out=16, + edge_scalar_out=NUM_RBF, update_w_distance=True, pooled_dim=args.get("pooled_dim", default_pooled_dim), radius=args.get("radius", default_radius), @@ -521,7 +520,7 @@ def from_config(cls, config: dict, device: torch.device) -> GVPEncoder: encoder = ProteinGVPEncoder( node_scalar_in=node_scalar_in, hidden_dims=(hidden_s, hidden_v), - edge_scalar_in=16, + edge_scalar_in=NUM_RBF, use_edge_update=use_edge_update, ).to(device) @@ -552,13 +551,13 @@ def from_checkpoint( encoder = ProteinGVPEncoder( node_scalar_in=args["node_scalar_in"], hidden_dims=tuple(args["hidden_dims"]), - edge_scalar_in=args.get("num_edge_rbf", 16), + edge_scalar_in=args.get("num_edge_rbf", NUM_RBF), edge_vec_in=1, - edge_scalar_out=16, + edge_scalar_out=NUM_RBF, update_w_distance=True, pooled_dim=args.get("pooled_dim", 128), - radius=args.get("radius", 8.0), - num_edge_rbf=args.get("num_edge_rbf", 16), + radius=args.get("radius", RBF_CUTOFF), + num_edge_rbf=args.get("num_edge_rbf", NUM_RBF), ).to(device) encoder.load_state_dict(state_dict) From 9bac89d25c8096325f4c48643d6203f28cab0840 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Mon, 16 Mar 2026 11:32:44 -0500 Subject: [PATCH 04/19] addressing code review comments --- src/encoder_base.py | 198 +++++++++++++++++++----------- src/gvp_encoder.py | 206 ++++++++++++++++--------------- tests/test_encoder.py | 278 +++++++++++++++++++++++------------------- 3 files changed, 385 insertions(+), 297 deletions(-) diff --git a/src/encoder_base.py b/src/encoder_base.py index 579411e..4b4d752 100644 --- a/src/encoder_base.py +++ b/src/encoder_base.py @@ -4,6 +4,7 @@ This module provides: - BaseProteinEncoder: Abstract base class that all encoders must implement - Registry pattern: Decorator-based registration and build_encoder() function +- CachedEmbeddingEncoder: Concrete encoder for pre-computed embeddings (ESM, SLAE) """ from __future__ import annotations @@ -17,6 +18,67 @@ if TYPE_CHECKING: from torch_geometric.data import HeteroData +# global encoder registry +_ENCODER_REGISTRY: dict[str, type[BaseProteinEncoder]] = {} + + +def register_encoder(name: str): + """ + Decorator to register an encoder class. + + Usage: + @register_encoder('my_encoder') + class MyEncoder(BaseProteinEncoder): + ... + """ + + def decorator(cls: type[BaseProteinEncoder]) -> type[BaseProteinEncoder]: + if name in _ENCODER_REGISTRY: + raise ValueError(f"Encoder '{name}' is already registered") + _ENCODER_REGISTRY[name] = cls + return cls + + return decorator + + +def get_encoder_class(name: str) -> type[BaseProteinEncoder]: + """ + Get encoder class by name. + + Args: + name: Encoder type identifier + + Returns: + Encoder class + + Raises: + KeyError: If encoder name is not registered + """ + if name not in _ENCODER_REGISTRY: + available = list(_ENCODER_REGISTRY.keys()) + raise KeyError(f"Unknown encoder type '{name}'. Available: {available}") + return _ENCODER_REGISTRY[name] + + +def build_encoder(config: dict, device: torch.device) -> BaseProteinEncoder: + """ + Build encoder from configuration dict. + + Args: + config: Configuration dictionary containing: + - encoder_type: 'gvp', 'slae', or 'esm' (required) + - Other encoder-specific parameters + device: Device to place the encoder on + + Returns: + Instantiated encoder implementing BaseProteinEncoder + """ + if "encoder_type" not in config: + raise ValueError("'encoder_type' must be specified in config") + encoder_type = config["encoder_type"] + encoder_cls = get_encoder_class(encoder_type) + return encoder_cls.from_config(config, device) + class BaseProteinEncoder(ABC, nn.Module): """ Abstract base class for protein encoders. @@ -71,37 +133,71 @@ def from_config(cls, config: dict, device: torch.device) -> BaseProteinEncoder: """ raise NotImplementedError("Subclasses must implement from_config") - +@register_encoder("esm") +@register_encoder("slae") class CachedEmbeddingEncoder(BaseProteinEncoder): """ - Base class for encoders that read pre-computed embeddings from data. + Encoder for pre-computed protein embeddings (ESM, SLAE, etc.). + + This pass-through encoder reads embeddings stored in HeteroData under a + specified key and returns them as scalar features. No neural network + computation occurs; all geometric processing happens in downstream layers. + + Supported embedding types: + - ESM: Evolutionary Scale Modeling embeddings (https://github.com/evolutionaryscale/esm) + - SLAE: Strictly Local Atom-level Environment Embeddings (https://www.biorxiv.org/content/10.1101/2025.10.03.680398v1) - Subclasses only need to define: - - encoder_type property (return string like 'slae', 'esm') - - from_config class method + Embedding dimension is inferred from the data on first forward pass. + Accessing output_dims before forward() raises RuntimeError. + + Memory: Embeddings are NOT loaded at initialization. The encoder stores + only the key name; actual embeddings are read from data at forward time, + allowing standard PyTorch batching/streaming. + + Note: Returns empty vector features (shape Nx0x3) since cached embeddings + are scalar-only. """ - def __init__(self, embedding_dim: int, embedding_key: str): + def __init__(self, embedding_key: str, encoder_type: str, embedding_dim: int | None = None): """ Initialize CachedEmbeddingEncoder. Args: - embedding_dim: Dimension of the cached embeddings embedding_key: Key to look up embeddings in data['protein'] + encoder_type: Encoder type identifier ('esm' or 'slae') + embedding_dim: Optional embedding dimension. If provided, output_dims is + available immediately. If None, dimension is inferred on first forward. """ super().__init__() - self._embedding_dim = embedding_dim + self._embedding_dim: int | None = embedding_dim self._embedding_key = embedding_key + self._encoder_type = encoder_type @property def output_dims(self) -> tuple[int, int]: - """Return (embedding_dim, 0) — scalars only.""" + """Return (embedding_dim, 0) — scalars only. + + Raises: + RuntimeError: If accessed before first forward pass (dimension not yet inferred) + """ + if self._embedding_dim is None: + raise RuntimeError( + f"{self._encoder_type.upper()} encoder dimension not yet known. " + "Run a forward pass first to infer dimension from data." + ) return self._embedding_dim, 0 - def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, None]: + @property + def encoder_type(self) -> str: + """Return encoder type identifier.""" + return self._encoder_type + + def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, tuple | None]: """ Read cached embeddings and return (s, V, None). + On first call, infers embedding dimension from the data. + Args: data: HeteroData with cached embeddings in data['protein'] @@ -111,69 +207,35 @@ def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, None]: pp_edge_attr: None — cached embedding encoders don't process edges """ if self._embedding_key not in data['protein']: - raise NotImplementedError( - f"{self.encoder_type.upper()} encoder requires cached embeddings. " - f"Please provide pre-computed {self._embedding_key} in data['protein']." + raise KeyError( + f"{self._encoder_type.upper()} encoder requires cached embeddings. " + f"Please provide pre-computed '{self._embedding_key}' in data['protein']." ) embeddings = data['protein'][self._embedding_key] - V = embeddings.new_empty(embeddings.size(0), 0, 3) - return embeddings, V, None - -# global encoder registry -_ENCODER_REGISTRY: dict[str, BaseProteinEncoder] = {} - -def register_encoder(name: str): - """ - Decorator to register an encoder class. - - Usage: - @register_encoder('my_encoder') - class MyEncoder(BaseProteinEncoder): - ... - """ - - def decorator(cls: BaseProteinEncoder) -> BaseProteinEncoder: - if name in _ENCODER_REGISTRY: - raise ValueError(f"Encoder '{name}' is already registered") - _ENCODER_REGISTRY[name] = cls - return cls - - return decorator - -def get_encoder_class(name: str) -> BaseProteinEncoder: - """ - Get encoder class by name. - - Args: - name: Encoder type identifier - Returns: - Encoder class + # Infer dimension on first forward + if self._embedding_dim is None: + self._embedding_dim = embeddings.size(-1) - Raises: - KeyError: If encoder name is not registered - """ - if name not in _ENCODER_REGISTRY: - available = list(_ENCODER_REGISTRY.keys()) - raise KeyError(f"Unknown encoder type '{name}'. Available: {available}") - return _ENCODER_REGISTRY[name] + V = embeddings.new_empty(embeddings.size(0), 0, 3) + return embeddings, V, None -def build_encoder(config: dict, device: torch.device) -> BaseProteinEncoder: - """ - Build encoder from configuration dict. + @classmethod + def from_config(cls, config: dict, device: torch.device) -> CachedEmbeddingEncoder: + """ + Construct CachedEmbeddingEncoder from config dict. - Args: - config: Configuration dictionary containing: - - encoder_type: 'gvp' or 'slae' (required) - - Other encoder-specific parameters - device: Device to place the encoder on + Args: + config: Configuration dictionary with: + - encoder_type: 'esm' or 'slae' (required) + - embedding_dim: Optional embedding dimension (if known upfront) + device: Device to place the encoder on - Returns: - Instantiated encoder implementing BaseProteinEncoder - """ - if "encoder_type" not in config: - raise ValueError("'encoder_type' must be specified in config") - encoder_type = config["encoder_type"] - encoder_cls = get_encoder_class(encoder_type) - return encoder_cls.from_config(config, device) + Returns: + Instantiated CachedEmbeddingEncoder + """ + encoder_type = config["encoder_type"] # "esm" or "slae" + embedding_key = f"{encoder_type}_embedding" + embedding_dim = config.get("embedding_dim") + return cls(embedding_key, encoder_type, embedding_dim).to(device) diff --git a/src/gvp_encoder.py b/src/gvp_encoder.py index 66d7293..9d71206 100644 --- a/src/gvp_encoder.py +++ b/src/gvp_encoder.py @@ -23,6 +23,67 @@ from src.utils import rbf +def edge_vectors(pos: torch.Tensor, edge_index: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """ + Compute edge distances and unit vectors from node positions. + + Args: + pos: (N, 3) node position coordinates + edge_index: (2, E) edge indices with source in row 0, destination in row 1 + + Returns: + rij: (E,) edge distances, clamped to minimum 1e-4 to avoid division by zero + r_hat: (E, 3) unit vectors pointing from source to destination, + computed as vec / rij (using clamped distances) + """ + src, dst = edge_index[0], edge_index[1] + vec = pos[dst] - pos[src] + rij = torch.linalg.norm(vec, dim=-1).clamp(min=1e-4) + r_hat = vec / rij[:, None] + return rij, r_hat + + +def make_gvp_encoder_data(data: HeteroData) -> Data: + """ + Build homogeneous Data from HeteroData for GVP encoder. + + Extracts protein subgraph. Can be used independently of model instantiation. + + Args: + data: HeteroData with protein nodes + + Returns: + enc_data: Data with x, pos, edge_index, and optionally cached edge features + """ + device = data['protein'].pos.device + prot = data['protein'] + + x = prot.x + pos = prot.pos + + # protein-protein edges + if EDGE_PP in data.edge_types: + edge_index = data[EDGE_PP].edge_index + else: + edge_index = torch.empty(2, 0, dtype=torch.long, device=device) + + enc_data = Data(x=x, pos=pos, edge_index=edge_index) + + # Copy cached edge features if available + if EDGE_PP in data.edge_types: + pp_edge = data[EDGE_PP] + if hasattr(pp_edge, 'edge_rbf'): + enc_data.edge_rbf = pp_edge.edge_rbf + if hasattr(pp_edge, 'edge_unit'): + enc_data.edge_unit = pp_edge.edge_unit + + # batch for multi-complex batches + if hasattr(prot, "batch"): + enc_data.batch = prot.batch + + return enc_data + + class ProteinGVPEncoder(nn.Module): """ Core GVP encoder architecture for protein structures. @@ -36,9 +97,9 @@ def __init__( node_scalar_in: int = NODE_FEATURE_DIM, node_vec_in: int = 1, hidden_dims: tuple[int, int] = (256, 32), - edge_scalar_in: int = NUM_RBF, - edge_vec_in: int = 1, - edge_scalar_out: int = NUM_RBF, + n_edge_scalar_in: int = NUM_RBF, + n_edge_vec_in: int = 1, + n_edge_scalar_out: int = NUM_RBF, n_layers: int = 3, n_message: int = 2, n_feedforward: int = 2, @@ -50,7 +111,7 @@ def __init__( pooled_dim: int = 128, pool_residue: bool = False, pool_aggr: Literal["mean", "sum", "max"] = "mean", - update_w_distance: bool = True, + update_w_distance_features: bool = True, distance_dim: int | None = None, radius: float = RBF_CUTOFF, num_edge_rbf: int = NUM_RBF, @@ -60,12 +121,12 @@ def __init__( Initialize GVP encoder for protein structure processing. Args: - node_scalar_in: Input scalar feature dimension (e.g., element one-hot) - node_vec_in: Input vector feature channels (typically 1 for orientation) + node_scalar_in: Number of input scalar features (e.g., 16 for element one-hot) + node_vec_in: Number of input vector channels (typically 1 for orientation) hidden_dims: (scalar_dim, vector_dim) hidden layer dimensions - edge_scalar_in: Input edge scalar dimension (RBF features) - edge_vec_in: Input edge vector channels (unit vectors) - edge_scalar_out: Output edge scalar dimension + n_edge_scalar_in: Number of input edge scalar features (RBF features) + n_edge_vec_in: Number of input edge vector channels (e.g., 1 for unit displacement vectors) + n_edge_scalar_out: Number of output edge scalar features n_layers: Number of GVP convolution layers n_message: Number of GVPs in message function n_feedforward: Number of GVPs in feedforward function @@ -77,7 +138,7 @@ def __init__( pooled_dim: Output dimension when pooling by residue pool_residue: If True, pool atom features to residue level pool_aggr: Aggregation method for residue pooling ('mean', 'sum', 'max') - update_w_distance: Include distance features in edge updates + update_w_distance_features: Include distance features in edge updates distance_dim: Dimension for distance conditioning, defaults to edge_scalar_in radius: Distance cutoff in Angstroms for RBF encoding num_edge_rbf: Number of RBF basis functions @@ -87,22 +148,22 @@ def __init__( self.node_scalar_in = node_scalar_in self.node_vec_in = node_vec_in self.hidden_dims = hidden_dims - self.edge_scalar_in = edge_scalar_in - self.edge_vec_in = edge_vec_in - self.edge_scalar_out = edge_scalar_out + self.n_edge_scalar_in = n_edge_scalar_in + self.n_edge_vec_in = n_edge_vec_in + self.n_edge_scalar_out = n_edge_scalar_out self.n_layers = n_layers self.drop_rate = drop_rate self.vector_gate = vector_gate self.init_vec_zero = init_vec_zero self.pool_residue = pool_residue self.pool_aggr = pool_aggr - self.update_w_distance = update_w_distance + self.update_w_distance_features = update_w_distance_features self.radius = radius self.num_edge_rbf = num_edge_rbf self.pooled_dim = pooled_dim self.use_edge_update = use_edge_update - distance_dim = distance_dim or edge_scalar_in + distance_dim = distance_dim or n_edge_scalar_in self.distance_dim = distance_dim activations = (scalar_activation, vector_activation) @@ -123,13 +184,13 @@ def __init__( vector_gate=vector_gate, ) - self.s_edge_width = edge_scalar_out - if edge_scalar_in != edge_scalar_out: - self.edge_in_proj = nn.Linear(edge_scalar_in, edge_scalar_out, bias=False) + self.s_edge_width = n_edge_scalar_out + if n_edge_scalar_in != n_edge_scalar_out: + self.edge_in_proj = nn.Linear(n_edge_scalar_in, n_edge_scalar_out, bias=False) else: self.edge_in_proj = nn.Identity() - edge_dims = (self.s_edge_width, edge_vec_in) + edge_dims = (self.s_edge_width, n_edge_vec_in) self.layers = nn.ModuleList([ GVPConvLayer( node_dims=hidden_dims, @@ -147,7 +208,7 @@ def __init__( self.edge_update = EdgeUpdate( n_node_scalars=S_hid, s_edge_width=self.s_edge_width, - update_w_distance=update_w_distance, + update_w_distance_features=update_w_distance_features, distance_dim=distance_dim, ) else: @@ -166,11 +227,13 @@ def _tuple_to_scalar_dense(x_tuple: tuple) -> torch.Tensor: """ Convert GVP tuple to dense scalar representation. + Computes L2 norms of vector features and concatenates with scalars. + Args: x_tuple: (s, V) where s is (N, scalar_dim) and V is (N, vector_dim, 3) Returns: - (N, scalar_dim + vector_dim) concatenation of scalars and vector norms + (N, scalar_dim + vector_dim) concatenation of scalars and vector L2 norms """ s, V = x_tuple vnorm = torch.linalg.norm(V, dim=-1) @@ -208,30 +271,11 @@ def _initial_node_tuple( zeros = torch.zeros(x_scalar.size(0), 1, 3, device=x_scalar.device if device is None else device) return (x_scalar, zeros) - @staticmethod - def _edge_vectors(pos: torch.Tensor, edge_index: torch.Tensor): - """ - Compute edge vectors and distances from node positions. - - Args: - pos: (N, 3) node position coordinates - edge_index: (2, E) edge indices with source in row 0, destination in row 1 - - Returns: - rij: (E,) edge distances clamped to minimum 1e-4 - r_hat: (E, 3) unit vectors pointing from source to destination - """ - src, dst = edge_index[0], edge_index[1] - vec = pos[dst] - pos[src] - rij = torch.linalg.norm(vec, dim=-1).clamp(min=1e-4) - r_hat = vec / rij[:, None] - return rij, r_hat - def _compute_edge_attr(self, data: Batch): """ Build edge attributes from positions or cached features. - If cached edge features (edge_rbf, edge_unit) are available in data, + If cached edge features (edge_rbf, edge_unit) are both available in data, use them directly. Otherwise, compute from positions. Args: @@ -247,8 +291,8 @@ def _compute_edge_attr(self, data: Batch): u = data.edge_unit else: # Fallback: compute from positions - d, u = self._edge_vectors(data.pos, data.edge_index) - s_edge_raw = rbf(d, num_gaussians=self.num_edge_rbf, cutoff=self.radius) + rij, u = edge_vectors(data.pos, data.edge_index) + s_edge_raw = rbf(rij, num_gaussians=self.num_edge_rbf, cutoff=self.radius) s_edge = self.edge_in_proj(s_edge_raw) V_edge = u.unsqueeze(1) @@ -259,7 +303,13 @@ def forward(self, data: Batch) -> tuple[tuple, tuple | None]: Forward pass through the GVP encoder. Args: - data: Batch with node features, positions, and edge indices + data: PyG Batch/Data object with required attributes: + - x: (N, node_scalar_in) node scalar features + - pos: (N, 3) node position coordinates + - edge_index: (2, E) edge indices + Optional cached edge features (if absent, computed from pos): + - edge_rbf: (E, num_rbf) RBF distance features + - edge_unit: (E, 3) unit edge vectors Returns: x: tuple (s, V) of node scalar and vector features @@ -282,7 +332,7 @@ def forward(self, data: Batch) -> tuple[tuple, tuple | None]: node_tuple=x, edge_index=data.edge_index, edge_attr=edge_attr, - distance_feat=(dist_feat if self.update_w_distance else None), + distance_feat=(dist_feat if self.update_w_distance_features else None), ) if self.pool_residue: @@ -363,10 +413,10 @@ def load_encoder_from_checkpoint( encoder = ProteinGVPEncoder( node_scalar_in=node_scalar_in, hidden_dims=hidden_dims, - edge_scalar_in=args.get("num_edge_rbf", default_num_edge_rbf), - edge_vec_in=1, - edge_scalar_out=NUM_RBF, - update_w_distance=True, + n_edge_scalar_in=args.get("num_edge_rbf", default_num_edge_rbf), + n_edge_vec_in=1, + n_edge_scalar_out=NUM_RBF, + update_w_distance_features=True, pooled_dim=args.get("pooled_dim", default_pooled_dim), radius=args.get("radius", default_radius), num_edge_rbf=args.get("num_edge_rbf", default_num_edge_rbf), @@ -420,52 +470,6 @@ def encoder_type(self) -> str: """Return encoder type identifier.""" return 'gvp' - @staticmethod - def make_encoder_data(data: HeteroData) -> Data: - """ - Build a homogeneous Data with protein nodes for GVP encoder. - - Extracts protein subgraph from HeteroData for use with GVP encoder. - If PP edge features are cached in the HeteroData, they are copied through. - - Args: - data: HeteroData with protein nodes - - Returns: - enc_data: Data with x, pos, edge_index, and optionally cached edge features - """ - device = data['protein'].pos.device - prot = data['protein'] - - x = prot.x - pos = prot.pos - - # protein-protein edges - if EDGE_PP in data.edge_types: - edge_index = data[EDGE_PP].edge_index - else: - edge_index = torch.empty(2, 0, dtype=torch.long, device=device) - - enc_data = Data( - x=x, - pos=pos, - edge_index=edge_index, - ) - - # Copy cached edge features if available - if EDGE_PP in data.edge_types: - pp_edge = data[EDGE_PP] - if hasattr(pp_edge, 'edge_rbf'): - enc_data.edge_rbf = pp_edge.edge_rbf - if hasattr(pp_edge, 'edge_unit'): - enc_data.edge_unit = pp_edge.edge_unit - - # batch for multi-complex batches - if hasattr(prot, "batch"): - enc_data.batch = prot.batch - - return enc_data - def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, tuple | None]: """ Encode protein data. @@ -479,7 +483,7 @@ def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, tuple | pp_edge_attr: tuple (s_edge, V_edge) for PP edges, or None if edge updates disabled """ # Convert HeteroData to homogeneous Data for GVP encoder - enc_data = self.make_encoder_data(data) + enc_data = make_gvp_encoder_data(data) with torch.set_grad_enabled(not self._freeze): (s, V), edge_attr = self.encoder(enc_data) @@ -520,7 +524,7 @@ def from_config(cls, config: dict, device: torch.device) -> GVPEncoder: encoder = ProteinGVPEncoder( node_scalar_in=node_scalar_in, hidden_dims=(hidden_s, hidden_v), - edge_scalar_in=NUM_RBF, + n_edge_scalar_in=NUM_RBF, use_edge_update=use_edge_update, ).to(device) @@ -551,10 +555,10 @@ def from_checkpoint( encoder = ProteinGVPEncoder( node_scalar_in=args["node_scalar_in"], hidden_dims=tuple(args["hidden_dims"]), - edge_scalar_in=args.get("num_edge_rbf", NUM_RBF), - edge_vec_in=1, - edge_scalar_out=NUM_RBF, - update_w_distance=True, + n_edge_scalar_in=args.get("num_edge_rbf", NUM_RBF), + n_edge_vec_in=1, + n_edge_scalar_out=NUM_RBF, + update_w_distance_features=True, pooled_dim=args.get("pooled_dim", 128), radius=args.get("radius", RBF_CUTOFF), num_edge_rbf=args.get("num_edge_rbf", NUM_RBF), diff --git a/tests/test_encoder.py b/tests/test_encoder.py index 3b47f41..813976e 100644 --- a/tests/test_encoder.py +++ b/tests/test_encoder.py @@ -5,7 +5,7 @@ 1. Registry pattern (register, get, build) 2. Base encoder interface contract 3. GVP encoder (ProteinGVPEncoder + GVPEncoder wrapper) -4. SLAE encoder and projection +4. Cached embedding encoder (SLAE, ESM) 5. Encoder interoperability (both work with flow model) """ @@ -14,9 +14,8 @@ from torch_cluster import radius_graph from torch_geometric.data import Data, HeteroData -from src.encoder_base import build_encoder, get_encoder_class +from src.encoder_base import build_encoder, get_encoder_class, CachedEmbeddingEncoder from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder -from src.slae_encoder import SLAEEncoder # ============== Fixtures ============== @@ -90,7 +89,7 @@ def test_build_encoder_works_with_package_import(self): gvp_encoder = pkg_build_encoder({'encoder_type': 'gvp', 'node_scalar_in': 16}, device) assert gvp_encoder.encoder_type == 'gvp' - slae_encoder = pkg_build_encoder({'encoder_type': 'slae', 'slae_dim': 128}, device) + slae_encoder = pkg_build_encoder({'encoder_type': 'slae'}, device) assert slae_encoder.encoder_type == 'slae' @@ -105,7 +104,12 @@ def test_gvp_registered(self): def test_slae_registered(self): """SLAE encoder should be registered.""" cls = get_encoder_class('slae') - assert cls is SLAEEncoder + assert cls is CachedEmbeddingEncoder + + def test_esm_registered(self): + """ESM encoder should be registered.""" + cls = get_encoder_class('esm') + assert cls is CachedEmbeddingEncoder def test_unknown_encoder_raises(self): """Unknown encoder type should raise KeyError.""" @@ -130,13 +134,23 @@ def test_build_encoder_slae(self, device): """build_encoder should construct SLAE encoder from config.""" config = { 'encoder_type': 'slae', - 'slae_dim': 128, } encoder = build_encoder(config, device) - assert isinstance(encoder, SLAEEncoder) + assert isinstance(encoder, CachedEmbeddingEncoder) assert encoder.encoder_type == 'slae' - assert encoder.output_dims == (128, 0) + # output_dims not available until forward pass + + def test_build_encoder_esm(self, device): + """build_encoder should construct ESM encoder from config.""" + config = { + 'encoder_type': 'esm', + } + encoder = build_encoder(config, device) + + assert isinstance(encoder, CachedEmbeddingEncoder) + assert encoder.encoder_type == 'esm' + # output_dims not available until forward pass # ============== Base Interface Tests ============== @@ -150,7 +164,7 @@ def test_gvp_implements_interface(self, device, sample_hetero_data): encoder=ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 16), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device), freeze=False, @@ -169,14 +183,12 @@ def test_gvp_implements_interface(self, device, sample_hetero_data): # GVP encoder should return edge features assert pp_edge_attr is not None or encoder.encoder.edge_update is None - def test_slae_implements_interface(self, device, sample_hetero_data_with_slae): - """SLAEEncoder should implement all required interface methods.""" - encoder = SLAEEncoder(slae_dim=128).to(device) + def test_cached_embedding_implements_interface(self, device, sample_hetero_data_with_slae): + """CachedEmbeddingEncoder should implement all required interface methods.""" + encoder = CachedEmbeddingEncoder( + embedding_key='slae_embedding', encoder_type='slae' + ).to(device) - # Check properties - assert isinstance(encoder.output_dims, tuple) - assert len(encoder.output_dims) == 2 - assert encoder.output_dims == (128, 0) assert isinstance(encoder.encoder_type, str) # Check forward returns (s, V, pp_edge_attr) tuple @@ -184,20 +196,25 @@ def test_slae_implements_interface(self, device, sample_hetero_data_with_slae): assert s.shape[0] == sample_hetero_data_with_slae['protein'].num_nodes assert s.shape[1] == 128 assert V.shape == (sample_hetero_data_with_slae['protein'].num_nodes, 0, 3) - # SLAE encoder should return None for edge features + # Cached embedding encoder should return None for edge features assert pp_edge_attr is None + # output_dims available after forward + assert isinstance(encoder.output_dims, tuple) + assert len(encoder.output_dims) == 2 + assert encoder.output_dims == (128, 0) + def test_from_config_class_method(self, device): """Both encoders should have from_config class method.""" gvp_config = {'encoder_type': 'gvp', 'node_scalar_in': 16, 'hidden_s': 64, 'hidden_v': 16} - slae_config = {'encoder_type': 'slae', 'slae_dim': 128} + slae_config = {'encoder_type': 'slae'} gvp_encoder = GVPEncoder.from_config(gvp_config, device) - slae_encoder = SLAEEncoder.from_config(slae_config, device) + slae_encoder = CachedEmbeddingEncoder.from_config(slae_config, device) assert isinstance(gvp_encoder, GVPEncoder) - assert isinstance(slae_encoder, SLAEEncoder) - assert slae_encoder.output_dims == (128, 0) + assert isinstance(slae_encoder, CachedEmbeddingEncoder) + # output_dims not available until forward pass for cached encoders # ============== GVP Encoder Tests ============== @@ -211,7 +228,7 @@ def simple_encoder(self): return ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 16), - edge_scalar_in=16, + n_edge_scalar_in=16, n_layers=2, pooled_dim=32, pool_residue=True, @@ -235,7 +252,7 @@ def test_encoder_forward_no_pooling(self, sample_homogeneous_data): encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 16), - edge_scalar_in=16, + n_edge_scalar_in=16, n_layers=1, pool_residue=False, num_edge_rbf=16, @@ -275,7 +292,7 @@ def test_wrapper_output_dims(self, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(128, 32), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device) @@ -288,7 +305,7 @@ def test_wrapper_encoder_type(self, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 16), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device) @@ -301,7 +318,7 @@ def test_wrapper_forward(self, device, sample_hetero_data): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 16), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device) @@ -322,7 +339,7 @@ def test_wrapper_freeze(self, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 16), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device) @@ -332,24 +349,56 @@ def test_wrapper_freeze(self, device): assert not p.requires_grad -# ============== SLAE Encoder Tests ============== +# ============== Cached Embedding Encoder Tests ============== -class TestSLAEEncoder: - """Tests for SLAEEncoder (BaseProteinEncoder implementation).""" +class TestCachedEmbeddingEncoder: + """Tests for CachedEmbeddingEncoder (handles both SLAE and ESM).""" - def test_encoder_output_dims(self, device): - """Encoder should expose correct output_dims.""" - encoder = SLAEEncoder(slae_dim=128).to(device) + def test_output_dims_before_forward_raises(self, device): + """output_dims should raise RuntimeError before forward pass.""" + encoder = CachedEmbeddingEncoder( + embedding_key='slae_embedding', encoder_type='slae' + ).to(device) + with pytest.raises(RuntimeError, match="dimension not yet known"): + _ = encoder.output_dims + + def test_slae_output_dims_after_forward(self, device, sample_hetero_data_with_slae): + """SLAE encoder should infer output_dims from data.""" + encoder = CachedEmbeddingEncoder( + embedding_key='slae_embedding', encoder_type='slae' + ).to(device) + encoder(sample_hetero_data_with_slae) assert encoder.output_dims == (128, 0) - def test_encoder_type(self, device): - """Encoder should return 'slae' as encoder_type.""" - encoder = SLAEEncoder(slae_dim=128).to(device) + def test_esm_output_dims_after_forward(self, device, sample_hetero_data): + """ESM encoder should infer output_dims from data.""" + encoder = CachedEmbeddingEncoder( + embedding_key='esm_embedding', encoder_type='esm' + ).to(device) + n_atoms = sample_hetero_data['protein'].num_nodes + sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 1536, device=device) + encoder(sample_hetero_data) + assert encoder.output_dims == (1536, 0) + + def test_slae_encoder_type(self, device): + """SLAE encoder should return 'slae' as encoder_type.""" + encoder = CachedEmbeddingEncoder( + embedding_key='slae_embedding', encoder_type='slae' + ).to(device) assert encoder.encoder_type == 'slae' - def test_encoder_forward(self, device, sample_hetero_data_with_slae): - """Forward pass should return (s, V, None) tuple with raw embeddings.""" - encoder = SLAEEncoder(slae_dim=128).to(device) + def test_esm_encoder_type(self, device): + """ESM encoder should return 'esm' as encoder_type.""" + encoder = CachedEmbeddingEncoder( + embedding_key='esm_embedding', encoder_type='esm' + ).to(device) + assert encoder.encoder_type == 'esm' + + def test_slae_forward(self, device, sample_hetero_data_with_slae): + """SLAE forward pass should return (s, V, None) tuple with raw embeddings.""" + encoder = CachedEmbeddingEncoder( + embedding_key='slae_embedding', encoder_type='slae' + ).to(device) s, V, pp_edge_attr = encoder(sample_hetero_data_with_slae) @@ -358,65 +407,14 @@ def test_encoder_forward(self, device, sample_hetero_data_with_slae): assert V.shape == (n_atoms, 0, 3) # Raw embeddings should be identical to input assert torch.allclose(s, sample_hetero_data_with_slae['protein'].slae_embedding) - # SLAE encoder doesn't return edge features + # Cached embedding encoder doesn't return edge features assert pp_edge_attr is None - def test_encoder_missing_embeddings_error(self, device, sample_hetero_data): - """Should raise NotImplementedError when embeddings are missing.""" - encoder = SLAEEncoder(slae_dim=128).to(device) - - # sample_hetero_data does NOT have slae_embedding - with pytest.raises(NotImplementedError, match="requires cached embeddings"): - encoder(sample_hetero_data) - - def test_encoder_no_nans(self, device, sample_hetero_data_with_slae): - """Output should not contain NaNs or Infs.""" - encoder = SLAEEncoder(slae_dim=128).to(device) - - s, V, _ = encoder(sample_hetero_data_with_slae) - - assert not torch.isnan(s).any(), "Scalar output contains NaNs" - assert not torch.isinf(s).any(), "Scalar output contains Infs" - - def test_encoder_no_learnable_params(self, device): - """SLAE encoder should have no learnable parameters.""" - encoder = SLAEEncoder(slae_dim=128).to(device) - assert sum(p.numel() for p in encoder.parameters()) == 0 - - -# ============== ESM Encoder Tests ============== - -class TestESMEncoder: - """Tests for ESMEncoder (BaseProteinEncoder implementation).""" - - def test_esm_registered(self): - """ESM encoder should be registered.""" - from src.esm_encoder import ESMEncoder - cls = get_encoder_class('esm') - assert cls is ESMEncoder - - def test_build_encoder_esm(self, device): - """Should build ESM encoder from config.""" - from src import build_encoder - encoder = build_encoder({'encoder_type': 'esm'}, device) - assert encoder.encoder_type == 'esm' - - def test_encoder_output_dims(self, device): - """Encoder should expose correct output_dims.""" - from src.esm_encoder import ESMEncoder - encoder = ESMEncoder(esm_dim=1536).to(device) - assert encoder.output_dims == (1536, 0) - - def test_encoder_type(self, device): - """Encoder should return 'esm' as encoder_type.""" - from src.esm_encoder import ESMEncoder - encoder = ESMEncoder(esm_dim=1536).to(device) - assert encoder.encoder_type == 'esm' - - def test_encoder_forward(self, device, sample_hetero_data): - """Forward pass should return (s, V, None) tuple with raw embeddings.""" - from src.esm_encoder import ESMEncoder - encoder = ESMEncoder(esm_dim=1536).to(device) + def test_esm_forward(self, device, sample_hetero_data): + """ESM forward pass should return (s, V, None) tuple with raw embeddings.""" + encoder = CachedEmbeddingEncoder( + embedding_key='esm_embedding', encoder_type='esm' + ).to(device) # Add mock ESM embeddings n_atoms = sample_hetero_data['protein'].num_nodes @@ -428,49 +426,70 @@ def test_encoder_forward(self, device, sample_hetero_data): assert V.shape == (n_atoms, 0, 3) # Raw embeddings should be identical to input assert torch.allclose(s, sample_hetero_data['protein'].esm_embedding) - # ESM encoder doesn't return edge features + # Cached embedding encoder doesn't return edge features assert pp_edge_attr is None - def test_encoder_missing_embeddings_error(self, device, sample_hetero_data): - """Should raise NotImplementedError when embeddings are missing.""" - from src.esm_encoder import ESMEncoder - encoder = ESMEncoder(esm_dim=1536).to(device) + def test_slae_missing_embeddings_error(self, device, sample_hetero_data): + """Should raise KeyError when SLAE embeddings are missing.""" + encoder = CachedEmbeddingEncoder( + embedding_key='slae_embedding', encoder_type='slae' + ).to(device) - # sample_hetero_data does NOT have esm_embedding - with pytest.raises(NotImplementedError, match="requires cached embeddings"): + # sample_hetero_data does NOT have slae_embedding + with pytest.raises(KeyError, match="requires cached embeddings"): encoder(sample_hetero_data) - def test_encoder_no_learnable_params(self, device): - """ESM encoder should have no learnable parameters.""" - from src.esm_encoder import ESMEncoder - encoder = ESMEncoder(esm_dim=1536).to(device) - assert sum(p.numel() for p in encoder.parameters()) == 0 + def test_esm_missing_embeddings_error(self, device, sample_hetero_data): + """Should raise KeyError when ESM embeddings are missing.""" + encoder = CachedEmbeddingEncoder( + embedding_key='esm_embedding', encoder_type='esm' + ).to(device) - def test_from_config(self, device): - """Should construct from config dict.""" - from src.esm_encoder import ESMEncoder - config = {'esm_dim': 2048} - encoder = ESMEncoder.from_config(config, device) - assert encoder.output_dims == (2048, 0) + # sample_hetero_data does NOT have esm_embedding + with pytest.raises(KeyError, match="requires cached embeddings"): + encoder(sample_hetero_data) - def test_esm_encoder_no_nans(self, device, sample_hetero_data): + def test_encoder_no_nans(self, device, sample_hetero_data_with_slae): """Output should not contain NaNs or Infs.""" - from src.esm_encoder import ESMEncoder - encoder = ESMEncoder(esm_dim=1536).to(device) - - # Add mock ESM embeddings - n_atoms = sample_hetero_data['protein'].num_nodes - sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 1536, device=device) + encoder = CachedEmbeddingEncoder( + embedding_key='slae_embedding', encoder_type='slae' + ).to(device) - s, V, _ = encoder(sample_hetero_data) + s, V, _ = encoder(sample_hetero_data_with_slae) assert not torch.isnan(s).any(), "Scalar output contains NaNs" assert not torch.isinf(s).any(), "Scalar output contains Infs" - def test_esm_encoder_device_placement(self, device, sample_hetero_data): + def test_encoder_no_learnable_params(self, device): + """Cached embedding encoder should have no learnable parameters.""" + encoder = CachedEmbeddingEncoder( + embedding_key='slae_embedding', encoder_type='slae' + ).to(device) + assert sum(p.numel() for p in encoder.parameters()) == 0 + + def test_slae_from_config(self, device, sample_hetero_data_with_slae): + """Should construct SLAE from config and infer dim from data.""" + config = {'encoder_type': 'slae'} + encoder = CachedEmbeddingEncoder.from_config(config, device) + assert encoder.encoder_type == 'slae' + encoder(sample_hetero_data_with_slae) + assert encoder.output_dims == (128, 0) + + def test_esm_from_config(self, device, sample_hetero_data): + """Should construct ESM from config and infer dim from data.""" + config = {'encoder_type': 'esm'} + encoder = CachedEmbeddingEncoder.from_config(config, device) + assert encoder.encoder_type == 'esm' + n_atoms = sample_hetero_data['protein'].num_nodes + sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 2048, device=device) + encoder(sample_hetero_data) + assert encoder.output_dims == (2048, 0) + + def test_device_placement(self, device, sample_hetero_data): """Verify tensors are on the correct device.""" - from src.esm_encoder import ESMEncoder - encoder = ESMEncoder(esm_dim=1536).to(device) + encoder = CachedEmbeddingEncoder( + embedding_key='esm_embedding', encoder_type='esm' + ).to(device) # Add mock ESM embeddings on correct device n_atoms = sample_hetero_data['protein'].num_nodes @@ -499,14 +518,17 @@ def test_both_encoders_work_with_flow(self, device, sample_hetero_data_with_slae encoder=ProteinGVPEncoder( node_scalar_in=16, hidden_dims=hidden_dims, - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device), freeze=False, ) - # SLAE encoder (output_dims = (128, 0), bridged by encoder_to_flow) - slae_encoder = SLAEEncoder(slae_dim=128).to(device) + # SLAE encoder via CachedEmbeddingEncoder + # embedding_dim=128 matches the fixture's slae_embedding shape + slae_encoder = CachedEmbeddingEncoder( + embedding_key='slae_embedding', encoder_type='slae', embedding_dim=128 + ).to(device) # Create flow models with each encoder flow_gvp = FlowWaterGVP( From 3a30c88ea52fbdd6e6ffab669c72c6db6c224b43 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Mon, 16 Mar 2026 11:46:12 -0500 Subject: [PATCH 05/19] github actions --- .github/workflows/build.yml | 8 ++-- .github/workflows/lint.yml | 4 +- src/__init__.py | 3 +- src/esm_encoder.py | 55 -------------------------- src/gvp.py | 8 ++-- src/slae.py | 78 ------------------------------------- src/slae_encoder.py | 53 ------------------------- 7 files changed, 12 insertions(+), 197 deletions(-) delete mode 100644 src/esm_encoder.py delete mode 100644 src/slae.py delete mode 100644 src/slae_encoder.py diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 4fc5448..4d5941d 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -1,10 +1,10 @@ name: Build on: - # push: - # branches: [main] - # pull_request: - # branches: [main] + push: + branches: [main] + pull_request: + branches: [main] workflow_dispatch: jobs: diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml index d08821d..08862ed 100644 --- a/.github/workflows/lint.yml +++ b/.github/workflows/lint.yml @@ -1,8 +1,8 @@ name: Lint on: - # pull_request: - # branches: [main] + pull_request: + branches: [main] workflow_dispatch: jobs: diff --git a/src/__init__.py b/src/__init__.py index d7a8a58..4492b19 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -6,9 +6,10 @@ Importing this module triggers encoder registration. """ -from src import esm_encoder, gvp_encoder, slae_encoder +from src import gvp_encoder from src.encoder_base import ( BaseProteinEncoder as BaseProteinEncoder, + CachedEmbeddingEncoder as CachedEmbeddingEncoder, build_encoder as build_encoder, register_encoder as register_encoder, ) diff --git a/src/esm_encoder.py b/src/esm_encoder.py deleted file mode 100644 index 5c3bcb8..0000000 --- a/src/esm_encoder.py +++ /dev/null @@ -1,55 +0,0 @@ -# esm_encoder.py -""" -ESM embeddings wrapper. - -This encoder reads pre-computed ESM3 embeddings from data and returns -them directly as scalar features with zero vector channels. Downstream -GVP message-passing layers (including protein-protein edges) provide all -geometric processing. -""" - -from __future__ import annotations - -import torch - -from src.encoder_base import CachedEmbeddingEncoder, register_encoder - - -@register_encoder("esm") -class ESMEncoder(CachedEmbeddingEncoder): - """ - ESM encoder that reads cached embeddings from data. - - Returns (esm_embedding, empty_vectors) with output_dims = (esm_dim, 0). - No learnable parameters — all geometric processing happens in the - downstream ProteinWaterUpdate layers (pp, wp, pw, ww edges). - """ - - def __init__(self, esm_dim: int = 1536): - """ - Initialize ESMEncoder. - - Args: - esm_dim: Dimension of ESM embeddings (default: 1536 for ESM3-open) - """ - super().__init__(embedding_dim=esm_dim, embedding_key="esm_embedding") - - @property - def encoder_type(self) -> str: - return "esm" - - @classmethod - def from_config(cls, config: dict, device: torch.device) -> ESMEncoder: - """ - Construct ESMEncoder from config dict. - - Args: - config: Configuration dictionary with: - - esm_dim: ESM embedding dimension (default: 1536) - device: Device to place the encoder on - - Returns: - Instantiated ESMEncoder - """ - esm_dim = config.get("esm_dim", 1536) - return cls(esm_dim=esm_dim).to(device) diff --git a/src/gvp.py b/src/gvp.py index 5abfbe2..4cbfef0 100644 --- a/src/gvp.py +++ b/src/gvp.py @@ -440,14 +440,14 @@ def __init__( self, n_node_scalars: int, # S_node (e.g., 256) s_edge_width: int, # fixed model edge width used everywhere - update_w_distance: bool = False, + update_w_distance_features: bool = False, distance_dim: int = 0, # e.g., RBF size ): super().__init__() - self.update_w_distance = update_w_distance + self.update_w_distance_features = update_w_distance_features self.s_edge_width = s_edge_width - in_dim = (2 * n_node_scalars) + s_edge_width + (distance_dim if update_w_distance else 0) + in_dim = (2 * n_node_scalars) + s_edge_width + (distance_dim if update_w_distance_features else 0) self.edge_mlp = nn.Sequential( nn.Linear(in_dim, s_edge_width), @@ -474,7 +474,7 @@ def forward( src, dst = edge_index[0], edge_index[1] parts = [s_node[src], s_node[dst], s_edge] - if self.update_w_distance: + if self.update_w_distance_features: parts.append(distance_feat) h = torch.cat(parts, dim=-1) # (E, 2*S_node + s_edge_width (+D)) diff --git a/src/slae.py b/src/slae.py deleted file mode 100644 index 502074d..0000000 --- a/src/slae.py +++ /dev/null @@ -1,78 +0,0 @@ -# slae.py -from __future__ import annotations - -""" -SLAE (Strictly Local All-Atom Environment) base encoder implementation. - -This encoder reads pre-computed SLAE embeddings from the data and returns -them directly as scalar features with zero vector channels. Downstream -GVP message-passing layers (including protein-protein edges) provide all -geometric processing. -""" - -import torch -from torch_geometric.data import HeteroData - -from src.encoder_base import BaseProteinEncoder, register_encoder - - -@register_encoder('slae') -class SLAEEncoder(BaseProteinEncoder): - """ - SLAE encoder that reads cached embeddings from data. - - Returns (slae_embedding, empty_vectors) with output_dims = (slae_dim, 0). - No learnable parameters — all geometric processing happens in the - downstream ProteinWaterUpdate layers (pp, wp, pw, ww edges). - """ - - def __init__(self, slae_dim: int = 128): - super().__init__() - self._slae_dim = slae_dim - - @property - def output_dims(self) -> tuple[int, int]: - """Return (slae_dim, 0) — scalars only.""" - return self._slae_dim, 0 - - @property - def encoder_type(self) -> str: - return 'slae' - - def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor]: - """ - Read cached SLAE embeddings and return (s, V). - - Args: - data: HeteroData with data['protein'].slae_embedding - - Returns: - s: (N, slae_dim) — raw SLAE embeddings - V: (N, 0, 3) — empty vector features - """ - if 'slae_embedding' not in data['protein']: - raise NotImplementedError( - "SLAE encoder requires cached embeddings. " - "Please provide pre-computed slae_embedding in data['protein']. " - "Run scripts/precompute_slae_embeddings.py first." - ) - - embeddings = data['protein'].slae_embedding - V = embeddings.new_empty(embeddings.size(0), 0, 3) - return embeddings, V - - @classmethod - def from_config(cls, config: dict, device: torch.device) -> SLAEEncoder: - """ - Construct SLAEEncoder from config dict. - - Args: - config: Configuration dictionary with: - - slae_dim: SLAE embedding dimension (default: 128) - device: Device to place the encoder on - - Returns: - Instantiated SLAEEncoder - """ - slae_dim = config.get('slae_dim', 128) - return cls(slae_dim=slae_dim).to(device) diff --git a/src/slae_encoder.py b/src/slae_encoder.py deleted file mode 100644 index bf9331f..0000000 --- a/src/slae_encoder.py +++ /dev/null @@ -1,53 +0,0 @@ -""" -SLAE (Strictly Local All-Atom Environment) base encoder implementation. - -This encoder reads pre-computed SLAE embeddings from the data and returns -them directly as scalar features with zero vector channels. Downstream -GVP message-passing layers (including protein-protein edges) provide all -geometric processing. -""" -from __future__ import annotations - -import torch - -from src.encoder_base import CachedEmbeddingEncoder, register_encoder - - -@register_encoder('slae') -class SLAEEncoder(CachedEmbeddingEncoder): - """ - SLAE encoder that reads cached embeddings from data. - - Returns (slae_embedding, empty_vectors) with output_dims = (slae_dim, 0). - No learnable parameters — all geometric processing happens in the - downstream ProteinWaterUpdate layers (pp, wp, pw, ww edges). - """ - - def __init__(self, slae_dim: int = 128): - """ - Initialize SLAEEncoder. - - Args: - slae_dim: Dimension of SLAE embeddings (default: 128) - """ - super().__init__(embedding_dim=slae_dim, embedding_key='slae_embedding') - - @property - def encoder_type(self) -> str: - return 'slae' - - @classmethod - def from_config(cls, config: dict, device: torch.device) -> SLAEEncoder: - """ - Construct SLAEEncoder from config dict. - - Args: - config: Configuration dictionary with: - - slae_dim: SLAE embedding dimension (default: 128) - device: Device to place the encoder on - - Returns: - Instantiated SLAEEncoder - """ - slae_dim = config.get('slae_dim', 128) - return cls(slae_dim=slae_dim).to(device) From b6c463e00a3dcf178922ea3e3c4301688b37a42e Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Mon, 16 Mar 2026 12:02:26 -0500 Subject: [PATCH 06/19] tests and ignoring E501 in linting --- pyproject.toml | 5 ++++- tests/test_flow.py | 30 +++++++++++++++--------------- tests/test_forward.py | 20 ++++++++++---------- 3 files changed, 29 insertions(+), 26 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 1234005..cc6eb86 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -230,9 +230,12 @@ python_version = "3.12" warn_return_any = true warn_unused_configs = true +[tool.ty.rules] +unresolved-import = "ignore" + [tool.ruff.lint] fixable = ["I001", "F401", "UP"] -ignore = ["E402", "E721", "E731", "E741", "F722", "F821", "UP015", "UP037"] +ignore = ["E402", "E501", "E721", "E731", "E741", "F722", "F821", "UP015", "UP037"] select = ["E", "F", "I001", "UP"] [tool.ruff.lint.flake8-import-conventions.extend-aliases] diff --git a/tests/test_flow.py b/tests/test_flow.py index a42fcf8..411fb53 100644 --- a/tests/test_flow.py +++ b/tests/test_flow.py @@ -16,7 +16,7 @@ ProteinWaterUpdate, build_knn_edges, ) -from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder +from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder, make_gvp_encoder_data @pytest.fixture @@ -143,15 +143,15 @@ def test_with_batch(self, device): class TestMakeEncoderData: def test_basic_output(self, simple_hetero_data): - enc_data = GVPEncoder.make_encoder_data(simple_hetero_data) + enc_data = make_gvp_encoder_data(simple_hetero_data) assert isinstance(enc_data, Data) assert hasattr(enc_data, 'x') assert hasattr(enc_data, 'pos') assert hasattr(enc_data, 'edge_index') - + def test_shapes(self, simple_hetero_data): - enc_data = GVPEncoder.make_encoder_data(simple_hetero_data) + enc_data = make_gvp_encoder_data(simple_hetero_data) n_nodes = simple_hetero_data['protein'].pos.size(0) n_edges = simple_hetero_data['protein', 'pp', 'protein'].edge_index.size(1) @@ -159,20 +159,20 @@ def test_shapes(self, simple_hetero_data): assert enc_data.x.shape[0] == n_nodes assert enc_data.pos.shape == (n_nodes, 3) assert enc_data.edge_index.shape == (2, n_edges) - + def test_batch_preserved(self, batched_hetero_data): - enc_data = GVPEncoder.make_encoder_data(batched_hetero_data) + enc_data = make_gvp_encoder_data(batched_hetero_data) assert hasattr(enc_data, 'batch') assert enc_data.batch.shape[0] == batched_hetero_data['protein'].pos.size(0) - + def test_no_edges(self, device): data = HeteroData() data['protein'].pos = torch.randn(10, 3, device=device) data['protein'].x = torch.randn(10, 16, device=device) # No edges defined - enc_data = GVPEncoder.make_encoder_data(data) + enc_data = make_gvp_encoder_data(data) assert enc_data.edge_index.shape == (2, 0) @@ -262,7 +262,7 @@ def test_forward_output_shape(self, simple_hetero_data, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 8), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) @@ -283,7 +283,7 @@ def test_forward_no_water(self, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 8), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) @@ -312,7 +312,7 @@ def test_self_conditioning(self, simple_hetero_data, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 8), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) @@ -342,7 +342,7 @@ def flow_matcher(self, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 8), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) @@ -444,7 +444,7 @@ def test_distortion_enabled(self, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 8), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) @@ -476,7 +476,7 @@ def test_single_water_molecule(self, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 8), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) @@ -508,7 +508,7 @@ def test_frozen_gvp_encoder(self, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 8), - edge_scalar_in=16, + n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=True) diff --git a/tests/test_forward.py b/tests/test_forward.py index d60e0bf..9553ed4 100644 --- a/tests/test_forward.py +++ b/tests/test_forward.py @@ -11,7 +11,7 @@ from torch_geometric.data import HeteroData from src.flow import FlowMatcher, FlowWaterGVP -from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder +from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder, make_gvp_encoder_data def _iter_tensors(obj): @@ -165,7 +165,7 @@ def test_forward_pass_no_nan_with_module_hooks(device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 8), - edge_scalar_in=16, + n_edge_scalar_in=16, n_layers=2, pool_residue=False, num_edge_rbf=16, @@ -182,7 +182,7 @@ def test_forward_pass_no_nan_with_module_hooks(device): ).to(device) # Quick pre-check: protein encoder input features created from pp edges - enc_data = GVPEncoder.make_encoder_data(data) + enc_data = make_gvp_encoder_data(data) assert_edge_index_in_range(enc_data.edge_index, enc_data.x.size(0), enc_data.x.size(0), "pp edge_index") # Also validate knn edges are sane (catches orientation / k issues) @@ -231,7 +231,7 @@ def test_training_step_no_nan_tripwire(device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 8), - edge_scalar_in=16, + n_edge_scalar_in=16, n_layers=2, pool_residue=False, num_edge_rbf=16, @@ -295,7 +295,7 @@ def test_forward_with_duplicate_protein_coords_catches_nan(device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 8), - edge_scalar_in=16, + n_edge_scalar_in=16, n_layers=2, pool_residue=False, num_edge_rbf=16, @@ -339,7 +339,7 @@ def test_forward_with_duplicate_protein_coords_localizes_nan(device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, hidden_dims=(64, 8), - edge_scalar_in=16, + n_edge_scalar_in=16, n_layers=2, pool_residue=False, num_edge_rbf=16, @@ -356,7 +356,7 @@ def test_forward_with_duplicate_protein_coords_localizes_nan(device): ).to(device) # ---- Pre-check: encoder input edges ---- - enc_data = GVPEncoder.make_encoder_data(data) + enc_data = make_gvp_encoder_data(data) for tensor in _iter_tensors(enc_data): assert torch.isfinite(tensor).all() @@ -496,7 +496,7 @@ def test_integration_trajectory_length(self, device): data = make_batched_hetero(device, n_graphs=1, n_protein_per=24, n_water_per=12) base_encoder = ProteinGVPEncoder( - node_scalar_in=16, hidden_dims=(64, 8), edge_scalar_in=16, + node_scalar_in=16, hidden_dims=(64, 8), n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) @@ -555,7 +555,7 @@ def test_velocity_field_finite(self, device): data = make_batched_hetero(device, n_graphs=1, n_protein_per=24, n_water_per=12) base_encoder = ProteinGVPEncoder( - node_scalar_in=16, hidden_dims=(64, 8), edge_scalar_in=16, + node_scalar_in=16, hidden_dims=(64, 8), n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) @@ -585,7 +585,7 @@ def test_velocity_field_changes_with_t(self, device): data = make_batched_hetero(device, n_graphs=1, n_protein_per=24, n_water_per=12) base_encoder = ProteinGVPEncoder( - node_scalar_in=16, hidden_dims=(64, 8), edge_scalar_in=16, + node_scalar_in=16, hidden_dims=(64, 8), n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) From c58125877a48f626b93ebf5534e7dc5a8bcacec9 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Mon, 16 Mar 2026 15:34:35 -0500 Subject: [PATCH 07/19] fixing ruff checks --- scripts/generate_water_plots.py | 2 +- scripts/inference.py | 1 + scripts/parse_edia_json.py | 94 ------ scripts/qc_symmetry_mates.py | 311 -------------------- scripts/qc_waters.py | 444 ----------------------------- scripts/run_edia_parallel.sh | 145 ---------- scripts/train.py | 7 +- tests/conftest.py | 1 + tests/test_dataset.py | 10 +- tests/test_embedding_generation.py | 4 +- tests/test_encoder.py | 3 +- tests/test_flow.py | 6 +- tests/test_forward.py | 2 +- tests/test_gvp.py | 3 +- tests/test_utils.py | 1 + 15 files changed, 21 insertions(+), 1013 deletions(-) delete mode 100755 scripts/parse_edia_json.py delete mode 100755 scripts/qc_symmetry_mates.py delete mode 100755 scripts/qc_waters.py delete mode 100755 scripts/run_edia_parallel.sh diff --git a/scripts/generate_water_plots.py b/scripts/generate_water_plots.py index 0e95ce0..0a0cfbb 100644 --- a/scripts/generate_water_plots.py +++ b/scripts/generate_water_plots.py @@ -14,7 +14,7 @@ """ import argparse -from concurrent.futures import ProcessPoolExecutor, as_completed +from concurrent.futures import as_completed, ProcessPoolExecutor from pathlib import Path import matplotlib.pyplot as plt diff --git a/scripts/inference.py b/scripts/inference.py index be3f466..c5c1fb8 100644 --- a/scripts/inference.py +++ b/scripts/inference.py @@ -39,6 +39,7 @@ setup_logging_for_tqdm, ) + # Configure logging to work with tqdm progress bars setup_logging_for_tqdm() diff --git a/scripts/parse_edia_json.py b/scripts/parse_edia_json.py deleted file mode 100755 index 5de0a77..0000000 --- a/scripts/parse_edia_json.py +++ /dev/null @@ -1,94 +0,0 @@ -#!/usr/bin/env python3 -""" -Parse density-fitness (EDIA) JSON output to CSV format. -Usage: python parse_edia_json.py [output_file] -""" - -import json -import sys -from pathlib import Path - -import pandas as pd -from loguru import logger - - -def parse_protein_data(file_path): - """ - Parse the JSON file containing protein residue data. - - Args: - file_path (str): Path to the JSON file - - Returns: - pandas.DataFrame: Parsed data as a DataFrame - """ - try: - with open(file_path, 'r') as f: - content = f.read().strip() - - # Handle malformed JSON (missing opening/closing brackets) - if not content.startswith('['): - content = '[' + content - if not content.endswith(']'): - content = content + ']' - - data = json.loads(content) - - df = pd.DataFrame(data) - - # Flatten the nested 'pdb' column if it exists - if 'pdb' in df.columns: - pdb_df = pd.json_normalize(df['pdb']) - pdb_df.columns = ['pdb_' + col for col in pdb_df.columns] - df = pd.concat([df.drop('pdb', axis=1), pdb_df], axis=1) - - return df - - except json.JSONDecodeError as e: - logger.error(f"JSON parsing error: {e}") - return None - except FileNotFoundError: - logger.warning(f"File not found: {file_path}") - return None - except Exception as e: - logger.error(f"Error parsing file: {e}") - return None - - -def save_to_csv(df, output_path): - """ - Save the DataFrame to CSV. - - Args: - df (pandas.DataFrame): Data to save - output_path (str): Output CSV file path - """ - if df is not None: - df.to_csv(output_path, index=False) - logger.info(f"Data saved to: {output_path}") - - -def main(): - if len(sys.argv) < 2: - logger.info("Usage: python parse_edia_json.py [output_file]") - sys.exit(1) - - input_file = sys.argv[1] - output_file = sys.argv[2] if len(sys.argv) > 2 else None - - df = parse_protein_data(input_file) - - if df is not None: - if output_file: - save_to_csv(df, output_file) - else: - # Print summary if no output file specified - logger.info(f"Total residues: {len(df)}") - logger.info(f"Columns: {list(df.columns)}") - else: - logger.warning("Failed to parse the file.") - sys.exit(1) - - -if __name__ == "__main__": - main() diff --git a/scripts/qc_symmetry_mates.py b/scripts/qc_symmetry_mates.py deleted file mode 100755 index 116d88a..0000000 --- a/scripts/qc_symmetry_mates.py +++ /dev/null @@ -1,311 +0,0 @@ -""" -Quality control script for symmetry mates. Completely generated by claude code and lightly edited. - -This script checks: -1. Are mates being generated correctly by PyMOL? -2. Are mate atoms actually near ASU (within cutoff)? -3. Are the right number of mates generated? -4. Visualize ASU + mates for manual inspection - -Usage: - python scripts/qc_symmetry_mates.py \ - --processed_dir /path/to/cache \ - --base_pdb_dir /path/to/pdbs \ - --num_samples 5 \ - --output_dir qc_mates -""" - -import argparse -import sys -from pathlib import Path - -sys.path.insert(0, str(Path(__file__).parent.parent)) - -import matplotlib.pyplot as plt -import numpy as np -import pandas as pd - -# For re-running mate detection -import pymol2 -import torch -from loguru import logger -from mpl_toolkits.mplot3d import Axes3D -from tqdm import tqdm - -from src.dataset import get_crystal_contacts_pymol - - -def analyze_mate_distances(cached_data, cutoff=5.0): - """ - Analyze distances between ASU and symmetry mates. - - Returns: - dict with statistics about mate distances - """ - protein_pos = cached_data['protein_pos'].numpy() - mate_pos = cached_data['mate_pos'].numpy() - - if mate_pos.shape[0] == 0: - return { - 'num_mates': 0, - 'min_dist': None, - 'max_dist': None, - 'mean_dist': None, - 'within_cutoff': 0, - } - - # Compute pairwise distances between ASU and mates - # This is expensive for large structures, so we sample - n_asu = min(protein_pos.shape[0], 1000) - n_mate = min(mate_pos.shape[0], 1000) - - asu_sample = protein_pos[np.random.choice(protein_pos.shape[0], n_asu, replace=False)] - mate_sample = mate_pos[np.random.choice(mate_pos.shape[0], n_mate, replace=False)] - - # Compute distances - dists = np.linalg.norm(asu_sample[:, None] - mate_sample[None, :], axis=-1) - min_dists_per_mate = dists.min(axis=0) - - return { - 'num_mates': mate_pos.shape[0], - 'num_asu': protein_pos.shape[0], - 'min_dist': float(min_dists_per_mate.min()), - 'max_dist': float(min_dists_per_mate.max()), - 'mean_dist': float(min_dists_per_mate.mean()), - 'median_dist': float(np.median(min_dists_per_mate)), - 'within_cutoff': int((min_dists_per_mate <= cutoff).sum()), - 'percent_within_cutoff': float((min_dists_per_mate <= cutoff).sum() / len(min_dists_per_mate) * 100), - } - - -def visualize_mates(cached_data, pdb_id, output_dir): - """ - Create 3D visualization of ASU + mates. - - Saves plot to output_dir. - """ - protein_pos = cached_data['protein_pos'].numpy() - mate_pos = cached_data['mate_pos'].numpy() - water_pos = cached_data['water_pos'].numpy() - - fig = plt.figure(figsize=(12, 10)) - ax = fig.add_subplot(111, projection='3d') - - # Subsample for visualization - n_asu_viz = min(protein_pos.shape[0], 500) - n_mate_viz = min(mate_pos.shape[0], 500) - n_water_viz = min(water_pos.shape[0], 100) - - asu_viz = protein_pos[np.random.choice(protein_pos.shape[0], n_asu_viz, replace=False)] - if mate_pos.shape[0] > 0: - mate_viz = mate_pos[np.random.choice(mate_pos.shape[0], n_mate_viz, replace=False)] - else: - mate_viz = np.zeros((0, 3)) - - if water_pos.shape[0] > 0: - water_viz = water_pos[np.random.choice(water_pos.shape[0], n_water_viz, replace=False)] - else: - water_viz = np.zeros((0, 3)) - - # Plot ASU (blue), mates (red), waters (cyan) - ax.scatter(asu_viz[:, 0], asu_viz[:, 1], asu_viz[:, 2], - c='blue', marker='o', s=20, alpha=0.6, label=f'ASU ({protein_pos.shape[0]} atoms)') - - if mate_viz.shape[0] > 0: - ax.scatter(mate_viz[:, 0], mate_viz[:, 1], mate_viz[:, 2], - c='red', marker='^', s=20, alpha=0.6, label=f'Mates ({mate_pos.shape[0]} atoms)') - - if water_viz.shape[0] > 0: - ax.scatter(water_viz[:, 0], water_viz[:, 1], water_viz[:, 2], - c='cyan', marker='*', s=10, alpha=0.4, label=f'Waters ({water_pos.shape[0]})') - - ax.set_xlabel('X (Å)') - ax.set_ylabel('Y (Å)') - ax.set_zlabel('Z (Å)') - ax.set_title(f'{pdb_id} - ASU + Symmetry Mates') - ax.legend() - - # Equal aspect ratio - max_range = np.array([asu_viz.max(axis=0) - asu_viz.min(axis=0)]).max() / 2.0 - mid_x = (asu_viz[:, 0].max() + asu_viz[:, 0].min()) * 0.5 - mid_y = (asu_viz[:, 1].max() + asu_viz[:, 1].min()) * 0.5 - mid_z = (asu_viz[:, 2].max() + asu_viz[:, 2].min()) * 0.5 - ax.set_xlim(mid_x - max_range, mid_x + max_range) - ax.set_ylim(mid_y - max_range, mid_y + max_range) - ax.set_zlim(mid_z - max_range, mid_z + max_range) - - output_path = Path(output_dir) / f"{pdb_id}_mates.png" - plt.savefig(output_path, dpi=150, bbox_inches='tight') - plt.close() - - logger.info(f" Saved visualization to {output_path}") - - -def compare_mate_detection(pdb_path, cache_path, cutoff=5.0): - """ - Re-run mate detection and compare with cached results. - - Returns: - dict with comparison statistics - """ - # Load cached mates - cached = torch.load(cache_path, weights_only=False) - cached_mate_pos = cached['mate_pos'].numpy() - - # Re-run PyMOL mate detection - crystal_data = get_crystal_contacts_pymol(str(pdb_path), cutoff=cutoff) - fresh_mate_coords = crystal_data['mate_coords'] - - return { - 'cached_num_mates': cached_mate_pos.shape[0], - 'fresh_num_mates': fresh_mate_coords.shape[0], - 'match': cached_mate_pos.shape[0] == fresh_mate_coords.shape[0], - } - - -def main(): - parser = argparse.ArgumentParser(description="QC symmetry mates") - parser.add_argument("--processed_dir", type=str, required=True) - parser.add_argument("--base_pdb_dir", type=str, - default="/sb/wankowicz_lab/data/srivasv/pdb_redo_data") - parser.add_argument("--num_samples", type=int, default=10, - help="Number of samples to analyze in detail") - parser.add_argument("--output_dir", type=str, default="qc_mates") - parser.add_argument("--cutoff", type=float, default=5.0) - parser.add_argument("--recompute_mates", action="store_true", - help="Re-run PyMOL and compare with cached mates") - - args = parser.parse_args() - - processed_dir = Path(args.processed_dir) - output_dir = Path(args.output_dir) - output_dir.mkdir(parents=True, exist_ok=True) - - # Find all cache files - cache_files = sorted(processed_dir.glob("*.pt")) - logger.info(f"Found {len(cache_files)} cache files") - - # Sample for detailed analysis - if args.num_samples < len(cache_files): - sample_files = np.random.choice(cache_files, args.num_samples, replace=False) - else: - sample_files = cache_files - - # Statistics - all_stats = [] - - logger.info("\n" + "="*80) - logger.info("SYMMETRY MATE QC ANALYSIS") - logger.info("="*80) - - for cache_path in tqdm(sample_files, desc="Analyzing samples"): - cache_key = cache_path.stem - parts = cache_key.split('_') - pdb_id = parts[0] - - # Load cached data - cached = torch.load(cache_path, weights_only=False) - - # Analyze mate distances - stats = analyze_mate_distances(cached, cutoff=args.cutoff) - stats['pdb_id'] = cache_key - - # Visualize - visualize_mates(cached, cache_key, output_dir) - - # Optionally recompute mates - if args.recompute_mates: - pdb_path = Path(args.base_pdb_dir) / pdb_id / f"{pdb_id}_final.pdb" - if pdb_path.exists(): - comparison = compare_mate_detection(pdb_path, cache_path, args.cutoff) - stats.update(comparison) - - all_stats.append(stats) - - # Create summary report - df = pd.DataFrame(all_stats) - - logger.info("\n" + "="*80) - logger.info("SUMMARY STATISTICS") - logger.info("="*80) - - logger.info(f"\nTotal structures analyzed: {len(df)}") - logger.info(f"\nStructures with mates: {(df['num_mates'] > 0).sum()} ({(df['num_mates'] > 0).sum() / len(df) * 100:.1f}%)") - logger.info(f"Structures without mates: {(df['num_mates'] == 0).sum()}") - - if (df['num_mates'] > 0).any(): - mate_df = df[df['num_mates'] > 0] - logger.info(f"\nFor structures WITH mates:") - logger.info(f" Number of mate atoms: {mate_df['num_mates'].describe()}") - logger.info(f"\n Distance to nearest ASU atom (Å):") - logger.info(f" Min: {mate_df['min_dist'].min():.2f}") - logger.info(f" Max: {mate_df['max_dist'].max():.2f}") - logger.info(f" Mean: {mate_df['mean_dist'].mean():.2f}") - logger.info(f" Median: {mate_df['median_dist'].median():.2f}") - logger.info(f"\n Percent of mates within cutoff ({args.cutoff}Å): {mate_df['percent_within_cutoff'].mean():.1f}%") - - # Check for issues - logger.warning(f"\n⚠️ POTENTIAL ISSUES:") - far_mates = mate_df[mate_df['min_dist'] > args.cutoff] - if len(far_mates) > 0: - logger.info(f" {len(far_mates)} structures have mates farther than cutoff:") - for _, row in far_mates.iterrows(): - logger.info(f" {row['pdb_id']}: min_dist = {row['min_dist']:.2f}Å") - - if args.recompute_mates and 'match' in df.columns: - logger.info(f"\n MATE RECOMPUTATION CHECK:") - logger.info(f" Matches cached: {df['match'].sum()} / {len(df)}") - if not df['match'].all(): - logger.warning(f" ⚠️ MISMATCHES FOUND:") - mismatches = df[~df['match']] - for _, row in mismatches.iterrows(): - logger.info(f" {row['pdb_id']}: cached={row['cached_num_mates']}, fresh={row['fresh_num_mates']}") - - # Save report - report_path = output_dir / "mate_qc_report.csv" - df.to_csv(report_path, index=False) - logger.info(f"\n✓ Full report saved to {report_path}") - - # Save summary plot - if (df['num_mates'] > 0).any(): - fig, axes = plt.subplots(2, 2, figsize=(12, 10)) - - # Number of mates histogram - axes[0, 0].hist(df['num_mates'], bins=30, edgecolor='black') - axes[0, 0].set_xlabel('Number of mate atoms') - axes[0, 0].set_ylabel('Count') - axes[0, 0].set_title('Distribution of Mate Atoms') - - # Distance to ASU - mate_df = df[df['num_mates'] > 0] - axes[0, 1].hist(mate_df['min_dist'], bins=30, edgecolor='black') - axes[0, 1].axvline(args.cutoff, color='red', linestyle='--', label=f'Cutoff ({args.cutoff}Å)') - axes[0, 1].set_xlabel('Min distance to ASU (Å)') - axes[0, 1].set_ylabel('Count') - axes[0, 1].set_title('Mate-ASU Distances') - axes[0, 1].legend() - - # Percent within cutoff - axes[1, 0].hist(mate_df['percent_within_cutoff'], bins=30, edgecolor='black') - axes[1, 0].set_xlabel('% mates within cutoff') - axes[1, 0].set_ylabel('Count') - axes[1, 0].set_title('Mates Within Cutoff Distribution') - - # ASU vs Mate size - axes[1, 1].scatter(mate_df['num_asu'], mate_df['num_mates'], alpha=0.6) - axes[1, 1].set_xlabel('Number of ASU atoms') - axes[1, 1].set_ylabel('Number of mate atoms') - axes[1, 1].set_title('ASU Size vs Mate Size') - - plt.tight_layout() - summary_plot_path = output_dir / "mate_qc_summary.png" - plt.savefig(summary_plot_path, dpi=150, bbox_inches='tight') - logger.info(f"✓ Summary plot saved to {summary_plot_path}") - - logger.info("\n" + "="*80) - logger.info("QC COMPLETE") - logger.info("="*80) - - -if __name__ == "__main__": - main() diff --git a/scripts/qc_waters.py b/scripts/qc_waters.py deleted file mode 100755 index 0c58a89..0000000 --- a/scripts/qc_waters.py +++ /dev/null @@ -1,444 +0,0 @@ -""" -Quality control script for water molecules. Completely generated by claude code and lightly edited. - -This script checks: -1. Distance of waters to nearest protein atom -2. Distribution of waters in ASU -3. Waters near crystal contacts (interface waters) -4. Water clustering -5. Potential issues (waters too far, overlapping waters, etc.) - -Usage: - python scripts/qc_waters.py \ - --processed_dir /path/to/cache \ - --base_pdb_dir /path/to/pdbs \ - --num_samples 10 \ - --output_dir qc_waters -""" - -import argparse -import sys -from pathlib import Path - -sys.path.insert(0, str(Path(__file__).parent.parent)) - -# For parsing PDB metadata -import biotite.structure as bts -import matplotlib.pyplot as plt -import numpy as np -import pandas as pd -import torch -from biotite.structure.io.pdb import PDBFile, get_structure -from loguru import logger -from mpl_toolkits.mplot3d import Axes3D -from scipy.spatial import cKDTree -from tqdm import tqdm - - -def analyze_water_protein_distances(protein_pos, water_pos): - """ - Compute distances between waters and nearest protein atoms. - - Returns: - dict with distance statistics - """ - if water_pos.shape[0] == 0: - return { - 'num_waters': 0, - 'min_dist': None, - 'max_dist': None, - 'mean_dist': None, - } - - # Build KD-tree for protein - tree = cKDTree(protein_pos) - - # Find nearest protein atom for each water - dists, _ = tree.query(water_pos, k=1) - - return { - 'num_waters': water_pos.shape[0], - 'min_dist': float(dists.min()), - 'max_dist': float(dists.max()), - 'mean_dist': float(dists.mean()), - 'median_dist': float(np.median(dists)), - 'std_dist': float(dists.std()), - 'within_3A': int((dists <= 3.0).sum()), - 'within_4A': int((dists <= 4.0).sum()), - 'beyond_6A': int((dists > 6.0).sum()), - } - - -def analyze_water_clustering(water_pos, cluster_threshold=3.5): - """ - Detect clusters of water molecules. - - Waters within cluster_threshold of each other form clusters. - - Returns: - dict with clustering statistics - """ - if water_pos.shape[0] == 0: - return {'num_clusters': 0, 'largest_cluster': 0} - - # Build KD-tree for waters - tree = cKDTree(water_pos) - - # Find neighbors within threshold - neighbor_counts = tree.query_ball_tree(tree, r=cluster_threshold) - neighbor_counts = [len(n) - 1 for n in neighbor_counts] # -1 to exclude self - - return { - 'isolated_waters': int((np.array(neighbor_counts) == 0).sum()), - 'mean_neighbors': float(np.mean(neighbor_counts)), - 'max_neighbors': int(np.max(neighbor_counts)), - } - - -def detect_water_issues(protein_pos, water_pos): - """ - Detect potential issues with water placement. - - Returns: - dict with issue counts - """ - issues = { - 'too_far': [], # Waters > 6Å from protein - 'too_close': [], # Waters < 2.0Å from protein - 'overlapping': [], # Waters < 1.5Å from each other - } - - if water_pos.shape[0] == 0: - return issues - - # Check protein distances - protein_tree = cKDTree(protein_pos) - dists, _ = protein_tree.query(water_pos, k=1) - - too_far_idx = np.where(dists > 6.0)[0] - too_close_idx = np.where(dists < 2.0)[0] - - issues['too_far'] = too_far_idx.tolist() - issues['too_close'] = too_close_idx.tolist() - - # Check water-water distances - if water_pos.shape[0] > 1: - water_tree = cKDTree(water_pos) - pairs = water_tree.query_pairs(r=1.5, output_type='ndarray') - issues['overlapping'] = pairs.tolist() - - return issues - - -def get_water_metadata(pdb_path, chain_filter=None): - """ - Extract water metadata from PDB file (B-factors, occupancy, etc.). - - Returns: - dict with metadata arrays - """ - try: - pdb_file = PDBFile.read(pdb_path) - atoms = get_structure(pdb_file, model=1, altloc="occupancy") - - if chain_filter is not None: - mask = np.isin(atoms.chain_id, np.array(chain_filter, dtype=atoms.chain_id.dtype)) - atoms = atoms[mask] - - # Filter for waters - water_mask = (atoms.res_name == "HOH") | (atoms.res_name == "WAT") - water_atoms = atoms[water_mask] - - if len(water_atoms) == 0: - return None - - return { - 'b_factors': water_atoms.b_factor, - 'occupancy': water_atoms.occupancy, - 'num_waters': len(water_atoms), - } - except Exception as e: - logger.warning(f" Warning: Could not parse PDB metadata: {e}") - return None - - -def visualize_waters(cached_data, pdb_id, output_dir, show_issues=None): - """ - Create 3D visualization of protein + waters with issues highlighted. - - Args: - show_issues: dict from detect_water_issues() - """ - protein_pos = cached_data['protein_pos'].numpy() - water_pos = cached_data['water_pos'].numpy() - - if water_pos.shape[0] == 0: - logger.info(f" Skipping {pdb_id}: no waters") - return - - fig = plt.figure(figsize=(14, 10)) - ax = fig.add_subplot(111, projection='3d') - - # Subsample protein for visualization - n_protein_viz = min(protein_pos.shape[0], 500) - protein_viz = protein_pos[np.random.choice(protein_pos.shape[0], n_protein_viz, replace=False)] - - # Plot protein (gray) - ax.scatter(protein_viz[:, 0], protein_viz[:, 1], protein_viz[:, 2], - c='gray', marker='o', s=1, alpha=0.3, label='Protein') - - # Plot waters - if show_issues is not None: - # Color waters by issue type - normal_idx = set(range(water_pos.shape[0])) - set(show_issues['too_far']) - set(show_issues['too_close']) - normal_idx = list(normal_idx) - - if normal_idx: - ax.scatter(water_pos[normal_idx, 0], water_pos[normal_idx, 1], water_pos[normal_idx, 2], - c='cyan', marker='*', s=50, alpha=0.8, label=f'Normal waters ({len(normal_idx)})') - - if show_issues['too_far']: - ax.scatter(water_pos[show_issues['too_far'], 0], - water_pos[show_issues['too_far'], 1], - water_pos[show_issues['too_far'], 2], - c='red', marker='X', s=100, alpha=1.0, - label=f'Too far (>{6.0}Å) ({len(show_issues["too_far"])})') - - if show_issues['too_close']: - ax.scatter(water_pos[show_issues['too_close'], 0], - water_pos[show_issues['too_close'], 1], - water_pos[show_issues['too_close'], 2], - c='orange', marker='D', s=100, alpha=1.0, - label=f'Too close (<2Å) ({len(show_issues["too_close"])})') - else: - ax.scatter(water_pos[:, 0], water_pos[:, 1], water_pos[:, 2], - c='cyan', marker='*', s=50, alpha=0.8, label=f'Waters ({water_pos.shape[0]})') - - ax.set_xlabel('X (Å)') - ax.set_ylabel('Y (Å)') - ax.set_zlabel('Z (Å)') - ax.set_title(f'{pdb_id} - Water Placement QC') - ax.legend() - - output_path = Path(output_dir) / f"{pdb_id}_waters.png" - plt.savefig(output_path, dpi=150, bbox_inches='tight') - plt.close() - - logger.info(f" Saved visualization to {output_path}") - - -def main(): - parser = argparse.ArgumentParser(description="QC water molecules") - parser.add_argument("--processed_dir", type=str, required=True) - parser.add_argument("--base_pdb_dir", type=str, - default="/sb/wankowicz_lab/data/srivasv/pdb_redo_data") - parser.add_argument("--num_samples", type=int, default=10, - help="Number of samples to visualize") - parser.add_argument("--output_dir", type=str, default="qc_waters") - parser.add_argument("--analyze_all", action="store_true", - help="Analyze all cache files (not just samples)") - parser.add_argument("--check_pdb_metadata", action="store_true", - help="Parse PDB files for B-factor/occupancy (slower)") - - args = parser.parse_args() - - processed_dir = Path(args.processed_dir) - output_dir = Path(args.output_dir) - output_dir.mkdir(parents=True, exist_ok=True) - - # Find all cache files - cache_files = sorted(processed_dir.glob("*.pt")) - logger.info(f"Found {len(cache_files)} cache files") - - # Determine which files to analyze - if args.analyze_all: - analyze_files = cache_files - else: - analyze_files = cache_files - - # Sample for visualization - if args.num_samples < len(cache_files): - viz_files = np.random.choice(cache_files, args.num_samples, replace=False) - else: - viz_files = cache_files - - # Statistics - all_stats = [] - all_issues = [] - - logger.info("\n" + "="*80) - logger.info("WATER MOLECULE QC ANALYSIS") - logger.info("="*80) - - for cache_path in tqdm(analyze_files, desc="Analyzing waters"): - cache_key = cache_path.stem - parts = cache_key.split('_') - pdb_id = parts[0] - chain_id = parts[-1] if len(parts) >= 3 else None - - # Load cached data - cached = torch.load(cache_path, weights_only=False) - - protein_pos = cached['protein_pos'].numpy() - water_pos = cached['water_pos'].numpy() - - # Analyze distances - dist_stats = analyze_water_protein_distances(protein_pos, water_pos) - dist_stats['pdb_id'] = cache_key - - # Analyze clustering - cluster_stats = analyze_water_clustering(water_pos) - dist_stats.update(cluster_stats) - - # Detect issues - issues = detect_water_issues(protein_pos, water_pos) - dist_stats['num_too_far'] = len(issues['too_far']) - dist_stats['num_too_close'] = len(issues['too_close']) - dist_stats['num_overlapping'] = len(issues['overlapping']) - - # Get PDB metadata if requested - if args.check_pdb_metadata and water_pos.shape[0] > 0: - pdb_path = Path(args.base_pdb_dir) / pdb_id / f"{pdb_id}_final.pdb" - if pdb_path.exists(): - metadata = get_water_metadata(str(pdb_path), chain_filter=[chain_id] if chain_id else None) - if metadata is not None: - dist_stats['mean_b_factor'] = float(metadata['b_factors'].mean()) - dist_stats['median_b_factor'] = float(np.median(metadata['b_factors'])) - dist_stats['mean_occupancy'] = float(metadata['occupancy'].mean()) - - all_stats.append(dist_stats) - - # Store issues for problematic structures - if dist_stats['num_too_far'] > 0 or dist_stats['num_too_close'] > 0: - all_issues.append({'pdb_id': cache_key, **issues}) - - # Visualize samples - if cache_path in viz_files: - visualize_waters(cached, cache_key, output_dir, show_issues=issues) - - # Create summary report - df = pd.DataFrame(all_stats) - - logger.info("\n" + "="*80) - logger.info("SUMMARY STATISTICS") - logger.info("="*80) - - logger.info(f"\nTotal structures analyzed: {len(df)}") - logger.info(f"Structures with waters: {(df['num_waters'] > 0).sum()} ({(df['num_waters'] > 0).sum() / len(df) * 100:.1f}%)") - logger.info(f"Structures without waters: {(df['num_waters'] == 0).sum()}") - - if (df['num_waters'] > 0).any(): - water_df = df[df['num_waters'] > 0] - - logger.info(f"\nFor structures WITH waters:") - logger.info(f" Number of waters: {water_df['num_waters'].describe()}") - - logger.info(f"\n Distance to nearest protein atom (Å):") - logger.info(f" Min: {water_df['min_dist'].min():.2f}") - logger.info(f" Max: {water_df['max_dist'].max():.2f}") - logger.info(f" Mean: {water_df['mean_dist'].mean():.2f} ± {water_df['std_dist'].mean():.2f}") - logger.info(f" Median: {water_df['median_dist'].median():.2f}") - - logger.info(f"\n Water placement:") - logger.info(f" Within 3Å: {water_df['within_3A'].sum()} ({water_df['within_3A'].sum() / water_df['num_waters'].sum() * 100:.1f}% of all waters)") - logger.info(f" Within 4Å: {water_df['within_4A'].sum()} ({water_df['within_4A'].sum() / water_df['num_waters'].sum() * 100:.1f}%)") - logger.info(f" Beyond 6Å: {water_df['beyond_6A'].sum()} ({water_df['beyond_6A'].sum() / water_df['num_waters'].sum() * 100:.1f}%)") - - logger.info(f"\n Water clustering:") - logger.info(f" Isolated waters (no neighbors within 3.5Å): {water_df['isolated_waters'].sum()}") - logger.info(f" Mean neighbors per water: {water_df['mean_neighbors'].mean():.2f}") - logger.info(f" Max neighbors: {water_df['max_neighbors'].max()}") - - if args.check_pdb_metadata and 'mean_b_factor' in water_df.columns: - logger.info(f"\n B-factors:") - logger.info(f" Mean: {water_df['mean_b_factor'].mean():.2f}") - logger.info(f" Median: {water_df['median_b_factor'].median():.2f}") - logger.info(f"\n Occupancy:") - logger.info(f" Mean: {water_df['mean_occupancy'].mean():.3f}") - - # Check for issues - logger.warning(f"\n⚠️ POTENTIAL ISSUES:") - logger.info(f" Structures with waters too far (>6Å): {(water_df['num_too_far'] > 0).sum()}") - logger.info(f" Structures with waters too close (<2Å): {(water_df['num_too_close'] > 0).sum()}") - logger.info(f" Structures with overlapping waters (<1.5Å): {(water_df['num_overlapping'] > 0).sum()}") - - if len(all_issues) > 0: - logger.info(f"\n Top problematic structures:") - issue_counts = [(issue['pdb_id'], - len(issue['too_far']) + len(issue['too_close']) + len(issue['overlapping'])) - for issue in all_issues] - issue_counts.sort(key=lambda x: x[1], reverse=True) - for pdb_id, count in issue_counts[:10]: - logger.info(f" {pdb_id}: {count} issues") - - # Save report - report_path = output_dir / "water_qc_report.csv" - df.to_csv(report_path, index=False) - logger.info(f"\n✓ Full report saved to {report_path}") - - # Save summary plots - if (df['num_waters'] > 0).any(): - water_df = df[df['num_waters'] > 0] - - fig, axes = plt.subplots(2, 3, figsize=(15, 10)) - - # Number of waters - axes[0, 0].hist(water_df['num_waters'], bins=30, edgecolor='black') - axes[0, 0].set_xlabel('Number of waters') - axes[0, 0].set_ylabel('Count') - axes[0, 0].set_title('Distribution of Water Count') - - # Distance to protein - axes[0, 1].hist(water_df['mean_dist'], bins=30, edgecolor='black') - axes[0, 1].axvline(3.0, color='green', linestyle='--', label='3Å (H-bond)') - axes[0, 1].axvline(6.0, color='red', linestyle='--', label='6Å (far)') - axes[0, 1].set_xlabel('Mean distance to protein (Å)') - axes[0, 1].set_ylabel('Count') - axes[0, 1].set_title('Water-Protein Distances') - axes[0, 1].legend() - - # Isolated waters - axes[0, 2].hist(water_df['isolated_waters'], bins=20, edgecolor='black') - axes[0, 2].set_xlabel('Number of isolated waters') - axes[0, 2].set_ylabel('Count') - axes[0, 2].set_title('Isolated Waters per Structure') - - # Mean neighbors - axes[1, 0].hist(water_df['mean_neighbors'], bins=30, edgecolor='black') - axes[1, 0].set_xlabel('Mean neighbors per water') - axes[1, 0].set_ylabel('Count') - axes[1, 0].set_title('Water Clustering') - - # Issues - issue_data = [water_df['num_too_far'].sum(), - water_df['num_too_close'].sum(), - water_df['num_overlapping'].sum()] - axes[1, 1].bar(['Too far\n(>6Å)', 'Too close\n(<2Å)', 'Overlapping\n(<1.5Å)'], - issue_data, color=['red', 'orange', 'purple']) - axes[1, 1].set_ylabel('Total count') - axes[1, 1].set_title('Water Placement Issues') - - # B-factors if available - if 'mean_b_factor' in water_df.columns: - axes[1, 2].hist(water_df['mean_b_factor'].dropna(), bins=30, edgecolor='black') - axes[1, 2].set_xlabel('Mean B-factor') - axes[1, 2].set_ylabel('Count') - axes[1, 2].set_title('Water B-factor Distribution') - else: - axes[1, 2].text(0.5, 0.5, 'B-factors not analyzed\n(use --check_pdb_metadata)', - ha='center', va='center') - axes[1, 2].set_xticks([]) - axes[1, 2].set_yticks([]) - - plt.tight_layout() - summary_plot_path = output_dir / "water_qc_summary.png" - plt.savefig(summary_plot_path, dpi=150, bbox_inches='tight') - logger.info(f"✓ Summary plot saved to {summary_plot_path}") - - logger.info("\n" + "="*80) - logger.info("QC COMPLETE") - logger.info("="*80) - - -if __name__ == "__main__": - main() diff --git a/scripts/run_edia_parallel.sh b/scripts/run_edia_parallel.sh deleted file mode 100755 index 24be9f0..0000000 --- a/scripts/run_edia_parallel.sh +++ /dev/null @@ -1,145 +0,0 @@ -#!/bin/bash -# run_edia_parallel.sh -# Runs EDIA (density-fitness) in parallel across all PDBs in water_pdbs.txt -# Usage: ./run_edia_parallel.sh [num_jobs] - -set -euo pipefail - -# Source sbgrid for density-fitness (EDIA) -# Temporarily disable 'nounset' (-u) and 'errexit' (-e) because sbgrid.shrc -# uses unset variables and may return non-zero exit status -set +eu -source /programs/sbgrid.shrc -set -eu - -# Configuration -PARENT_DIR="/sb/wankowicz_lab/data/srivasv/pdb_redo_data" -OUTPUT_DIR="/sb/wankowicz_lab/data/srivasv/edia_results" -SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" -FULL_ID_FILE="$SCRIPT_DIR/../splits/water_pdbs.txt" -PARSER="$SCRIPT_DIR/parse_edia_json.py" - -# Number of parallel jobs (default: 32, or pass as argument) -NUM_JOBS="${1:-32}" - -# Validate input file exists -if [[ ! -f "$FULL_ID_FILE" ]]; then - echo "Error: Input file not found: $FULL_ID_FILE" - exit 1 -fi - -# Create output directory -mkdir -p "$OUTPUT_DIR" - -# Count total PDBs (excluding comments and empty lines) -TOTAL_PDBS=$(grep -cvE '^\s*$|^\s*#' "$FULL_ID_FILE" || echo "0") - -echo "==============================================" -echo "EDIA (density-fitness) Parallel Pipeline" -echo "==============================================" -echo "Data directory: $PARENT_DIR" -echo "Output directory: $OUTPUT_DIR" -echo "Input file: $FULL_ID_FILE" -echo "Total PDBs: $TOTAL_PDBS" -echo "Parallel jobs: $NUM_JOBS" -echo "CPUs available: $(nproc)" -echo "==============================================" -echo "" - -# Export variables for use in subshells -export PARENT_DIR OUTPUT_DIR PARSER - -# Define the worker function -process_pdb() { - local pdb_line="$1" - - # Strip whitespace, remove .pdb extension, remove _final or _final_X suffix - local pdb=$(echo "$pdb_line" | tr -d '[:space:]' | sed -E 's/\.pdb$//; s/_final(_[A-Za-z0-9])?$//') - - # Skip empty lines or comments - [[ -z "$pdb" || "${pdb:0:1}" == "#" ]] && return 0 - - # Set up paths - local pdb_lower=$(echo "$pdb" | tr '[:upper:]' '[:lower:]') - local pdb_path="$PARENT_DIR/$pdb_lower" - - # Check if directory exists - if [[ ! -d "$pdb_path" ]]; then - echo "[SKIP] Path not found for $pdb_lower" - return 0 - fi - - # Define input file paths - local FINAL_MTZ_FILE="$pdb_path/${pdb_lower}_final.mtz" - local FINAL_PDB_FILE="$pdb_path/${pdb_lower}_final.pdb" - - # Define output file paths - local output_path="$OUTPUT_DIR/$pdb_lower" - local FINAL_OUTPUT_FILE="$output_path/${pdb_lower}_edia.json" - local FINAL_CSV_FILE="$output_path/${pdb_lower}_residue_stats.csv" - - # Check for required files - if [[ ! -f "$FINAL_PDB_FILE" ]]; then - echo "[SKIP] PDB not found: $FINAL_PDB_FILE" - return 0 - fi - - if [[ ! -f "$FINAL_MTZ_FILE" ]]; then - echo "[SKIP] MTZ not found: $FINAL_MTZ_FILE" - return 0 - fi - - # Skip if already processed (check for CSV as final output) - if [[ -f "$FINAL_CSV_FILE" ]]; then - echo "[SKIP] Already processed: $pdb_lower" - return 0 - fi - - # Create output directory for this PDB - mkdir -p "$output_path" - - # Run density-fitness (EDIA) - if density-fitness "$FINAL_MTZ_FILE" "$FINAL_PDB_FILE" -o "$FINAL_OUTPUT_FILE" 2>/dev/null; then - # Parse output to CSV - if [[ -f "$FINAL_OUTPUT_FILE" ]]; then - if uv run "$PARSER" "$FINAL_OUTPUT_FILE" "$FINAL_CSV_FILE" 2>/dev/null; then - echo "[OK] $pdb_lower" - else - echo "[OK-NOPARSED] $pdb_lower" - fi - else - echo "[OK-NOFILE] $pdb_lower" - fi - else - echo "[FAIL] $pdb_lower" - return 0 - fi -} - -# Export function for parallel -export -f process_pdb - -# Run with GNU Parallel if available -if command -v parallel &> /dev/null; then - echo "Using GNU Parallel..." - echo "" - - grep -vE '^\s*$|^\s*#' "$FULL_ID_FILE" | \ - parallel --bar \ - --jobs "$NUM_JOBS" \ - --joblog "$OUTPUT_DIR/edia_joblog_$(date +%Y%m%d_%H%M%S).txt" \ - process_pdb {} -else - echo "GNU Parallel not found. Using xargs instead..." - echo "(Install GNU Parallel for better progress tracking)" - echo "" - - grep -vE '^\s*$|^\s*#' "$FULL_ID_FILE" | \ - xargs -P "$NUM_JOBS" -I {} bash -c 'process_pdb "$@"' _ {} -fi - -echo "" -echo "==============================================" -echo "EDIA Pipeline complete!" -echo "Results in: $OUTPUT_DIR" -echo "==============================================" diff --git a/scripts/train.py b/scripts/train.py index 84e3a46..092bd15 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -29,14 +29,13 @@ import wandb from loguru import logger from torch.optim import AdamW -from torch.optim.lr_scheduler import CosineAnnealingLR, StepLR, LinearLR -from tqdm import tqdm - +from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, StepLR from torch.utils.data import DataLoader from torch_geometric.data import HeteroData +from tqdm import tqdm -from src.encoder_base import build_encoder from src.dataset import get_dataloader +from src.encoder_base import build_encoder from src.flow import FlowMatcher, FlowWaterGVP from src.utils import ( compute_placement_metrics, diff --git a/tests/conftest.py b/tests/conftest.py index 474ef68..783d5cd 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -4,6 +4,7 @@ import pytest import torch + TEST_DIR = Path(__file__).parent ENV_PDB_DIR = os.environ.get("ENV_PDB_DIR") PDB_BASE_DIR = Path(ENV_PDB_DIR) if ENV_PDB_DIR else TEST_DIR / "test_files" diff --git a/tests/test_dataset.py b/tests/test_dataset.py index 80c5975..4bfc258 100644 --- a/tests/test_dataset.py +++ b/tests/test_dataset.py @@ -24,21 +24,21 @@ import torch from src.dataset import ( - ELEM_IDX, - ELEMENT_VOCAB, - ProteinWaterDataset, _make_undirected, check_chain_interactions, check_com_distance, check_water_clashes, compute_normalized_bfactors, + ELEM_IDX, element_onehot, + ELEMENT_VOCAB, filter_waters_by_quality, get_crystal_contacts_pymol, get_dataloader, load_edia_for_pdb, match_atoms_to_coords, parse_asu_with_biotite, + ProteinWaterDataset, ) @@ -599,12 +599,12 @@ def test_duplicate_single_sample(self, single_pdb_list_file, tmp_processed_dir, def test_cached_file_created(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): """Preprocessing should create cached .pt file.""" - dataset = ProteinWaterDataset( + _ = ProteinWaterDataset( pdb_list_file=single_pdb_list_file, processed_dir=str(tmp_processed_dir), base_pdb_dir=str(pdb_base_dir), preprocess=True, - ) + ) # need to call this to trigger the processing cache_file = tmp_processed_dir / "6eey_final_A.pt" assert cache_file.exists() diff --git a/tests/test_embedding_generation.py b/tests/test_embedding_generation.py index 83bc700..4456df2 100644 --- a/tests/test_embedding_generation.py +++ b/tests/test_embedding_generation.py @@ -164,7 +164,6 @@ def _make_data(atom_specs, embedding_dim=128): # Build arrays n_atoms = len(atom_specs) - n_residues = len(residue_keys) slae_residue_idx = torch.zeros(n_atoms, dtype=torch.long) slae_atom_type = torch.zeros(n_atoms, dtype=torch.long) @@ -427,9 +426,10 @@ def test_empty_geometry_list(self, make_slae_test_data): @pytest.mark.unit def test_empty_slae_embeddings(self): """All geometry atoms get zero vectors when SLAE is empty.""" - import numpy as np from unittest.mock import patch + import numpy as np + # Empty SLAE data slae_emb = torch.zeros(0, 128) slae_residue_idx = torch.zeros(0, dtype=torch.long) diff --git a/tests/test_encoder.py b/tests/test_encoder.py index 813976e..4711fde 100644 --- a/tests/test_encoder.py +++ b/tests/test_encoder.py @@ -14,9 +14,10 @@ from torch_cluster import radius_graph from torch_geometric.data import Data, HeteroData -from src.encoder_base import build_encoder, get_encoder_class, CachedEmbeddingEncoder +from src.encoder_base import build_encoder, CachedEmbeddingEncoder, get_encoder_class from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder + # ============== Fixtures ============== @pytest.fixture diff --git a/tests/test_flow.py b/tests/test_flow.py index 411fb53..ebcf962 100644 --- a/tests/test_flow.py +++ b/tests/test_flow.py @@ -3,7 +3,7 @@ All test cases created with assistance from Claude Code and refined. """ -from unittest.mock import MagicMock, Mock, patch +from unittest.mock import Mock import numpy as np import pytest @@ -11,12 +11,12 @@ from torch_geometric.data import Data, HeteroData from src.flow import ( + build_knn_edges, FlowMatcher, FlowWaterGVP, ProteinWaterUpdate, - build_knn_edges, ) -from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder, make_gvp_encoder_data +from src.gvp_encoder import GVPEncoder, make_gvp_encoder_data, ProteinGVPEncoder @pytest.fixture diff --git a/tests/test_forward.py b/tests/test_forward.py index 9553ed4..ca96ae3 100644 --- a/tests/test_forward.py +++ b/tests/test_forward.py @@ -11,7 +11,7 @@ from torch_geometric.data import HeteroData from src.flow import FlowMatcher, FlowWaterGVP -from src.gvp_encoder import GVPEncoder, ProteinGVPEncoder, make_gvp_encoder_data +from src.gvp_encoder import GVPEncoder, make_gvp_encoder_data, ProteinGVPEncoder def _iter_tensors(obj): diff --git a/tests/test_gvp.py b/tests/test_gvp.py index c6442be..bf7197e 100644 --- a/tests/test_gvp.py +++ b/tests/test_gvp.py @@ -1,8 +1,7 @@ -import pytest import torch import torch.nn.functional as F -from src.gvp import GVP, Dropout, LayerNorm, _merge, _split, tuple_cat, tuple_sum +from src.gvp import _merge, _split, Dropout, GVP, LayerNorm, tuple_cat, tuple_sum class TestGVPHelpers: diff --git a/tests/test_utils.py b/tests/test_utils.py index 1691878..cf1a20b 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -17,6 +17,7 @@ import pytest import torch + matplotlib.use('Agg') from pathlib import Path From 25dfdfbe750074bd1af84baa46e77ef5d5072714 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Mon, 16 Mar 2026 16:13:44 -0500 Subject: [PATCH 08/19] making fixes so that ruff, ty, and build checks pass --- pyproject.toml | 8 ++++++++ src/__init__.py | 4 ++-- src/dataset.py | 5 +++-- src/encoder_base.py | 3 ++- src/flow.py | 14 ++++++-------- src/gvp.py | 18 +++++++++--------- src/gvp_encoder.py | 2 +- src/utils.py | 25 +++++++++++++------------ 8 files changed, 44 insertions(+), 35 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index cc6eb86..ef17725 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -232,6 +232,14 @@ warn_unused_configs = true [tool.ty.rules] unresolved-import = "ignore" +# PyTorch Geometric has incomplete type stubs - dynamic attributes on Batch/Data +unresolved-attribute = "ignore" +# PyG's MessagePassing.message() override pattern is intentional +invalid-method-override = "ignore" +# PyG's propagate() uses **kwargs which ty doesn't understand +missing-argument = "ignore" +# Dict.get() returns union types that don't narrow well +invalid-argument-type = "ignore" [tool.ruff.lint] fixable = ["I001", "F401", "UP"] diff --git a/src/__init__.py b/src/__init__.py index 4492b19..b556fa2 100644 --- a/src/__init__.py +++ b/src/__init__.py @@ -6,10 +6,10 @@ Importing this module triggers encoder registration. """ -from src import gvp_encoder +from src import gvp_encoder as gvp_encoder from src.encoder_base import ( BaseProteinEncoder as BaseProteinEncoder, - CachedEmbeddingEncoder as CachedEmbeddingEncoder, build_encoder as build_encoder, + CachedEmbeddingEncoder as CachedEmbeddingEncoder, register_encoder as register_encoder, ) diff --git a/src/dataset.py b/src/dataset.py index 0505efc..151ecb7 100644 --- a/src/dataset.py +++ b/src/dataset.py @@ -11,7 +11,7 @@ import pymol2 import torch import torch.nn.functional as F -from biotite.structure.io.pdb import PDBFile, get_structure +from biotite.structure.io.pdb import get_structure, PDBFile from loguru import logger from scipy.spatial.distance import cdist from torch import Tensor @@ -23,6 +23,7 @@ from src.constants import EDGE_PP from src.utils import atom37_to_atoms + ELEMENT_VOCAB = [ "C", "N", "O", "S", "P", "SE", "MG", "ZN", "CA", "FE", "NA", "K", "CL", "F", "BR", ] @@ -681,7 +682,7 @@ def _preprocess_one(self, entry: dict, cache_path: Path): pdb_path = str(entry['pdb_path']) chain_filter = [entry['chain_id']] if entry['chain_id'] is not None else None - protein_atoms, water_atoms = parse_asu_with_biotite(pdb_path, chain_filter) + protein_atoms, water_atoms = parse_asu_with_biotite(pdb_path) # check inter-chain interactions for multi-chain proteins chain_valid, chain_reason, _ = check_chain_interactions( diff --git a/src/encoder_base.py b/src/encoder_base.py index 4b4d752..ee36f4b 100644 --- a/src/encoder_base.py +++ b/src/encoder_base.py @@ -15,11 +15,12 @@ import torch import torch.nn as nn + if TYPE_CHECKING: from torch_geometric.data import HeteroData # global encoder registry -_ENCODER_REGISTRY: dict[str, type[BaseProteinEncoder]] = {} +_ENCODER_REGISTRY: dict[str, type["BaseProteinEncoder"]] = {} def register_encoder(name: str): diff --git a/src/flow.py b/src/flow.py index 501623e..10b9507 100644 --- a/src/flow.py +++ b/src/flow.py @@ -1,14 +1,12 @@ from __future__ import annotations import copy -import math -from pathlib import Path import numpy as np import torch import torch.nn.functional as F -from torch import Tensor, nn -from torch_geometric.data import Batch, Data, HeteroData +from torch import nn, Tensor +from torch_geometric.data import Batch, HeteroData from torch_geometric.nn import knn from torch_scatter import scatter_mean from tqdm.auto import tqdm @@ -378,7 +376,7 @@ def compute_sigma(data: HeteroData) -> float: return float(pos.std().item()) @staticmethod - def compute_sigma_per_graph(data: HeteroData, device: torch.device) -> torch.Tensor: + def compute_sigma_per_graph(data: HeteroData | Batch, device: torch.device) -> torch.Tensor: """ Compute sigma (std of protein coordinates) per graph in a batch. @@ -402,7 +400,7 @@ def training_step( optimizer: torch.optim.Optimizer, grad_clip: float = 1.0, use_self_conditioning: bool = True, - ) -> dict[str, float]: + ) -> dict[str, float | dict | None]: """ Single flow matching training step. @@ -578,7 +576,7 @@ def euler_integrate( num_steps: int = 100, use_sc: bool = True, sc_ema_alpha: float = 0.2, - device: str = "cuda", + device: str | torch.device = "cuda", water_ratio: float | None = None, ) -> list[np.ndarray]: """ @@ -656,7 +654,7 @@ def rk4_integrate( num_steps: int = 500, use_sc: bool = True, sc_ema_alpha: float = 0.2, - device: str = "cuda", + device: str | torch.device = "cuda", return_trajectory: bool = True, water_ratio: float | None = None, ) -> list[dict[str, np.ndarray]]: diff --git a/src/gvp.py b/src/gvp.py index 4cbfef0..2496ae6 100644 --- a/src/gvp.py +++ b/src/gvp.py @@ -113,7 +113,7 @@ def __init__( activations=(F.relu, torch.sigmoid), vector_gate=False, ): - super(GVP, self).__init__() + super().__init__() self.si, self.vi = in_dims self.so, self.vo = out_dims self.vector_gate = vector_gate @@ -172,7 +172,7 @@ class _VDropout(nn.Module): """ def __init__(self, drop_rate): - super(_VDropout, self).__init__() + super().__init__() self.drop_rate = drop_rate self.dummy_param = nn.Parameter(torch.empty(0)) @@ -197,7 +197,7 @@ class Dropout(nn.Module): """ def __init__(self, drop_rate): - super(Dropout, self).__init__() + super().__init__() self.sdropout = nn.Dropout(drop_rate) self.vdropout = _VDropout(drop_rate) @@ -220,7 +220,7 @@ class LayerNorm(nn.Module): """ def __init__(self, dims): - super(LayerNorm, self).__init__() + super().__init__() self.s, self.v = dims self.scalar_norm = nn.LayerNorm(self.s) @@ -270,7 +270,7 @@ def __init__( activations=(F.relu, torch.sigmoid), vector_gate=False, ): - super(GVPConv, self).__init__(aggr=aggr) + super().__init__(aggr=aggr) self.si, self.vi = in_dims self.so, self.vo = out_dims self.se, self.ve = edge_dims @@ -352,7 +352,7 @@ def __init__( activations=(F.relu, torch.sigmoid), vector_gate=False, ): - super(GVPConvLayer, self).__init__() + super().__init__() self.conv = GVPConv( node_dims, node_dims, @@ -462,7 +462,7 @@ def forward( node_tuple: tuple, # (s_node, V_node) with s_node: (N, S_node) edge_index: torch.Tensor, # (2, E) edge_attr: tuple, # (s_edge, V_edge) with s_edge: (E, s_edge_width) - distance_feat: torch.Tensor = None, # (E, D) if enabled + distance_feat: torch.Tensor | None = None, # (E, D) if enabled ) -> tuple: s_node, _ = node_tuple @@ -514,7 +514,7 @@ def __init__( for i in range(n_message_gvps): vin = v_dim + (1 if i == 0 else 0) + (v_dim if (i == 0 and use_dst_feats) else 0) sin = s_dim + (rbf_dim if i == 0 else 0) + (s_dim if (i == 0 and use_dst_feats) else 0) - msg_layers.append(GVP( + msg_layers.append(GVP_( in_dims=(sin, vin), out_dims=(s_dim, v_dim), vector_gate=True, @@ -622,7 +622,7 @@ def __init__( for nt in dst_ntypes: upd_layers = [] for _ in range(n_update_gvps): - upd_layers.append(GVP((s_dim, v_dim), (s_dim, v_dim))) + upd_layers.append(GVP_((s_dim, v_dim), (s_dim, v_dim))) self.node_updates[nt] = nn.Sequential(*upd_layers) # per-edge-type message convs feeding a HeteroConv(aggr='sum' across relations) diff --git a/src/gvp_encoder.py b/src/gvp_encoder.py index 9d71206..0ff1c2f 100644 --- a/src/gvp_encoder.py +++ b/src/gvp_encoder.py @@ -19,7 +19,7 @@ from src.constants import EDGE_PP, NODE_FEATURE_DIM, NUM_RBF, RBF_CUTOFF from src.encoder_base import BaseProteinEncoder, register_encoder -from src.gvp import GVP, EdgeUpdate, GVPConvLayer +from src.gvp import EdgeUpdate, GVP, GVPConvLayer from src.utils import rbf diff --git a/src/utils.py b/src/utils.py index a5f6fc7..52886e4 100644 --- a/src/utils.py +++ b/src/utils.py @@ -1,6 +1,7 @@ # utils.py from __future__ import annotations + """ Utility functions organized by category: 1. Feature encoding (rbf, atom37_to_atoms, normalize_ins_code) @@ -23,11 +24,11 @@ from PIL import Image from scipy.optimize import linear_sum_assignment from torch import Tensor - from tqdm import tqdm from src.constants import NUM_RBF, RBF_CUTOFF + def setup_logging_for_tqdm( level: str = "INFO", log_file: str | None = None, @@ -235,15 +236,15 @@ def recall_precision( recall: fraction of true points with a prediction within thresh precision: fraction of predictions within thresh of a true point """ - # handle empty inputs + # convert numpy arrays to tensors first if isinstance(pred, np.ndarray): - if pred.size == 0 or true.size == 0: - return 0.0, 0.0 pred = torch.from_numpy(pred) + if isinstance(true, np.ndarray): true = torch.from_numpy(true) - else: - if pred.numel() == 0 or true.numel() == 0: - return 0.0, 0.0 + + # handle empty inputs + if pred.numel() == 0 or true.numel() == 0: + return 0.0, 0.0 # ensure same device if pred.device != true.device: @@ -347,13 +348,13 @@ def compute_placement_metrics( def plot_3d_frame( ax, protein_pos: np.ndarray, - mate_pos: np.ndarray, + mate_pos: np.ndarray | None, water_pred: np.ndarray, water_true: np.ndarray, title: str = "", - xlim: tuple[float, float] = None, - ylim: tuple[float, float] = None, - zlim: tuple[float, float] = None, + xlim: tuple[float, float] | None = None, + ylim: tuple[float, float] | None = None, + zlim: tuple[float, float] | None = None, ): """ Plot a single 3D frame showing protein structure and water positions. @@ -434,7 +435,7 @@ def create_trajectory_gif( save_path: str, title: str = "", fps: int = 10, - pdb_id: str = None, + pdb_id: str | None = None, ): """ Create a GIF from a trajectory of water positions. From bd8334924ec595b0f41e712a9a23c3ebe31903a0 Mon Sep 17 00:00:00 2001 From: vratins <114123331+vratins@users.noreply.github.com> Date: Mon, 16 Mar 2026 21:14:09 +0000 Subject: [PATCH 09/19] Auto-commit ruff fixes [skip ci] --- scripts/generate_esm_embeddings.py | 8 +- scripts/generate_slae_embeddings.py | 26 +- scripts/generate_water_plots.py | 231 ++++++++++++----- scripts/inference.py | 4 +- scripts/train.py | 2 +- src/constants.py | 34 ++- src/dataset.py | 270 +++++++++++--------- src/encoder_base.py | 14 +- src/flow.py | 262 ++++++++++--------- src/gvp.py | 157 +++++++----- src/gvp_encoder.py | 93 ++++--- src/utils.py | 7 + tests/conftest.py | 4 +- tests/test_dataset.py | 223 ++++++++++------ tests/test_embedding_generation.py | 69 ++--- tests/test_encoder.py | 180 +++++++------ tests/test_flow.py | 379 ++++++++++++++-------------- tests/test_forward.py | 164 ++++++++---- tests/test_gvp.py | 72 +++--- tests/test_utils.py | 121 +++++---- 20 files changed, 1383 insertions(+), 937 deletions(-) diff --git a/scripts/generate_esm_embeddings.py b/scripts/generate_esm_embeddings.py index 6e7aaae..cb58253 100644 --- a/scripts/generate_esm_embeddings.py +++ b/scripts/generate_esm_embeddings.py @@ -54,8 +54,8 @@ def compute_esm_embeddings( Compute ESM3 residue-level embeddings using an in-memory sanitized structure. This bypasses ESM's default behavior of dropping HETATM non-canonicals. - It extracts the sequence, strips HETATM flags, renames modified residues - (e.g., MSE -> MET) and unknowns (-> UNK), and feeds the buffer directly + It extracts the sequence, strips HETATM flags, renames modified residues + (e.g., MSE -> MET) and unknowns (-> UNK), and feeds the buffer directly to ESM3 so that all residues receive a structural embedding. How ESM parses: (https://github.com/evolutionaryscale/esm/blob/main/esm/utils/structure/protein_chain.py) @@ -89,7 +89,7 @@ def compute_esm_embeddings( key_to_resname[key] = res_name_arr[i] unique_res_keys = list(key_to_resname.keys()) - + biotite_seq = [ THREE_TO_ONE.get(key_to_resname[key], "X") for key in unique_res_keys ] @@ -270,4 +270,4 @@ def main() -> None: if __name__ == "__main__": - main() \ No newline at end of file + main() diff --git a/scripts/generate_slae_embeddings.py b/scripts/generate_slae_embeddings.py index 8b1ced3..c47f4f5 100644 --- a/scripts/generate_slae_embeddings.py +++ b/scripts/generate_slae_embeddings.py @@ -27,6 +27,7 @@ 'pdb_id': str, } """ + from __future__ import annotations import argparse @@ -193,8 +194,8 @@ def compute_slae_embeddings_batch( """ Compute and align SLAE node embeddings from atom37 coords for batched structures. - Embeddings are aligned to the target geometry's atom order. Non-canonical atoms - lacking SLAE representations are zero-padded. All input lists share the + Embeddings are aligned to the target geometry's atom order. Non-canonical atoms + lacking SLAE representations are zero-padded. All input lists share the same length (batch size). Args: @@ -217,7 +218,11 @@ def compute_slae_embeddings_batch( slae_atom_info_list = [] for coords, residue_type, residue_id, chains, ins_code in zip( - coords_list, residue_type_list, residue_id_list, chains_list, ins_code_list, + coords_list, + residue_type_list, + residue_id_list, + chains_list, + ins_code_list, strict=True, ): # create PyG Data object with atom37 coords (featurizer will convert to flat) @@ -430,13 +435,14 @@ def main() -> None: # compute embeddings for batch if batch_data: try: - (coords_list, - residue_type_list, - residue_id_list, - chains_list, - ins_code_list, - geometry_atom_info_list - ) = zip(*batch_data) + ( + coords_list, + residue_type_list, + residue_id_list, + chains_list, + ins_code_list, + geometry_atom_info_list, + ) = zip(*batch_data) embeddings_list = compute_slae_embeddings_batch( list(coords_list), diff --git a/scripts/generate_water_plots.py b/scripts/generate_water_plots.py index 0a0cfbb..f9fa269 100644 --- a/scripts/generate_water_plots.py +++ b/scripts/generate_water_plots.py @@ -75,7 +75,9 @@ def extract_water_bfactors_from_pdb( """ try: pdb_file = PDBFile.read(pdb_path) - atoms = pdb_file.get_structure(model=1, altloc="occupancy", extra_fields=["b_factor"]) + atoms = pdb_file.get_structure( + model=1, altloc="occupancy", extra_fields=["b_factor"] + ) # Filter for water molecules (HOH or WAT) water_mask = (atoms.res_name == "HOH") | (atoms.res_name == "WAT") @@ -85,8 +87,26 @@ def extract_water_bfactors_from_pdb( if normalization == "protein": # Standard amino acid residue names protein_residues = { - "ALA", "ARG", "ASN", "ASP", "CYS", "GLN", "GLU", "GLY", "HIS", "ILE", - "LEU", "LYS", "MET", "PHE", "PRO", "SER", "THR", "TRP", "TYR", "VAL", + "ALA", + "ARG", + "ASN", + "ASP", + "CYS", + "GLN", + "GLU", + "GLY", + "HIS", + "ILE", + "LEU", + "LYS", + "MET", + "PHE", + "PRO", + "SER", + "THR", + "TRP", + "TYR", + "VAL", } protein_mask = np.isin(atoms.res_name, list(protein_residues)) norm_bfactors = atoms.b_factor[protein_mask] @@ -104,11 +124,11 @@ def extract_water_bfactors_from_pdb( if len(water_atoms) == 0: return None - #extract PDB ID from filename (e.g., "3ilf_final.pdb" -> "3ilf") + # extract PDB ID from filename (e.g., "3ilf_final.pdb" -> "3ilf") pdb_id = pdb_path.stem.replace("_final", "") - #build DataFrame with one row per unique water residue - #water molecules have one oxygen atom, so we take unique (chain, res_id) pairs + # build DataFrame with one row per unique water residue + # water molecules have one oxygen atom, so we take unique (chain, res_id) pairs records = [] seen = set() for i in range(len(water_atoms)): @@ -118,15 +138,17 @@ def extract_water_bfactors_from_pdb( if key not in seen: seen.add(key) raw_bfactor = water_atoms.b_factor[i] - #z-score using whole-PDB statistics + # z-score using whole-PDB statistics normalized = (raw_bfactor - pdb_mean) / pdb_std if pdb_std > 0 else 0.0 - records.append({ - "pdb_id": pdb_id, - "chain_id": chain_id, - "res_id": res_id, - "b_factor": raw_bfactor, - "b_factor_normalized": normalized, - }) + records.append( + { + "pdb_id": pdb_id, + "chain_id": chain_id, + "res_id": res_id, + "b_factor": raw_bfactor, + "b_factor_normalized": normalized, + } + ) return pd.DataFrame(records) @@ -170,9 +192,13 @@ def load_all_bfactors( all_bfactors = [] with ProcessPoolExecutor(max_workers=num_workers) as executor: - futures = {executor.submit(_extract_bfactors_worker, task): task[1] for task in tasks} + futures = { + executor.submit(_extract_bfactors_worker, task): task[1] for task in tasks + } - for future in tqdm(as_completed(futures), total=len(futures), desc="Extracting B-factors"): + for future in tqdm( + as_completed(futures), total=len(futures), desc="Extracting B-factors" + ): result = future.result() if result is not None: all_bfactors.append(result) @@ -181,11 +207,15 @@ def load_all_bfactors( raise ValueError("No B-factor data extracted from any PDB files") combined = pd.concat(all_bfactors, ignore_index=True) - logger.info(f"Extracted B-factors for {len(combined)} water molecules from {len(all_bfactors)} PDBs") + logger.info( + f"Extracted B-factors for {len(combined)} water molecules from {len(all_bfactors)} PDBs" + ) return combined -def merge_edia_with_bfactors(edia_df: pd.DataFrame, bfactor_df: pd.DataFrame) -> pd.DataFrame: +def merge_edia_with_bfactors( + edia_df: pd.DataFrame, bfactor_df: pd.DataFrame +) -> pd.DataFrame: """Merge EDIA data with B-factor data. Matching is done on (pdb_id, chain, residue_number): @@ -200,20 +230,24 @@ def merge_edia_with_bfactors(edia_df: pd.DataFrame, bfactor_df: pd.DataFrame) -> Returns: Merged DataFrame with b_factor and b_factor_normalized columns added """ - #rename B-factor columns to match EDIA column names - bfactor_renamed = bfactor_df.rename(columns={ - "chain_id": "pdb_strandID", - "res_id": "pdb_seqNum", - }) + # rename B-factor columns to match EDIA column names + bfactor_renamed = bfactor_df.rename( + columns={ + "chain_id": "pdb_strandID", + "res_id": "pdb_seqNum", + } + ) - #merge on the matching key + # merge on the matching key merged = edia_df.merge( - bfactor_renamed[["pdb_id", "pdb_strandID", "pdb_seqNum", "b_factor", "b_factor_normalized"]], + bfactor_renamed[ + ["pdb_id", "pdb_strandID", "pdb_seqNum", "b_factor", "b_factor_normalized"] + ], on=["pdb_id", "pdb_strandID", "pdb_seqNum"], how="left", ) - #report match statistics + # report match statistics n_total = len(merged) n_matched = merged["b_factor"].notna().sum() match_rate = 100 * n_matched / n_total if n_total > 0 else 0 @@ -226,7 +260,9 @@ def plot_ediam_waters(df: pd.DataFrame, output_dir: Path): """Plot histogram of EDIAm for all water molecules.""" fig, ax = plt.subplots(figsize=(10, 6)) - ax.hist(df["EDIAm"].dropna(), bins=50, edgecolor="black", alpha=0.7, color="steelblue") + ax.hist( + df["EDIAm"].dropna(), bins=50, edgecolor="black", alpha=0.7, color="steelblue" + ) # Add threshold lines ax.axvline(x=0.4, color="red", linestyle="--", linewidth=2, label="EDIAm = 0.4") @@ -247,12 +283,19 @@ def plot_ediam_waters(df: pd.DataFrame, output_dir: Path): f"mean = {mean_val:.3f}\n" f"median = {median_val:.3f}\n" f"─────────────\n" - f"< 0.4: {n_low:,} ({100*n_low/n_total:.1f}%)\n" - f"0.4–0.8: {n_mid:,} ({100*n_mid/n_total:.1f}%)\n" - f"≥ 0.8: {n_high:,} ({100*n_high/n_total:.1f}%)" + f"< 0.4: {n_low:,} ({100 * n_low / n_total:.1f}%)\n" + f"0.4–0.8: {n_mid:,} ({100 * n_mid / n_total:.1f}%)\n" + f"≥ 0.8: {n_high:,} ({100 * n_high / n_total:.1f}%)" + ) + ax.text( + 0.02, + 0.98, + textstr, + transform=ax.transAxes, + fontsize=10, + verticalalignment="top", + bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), ) - ax.text(0.02, 0.98, textstr, transform=ax.transAxes, fontsize=10, - verticalalignment="top", bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5)) ax.set_xlabel("EDIAm Score", fontsize=12) ax.set_ylabel("Count", fontsize=12) @@ -271,7 +314,9 @@ def plot_ediam_pdbs(df: pd.DataFrame, output_dir: Path): fig, ax = plt.subplots(figsize=(10, 6)) - ax.hist(pdb_means.dropna(), bins=50, edgecolor="black", alpha=0.7, color="steelblue") + ax.hist( + pdb_means.dropna(), bins=50, edgecolor="black", alpha=0.7, color="steelblue" + ) # Add threshold lines ax.axvline(x=0.4, color="red", linestyle="--", linewidth=2, label="EDIAm = 0.4") @@ -292,12 +337,19 @@ def plot_ediam_pdbs(df: pd.DataFrame, output_dir: Path): f"mean = {mean_val:.3f}\n" f"median = {median_val:.3f}\n" f"─────────────\n" - f"< 0.4: {n_low:,} ({100*n_low/n_pdbs:.1f}%)\n" - f"0.4–0.8: {n_mid:,} ({100*n_mid/n_pdbs:.1f}%)\n" - f"≥ 0.8: {n_high:,} ({100*n_high/n_pdbs:.1f}%)" + f"< 0.4: {n_low:,} ({100 * n_low / n_pdbs:.1f}%)\n" + f"0.4–0.8: {n_mid:,} ({100 * n_mid / n_pdbs:.1f}%)\n" + f"≥ 0.8: {n_high:,} ({100 * n_high / n_pdbs:.1f}%)" + ) + ax.text( + 0.02, + 0.98, + textstr, + transform=ax.transAxes, + fontsize=10, + verticalalignment="top", + bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), ) - ax.text(0.02, 0.98, textstr, transform=ax.transAxes, fontsize=10, - verticalalignment="top", bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5)) ax.set_xlabel("Mean EDIAm Score", fontsize=12) ax.set_ylabel("Number of PDBs", fontsize=12) @@ -321,8 +373,15 @@ def plot_rsccs_waters(df: pd.DataFrame, output_dir: Path): mean_val = df["RSCCS"].mean() median_val = df["RSCCS"].median() textstr = f"n = {n_total:,}\nmean = {mean_val:.3f}\nmedian = {median_val:.3f}" - ax.text(0.02, 0.98, textstr, transform=ax.transAxes, fontsize=10, - verticalalignment="top", bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5)) + ax.text( + 0.02, + 0.98, + textstr, + transform=ax.transAxes, + fontsize=10, + verticalalignment="top", + bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), + ) ax.set_xlabel("RSCCS Score", fontsize=12) ax.set_ylabel("Count", fontsize=12) @@ -347,8 +406,15 @@ def plot_rsccs_pdbs(df: pd.DataFrame, output_dir: Path): mean_val = pdb_means.mean() median_val = pdb_means.median() textstr = f"n = {n_pdbs:,} PDBs\nmean = {mean_val:.3f}\nmedian = {median_val:.3f}" - ax.text(0.02, 0.98, textstr, transform=ax.transAxes, fontsize=10, - verticalalignment="top", bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5)) + ax.text( + 0.02, + 0.98, + textstr, + transform=ax.transAxes, + fontsize=10, + verticalalignment="top", + bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), + ) ax.set_xlabel("Mean RSCCS Score", fontsize=12) ax.set_ylabel("Number of PDBs", fontsize=12) @@ -378,7 +444,9 @@ def plot_bfactor_waters(df: pd.DataFrame, output_dir: Path): # Add cutoff lines at +1.5 and -1.5 cutoff = 1.5 - ax.axvline(x=cutoff, color="red", linestyle="--", linewidth=2, label=f"cutoff = ±{cutoff}") + ax.axvline( + x=cutoff, color="red", linestyle="--", linewidth=2, label=f"cutoff = ±{cutoff}" + ) ax.axvline(x=-cutoff, color="red", linestyle="--", linewidth=2) # Add statistics @@ -396,13 +464,20 @@ def plot_bfactor_waters(df: pd.DataFrame, output_dir: Path): f"mean = {mean_val:.2f}\n" f"median = {median_val:.2f}\n" f"─────────────\n" - f"< -{cutoff}: {n_below:,} ({100*n_below/n_total:.1f}%)\n" - f"-{cutoff} to {cutoff}: {n_within:,} ({100*n_within/n_total:.1f}%)\n" - f"> {cutoff}: {n_above:,} ({100*n_above/n_total:.1f}%)" + f"< -{cutoff}: {n_below:,} ({100 * n_below / n_total:.1f}%)\n" + f"-{cutoff} to {cutoff}: {n_within:,} ({100 * n_within / n_total:.1f}%)\n" + f"> {cutoff}: {n_above:,} ({100 * n_above / n_total:.1f}%)" + ) + ax.text( + 0.98, + 0.98, + textstr, + transform=ax.transAxes, + fontsize=10, + verticalalignment="top", + horizontalalignment="right", + bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), ) - ax.text(0.98, 0.98, textstr, transform=ax.transAxes, fontsize=10, - verticalalignment="top", horizontalalignment="right", - bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5)) ax.set_xlabel("Normalized B-factor (z-score)", fontsize=12) ax.set_ylabel("Count", fontsize=12) @@ -436,7 +511,9 @@ def plot_bfactor_pdbs(df: pd.DataFrame, output_dir: Path): # Add cutoff line for high variability cutoff = 1.5 - ax.axvline(x=cutoff, color="red", linestyle="--", linewidth=2, label=f"cutoff = {cutoff}") + ax.axvline( + x=cutoff, color="red", linestyle="--", linewidth=2, label=f"cutoff = {cutoff}" + ) # Add statistics n_pdbs = len(pdb_stds) @@ -452,12 +529,19 @@ def plot_bfactor_pdbs(df: pd.DataFrame, output_dir: Path): f"mean = {mean_val:.2f}\n" f"median = {median_val:.2f}\n" f"─────────────\n" - f"≤ {cutoff}: {n_below:,} ({100*n_below/n_pdbs:.1f}%)\n" - f"> {cutoff}: {n_above:,} ({100*n_above/n_pdbs:.1f}%)" + f"≤ {cutoff}: {n_below:,} ({100 * n_below / n_pdbs:.1f}%)\n" + f"> {cutoff}: {n_above:,} ({100 * n_above / n_pdbs:.1f}%)" + ) + ax.text( + 0.98, + 0.98, + textstr, + transform=ax.transAxes, + fontsize=10, + verticalalignment="top", + horizontalalignment="right", + bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), ) - ax.text(0.98, 0.98, textstr, transform=ax.transAxes, fontsize=10, - verticalalignment="top", horizontalalignment="right", - bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5)) ax.set_xlabel("Std Dev of Normalized B-factor (z-score)", fontsize=12) ax.set_ylabel("Number of PDBs", fontsize=12) @@ -483,13 +567,25 @@ def plot_ediam_bfactor_correlation(df: pd.DataFrame, output_dir: Path): # Left: Scatter plot ax1 = axes[0] - ax1.scatter(df_valid["b_factor_normalized"], df_valid["EDIAm"], alpha=0.1, s=5, c="steelblue") + ax1.scatter( + df_valid["b_factor_normalized"], + df_valid["EDIAm"], + alpha=0.1, + s=5, + c="steelblue", + ) # Add correlation coefficient corr = df_valid["EDIAm"].corr(df_valid["b_factor_normalized"]) - ax1.text(0.02, 0.98, f"r = {corr:.3f}\nn = {len(df_valid):,}", - transform=ax1.transAxes, fontsize=12, - verticalalignment="top", bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5)) + ax1.text( + 0.02, + 0.98, + f"r = {corr:.3f}\nn = {len(df_valid):,}", + transform=ax1.transAxes, + fontsize=12, + verticalalignment="top", + bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), + ) ax1.set_xlabel("Normalized B-factor (z-score)", fontsize=12) ax1.set_ylabel("EDIAm Score", fontsize=12) @@ -497,7 +593,13 @@ def plot_ediam_bfactor_correlation(df: pd.DataFrame, output_dir: Path): # Right: Hexbin density plot ax2 = axes[1] - hb = ax2.hexbin(df_valid["b_factor_normalized"], df_valid["EDIAm"], gridsize=50, cmap="YlOrRd", mincnt=1) + hb = ax2.hexbin( + df_valid["b_factor_normalized"], + df_valid["EDIAm"], + gridsize=50, + cmap="YlOrRd", + mincnt=1, + ) fig.colorbar(hb, ax=ax2, label="Count") ax2.set_xlabel("Normalized B-factor (z-score)", fontsize=12) @@ -613,7 +715,9 @@ def main(): # Extract B-factors if needed if need_bfactor: - logger.info(f"\nExtracting B-factors from PDB files (normalization: {args.bfactor_normalization})...") + logger.info( + f"\nExtracting B-factors from PDB files (normalization: {args.bfactor_normalization})..." + ) if args.bfactor_only: # Get PDB IDs from text file @@ -623,7 +727,10 @@ def main(): pdb_ids = df["pdb_id"].unique().tolist() bfactor_df = load_all_bfactors( - args.pdb_dir, pdb_ids, args.num_workers, normalization=args.bfactor_normalization + args.pdb_dir, + pdb_ids, + args.num_workers, + normalization=args.bfactor_normalization, ) # Merge with EDIA data if both are available diff --git a/scripts/inference.py b/scripts/inference.py index c5c1fb8..324ed4f 100644 --- a/scripts/inference.py +++ b/scripts/inference.py @@ -576,7 +576,9 @@ def main(): logger.info(f" Samples processed: {summary['n_samples']}") logger.info(f" Avg waters (true): {summary['avg_n_waters_true']:.1f}") logger.info(f" Avg waters (pred): {summary['avg_n_waters_pred']:.1f}") - logger.info(f" Avg RMSD: {summary['avg_rmsd']:.3f} ± {summary['std_rmsd']:.3f} Å") + logger.info( + f" Avg RMSD: {summary['avg_rmsd']:.3f} ± {summary['std_rmsd']:.3f} Å" + ) logger.info(f" Avg Precision: {summary['avg_precision']:.3%}") logger.info(f" Avg Recall: {summary['avg_recall']:.3%}") logger.info(f" Avg F1: {summary['avg_f1']:.4f}") diff --git a/scripts/train.py b/scripts/train.py index 092bd15..fe50412 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -1088,7 +1088,7 @@ def main(): epoch, device, global_step, - eval_indices, + eval_indices, run_dir, ) if eval_metrics: diff --git a/src/constants.py b/src/constants.py index e2057cc..b260bab 100644 --- a/src/constants.py +++ b/src/constants.py @@ -7,8 +7,8 @@ NODE_FEATURE_DIM = 16 # Default node scalar feature dimension # RBF (Radial Basis Function) parameters -NUM_RBF = 16 # Number of RBF basis functions -RBF_CUTOFF = 8.0 # Distance cutoff in Angstroms for RBF encoding +NUM_RBF = 16 # Number of RBF basis functions +RBF_CUTOFF = 8.0 # Distance cutoff in Angstroms for RBF encoding # Edge type tuples: (src_node_type, edge_name, dst_node_type) EDGE_PP = ("protein", "pp", "protein") # protein -> protein @@ -51,12 +51,30 @@ } # Standard 1-to-3 letter mapping to feed sanitized residues back to ESM3. -# This acts as the inverse of THREE_TO_ONE, ensuring ESM3 recognizes +# This acts as the inverse of THREE_TO_ONE, ensuring ESM3 recognizes # the atoms and safely maps true unknowns to 'UNK'. (probably a more efficient way to do this I know) ONE_TO_THREE = { - "A": "ALA", "C": "CYS", "D": "ASP", "E": "GLU", "F": "PHE", - "G": "GLY", "H": "HIS", "I": "ILE", "K": "LYS", "L": "LEU", - "M": "MET", "N": "ASN", "P": "PRO", "Q": "GLN", "R": "ARG", - "S": "SER", "T": "THR", "V": "VAL", "W": "TRP", "Y": "TYR", - "X": "UNK", "U": "SEC", "O": "PYL" + "A": "ALA", + "C": "CYS", + "D": "ASP", + "E": "GLU", + "F": "PHE", + "G": "GLY", + "H": "HIS", + "I": "ILE", + "K": "LYS", + "L": "LEU", + "M": "MET", + "N": "ASN", + "P": "PRO", + "Q": "GLN", + "R": "ARG", + "S": "SER", + "T": "THR", + "V": "VAL", + "W": "TRP", + "Y": "TYR", + "X": "UNK", + "U": "SEC", + "O": "PYL", } diff --git a/src/dataset.py b/src/dataset.py index 151ecb7..1bae99b 100644 --- a/src/dataset.py +++ b/src/dataset.py @@ -1,4 +1,4 @@ -#dataset.py +# dataset.py from __future__ import annotations @@ -25,7 +25,21 @@ ELEMENT_VOCAB = [ - "C", "N", "O", "S", "P", "SE", "MG", "ZN", "CA", "FE", "NA", "K", "CL", "F", "BR", + "C", + "N", + "O", + "S", + "P", + "SE", + "MG", + "ZN", + "CA", + "FE", + "NA", + "K", + "CL", + "F", + "BR", ] ELEM_IDX = {e: i for i, e in enumerate(ELEMENT_VOCAB)} @@ -33,7 +47,9 @@ def element_onehot(symbols: list[str]) -> Tensor: """One-hot encoding with 'other' bucket at end.""" other_idx = len(ELEMENT_VOCAB) - indices = torch.tensor([ELEM_IDX.get(s.upper(), other_idx) for s in symbols], dtype=torch.long) + indices = torch.tensor( + [ELEM_IDX.get(s.upper(), other_idx) for s in symbols], dtype=torch.long + ) return F.one_hot(indices, num_classes=other_idx + 1).float() @@ -80,29 +96,31 @@ def get_crystal_contacts_pymol(pdb_path: str, cutoff: float = 5.0) -> dict: cmd.load(pdb_path, obj) cmd.symexp("sym", obj, obj, cutoff) cmd.select("interface", f"byres (sym* within {cutoff} of {obj})") - + asu_coords = cmd.get_coords(obj, state=1) mate_coords = cmd.get_coords("sym* and interface", state=1) asu_atoms = cmd.get_model(obj, state=1).atom mate_atoms = cmd.get_model("sym* and interface", state=1).atom - + return { - "asu_coords": asu_coords if asu_coords is not None else np.zeros((0, 3), dtype=float), - "mate_coords": mate_coords if mate_coords is not None else np.zeros((0, 3), dtype=float), + "asu_coords": asu_coords + if asu_coords is not None + else np.zeros((0, 3), dtype=float), + "mate_coords": mate_coords + if mate_coords is not None + else np.zeros((0, 3), dtype=float), "asu_atoms": asu_atoms, "mate_atoms": mate_atoms, } def match_atoms_to_coords( - atoms: bts.AtomArray, - target_coords: np.ndarray, - tolerance: float = 0.01 + atoms: bts.AtomArray, target_coords: np.ndarray, tolerance: float = 0.01 ) -> list[int]: """Match biotite atoms to PyMOL coordinates, return indices.""" if target_coords.shape[0] == 0: return [] - + matched = [] for i, coord in enumerate(target_coords): dists = np.linalg.norm(atoms.coord - coord, axis=1) @@ -120,6 +138,7 @@ def _make_undirected(edge_index: torch.Tensor) -> torch.Tensor: ei = torch.unique(ei.T, dim=0).T # drop duplicates return ei + def check_com_distance( protein_coords: torch.Tensor, water_coords: torch.Tensor, @@ -215,11 +234,13 @@ def check_chain_interactions( # get coordinates per chain chain_coords = { - cid: torch.tensor(protein_atoms[protein_atoms.chain_id == cid].coord, dtype=torch.float32) + cid: torch.tensor( + protein_atoms[protein_atoms.chain_id == cid].coord, dtype=torch.float32 + ) for cid in chain_ids } - min_interface_dist = float('inf') + min_interface_dist = float("inf") for chain_a, chain_b in itertools.combinations(chain_ids, 2): coords_a = chain_coords[chain_a] @@ -237,7 +258,7 @@ def check_chain_interactions( False, f"Multi-chain ({num_chains} chains) min interface distance {min_interface_dist:.1f}A " f"> {interface_dist_threshold}A (likely ASU copies, not PPI)", - "Non-Interacting (ASU Copies)" + "Non-Interacting (ASU Copies)", ) return True, "", "Interacting" @@ -339,9 +360,7 @@ def compute_normalized_bfactors( try: pdb_file = PDBFile.read(pdb_path) atoms = pdb_file.get_structure( - model=1, - altloc="occupancy", - extra_fields=["b_factor"] + model=1, altloc="occupancy", extra_fields=["b_factor"] ) # compute B-factor statistics for normalization from PDB entry (including non-water atoms) @@ -353,7 +372,9 @@ def compute_normalized_bfactors( # apply chain filter if specified if chain_filter is not None: - mask = np.isin(atoms.chain_id, np.array(chain_filter, dtype=atoms.chain_id.dtype)) + mask = np.isin( + atoms.chain_id, np.array(chain_filter, dtype=atoms.chain_id.dtype) + ) atoms = atoms[mask] # filter for water molecules @@ -471,7 +492,9 @@ def filter_waters_by_quality( lookup_fail = np.zeros(n_waters, dtype=bool) for lookup, threshold, fail_if_below, name in lookup_filters: if lookup is not None: - fail_mask = apply_threshold_filter(water_keys, lookup, threshold, fail_if_below) + fail_mask = apply_threshold_filter( + water_keys, lookup, threshold, fail_if_below + ) stats[f"removed_{name}"] = int(fail_mask.sum()) lookup_fail |= fail_mask @@ -483,10 +506,12 @@ def filter_waters_by_quality( if cache_key is not None and stats["total"] > 0: removed = stats["total"] - stats["kept"] if removed > 0: - logger.info(f" {cache_key}: Filtered {removed}/{stats['total']} waters " - f"(dist:{stats['removed_distance']}, " - f"edia:{stats['removed_edia']}, " - f"bfactor:{stats['removed_bfactor']})") + logger.info( + f" {cache_key}: Filtered {removed}/{stats['total']} waters " + f"(dist:{stats['removed_distance']}, " + f"edia:{stats['removed_edia']}, " + f"bfactor:{stats['removed_bfactor']})" + ) return keep_mask @@ -494,13 +519,13 @@ def filter_waters_by_quality( class ProteinWaterDataset(Dataset): """ Dataset for protein crystal contact prediction. - + Returns HeteroData with: - 'protein' node type: ASU protein atoms + optionally symmetry mates - 'water' node type: water molecules - - ('protein', 'pp', 'protein') edges + - ('protein', 'pp', 'protein') edges """ - + def __init__( self, pdb_list_file: str, @@ -579,10 +604,12 @@ def __init__( # if single sample and duplication requested, set effective length [this is for experiments to check if the model can memorize a sample] if len(self.entries) == 1 and duplicate_single_sample > 1: self._effective_length = duplicate_single_sample - logger.info(f"Single sample detected. Duplicating {duplicate_single_sample}x ") + logger.info( + f"Single sample detected. Duplicating {duplicate_single_sample}x " + ) else: self._effective_length = len(self.entries) - + def _parse_pdb_list(self, pdb_list_file: str) -> list[dict]: """ Parse PDB list file and construct entries with paths. @@ -594,13 +621,13 @@ def _parse_pdb_list(self, pdb_list_file: str) -> list[dict]: Constructs path: {base_pdb_dir}/{pdb_id}/{pdb_id}_final.pdb """ entries = [] - with open(pdb_list_file, 'r') as f: + with open(pdb_list_file, "r") as f: for line in f: line = line.strip() if not line: continue - parts = line.split('_') + parts = line.split("_") if len(parts) < 2: logger.warning(f"Warning: Skipping malformed line: {line}") continue @@ -617,22 +644,25 @@ def _parse_pdb_list(self, pdb_list_file: str) -> list[dict]: pdb_path = self.base_pdb_dir / pdb_id / f"{pdb_id}_final.pdb" - entries.append({ - 'pdb_id': pdb_id, - 'chain_id': chain_id, - 'pdb_path': pdb_path, - 'cache_key': line, - }) + entries.append( + { + "pdb_id": pdb_id, + "chain_id": chain_id, + "pdb_path": pdb_path, + "cache_key": line, + } + ) logger.info(f"Loaded {len(entries)} entries from {pdb_list_file}") return entries - + def _preprocess_all(self): """Preprocess all PDB files that don't have cached results.""" self.processed_dir.mkdir(parents=True, exist_ok=True) to_process = [ - e for e in self.entries + e + for e in self.entries if not (self.processed_dir / f"{e['cache_key']}.pt").exists() ] @@ -648,18 +678,19 @@ def _preprocess_all(self): self._preprocess_one(entry, cache_path) except Exception as e: logger.warning(f"\nFailed to preprocess {entry['cache_key']}: {e}") - failures.append((entry['cache_key'], str(e))) + failures.append((entry["cache_key"], str(e))) # write failures to log file if failures: failure_log_path = self.processed_dir / "preprocessing_failures.log" - with open(failure_log_path, 'a') as f: + with open(failure_log_path, "a") as f: for pdb_id, reason in failures: f.write(f"{pdb_id}\t{reason}\n") logger.info(f"Logged {len(failures)} failures to {failure_log_path}") valid_entries = [ - e for e in self.entries + e + for e in self.entries if (self.processed_dir / f"{e['cache_key']}.pt").exists() ] n_removed = len(self.entries) - len(valid_entries) @@ -667,7 +698,7 @@ def _preprocess_all(self): logger.info(f"Filtered out {n_removed} entries without valid cache files.") self.entries = valid_entries logger.info(f"Dataset contains {len(self.entries)} valid entries.") - + def _preprocess_one(self, entry: dict, cache_path: Path): """ Preprocess a single PDB file. @@ -679,8 +710,8 @@ def _preprocess_one(self, entry: dict, cache_path: Path): Raises ValueError if structure fails quality filters. """ - pdb_path = str(entry['pdb_path']) - chain_filter = [entry['chain_id']] if entry['chain_id'] is not None else None + pdb_path = str(entry["pdb_path"]) + chain_filter = [entry["chain_id"]] if entry["chain_id"] is not None else None protein_atoms, water_atoms = parse_asu_with_biotite(pdb_path) @@ -714,23 +745,23 @@ def _preprocess_one(self, entry: dict, cache_path: Path): # load EDIA data if directory provided and EDIA filtering enabled edia_lookup = None if self.filter_by_edia and self.edia_dir is not None: - edia_lookup = load_edia_for_pdb(self.edia_dir, entry['pdb_id']) + edia_lookup = load_edia_for_pdb(self.edia_dir, entry["pdb_id"]) if edia_lookup is None: - logger.warning(f"Warning: EDIA file not found for {entry['pdb_id']}, skipping EDIA filtering") + logger.warning( + f"Warning: EDIA file not found for {entry['pdb_id']}, skipping EDIA filtering" + ) # compute normalized B-factors if B-factor filtering enabled bfactor_lookup = None if self.filter_by_bfactor: bfactor_lookup, _ = compute_normalized_bfactors( - pdb_path, - chain_filter=chain_filter + pdb_path, chain_filter=chain_filter ) # build water keys for filtering - water_keys = list(zip( - water_atoms.chain_id.astype(str), - water_atoms.res_id.astype(int) - )) + water_keys = list( + zip(water_atoms.chain_id.astype(str), water_atoms.res_id.astype(int)) + ) # apply quality filters keep_mask = filter_waters_by_quality( @@ -742,7 +773,7 @@ def _preprocess_one(self, entry: dict, cache_path: Path): max_protein_dist=self.max_protein_dist, min_edia=self.min_edia, max_bfactor_zscore=self.max_bfactor_zscore, - cache_key=entry['cache_key'], + cache_key=entry["cache_key"], ) water_atoms = water_atoms[keep_mask] @@ -762,7 +793,7 @@ def _preprocess_one(self, entry: dict, cache_path: Path): if not com_valid: raise ValueError(f"Quality filter failed: {com_reason}") - # check water clashes with protein atoms + # check water clashes with protein atoms clash_valid, clash_reason = check_water_clashes( protein_pos, water_pos_raw, @@ -775,10 +806,10 @@ def _preprocess_one(self, entry: dict, cache_path: Path): # center protein positions center = protein_pos.mean(dim=0, keepdim=True) protein_pos = protein_pos - center - + protein_elements = [str(e).upper() for e in protein_atoms.element] protein_x = element_onehot(protein_elements) - + # compute residue indices (using chain_id, res_id only - matches SLAE's atomarray_to_tensors) res_id = protein_atoms.res_id chain_id_arr = protein_atoms.chain_id @@ -807,7 +838,7 @@ def _preprocess_one(self, entry: dict, cache_path: Path): else: water_pos = torch.zeros((0, 3), dtype=torch.float32) water_x = torch.zeros((0, len(ELEMENT_VOCAB) + 1), dtype=torch.float32) - + # process symmetry mate atoms mate_coords = crystal_data["mate_coords"] if mate_coords.shape[0] > 0: @@ -828,20 +859,23 @@ def _preprocess_one(self, entry: dict, cache_path: Path): mate_res_idx = torch.empty(0, dtype=torch.long) # cache all data - torch.save({ - 'protein_pos': protein_pos, - 'protein_x': protein_x, - 'protein_res_idx': protein_res_idx, - 'water_pos': water_pos, - 'water_x': water_x, - 'mate_pos': mate_pos, - 'mate_x': mate_x, - 'mate_res_idx': mate_res_idx, - }, cache_path) - + torch.save( + { + "protein_pos": protein_pos, + "protein_x": protein_x, + "protein_res_idx": protein_res_idx, + "water_pos": water_pos, + "water_x": water_x, + "mate_pos": mate_pos, + "mate_x": mate_x, + "mate_res_idx": mate_res_idx, + }, + cache_path, + ) + def __len__(self) -> int: return self._effective_length - + def __getitem__(self, idx: int) -> HeteroData: """ Load cached data and build graph on-the-fly. @@ -856,18 +890,20 @@ def __getitem__(self, idx: int) -> HeteroData: actual_idx = idx % len(self.entries) entry = self.entries[actual_idx] cache_path = self.processed_dir / f"{entry['cache_key']}.pt" - + if not cache_path.exists(): raise FileNotFoundError( f"Cached file not found: {cache_path}. " f"Run with preprocess=True to generate it." ) - + cached = torch.load(cache_path, weights_only=False) - if 'protein_slae_embedding' in cached and 'protein_atom37_coords' in cached: - atom37_coords = cached['protein_atom37_coords'] - protein_pos, residue_idx_per_atom, atom_types = atom37_to_atoms(atom37_coords) + if "protein_slae_embedding" in cached and "protein_atom37_coords" in cached: + atom37_coords = cached["protein_atom37_coords"] + protein_pos, residue_idx_per_atom, atom_types = atom37_to_atoms( + atom37_coords + ) # recenter (TODO: optimize by centering in precompute script) center = protein_pos.mean(dim=0, keepdim=True) @@ -877,76 +913,82 @@ def __getitem__(self, idx: int) -> HeteroData: protein_res_idx = residue_idx_per_atom num_asu_protein = protein_pos.size(0) else: - # use original protein atoms from cache - protein_pos = cached['protein_pos'] - protein_x = cached['protein_x'] - protein_res_idx = cached['protein_res_idx'] + # use original protein atoms from cache + protein_pos = cached["protein_pos"] + protein_x = cached["protein_x"] + protein_res_idx = cached["protein_res_idx"] num_asu_protein = protein_pos.size(0) - + # compute num_residues for protein (before adding mates) - num_protein_residues = int(protein_res_idx.max().item() + 1) if protein_res_idx.numel() > 0 else 0 + num_protein_residues = ( + int(protein_res_idx.max().item() + 1) if protein_res_idx.numel() > 0 else 0 + ) # concatenate symmetry mate atoms to protein if mates are included - if self.include_mates and cached['mate_pos'].size(0) > 0: - mate_pos = cached['mate_pos'] - mate_x = cached['mate_x'] + if self.include_mates and cached["mate_pos"].size(0) > 0: + mate_pos = cached["mate_pos"] + mate_x = cached["mate_x"] protein_pos = torch.cat([protein_pos, mate_pos], dim=0) protein_x = torch.cat([protein_x, mate_x], dim=0) # load mate residue indices (properly grouped by residue) # offset by max protein residue index - max_res_idx = protein_res_idx.max().item() if protein_res_idx.numel() > 0 else -1 - if 'mate_res_idx' in cached: - mate_res_idx = cached['mate_res_idx'] + max_res_idx + 1 + max_res_idx = ( + protein_res_idx.max().item() if protein_res_idx.numel() > 0 else -1 + ) + if "mate_res_idx" in cached: + mate_res_idx = cached["mate_res_idx"] + max_res_idx + 1 else: # fallback for old cache files without mate_res_idx mate_res_idx = torch.arange( max_res_idx + 1, max_res_idx + 1 + mate_pos.size(0), - dtype=torch.long + dtype=torch.long, ) protein_res_idx = torch.cat([protein_res_idx, mate_res_idx], dim=0) - - water_pos = cached['water_pos'] - water_x = cached['water_x'] + + water_pos = cached["water_pos"] + water_x = cached["water_x"] data = HeteroData() # compute total num_residues (protein + mates) - num_residues = int(protein_res_idx.max().item() + 1) if protein_res_idx.numel() > 0 else 0 + num_residues = ( + int(protein_res_idx.max().item() + 1) if protein_res_idx.numel() > 0 else 0 + ) - data['protein'].x = protein_x - data['protein'].pos = protein_pos - data['protein'].residue_index = protein_res_idx - data['protein'].num_nodes = protein_pos.size(0) - data['protein'].num_residues = num_residues - data['protein'].num_protein_residues = num_protein_residues # excludes mates + data["protein"].x = protein_x + data["protein"].pos = protein_pos + data["protein"].residue_index = protein_res_idx + data["protein"].num_nodes = protein_pos.size(0) + data["protein"].num_residues = num_residues + data["protein"].num_protein_residues = num_protein_residues # excludes mates # load SLAE embeddings if available (precomputed by scripts/precompute_slae_embeddings.py) - if 'protein_slae_embedding' in cached: - slae_emb = cached['protein_slae_embedding'] + if "protein_slae_embedding" in cached: + slae_emb = cached["protein_slae_embedding"] # handle mates: if mates were concatenated during preprocessing, embeddings include them - if self.include_mates and 'mate_slae_embedding' in cached: - mate_emb = cached['mate_slae_embedding'] + if self.include_mates and "mate_slae_embedding" in cached: + mate_emb = cached["mate_slae_embedding"] slae_emb = torch.cat([slae_emb, mate_emb], dim=0) - data['protein'].slae_embedding = slae_emb - - data['water'].x = water_x - data['water'].pos = water_pos - data['water'].num_nodes = water_pos.size(0) - + data["protein"].slae_embedding = slae_emb + + data["water"].x = water_x + data["water"].pos = water_pos + data["water"].num_nodes = water_pos.size(0) + if protein_pos.size(0) > 0: pp_edge_index = radius_graph(protein_pos, r=self.cutoff, loop=False) pp_edge_index = _make_undirected(pp_edge_index) data[EDGE_PP].edge_index = pp_edge_index else: data[EDGE_PP].edge_index = torch.empty((2, 0), dtype=torch.long) - + # store metadata - data.pdb_id = entry['cache_key'] + data.pdb_id = entry["cache_key"] data.num_asu_protein_atoms = num_asu_protein - + return data @@ -957,7 +999,7 @@ def get_dataloader( shuffle: bool = True, num_workers: int = 4, pin_memory: bool = False, - **dataset_kwargs + **dataset_kwargs, ) -> DataLoader: """ Create a DataLoader for crystal contact dataset. @@ -975,12 +1017,10 @@ def get_dataloader( Note: For single-protein overfitting, use duplicate_single_sample parameter: - duplicate_single_sample=100 creates 100 copies of the sample in the dataset - - Then batch_size works normally + - Then batch_size works normally """ dataset = ProteinWaterDataset( - pdb_list_file=pdb_list_file, - processed_dir=processed_dir, - **dataset_kwargs + pdb_list_file=pdb_list_file, processed_dir=processed_dir, **dataset_kwargs ) loader = DataLoader( diff --git a/src/encoder_base.py b/src/encoder_base.py index ee36f4b..49786d1 100644 --- a/src/encoder_base.py +++ b/src/encoder_base.py @@ -80,6 +80,7 @@ def build_encoder(config: dict, device: torch.device) -> BaseProteinEncoder: encoder_cls = get_encoder_class(encoder_type) return encoder_cls.from_config(config, device) + class BaseProteinEncoder(ABC, nn.Module): """ Abstract base class for protein encoders. @@ -134,6 +135,7 @@ def from_config(cls, config: dict, device: torch.device) -> BaseProteinEncoder: """ raise NotImplementedError("Subclasses must implement from_config") + @register_encoder("esm") @register_encoder("slae") class CachedEmbeddingEncoder(BaseProteinEncoder): @@ -159,7 +161,9 @@ class CachedEmbeddingEncoder(BaseProteinEncoder): are scalar-only. """ - def __init__(self, embedding_key: str, encoder_type: str, embedding_dim: int | None = None): + def __init__( + self, embedding_key: str, encoder_type: str, embedding_dim: int | None = None + ): """ Initialize CachedEmbeddingEncoder. @@ -193,7 +197,9 @@ def encoder_type(self) -> str: """Return encoder type identifier.""" return self._encoder_type - def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, tuple | None]: + def forward( + self, data: HeteroData + ) -> tuple[torch.Tensor, torch.Tensor, tuple | None]: """ Read cached embeddings and return (s, V, None). @@ -207,13 +213,13 @@ def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, tuple | V: (N, 0, 3) — empty vector features pp_edge_attr: None — cached embedding encoders don't process edges """ - if self._embedding_key not in data['protein']: + if self._embedding_key not in data["protein"]: raise KeyError( f"{self._encoder_type.upper()} encoder requires cached embeddings. " f"Please provide pre-computed '{self._embedding_key}' in data['protein']." ) - embeddings = data['protein'][self._embedding_key] + embeddings = data["protein"][self._embedding_key] # Infer dimension on first forward if self._embedding_dim is None: diff --git a/src/flow.py b/src/flow.py index 10b9507..104d3a4 100644 --- a/src/flow.py +++ b/src/flow.py @@ -17,11 +17,13 @@ from src.utils import ot_coupling -def build_knn_edges(src_pos: torch.Tensor, - dst_pos: torch.Tensor, - k: int, - batch_src: torch.Tensor | None = None, - batch_dst: torch.Tensor | None = None) -> torch.Tensor: +def build_knn_edges( + src_pos: torch.Tensor, + dst_pos: torch.Tensor, + k: int, + batch_src: torch.Tensor | None = None, + batch_dst: torch.Tensor | None = None, +) -> torch.Tensor: """ KNN edges from src -> dst (source indices in row 0, dest in row 1). """ @@ -62,28 +64,29 @@ def __init__( etypes = ALL_EDGE_TYPES - self.blocks = nn.ModuleList([ - GVPMultiEdgeConv( - etypes=etypes, - s_dim=s_h, v_dim=v_h, - rbf_dim=rbf_dim, - n_message_gvps=3, - n_update_gvps=3, - use_dst_feats=use_dst_feats, - drop_rate=drop_rate, - aggr_edges=aggr_edges, - activations=(F.relu, torch.sigmoid), - vector_gate=vector_gate, - ) - for _ in range(layers) - ]) + self.blocks = nn.ModuleList( + [ + GVPMultiEdgeConv( + etypes=etypes, + s_dim=s_h, + v_dim=v_h, + rbf_dim=rbf_dim, + n_message_gvps=3, + n_update_gvps=3, + use_dst_feats=use_dst_feats, + drop_rate=drop_rate, + aggr_edges=aggr_edges, + activations=(F.relu, torch.sigmoid), + vector_gate=vector_gate, + ) + for _ in range(layers) + ] + ) self.etypes = etypes - def build_edges(self, - data: HeteroData, - k_pw: int = 12, - k_ww: int = 8, - k_wp: int = 8) -> dict[tuple[str, str, str], torch.Tensor]: + def build_edges( + self, data: HeteroData, k_pw: int = 12, k_ww: int = 8, k_wp: int = 8 + ) -> dict[tuple[str, str, str], torch.Tensor]: """ Build KNN edges for protein-water interactions. @@ -104,26 +107,32 @@ def build_edges(self, Dict mapping edge type tuples to (2, E) edge index tensors """ edge_index_dict: dict[tuple[str, str, str], torch.Tensor] = {} - device = data['protein'].pos.device + device = data["protein"].pos.device - batch_p = data['protein'].batch if 'batch' in data['protein'] else None - batch_w = data['water'].batch if 'batch' in data['water'] else None + batch_p = data["protein"].batch if "batch" in data["protein"] else None + batch_w = data["water"].batch if "batch" in data["water"] else None - pos_p = data['protein'].pos - pos_w = data['water'].pos + pos_p = data["protein"].pos + pos_w = data["water"].pos # protein -> water if pos_p.numel() > 0 and pos_w.numel() > 0: # p->w - ei_pw = build_knn_edges(pos_p, pos_w, k=k_pw, batch_src=batch_p, batch_dst=batch_w) + ei_pw = build_knn_edges( + pos_p, pos_w, k=k_pw, batch_src=batch_p, batch_dst=batch_w + ) # w->p then reverse - ei_wp = build_knn_edges(pos_w, pos_p, k=k_pw, batch_src=batch_w, batch_dst=batch_p) + ei_wp = build_knn_edges( + pos_w, pos_p, k=k_pw, batch_src=batch_w, batch_dst=batch_p + ) ei_wp_reversed = ei_wp.flip(0) # union ei_pw_union = torch.cat([ei_pw, ei_wp_reversed], dim=1).unique(dim=1) edge_index_dict[EDGE_PW] = ei_pw_union else: - edge_index_dict[EDGE_PW] = torch.empty(2, 0, dtype=torch.long, device=device) + edge_index_dict[EDGE_PW] = torch.empty( + 2, 0, dtype=torch.long, device=device + ) # water -> water if pos_w.numel() > 0: @@ -131,7 +140,9 @@ def build_edges(self, pos_w, pos_w, k=k_ww, batch_src=batch_w, batch_dst=batch_w ) else: - edge_index_dict[EDGE_WW] = torch.empty(2, 0, dtype=torch.long, device=device) + edge_index_dict[EDGE_WW] = torch.empty( + 2, 0, dtype=torch.long, device=device + ) # protein-protein edges (cached from dataset) if EDGE_PP in data.edge_types: @@ -147,7 +158,9 @@ def build_edges(self, pos_w, pos_p, k=k_wp, batch_src=batch_w, batch_dst=batch_p ) else: - edge_index_dict[EDGE_WP] = torch.empty(2, 0, dtype=torch.long, device=device) + edge_index_dict[EDGE_WP] = torch.empty( + 2, 0, dtype=torch.long, device=device + ) for et in self.etypes: if et not in edge_index_dict: @@ -155,23 +168,27 @@ def build_edges(self, return edge_index_dict - def forward(self, - x_dict: dict[str, tuple[torch.Tensor, torch.Tensor]], - data: HeteroData, - k_pw: int = 12, - k_ww: int = 8, - k_wp: int = 8): + def forward( + self, + x_dict: dict[str, tuple[torch.Tensor, torch.Tensor]], + data: HeteroData, + k_pw: int = 12, + k_ww: int = 8, + k_wp: int = 8, + ): """ x_dict: { 'protein': (s_p, v_p), 'water': (s_w, v_w) } """ - pos_dict = {nt: data[nt].pos for nt in data.node_types if 'pos' in data[nt]} + pos_dict = {nt: data[nt].pos for nt in data.node_types if "pos" in data[nt]} edge_index_dict = self.build_edges( data, - k_pw=k_pw, k_ww=k_ww, k_wp=k_wp, + k_pw=k_pw, + k_ww=k_ww, + k_wp=k_wp, ) for block in self.blocks: @@ -269,10 +286,12 @@ def __init__( vector_gate=True, ) - def forward(self, - data: HeteroData, - t: torch.Tensor, - sc: dict[str, torch.Tensor] | None = None) -> torch.Tensor: + def forward( + self, + data: HeteroData, + t: torch.Tensor, + sc: dict[str, torch.Tensor] | None = None, + ) -> torch.Tensor: """ data: HeteroData with node types 'protein' (may include mates), 'water' t: (B,) diffusion time per complex @@ -281,7 +300,7 @@ def forward(self, Returns: v_pred: (N_water, 3) vector field at each water node. """ - device = data['protein'].pos.device + device = data["protein"].pos.device # Single unified encoder call - works for ANY encoder type s_all, v_all, _edge_attr = self.encoder(data) @@ -291,33 +310,33 @@ def forward(self, encoder_input = (s_all, v_all) if self.encoder.output_dims[1] > 0 else s_all s_p_latent, v_p_latent = self.encoder_to_flow(encoder_input) - if 'water' not in data.node_types or data['water'].num_nodes == 0: + if "water" not in data.node_types or data["water"].num_nodes == 0: return torch.zeros(0, 3, device=device) - batch_p = data['protein'].batch - batch_w = data['water'].batch + batch_p = data["protein"].batch + batch_w = data["water"].batch t_p = t[batch_p].unsqueeze(-1) t_w = t[batch_w].unsqueeze(-1) s_p = self.protein_scalar_encoder(torch.cat([s_p_latent, t_p], dim=-1)) - s_w = self.water_scalar_encoder(torch.cat([data['water'].x, t_w], dim=-1)) + s_w = self.water_scalar_encoder(torch.cat([data["water"].x, t_w], dim=-1)) # initial water vectors (all zeros to start) v_w = torch.zeros( - data['water'].num_nodes, + data["water"].num_nodes, self.hidden_dims[1], 3, device=device, ) # self conditioning - if sc is not None and ('x1_pred' in sc) and sc['x1_pred'] is not None: - delta = (sc['x1_pred'] - data['water'].pos) + if sc is not None and ("x1_pred" in sc) and sc["x1_pred"] is not None: + delta = sc["x1_pred"] - data["water"].pos delta_vec = delta.unsqueeze(1) # vector conditioning (equivariant) - s_empty = torch.empty(data['water'].num_nodes, 0, device=device) + s_empty = torch.empty(data["water"].num_nodes, 0, device=device) _, v_sc = self.sc_vec_encoder((s_empty, delta_vec)) v_w = v_w + v_sc @@ -328,8 +347,8 @@ def forward(self, # build hetero feature dict for GVP multi-edge updates x_dict = { - 'protein': (s_p, v_p_latent), - 'water': (s_w, v_w), + "protein": (s_p, v_p_latent), + "water": (s_w, v_w), } # hetero update (protein+water graph) @@ -342,7 +361,7 @@ def forward(self, ) # water vector field head - _, v_pred = self.vfield_head(x_dict['water']) + _, v_pred = self.vfield_head(x_dict["water"]) return v_pred.squeeze(1) @@ -372,24 +391,26 @@ def __init__( @staticmethod def compute_sigma(data: HeteroData) -> float: """Compute sigma as std of protein coordinates.""" - pos = data['protein'].pos + pos = data["protein"].pos return float(pos.std().item()) @staticmethod - def compute_sigma_per_graph(data: HeteroData | Batch, device: torch.device) -> torch.Tensor: + def compute_sigma_per_graph( + data: HeteroData | Batch, device: torch.device + ) -> torch.Tensor: """ Compute sigma (std of protein coordinates) per graph in a batch. Returns: sigma: (num_graphs,) tensor of sigma values per graph """ - pos = data['protein'].pos # (N_total, 3) - batch_p = data['protein'].batch # (N_total,) + pos = data["protein"].pos # (N_total, 3) + batch_p = data["protein"].batch # (N_total,) # Var(X) = E[X^2] - E[X]^2 mean_pos = scatter_mean(pos, batch_p, dim=0) # (num_graphs, 3) - mean_sq = scatter_mean(pos ** 2, batch_p, dim=0) # (num_graphs, 3) - var_per_dim = mean_sq - mean_pos ** 2 # (num_graphs, 3) + mean_sq = scatter_mean(pos**2, batch_p, dim=0) # (num_graphs, 3) + var_per_dim = mean_sq - mean_pos**2 # (num_graphs, 3) sigma = torch.sqrt(var_per_dim.mean(dim=-1).clamp(min=1e-8)) # (num_graphs,) return sigma @@ -403,15 +424,15 @@ def training_step( ) -> dict[str, float | dict | None]: """ Single flow matching training step. - + Returns dict with 'loss', 'rmsd', 'sigma'. """ self.model.train() - device = batch['protein'].pos.device + device = batch["protein"].pos.device - x1 = batch['water'].pos - batch_w = batch['water'].batch - batch_p = batch['protein'].batch + x1 = batch["water"].pos + batch_w = batch["water"].batch + batch_p = batch["protein"].batch num_graphs = int(batch_p.max().item()) + 1 sigma = self.compute_sigma(batch) @@ -436,13 +457,13 @@ def training_step( sc = None if use_self_conditioning and torch.rand(1).item() < self.p_self_cond: with torch.no_grad(): - batch['water'].pos = x_t + batch["water"].pos = x_t v_sc = self.model(batch, t, sc=None) x1_sc = x_t + (1.0 - t_per_atom) * v_sc - sc = {'x1_pred': x1_sc} + sc = {"x1_pred": x1_sc} # forward pass - batch['water'].pos = x_t + batch["water"].pos = x_t v_pred = self.model(batch, t, sc=sc) # target velocity @@ -458,12 +479,13 @@ def training_step( if loss.item() > 100.0: with torch.no_grad(): from torch_scatter import scatter_add + weighted_mse = (w * per_atom_mse).squeeze(-1) # compute per-graph loss: sum(weighted_mse) / sum(w) for each graph numerator = scatter_add(weighted_mse, batch_w, dim=0) denominator = scatter_add(w.squeeze(-1), batch_w, dim=0) per_sample_loss = numerator / (denominator + 1e-8) - per_sample_info = {'losses': per_sample_loss, 'num_graphs': num_graphs} + per_sample_info = {"losses": per_sample_loss, "num_graphs": num_graphs} # backward optimizer.zero_grad(set_to_none=True) @@ -471,30 +493,35 @@ def training_step( if grad_clip > 0: torch.nn.utils.clip_grad_norm_( [p for p in self.model.parameters() if p.requires_grad], - max_norm=grad_clip + max_norm=grad_clip, ) optimizer.step() # training RMSD with torch.no_grad(): x1_hat = x_t + (1.0 - t_per_atom) * v_pred - # rmsd = compute_rmsd(x1_hat, x1_star) + # rmsd = compute_rmsd(x1_hat, x1_star) # on-gpu version of rmsd diff2 = ((x1_hat - x1_star) ** 2).sum(-1) # (Nw,) rmsd = torch.sqrt(scatter_mean(diff2, batch_w, dim=0)).mean().item() - return {'loss': loss.item(), 'rmsd': rmsd, 'sigma': sigma, 'per_sample_info': per_sample_info} + return { + "loss": loss.item(), + "rmsd": rmsd, + "sigma": sigma, + "per_sample_info": per_sample_info, + } @torch.no_grad() def validation_step(self, batch: HeteroData) -> dict[str, float]: """Single validation step (no gradient, no optimizer).""" self.model.eval() - device = batch['protein'].pos.device + device = batch["protein"].pos.device - x1 = batch['water'].pos - batch_w = batch['water'].batch - batch_p = batch['protein'].batch + x1 = batch["water"].pos + batch_w = batch["water"].batch + batch_p = batch["protein"].batch num_graphs = int(batch_p.max().item()) + 1 sigma = self.compute_sigma(batch) @@ -505,7 +532,7 @@ def validation_step(self, batch: HeteroData) -> dict[str, float]: t_per_atom = t[batch_w].unsqueeze(-1) x_t = (1.0 - t_per_atom) * x0_star + t_per_atom * x1_star - batch['water'].pos = x_t + batch["water"].pos = x_t v_pred = self.model(batch, t, sc=None) v_target = x1_star - x0_star @@ -518,7 +545,7 @@ def validation_step(self, batch: HeteroData) -> dict[str, float]: diff2 = ((x1_hat - x1_star) ** 2).sum(-1) # (Nw,) rmsd = torch.sqrt(scatter_mean(diff2, batch_w, dim=0)).mean().item() - return {'loss': loss.item(), 'rmsd': rmsd} + return {"loss": loss.item(), "rmsd": rmsd} def _setup_water_nodes_from_ratio( self, @@ -538,7 +565,7 @@ def _setup_water_nodes_from_ratio( x: (N_water_total, 3) initial noise positions batch_w: (N_water_total,) batch indices """ - num_residues = g['protein'].num_residues # (num_graphs,) + num_residues = g["protein"].num_residues # (num_graphs,) num_graphs = num_residues.size(0) # compute waters per graph: num_residues * ratio, minimum 1 @@ -562,10 +589,10 @@ def _setup_water_nodes_from_ratio( water_x[:, 2] = 1.0 # oxygen is index 2 in ELEMENT_VOCAB # update graph with new water nodes - g['water'].pos = x - g['water'].x = water_x - g['water'].batch = batch_w - g['water'].num_nodes = total_waters + g["water"].pos = x + g["water"].x = water_x + g["water"].batch = batch_w + g["water"].num_nodes = total_waters return x, batch_w @@ -607,14 +634,16 @@ def euler_integrate( if water_ratio is not None: # sample waters based on residue count x, batch_w = self._setup_water_nodes_from_ratio(g, water_ratio, device) - num_graphs = g['protein'].num_residues.size(0) + num_graphs = g["protein"].num_residues.size(0) else: # use existing water nodes - batch_w = g['water'].batch + batch_w = g["water"].batch num_graphs = int(batch_w.max().item()) + 1 sigma_per_graph = self.compute_sigma_per_graph(g, device) sigma_per_water = sigma_per_graph[batch_w] - x = torch.randn(g['water'].num_nodes, 3, device=device) * sigma_per_water.unsqueeze(-1) + x = torch.randn( + g["water"].num_nodes, 3, device=device + ) * sigma_per_water.unsqueeze(-1) x1_pred_ema = x.clone() @@ -625,18 +654,20 @@ def euler_integrate( t_scalar = ts[i] t = t_scalar.expand(num_graphs) # (num_graphs,) all same value - g['water'].pos = x - sc = {'x1_pred': x1_pred_ema} if use_sc else None + g["water"].pos = x + sc = {"x1_pred": x1_pred_ema} if use_sc else None v = self.model(g, t, sc=sc) x = x + dt * v if use_sc: t_next_scalar = ts[i + 1] t_next = t_next_scalar.expand(num_graphs) - g['water'].pos = x - v_next = self.model(g, t_next, sc={'x1_pred': x1_pred_ema}) + g["water"].pos = x + v_next = self.model(g, t_next, sc={"x1_pred": x1_pred_ema}) x1_pred_now = x + (1.0 - t_next_scalar) * v_next - x1_pred_ema = (1.0 - sc_ema_alpha) * x1_pred_ema + sc_ema_alpha * x1_pred_now + x1_pred_ema = ( + 1.0 - sc_ema_alpha + ) * x1_pred_ema + sc_ema_alpha * x1_pred_now # split results by graph x_cpu = x.detach().cpu() @@ -686,24 +717,24 @@ def rk4_integrate( graphs = [graphs] # store original pdb_ids before batching - pdb_ids = [getattr(g, 'pdb_id', None) for g in graphs] + pdb_ids = [getattr(g, "pdb_id", None) for g in graphs] # batch graphs together g = Batch.from_data_list([copy.deepcopy(graph) for graph in graphs]).to(device) - batch_p = g['protein'].batch + batch_p = g["protein"].batch # store ground truth water positions and batch indices before modifying - x1_true = g['water'].pos.clone() - batch_w_true = g['water'].batch.clone() + x1_true = g["water"].pos.clone() + batch_w_true = g["water"].batch.clone() if water_ratio is not None: # sample waters based on residue count x, batch_w = self._setup_water_nodes_from_ratio(g, water_ratio, device) - num_graphs = g['protein'].num_residues.size(0) + num_graphs = g["protein"].num_residues.size(0) else: # use existing water nodes - batch_w = g['water'].batch + batch_w = g["water"].batch num_graphs = int(batch_w.max().item()) + 1 sigma_per_graph = self.compute_sigma_per_graph(g, device) sigma_per_water = sigma_per_graph[batch_w] @@ -729,8 +760,8 @@ def rk4_integrate( t0 = t0_scalar.expand(num_graphs) # (num_graphs,) all same value def f(xpos, t_tensor): - g['water'].pos = xpos - sc = {'x1_pred': x1_pred_ema} if use_sc else None + g["water"].pos = xpos + sc = {"x1_pred": x1_pred_ema} if use_sc else None return self.model(g, t_tensor, sc=sc) k1 = f(x, t0) @@ -743,10 +774,12 @@ def f(xpos, t_tensor): if use_sc: t1_scalar = ts[step + 1] t1 = t1_scalar.expand(num_graphs) - g['water'].pos = x - v_next = self.model(g, t1, sc={'x1_pred': x1_pred_ema}) + g["water"].pos = x + v_next = self.model(g, t1, sc={"x1_pred": x1_pred_ema}) x1_pred_now = x + (1.0 - t1_scalar) * v_next - x1_pred_ema = (1.0 - sc_ema_alpha) * x1_pred_ema + sc_ema_alpha * x1_pred_now + x1_pred_ema = ( + 1.0 - sc_ema_alpha + ) * x1_pred_ema + sc_ema_alpha * x1_pred_now if return_trajectory: x_cpu = x.detach().cpu() @@ -756,7 +789,7 @@ def f(xpos, t_tensor): # split results by graph x_cpu = x.detach().cpu() - protein_pos_cpu = g['protein'].pos.detach().cpu() + protein_pos_cpu = g["protein"].pos.detach().cpu() x1_true_cpu = x1_true.detach().cpu() batch_w_cpu = batch_w.cpu() batch_w_true_cpu = batch_w_true.cpu() @@ -769,14 +802,14 @@ def f(xpos, t_tensor): mask_p = batch_p_cpu == i result = { - 'protein_pos': protein_pos_cpu[mask_p].numpy(), - 'water_true': x1_true_cpu[mask_w_true].numpy(), - 'water_pred': x_cpu[mask_w].numpy(), - 'pdb_id': pdb_ids[i], + "protein_pos": protein_pos_cpu[mask_p].numpy(), + "water_true": x1_true_cpu[mask_w_true].numpy(), + "water_pred": x_cpu[mask_w].numpy(), + "pdb_id": pdb_ids[i], } if return_trajectory: - result['trajectory'] = trajectories[i] + result["trajectory"] = trajectories[i] results.append(result) @@ -812,7 +845,7 @@ def sample( results = self.rk4_integrate( graphs, num_steps, use_sc, device=device, return_trajectory=False ) - results = [r['water_pred'] for r in results] + results = [r["water_pred"] for r in results] else: raise ValueError(f"Unknown method: {method}") @@ -820,4 +853,3 @@ def sample( if single_input: return results[0] return results - diff --git a/src/gvp.py b/src/gvp.py index 2496ae6..96e80d8 100644 --- a/src/gvp.py +++ b/src/gvp.py @@ -1,4 +1,4 @@ -# gvp.py +# gvp.py # generic GVP and GVPConv layers adapted from Jing et al. (2021) import functools @@ -430,24 +430,30 @@ def forward(self, x, edge_index, edge_attr, autoregressive_x=None, node_mask=Non x = x_ return x + class EdgeUpdate(nn.Module): """ Residual edge update that keeps a fixed scalar edge width across layers. Input to the MLP is [s_src, s_dst, s_edge( fixed width ), (optional) distance/RBF]. Output shape == input edge width. No width drift. """ + def __init__( self, - n_node_scalars: int, # S_node (e.g., 256) - s_edge_width: int, # fixed model edge width used everywhere + n_node_scalars: int, # S_node (e.g., 256) + s_edge_width: int, # fixed model edge width used everywhere update_w_distance_features: bool = False, - distance_dim: int = 0, # e.g., RBF size + distance_dim: int = 0, # e.g., RBF size ): super().__init__() self.update_w_distance_features = update_w_distance_features self.s_edge_width = s_edge_width - in_dim = (2 * n_node_scalars) + s_edge_width + (distance_dim if update_w_distance_features else 0) + in_dim = ( + (2 * n_node_scalars) + + s_edge_width + + (distance_dim if update_w_distance_features else 0) + ) self.edge_mlp = nn.Sequential( nn.Linear(in_dim, s_edge_width), @@ -459,47 +465,53 @@ def __init__( def forward( self, - node_tuple: tuple, # (s_node, V_node) with s_node: (N, S_node) - edge_index: torch.Tensor, # (2, E) - edge_attr: tuple, # (s_edge, V_edge) with s_edge: (E, s_edge_width) + node_tuple: tuple, # (s_node, V_node) with s_node: (N, S_node) + edge_index: torch.Tensor, # (2, E) + edge_attr: tuple, # (s_edge, V_edge) with s_edge: (E, s_edge_width) distance_feat: torch.Tensor | None = None, # (E, D) if enabled ) -> tuple: s_node, _ = node_tuple s_edge, V_edge = edge_attr - assert s_edge.shape[-1] == self.s_edge_width, \ + assert s_edge.shape[-1] == self.s_edge_width, ( f"EdgeUpdate expected width {self.s_edge_width}, got {s_edge.shape[-1]}" + ) src, dst = edge_index[0], edge_index[1] parts = [s_node[src], s_node[dst], s_edge] - + if self.update_w_distance_features: parts.append(distance_feat) - h = torch.cat(parts, dim=-1) # (E, 2*S_node + s_edge_width (+D)) - upd = self.edge_mlp(h) # (E, s_edge_width) - s_edge = self.edge_norm(s_edge + upd) # residual, fixed width - return (s_edge, V_edge) # vectors unchanged - + h = torch.cat(parts, dim=-1) # (E, 2*S_node + s_edge_width (+D)) + upd = self.edge_mlp(h) # (E, s_edge_width) + s_edge = self.edge_norm(s_edge + upd) # residual, fixed width + return (s_edge, V_edge) # vectors unchanged + + # multi edge gvp + class GVPMultiEdge(MessagePassing): """ Per-edge-type GVP message function (messages only, no residual/update). Outputs a merged tensor for PyG aggregation: _merge(s_msg, v_msg). """ + def __init__( self, - src_type: str, dst_type: str, - s_dim: int, v_dim: int, - rbf_dim: int = 16, + src_type: str, + dst_type: str, + s_dim: int, + v_dim: int, + rbf_dim: int = 16, use_dst_feats: bool = False, n_message_gvps: int = 1, activations=(F.relu, torch.sigmoid), vector_gate=True, - aggr="sum", - rbf_dmax: float = 20.0 + aggr="sum", + rbf_dmax: float = 20.0, ): super().__init__(aggr=aggr) self.src_type, self.dst_type = src_type, dst_type @@ -512,24 +524,36 @@ def __init__( # message GVP stack; first layer takes [unit_vec] and [rbf] extras msg_layers = [] for i in range(n_message_gvps): - vin = v_dim + (1 if i == 0 else 0) + (v_dim if (i == 0 and use_dst_feats) else 0) - sin = s_dim + (rbf_dim if i == 0 else 0) + (s_dim if (i == 0 and use_dst_feats) else 0) - msg_layers.append(GVP_( - in_dims=(sin, vin), - out_dims=(s_dim, v_dim), - vector_gate=True, - activations=activations - )) + vin = ( + v_dim + + (1 if i == 0 else 0) + + (v_dim if (i == 0 and use_dst_feats) else 0) + ) + sin = ( + s_dim + + (rbf_dim if i == 0 else 0) + + (s_dim if (i == 0 and use_dst_feats) else 0) + ) + msg_layers.append( + GVP_( + in_dims=(sin, vin), + out_dims=(s_dim, v_dim), + vector_gate=True, + activations=activations, + ) + ) self.edge_message = nn.Sequential(*msg_layers) @staticmethod def _unit_and_rbf(pos_src, pos_dst, edge_index, rbf_dim, rbf_dmax): # displacement from src->dst (dst - src) to match your original convention s, d = edge_index - disp = pos_dst[d] - pos_src[s] # (E, 3) + disp = pos_dst[d] - pos_src[s] # (E, 3) dist = torch.clamp(disp.norm(dim=-1, keepdim=True), min=1e-5) # (E,1) - unit = disp / dist # (E, 3) - rbf_e = rbf(dist.squeeze(-1), num_gaussians=rbf_dim, cutoff=rbf_dmax) # (E, rbf_dim) + unit = disp / dist # (E, 3) + rbf_e = rbf( + dist.squeeze(-1), num_gaussians=rbf_dim, cutoff=rbf_dmax + ) # (E, rbf_dim) return unit, rbf_e def forward(self, x, edge_index, pos_pair): @@ -546,7 +570,7 @@ def forward(self, x, edge_index, pos_pair): s_src = s_dst = x[0] v_src = v_dst = x[1] - #edge case for no edges being made for weird reason (debugs in flow) + # edge case for no edges being made for weird reason (debugs in flow) if edge_index.numel() == 0: N_dst = s_dst.size(0) return s_dst.new_zeros(N_dst, self.s_dim + 3 * self.v_dim) @@ -558,7 +582,9 @@ def forward(self, x, edge_index, pos_pair): v_dst_f = v_dst.reshape(v_dst.size(0), -1) # (N_dst, 3*v) # compute geometric edge attributes on the fly - unit, rbf_e = self._unit_and_rbf(pos_src, pos_dst, edge_index, self.rbf_dim, self.rbf_dmax) + unit, rbf_e = self._unit_and_rbf( + pos_src, pos_dst, edge_index, self.rbf_dim, self.rbf_dmax + ) # pack as simple tensors for message() # edge_attr = (rbf_e, unit) -- pass separately to keep shapes explicit @@ -568,18 +594,18 @@ def forward(self, x, edge_index, pos_pair): s=(s_src, s_dst), vf=(v_src_f, v_dst_f), rbf_e=rbf_e, - unit=unit + unit=unit, ) return out # merged messages (s_msg + flattened v_msg) def message(self, s_i, s_j, vf_i, vf_j, rbf_e, unit): # Unflatten vectors - v_j = vf_j.view(vf_j.size(0), -1, 3) # (E, v, 3) - v_i = vf_i.view(vf_i.size(0), -1, 3) # (E, v, 3) + v_j = vf_j.view(vf_j.size(0), -1, 3) # (E, v, 3) + v_i = vf_i.view(vf_i.size(0), -1, 3) # (E, v, 3) # Build vector features: [unit_vec] (+ src v) (+ dst v if use_dst_feats) - v_list = [unit.unsqueeze(1), v_j] # (E, 1,3) + (E, v,3) - s_list = [s_j, rbf_e] # (E, s_dim) + (E, rbf_dim) + v_list = [unit.unsqueeze(1), v_j] # (E, 1,3) + (E, v,3) + s_list = [s_j, rbf_e] # (E, s_dim) + (E, rbf_dim) if self.use_dst_feats: v_list.append(v_i) s_list.append(s_i) @@ -588,23 +614,26 @@ def message(self, s_i, s_j, vf_i, vf_j, rbf_e, unit): s_in = torch.cat(s_list, dim=1) s_msg, v_msg = self.edge_message((s_in, v_in)) # (E, s_dim), (E, v, 3) - return _merge(s_msg, v_msg) # (E, s_dim + 3*v) + return _merge(s_msg, v_msg) # (E, s_dim + 3*v) + class GVPMultiEdgeConv(nn.Module): """ Hetero multi-edge message passing – messages only per relation, then single residual+FFN+layernorm per destination node type. """ + def __init__( self, - etypes, # List[EdgeType] - s_dim: int, v_dim: int, + etypes, # List[EdgeType] + s_dim: int, + v_dim: int, rbf_dim: int = 16, n_message_gvps: int = 1, n_update_gvps: int = 1, use_dst_feats: bool = False, drop_rate: float = 0.1, - aggr_edges: str = "sum", # 'mean' or 'add' per edge aggregation + aggr_edges: str = "sum", # 'mean' or 'add' per edge aggregation activations=(F.relu, torch.sigmoid), vector_gate=True, ): @@ -614,8 +643,12 @@ def __init__( # per-dst-type norms and update stacks dst_ntypes = sorted({dst for (_, _, dst) in etypes}) - self.msg_norms = nn.ModuleDict({nt: LayerNorm((s_dim, v_dim)) for nt in dst_ntypes}) - self.upd_norms = nn.ModuleDict({nt: LayerNorm((s_dim, v_dim)) for nt in dst_ntypes}) + self.msg_norms = nn.ModuleDict( + {nt: LayerNorm((s_dim, v_dim)) for nt in dst_ntypes} + ) + self.upd_norms = nn.ModuleDict( + {nt: LayerNorm((s_dim, v_dim)) for nt in dst_ntypes} + ) GVP_ = functools.partial(GVP, activations=activations, vector_gate=vector_gate) self.node_updates = nn.ModuleDict() @@ -627,13 +660,18 @@ def __init__( # per-edge-type message convs feeding a HeteroConv(aggr='sum' across relations) rel_convs = {} - for (src, rel, dst) in etypes: + for src, rel, dst in etypes: rel_convs[(src, rel, dst)] = GVPMultiEdge( - src, dst, s_dim, v_dim, - rbf_dim=rbf_dim, use_dst_feats=use_dst_feats, + src, + dst, + s_dim, + v_dim, + rbf_dim=rbf_dim, + use_dst_feats=use_dst_feats, n_message_gvps=n_message_gvps, - activations=activations, vector_gate=vector_gate, - aggr=("mean" if aggr_edges == "mean" else "sum") + activations=activations, + vector_gate=vector_gate, + aggr=("mean" if aggr_edges == "mean" else "sum"), ) self.hconv = HeteroConv(rel_convs, aggr="sum") # sum messages across edge types @@ -642,25 +680,28 @@ def forward(self, x_dict, edge_index_dict, pos_dict): x_dict: {ntype: (s, v)} returns updated x_dict """ - #build per-edge-type (pos_src, pos_dst) tuples + # build per-edge-type (pos_src, pos_dst) tuples pos_pair_dict = { - et: (pos_dict[et[0]], pos_dict[et[2]]) - for et in edge_index_dict.keys() + et: (pos_dict[et[0]], pos_dict[et[2]]) for et in edge_index_dict.keys() } - #heteroConv will forward kwarg name without '_dict' into each conv + # heteroConv will forward kwarg name without '_dict' into each conv merged_msgs = self.hconv(x_dict, edge_index_dict, pos_pair_dict=pos_pair_dict) - #merged_msgs[ntype] is a merged tensor of total messages: [N, s_dim + 3*v_dim] + # merged_msgs[ntype] is a merged tensor of total messages: [N, s_dim + 3*v_dim] for ntype, merged in merged_msgs.items(): s_msg, v_msg = _split(merged, self.v_dim) - #apply dropout, residual add, layernorm (message stage) + # apply dropout, residual add, layernorm (message stage) s_old, v_old = x_dict[ntype] s_msg, v_msg = self.drop((s_msg, v_msg)) - s_mid, v_mid = self.msg_norms[ntype](tuple_sum((s_old, v_old), (s_msg, v_msg))) + s_mid, v_mid = self.msg_norms[ntype]( + tuple_sum((s_old, v_old), (s_msg, v_msg)) + ) - #per-node update stack (GVP x N), then dropout + residual + layernorm + # per-node update stack (GVP x N), then dropout + residual + layernorm s_res, v_res = self.node_updates[ntype]((s_mid, v_mid)) s_res, v_res = self.drop((s_res, v_res)) - x_dict[ntype] = self.upd_norms[ntype](tuple_sum((s_mid, v_mid), (s_res, v_res))) + x_dict[ntype] = self.upd_norms[ntype]( + tuple_sum((s_mid, v_mid), (s_res, v_res)) + ) return x_dict diff --git a/src/gvp_encoder.py b/src/gvp_encoder.py index 0ff1c2f..7aed3d4 100644 --- a/src/gvp_encoder.py +++ b/src/gvp_encoder.py @@ -4,6 +4,7 @@ This encoder processes protein structure directly using GVP layers to produce geometric features for the flow model. """ + from __future__ import annotations from collections.abc import Callable @@ -23,7 +24,9 @@ from src.utils import rbf -def edge_vectors(pos: torch.Tensor, edge_index: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: +def edge_vectors( + pos: torch.Tensor, edge_index: torch.Tensor +) -> tuple[torch.Tensor, torch.Tensor]: """ Compute edge distances and unit vectors from node positions. @@ -55,8 +58,8 @@ def make_gvp_encoder_data(data: HeteroData) -> Data: Returns: enc_data: Data with x, pos, edge_index, and optionally cached edge features """ - device = data['protein'].pos.device - prot = data['protein'] + device = data["protein"].pos.device + prot = data["protein"] x = prot.x pos = prot.pos @@ -72,9 +75,9 @@ def make_gvp_encoder_data(data: HeteroData) -> Data: # Copy cached edge features if available if EDGE_PP in data.edge_types: pp_edge = data[EDGE_PP] - if hasattr(pp_edge, 'edge_rbf'): + if hasattr(pp_edge, "edge_rbf"): enc_data.edge_rbf = pp_edge.edge_rbf - if hasattr(pp_edge, 'edge_unit'): + if hasattr(pp_edge, "edge_unit"): enc_data.edge_unit = pp_edge.edge_unit # batch for multi-complex batches @@ -186,23 +189,27 @@ def __init__( self.s_edge_width = n_edge_scalar_out if n_edge_scalar_in != n_edge_scalar_out: - self.edge_in_proj = nn.Linear(n_edge_scalar_in, n_edge_scalar_out, bias=False) + self.edge_in_proj = nn.Linear( + n_edge_scalar_in, n_edge_scalar_out, bias=False + ) else: self.edge_in_proj = nn.Identity() edge_dims = (self.s_edge_width, n_edge_vec_in) - self.layers = nn.ModuleList([ - GVPConvLayer( - node_dims=hidden_dims, - edge_dims=edge_dims, - n_message=n_message, - n_feedforward=n_feedforward, - drop_rate=drop_rate, - activations=activations, - vector_gate=vector_gate, - ) - for _ in range(n_layers) - ]) + self.layers = nn.ModuleList( + [ + GVPConvLayer( + node_dims=hidden_dims, + edge_dims=edge_dims, + n_message=n_message, + n_feedforward=n_feedforward, + drop_rate=drop_rate, + activations=activations, + vector_gate=vector_gate, + ) + for _ in range(n_layers) + ] + ) if use_edge_update: self.edge_update = EdgeUpdate( @@ -259,7 +266,9 @@ def _pool_by_residue( elif aggr == "sum": return scatter_add(atom_embed, residue_index, dim=0, dim_size=num_residues) elif aggr == "max": - out, _ = scatter_max(atom_embed, residue_index, dim=0, dim_size=num_residues) + out, _ = scatter_max( + atom_embed, residue_index, dim=0, dim_size=num_residues + ) return out else: raise ValueError(f"Unknown pool_aggr={aggr!r}") @@ -268,7 +277,9 @@ def _pool_by_residue( def _initial_node_tuple( x_scalar: torch.Tensor, device: torch.device | None = None ) -> tuple[torch.Tensor, torch.Tensor]: - zeros = torch.zeros(x_scalar.size(0), 1, 3, device=x_scalar.device if device is None else device) + zeros = torch.zeros( + x_scalar.size(0), 1, 3, device=x_scalar.device if device is None else device + ) return (x_scalar, zeros) def _compute_edge_attr(self, data: Batch): @@ -286,7 +297,7 @@ def _compute_edge_attr(self, data: Batch): s_edge_raw: Raw RBF features (for distance conditioning) """ # Use cached features if available - if hasattr(data, 'edge_rbf') and hasattr(data, 'edge_unit'): + if hasattr(data, "edge_rbf") and hasattr(data, "edge_unit"): s_edge_raw = data.edge_rbf u = data.edge_unit else: @@ -332,15 +343,21 @@ def forward(self, data: Batch) -> tuple[tuple, tuple | None]: node_tuple=x, edge_index=data.edge_index, edge_attr=edge_attr, - distance_feat=(dist_feat if self.update_w_distance_features else None), + distance_feat=( + dist_feat if self.update_w_distance_features else None + ), ) if self.pool_residue: if not (hasattr(data, "residue_index") and hasattr(data, "num_residues")): - raise ValueError("Pooling requires data.residue_index and data.num_residues") + raise ValueError( + "Pooling requires data.residue_index and data.num_residues" + ) atom_dense = self._tuple_to_scalar_dense(x) atom_embed = self.atom_readout(atom_dense) - res_embed = self._pool_by_residue(atom_embed, data.residue_index, int(data.num_residues)) + res_embed = self._pool_by_residue( + atom_embed, data.residue_index, int(data.num_residues) + ) return res_embed, None # No edge features when pooling # Return edge_attr only if edge_update was used @@ -390,7 +407,9 @@ def load_encoder_from_checkpoint( # use node_scalar_in from checkpoint if not specified if node_scalar_in is None: if "node_scalar_in" not in args: - raise ValueError("node_scalar_in not in checkpoint and not provided") + raise ValueError( + "node_scalar_in not in checkpoint and not provided" + ) node_scalar_in = args["node_scalar_in"] elif "node_scalar_in" in args and args["node_scalar_in"] != node_scalar_in: raise ValueError( @@ -403,7 +422,9 @@ def load_encoder_from_checkpoint( logger.warning(f"Failed to load checkpoint {checkpoint_path}: {e}") logger.info("Initializing blank encoder instead.") else: - logger.warning(f"Checkpoint not found at {checkpoint_path}, initializing blank encoder.") + logger.warning( + f"Checkpoint not found at {checkpoint_path}, initializing blank encoder." + ) if node_scalar_in is None: raise ValueError("node_scalar_in required when checkpoint doesn't exist") @@ -432,7 +453,7 @@ def load_encoder_from_checkpoint( return encoder, args -@register_encoder('gvp') +@register_encoder("gvp") class GVPEncoder(BaseProteinEncoder): """ GVP encoder implementing the BaseProteinEncoder interface. @@ -468,9 +489,11 @@ def output_dims(self) -> tuple[int, int]: @property def encoder_type(self) -> str: """Return encoder type identifier.""" - return 'gvp' + return "gvp" - def forward(self, data: HeteroData) -> tuple[torch.Tensor, torch.Tensor, tuple | None]: + def forward( + self, data: HeteroData + ) -> tuple[torch.Tensor, torch.Tensor, tuple | None]: """ Encode protein data. @@ -507,12 +530,12 @@ def from_config(cls, config: dict, device: torch.device) -> GVPEncoder: Returns: Instantiated GVPEncoder """ - encoder_ckpt = config.get('encoder_ckpt') - node_scalar_in = config.get('node_scalar_in', 16) - hidden_s = config.get('hidden_s', 256) - hidden_v = config.get('hidden_v', 32) - freeze = config.get('freeze_encoder', False) - use_edge_update = config.get('use_edge_update', True) + encoder_ckpt = config.get("encoder_ckpt") + node_scalar_in = config.get("node_scalar_in", 16) + hidden_s = config.get("hidden_s", 256) + hidden_v = config.get("hidden_v", 32) + freeze = config.get("freeze_encoder", False) + use_edge_update = config.get("use_edge_update", True) if encoder_ckpt: encoder, _ = load_encoder_from_checkpoint( diff --git a/src/utils.py b/src/utils.py index 52886e4..51fa714 100644 --- a/src/utils.py +++ b/src/utils.py @@ -131,6 +131,7 @@ def parse_split_file(split_file: Path, base_pdb_dir: Path) -> list[dict]: ATOM37_FILL = 1e-5 + def rbf(r: Tensor, num_gaussians: int = NUM_RBF, cutoff: float = RBF_CUTOFF) -> Tensor: """ Compute radial basis function encoding of distances. @@ -215,6 +216,7 @@ def ot_coupling( return x0_star, x1_star + # eval metric functions @torch.no_grad() def recall_precision( @@ -257,6 +259,7 @@ def recall_precision( return recall, precision + @torch.no_grad() def compute_rmsd( pred: torch.Tensor | np.ndarray, @@ -296,6 +299,7 @@ def compute_rmsd( diff = pred[row_ind] - target[col_ind] return float(np.sqrt(np.mean(np.sum(diff**2, axis=1)))) + def compute_placement_metrics( pred: torch.Tensor | np.ndarray, true: torch.Tensor | np.ndarray, @@ -344,6 +348,7 @@ def compute_placement_metrics( return {"precision": precision, "recall": recall, "f1": f1, "auc_pr": auc_pr} + # viz functions def plot_3d_frame( ax, @@ -428,6 +433,7 @@ def plot_3d_frame( if zlim is not None: ax.set_zlim(zlim) + def create_trajectory_gif( trajectory: Sequence[np.ndarray], protein_pos: np.ndarray, @@ -507,6 +513,7 @@ def create_trajectory_gif( loop=0, ) + def save_protein_plot( pred_ca: torch.Tensor, true_ca: torch.Tensor, diff --git a/tests/conftest.py b/tests/conftest.py index 783d5cd..b1b2b42 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -14,10 +14,10 @@ def device(): return torch.device("cuda" if torch.cuda.is_available() else "cpu") + @pytest.fixture(scope="session") def pdb_base_dir(): - """Wrapper of constant PDB_BASE_DIR. - """ + """Wrapper of constant PDB_BASE_DIR.""" return PDB_BASE_DIR diff --git a/tests/test_dataset.py b/tests/test_dataset.py index 4bfc258..3cf2c2c 100644 --- a/tests/test_dataset.py +++ b/tests/test_dataset.py @@ -179,9 +179,9 @@ def test_perfect_match(self): # Create atom array with known coords atoms = bts.AtomArray(3) - atoms.coord = np.array([[0., 0., 0.], [1., 0., 0.], [2., 0., 0.]]) + atoms.coord = np.array([[0.0, 0.0, 0.0], [1.0, 0.0, 0.0], [2.0, 0.0, 0.0]]) - target_coords = np.array([[0., 0., 0.], [2., 0., 0.]]) + target_coords = np.array([[0.0, 0.0, 0.0], [2.0, 0.0, 0.0]]) matched = match_atoms_to_coords(atoms, target_coords, tolerance=0.01) @@ -194,9 +194,9 @@ def test_no_match_outside_tolerance(self): import biotite.structure as bts atoms = bts.AtomArray(2) - atoms.coord = np.array([[0., 0., 0.], [1., 0., 0.]]) + atoms.coord = np.array([[0.0, 0.0, 0.0], [1.0, 0.0, 0.0]]) - target_coords = np.array([[0.5, 0., 0.]]) # Not close to any atom + target_coords = np.array([[0.5, 0.0, 0.0]]) # Not close to any atom matched = match_atoms_to_coords(atoms, target_coords, tolerance=0.01) @@ -207,7 +207,7 @@ def test_empty_target_coords(self): import biotite.structure as bts atoms = bts.AtomArray(3) - atoms.coord = np.array([[0., 0., 0.], [1., 0., 0.], [2., 0., 0.]]) + atoms.coord = np.array([[0.0, 0.0, 0.0], [1.0, 0.0, 0.0], [2.0, 0.0, 0.0]]) target_coords = np.zeros((0, 3)) @@ -220,9 +220,9 @@ def test_tolerance_parameter(self): import biotite.structure as bts atoms = bts.AtomArray(1) - atoms.coord = np.array([[0., 0., 0.]]) + atoms.coord = np.array([[0.0, 0.0, 0.0]]) - target_coords = np.array([[0.05, 0., 0.]]) + target_coords = np.array([[0.05, 0.0, 0.0]]) # Should not match with tight tolerance matched_tight = match_atoms_to_coords(atoms, target_coords, tolerance=0.01) @@ -242,7 +242,9 @@ def test_same_com_passes(self): protein_coords = torch.randn(100, 3) water_coords = protein_coords[:10].clone() # Same center region - is_valid, reason = check_com_distance(protein_coords, water_coords, max_com_dist=25.0) + is_valid, reason = check_com_distance( + protein_coords, water_coords, max_com_dist=25.0 + ) assert is_valid is True assert reason == "" @@ -252,7 +254,9 @@ def test_close_com_passes(self): protein_coords = torch.zeros(100, 3) water_coords = torch.zeros(10, 3) + 5.0 # 5A offset in all dims ~8.7A distance - is_valid, reason = check_com_distance(protein_coords, water_coords, max_com_dist=25.0) + is_valid, reason = check_com_distance( + protein_coords, water_coords, max_com_dist=25.0 + ) assert is_valid is True @@ -261,7 +265,9 @@ def test_far_com_fails(self): protein_coords = torch.zeros(100, 3) water_coords = torch.zeros(10, 3) + 100.0 # 100A offset - is_valid, reason = check_com_distance(protein_coords, water_coords, max_com_dist=25.0) + is_valid, reason = check_com_distance( + protein_coords, water_coords, max_com_dist=25.0 + ) assert is_valid is False assert "CoM distance" in reason @@ -272,7 +278,9 @@ def test_empty_water_passes(self): protein_coords = torch.randn(100, 3) water_coords = torch.zeros(0, 3) - is_valid, reason = check_com_distance(protein_coords, water_coords, max_com_dist=25.0) + is_valid, reason = check_com_distance( + protein_coords, water_coords, max_com_dist=25.0 + ) assert is_valid is True @@ -282,11 +290,15 @@ def test_custom_threshold(self): water_coords = torch.zeros(10, 3) + 10.0 # ~17.3A distance # Tight threshold should fail - is_valid_tight, _ = check_com_distance(protein_coords, water_coords, max_com_dist=10.0) + is_valid_tight, _ = check_com_distance( + protein_coords, water_coords, max_com_dist=10.0 + ) assert is_valid_tight is False # Loose threshold should pass - is_valid_loose, _ = check_com_distance(protein_coords, water_coords, max_com_dist=50.0) + is_valid_loose, _ = check_com_distance( + protein_coords, water_coords, max_com_dist=50.0 + ) assert is_valid_loose is True @@ -359,7 +371,7 @@ def test_empty_water_passes(self): def test_custom_clash_distance(self): """Custom clash distance should be respected.""" protein_coords = torch.zeros(10, 3) - water_coords = torch.tensor([[1.5, 0., 0.]]) # 1.5A from origin + water_coords = torch.tensor([[1.5, 0.0, 0.0]]) # 1.5A from origin # 1A clash distance - not clashing is_valid_1a, _ = check_water_clashes( @@ -373,6 +385,7 @@ def test_custom_clash_distance(self): ) assert is_valid_2a is False + @pytest.mark.unit class TestCheckChainInteractions: """Tests for chain interaction quality filter.""" @@ -385,7 +398,9 @@ def test_single_chain_passes(self): atoms.chain_id = np.array(["A"] * 10) atoms.coord = np.random.randn(10, 3) - is_valid, reason, status = check_chain_interactions(atoms, interface_dist_threshold=4.0) + is_valid, reason, status = check_chain_interactions( + atoms, interface_dist_threshold=4.0 + ) assert is_valid is True assert status == "Single Chain" @@ -399,9 +414,13 @@ def test_interacting_chains_pass(self): # Place chains close together atoms.coord = np.zeros((20, 3)) atoms.coord[:10] = np.random.randn(10, 3) - atoms.coord[10:] = np.random.randn(10, 3) + np.array([2., 0., 0.]) # 2A offset + atoms.coord[10:] = np.random.randn(10, 3) + np.array( + [2.0, 0.0, 0.0] + ) # 2A offset - is_valid, reason, status = check_chain_interactions(atoms, interface_dist_threshold=4.0) + is_valid, reason, status = check_chain_interactions( + atoms, interface_dist_threshold=4.0 + ) assert is_valid is True assert status == "Interacting" @@ -415,9 +434,13 @@ def test_non_interacting_chains_fail(self): # Place chains far apart atoms.coord = np.zeros((20, 3)) atoms.coord[:10] = np.zeros((10, 3)) - atoms.coord[10:] = np.zeros((10, 3)) + np.array([100., 0., 0.]) # 100A offset + atoms.coord[10:] = np.zeros((10, 3)) + np.array( + [100.0, 0.0, 0.0] + ) # 100A offset - is_valid, reason, status = check_chain_interactions(atoms, interface_dist_threshold=4.0) + is_valid, reason, status = check_chain_interactions( + atoms, interface_dist_threshold=4.0 + ) assert is_valid is False assert "ASU copies" in reason or "not PPI" in reason @@ -431,10 +454,16 @@ def test_three_chains_any_pair_interacting(self): atoms.chain_id = np.array(["A"] * 10 + ["B"] * 10 + ["C"] * 10) atoms.coord = np.zeros((30, 3)) atoms.coord[:10] = np.zeros((10, 3)) - atoms.coord[10:20] = np.zeros((10, 3)) + np.array([2., 0., 0.]) # B close to A - atoms.coord[20:] = np.zeros((10, 3)) + np.array([100., 0., 0.]) # C far from all - - is_valid, reason, status = check_chain_interactions(atoms, interface_dist_threshold=4.0) + atoms.coord[10:20] = np.zeros((10, 3)) + np.array( + [2.0, 0.0, 0.0] + ) # B close to A + atoms.coord[20:] = np.zeros((10, 3)) + np.array( + [100.0, 0.0, 0.0] + ) # C far from all + + is_valid, reason, status = check_chain_interactions( + atoms, interface_dist_threshold=4.0 + ) # Should pass because A and B interact assert is_valid is True @@ -470,7 +499,9 @@ def test_chain_filter(self, pdb_6eey): if len(all_chains) > 1: first_chain = list(all_chains)[0] - protein_filtered, _ = parse_asu_with_biotite(pdb_6eey, chain_filter=[first_chain]) + protein_filtered, _ = parse_asu_with_biotite( + pdb_6eey, chain_filter=[first_chain] + ) assert set(protein_filtered.chain_id) == {first_chain} assert len(protein_filtered) < len(protein_all) @@ -510,14 +541,18 @@ def test_different_cutoffs(self, pdb_6eey): result_large = get_crystal_contacts_pymol(pdb_6eey, cutoff=8.0) # Larger cutoff should generally find more interface atoms - assert result_large["mate_coords"].shape[0] >= result_small["mate_coords"].shape[0] + assert ( + result_large["mate_coords"].shape[0] >= result_small["mate_coords"].shape[0] + ) @pytest.mark.integration class TestProteinWaterDataset: """Tests for the main dataset class.""" - def test_dataset_creation(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): + def test_dataset_creation( + self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir + ): """Dataset should be created successfully.""" dataset = ProteinWaterDataset( pdb_list_file=single_pdb_list_file, @@ -528,7 +563,9 @@ def test_dataset_creation(self, single_pdb_list_file, tmp_processed_dir, pdb_bas assert len(dataset) >= 1 - def test_getitem_returns_heterodata(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): + def test_getitem_returns_heterodata( + self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir + ): """__getitem__ should return HeteroData.""" from torch_geometric.data import HeteroData @@ -542,7 +579,9 @@ def test_getitem_returns_heterodata(self, single_pdb_list_file, tmp_processed_di data = dataset[0] assert isinstance(data, HeteroData) - def test_heterodata_has_required_fields(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): + def test_heterodata_has_required_fields( + self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir + ): """HeteroData should have required node types and fields.""" dataset = ProteinWaterDataset( pdb_list_file=single_pdb_list_file, @@ -554,18 +593,20 @@ def test_heterodata_has_required_fields(self, single_pdb_list_file, tmp_processe data = dataset[0] # Check protein nodes - assert hasattr(data['protein'], 'pos') - assert hasattr(data['protein'], 'x') - assert hasattr(data['protein'], 'residue_index') + assert hasattr(data["protein"], "pos") + assert hasattr(data["protein"], "x") + assert hasattr(data["protein"], "residue_index") # Check water nodes - assert hasattr(data['water'], 'pos') - assert hasattr(data['water'], 'x') + assert hasattr(data["water"], "pos") + assert hasattr(data["water"], "x") # Check edges - assert ('protein', 'pp', 'protein') in data.edge_types + assert ("protein", "pp", "protein") in data.edge_types - def test_protein_positions_centered(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): + def test_protein_positions_centered( + self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir + ): """ASU protein positions should be centered (mean ~ 0).""" dataset = ProteinWaterDataset( pdb_list_file=single_pdb_list_file, @@ -576,11 +617,13 @@ def test_protein_positions_centered(self, single_pdb_list_file, tmp_processed_di ) data = dataset[0] - protein_center = data['protein'].pos.mean(dim=0) + protein_center = data["protein"].pos.mean(dim=0) assert torch.allclose(protein_center, torch.zeros(3), atol=1e-3) - def test_duplicate_single_sample(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): + def test_duplicate_single_sample( + self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir + ): """duplicate_single_sample should multiply dataset length.""" dataset = ProteinWaterDataset( pdb_list_file=single_pdb_list_file, @@ -595,21 +638,25 @@ def test_duplicate_single_sample(self, single_pdb_list_file, tmp_processed_dir, # All items should be the same data_0 = dataset[0] data_5 = dataset[5] - assert torch.allclose(data_0['protein'].pos, data_5['protein'].pos) + assert torch.allclose(data_0["protein"].pos, data_5["protein"].pos) - def test_cached_file_created(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): + def test_cached_file_created( + self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir + ): """Preprocessing should create cached .pt file.""" _ = ProteinWaterDataset( pdb_list_file=single_pdb_list_file, processed_dir=str(tmp_processed_dir), base_pdb_dir=str(pdb_base_dir), preprocess=True, - ) # need to call this to trigger the processing + ) # need to call this to trigger the processing cache_file = tmp_processed_dir / "6eey_final_A.pt" assert cache_file.exists() - def test_no_reprocess_if_cached(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): + def test_no_reprocess_if_cached( + self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir + ): """Should not reprocess if cache exists.""" # First creation ProteinWaterDataset( @@ -707,7 +754,9 @@ def test_8dzt_passes_clash_check_at_5_percent(self, pdb_8dzt): class TestGetDataloader: """Tests for dataloader creation.""" - def test_dataloader_creation(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): + def test_dataloader_creation( + self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir + ): """Dataloader should be created successfully.""" loader = get_dataloader( pdb_list_file=single_pdb_list_file, @@ -721,7 +770,9 @@ def test_dataloader_creation(self, single_pdb_list_file, tmp_processed_dir, pdb_ assert loader is not None assert len(loader) >= 1 - def test_dataloader_iteration(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): + def test_dataloader_iteration( + self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir + ): """Should be able to iterate over dataloader.""" loader = get_dataloader( pdb_list_file=single_pdb_list_file, @@ -734,9 +785,11 @@ def test_dataloader_iteration(self, single_pdb_list_file, tmp_processed_dir, pdb batch = next(iter(loader)) assert batch is not None - assert hasattr(batch['protein'], 'pos') + assert hasattr(batch["protein"], "pos") - def test_dataloader_batching(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): + def test_dataloader_batching( + self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir + ): """Dataloader should support batching with duplicate_single_sample.""" loader = get_dataloader( pdb_list_file=single_pdb_list_file, @@ -753,11 +806,12 @@ def test_dataloader_batching(self, single_pdb_list_file, tmp_processed_dir, pdb_ batch = next(iter(loader)) # Batch should have batch indices - assert hasattr(batch['protein'], 'batch') + assert hasattr(batch["protein"], "batch") # ============== Tests for PDB list parsing ============== + @pytest.mark.unit class TestPdbListParsing: """Tests for PDB list file parsing.""" @@ -775,8 +829,8 @@ def test_chain_specific_format(self, tmp_path, pdb_base_dir): ) assert len(dataset.entries) == 1 - assert dataset.entries[0]['pdb_id'] == '6eey' - assert dataset.entries[0]['chain_id'] == 'A' + assert dataset.entries[0]["pdb_id"] == "6eey" + assert dataset.entries[0]["chain_id"] == "A" def test_whole_pdb_format(self, tmp_path, pdb_base_dir): """Should parse whole PDB format: pdb_id_final""" @@ -791,8 +845,8 @@ def test_whole_pdb_format(self, tmp_path, pdb_base_dir): ) assert len(dataset.entries) == 1 - assert dataset.entries[0]['pdb_id'] == '6eey' - assert dataset.entries[0]['chain_id'] is None + assert dataset.entries[0]["pdb_id"] == "6eey" + assert dataset.entries[0]["chain_id"] is None def test_multiple_entries(self, tmp_path, pdb_base_dir): """Should parse multiple entries.""" @@ -827,7 +881,9 @@ def test_empty_lines_ignored(self, tmp_path, pdb_base_dir): class TestDatasetEdgeCases: """Tests for edge cases in dataset handling.""" - def test_include_mates_flag(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): + def test_include_mates_flag( + self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir + ): """include_mates flag should affect protein node count.""" # Create dataset with mates dataset_with_mates = ProteinWaterDataset( @@ -851,7 +907,7 @@ def test_include_mates_flag(self, single_pdb_list_file, tmp_processed_dir, pdb_b data_without = dataset_no_mates[0] # With mates should have >= atoms - assert data_with['protein'].num_nodes >= data_without['protein'].num_nodes + assert data_with["protein"].num_nodes >= data_without["protein"].num_nodes def test_custom_cutoff(self, single_pdb_list_file, tmp_processed_dir, pdb_base_dir): """Custom cutoff should affect edge connectivity.""" @@ -875,14 +931,15 @@ def test_custom_cutoff(self, single_pdb_list_file, tmp_processed_dir, pdb_base_d data_large = dataset_large[0] # Larger cutoff should have more edges - n_edges_small = data_small['protein', 'pp', 'protein'].edge_index.shape[1] - n_edges_large = data_large['protein', 'pp', 'protein'].edge_index.shape[1] + n_edges_small = data_small["protein", "pp", "protein"].edge_index.shape[1] + n_edges_large = data_large["protein", "pp", "protein"].edge_index.shape[1] assert n_edges_large >= n_edges_small # ============== Tests for water quality filtering ============== + @pytest.mark.unit class TestLoadEdiaForPdb: """Tests for EDIA data loading from CSV files.""" @@ -981,13 +1038,15 @@ class TestFilterWatersByQuality: @pytest.fixture def mock_water_coords(self): """Create mock water coordinates for testing.""" - return np.array([ - [0., 0., 0.], # Water 0 - close to protein - [5., 0., 0.], # Water 1 - medium distance - [15., 0., 0.], # Water 2 - far from protein - [3., 0., 0.], # Water 3 - close - [20., 0., 0.], # Water 4 - very far - ]) + return np.array( + [ + [0.0, 0.0, 0.0], # Water 0 - close to protein + [5.0, 0.0, 0.0], # Water 1 - medium distance + [15.0, 0.0, 0.0], # Water 2 - far from protein + [3.0, 0.0, 0.0], # Water 3 - close + [20.0, 0.0, 0.0], # Water 4 - very far + ] + ) @pytest.fixture def mock_water_keys(self): @@ -999,7 +1058,9 @@ def mock_protein_coords(self): """Create mock protein coordinates for testing.""" return np.zeros((10, 3)) # Protein centered at origin - def test_distance_filtering(self, mock_water_coords, mock_water_keys, mock_protein_coords): + def test_distance_filtering( + self, mock_water_coords, mock_water_keys, mock_protein_coords + ): """Waters far from protein should be removed.""" keep_mask = filter_waters_by_quality( mock_water_coords, @@ -1037,11 +1098,11 @@ def test_edia_filtering(self, mock_water_coords, mock_water_keys): def test_bfactor_filtering(self, mock_water_coords, mock_water_keys): """Waters with high B-factor z-score should be removed.""" bfactor_lookup = { - ("A", 101): 1.0, # Pass - ("A", 102): 6.0, # Fail (> 5.0) - ("A", 103): 2.5, # Pass - ("B", 201): 7.0, # Fail - ("B", 202): 0.5, # Pass + ("A", 101): 1.0, # Pass + ("A", 102): 6.0, # Fail (> 5.0) + ("A", 103): 2.5, # Pass + ("B", 201): 7.0, # Fail + ("B", 202): 0.5, # Pass } keep_mask = filter_waters_by_quality( @@ -1055,7 +1116,9 @@ def test_bfactor_filtering(self, mock_water_coords, mock_water_keys): assert keep_mask.sum() == 3 - def test_combined_filters(self, mock_water_coords, mock_water_keys, mock_protein_coords): + def test_combined_filters( + self, mock_water_coords, mock_water_keys, mock_protein_coords + ): """Waters failing ANY criterion should be removed.""" edia_lookup = { ("A", 101): 0.85, # Pass EDIA @@ -1065,11 +1128,11 @@ def test_combined_filters(self, mock_water_coords, mock_water_keys, mock_protein ("B", 202): 0.85, # Pass EDIA, but will fail distance } bfactor_lookup = { - ("A", 101): 1.0, # Pass B-factor - ("A", 102): 6.0, # Fail B-factor - ("A", 103): 1.0, # Pass B-factor - ("B", 201): 1.0, # Pass B-factor - ("B", 202): 1.0, # Pass B-factor + ("A", 101): 1.0, # Pass B-factor + ("A", 102): 6.0, # Fail B-factor + ("A", 103): 1.0, # Pass B-factor + ("B", 201): 1.0, # Pass B-factor + ("B", 202): 1.0, # Pass B-factor } keep_mask = filter_waters_by_quality( @@ -1132,7 +1195,6 @@ def test_all_filters_disabled(self, mock_water_coords, mock_water_keys): assert keep_mask.sum() == 5 - @pytest.mark.integration class TestWaterFilteringIntegration: """Integration tests for water filtering with real PDB files.""" @@ -1156,10 +1218,9 @@ def test_filtering_with_real_pdb(self, pdb_6eey): bfactor_lookup, _ = compute_normalized_bfactors(pdb_6eey) # Build water keys - water_keys = list(zip( - water_atoms.chain_id.astype(str), - water_atoms.res_id.astype(int) - )) + water_keys = list( + zip(water_atoms.chain_id.astype(str), water_atoms.res_id.astype(int)) + ) # Apply filtering with distance and bfactor keep_mask = filter_waters_by_quality( @@ -1175,7 +1236,9 @@ def test_filtering_with_real_pdb(self, pdb_6eey): assert len(keep_mask) == len(water_keys) assert keep_mask.dtype == bool - def test_dataset_with_filtering_disabled(self, single_pdb_list_file, tmp_path, pdb_base_dir): + def test_dataset_with_filtering_disabled( + self, single_pdb_list_file, tmp_path, pdb_base_dir + ): """Dataset with filtering disabled should have same waters.""" # Create dataset with filtering disabled dataset = ProteinWaterDataset( @@ -1190,4 +1253,4 @@ def test_dataset_with_filtering_disabled(self, single_pdb_list_file, tmp_path, p data = dataset[0] # Should have water nodes - assert data['water'].num_nodes >= 0 + assert data["water"].num_nodes >= 0 diff --git a/tests/test_embedding_generation.py b/tests/test_embedding_generation.py index 4456df2..95871a7 100644 --- a/tests/test_embedding_generation.py +++ b/tests/test_embedding_generation.py @@ -24,15 +24,15 @@ def test_slae_embedding_loading(self): with tempfile.TemporaryDirectory() as tmpdir: # Create a dummy cache file with SLAE embeddings cache_data = { - 'protein_pos': torch.randn(50, 3), - 'protein_x': torch.randn(50, 16), - 'protein_res_idx': torch.arange(50), - 'protein_slae_embedding': torch.randn(50, 128), # SLAE embeddings - 'water_pos': torch.randn(10, 3), - 'water_x': torch.randn(10, 16), - 'mate_pos': torch.zeros(0, 3), - 'mate_x': torch.zeros(0, 16), - 'mate_res_idx': torch.zeros(0, dtype=torch.long), + "protein_pos": torch.randn(50, 3), + "protein_x": torch.randn(50, 16), + "protein_res_idx": torch.arange(50), + "protein_slae_embedding": torch.randn(50, 128), # SLAE embeddings + "water_pos": torch.randn(10, 3), + "water_x": torch.randn(10, 16), + "mate_pos": torch.zeros(0, 3), + "mate_x": torch.zeros(0, 16), + "mate_res_idx": torch.zeros(0, dtype=torch.long), } cache_path = Path(tmpdir) / "test_final_A.pt" @@ -53,9 +53,12 @@ def test_slae_embedding_loading(self): data = dataset[0] # Check SLAE embeddings are loaded - assert 'slae_embedding' in data['protein'], "SLAE embeddings should be loaded" - assert data['protein'].slae_embedding.shape == (50, 128), \ + assert "slae_embedding" in data["protein"], ( + "SLAE embeddings should be loaded" + ) + assert data["protein"].slae_embedding.shape == (50, 128), ( f"SLAE embedding shape mismatch: {data['protein'].slae_embedding.shape}" + ) def test_embedding_optional_backward_compat(self): """Test that dataset works without SLAE embeddings (backward compatibility).""" @@ -64,14 +67,14 @@ def test_embedding_optional_backward_compat(self): with tempfile.TemporaryDirectory() as tmpdir: # Create cache file WITHOUT SLAE embeddings cache_data = { - 'protein_pos': torch.randn(30, 3), - 'protein_x': torch.randn(30, 16), - 'protein_res_idx': torch.arange(30), - 'water_pos': torch.randn(5, 3), - 'water_x': torch.randn(5, 16), - 'mate_pos': torch.zeros(0, 3), - 'mate_x': torch.zeros(0, 16), - 'mate_res_idx': torch.zeros(0, dtype=torch.long), + "protein_pos": torch.randn(30, 3), + "protein_x": torch.randn(30, 16), + "protein_res_idx": torch.arange(30), + "water_pos": torch.randn(5, 3), + "water_x": torch.randn(5, 16), + "mate_pos": torch.zeros(0, 3), + "mate_x": torch.zeros(0, 16), + "mate_res_idx": torch.zeros(0, dtype=torch.long), } cache_path = Path(tmpdir) / "test_final_B.pt" @@ -89,8 +92,9 @@ def test_embedding_optional_backward_compat(self): data = dataset[0] # Should work without SLAE embeddings - assert 'slae_embedding' not in data['protein'], \ + assert "slae_embedding" not in data["protein"], ( "SLAE embeddings should not be present when not in cache" + ) def test_embedding_with_mates(self): """Test that embeddings are correctly concatenated with mate embeddings.""" @@ -99,16 +103,16 @@ def test_embedding_with_mates(self): with tempfile.TemporaryDirectory() as tmpdir: # Create cache file with protein and mate SLAE embeddings cache_data = { - 'protein_pos': torch.randn(30, 3), - 'protein_x': torch.randn(30, 16), - 'protein_res_idx': torch.arange(30), - 'protein_slae_embedding': torch.randn(30, 128), - 'water_pos': torch.randn(5, 3), - 'water_x': torch.randn(5, 16), - 'mate_pos': torch.randn(10, 3), - 'mate_x': torch.randn(10, 16), - 'mate_res_idx': torch.arange(10), - 'mate_slae_embedding': torch.randn(10, 128), + "protein_pos": torch.randn(30, 3), + "protein_x": torch.randn(30, 16), + "protein_res_idx": torch.arange(30), + "protein_slae_embedding": torch.randn(30, 128), + "water_pos": torch.randn(5, 3), + "water_x": torch.randn(5, 16), + "mate_pos": torch.randn(10, 3), + "mate_x": torch.randn(10, 16), + "mate_res_idx": torch.arange(10), + "mate_slae_embedding": torch.randn(10, 128), } cache_path = Path(tmpdir) / "test_final_C.pt" @@ -127,9 +131,10 @@ def test_embedding_with_mates(self): data = dataset[0] # Embeddings should be concatenated (protein + mate) - assert 'slae_embedding' in data['protein'] - assert data['protein'].slae_embedding.shape == (40, 128), \ + assert "slae_embedding" in data["protein"] + assert data["protein"].slae_embedding.shape == (40, 128), ( f"Expected (40, 128), got {data['protein'].slae_embedding.shape}" + ) class TestAlignSlaeToGeometry: diff --git a/tests/test_encoder.py b/tests/test_encoder.py index 4711fde..872f246 100644 --- a/tests/test_encoder.py +++ b/tests/test_encoder.py @@ -20,6 +20,7 @@ # ============== Fixtures ============== + @pytest.fixture def sample_homogeneous_data(): """Sample Data for testing ProteinGVPEncoder directly.""" @@ -44,19 +45,19 @@ def sample_hetero_data(device): data = HeteroData() # Protein nodes - data['protein'].x = torch.randn(num_protein, 16, device=device) - data['protein'].pos = torch.randn(num_protein, 3, device=device) - data['protein'].batch = torch.zeros(num_protein, dtype=torch.long, device=device) - data['protein'].num_nodes = num_protein + data["protein"].x = torch.randn(num_protein, 16, device=device) + data["protein"].pos = torch.randn(num_protein, 3, device=device) + data["protein"].batch = torch.zeros(num_protein, dtype=torch.long, device=device) + data["protein"].num_nodes = num_protein # Water nodes - data['water'].x = torch.randn(num_water, 16, device=device) - data['water'].pos = torch.randn(num_water, 3, device=device) - data['water'].batch = torch.zeros(num_water, dtype=torch.long, device=device) + data["water"].x = torch.randn(num_water, 16, device=device) + data["water"].pos = torch.randn(num_water, 3, device=device) + data["water"].batch = torch.zeros(num_water, dtype=torch.long, device=device) # Protein-protein edges - pp_edges = radius_graph(data['protein'].pos, r=8.0, loop=False) - data['protein', 'pp', 'protein'].edge_index = pp_edges + pp_edges = radius_graph(data["protein"].pos, r=8.0, loop=False) + data["protein", "pp", "protein"].edge_index = pp_edges return data @@ -65,13 +66,16 @@ def sample_hetero_data(device): def sample_hetero_data_with_slae(sample_hetero_data): """Sample HeteroData with SLAE embeddings.""" data = sample_hetero_data - num_protein = data['protein'].num_nodes - data['protein'].slae_embedding = torch.randn(num_protein, 128, device=data['protein'].pos.device) + num_protein = data["protein"].num_nodes + data["protein"].slae_embedding = torch.randn( + num_protein, 128, device=data["protein"].pos.device + ) return data # ============== Registry Tests ============== + class TestPackageLevelImport: """Tests that importing from src package triggers encoder registration.""" @@ -87,11 +91,13 @@ def test_build_encoder_works_with_package_import(self): # Should work without any explicit encoder module imports device = torch.device("cpu") - gvp_encoder = pkg_build_encoder({'encoder_type': 'gvp', 'node_scalar_in': 16}, device) - assert gvp_encoder.encoder_type == 'gvp' + gvp_encoder = pkg_build_encoder( + {"encoder_type": "gvp", "node_scalar_in": 16}, device + ) + assert gvp_encoder.encoder_type == "gvp" - slae_encoder = pkg_build_encoder({'encoder_type': 'slae'}, device) - assert slae_encoder.encoder_type == 'slae' + slae_encoder = pkg_build_encoder({"encoder_type": "slae"}, device) + assert slae_encoder.encoder_type == "slae" class TestEncoderRegistry: @@ -99,63 +105,64 @@ class TestEncoderRegistry: def test_gvp_registered(self): """GVP encoder should be registered.""" - cls = get_encoder_class('gvp') + cls = get_encoder_class("gvp") assert cls is GVPEncoder def test_slae_registered(self): """SLAE encoder should be registered.""" - cls = get_encoder_class('slae') + cls = get_encoder_class("slae") assert cls is CachedEmbeddingEncoder def test_esm_registered(self): """ESM encoder should be registered.""" - cls = get_encoder_class('esm') + cls = get_encoder_class("esm") assert cls is CachedEmbeddingEncoder def test_unknown_encoder_raises(self): """Unknown encoder type should raise KeyError.""" with pytest.raises(KeyError, match="Unknown encoder type"): - get_encoder_class('nonexistent_encoder') + get_encoder_class("nonexistent_encoder") def test_build_encoder_gvp(self, device): """build_encoder should construct GVP encoder from config.""" config = { - 'encoder_type': 'gvp', - 'node_scalar_in': 16, - 'hidden_s': 64, - 'hidden_v': 16, + "encoder_type": "gvp", + "node_scalar_in": 16, + "hidden_s": 64, + "hidden_v": 16, } encoder = build_encoder(config, device) assert isinstance(encoder, GVPEncoder) - assert encoder.encoder_type == 'gvp' + assert encoder.encoder_type == "gvp" assert encoder.output_dims == (64, 16) def test_build_encoder_slae(self, device): """build_encoder should construct SLAE encoder from config.""" config = { - 'encoder_type': 'slae', + "encoder_type": "slae", } encoder = build_encoder(config, device) assert isinstance(encoder, CachedEmbeddingEncoder) - assert encoder.encoder_type == 'slae' + assert encoder.encoder_type == "slae" # output_dims not available until forward pass def test_build_encoder_esm(self, device): """build_encoder should construct ESM encoder from config.""" config = { - 'encoder_type': 'esm', + "encoder_type": "esm", } encoder = build_encoder(config, device) assert isinstance(encoder, CachedEmbeddingEncoder) - assert encoder.encoder_type == 'esm' + assert encoder.encoder_type == "esm" # output_dims not available until forward pass # ============== Base Interface Tests ============== + class TestBaseEncoderInterface: """Tests for BaseProteinEncoder interface contract.""" @@ -178,25 +185,27 @@ def test_gvp_implements_interface(self, device, sample_hetero_data): # Check forward returns (s, V, pp_edge_attr) tuple s, V, pp_edge_attr = encoder(sample_hetero_data) - assert s.shape[0] == sample_hetero_data['protein'].num_nodes - assert V.shape[0] == sample_hetero_data['protein'].num_nodes + assert s.shape[0] == sample_hetero_data["protein"].num_nodes + assert V.shape[0] == sample_hetero_data["protein"].num_nodes assert V.shape[2] == 3 # GVP encoder should return edge features assert pp_edge_attr is not None or encoder.encoder.edge_update is None - def test_cached_embedding_implements_interface(self, device, sample_hetero_data_with_slae): + def test_cached_embedding_implements_interface( + self, device, sample_hetero_data_with_slae + ): """CachedEmbeddingEncoder should implement all required interface methods.""" encoder = CachedEmbeddingEncoder( - embedding_key='slae_embedding', encoder_type='slae' + embedding_key="slae_embedding", encoder_type="slae" ).to(device) assert isinstance(encoder.encoder_type, str) # Check forward returns (s, V, pp_edge_attr) tuple s, V, pp_edge_attr = encoder(sample_hetero_data_with_slae) - assert s.shape[0] == sample_hetero_data_with_slae['protein'].num_nodes + assert s.shape[0] == sample_hetero_data_with_slae["protein"].num_nodes assert s.shape[1] == 128 - assert V.shape == (sample_hetero_data_with_slae['protein'].num_nodes, 0, 3) + assert V.shape == (sample_hetero_data_with_slae["protein"].num_nodes, 0, 3) # Cached embedding encoder should return None for edge features assert pp_edge_attr is None @@ -207,8 +216,13 @@ def test_cached_embedding_implements_interface(self, device, sample_hetero_data_ def test_from_config_class_method(self, device): """Both encoders should have from_config class method.""" - gvp_config = {'encoder_type': 'gvp', 'node_scalar_in': 16, 'hidden_s': 64, 'hidden_v': 16} - slae_config = {'encoder_type': 'slae'} + gvp_config = { + "encoder_type": "gvp", + "node_scalar_in": 16, + "hidden_s": 64, + "hidden_v": 16, + } + slae_config = {"encoder_type": "slae"} gvp_encoder = GVPEncoder.from_config(gvp_config, device) slae_encoder = CachedEmbeddingEncoder.from_config(slae_config, device) @@ -220,6 +234,7 @@ def test_from_config_class_method(self, device): # ============== GVP Encoder Tests ============== + class TestProteinGVPEncoder: """Tests for the core ProteinGVPEncoder.""" @@ -241,10 +256,15 @@ def test_encoder_initialization(self, simple_encoder): assert simple_encoder is not None assert simple_encoder.n_layers == 2 - def test_encoder_forward_with_pooling(self, simple_encoder, sample_homogeneous_data): + def test_encoder_forward_with_pooling( + self, simple_encoder, sample_homogeneous_data + ): """Test forward pass with residue pooling.""" output, edge_attr = simple_encoder(sample_homogeneous_data) - assert output.shape == (sample_homogeneous_data.num_residues, simple_encoder.pooled_dim) + assert output.shape == ( + sample_homogeneous_data.num_residues, + simple_encoder.pooled_dim, + ) # Pooling mode returns None for edge features assert edge_attr is None @@ -312,7 +332,7 @@ def test_wrapper_encoder_type(self, device): encoder = GVPEncoder(encoder=base_encoder, freeze=False) - assert encoder.encoder_type == 'gvp' + assert encoder.encoder_type == "gvp" def test_wrapper_forward(self, device, sample_hetero_data): """Wrapper forward should return (s, V, edge_attr) from HeteroData.""" @@ -327,8 +347,8 @@ def test_wrapper_forward(self, device, sample_hetero_data): s, V, edge_attr = encoder(sample_hetero_data) - assert s.shape == (sample_hetero_data['protein'].num_nodes, 64) - assert V.shape == (sample_hetero_data['protein'].num_nodes, 16, 3) + assert s.shape == (sample_hetero_data["protein"].num_nodes, 64) + assert V.shape == (sample_hetero_data["protein"].num_nodes, 16, 3) # edge_attr should be a tuple (s_edge, V_edge) when edge_update is enabled assert edge_attr is not None s_edge, V_edge = edge_attr @@ -352,13 +372,14 @@ def test_wrapper_freeze(self, device): # ============== Cached Embedding Encoder Tests ============== + class TestCachedEmbeddingEncoder: """Tests for CachedEmbeddingEncoder (handles both SLAE and ESM).""" def test_output_dims_before_forward_raises(self, device): """output_dims should raise RuntimeError before forward pass.""" encoder = CachedEmbeddingEncoder( - embedding_key='slae_embedding', encoder_type='slae' + embedding_key="slae_embedding", encoder_type="slae" ).to(device) with pytest.raises(RuntimeError, match="dimension not yet known"): _ = encoder.output_dims @@ -366,7 +387,7 @@ def test_output_dims_before_forward_raises(self, device): def test_slae_output_dims_after_forward(self, device, sample_hetero_data_with_slae): """SLAE encoder should infer output_dims from data.""" encoder = CachedEmbeddingEncoder( - embedding_key='slae_embedding', encoder_type='slae' + embedding_key="slae_embedding", encoder_type="slae" ).to(device) encoder(sample_hetero_data_with_slae) assert encoder.output_dims == (128, 0) @@ -374,66 +395,70 @@ def test_slae_output_dims_after_forward(self, device, sample_hetero_data_with_sl def test_esm_output_dims_after_forward(self, device, sample_hetero_data): """ESM encoder should infer output_dims from data.""" encoder = CachedEmbeddingEncoder( - embedding_key='esm_embedding', encoder_type='esm' + embedding_key="esm_embedding", encoder_type="esm" ).to(device) - n_atoms = sample_hetero_data['protein'].num_nodes - sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 1536, device=device) + n_atoms = sample_hetero_data["protein"].num_nodes + sample_hetero_data["protein"].esm_embedding = torch.randn( + n_atoms, 1536, device=device + ) encoder(sample_hetero_data) assert encoder.output_dims == (1536, 0) def test_slae_encoder_type(self, device): """SLAE encoder should return 'slae' as encoder_type.""" encoder = CachedEmbeddingEncoder( - embedding_key='slae_embedding', encoder_type='slae' + embedding_key="slae_embedding", encoder_type="slae" ).to(device) - assert encoder.encoder_type == 'slae' + assert encoder.encoder_type == "slae" def test_esm_encoder_type(self, device): """ESM encoder should return 'esm' as encoder_type.""" encoder = CachedEmbeddingEncoder( - embedding_key='esm_embedding', encoder_type='esm' + embedding_key="esm_embedding", encoder_type="esm" ).to(device) - assert encoder.encoder_type == 'esm' + assert encoder.encoder_type == "esm" def test_slae_forward(self, device, sample_hetero_data_with_slae): """SLAE forward pass should return (s, V, None) tuple with raw embeddings.""" encoder = CachedEmbeddingEncoder( - embedding_key='slae_embedding', encoder_type='slae' + embedding_key="slae_embedding", encoder_type="slae" ).to(device) s, V, pp_edge_attr = encoder(sample_hetero_data_with_slae) - n_atoms = sample_hetero_data_with_slae['protein'].num_nodes + n_atoms = sample_hetero_data_with_slae["protein"].num_nodes assert s.shape == (n_atoms, 128) assert V.shape == (n_atoms, 0, 3) # Raw embeddings should be identical to input - assert torch.allclose(s, sample_hetero_data_with_slae['protein'].slae_embedding) + assert torch.allclose(s, sample_hetero_data_with_slae["protein"].slae_embedding) # Cached embedding encoder doesn't return edge features assert pp_edge_attr is None def test_esm_forward(self, device, sample_hetero_data): """ESM forward pass should return (s, V, None) tuple with raw embeddings.""" encoder = CachedEmbeddingEncoder( - embedding_key='esm_embedding', encoder_type='esm' + embedding_key="esm_embedding", encoder_type="esm" ).to(device) # Add mock ESM embeddings - n_atoms = sample_hetero_data['protein'].num_nodes - sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 1536, device=device) + n_atoms = sample_hetero_data["protein"].num_nodes + sample_hetero_data["protein"].esm_embedding = torch.randn( + n_atoms, 1536, device=device + ) s, V, pp_edge_attr = encoder(sample_hetero_data) assert s.shape == (n_atoms, 1536) assert V.shape == (n_atoms, 0, 3) # Raw embeddings should be identical to input - assert torch.allclose(s, sample_hetero_data['protein'].esm_embedding) + assert torch.allclose(s, sample_hetero_data["protein"].esm_embedding) # Cached embedding encoder doesn't return edge features assert pp_edge_attr is None def test_slae_missing_embeddings_error(self, device, sample_hetero_data): """Should raise KeyError when SLAE embeddings are missing.""" encoder = CachedEmbeddingEncoder( - embedding_key='slae_embedding', encoder_type='slae' + embedding_key="slae_embedding", encoder_type="slae" ).to(device) # sample_hetero_data does NOT have slae_embedding @@ -443,7 +468,7 @@ def test_slae_missing_embeddings_error(self, device, sample_hetero_data): def test_esm_missing_embeddings_error(self, device, sample_hetero_data): """Should raise KeyError when ESM embeddings are missing.""" encoder = CachedEmbeddingEncoder( - embedding_key='esm_embedding', encoder_type='esm' + embedding_key="esm_embedding", encoder_type="esm" ).to(device) # sample_hetero_data does NOT have esm_embedding @@ -453,7 +478,7 @@ def test_esm_missing_embeddings_error(self, device, sample_hetero_data): def test_encoder_no_nans(self, device, sample_hetero_data_with_slae): """Output should not contain NaNs or Infs.""" encoder = CachedEmbeddingEncoder( - embedding_key='slae_embedding', encoder_type='slae' + embedding_key="slae_embedding", encoder_type="slae" ).to(device) s, V, _ = encoder(sample_hetero_data_with_slae) @@ -464,47 +489,56 @@ def test_encoder_no_nans(self, device, sample_hetero_data_with_slae): def test_encoder_no_learnable_params(self, device): """Cached embedding encoder should have no learnable parameters.""" encoder = CachedEmbeddingEncoder( - embedding_key='slae_embedding', encoder_type='slae' + embedding_key="slae_embedding", encoder_type="slae" ).to(device) assert sum(p.numel() for p in encoder.parameters()) == 0 def test_slae_from_config(self, device, sample_hetero_data_with_slae): """Should construct SLAE from config and infer dim from data.""" - config = {'encoder_type': 'slae'} + config = {"encoder_type": "slae"} encoder = CachedEmbeddingEncoder.from_config(config, device) - assert encoder.encoder_type == 'slae' + assert encoder.encoder_type == "slae" encoder(sample_hetero_data_with_slae) assert encoder.output_dims == (128, 0) def test_esm_from_config(self, device, sample_hetero_data): """Should construct ESM from config and infer dim from data.""" - config = {'encoder_type': 'esm'} + config = {"encoder_type": "esm"} encoder = CachedEmbeddingEncoder.from_config(config, device) - assert encoder.encoder_type == 'esm' - n_atoms = sample_hetero_data['protein'].num_nodes - sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 2048, device=device) + assert encoder.encoder_type == "esm" + n_atoms = sample_hetero_data["protein"].num_nodes + sample_hetero_data["protein"].esm_embedding = torch.randn( + n_atoms, 2048, device=device + ) encoder(sample_hetero_data) assert encoder.output_dims == (2048, 0) def test_device_placement(self, device, sample_hetero_data): """Verify tensors are on the correct device.""" encoder = CachedEmbeddingEncoder( - embedding_key='esm_embedding', encoder_type='esm' + embedding_key="esm_embedding", encoder_type="esm" ).to(device) # Add mock ESM embeddings on correct device - n_atoms = sample_hetero_data['protein'].num_nodes - sample_hetero_data['protein'].esm_embedding = torch.randn(n_atoms, 1536, device=device) + n_atoms = sample_hetero_data["protein"].num_nodes + sample_hetero_data["protein"].esm_embedding = torch.randn( + n_atoms, 1536, device=device + ) s, V, _ = encoder(sample_hetero_data) # Compare device types (handles cuda vs cuda:0) - assert s.device.type == device.type, f"Expected device type {device.type}, got {s.device.type}" - assert V.device.type == device.type, f"Expected device type {device.type}, got {V.device.type}" + assert s.device.type == device.type, ( + f"Expected device type {device.type}, got {s.device.type}" + ) + assert V.device.type == device.type, ( + f"Expected device type {device.type}, got {V.device.type}" + ) # ============== Encoder Interoperability Tests ============== + class TestEncoderInteroperability: """Tests that both encoders work interchangeably with flow model.""" @@ -528,7 +562,7 @@ def test_both_encoders_work_with_flow(self, device, sample_hetero_data_with_slae # SLAE encoder via CachedEmbeddingEncoder # embedding_dim=128 matches the fixture's slae_embedding shape slae_encoder = CachedEmbeddingEncoder( - embedding_key='slae_embedding', encoder_type='slae', embedding_dim=128 + embedding_key="slae_embedding", encoder_type="slae", embedding_dim=128 ).to(device) # Create flow models with each encoder diff --git a/tests/test_flow.py b/tests/test_flow.py index ebcf962..c1daee8 100644 --- a/tests/test_flow.py +++ b/tests/test_flow.py @@ -23,22 +23,22 @@ def simple_hetero_data(device): """Minimal HeteroData with protein and water nodes.""" data = HeteroData() - + # Protein: 10 atoms - data['protein'].pos = torch.randn(10, 3, device=device) - data['protein'].x = torch.randn(10, 16, device=device) - data['protein'].batch = torch.zeros(10, dtype=torch.long, device=device) - + data["protein"].pos = torch.randn(10, 3, device=device) + data["protein"].x = torch.randn(10, 16, device=device) + data["protein"].batch = torch.zeros(10, dtype=torch.long, device=device) + # Water: 5 molecules - data['water'].pos = torch.randn(5, 3, device=device) - data['water'].x = torch.randn(5, 16, device=device) - data['water'].batch = torch.zeros(5, dtype=torch.long, device=device) - + data["water"].pos = torch.randn(5, 3, device=device) + data["water"].x = torch.randn(5, 16, device=device) + data["water"].batch = torch.zeros(5, dtype=torch.long, device=device) + # Protein-protein edges - data['protein', 'pp', 'protein'].edge_index = torch.tensor( + data["protein", "pp", "protein"].edge_index = torch.tensor( [[0, 1, 2, 3], [1, 2, 3, 4]], dtype=torch.long, device=device ) - + return data @@ -46,27 +46,25 @@ def simple_hetero_data(device): def batched_hetero_data(device): """HeteroData with 2 graphs batched.""" data = HeteroData() - + # Protein: 20 atoms (10 per graph) - data['protein'].pos = torch.randn(20, 3, device=device) - data['protein'].x = torch.randn(20, 16, device=device) - data['protein'].batch = torch.cat([ - torch.zeros(10, dtype=torch.long), - torch.ones(10, dtype=torch.long) - ]).to(device) - + data["protein"].pos = torch.randn(20, 3, device=device) + data["protein"].x = torch.randn(20, 16, device=device) + data["protein"].batch = torch.cat( + [torch.zeros(10, dtype=torch.long), torch.ones(10, dtype=torch.long)] + ).to(device) + # Water: 8 molecules (4 per graph) - data['water'].pos = torch.randn(8, 3, device=device) - data['water'].x = torch.randn(8, 16, device=device) - data['water'].batch = torch.cat([ - torch.zeros(4, dtype=torch.long), - torch.ones(4, dtype=torch.long) - ]).to(device) - - data['protein', 'pp', 'protein'].edge_index = torch.tensor( + data["water"].pos = torch.randn(8, 3, device=device) + data["water"].x = torch.randn(8, 16, device=device) + data["water"].batch = torch.cat( + [torch.zeros(4, dtype=torch.long), torch.ones(4, dtype=torch.long)] + ).to(device) + + data["protein", "pp", "protein"].edge_index = torch.tensor( [[0, 1, 10, 11], [1, 2, 11, 12]], dtype=torch.long, device=device ) - + return data @@ -75,12 +73,12 @@ def mock_encoder(device): """Mock BaseProteinEncoder.""" encoder = Mock() encoder.output_dims = (256, 32) # Required by FlowWaterGVP - encoder.encoder_type = 'mock' + encoder.encoder_type = "mock" encoder.parameters = Mock(return_value=iter([torch.nn.Parameter(torch.randn(1))])) encoder.eval = Mock() def mock_forward(data): - n = data['protein'].pos.size(0) + n = data["protein"].pos.size(0) s = torch.randn(n, 256, device=device) v = torch.randn(n, 32, 3, device=device) return s, v @@ -92,69 +90,69 @@ def mock_forward(data): @pytest.mark.unit class TestBuildKnnEdges: - def test_basic_knn(self, device): - src = torch.tensor([[0., 0., 0.], [1., 0., 0.], [2., 0., 0.]], device=device) - dst = torch.tensor([[0.5, 0., 0.], [1.5, 0., 0.]], device=device) - + src = torch.tensor( + [[0.0, 0.0, 0.0], [1.0, 0.0, 0.0], [2.0, 0.0, 0.0]], device=device + ) + dst = torch.tensor([[0.5, 0.0, 0.0], [1.5, 0.0, 0.0]], device=device) + edges = build_knn_edges(src, dst, k=2) - + assert edges.shape[0] == 2 assert edges.shape[1] > 0 assert edges.dtype == torch.long - + def test_empty_src(self, device): src = torch.empty(0, 3, device=device) dst = torch.randn(5, 3, device=device) - + edges = build_knn_edges(src, dst, k=3) - + assert edges.shape == (2, 0) - + def test_empty_dst(self, device): src = torch.randn(5, 3, device=device) dst = torch.empty(0, 3, device=device) - + edges = build_knn_edges(src, dst, k=3) - + assert edges.shape == (2, 0) - + def test_self_edges_removed(self, device): pos = torch.randn(10, 3, device=device) - + edges = build_knn_edges(pos, pos, k=5) - + # No self-loops assert (edges[0] != edges[1]).all() - + def test_with_batch(self, device): src = torch.randn(10, 3, device=device) dst = torch.randn(8, 3, device=device) batch_src = torch.cat([torch.zeros(5), torch.ones(5)]).long().to(device) batch_dst = torch.cat([torch.zeros(4), torch.ones(4)]).long().to(device) - + edges = build_knn_edges(src, dst, k=3, batch_src=batch_src, batch_dst=batch_dst) - + assert edges.shape[0] == 2 assert edges.shape[1] > 0 @pytest.mark.unit class TestMakeEncoderData: - def test_basic_output(self, simple_hetero_data): enc_data = make_gvp_encoder_data(simple_hetero_data) assert isinstance(enc_data, Data) - assert hasattr(enc_data, 'x') - assert hasattr(enc_data, 'pos') - assert hasattr(enc_data, 'edge_index') + assert hasattr(enc_data, "x") + assert hasattr(enc_data, "pos") + assert hasattr(enc_data, "edge_index") def test_shapes(self, simple_hetero_data): enc_data = make_gvp_encoder_data(simple_hetero_data) - n_nodes = simple_hetero_data['protein'].pos.size(0) - n_edges = simple_hetero_data['protein', 'pp', 'protein'].edge_index.size(1) + n_nodes = simple_hetero_data["protein"].pos.size(0) + n_edges = simple_hetero_data["protein", "pp", "protein"].edge_index.size(1) assert enc_data.x.shape[0] == n_nodes assert enc_data.pos.shape == (n_nodes, 3) @@ -163,13 +161,13 @@ def test_shapes(self, simple_hetero_data): def test_batch_preserved(self, batched_hetero_data): enc_data = make_gvp_encoder_data(batched_hetero_data) - assert hasattr(enc_data, 'batch') - assert enc_data.batch.shape[0] == batched_hetero_data['protein'].pos.size(0) + assert hasattr(enc_data, "batch") + assert enc_data.batch.shape[0] == batched_hetero_data["protein"].pos.size(0) def test_no_edges(self, device): data = HeteroData() - data['protein'].pos = torch.randn(10, 3, device=device) - data['protein'].x = torch.randn(10, 16, device=device) + data["protein"].pos = torch.randn(10, 3, device=device) + data["protein"].x = torch.randn(10, 16, device=device) # No edges defined enc_data = make_gvp_encoder_data(data) @@ -179,18 +177,17 @@ def test_no_edges(self, device): @pytest.mark.unit class TestProteinWaterUpdate: - def test_init(self): updater = ProteinWaterUpdate( hidden_dims=(128, 16), rbf_dim=16, layers=2, ) - + assert len(updater.blocks) == 2 - assert ('protein', 'pw', 'water') in updater.etypes - assert ('water', 'ww', 'water') in updater.etypes - + assert ("protein", "pw", "water") in updater.etypes + assert ("water", "ww", "water") in updater.etypes + def test_init_always_includes_all_edge_types(self): updater = ProteinWaterUpdate( hidden_dims=(128, 16), @@ -198,66 +195,70 @@ def test_init_always_includes_all_edge_types(self): layers=2, ) - assert ('protein', 'pp', 'protein') in updater.etypes - assert ('water', 'wp', 'protein') in updater.etypes - + assert ("protein", "pp", "protein") in updater.etypes + assert ("water", "wp", "protein") in updater.etypes + def test_build_edges(self, simple_hetero_data): updater = ProteinWaterUpdate(hidden_dims=(128, 16), layers=1) - + edge_dict = updater.build_edges(simple_hetero_data, k_pw=4, k_ww=3) - - assert ('protein', 'pw', 'water') in edge_dict - assert ('water', 'ww', 'water') in edge_dict - assert edge_dict[('protein', 'pw', 'water')].shape[0] == 2 - + + assert ("protein", "pw", "water") in edge_dict + assert ("water", "ww", "water") in edge_dict + assert edge_dict[("protein", "pw", "water")].shape[0] == 2 + def test_build_edges_empty_water(self, device): data = HeteroData() - data['protein'].pos = torch.randn(10, 3, device=device) - data['protein'].x = torch.randn(10, 16, device=device) - data['water'].pos = torch.empty(0, 3, device=device) - data['water'].x = torch.empty(0, 16, device=device) - + data["protein"].pos = torch.randn(10, 3, device=device) + data["protein"].x = torch.randn(10, 16, device=device) + data["water"].pos = torch.empty(0, 3, device=device) + data["water"].x = torch.empty(0, 16, device=device) + updater = ProteinWaterUpdate(hidden_dims=(128, 16), layers=1) edge_dict = updater.build_edges(data) - - assert edge_dict[('protein', 'pw', 'water')].shape == (2, 0) - assert edge_dict[('water', 'ww', 'water')].shape == (2, 0) - + + assert edge_dict[("protein", "pw", "water")].shape == (2, 0) + assert edge_dict[("water", "ww", "water")].shape == (2, 0) + def test_forward_shapes(self, simple_hetero_data, device): s_h, v_h = 128, 16 updater = ProteinWaterUpdate(hidden_dims=(s_h, v_h), layers=1).to(device) - - n_p = simple_hetero_data['protein'].pos.size(0) - n_w = simple_hetero_data['water'].pos.size(0) - + + n_p = simple_hetero_data["protein"].pos.size(0) + n_w = simple_hetero_data["water"].pos.size(0) + x_dict = { - 'protein': (torch.randn(n_p, s_h, device=device), - torch.randn(n_p, v_h, 3, device=device)), - 'water': (torch.randn(n_w, s_h, device=device), - torch.randn(n_w, v_h, 3, device=device)), + "protein": ( + torch.randn(n_p, s_h, device=device), + torch.randn(n_p, v_h, 3, device=device), + ), + "water": ( + torch.randn(n_w, s_h, device=device), + torch.randn(n_w, v_h, 3, device=device), + ), } - + out = updater(x_dict, simple_hetero_data) - - assert out['water'][0].shape == (n_w, s_h) - assert out['water'][1].shape == (n_w, v_h, 3) + + assert out["water"][0].shape == (n_w, s_h) + assert out["water"][1].shape == (n_w, v_h, 3) # ============== Tests for FlowWaterGVP ============== + @pytest.mark.unit class TestFlowWaterGVP: - def test_init(self, mock_encoder, device): model = FlowWaterGVP( encoder=mock_encoder, hidden_dims=(128, 16), layers=2, ).to(device) - + assert model.hidden_dims == (128, 16) assert model.layers == 2 - + def test_forward_output_shape(self, simple_hetero_data, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, @@ -272,13 +273,13 @@ def test_forward_output_shape(self, simple_hetero_data, device): hidden_dims=(64, 8), layers=1, ).to(device) - + t = torch.tensor([0.5], device=device) v_pred = model(simple_hetero_data, t) - - n_water = simple_hetero_data['water'].num_nodes + + n_water = simple_hetero_data["water"].num_nodes assert v_pred.shape == (n_water, 3) - + def test_forward_no_water(self, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, @@ -293,21 +294,21 @@ def test_forward_no_water(self, device): hidden_dims=(64, 8), layers=1, ).to(device) - + data = HeteroData() - data['protein'].pos = torch.randn(10, 3, device=device) - data['protein'].x = torch.randn(10, 16, device=device) - data['protein'].batch = torch.zeros(10, dtype=torch.long, device=device) - data['protein', 'pp', 'protein'].edge_index = torch.tensor( + data["protein"].pos = torch.randn(10, 3, device=device) + data["protein"].x = torch.randn(10, 16, device=device) + data["protein"].batch = torch.zeros(10, dtype=torch.long, device=device) + data["protein", "pp", "protein"].edge_index = torch.tensor( [[0, 1], [1, 2]], dtype=torch.long, device=device ) # No water nodes - + t = torch.tensor([0.5], device=device) v_pred = model(data, t) - + assert v_pred.shape == (0, 3) - + def test_self_conditioning(self, simple_hetero_data, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, @@ -322,21 +323,21 @@ def test_self_conditioning(self, simple_hetero_data, device): hidden_dims=(64, 8), layers=1, ).to(device) - - n_water = simple_hetero_data['water'].num_nodes - sc = {'x1_pred': torch.randn(n_water, 3, device=device)} + + n_water = simple_hetero_data["water"].num_nodes + sc = {"x1_pred": torch.randn(n_water, 3, device=device)} t = torch.tensor([0.5], device=device) - + v_pred = model(simple_hetero_data, t, sc=sc) - + assert v_pred.shape == (n_water, 3) # ============== Tests for FlowMatcher ============== + @pytest.mark.unit class TestFlowMatcher: - @pytest.fixture def flow_matcher(self, device): base_encoder = ProteinGVPEncoder( @@ -354,43 +355,45 @@ def flow_matcher(self, device): ).to(device) return FlowMatcher(model, p_self_cond=0.5) - + def test_compute_sigma(self, simple_hetero_data): sigma = FlowMatcher.compute_sigma(simple_hetero_data) - + assert isinstance(sigma, float) assert sigma > 0 - + def test_training_step(self, flow_matcher, simple_hetero_data, device): optimizer = torch.optim.Adam(flow_matcher.model.parameters(), lr=1e-4) - + result = flow_matcher.training_step( simple_hetero_data, optimizer, use_self_conditioning=False ) - - assert 'loss' in result - assert 'rmsd' in result - assert 'sigma' in result - assert result['loss'] >= 0 - - def test_training_step_with_self_cond(self, flow_matcher, simple_hetero_data, device): + + assert "loss" in result + assert "rmsd" in result + assert "sigma" in result + assert result["loss"] >= 0 + + def test_training_step_with_self_cond( + self, flow_matcher, simple_hetero_data, device + ): optimizer = torch.optim.Adam(flow_matcher.model.parameters(), lr=1e-4) - + # Force self-conditioning flow_matcher.p_self_cond = 1.0 result = flow_matcher.training_step( simple_hetero_data, optimizer, use_self_conditioning=True ) - - assert 'loss' in result - + + assert "loss" in result + def test_validation_step(self, flow_matcher, simple_hetero_data): result = flow_matcher.validation_step(simple_hetero_data) - - assert 'loss' in result - assert 'rmsd' in result - assert result['loss'] >= 0 - + + assert "loss" in result + assert "rmsd" in result + assert result["loss"] >= 0 + @pytest.mark.slow def test_euler_integrate(self, flow_matcher, simple_hetero_data, device): results = flow_matcher.euler_integrate( @@ -399,47 +402,50 @@ def test_euler_integrate(self, flow_matcher, simple_hetero_data, device): # euler_integrate returns List[np.ndarray], one per input graph water_pred = results[0] - n_water = simple_hetero_data['water'].num_nodes + n_water = simple_hetero_data["water"].num_nodes assert water_pred.shape == (n_water, 3) assert isinstance(water_pred, np.ndarray) @pytest.mark.slow def test_rk4_integrate(self, flow_matcher, simple_hetero_data, device): results = flow_matcher.rk4_integrate( - simple_hetero_data, num_steps=5, use_sc=False, - device=str(device), return_trajectory=True + simple_hetero_data, + num_steps=5, + use_sc=False, + device=str(device), + return_trajectory=True, ) # rk4_integrate returns List[Dict], one per input graph result = results[0] - assert 'water_pred' in result - assert 'water_true' in result - assert 'protein_pos' in result - assert 'trajectory' in result - assert len(result['trajectory']) == 5 - + assert "water_pred" in result + assert "water_true" in result + assert "protein_pos" in result + assert "trajectory" in result + assert len(result["trajectory"]) == 5 + def test_sample_euler(self, flow_matcher, simple_hetero_data, device): water_pred = flow_matcher.sample( simple_hetero_data, num_steps=3, method="euler", device=str(device) ) - - n_water = simple_hetero_data['water'].num_nodes + + n_water = simple_hetero_data["water"].num_nodes assert water_pred.shape == (n_water, 3) - + def test_sample_rk4(self, flow_matcher, simple_hetero_data, device): water_pred = flow_matcher.sample( simple_hetero_data, num_steps=3, method="rk4", device=str(device) ) - - n_water = simple_hetero_data['water'].num_nodes + + n_water = simple_hetero_data["water"].num_nodes assert water_pred.shape == (n_water, 3) # ============== Tests for distortion ============== + @pytest.mark.unit class TestDistortion: - def test_distortion_enabled(self, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, @@ -460,7 +466,7 @@ def test_distortion_enabled(self, device): use_distortion=True, p_distort=1.0, # Always apply t_distort=0.0, # Apply at all times - sigma_distort=0.5 + sigma_distort=0.5, ) assert fm.use_distortion is True @@ -469,9 +475,9 @@ def test_distortion_enabled(self, device): # ============== Edge case tests ============== + @pytest.mark.unit class TestEdgeCases: - def test_single_water_molecule(self, device): base_encoder = ProteinGVPEncoder( node_scalar_in=16, @@ -486,23 +492,23 @@ def test_single_water_molecule(self, device): hidden_dims=(64, 8), layers=1, ).to(device) - + data = HeteroData() - data['protein'].pos = torch.randn(10, 3, device=device) - data['protein'].x = torch.randn(10, 16, device=device) - data['protein'].batch = torch.zeros(10, dtype=torch.long, device=device) - data['water'].pos = torch.randn(1, 3, device=device) # Single water - data['water'].x = torch.randn(1, 16, device=device) - data['water'].batch = torch.zeros(1, dtype=torch.long, device=device) - data['protein', 'pp', 'protein'].edge_index = torch.tensor( + data["protein"].pos = torch.randn(10, 3, device=device) + data["protein"].x = torch.randn(10, 16, device=device) + data["protein"].batch = torch.zeros(10, dtype=torch.long, device=device) + data["water"].pos = torch.randn(1, 3, device=device) # Single water + data["water"].x = torch.randn(1, 16, device=device) + data["water"].batch = torch.zeros(1, dtype=torch.long, device=device) + data["protein", "pp", "protein"].edge_index = torch.tensor( [[0, 1], [1, 2]], dtype=torch.long, device=device ) - + t = torch.tensor([0.5], device=device) v_pred = model(data, t) - + assert v_pred.shape == (1, 3) - + def test_frozen_gvp_encoder(self, device): """Freezing is handled by the encoder itself, not FlowWaterGVP.""" base_encoder = ProteinGVPEncoder( @@ -526,6 +532,7 @@ def test_frozen_gvp_encoder(self, device): # ============== Tests for edge connectivity ============== + @pytest.mark.unit class TestWaterEdgeConnectivity: """Tests to ensure all waters have edges (both protein-water and water-water).""" @@ -535,74 +542,80 @@ def test_all_waters_have_protein_edges(self, simple_hetero_data): updater = ProteinWaterUpdate(hidden_dims=(128, 16), layers=1) edge_dict = updater.build_edges(simple_hetero_data, k_pw=4, k_ww=3) - pw_edges = edge_dict[('protein', 'pw', 'water')] + pw_edges = edge_dict[("protein", "pw", "water")] - n_water = simple_hetero_data['water'].num_nodes + n_water = simple_hetero_data["water"].num_nodes # Check that all water nodes appear in the protein-water edges water_nodes_with_edges = torch.unique(pw_edges[1]) - assert len(water_nodes_with_edges) == n_water, \ + assert len(water_nodes_with_edges) == n_water, ( f"Only {len(water_nodes_with_edges)}/{n_water} waters have protein edges" + ) def test_all_waters_have_water_edges(self, simple_hetero_data): """Ensure every water has at least one water-water edge (if multiple waters exist).""" updater = ProteinWaterUpdate(hidden_dims=(128, 16), layers=1) edge_dict = updater.build_edges(simple_hetero_data, k_pw=4, k_ww=3) - ww_edges = edge_dict[('water', 'ww', 'water')] + ww_edges = edge_dict[("water", "ww", "water")] - n_water = simple_hetero_data['water'].num_nodes + n_water = simple_hetero_data["water"].num_nodes if n_water > 1: # Check that all water nodes appear in the water-water edges water_nodes_with_edges = torch.unique(ww_edges[0]) - assert len(water_nodes_with_edges) == n_water, \ + assert len(water_nodes_with_edges) == n_water, ( f"Only {len(water_nodes_with_edges)}/{n_water} waters have water-water edges" + ) def test_batched_waters_have_edges(self, batched_hetero_data): """Ensure all waters in a batched graph have edges.""" updater = ProteinWaterUpdate(hidden_dims=(128, 16), layers=1) edge_dict = updater.build_edges(batched_hetero_data, k_pw=4, k_ww=3) - pw_edges = edge_dict[('protein', 'pw', 'water')] - ww_edges = edge_dict[('water', 'ww', 'water')] + pw_edges = edge_dict[("protein", "pw", "water")] + ww_edges = edge_dict[("water", "ww", "water")] - n_water = batched_hetero_data['water'].num_nodes + n_water = batched_hetero_data["water"].num_nodes # Check protein-water edges water_nodes_with_pw_edges = torch.unique(pw_edges[1]) - assert len(water_nodes_with_pw_edges) == n_water, \ + assert len(water_nodes_with_pw_edges) == n_water, ( f"Only {len(water_nodes_with_pw_edges)}/{n_water} waters have protein edges in batched data" + ) # Check water-water edges if n_water > 1: water_nodes_with_ww_edges = torch.unique(ww_edges[0]) - assert len(water_nodes_with_ww_edges) == n_water, \ + assert len(water_nodes_with_ww_edges) == n_water, ( f"Only {len(water_nodes_with_ww_edges)}/{n_water} waters have water-water edges in batched data" + ) def test_single_water_has_protein_edges_no_water_edges(self, device): """A single water should have protein edges but no water-water edges.""" data = HeteroData() - data['protein'].pos = torch.randn(10, 3, device=device) - data['protein'].x = torch.randn(10, 16, device=device) - data['protein'].batch = torch.zeros(10, dtype=torch.long, device=device) - data['water'].pos = torch.randn(1, 3, device=device) # Single water - data['water'].x = torch.randn(1, 16, device=device) - data['water'].batch = torch.zeros(1, dtype=torch.long, device=device) - data['protein', 'pp', 'protein'].edge_index = torch.tensor( + data["protein"].pos = torch.randn(10, 3, device=device) + data["protein"].x = torch.randn(10, 16, device=device) + data["protein"].batch = torch.zeros(10, dtype=torch.long, device=device) + data["water"].pos = torch.randn(1, 3, device=device) # Single water + data["water"].x = torch.randn(1, 16, device=device) + data["water"].batch = torch.zeros(1, dtype=torch.long, device=device) + data["protein", "pp", "protein"].edge_index = torch.tensor( [[0, 1], [1, 2]], dtype=torch.long, device=device ) updater = ProteinWaterUpdate(hidden_dims=(128, 16), layers=1) edge_dict = updater.build_edges(data, k_pw=4, k_ww=3) - pw_edges = edge_dict[('protein', 'pw', 'water')] - ww_edges = edge_dict[('water', 'ww', 'water')] + pw_edges = edge_dict[("protein", "pw", "water")] + ww_edges = edge_dict[("water", "ww", "water")] # Single water should have protein edges - assert pw_edges.shape[1] > 0, "Single water should have at least one protein edge" + assert pw_edges.shape[1] > 0, ( + "Single water should have at least one protein edge" + ) water_nodes_with_edges = torch.unique(pw_edges[1]) assert len(water_nodes_with_edges) == 1, "Single water must have protein edges" # Single water should have no water-water edges (since k_ww excludes self-loops) - assert ww_edges.shape[1] == 0, "Single water should have no water-water edges" \ No newline at end of file + assert ww_edges.shape[1] == 0, "Single water should have no water-water edges" diff --git a/tests/test_forward.py b/tests/test_forward.py index ca96ae3..ed476e2 100644 --- a/tests/test_forward.py +++ b/tests/test_forward.py @@ -129,7 +129,9 @@ def make_batched_hetero( return data -def assert_edge_index_in_range(edge_index: torch.Tensor, n_src: int, n_dst: int, name: str): +def assert_edge_index_in_range( + edge_index: torch.Tensor, n_src: int, n_dst: int, name: str +): if edge_index.numel() == 0: return smax = int(edge_index[0].max().item()) @@ -137,12 +139,15 @@ def assert_edge_index_in_range(edge_index: torch.Tensor, n_src: int, n_dst: int, assert smax < n_src, f"{name}: src index out of range (max={smax}, n_src={n_src})" assert dmax < n_dst, f"{name}: dst index out of range (max={dmax}, n_dst={n_dst})" + def test_rbf_zero_is_finite(device): from src.utils import rbf + r = torch.zeros(128, device=device) out = rbf(r, num_gaussians=16, cutoff=8.0) assert torch.isfinite(out).all() + @pytest.mark.slow def test_forward_pass_no_nan_with_module_hooks(device): torch.manual_seed(0) @@ -177,24 +182,38 @@ def test_forward_pass_no_nan_with_module_hooks(device): encoder=encoder, hidden_dims=(64, 8), layers=2, - k_pw=8, # keep <= n_water_per - k_ww=8, # keep <= n_water_per + k_pw=8, # keep <= n_water_per + k_ww=8, # keep <= n_water_per ).to(device) # Quick pre-check: protein encoder input features created from pp edges enc_data = make_gvp_encoder_data(data) - assert_edge_index_in_range(enc_data.edge_index, enc_data.x.size(0), enc_data.x.size(0), "pp edge_index") + assert_edge_index_in_range( + enc_data.edge_index, enc_data.x.size(0), enc_data.x.size(0), "pp edge_index" + ) # Also validate knn edges are sane (catches orientation / k issues) edge_dict = model.updater.build_edges(data, k_pw=model.k_pw, k_ww=model.k_ww) - assert_edge_index_in_range(edge_dict[("protein", "pw", "water")], data["protein"].pos.size(0), data["water"].pos.size(0), "pw edge_index") - assert_edge_index_in_range(edge_dict[("water", "ww", "water")], data["water"].pos.size(0), data["water"].pos.size(0), "ww edge_index") + assert_edge_index_in_range( + edge_dict[("protein", "pw", "water")], + data["protein"].pos.size(0), + data["water"].pos.size(0), + "pw edge_index", + ) + assert_edge_index_in_range( + edge_dict[("water", "ww", "water")], + data["water"].pos.size(0), + data["water"].pos.size(0), + "ww edge_index", + ) t = torch.linspace(0.05, 0.95, steps=n_graphs, device=device) with FiniteHookManager() as hm: # Encoder internals (access via .encoder.encoder for wrapped GVPEncoder) - hm.watch(model.encoder.encoder.input_scalar_encoder, "encoder.input_scalar_encoder") + hm.watch( + model.encoder.encoder.input_scalar_encoder, "encoder.input_scalar_encoder" + ) hm.watch(model.encoder.encoder.input_gvp, "encoder.input_gvp") for i, layer in enumerate(model.encoder.encoder.layers): hm.watch(layer, f"encoder.layers[{i}]") @@ -249,7 +268,7 @@ def test_training_step_no_nan_tripwire(device): fm = FlowMatcher( model=model, - p_self_cond=0.0, # simpler/cleaner for tripwire + p_self_cond=0.0, # simpler/cleaner for tripwire use_distortion=False, loss_eps=1e-3, ) @@ -263,7 +282,9 @@ def test_training_step_no_nan_tripwire(device): hm.watch(model.vfield_head, "vfield_head") for step in range(5): - out = fm.training_step(data, opt, grad_clip=1.0, use_self_conditioning=False) + out = fm.training_step( + data, opt, grad_clip=1.0, use_self_conditioning=False + ) loss = out["loss"] assert isinstance(loss, float) assert loss == loss, "loss is NaN" @@ -322,6 +343,7 @@ def test_forward_with_duplicate_protein_coords_catches_nan(device): assert torch.isfinite(v_pred).all(), "v_pred has NaN/Inf under duplicate coords" + @pytest.mark.slow def test_forward_with_duplicate_protein_coords_localizes_nan(device): torch.manual_seed(2) @@ -333,7 +355,7 @@ def test_forward_with_duplicate_protein_coords_localizes_nan(device): n_graphs=1, n_protein_per=32, n_water_per=16, - duplicate_protein_coords=True, # force a zero-distance pair + duplicate_protein_coords=True, # force a zero-distance pair ) base_encoder = ProteinGVPEncoder( @@ -364,7 +386,9 @@ def test_forward_with_duplicate_protein_coords_localizes_nan(device): with FiniteHookManager() as hm: # encoder internals (access via .encoder.encoder for wrapped GVPEncoder) - hm.watch(model.encoder.encoder.input_scalar_encoder, "encoder.input_scalar_encoder") + hm.watch( + model.encoder.encoder.input_scalar_encoder, "encoder.input_scalar_encoder" + ) hm.watch(model.encoder.encoder.input_gvp, "encoder.input_gvp") for i, layer in enumerate(model.encoder.encoder.layers): hm.watch(layer, f"encoder.layers[{i}]") @@ -383,6 +407,7 @@ def test_forward_with_duplicate_protein_coords_localizes_nan(device): # ============== Tests for Flow Matching Fundamentals ============== + @pytest.mark.unit class TestHungarianMatchingConsistency: """Tests for Hungarian matching correctness.""" @@ -400,8 +425,12 @@ def test_hungarian_is_deterministic(self, device): x0_star_2, x1_star_2 = ot_coupling(x1, batch, x0) # Should be identical - assert torch.allclose(x0_star_1, x0_star_2, atol=1e-6), "Hungarian matching is non-deterministic!" - assert torch.allclose(x1_star_1, x1_star_2, atol=1e-6), "Hungarian matching is non-deterministic!" + assert torch.allclose(x0_star_1, x0_star_2, atol=1e-6), ( + "Hungarian matching is non-deterministic!" + ) + assert torch.allclose(x1_star_1, x1_star_2, atol=1e-6), ( + "Hungarian matching is non-deterministic!" + ) def test_hungarian_is_permutation(self, device): """Hungarian matching is a permutation (reordering).""" @@ -417,25 +446,27 @@ def test_hungarian_is_permutation(self, device): for i in range(len(x1_star)): dists = torch.norm(x1 - x1_star[i], dim=-1) min_dist = dists.min().item() - assert min_dist < 1e-5, f"x1_star[{i}] not found in x1 (min_dist={min_dist})" + assert min_dist < 1e-5, ( + f"x1_star[{i}] not found in x1 (min_dist={min_dist})" + ) def test_hungarian_batched_no_cross_contamination(self, device): """Hungarian matching doesn't cross batch boundaries.""" from src.utils import ot_coupling # Two separate graphs - x1 = torch.cat([ - torch.randn(5, 3, device=device), - torch.randn(7, 3, device=device) + 100.0 # Far away - ]) - x0 = torch.cat([ - torch.randn(5, 3, device=device), - torch.randn(7, 3, device=device) + 100.0 - ]) - batch = torch.cat([ - torch.zeros(5, dtype=torch.long), - torch.ones(7, dtype=torch.long) - ]).to(device) + x1 = torch.cat( + [ + torch.randn(5, 3, device=device), + torch.randn(7, 3, device=device) + 100.0, # Far away + ] + ) + x0 = torch.cat( + [torch.randn(5, 3, device=device), torch.randn(7, 3, device=device) + 100.0] + ) + batch = torch.cat( + [torch.zeros(5, dtype=torch.long), torch.ones(7, dtype=torch.long)] + ).to(device) x0_star, x1_star = ot_coupling(x1, batch, x0) @@ -455,13 +486,16 @@ class TestNoiseSamplingScale: def test_compute_sigma_matches_protein_std(self, device): """Sigma should equal std of protein coordinates.""" - data = make_batched_hetero(device, n_graphs=1, n_protein_per=100, n_water_per=20) + data = make_batched_hetero( + device, n_graphs=1, n_protein_per=100, n_water_per=20 + ) sigma = FlowMatcher.compute_sigma(data) - expected_sigma = data['protein'].pos.std().item() + expected_sigma = data["protein"].pos.std().item() - assert abs(sigma - expected_sigma) < 1e-5, \ + assert abs(sigma - expected_sigma) < 1e-5, ( f"Sigma {sigma:.6f} doesn't match protein std {expected_sigma:.6f}" + ) def test_sigma_consistent_across_calls(self, device): """compute_sigma should be deterministic.""" @@ -477,14 +511,15 @@ def test_noise_scale_reasonable(self, device): data = make_batched_hetero(device, n_graphs=1, n_protein_per=50, n_water_per=20) sigma = FlowMatcher.compute_sigma(data) - n_water = data['water'].num_nodes + n_water = data["water"].num_nodes x0 = torch.randn(n_water, 3, device=device) * sigma x0_std = x0.std().item() # Allow some variation due to random sampling - assert abs(x0_std - sigma) / sigma < 0.3, \ + assert abs(x0_std - sigma) / sigma < 0.3, ( f"x0 std {x0_std:.3f} too different from sigma {sigma:.3f}" + ) @pytest.mark.unit @@ -496,28 +531,37 @@ def test_integration_trajectory_length(self, device): data = make_batched_hetero(device, n_graphs=1, n_protein_per=24, n_water_per=12) base_encoder = ProteinGVPEncoder( - node_scalar_in=16, hidden_dims=(64, 8), n_edge_scalar_in=16, + node_scalar_in=16, + hidden_dims=(64, 8), + n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) model = FlowWaterGVP( - encoder=encoder, hidden_dims=(64, 8), layers=1, - k_pw=8, k_ww=8, + encoder=encoder, + hidden_dims=(64, 8), + layers=1, + k_pw=8, + k_ww=8, ).to(device) fm = FlowMatcher(model, p_self_cond=0.0) num_steps = 20 results = fm.rk4_integrate( - data, num_steps=num_steps, use_sc=False, - device=str(device), return_trajectory=True + data, + num_steps=num_steps, + use_sc=False, + device=str(device), + return_trajectory=True, ) # rk4_integrate returns List[Dict], one per input graph result = results[0] - assert len(result['trajectory']) == num_steps, \ + assert len(result["trajectory"]) == num_steps, ( f"Expected {num_steps} trajectory steps, got {len(result['trajectory'])}" + ) def test_interpolation_at_boundaries(self, device): """Interpolation x_t = (1-t)*x0 + t*x1 gives correct values at boundaries.""" @@ -534,16 +578,18 @@ def test_interpolation_at_boundaries(self, device): t_per_atom_0 = t0[batch].unsqueeze(-1) x_t_0 = (1.0 - t_per_atom_0) * x0_star + t_per_atom_0 * x1_star - assert torch.allclose(x_t_0, x0_star, atol=1e-6), \ + assert torch.allclose(x_t_0, x0_star, atol=1e-6), ( "Interpolation at t=0 doesn't match x0" + ) # At t=1: x_t should equal x1_star t1 = torch.ones(1, device=device) t_per_atom_1 = t1[batch].unsqueeze(-1) x_t_1 = (1.0 - t_per_atom_1) * x0_star + t_per_atom_1 * x1_star - assert torch.allclose(x_t_1, x1_star, atol=1e-6), \ + assert torch.allclose(x_t_1, x1_star, atol=1e-6), ( "Interpolation at t=1 doesn't match x1" + ) @pytest.mark.unit @@ -555,14 +601,19 @@ def test_velocity_field_finite(self, device): data = make_batched_hetero(device, n_graphs=1, n_protein_per=24, n_water_per=12) base_encoder = ProteinGVPEncoder( - node_scalar_in=16, hidden_dims=(64, 8), n_edge_scalar_in=16, + node_scalar_in=16, + hidden_dims=(64, 8), + n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) model = FlowWaterGVP( - encoder=encoder, hidden_dims=(64, 8), layers=1, - k_pw=8, k_ww=8, + encoder=encoder, + hidden_dims=(64, 8), + layers=1, + k_pw=8, + k_ww=8, ).to(device) model.eval() @@ -572,27 +623,32 @@ def test_velocity_field_finite(self, device): t = torch.tensor([t_val], device=device) v_pred = model(data, t, sc=None) - assert torch.isfinite(v_pred).all(), \ - f"Velocity has NaN/Inf at t={t_val}" + assert torch.isfinite(v_pred).all(), f"Velocity has NaN/Inf at t={t_val}" # Check magnitude is not absurdly large max_mag = torch.norm(v_pred, dim=-1).max().item() - assert max_mag < 1e6, \ + assert max_mag < 1e6, ( f"Velocity magnitude too large at t={t_val}: {max_mag:.3e}" + ) def test_velocity_field_changes_with_t(self, device): """Velocity field should depend on t (different outputs for different times).""" data = make_batched_hetero(device, n_graphs=1, n_protein_per=24, n_water_per=12) base_encoder = ProteinGVPEncoder( - node_scalar_in=16, hidden_dims=(64, 8), n_edge_scalar_in=16, + node_scalar_in=16, + hidden_dims=(64, 8), + n_edge_scalar_in=16, pool_residue=False, ).to(device) encoder = GVPEncoder(encoder=base_encoder, freeze=False) model = FlowWaterGVP( - encoder=encoder, hidden_dims=(64, 8), layers=1, - k_pw=8, k_ww=8, + encoder=encoder, + hidden_dims=(64, 8), + layers=1, + k_pw=8, + k_ww=8, ).to(device) model.eval() @@ -606,8 +662,7 @@ def test_velocity_field_changes_with_t(self, device): # Velocities should be different diff = torch.norm(v0 - v1, dim=-1).mean().item() - assert diff > 1e-4, \ - f"Velocity field doesn't change with t (diff={diff:.6f})" + assert diff > 1e-4, f"Velocity field doesn't change with t (diff={diff:.6f})" @pytest.mark.unit @@ -625,8 +680,8 @@ def test_velocity_target_scale(self, device): data = make_batched_hetero(device, n_graphs=1, n_protein_per=50, n_water_per=20) - x1 = data['water'].pos - batch = data['water'].batch + x1 = data["water"].pos + batch = data["water"].batch sigma = FlowMatcher.compute_sigma(data) x0 = torch.randn_like(x1) * sigma @@ -638,7 +693,6 @@ def test_velocity_target_scale(self, device): target_mag = torch.norm(v_target, dim=-1).mean().item() # Should be on order of sigma (could be sigma to 3*sigma depending on x1 spread) - assert 0.5 * sigma < target_mag < 5 * sigma, \ + assert 0.5 * sigma < target_mag < 5 * sigma, ( f"Target velocity magnitude {target_mag:.3f} seems off (sigma={sigma:.3f})" - - + ) diff --git a/tests/test_gvp.py b/tests/test_gvp.py index bf7197e..b203398 100644 --- a/tests/test_gvp.py +++ b/tests/test_gvp.py @@ -6,141 +6,141 @@ class TestGVPHelpers: """Tests for GVP helper functions.""" - + def test_tuple_sum(self): """Test tuple summation.""" t1 = (torch.ones(5, 3), torch.ones(5, 2, 3)) t2 = (torch.ones(5, 3) * 2, torch.ones(5, 2, 3) * 2) result = tuple_sum(t1, t2) - + assert torch.allclose(result[0], torch.ones(5, 3) * 3) assert torch.allclose(result[1], torch.ones(5, 2, 3) * 3) - + def test_tuple_cat(self): """Test tuple concatenation.""" t1 = (torch.ones(5, 3), torch.ones(5, 2, 3)) t2 = (torch.ones(5, 4), torch.ones(5, 3, 3)) result = tuple_cat(t1, t2, dim=-1) - + assert result[0].shape == (5, 7) # 3 + 4 assert result[1].shape == (5, 5, 3) # 2 + 3 - + def test_merge_split(self): """Test merge and split are inverses.""" s = torch.randn(10, 5) v = torch.randn(10, 3, 3) - + merged = _merge(s, v) s_out, v_out = _split(merged, nv=3) - + assert torch.allclose(s, s_out) assert torch.allclose(v, v_out) class TestGVP: """Tests for GVP layer.""" - + def test_scalar_only_forward(self): """Test GVP with only scalar inputs/outputs.""" gvp = GVP(in_dims=(10, 0), out_dims=(5, 0)) x = torch.randn(8, 10) - + out = gvp(x) - + assert isinstance(out, torch.Tensor) assert out.shape == (8, 5) - + def test_vector_forward(self): """Test GVP with vector inputs/outputs.""" gvp = GVP( - in_dims=(10, 5), + in_dims=(10, 5), out_dims=(8, 3), activations=(F.relu, torch.sigmoid), - vector_gate=True + vector_gate=True, ) s = torch.randn(16, 10) v = torch.randn(16, 5, 3) - + s_out, v_out = gvp((s, v)) - + assert s_out.shape == (16, 8) assert v_out.shape == (16, 3, 3) - + def test_no_vector_input(self): """Test GVP with scalar input, vector output.""" gvp = GVP(in_dims=(10, 0), out_dims=(8, 3)) x = torch.randn(16, 10) - + s_out, v_out = gvp(x) - + assert s_out.shape == (16, 8) assert v_out.shape == (16, 3, 3) - + def test_deterministic_forward(self): """Test forward pass is deterministic.""" gvp = GVP(in_dims=(5, 2), out_dims=(3, 2)) s = torch.randn(4, 5) v = torch.randn(4, 2, 3) - + out1 = gvp((s, v)) out2 = gvp((s, v)) - + assert torch.allclose(out1[0], out2[0]) assert torch.allclose(out1[1], out2[1]) class TestLayerNorm: """Tests for GVP LayerNorm.""" - + def test_scalar_only(self): """Test with scalar features only.""" ln = LayerNorm((10, 0)) x = torch.randn(8, 10) - + out = ln(x) - + assert out.shape == (8, 10) # Check normalization assert torch.allclose(out.mean(dim=-1), torch.zeros(8), atol=1e-6) - + def test_scalar_vector(self): """Test with both scalar and vector features.""" ln = LayerNorm((10, 5)) s = torch.randn(8, 10) v = torch.randn(8, 5, 3) - + s_out, v_out = ln((s, v)) - + assert s_out.shape == (8, 10) assert v_out.shape == (8, 5, 3) class TestDropout: """Tests for GVP Dropout.""" - + def test_dropout_training(self): """Test dropout in training mode.""" drop = Dropout(drop_rate=0.5) drop.train() - + s = torch.ones(100, 10) v = torch.ones(100, 5, 3) - + s_out, v_out = drop((s, v)) - + # Some values should be zeroed assert not torch.allclose(s_out, s) - + def test_dropout_eval(self): """Test dropout in eval mode (should be identity).""" drop = Dropout(drop_rate=0.5) drop.eval() - + s = torch.ones(100, 10) v = torch.ones(100, 5, 3) - + s_out, v_out = drop((s, v)) - + # Should be unchanged in eval mode assert torch.allclose(s_out, s) - assert torch.allclose(v_out, v) \ No newline at end of file + assert torch.allclose(v_out, v) diff --git a/tests/test_utils.py b/tests/test_utils.py index cf1a20b..8d3446c 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -18,7 +18,7 @@ import torch -matplotlib.use('Agg') +matplotlib.use("Agg") from pathlib import Path @@ -77,7 +77,7 @@ def test_at_cutoff(self): def test_batched_input(self): """RBF should handle batched inputs.""" - r = torch.rand(100, device='cpu') * 10.0 + r = torch.rand(100, device="cpu") * 10.0 out = rbf(r, num_gaussians=16, cutoff=8.0) assert out.shape == (100, 16) assert torch.isfinite(out).all() @@ -231,18 +231,11 @@ def test_is_permutation(self): def test_no_cross_batch_matching(self): """Matching should not cross batch boundaries.""" # Two graphs, far apart - x1 = torch.cat([ - torch.randn(5, 3), - torch.randn(5, 3) + 100.0 - ]) - x0 = torch.cat([ - torch.randn(5, 3), - torch.randn(5, 3) + 100.0 - ]) - batch = torch.cat([ - torch.zeros(5, dtype=torch.long), - torch.ones(5, dtype=torch.long) - ]) + x1 = torch.cat([torch.randn(5, 3), torch.randn(5, 3) + 100.0]) + x0 = torch.cat([torch.randn(5, 3), torch.randn(5, 3) + 100.0]) + batch = torch.cat( + [torch.zeros(5, dtype=torch.long), torch.ones(5, dtype=torch.long)] + ) _, x1_star = ot_coupling(x1, batch, x0) @@ -269,6 +262,7 @@ def test_optimal_matching(self): assert torch.allclose(x1_star[0], torch.tensor([0.0, 0.0, 0.0])) assert torch.allclose(x1_star[1], torch.tensor([1.0, 0.0, 0.0])) + @pytest.mark.unit class TestRecallPrecision: """Tests for recall_precision metric.""" @@ -336,8 +330,8 @@ def test_different_thresholds(self): @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") def test_gpu_tensors(self): """Should work with GPU tensors.""" - pred = torch.rand(10, 3, device='cuda') - true = torch.rand(10, 3, device='cuda') + pred = torch.rand(10, 3, device="cuda") + true = torch.rand(10, 3, device="cuda") recall, precision = recall_precision(pred, true, thresh=5.0) assert isinstance(recall, float) assert isinstance(precision, float) @@ -376,16 +370,20 @@ def test_numpy_input(self): def test_batched(self): """Batched RMSD should average over graphs.""" - pred = torch.tensor([ - [0.0, 0.0, 0.0], # graph 0 - [1.0, 0.0, 0.0], # graph 0 - [0.0, 0.0, 0.0], # graph 1 - ]) - target = torch.tensor([ - [0.0, 0.0, 0.0], - [1.0, 0.0, 0.0], - [2.0, 0.0, 0.0], # 2.0 away - ]) + pred = torch.tensor( + [ + [0.0, 0.0, 0.0], # graph 0 + [1.0, 0.0, 0.0], # graph 0 + [0.0, 0.0, 0.0], # graph 1 + ] + ) + target = torch.tensor( + [ + [0.0, 0.0, 0.0], + [1.0, 0.0, 0.0], + [2.0, 0.0, 0.0], # 2.0 away + ] + ) batch = torch.tensor([0, 0, 1]) rmsd = compute_rmsd(pred, target, batch=batch) @@ -403,9 +401,9 @@ def test_perfect_placement(self): pts = torch.rand(10, 3) metrics = compute_placement_metrics(pts, pts.clone(), threshold=0.1) - assert metrics['recall'] == pytest.approx(1.0) - assert metrics['precision'] == pytest.approx(1.0) - assert metrics['f1'] == pytest.approx(1.0) + assert metrics["recall"] == pytest.approx(1.0) + assert metrics["precision"] == pytest.approx(1.0) + assert metrics["f1"] == pytest.approx(1.0) def test_no_overlap(self): """No overlap should give zero metrics.""" @@ -414,19 +412,19 @@ def test_no_overlap(self): metrics = compute_placement_metrics(pred, true, threshold=1.0) - assert metrics['recall'] == 0.0 - assert metrics['precision'] == 0.0 - assert metrics['f1'] == 0.0 + assert metrics["recall"] == 0.0 + assert metrics["precision"] == 0.0 + assert metrics["f1"] == 0.0 def test_empty_inputs(self): """Empty inputs should return zero metrics.""" metrics = compute_placement_metrics( np.zeros((0, 3)), np.zeros((5, 3)), threshold=1.0 ) - assert metrics['recall'] == 0.0 - assert metrics['precision'] == 0.0 - assert metrics['f1'] == 0.0 - assert metrics['auc_pr'] == 0.0 + assert metrics["recall"] == 0.0 + assert metrics["precision"] == 0.0 + assert metrics["f1"] == 0.0 + assert metrics["auc_pr"] == 0.0 def test_auc_pr_range(self): """AUC-PR should be in [0, 1].""" @@ -435,7 +433,7 @@ def test_auc_pr_range(self): metrics = compute_placement_metrics(pred, true, threshold=1.0) - assert 0.0 <= metrics['auc_pr'] <= 1.0 + assert 0.0 <= metrics["auc_pr"] <= 1.0 def test_f1_formula(self): """F1 should be harmonic mean of precision and recall.""" @@ -444,10 +442,13 @@ def test_f1_formula(self): metrics = compute_placement_metrics(pred, true, threshold=1.0) - expected_f1 = 2 * metrics['precision'] * metrics['recall'] / ( - metrics['precision'] + metrics['recall'] + 1e-8 + expected_f1 = ( + 2 + * metrics["precision"] + * metrics["recall"] + / (metrics["precision"] + metrics["recall"] + 1e-8) ) - assert metrics['f1'] == pytest.approx(expected_f1, abs=1e-6) + assert metrics["f1"] == pytest.approx(expected_f1, abs=1e-6) @pytest.mark.unit @@ -457,51 +458,51 @@ class TestPlot3DFrame: def test_runs_without_error(self): """Basic plot should not raise.""" import matplotlib.pyplot as plt + fig = plt.figure() - ax = fig.add_subplot(111, projection='3d') + ax = fig.add_subplot(111, projection="3d") plot_3d_frame( ax, np.random.rand(10, 3), np.random.rand(3, 3), np.random.rand(5, 3), np.random.rand(5, 3), - title="Test" + title="Test", ) plt.close(fig) def test_no_mates(self): """Plot with no mates should work.""" import matplotlib.pyplot as plt + fig = plt.figure() - ax = fig.add_subplot(111, projection='3d') + ax = fig.add_subplot(111, projection="3d") plot_3d_frame( - ax, - np.random.rand(10, 3), - None, - np.random.rand(5, 3), - np.random.rand(5, 3) + ax, np.random.rand(10, 3), None, np.random.rand(5, 3), np.random.rand(5, 3) ) plt.close(fig) def test_empty_mates(self): """Plot with empty mates array should work.""" import matplotlib.pyplot as plt + fig = plt.figure() - ax = fig.add_subplot(111, projection='3d') + ax = fig.add_subplot(111, projection="3d") plot_3d_frame( ax, np.random.rand(10, 3), np.zeros((0, 3)), # Empty mates np.random.rand(5, 3), - np.random.rand(5, 3) + np.random.rand(5, 3), ) plt.close(fig) def test_with_axis_limits(self): """Plot with axis limits should work.""" import matplotlib.pyplot as plt + fig = plt.figure() - ax = fig.add_subplot(111, projection='3d') + ax = fig.add_subplot(111, projection="3d") plot_3d_frame( ax, np.random.rand(10, 3), @@ -510,7 +511,7 @@ def test_with_axis_limits(self): np.random.rand(5, 3), xlim=(-1, 1), ylim=(-1, 1), - zlim=(-1, 1) + zlim=(-1, 1), ) plt.close(fig) @@ -522,10 +523,7 @@ class TestSaveProteinPlot: def test_saves_file(self, tmp_path): """Plot should be saved to disk.""" save_protein_plot( - torch.rand(20, 3), - torch.rand(20, 3), - step=1, - save_dir=str(tmp_path) + torch.rand(20, 3), torch.rand(20, 3), step=1, save_dir=str(tmp_path) ) assert (tmp_path / "step_1.png").exists() @@ -533,10 +531,7 @@ def test_different_sizes(self, tmp_path): """Should work with different protein sizes.""" for n in [5, 20, 100]: save_protein_plot( - torch.rand(n, 3), - torch.rand(n, 3), - step=n, - save_dir=str(tmp_path) + torch.rand(n, 3), torch.rand(n, 3), step=n, save_dir=str(tmp_path) ) assert (tmp_path / f"step_{n}.png").exists() @@ -557,7 +552,7 @@ def test_creates_gif(self, tmp_path): protein_pos=protein_pos, water_true=water_true, save_path=gif_path, - fps=5 + fps=5, ) assert Path(gif_path).exists() @@ -574,7 +569,7 @@ def test_with_pdb_id(self, tmp_path): protein_pos=protein_pos, water_true=water_true, save_path=gif_path, - pdb_id="1ABC" + pdb_id="1ABC", ) assert Path(gif_path).exists() @@ -590,7 +585,7 @@ def test_long_trajectory_sampled(self, tmp_path): trajectory=trajectory, protein_pos=protein_pos, water_true=water_true, - save_path=gif_path + save_path=gif_path, ) assert Path(gif_path).exists() From 320fd4da83aa06313d731575719d8cf84640ae1f Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Mon, 16 Mar 2026 16:29:17 -0500 Subject: [PATCH 10/19] fixing failing tests in test_embedding_generation due to import issues --- tests/test_embedding_generation.py | 54 ++++++++++++++++++++---------- 1 file changed, 36 insertions(+), 18 deletions(-) diff --git a/tests/test_embedding_generation.py b/tests/test_embedding_generation.py index 4456df2..dbf2711 100644 --- a/tests/test_embedding_generation.py +++ b/tests/test_embedding_generation.py @@ -202,6 +202,8 @@ def test_perfect_alignment_single_residue(self, make_slae_test_data): """Same atoms in same order align exactly.""" from unittest.mock import patch + from scripts import generate_slae_embeddings + atom_specs = [ ("A", 1, "", "N"), ("A", 1, "", "CA"), @@ -217,8 +219,8 @@ def test_perfect_alignment_single_residue(self, make_slae_test_data): ("A", 1, "", "O"), ] - with patch( - "scripts.generate_slae_embeddings.PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS + with patch.object( + generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS ): from scripts.generate_slae_embeddings import align_slae_to_geometry @@ -241,6 +243,8 @@ def test_reordered_atoms(self, make_slae_test_data): """Atoms reordered correctly when geometry order differs.""" from unittest.mock import patch + from scripts import generate_slae_embeddings + atom_specs = [ ("A", 1, "", "N"), ("A", 1, "", "CA"), @@ -257,8 +261,8 @@ def test_reordered_atoms(self, make_slae_test_data): ("A", 1, "", "N"), ] - with patch( - "scripts.generate_slae_embeddings.PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS + with patch.object( + generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS ): from scripts.generate_slae_embeddings import align_slae_to_geometry @@ -281,6 +285,8 @@ def test_noncanonical_atom_gets_zero_vector(self, make_slae_test_data): """Non-SLAE atoms (e.g., 'XX1') get zero embeddings.""" from unittest.mock import patch + from scripts import generate_slae_embeddings + atom_specs = [ ("A", 1, "", "N"), ("A", 1, "", "CA"), @@ -294,8 +300,8 @@ def test_noncanonical_atom_gets_zero_vector(self, make_slae_test_data): ("A", 1, "", "CA"), ] - with patch( - "scripts.generate_slae_embeddings.PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS + with patch.object( + generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS ): from scripts.generate_slae_embeddings import align_slae_to_geometry @@ -319,6 +325,8 @@ def test_multiple_chains(self, make_slae_test_data): """Same res_id in different chains distinguished correctly.""" from unittest.mock import patch + from scripts import generate_slae_embeddings + atom_specs = [ ("A", 1, "", "N"), ("A", 1, "", "CA"), @@ -335,8 +343,8 @@ def test_multiple_chains(self, make_slae_test_data): ("A", 1, "", "CA"), ] - with patch( - "scripts.generate_slae_embeddings.PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS + with patch.object( + generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS ): from scripts.generate_slae_embeddings import align_slae_to_geometry @@ -359,6 +367,8 @@ def test_with_insertion_codes(self, make_slae_test_data): """Insertion codes differentiate atoms properly.""" from unittest.mock import patch + from scripts import generate_slae_embeddings + atom_specs = [ ("A", 1, "", "N"), ("A", 1, "", "CA"), @@ -374,8 +384,8 @@ def test_with_insertion_codes(self, make_slae_test_data): ("A", 1, "", "CA"), ] - with patch( - "scripts.generate_slae_embeddings.PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS + with patch.object( + generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS ): from scripts.generate_slae_embeddings import align_slae_to_geometry @@ -398,6 +408,8 @@ def test_empty_geometry_list(self, make_slae_test_data): """Returns shape (0, embedding_dim) for empty geometry.""" from unittest.mock import patch + from scripts import generate_slae_embeddings + atom_specs = [ ("A", 1, "", "N"), ("A", 1, "", "CA"), @@ -406,8 +418,8 @@ def test_empty_geometry_list(self, make_slae_test_data): geometry_atom_info = [] - with patch( - "scripts.generate_slae_embeddings.PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS + with patch.object( + generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS ): from scripts.generate_slae_embeddings import align_slae_to_geometry @@ -430,6 +442,8 @@ def test_empty_slae_embeddings(self): import numpy as np + from scripts import generate_slae_embeddings + # Empty SLAE data slae_emb = torch.zeros(0, 128) slae_residue_idx = torch.zeros(0, dtype=torch.long) @@ -443,8 +457,8 @@ def test_empty_slae_embeddings(self): ("A", 1, "", "CA"), ] - with patch( - "scripts.generate_slae_embeddings.PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS + with patch.object( + generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS ): from scripts.generate_slae_embeddings import align_slae_to_geometry @@ -466,6 +480,8 @@ def test_negative_residue_ids(self, make_slae_test_data): """Handles negative res_ids (like HIS A -1 in 6eey).""" from unittest.mock import patch + from scripts import generate_slae_embeddings + atom_specs = [ ("A", -1, "", "N"), ("A", -1, "", "CA"), @@ -481,8 +497,8 @@ def test_negative_residue_ids(self, make_slae_test_data): ("A", -1, "", "N"), ] - with patch( - "scripts.generate_slae_embeddings.PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS + with patch.object( + generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS ): from scripts.generate_slae_embeddings import align_slae_to_geometry @@ -505,6 +521,8 @@ def test_with_6eey_pdb_subset(self, make_slae_test_data): """Integration test with real PDB atom patterns (6eey style).""" from unittest.mock import patch + from scripts import generate_slae_embeddings + # Simulate 6eey-like structure: multiple chains, insertion codes, gaps atom_specs = [ # Chain A residue 1 @@ -536,8 +554,8 @@ def test_with_6eey_pdb_subset(self, make_slae_test_data): ("A", 1, "", "OXT"), # Non-canonical - not in SLAE ] - with patch( - "scripts.generate_slae_embeddings.PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS + with patch.object( + generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS ): from scripts.generate_slae_embeddings import align_slae_to_geometry From 190a3d87b484361ef87440eaf7aa936b07c0513d Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Mon, 16 Mar 2026 16:37:00 -0500 Subject: [PATCH 11/19] removing any SLAE related tests --- tests/test_embedding_generation.py | 446 +---------------------------- 1 file changed, 1 insertion(+), 445 deletions(-) diff --git a/tests/test_embedding_generation.py b/tests/test_embedding_generation.py index 095b862..87a9e94 100644 --- a/tests/test_embedding_generation.py +++ b/tests/test_embedding_generation.py @@ -2,7 +2,7 @@ Tests for embedding generation and loading. Tests: -1. Dataset loading with SLAE embeddings +1. Dataset loading with pre-computed embeddings 2. Backward compatibility (dataset works without embeddings) """ @@ -137,449 +137,5 @@ def test_embedding_with_mates(self): ) -class TestAlignSlaeToGeometry: - """Test the align_slae_to_geometry function.""" - - # Mock PROTEIN_ATOMS for testing (subset of actual 37 atom types) - MOCK_PROTEIN_ATOMS = ["N", "CA", "C", "O", "CB", "CG", "CD", "NE", "CZ", "NH1"] - - @pytest.fixture - def make_slae_test_data(self): - """Factory fixture to create SLAE-like test data from atom specs. - - Args: - atom_specs: List of (chain, res_id, ins_code, atom_name) tuples - - Returns: - Dict with slae_emb, slae_residue_idx, slae_atom_type, - slae_chains, slae_residue_ids, slae_ins_codes - """ - import numpy as np - - def _make_data(atom_specs, embedding_dim=128): - # Build residue info from unique (chain, res_id, ins_code) tuples - residue_keys = [] - for chain, res_id, ins_code, _ in atom_specs: - key = (chain, res_id, ins_code) - if key not in residue_keys: - residue_keys.append(key) - - # Map residue key to index - res_key_to_idx = {k: i for i, k in enumerate(residue_keys)} - - # Build arrays - n_atoms = len(atom_specs) - - slae_residue_idx = torch.zeros(n_atoms, dtype=torch.long) - slae_atom_type = torch.zeros(n_atoms, dtype=torch.long) - # Generate distinguishable embeddings: atom i gets value i+1 in first dim - slae_emb = torch.zeros(n_atoms, embedding_dim) - - for i, (chain, res_id, ins_code, atom_name) in enumerate(atom_specs): - res_key = (chain, res_id, ins_code) - slae_residue_idx[i] = res_key_to_idx[res_key] - # Map atom name to type index - if atom_name in self.MOCK_PROTEIN_ATOMS: - slae_atom_type[i] = self.MOCK_PROTEIN_ATOMS.index(atom_name) - else: - slae_atom_type[i] = 0 # Fallback - # Use index-based value for easy verification - slae_emb[i, 0] = float(i + 1) - - # Build per-residue arrays - slae_chains = np.array([k[0] for k in residue_keys]) - slae_residue_ids = torch.tensor([k[1] for k in residue_keys]) - slae_ins_codes = np.array([k[2] for k in residue_keys]) - - return { - "slae_emb": slae_emb, - "slae_residue_idx": slae_residue_idx, - "slae_atom_type": slae_atom_type, - "slae_chains": slae_chains, - "slae_residue_ids": slae_residue_ids, - "slae_ins_codes": slae_ins_codes, - } - - return _make_data - - @pytest.mark.unit - def test_perfect_alignment_single_residue(self, make_slae_test_data): - """Same atoms in same order align exactly.""" - from unittest.mock import patch - - from scripts import generate_slae_embeddings - - atom_specs = [ - ("A", 1, "", "N"), - ("A", 1, "", "CA"), - ("A", 1, "", "C"), - ("A", 1, "", "O"), - ] - data = make_slae_test_data(atom_specs) - - geometry_atom_info = [ - ("A", 1, "", "N"), - ("A", 1, "", "CA"), - ("A", 1, "", "C"), - ("A", 1, "", "O"), - ] - - with patch.object( - generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS - ): - from scripts.generate_slae_embeddings import align_slae_to_geometry - - aligned = align_slae_to_geometry( - data["slae_emb"], - data["slae_residue_idx"], - data["slae_atom_type"], - data["slae_chains"], - data["slae_residue_ids"], - data["slae_ins_codes"], - geometry_atom_info, - ) - - assert aligned.shape == (4, 128) - # Check first dimension values match expected order - assert torch.allclose(aligned[:, 0], torch.tensor([1.0, 2.0, 3.0, 4.0])) - - @pytest.mark.unit - def test_reordered_atoms(self, make_slae_test_data): - """Atoms reordered correctly when geometry order differs.""" - from unittest.mock import patch - - from scripts import generate_slae_embeddings - - atom_specs = [ - ("A", 1, "", "N"), - ("A", 1, "", "CA"), - ("A", 1, "", "C"), - ("A", 1, "", "O"), - ] - data = make_slae_test_data(atom_specs) - - # Geometry has different order - geometry_atom_info = [ - ("A", 1, "", "O"), - ("A", 1, "", "C"), - ("A", 1, "", "CA"), - ("A", 1, "", "N"), - ] - - with patch.object( - generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS - ): - from scripts.generate_slae_embeddings import align_slae_to_geometry - - aligned = align_slae_to_geometry( - data["slae_emb"], - data["slae_residue_idx"], - data["slae_atom_type"], - data["slae_chains"], - data["slae_residue_ids"], - data["slae_ins_codes"], - geometry_atom_info, - ) - - assert aligned.shape == (4, 128) - # O was index 3 (value 4), C was 2 (value 3), CA was 1 (value 2), N was 0 (value 1) - assert torch.allclose(aligned[:, 0], torch.tensor([4.0, 3.0, 2.0, 1.0])) - - @pytest.mark.unit - def test_noncanonical_atom_gets_zero_vector(self, make_slae_test_data): - """Non-SLAE atoms (e.g., 'XX1') get zero embeddings.""" - from unittest.mock import patch - - from scripts import generate_slae_embeddings - - atom_specs = [ - ("A", 1, "", "N"), - ("A", 1, "", "CA"), - ] - data = make_slae_test_data(atom_specs) - - # Geometry has a non-canonical atom - geometry_atom_info = [ - ("A", 1, "", "N"), - ("A", 1, "", "XX1"), # Non-canonical - ("A", 1, "", "CA"), - ] - - with patch.object( - generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS - ): - from scripts.generate_slae_embeddings import align_slae_to_geometry - - aligned = align_slae_to_geometry( - data["slae_emb"], - data["slae_residue_idx"], - data["slae_atom_type"], - data["slae_chains"], - data["slae_residue_ids"], - data["slae_ins_codes"], - geometry_atom_info, - ) - - assert aligned.shape == (3, 128) - assert torch.allclose(aligned[0, 0], torch.tensor(1.0)) # N - assert torch.allclose(aligned[1], torch.zeros(128)) # XX1 -> zero vector - assert torch.allclose(aligned[2, 0], torch.tensor(2.0)) # CA - - @pytest.mark.unit - def test_multiple_chains(self, make_slae_test_data): - """Same res_id in different chains distinguished correctly.""" - from unittest.mock import patch - - from scripts import generate_slae_embeddings - - atom_specs = [ - ("A", 1, "", "N"), - ("A", 1, "", "CA"), - ("B", 1, "", "N"), - ("B", 1, "", "CA"), - ] - data = make_slae_test_data(atom_specs) - - # Geometry requests atoms from both chains - geometry_atom_info = [ - ("B", 1, "", "CA"), - ("A", 1, "", "N"), - ("B", 1, "", "N"), - ("A", 1, "", "CA"), - ] - - with patch.object( - generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS - ): - from scripts.generate_slae_embeddings import align_slae_to_geometry - - aligned = align_slae_to_geometry( - data["slae_emb"], - data["slae_residue_idx"], - data["slae_atom_type"], - data["slae_chains"], - data["slae_residue_ids"], - data["slae_ins_codes"], - geometry_atom_info, - ) - - assert aligned.shape == (4, 128) - # B:1:CA=4, A:1:N=1, B:1:N=3, A:1:CA=2 - assert torch.allclose(aligned[:, 0], torch.tensor([4.0, 1.0, 3.0, 2.0])) - - @pytest.mark.unit - def test_with_insertion_codes(self, make_slae_test_data): - """Insertion codes differentiate atoms properly.""" - from unittest.mock import patch - - from scripts import generate_slae_embeddings - - atom_specs = [ - ("A", 1, "", "N"), - ("A", 1, "", "CA"), - ("A", 1, "A", "N"), # Same res_id but with insertion code - ("A", 1, "A", "CA"), - ] - data = make_slae_test_data(atom_specs) - - geometry_atom_info = [ - ("A", 1, "A", "CA"), - ("A", 1, "", "N"), - ("A", 1, "A", "N"), - ("A", 1, "", "CA"), - ] - - with patch.object( - generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS - ): - from scripts.generate_slae_embeddings import align_slae_to_geometry - - aligned = align_slae_to_geometry( - data["slae_emb"], - data["slae_residue_idx"], - data["slae_atom_type"], - data["slae_chains"], - data["slae_residue_ids"], - data["slae_ins_codes"], - geometry_atom_info, - ) - - assert aligned.shape == (4, 128) - # A:1:A:CA=4, A:1::N=1, A:1:A:N=3, A:1::CA=2 - assert torch.allclose(aligned[:, 0], torch.tensor([4.0, 1.0, 3.0, 2.0])) - - @pytest.mark.unit - def test_empty_geometry_list(self, make_slae_test_data): - """Returns shape (0, embedding_dim) for empty geometry.""" - from unittest.mock import patch - - from scripts import generate_slae_embeddings - - atom_specs = [ - ("A", 1, "", "N"), - ("A", 1, "", "CA"), - ] - data = make_slae_test_data(atom_specs) - - geometry_atom_info = [] - - with patch.object( - generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS - ): - from scripts.generate_slae_embeddings import align_slae_to_geometry - - aligned = align_slae_to_geometry( - data["slae_emb"], - data["slae_residue_idx"], - data["slae_atom_type"], - data["slae_chains"], - data["slae_residue_ids"], - data["slae_ins_codes"], - geometry_atom_info, - ) - - assert aligned.shape == (0, 128) - - @pytest.mark.unit - def test_empty_slae_embeddings(self): - """All geometry atoms get zero vectors when SLAE is empty.""" - from unittest.mock import patch - - import numpy as np - - from scripts import generate_slae_embeddings - - # Empty SLAE data - slae_emb = torch.zeros(0, 128) - slae_residue_idx = torch.zeros(0, dtype=torch.long) - slae_atom_type = torch.zeros(0, dtype=torch.long) - slae_chains = np.array([]) - slae_residue_ids = torch.zeros(0, dtype=torch.long) - slae_ins_codes = np.array([]) - - geometry_atom_info = [ - ("A", 1, "", "N"), - ("A", 1, "", "CA"), - ] - - with patch.object( - generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS - ): - from scripts.generate_slae_embeddings import align_slae_to_geometry - - aligned = align_slae_to_geometry( - slae_emb, - slae_residue_idx, - slae_atom_type, - slae_chains, - slae_residue_ids, - slae_ins_codes, - geometry_atom_info, - ) - - assert aligned.shape == (2, 128) - assert torch.allclose(aligned, torch.zeros(2, 128)) - - @pytest.mark.unit - def test_negative_residue_ids(self, make_slae_test_data): - """Handles negative res_ids (like HIS A -1 in 6eey).""" - from unittest.mock import patch - - from scripts import generate_slae_embeddings - - atom_specs = [ - ("A", -1, "", "N"), - ("A", -1, "", "CA"), - ("A", 0, "", "N"), - ("A", 1, "", "N"), - ] - data = make_slae_test_data(atom_specs) - - geometry_atom_info = [ - ("A", 1, "", "N"), - ("A", -1, "", "CA"), - ("A", 0, "", "N"), - ("A", -1, "", "N"), - ] - - with patch.object( - generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS - ): - from scripts.generate_slae_embeddings import align_slae_to_geometry - - aligned = align_slae_to_geometry( - data["slae_emb"], - data["slae_residue_idx"], - data["slae_atom_type"], - data["slae_chains"], - data["slae_residue_ids"], - data["slae_ins_codes"], - geometry_atom_info, - ) - - assert aligned.shape == (4, 128) - # A:1:N=4, A:-1:CA=2, A:0:N=3, A:-1:N=1 - assert torch.allclose(aligned[:, 0], torch.tensor([4.0, 2.0, 3.0, 1.0])) - - @pytest.mark.unit - def test_with_6eey_pdb_subset(self, make_slae_test_data): - """Integration test with real PDB atom patterns (6eey style).""" - from unittest.mock import patch - - from scripts import generate_slae_embeddings - - # Simulate 6eey-like structure: multiple chains, insertion codes, gaps - atom_specs = [ - # Chain A residue 1 - ("A", 1, "", "N"), - ("A", 1, "", "CA"), - ("A", 1, "", "C"), - ("A", 1, "", "O"), - ("A", 1, "", "CB"), - # Chain A residue 2 with insertion code - ("A", 2, "A", "N"), - ("A", 2, "A", "CA"), - # Chain B residue 1 (same res_id as chain A) - ("B", 1, "", "N"), - ("B", 1, "", "CA"), - ] - data = make_slae_test_data(atom_specs) - - # Geometry with different order and some missing atoms - geometry_atom_info = [ - ("B", 1, "", "CA"), # Chain B first - ("B", 1, "", "N"), - ("A", 2, "A", "N"), # Insertion code residue - ("A", 1, "", "CB"), - ("A", 1, "", "O"), - ("A", 1, "", "C"), - ("A", 1, "", "CA"), - ("A", 1, "", "N"), - ("A", 2, "A", "CA"), - ("A", 1, "", "OXT"), # Non-canonical - not in SLAE - ] - - with patch.object( - generate_slae_embeddings, "PROTEIN_ATOMS", self.MOCK_PROTEIN_ATOMS - ): - from scripts.generate_slae_embeddings import align_slae_to_geometry - - aligned = align_slae_to_geometry( - data["slae_emb"], - data["slae_residue_idx"], - data["slae_atom_type"], - data["slae_chains"], - data["slae_residue_ids"], - data["slae_ins_codes"], - geometry_atom_info, - ) - - assert aligned.shape == (10, 128) - # Verify specific mappings: - # B:1:CA=9, B:1:N=8, A:2A:N=6, A:1:CB=5, A:1:O=4, A:1:C=3, A:1:CA=2, A:1:N=1, A:2A:CA=7, OXT=0 - expected = torch.tensor([9.0, 8.0, 6.0, 5.0, 4.0, 3.0, 2.0, 1.0, 7.0, 0.0]) - assert torch.allclose(aligned[:, 0], expected) - - if __name__ == "__main__": pytest.main([__file__, "-v"]) From 1fc258db95bef14bf0d264d24b5b6cf35e349853 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava <114123331+vratins@users.noreply.github.com> Date: Mon, 16 Mar 2026 16:45:39 -0500 Subject: [PATCH 12/19] Delete scripts/generate_water_plots.py --- scripts/generate_water_plots.py | 767 -------------------------------- 1 file changed, 767 deletions(-) delete mode 100644 scripts/generate_water_plots.py diff --git a/scripts/generate_water_plots.py b/scripts/generate_water_plots.py deleted file mode 100644 index f9fa269..0000000 --- a/scripts/generate_water_plots.py +++ /dev/null @@ -1,767 +0,0 @@ -""" -Generate whole-dataset distribution and correlation plots for water quality metrics. - -Generates: -- EDIA distribution histograms (all waters + mean per PDB) -- RSCC distribution histograms (all waters + mean per PDB) -- B-factor distribution histograms (all waters + mean per PDB) -- EDIA vs B-factor correlation scatter/hexbin - -Usage: - uv run scripts/generate_water_plots.py --output-dir figures/water_quality - uv run scripts/generate_water_plots.py --skip-bfactor --output-dir figures/water_quality - uv run scripts/generate_water_plots.py --bfactor-only --bfactor-normalization water --output-dir figures/water_quality -""" - -import argparse -from concurrent.futures import as_completed, ProcessPoolExecutor -from pathlib import Path - -import matplotlib.pyplot as plt -import numpy as np -import pandas as pd -from biotite.structure.io.pdb import PDBFile -from loguru import logger -from tqdm import tqdm - - -def load_all_water_data(edia_dir: Path) -> pd.DataFrame: - """Load all EDIA CSV files and filter for water (HOH) residues only.""" - all_data = [] - - csv_files = list(edia_dir.rglob("*_residue_stats.csv")) - logger.info(f"Found {len(csv_files)} EDIA CSV files") - - for csv_file in tqdm(csv_files, desc="Loading EDIA data"): - try: - df = pd.read_csv(csv_file) - # Filter for water molecules only - water_df = df[df["compID"] == "HOH"].copy() - if len(water_df) > 0: - # Add PDB ID from filename - pdb_id = csv_file.stem.replace("_residue_stats", "") - water_df["pdb_id"] = pdb_id - all_data.append(water_df) - except Exception as e: - logger.error(f"Error reading {csv_file}: {e}") - - if not all_data: - raise ValueError("No water data found in any CSV files") - - combined = pd.concat(all_data, ignore_index=True) - logger.info(f"Loaded {len(combined)} water molecules from {len(all_data)} PDBs") - return combined - - -def extract_water_bfactors_from_pdb( - pdb_path: Path, - normalization: str = "all", -) -> pd.DataFrame | None: - """Extract B-factors for water molecules from a PDB file using biotite. - - B-factors are normalized using statistics from a chosen atom subset - to account for structure-to-structure variation. - - Args: - pdb_path: Path to the PDB file - normalization: Strategy for computing normalization statistics: - - "all": Use all atoms in the PDB (default) - - "protein": Use only protein atoms (excludes waters, ligands) - - "water": Use only water atoms (HOH/WAT) - - Returns: - DataFrame with columns: pdb_id, chain_id, res_id, b_factor, b_factor_normalized - Returns None if extraction fails - """ - try: - pdb_file = PDBFile.read(pdb_path) - atoms = pdb_file.get_structure( - model=1, altloc="occupancy", extra_fields=["b_factor"] - ) - - # Filter for water molecules (HOH or WAT) - water_mask = (atoms.res_name == "HOH") | (atoms.res_name == "WAT") - water_atoms = atoms[water_mask] - - # Compute B-factor statistics based on normalization strategy - if normalization == "protein": - # Standard amino acid residue names - protein_residues = { - "ALA", - "ARG", - "ASN", - "ASP", - "CYS", - "GLN", - "GLU", - "GLY", - "HIS", - "ILE", - "LEU", - "LYS", - "MET", - "PHE", - "PRO", - "SER", - "THR", - "TRP", - "TYR", - "VAL", - } - protein_mask = np.isin(atoms.res_name, list(protein_residues)) - norm_bfactors = atoms.b_factor[protein_mask] - elif normalization == "water": - norm_bfactors = water_atoms.b_factor - else: # "all" - norm_bfactors = atoms.b_factor - - if len(norm_bfactors) == 0: - return None - - pdb_mean = np.mean(norm_bfactors) - pdb_std = np.std(norm_bfactors) - - if len(water_atoms) == 0: - return None - - # extract PDB ID from filename (e.g., "3ilf_final.pdb" -> "3ilf") - pdb_id = pdb_path.stem.replace("_final", "") - - # build DataFrame with one row per unique water residue - # water molecules have one oxygen atom, so we take unique (chain, res_id) pairs - records = [] - seen = set() - for i in range(len(water_atoms)): - chain_id = water_atoms.chain_id[i] - res_id = water_atoms.res_id[i] - key = (chain_id, res_id) - if key not in seen: - seen.add(key) - raw_bfactor = water_atoms.b_factor[i] - # z-score using whole-PDB statistics - normalized = (raw_bfactor - pdb_mean) / pdb_std if pdb_std > 0 else 0.0 - records.append( - { - "pdb_id": pdb_id, - "chain_id": chain_id, - "res_id": res_id, - "b_factor": raw_bfactor, - "b_factor_normalized": normalized, - } - ) - - return pd.DataFrame(records) - - except Exception as e: - logger.error(f"Error extracting B-factors from {pdb_path}: {e}") - return None - - -def _extract_bfactors_worker(args: tuple) -> pd.DataFrame | None: - """Worker function for parallel B-factor extraction.""" - pdb_path, pdb_id, normalization = args - return extract_water_bfactors_from_pdb(pdb_path, normalization=normalization) - - -def load_all_bfactors( - pdb_dir: Path, - pdb_ids: list[str], - num_workers: int = 4, - normalization: str = "all", -) -> pd.DataFrame: - """Load B-factors for all PDB IDs in parallel. - - Args: - pdb_dir: Directory containing PDB files (organized as pdb_dir/pdb_id/pdb_id_final.pdb) - pdb_ids: List of PDB IDs to process - num_workers: Number of parallel workers - normalization: Strategy for B-factor normalization ("all", "protein", or "water") - - Returns: - DataFrame with columns: pdb_id, chain_id, res_id, b_factor, b_factor_normalized - """ - # Build list of (pdb_path, pdb_id, normalization) tuples - tasks = [] - for pdb_id in pdb_ids: - pdb_path = pdb_dir / pdb_id / f"{pdb_id}_final.pdb" - if pdb_path.exists(): - tasks.append((pdb_path, pdb_id, normalization)) - - logger.info(f"Found {len(tasks)} PDB files out of {len(pdb_ids)} requested") - - all_bfactors = [] - - with ProcessPoolExecutor(max_workers=num_workers) as executor: - futures = { - executor.submit(_extract_bfactors_worker, task): task[1] for task in tasks - } - - for future in tqdm( - as_completed(futures), total=len(futures), desc="Extracting B-factors" - ): - result = future.result() - if result is not None: - all_bfactors.append(result) - - if not all_bfactors: - raise ValueError("No B-factor data extracted from any PDB files") - - combined = pd.concat(all_bfactors, ignore_index=True) - logger.info( - f"Extracted B-factors for {len(combined)} water molecules from {len(all_bfactors)} PDBs" - ) - return combined - - -def merge_edia_with_bfactors( - edia_df: pd.DataFrame, bfactor_df: pd.DataFrame -) -> pd.DataFrame: - """Merge EDIA data with B-factor data. - - Matching is done on (pdb_id, chain, residue_number): - - EDIA: pdb_strandID (chain), pdb_seqNum (residue number) - - PDB: chain_id, res_id - - Args: - edia_df: DataFrame with EDIA data (must have pdb_id, pdb_strandID, pdb_seqNum) - bfactor_df: DataFrame with B-factor data (must have pdb_id, chain_id, res_id, - b_factor, b_factor_normalized) - - Returns: - Merged DataFrame with b_factor and b_factor_normalized columns added - """ - # rename B-factor columns to match EDIA column names - bfactor_renamed = bfactor_df.rename( - columns={ - "chain_id": "pdb_strandID", - "res_id": "pdb_seqNum", - } - ) - - # merge on the matching key - merged = edia_df.merge( - bfactor_renamed[ - ["pdb_id", "pdb_strandID", "pdb_seqNum", "b_factor", "b_factor_normalized"] - ], - on=["pdb_id", "pdb_strandID", "pdb_seqNum"], - how="left", - ) - - # report match statistics - n_total = len(merged) - n_matched = merged["b_factor"].notna().sum() - match_rate = 100 * n_matched / n_total if n_total > 0 else 0 - logger.info(f"B-factor match rate: {n_matched}/{n_total} ({match_rate:.1f}%)") - - return merged - - -def plot_ediam_waters(df: pd.DataFrame, output_dir: Path): - """Plot histogram of EDIAm for all water molecules.""" - fig, ax = plt.subplots(figsize=(10, 6)) - - ax.hist( - df["EDIAm"].dropna(), bins=50, edgecolor="black", alpha=0.7, color="steelblue" - ) - - # Add threshold lines - ax.axvline(x=0.4, color="red", linestyle="--", linewidth=2, label="EDIAm = 0.4") - ax.axvline(x=0.8, color="orange", linestyle="--", linewidth=2, label="EDIAm = 0.8") - - # Add statistics - n_total = len(df) - mean_val = df["EDIAm"].mean() - median_val = df["EDIAm"].median() - - # Count waters in each threshold region - n_low = (df["EDIAm"] < 0.4).sum() - n_mid = ((df["EDIAm"] >= 0.4) & (df["EDIAm"] < 0.8)).sum() - n_high = (df["EDIAm"] >= 0.8).sum() - - textstr = ( - f"n = {n_total:,}\n" - f"mean = {mean_val:.3f}\n" - f"median = {median_val:.3f}\n" - f"─────────────\n" - f"< 0.4: {n_low:,} ({100 * n_low / n_total:.1f}%)\n" - f"0.4–0.8: {n_mid:,} ({100 * n_mid / n_total:.1f}%)\n" - f"≥ 0.8: {n_high:,} ({100 * n_high / n_total:.1f}%)" - ) - ax.text( - 0.02, - 0.98, - textstr, - transform=ax.transAxes, - fontsize=10, - verticalalignment="top", - bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), - ) - - ax.set_xlabel("EDIAm Score", fontsize=12) - ax.set_ylabel("Count", fontsize=12) - ax.set_title("Distribution of Water EDIAm Scores (All Waters)", fontsize=14) - ax.legend(loc="upper right") - - plt.tight_layout() - fig.savefig(output_dir / "01_ediam_waters.png", dpi=150) - plt.close(fig) - logger.info("Saved: 01_ediam_waters.png") - - -def plot_ediam_pdbs(df: pd.DataFrame, output_dir: Path): - """Plot histogram of mean EDIAm per PDB.""" - pdb_means = df.groupby("pdb_id")["EDIAm"].mean() - - fig, ax = plt.subplots(figsize=(10, 6)) - - ax.hist( - pdb_means.dropna(), bins=50, edgecolor="black", alpha=0.7, color="steelblue" - ) - - # Add threshold lines - ax.axvline(x=0.4, color="red", linestyle="--", linewidth=2, label="EDIAm = 0.4") - ax.axvline(x=0.8, color="orange", linestyle="--", linewidth=2, label="EDIAm = 0.8") - - # Add statistics - n_pdbs = len(pdb_means) - mean_val = pdb_means.mean() - median_val = pdb_means.median() - - # Count PDBs in each threshold region - n_low = (pdb_means < 0.4).sum() - n_mid = ((pdb_means >= 0.4) & (pdb_means < 0.8)).sum() - n_high = (pdb_means >= 0.8).sum() - - textstr = ( - f"n = {n_pdbs:,} PDBs\n" - f"mean = {mean_val:.3f}\n" - f"median = {median_val:.3f}\n" - f"─────────────\n" - f"< 0.4: {n_low:,} ({100 * n_low / n_pdbs:.1f}%)\n" - f"0.4–0.8: {n_mid:,} ({100 * n_mid / n_pdbs:.1f}%)\n" - f"≥ 0.8: {n_high:,} ({100 * n_high / n_pdbs:.1f}%)" - ) - ax.text( - 0.02, - 0.98, - textstr, - transform=ax.transAxes, - fontsize=10, - verticalalignment="top", - bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), - ) - - ax.set_xlabel("Mean EDIAm Score", fontsize=12) - ax.set_ylabel("Number of PDBs", fontsize=12) - ax.set_title("Distribution of Mean Water EDIAm (Per PDB)", fontsize=14) - ax.legend(loc="upper right") - - plt.tight_layout() - fig.savefig(output_dir / "02_ediam_pdbs.png", dpi=150) - plt.close(fig) - logger.info("Saved: 02_ediam_pdbs.png") - - -def plot_rsccs_waters(df: pd.DataFrame, output_dir: Path): - """Plot histogram of RSCCS for all water molecules.""" - fig, ax = plt.subplots(figsize=(10, 6)) - - ax.hist(df["RSCCS"].dropna(), bins=50, edgecolor="black", alpha=0.7, color="teal") - - # Add statistics - n_total = df["RSCCS"].notna().sum() - mean_val = df["RSCCS"].mean() - median_val = df["RSCCS"].median() - textstr = f"n = {n_total:,}\nmean = {mean_val:.3f}\nmedian = {median_val:.3f}" - ax.text( - 0.02, - 0.98, - textstr, - transform=ax.transAxes, - fontsize=10, - verticalalignment="top", - bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), - ) - - ax.set_xlabel("RSCCS Score", fontsize=12) - ax.set_ylabel("Count", fontsize=12) - ax.set_title("Distribution of Water RSCCS Scores (All Waters)", fontsize=14) - - plt.tight_layout() - fig.savefig(output_dir / "03_rsccs_waters.png", dpi=150) - plt.close(fig) - logger.info("Saved: 03_rsccs_waters.png") - - -def plot_rsccs_pdbs(df: pd.DataFrame, output_dir: Path): - """Plot histogram of mean RSCCS per PDB.""" - pdb_means = df.groupby("pdb_id")["RSCCS"].mean() - - fig, ax = plt.subplots(figsize=(10, 6)) - - ax.hist(pdb_means.dropna(), bins=50, edgecolor="black", alpha=0.7, color="teal") - - # Add statistics - n_pdbs = len(pdb_means) - mean_val = pdb_means.mean() - median_val = pdb_means.median() - textstr = f"n = {n_pdbs:,} PDBs\nmean = {mean_val:.3f}\nmedian = {median_val:.3f}" - ax.text( - 0.02, - 0.98, - textstr, - transform=ax.transAxes, - fontsize=10, - verticalalignment="top", - bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), - ) - - ax.set_xlabel("Mean RSCCS Score", fontsize=12) - ax.set_ylabel("Number of PDBs", fontsize=12) - ax.set_title("Distribution of Mean Water RSCCS (Per PDB)", fontsize=14) - - plt.tight_layout() - fig.savefig(output_dir / "04_rsccs_pdbs.png", dpi=150) - plt.close(fig) - logger.info("Saved: 04_rsccs_pdbs.png") - - -def plot_bfactor_waters(df: pd.DataFrame, output_dir: Path): - """Plot histogram of normalized B-factor for all water molecules. - - B-factors are normalized per-PDB (z-score using whole-PDB mean/std). - Shows cutoff at 5.0 (high B-factor = worse quality). - """ - bfactor_data = df["b_factor_normalized"].dropna() - - if len(bfactor_data) == 0: - logger.warning("Skipped: 05_bfactor_waters.png (no B-factor data)") - return - - fig, ax = plt.subplots(figsize=(10, 6)) - - ax.hist(bfactor_data, bins=50, edgecolor="black", alpha=0.7, color="coral") - - # Add cutoff lines at +1.5 and -1.5 - cutoff = 1.5 - ax.axvline( - x=cutoff, color="red", linestyle="--", linewidth=2, label=f"cutoff = ±{cutoff}" - ) - ax.axvline(x=-cutoff, color="red", linestyle="--", linewidth=2) - - # Add statistics - n_total = len(bfactor_data) - mean_val = bfactor_data.mean() - median_val = bfactor_data.median() - - # Count waters in each region - n_below = (bfactor_data < -cutoff).sum() - n_within = ((bfactor_data >= -cutoff) & (bfactor_data <= cutoff)).sum() - n_above = (bfactor_data > cutoff).sum() - - textstr = ( - f"n = {n_total:,}\n" - f"mean = {mean_val:.2f}\n" - f"median = {median_val:.2f}\n" - f"─────────────\n" - f"< -{cutoff}: {n_below:,} ({100 * n_below / n_total:.1f}%)\n" - f"-{cutoff} to {cutoff}: {n_within:,} ({100 * n_within / n_total:.1f}%)\n" - f"> {cutoff}: {n_above:,} ({100 * n_above / n_total:.1f}%)" - ) - ax.text( - 0.98, - 0.98, - textstr, - transform=ax.transAxes, - fontsize=10, - verticalalignment="top", - horizontalalignment="right", - bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), - ) - - ax.set_xlabel("Normalized B-factor (z-score)", fontsize=12) - ax.set_ylabel("Count", fontsize=12) - ax.set_title("Distribution of Normalized Water B-factors (All Waters)", fontsize=14) - ax.legend(loc="upper left") - - plt.tight_layout() - fig.savefig(output_dir / "05_bfactor_waters.png", dpi=150) - plt.close(fig) - logger.info("Saved: 05_bfactor_waters.png") - - -def plot_bfactor_pdbs(df: pd.DataFrame, output_dir: Path): - """Plot histogram of std dev of normalized B-factor per PDB. - - B-factors are normalized per-PDB (z-score using whole-PDB mean/std). - Shows variation in water B-factors within each structure. - """ - # Filter to rows with B-factor data - df_with_bfactor = df[df["b_factor_normalized"].notna()] - - if len(df_with_bfactor) == 0: - logger.warning("Skipped: 06_bfactor_pdbs.png (no B-factor data)") - return - - pdb_stds = df_with_bfactor.groupby("pdb_id")["b_factor_normalized"].std() - - fig, ax = plt.subplots(figsize=(10, 6)) - - ax.hist(pdb_stds, bins=50, edgecolor="black", alpha=0.7, color="coral") - - # Add cutoff line for high variability - cutoff = 1.5 - ax.axvline( - x=cutoff, color="red", linestyle="--", linewidth=2, label=f"cutoff = {cutoff}" - ) - - # Add statistics - n_pdbs = len(pdb_stds) - mean_val = pdb_stds.mean() - median_val = pdb_stds.median() - - # Count PDBs in each region - n_below = (pdb_stds <= cutoff).sum() - n_above = (pdb_stds > cutoff).sum() - - textstr = ( - f"n = {n_pdbs:,} PDBs\n" - f"mean = {mean_val:.2f}\n" - f"median = {median_val:.2f}\n" - f"─────────────\n" - f"≤ {cutoff}: {n_below:,} ({100 * n_below / n_pdbs:.1f}%)\n" - f"> {cutoff}: {n_above:,} ({100 * n_above / n_pdbs:.1f}%)" - ) - ax.text( - 0.98, - 0.98, - textstr, - transform=ax.transAxes, - fontsize=10, - verticalalignment="top", - horizontalalignment="right", - bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), - ) - - ax.set_xlabel("Std Dev of Normalized B-factor (z-score)", fontsize=12) - ax.set_ylabel("Number of PDBs", fontsize=12) - ax.set_title("Distribution of Water B-factor Variability (Per PDB)", fontsize=14) - ax.legend(loc="upper left") - - plt.tight_layout() - fig.savefig(output_dir / "06_bfactor_pdbs.png", dpi=150) - plt.close(fig) - logger.info("Saved: 06_bfactor_pdbs.png") - - -def plot_ediam_bfactor_correlation(df: pd.DataFrame, output_dir: Path): - """Plot EDIA vs normalized B-factor scatter with hexbin overlay.""" - # Filter to rows with both EDIA and normalized B-factor data - df_valid = df[df["EDIAm"].notna() & df["b_factor_normalized"].notna()] - - if len(df_valid) == 0: - logger.warning("Skipped: 07_ediam_bfactor_correlation.png (no matched data)") - return - - fig, axes = plt.subplots(1, 2, figsize=(14, 6)) - - # Left: Scatter plot - ax1 = axes[0] - ax1.scatter( - df_valid["b_factor_normalized"], - df_valid["EDIAm"], - alpha=0.1, - s=5, - c="steelblue", - ) - - # Add correlation coefficient - corr = df_valid["EDIAm"].corr(df_valid["b_factor_normalized"]) - ax1.text( - 0.02, - 0.98, - f"r = {corr:.3f}\nn = {len(df_valid):,}", - transform=ax1.transAxes, - fontsize=12, - verticalalignment="top", - bbox=dict(boxstyle="round", facecolor="wheat", alpha=0.5), - ) - - ax1.set_xlabel("Normalized B-factor (z-score)", fontsize=12) - ax1.set_ylabel("EDIAm Score", fontsize=12) - ax1.set_title("EDIAm vs Normalized B-factor (Scatter)", fontsize=14) - - # Right: Hexbin density plot - ax2 = axes[1] - hb = ax2.hexbin( - df_valid["b_factor_normalized"], - df_valid["EDIAm"], - gridsize=50, - cmap="YlOrRd", - mincnt=1, - ) - fig.colorbar(hb, ax=ax2, label="Count") - - ax2.set_xlabel("Normalized B-factor (z-score)", fontsize=12) - ax2.set_ylabel("EDIAm Score", fontsize=12) - ax2.set_title("EDIAm vs Normalized B-factor (Density)", fontsize=14) - - plt.tight_layout() - fig.savefig(output_dir / "07_ediam_bfactor_correlation.png", dpi=150) - plt.close(fig) - logger.info("Saved: 07_ediam_bfactor_correlation.png") - - -def get_pdb_ids_from_directory(pdb_dir: Path) -> list[str]: - """Get list of PDB IDs by scanning the PDB directory structure. - - Assumes PDB files are organized as pdb_dir/pdb_id/pdb_id_final.pdb - """ - pdb_ids = [] - for subdir in pdb_dir.iterdir(): - if subdir.is_dir(): - pdb_file = subdir / f"{subdir.name}_final.pdb" - if pdb_file.exists(): - pdb_ids.append(subdir.name) - logger.info(f"Found {len(pdb_ids)} PDB IDs in directory") - return pdb_ids - - -def get_pdb_ids_from_file(pdb_list_file: Path) -> list[str]: - """Get list of PDB IDs from a text file. - - Expects each line to be in format '_final'. - Strips the '_final' suffix to return just the pdb_id. - """ - pdb_ids = [] - with open(pdb_list_file) as f: - for line in f: - line = line.strip() - if line: - # Strip '_final' suffix if present - pdb_id = line.replace("_final", "") - pdb_ids.append(pdb_id) - logger.info(f"Loaded {len(pdb_ids)} PDB IDs from {pdb_list_file}") - return pdb_ids - - -def main(): - parser = argparse.ArgumentParser( - description="Generate whole-dataset distribution and correlation plots for water quality metrics" - ) - parser.add_argument( - "--edia-dir", - type=Path, - default=Path("/sb/wankowicz_lab/data/srivasv/edia_results"), - help="Directory containing EDIA CSV files", - ) - parser.add_argument( - "--pdb-dir", - type=Path, - default=Path("/sb/wankowicz_lab/data/srivasv/pdb_redo_data"), - help="Directory containing PDB files", - ) - parser.add_argument( - "--output-dir", - type=Path, - default=Path("figures/water_quality"), - help="Directory to save output figures", - ) - parser.add_argument( - "--num-workers", - type=int, - default=4, - help="Number of parallel workers for B-factor extraction", - ) - parser.add_argument( - "--skip-bfactor", - action="store_true", - help="Skip B-factor extraction (generate only EDIA/RSCC plots)", - ) - parser.add_argument( - "--bfactor-normalization", - type=str, - choices=["all", "protein", "water"], - default="all", - help="B-factor normalization strategy: 'all' (all atoms), 'protein' (protein atoms only), 'water' (water atoms only). Default: all", - ) - parser.add_argument( - "--bfactor-only", - action="store_true", - help="Generate only B-factor plots (skip EDIA/RSCC plots)", - ) - parser.add_argument( - "--pdb-list", - type=Path, - default=Path("splits/water_pdbs.txt"), - help="Text file with PDB IDs (one per line, format: _final). Used with --bfactor-only.", - ) - args = parser.parse_args() - - # Create output directory - args.output_dir.mkdir(parents=True, exist_ok=True) - - # Determine what data we need based on flags - need_edia = not args.bfactor_only - need_bfactor = not args.skip_bfactor - - df = None - bfactor_df = None - - # Load EDIA water data only if needed - if need_edia: - logger.info("Loading EDIA water data...") - df = load_all_water_data(args.edia_dir) - - # Extract B-factors if needed - if need_bfactor: - logger.info( - f"\nExtracting B-factors from PDB files (normalization: {args.bfactor_normalization})..." - ) - - if args.bfactor_only: - # Get PDB IDs from text file - pdb_ids = get_pdb_ids_from_file(args.pdb_list) - else: - # Get PDB IDs from EDIA data - pdb_ids = df["pdb_id"].unique().tolist() - - bfactor_df = load_all_bfactors( - args.pdb_dir, - pdb_ids, - args.num_workers, - normalization=args.bfactor_normalization, - ) - - # Merge with EDIA data if both are available - if df is not None: - logger.info("\nMerging EDIA and B-factor data...") - df = merge_edia_with_bfactors(df, bfactor_df) - - # Generate plots - logger.info("\nGenerating plots...") - - # EDIA/RSCC plots (skip if --bfactor-only) - if need_edia and df is not None: - plot_ediam_waters(df, args.output_dir) - plot_ediam_pdbs(df, args.output_dir) - plot_rsccs_waters(df, args.output_dir) - plot_rsccs_pdbs(df, args.output_dir) - - # B-factor plots - if need_bfactor and bfactor_df is not None: - if args.bfactor_only: - # Use bfactor_df directly when no EDIA data - plot_bfactor_waters(bfactor_df, args.output_dir) - plot_bfactor_pdbs(bfactor_df, args.output_dir) - else: - # Use merged df when EDIA data is available - plot_bfactor_waters(df, args.output_dir) - plot_bfactor_pdbs(df, args.output_dir) - plot_ediam_bfactor_correlation(df, args.output_dir) - - logger.info(f"\nAll figures saved to: {args.output_dir}") - - -if __name__ == "__main__": - main() From db16593d03938574c04189277e3f3995f1ea04f1 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Wed, 18 Mar 2026 12:19:10 -0500 Subject: [PATCH 13/19] addresssing review comments --- .github/workflows/build.yml | 6 ++ pyproject.toml | 1 - src/dataset.py | 10 ++-- src/encoder_base.py | 4 +- src/gvp.py | 17 +++--- src/gvp_encoder.py | 16 +++--- src/utils.py | 16 +++--- tests/conftest.py | 24 ++++++++ tests/test_dataset.py | 12 +--- tests/test_encoder.py | 39 +++---------- tests/test_flow.py | 46 ++++++++-------- tests/test_forward.py | 41 ++++++-------- tests/test_utils.py | 107 ++++++++++++++++++------------------ 13 files changed, 164 insertions(+), 175 deletions(-) diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index 4d5941d..9d0a7ab 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -5,6 +5,12 @@ on: branches: [main] pull_request: branches: [main] + paths-ignore: + - '**.md' + - 'docs/**' + - 'figures/**' + - '.gitignore' + - 'LICENSE' workflow_dispatch: jobs: diff --git a/pyproject.toml b/pyproject.toml index ef17725..19f08d4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -231,7 +231,6 @@ warn_return_any = true warn_unused_configs = true [tool.ty.rules] -unresolved-import = "ignore" # PyTorch Geometric has incomplete type stubs - dynamic attributes on Batch/Data unresolved-attribute = "ignore" # PyG's MessagePassing.message() override pattern is intentional diff --git a/src/dataset.py b/src/dataset.py index 1bae99b..4f57a77 100644 --- a/src/dataset.py +++ b/src/dataset.py @@ -102,13 +102,11 @@ def get_crystal_contacts_pymol(pdb_path: str, cutoff: float = 5.0) -> dict: asu_atoms = cmd.get_model(obj, state=1).atom mate_atoms = cmd.get_model("sym* and interface", state=1).atom + asu_coords = asu_coords if asu_coords is not None else np.zeros((0, 3), dtype=float) + mate_coords = mate_coords if mate_coords is not None else np.zeros((0, 3), dtype=float) return { - "asu_coords": asu_coords - if asu_coords is not None - else np.zeros((0, 3), dtype=float), - "mate_coords": mate_coords - if mate_coords is not None - else np.zeros((0, 3), dtype=float), + "asu_coords": asu_coords, + "mate_coords": mate_coords, "asu_atoms": asu_atoms, "mate_atoms": mate_atoms, } diff --git a/src/encoder_base.py b/src/encoder_base.py index 49786d1..333ca69 100644 --- a/src/encoder_base.py +++ b/src/encoder_base.py @@ -35,7 +35,7 @@ class MyEncoder(BaseProteinEncoder): def decorator(cls: type[BaseProteinEncoder]) -> type[BaseProteinEncoder]: if name in _ENCODER_REGISTRY: - raise ValueError(f"Encoder '{name}' is already registered") + raise KeyError(f"Encoder '{name}' is already registered") _ENCODER_REGISTRY[name] = cls return cls @@ -75,7 +75,7 @@ def build_encoder(config: dict, device: torch.device) -> BaseProteinEncoder: Instantiated encoder implementing BaseProteinEncoder """ if "encoder_type" not in config: - raise ValueError("'encoder_type' must be specified in config") + raise KeyError("'encoder_type' must be specified in config") encoder_type = config["encoder_type"] encoder_cls = get_encoder_class(encoder_type) return encoder_cls.from_config(config, device) diff --git a/src/gvp.py b/src/gvp.py index 96e80d8..b4bf384 100644 --- a/src/gvp.py +++ b/src/gvp.py @@ -474,9 +474,10 @@ def forward( s_node, _ = node_tuple s_edge, V_edge = edge_attr - assert s_edge.shape[-1] == self.s_edge_width, ( - f"EdgeUpdate expected width {self.s_edge_width}, got {s_edge.shape[-1]}" - ) + if s_edge.shape[-1] != self.s_edge_width: + raise ValueError( + f"EdgeUpdate expected width {self.s_edge_width}, got {s_edge.shape[-1]}" + ) src, dst = edge_index[0], edge_index[1] parts = [s_node[src], s_node[dst], s_edge] @@ -524,19 +525,19 @@ def __init__( # message GVP stack; first layer takes [unit_vec] and [rbf] extras msg_layers = [] for i in range(n_message_gvps): - vin = ( + vector_input_dim = ( v_dim - + (1 if i == 0 else 0) + + (1 if i == 0 else 0) # +1 for unit displacement vector on first layer + (v_dim if (i == 0 and use_dst_feats) else 0) ) - sin = ( + scalar_input_dim = ( s_dim + (rbf_dim if i == 0 else 0) + (s_dim if (i == 0 and use_dst_feats) else 0) ) msg_layers.append( GVP_( - in_dims=(sin, vin), + in_dims=(scalar_input_dim, vector_input_dim), out_dims=(s_dim, v_dim), vector_gate=True, activations=activations, @@ -625,7 +626,7 @@ class GVPMultiEdgeConv(nn.Module): def __init__( self, - etypes, # List[EdgeType] + etypes: list[tuple[str, str, str]], # (src_type, edge_type, dst_type) s_dim: int, v_dim: int, rbf_dim: int = 16, diff --git a/src/gvp_encoder.py b/src/gvp_encoder.py index 7aed3d4..87b229f 100644 --- a/src/gvp_encoder.py +++ b/src/gvp_encoder.py @@ -77,8 +77,8 @@ def make_gvp_encoder_data(data: HeteroData) -> Data: pp_edge = data[EDGE_PP] if hasattr(pp_edge, "edge_rbf"): enc_data.edge_rbf = pp_edge.edge_rbf - if hasattr(pp_edge, "edge_unit"): - enc_data.edge_unit = pp_edge.edge_unit + if hasattr(pp_edge, "edge_unit_vectors"): + enc_data.edge_unit_vectors = pp_edge.edge_unit_vectors # batch for multi-complex batches if hasattr(prot, "batch"): @@ -271,7 +271,7 @@ def _pool_by_residue( ) return out else: - raise ValueError(f"Unknown pool_aggr={aggr!r}") + raise ValueError(f"Unknown pool_aggr={aggr}") @staticmethod def _initial_node_tuple( @@ -286,20 +286,20 @@ def _compute_edge_attr(self, data: Batch): """ Build edge attributes from positions or cached features. - If cached edge features (edge_rbf, edge_unit) are both available in data, + If cached edge features (edge_rbf, edge_unit_vectors) are both available in data, use them directly. Otherwise, compute from positions. Args: - data: Batch with pos, edge_index, and optionally edge_rbf, edge_unit + data: Batch with pos, edge_index, and optionally edge_rbf, edge_unit_vectors Returns: (s_edge, V_edge): Tuple of scalar and vector edge features s_edge_raw: Raw RBF features (for distance conditioning) """ # Use cached features if available - if hasattr(data, "edge_rbf") and hasattr(data, "edge_unit"): + if hasattr(data, "edge_rbf") and hasattr(data, "edge_unit_vectors"): s_edge_raw = data.edge_rbf - u = data.edge_unit + u = data.edge_unit_vectors else: # Fallback: compute from positions rij, u = edge_vectors(data.pos, data.edge_index) @@ -320,7 +320,7 @@ def forward(self, data: Batch) -> tuple[tuple, tuple | None]: - edge_index: (2, E) edge indices Optional cached edge features (if absent, computed from pos): - edge_rbf: (E, num_rbf) RBF distance features - - edge_unit: (E, 3) unit edge vectors + - edge_unit_vectors: (E, 3) unit edge vectors Returns: x: tuple (s, V) of node scalar and vector features diff --git a/src/utils.py b/src/utils.py index 51fa714..82d77d9 100644 --- a/src/utils.py +++ b/src/utils.py @@ -221,7 +221,7 @@ def ot_coupling( @torch.no_grad() def recall_precision( pred: torch.Tensor | np.ndarray, - true: torch.Tensor | np.ndarray, + ground_truth: torch.Tensor | np.ndarray, thresh: float = 1.0, ) -> tuple[float, float]: """ @@ -231,7 +231,7 @@ def recall_precision( Args: pred: (N_pred, 3) predicted positions - true: (N_true, 3) ground truth positions + ground_truth: (N_true, 3) ground truth positions thresh: distance threshold in Angstroms Returns: @@ -241,18 +241,18 @@ def recall_precision( # convert numpy arrays to tensors first if isinstance(pred, np.ndarray): pred = torch.from_numpy(pred) - if isinstance(true, np.ndarray): - true = torch.from_numpy(true) + if isinstance(ground_truth, np.ndarray): + ground_truth = torch.from_numpy(ground_truth) # handle empty inputs - if pred.numel() == 0 or true.numel() == 0: + if pred.numel() == 0 or ground_truth.numel() == 0: return 0.0, 0.0 # ensure same device - if pred.device != true.device: - true = true.to(pred.device) + if pred.device != ground_truth.device: + ground_truth = ground_truth.to(pred.device) - D = torch.cdist(true.float(), pred.float(), p=2) + D = torch.cdist(ground_truth.float(), pred.float(), p=2) recall = (D.min(dim=1)[0] <= thresh).float().mean().item() precision = (D.min(dim=0)[0] <= thresh).float().mean().item() diff --git a/tests/conftest.py b/tests/conftest.py index b1b2b42..e16064c 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -51,3 +51,27 @@ def pdb_8dzt(): def pdb_1deu(): """1deu - has insertion codes (52 residues with ins_code='P').""" return _resolve_pdb_path("1deu") + + +# ============== Shared encoder fixtures ============== + + +@pytest.fixture +def base_encoder(device): + """Base ProteinGVPEncoder for flow model tests.""" + from src.gvp_encoder import ProteinGVPEncoder + + return ProteinGVPEncoder( + node_scalar_in=16, + hidden_dims=(64, 8), + n_edge_scalar_in=16, + pool_residue=False, + ).to(device) + + +@pytest.fixture +def gvp_encoder(base_encoder): + """Wrapped GVPEncoder for flow model tests.""" + from src.gvp_encoder import GVPEncoder + + return GVPEncoder(encoder=base_encoder, freeze=False) diff --git a/tests/test_dataset.py b/tests/test_dataset.py index 3cf2c2c..a1af9bc 100644 --- a/tests/test_dataset.py +++ b/tests/test_dataset.py @@ -412,11 +412,8 @@ def test_interacting_chains_pass(self): atoms = bts.AtomArray(20) atoms.chain_id = np.array(["A"] * 10 + ["B"] * 10) # Place chains close together - atoms.coord = np.zeros((20, 3)) - atoms.coord[:10] = np.random.randn(10, 3) - atoms.coord[10:] = np.random.randn(10, 3) + np.array( - [2.0, 0.0, 0.0] - ) # 2A offset + atoms.coord = np.random.randn(20, 3) + atoms.coord[10:] += np.array([2.0, 0.0, 0.0]) is_valid, reason, status = check_chain_interactions( atoms, interface_dist_threshold=4.0 @@ -433,10 +430,7 @@ def test_non_interacting_chains_fail(self): atoms.chain_id = np.array(["A"] * 10 + ["B"] * 10) # Place chains far apart atoms.coord = np.zeros((20, 3)) - atoms.coord[:10] = np.zeros((10, 3)) - atoms.coord[10:] = np.zeros((10, 3)) + np.array( - [100.0, 0.0, 0.0] - ) # 100A offset + atoms.coord[10:] += np.array([100.0, 0.0, 0.0]) is_valid, reason, status = check_chain_interactions( atoms, interface_dist_threshold=4.0 diff --git a/tests/test_encoder.py b/tests/test_encoder.py index 872f246..4925384 100644 --- a/tests/test_encoder.py +++ b/tests/test_encoder.py @@ -210,8 +210,6 @@ def test_cached_embedding_implements_interface( assert pp_edge_attr is None # output_dims available after forward - assert isinstance(encoder.output_dims, tuple) - assert len(encoder.output_dims) == 2 assert encoder.output_dims == (128, 0) def test_from_config_class_method(self, device): @@ -283,8 +281,13 @@ def test_encoder_forward_no_pooling(self, sample_homogeneous_data): assert s.shape == (sample_homogeneous_data.num_nodes, 64) assert v.shape == (sample_homogeneous_data.num_nodes, 16, 3) - # edge_attr should be a tuple (s_edge, V_edge) - assert edge_attr is not None + # edge_attr should be (s_edge, V_edge) tuple when using GVP encoder with + # pool_residue=False and use_edge_update=True; it's None for: + # (1) cached embedding encoders (ESM/SLAE), (2) pooled outputs, or (3) edge updates disabled + assert edge_attr is not None, ( + "edge_attr should be (s_edge, V_edge) tuple when using GVP encoder with " + "pool_residue=False and use_edge_update=True" + ) s_edge, V_edge = edge_attr assert s_edge.dim() == 2 assert V_edge.dim() == 3 @@ -384,14 +387,6 @@ def test_output_dims_before_forward_raises(self, device): with pytest.raises(RuntimeError, match="dimension not yet known"): _ = encoder.output_dims - def test_slae_output_dims_after_forward(self, device, sample_hetero_data_with_slae): - """SLAE encoder should infer output_dims from data.""" - encoder = CachedEmbeddingEncoder( - embedding_key="slae_embedding", encoder_type="slae" - ).to(device) - encoder(sample_hetero_data_with_slae) - assert encoder.output_dims == (128, 0) - def test_esm_output_dims_after_forward(self, device, sample_hetero_data): """ESM encoder should infer output_dims from data.""" encoder = CachedEmbeddingEncoder( @@ -493,26 +488,6 @@ def test_encoder_no_learnable_params(self, device): ).to(device) assert sum(p.numel() for p in encoder.parameters()) == 0 - def test_slae_from_config(self, device, sample_hetero_data_with_slae): - """Should construct SLAE from config and infer dim from data.""" - config = {"encoder_type": "slae"} - encoder = CachedEmbeddingEncoder.from_config(config, device) - assert encoder.encoder_type == "slae" - encoder(sample_hetero_data_with_slae) - assert encoder.output_dims == (128, 0) - - def test_esm_from_config(self, device, sample_hetero_data): - """Should construct ESM from config and infer dim from data.""" - config = {"encoder_type": "esm"} - encoder = CachedEmbeddingEncoder.from_config(config, device) - assert encoder.encoder_type == "esm" - n_atoms = sample_hetero_data["protein"].num_nodes - sample_hetero_data["protein"].esm_embedding = torch.randn( - n_atoms, 2048, device=device - ) - encoder(sample_hetero_data) - assert encoder.output_dims == (2048, 0) - def test_device_placement(self, device, sample_hetero_data): """Verify tensors are on the correct device.""" encoder = CachedEmbeddingEncoder( diff --git a/tests/test_flow.py b/tests/test_flow.py index c1daee8..f7fc9c6 100644 --- a/tests/test_flow.py +++ b/tests/test_flow.py @@ -8,6 +8,7 @@ import numpy as np import pytest import torch +import torch.nn.functional as F from torch_geometric.data import Data, HeteroData from src.flow import ( @@ -49,14 +50,18 @@ def batched_hetero_data(device): # Protein: 20 atoms (10 per graph) data["protein"].pos = torch.randn(20, 3, device=device) - data["protein"].x = torch.randn(20, 16, device=device) + # One-hot encoded element types (16 classes: 15 elements + 1 "other") + protein_elem_indices = torch.randint(0, 16, (20,), device=device) + data["protein"].x = F.one_hot(protein_elem_indices, num_classes=16).float() data["protein"].batch = torch.cat( [torch.zeros(10, dtype=torch.long), torch.ones(10, dtype=torch.long)] ).to(device) # Water: 8 molecules (4 per graph) data["water"].pos = torch.randn(8, 3, device=device) - data["water"].x = torch.randn(8, 16, device=device) + # One-hot encoded element types (water is oxygen, index 2 in ELEMENT_VOCAB) + water_elem_indices = torch.full((8,), 2, dtype=torch.long, device=device) + data["water"].x = F.one_hot(water_elem_indices, num_classes=16).float() data["water"].batch = torch.cat( [torch.zeros(4, dtype=torch.long), torch.ones(4, dtype=torch.long)] ).to(device) @@ -88,6 +93,9 @@ def mock_forward(data): return encoder +# base_encoder and gvp_encoder fixtures are defined in conftest.py + + @pytest.mark.unit class TestBuildKnnEdges: def test_basic_knn(self, device): @@ -99,9 +107,17 @@ def test_basic_knn(self, device): edges = build_knn_edges(src, dst, k=2) assert edges.shape[0] == 2 - assert edges.shape[1] > 0 + assert edges.shape[1] >= 4 # At least 2 dst points × 2 neighbors each assert edges.dtype == torch.long + # dst[0] at 0.5 should connect to src[0] (dist=0.5) and src[1] (dist=0.5) + # dst[1] at 1.5 should connect to src[1] (dist=0.5) and src[2] (dist=0.5) + edge_set = set(zip(edges[0].tolist(), edges[1].tolist())) + assert (0, 0) in edge_set, f"Missing edge src[0]->dst[0], got {edge_set}" + assert (1, 0) in edge_set, f"Missing edge src[1]->dst[0], got {edge_set}" + assert (1, 1) in edge_set, f"Missing edge src[1]->dst[1], got {edge_set}" + assert (2, 1) in edge_set, f"Missing edge src[2]->dst[1], got {edge_set}" + def test_empty_src(self, device): src = torch.empty(0, 3, device=device) dst = torch.randn(5, 3, device=device) @@ -309,17 +325,9 @@ def test_forward_no_water(self, device): assert v_pred.shape == (0, 3) - def test_self_conditioning(self, simple_hetero_data, device): - base_encoder = ProteinGVPEncoder( - node_scalar_in=16, - hidden_dims=(64, 8), - n_edge_scalar_in=16, - pool_residue=False, - ).to(device) - encoder = GVPEncoder(encoder=base_encoder, freeze=False) - + def test_self_conditioning(self, simple_hetero_data, device, gvp_encoder): model = FlowWaterGVP( - encoder=encoder, + encoder=gvp_encoder, hidden_dims=(64, 8), layers=1, ).to(device) @@ -339,17 +347,9 @@ def test_self_conditioning(self, simple_hetero_data, device): @pytest.mark.unit class TestFlowMatcher: @pytest.fixture - def flow_matcher(self, device): - base_encoder = ProteinGVPEncoder( - node_scalar_in=16, - hidden_dims=(64, 8), - n_edge_scalar_in=16, - pool_residue=False, - ).to(device) - encoder = GVPEncoder(encoder=base_encoder, freeze=False) - + def flow_matcher(self, device, gvp_encoder): model = FlowWaterGVP( - encoder=encoder, + encoder=gvp_encoder, hidden_dims=(64, 8), layers=1, ).to(device) diff --git a/tests/test_forward.py b/tests/test_forward.py index ed476e2..d708379 100644 --- a/tests/test_forward.py +++ b/tests/test_forward.py @@ -4,6 +4,7 @@ All test cases created with assistance from Claude Code and refined. """ +import math import os import pytest @@ -209,6 +210,8 @@ def test_forward_pass_no_nan_with_module_hooks(device): t = torch.linspace(0.05, 0.95, steps=n_graphs, device=device) + # FiniteHookManager registers forward hooks on modules to catch NaN/Inf + # outputs early. This helps identify exactly which layer produces invalid values. with FiniteHookManager() as hm: # Encoder internals (access via .encoder.encoder for wrapped GVPEncoder) hm.watch( @@ -287,7 +290,7 @@ def test_training_step_no_nan_tripwire(device): ) loss = out["loss"] assert isinstance(loss, float) - assert loss == loss, "loss is NaN" + assert not math.isnan(loss), "loss is NaN" # ensure params stayed finite for name, p in model.named_parameters(): @@ -596,20 +599,12 @@ def test_interpolation_at_boundaries(self, device): class TestVelocityFieldProperties: """Test velocity field sanity checks.""" - def test_velocity_field_finite(self, device): + def test_velocity_field_finite(self, device, gvp_encoder): """Velocity predictions should be finite (no NaN or Inf).""" data = make_batched_hetero(device, n_graphs=1, n_protein_per=24, n_water_per=12) - base_encoder = ProteinGVPEncoder( - node_scalar_in=16, - hidden_dims=(64, 8), - n_edge_scalar_in=16, - pool_residue=False, - ).to(device) - encoder = GVPEncoder(encoder=base_encoder, freeze=False) - model = FlowWaterGVP( - encoder=encoder, + encoder=gvp_encoder, hidden_dims=(64, 8), layers=1, k_pw=8, @@ -631,20 +626,12 @@ def test_velocity_field_finite(self, device): f"Velocity magnitude too large at t={t_val}: {max_mag:.3e}" ) - def test_velocity_field_changes_with_t(self, device): + def test_velocity_field_changes_with_t(self, device, gvp_encoder): """Velocity field should depend on t (different outputs for different times).""" data = make_batched_hetero(device, n_graphs=1, n_protein_per=24, n_water_per=12) - base_encoder = ProteinGVPEncoder( - node_scalar_in=16, - hidden_dims=(64, 8), - n_edge_scalar_in=16, - pool_residue=False, - ).to(device) - encoder = GVPEncoder(encoder=base_encoder, freeze=False) - model = FlowWaterGVP( - encoder=encoder, + encoder=gvp_encoder, hidden_dims=(64, 8), layers=1, k_pw=8, @@ -659,10 +646,14 @@ def test_velocity_field_changes_with_t(self, device): v0 = model(data, t0, sc=None) v1 = model(data, t1, sc=None) - # Velocities should be different - diff = torch.norm(v0 - v1, dim=-1).mean().item() - - assert diff > 1e-4, f"Velocity field doesn't change with t (diff={diff:.6f})" + # Velocities should be different - check that MOST individual waters show change + per_water_diff = torch.norm(v0 - v1, dim=-1) # per-water norm + n_different = (per_water_diff > 1e-5).sum().item() + n_water = v0.shape[0] + assert n_different >= n_water * 0.9, ( + f"Velocity field doesn't change with t for most waters " + f"({n_different}/{n_water} changed)" + ) @pytest.mark.unit diff --git a/tests/test_utils.py b/tests/test_utils.py index 8d3446c..c6d85d8 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -536,56 +536,57 @@ def test_different_sizes(self, tmp_path): assert (tmp_path / f"step_{n}.png").exists() -@pytest.mark.unit -class TestCreateTrajectoryGif: - """Tests for GIF creation from trajectory.""" - - def test_creates_gif(self, tmp_path): - """GIF should be created from trajectory.""" - trajectory = [np.random.rand(5, 3) for _ in range(10)] - protein_pos = np.random.rand(20, 3) - water_true = np.random.rand(5, 3) - - gif_path = str(tmp_path / "test.gif") - create_trajectory_gif( - trajectory=trajectory, - protein_pos=protein_pos, - water_true=water_true, - save_path=gif_path, - fps=5, - ) - - assert Path(gif_path).exists() - - def test_with_pdb_id(self, tmp_path): - """GIF should work with pdb_id parameter.""" - trajectory = [np.random.rand(3, 3) for _ in range(5)] - protein_pos = np.random.rand(10, 3) - water_true = np.random.rand(3, 3) - - gif_path = str(tmp_path / "test_pdb.gif") - create_trajectory_gif( - trajectory=trajectory, - protein_pos=protein_pos, - water_true=water_true, - save_path=gif_path, - pdb_id="1ABC", - ) - - assert Path(gif_path).exists() - - def test_long_trajectory_sampled(self, tmp_path): - """Long trajectories should be sampled to max 100 frames.""" - trajectory = [np.random.rand(3, 3) for _ in range(200)] - protein_pos = np.random.rand(10, 3) - water_true = np.random.rand(3, 3) - - gif_path = str(tmp_path / "long.gif") - create_trajectory_gif( - trajectory=trajectory, - protein_pos=protein_pos, - water_true=water_true, - save_path=gif_path, - ) - - assert Path(gif_path).exists() +# commenting the test below out as gif creation is just a viz tool and this test takes too long to run +# @pytest.mark.unit +# class TestCreateTrajectoryGif: +# """Tests for GIF creation from trajectory.""" + +# def test_creates_gif(self, tmp_path): +# """GIF should be created from trajectory.""" +# trajectory = [np.random.rand(5, 3) for _ in range(10)] +# protein_pos = np.random.rand(20, 3) +# water_true = np.random.rand(5, 3) + +# gif_path = str(tmp_path / "test.gif") +# create_trajectory_gif( +# trajectory=trajectory, +# protein_pos=protein_pos, +# water_true=water_true, +# save_path=gif_path, +# fps=5, +# ) + +# assert Path(gif_path).exists() + +# def test_with_pdb_id(self, tmp_path): +# """GIF should work with pdb_id parameter.""" +# trajectory = [np.random.rand(3, 3) for _ in range(5)] +# protein_pos = np.random.rand(10, 3) +# water_true = np.random.rand(3, 3) + +# gif_path = str(tmp_path / "test_pdb.gif") +# create_trajectory_gif( +# trajectory=trajectory, +# protein_pos=protein_pos, +# water_true=water_true, +# save_path=gif_path, +# pdb_id="1ABC", +# ) + +# assert Path(gif_path).exists() + +# def test_long_trajectory_sampled(self, tmp_path): +# """Long trajectories should be sampled to max 100 frames.""" +# trajectory = [np.random.rand(3, 3) for _ in range(200)] +# protein_pos = np.random.rand(10, 3) +# water_true = np.random.rand(3, 3) + +# gif_path = str(tmp_path / "long.gif") +# create_trajectory_gif( +# trajectory=trajectory, +# protein_pos=protein_pos, +# water_true=water_true, +# save_path=gif_path, +# ) + +# assert Path(gif_path).exists() From 99a962c38a3d43b4b26380ec00e740f0fe8acc2a Mon Sep 17 00:00:00 2001 From: vratins <114123331+vratins@users.noreply.github.com> Date: Wed, 18 Mar 2026 17:19:36 +0000 Subject: [PATCH 14/19] Auto-commit ruff fixes [skip ci] --- src/dataset.py | 8 ++++++-- tests/test_utils.py | 3 --- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/src/dataset.py b/src/dataset.py index 4f57a77..f86784c 100644 --- a/src/dataset.py +++ b/src/dataset.py @@ -102,8 +102,12 @@ def get_crystal_contacts_pymol(pdb_path: str, cutoff: float = 5.0) -> dict: asu_atoms = cmd.get_model(obj, state=1).atom mate_atoms = cmd.get_model("sym* and interface", state=1).atom - asu_coords = asu_coords if asu_coords is not None else np.zeros((0, 3), dtype=float) - mate_coords = mate_coords if mate_coords is not None else np.zeros((0, 3), dtype=float) + asu_coords = ( + asu_coords if asu_coords is not None else np.zeros((0, 3), dtype=float) + ) + mate_coords = ( + mate_coords if mate_coords is not None else np.zeros((0, 3), dtype=float) + ) return { "asu_coords": asu_coords, "mate_coords": mate_coords, diff --git a/tests/test_utils.py b/tests/test_utils.py index c6d85d8..9b70655 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -20,15 +20,12 @@ matplotlib.use("Agg") -from pathlib import Path from src.utils import ( ATOM37_FILL, atom37_to_atoms, compute_placement_metrics, compute_rmsd, - create_trajectory_gif, - # Optimal transport ot_coupling, # Visualization plot_3d_frame, From 03546b2f98ef42ebf78b0fefa778547649539d79 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Wed, 18 Mar 2026 13:26:00 -0500 Subject: [PATCH 15/19] fixing gvp encoder forward return type signature, making sure uv venv is created for ty check --- .github/workflows/lint.yml | 7 +++++-- src/gvp_encoder.py | 21 +++++++++++++++++++-- 2 files changed, 24 insertions(+), 4 deletions(-) diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml index 08862ed..2e210d5 100644 --- a/.github/workflows/lint.yml +++ b/.github/workflows/lint.yml @@ -37,9 +37,12 @@ jobs: steps: - uses: actions/checkout@v6 - + - name: Setup UV and python version uses: astral-sh/setup-uv@v7 + - name: Install dependencies + run: uv sync + - name: Run ty - run: uvx ty check src/ + run: uv run ty check src/ diff --git a/src/gvp_encoder.py b/src/gvp_encoder.py index 87b229f..be396e7 100644 --- a/src/gvp_encoder.py +++ b/src/gvp_encoder.py @@ -9,7 +9,7 @@ from collections.abc import Callable from pathlib import Path -from typing import Any, Literal +from typing import Any, Literal, TypeAlias import torch import torch.nn as nn @@ -23,6 +23,23 @@ from src.gvp import EdgeUpdate, GVP, GVPConvLayer from src.utils import rbf +# Type aliases for GVP feature tuples +GVPTuple: TypeAlias = tuple[torch.Tensor, torch.Tensor] +"""(scalar, vector) feature pair. Scalar: (N, dim), Vector: (N, dim, 3).""" + +NodeFeatures: TypeAlias = GVPTuple +"""Node (scalar, vector) features from GVP layers.""" + +EdgeAttr: TypeAlias = GVPTuple +"""Edge (scalar, vector) attributes.""" + +# Forward return type: either pooled residue embeddings or full GVP features +ForwardOutput: TypeAlias = tuple[NodeFeatures | torch.Tensor, EdgeAttr | None] +"""Return type of forward(): +- When pool_residue=True: (residue_embed: Tensor, None) +- When pool_residue=False: (node_features: GVPTuple, edge_attr: GVPTuple | None) +""" + def edge_vectors( pos: torch.Tensor, edge_index: torch.Tensor @@ -309,7 +326,7 @@ def _compute_edge_attr(self, data: Batch): V_edge = u.unsqueeze(1) return (s_edge, V_edge), s_edge_raw - def forward(self, data: Batch) -> tuple[tuple, tuple | None]: + def forward(self, data: Batch) -> ForwardOutput: """ Forward pass through the GVP encoder. From 93399e6b30b74e5aa58bf32253c3f627e72dae4f Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Wed, 18 Mar 2026 13:29:51 -0500 Subject: [PATCH 16/19] fixing typealias bug for ty --- src/gvp_encoder.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/src/gvp_encoder.py b/src/gvp_encoder.py index be396e7..3d6c84f 100644 --- a/src/gvp_encoder.py +++ b/src/gvp_encoder.py @@ -9,7 +9,7 @@ from collections.abc import Callable from pathlib import Path -from typing import Any, Literal, TypeAlias +from typing import Any, Literal import torch import torch.nn as nn @@ -24,17 +24,17 @@ from src.utils import rbf # Type aliases for GVP feature tuples -GVPTuple: TypeAlias = tuple[torch.Tensor, torch.Tensor] +type GVPTuple = tuple[torch.Tensor, torch.Tensor] """(scalar, vector) feature pair. Scalar: (N, dim), Vector: (N, dim, 3).""" -NodeFeatures: TypeAlias = GVPTuple +type NodeFeatures = GVPTuple """Node (scalar, vector) features from GVP layers.""" -EdgeAttr: TypeAlias = GVPTuple +type EdgeAttr = GVPTuple """Edge (scalar, vector) attributes.""" # Forward return type: either pooled residue embeddings or full GVP features -ForwardOutput: TypeAlias = tuple[NodeFeatures | torch.Tensor, EdgeAttr | None] +type ForwardOutput = tuple[NodeFeatures | torch.Tensor, EdgeAttr | None] """Return type of forward(): - When pool_residue=True: (residue_embed: Tensor, None) - When pool_residue=False: (node_features: GVPTuple, edge_attr: GVPTuple | None) From 4f7e0149bc6caf4c7e5a2b55eb51b4e184a8a83a Mon Sep 17 00:00:00 2001 From: vratins <114123331+vratins@users.noreply.github.com> Date: Wed, 18 Mar 2026 18:30:12 +0000 Subject: [PATCH 17/19] Auto-commit ruff fixes [skip ci] --- src/gvp_encoder.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/gvp_encoder.py b/src/gvp_encoder.py index 3d6c84f..82000f6 100644 --- a/src/gvp_encoder.py +++ b/src/gvp_encoder.py @@ -23,6 +23,7 @@ from src.gvp import EdgeUpdate, GVP, GVPConvLayer from src.utils import rbf + # Type aliases for GVP feature tuples type GVPTuple = tuple[torch.Tensor, torch.Tensor] """(scalar, vector) feature pair. Scalar: (N, dim), Vector: (N, dim, 3).""" From a11ad8beb05d54dc8eee80003a34fe9f51652057 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Wed, 18 Mar 2026 13:33:01 -0500 Subject: [PATCH 18/19] use uvx instead of uv to run ty --- .github/workflows/lint.yml | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml index 2e210d5..3f89bfb 100644 --- a/.github/workflows/lint.yml +++ b/.github/workflows/lint.yml @@ -41,8 +41,5 @@ jobs: - name: Setup UV and python version uses: astral-sh/setup-uv@v7 - - name: Install dependencies - run: uv sync - - name: Run ty - run: uv run ty check src/ + run: uvx ty check src/ From 5efc687ced582984e4ad28f633973c7c48bf2706 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Wed, 18 Mar 2026 13:39:38 -0500 Subject: [PATCH 19/19] pyproject --- pyproject.toml | 2 ++ 1 file changed, 2 insertions(+) diff --git a/pyproject.toml b/pyproject.toml index 19f08d4..7e08122 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -231,6 +231,8 @@ warn_return_any = true warn_unused_configs = true [tool.ty.rules] +# uvx runs ty in isolation without project dependencies - ignore missing imports +unresolved-import = "ignore" # PyTorch Geometric has incomplete type stubs - dynamic attributes on Batch/Data unresolved-attribute = "ignore" # PyG's MessagePassing.message() override pattern is intentional