From a49337917e14139f16b4b0ff2f9c9b5d7cb2fcb5 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Thu, 26 Feb 2026 16:09:57 -0600 Subject: [PATCH] adding graph creation constants, updating utils with constants --- src/constants.py | 18 ++--- src/utils.py | 172 ++++++++++++++++++++++++++++++++------------ tests/test_utils.py | 3 +- 3 files changed, 139 insertions(+), 54 deletions(-) diff --git a/src/constants.py b/src/constants.py index dbb7481..0033aac 100644 --- a/src/constants.py +++ b/src/constants.py @@ -1,16 +1,18 @@ # constants.py """ Constants for edge types and other shared values. - -Using constants instead of string literals helps prevent typos and enables -IDE autocompletion and refactoring support. """ +# feature dimensions +NUM_RBF = 16 # RBF basis functions for edge features +NODE_FEATURE_DIM = 16 # Element one-hot encoding dimension (15 elements + 1 "other") +RBF_CUTOFF = 8.0 # Distance cutoff for RBF encoding (Angstroms) + # 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 -EDGE_PW = ('protein', 'pw', 'water') # protein -> water -EDGE_WP = ('water', 'wp', 'protein') # water -> protein +EDGE_PP = ("protein", "pp", "protein") # protein -> protein +EDGE_WW = ("water", "ww", "water") # water -> water +EDGE_PW = ("protein", "pw", "water") # protein -> water +EDGE_WP = ("water", "wp", "protein") # water -> protein -# All edge types used in the model +# all edge types used in the model (future support: het atoms) ALL_EDGE_TYPES = [EDGE_PW, EDGE_WW, EDGE_PP, EDGE_WP] diff --git a/src/utils.py b/src/utils.py index fd0b12d..3e1fe3d 100644 --- a/src/utils.py +++ b/src/utils.py @@ -20,42 +20,67 @@ 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(): +def setup_logging_for_tqdm( + level: str = "INFO", + log_file: str | None = None, +): """ Configure loguru to work with tqdm progress bars. Redirects log output through tqdm.write() so log messages don't break progress bar rendering. Call this once at the start of scripts that use both logging and tqdm. + + Args: + level: Log level (DEBUG, INFO, WARNING, ERROR) + log_file: Optional path to log file for persistent logging """ + from pathlib import Path + from loguru import logger logger.remove() logger.add( lambda msg: tqdm.write(msg, end=""), + level=level.upper(), colorize=True, format="{level: <8} | {message}", ) + if log_file is not None: + log_path = Path(log_file) + log_path.parent.mkdir(parents=True, exist_ok=True) + logger.add(str(log_path), level=level.upper(), enqueue=True) ATOM37_FILL = 1e-5 -def rbf(r: Tensor, num_gaussians: int = 16, cutoff: float = 8.0) -> Tensor: - """Radial basis function encoding of distances using Bessel functions.""" +def rbf(r: Tensor, num_gaussians: int = NUM_RBF, cutoff: float = RBF_CUTOFF) -> Tensor: + """ + Compute radial basis function encoding of distances. + + Uses Bessel basis functions with smooth cutoff for distance encoding. + + Args: + r: (N,) distance tensor in Angstroms + num_gaussians: Number of RBF basis functions + cutoff: Maximum distance in Angstroms (values beyond are zeroed) + + Returns: + (N, num_gaussians) RBF feature tensor + """ r = r.clamp(min=1e-4) return soft_one_hot_linspace( - r, - start=0.0, - end=cutoff, - number=num_gaussians, - basis='bessel', - cutoff=True + r, start=0.0, end=cutoff, number=num_gaussians, basis="bessel", cutoff=True ) -def atom37_to_atoms(atom_tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: +def atom37_to_atoms( + atom_tensor: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """ Convert atom37 representation to flat atom list. @@ -78,6 +103,7 @@ def atom37_to_atoms(atom_tensor: torch.Tensor) -> tuple[torch.Tensor, torch.Tens return coords, residue_index, atom_type + @torch.no_grad() def ot_coupling( x1: torch.Tensor, @@ -96,12 +122,11 @@ def ot_coupling( x0_star: (M, 3) source positions (unchanged) x1_star: (M, 3) x1 permuted to match x0 per graph """ - device = x1.device x0_star = torch.empty_like(x1) x1_star = torch.empty_like(x1) for g in torch.unique(batch): - m = (batch == g) + m = batch == g if m.sum() == 0: continue @@ -117,8 +142,7 @@ def ot_coupling( return x0_star, x1_star -#eval metric functions - +# eval metric functions @torch.no_grad() def recall_precision( pred: torch.Tensor | np.ndarray, @@ -160,7 +184,6 @@ def recall_precision( return recall, precision - @torch.no_grad() def compute_rmsd( pred: torch.Tensor | np.ndarray, @@ -200,7 +223,6 @@ 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, @@ -225,7 +247,7 @@ def compute_placement_metrics( true = true.detach().cpu().numpy() if pred.size == 0 or true.size == 0: - return {'precision': 0.0, 'recall': 0.0, 'f1': 0.0, 'auc_pr': 0.0} + return {"precision": 0.0, "recall": 0.0, "f1": 0.0, "auc_pr": 0.0} D = spdist.cdist(true, pred) @@ -247,9 +269,9 @@ def compute_placement_metrics( sorted_idx = np.argsort(recalls) auc_pr = auc(np.array(recalls)[sorted_idx], np.array(precisions)[sorted_idx]) - return {'precision': precision, 'recall': recall, 'f1': f1, 'auc_pr': auc_pr} + return {"precision": precision, "recall": recall, "f1": f1, "auc_pr": auc_pr} -#viz functions +# viz functions def plot_3d_frame( ax, protein_pos: np.ndarray, @@ -261,36 +283,70 @@ def plot_3d_frame( ylim: tuple[float, float] = None, zlim: tuple[float, float] = None, ): - """Plot a single 3D frame showing protein, mates, and waters.""" + """ + Plot a single 3D frame showing protein structure and water positions. + + Args: + ax: Matplotlib 3D axes object + protein_pos: (N_p, 3) protein atom coordinates + mate_pos: (N_m, 3) symmetry mate coordinates, or None + water_pred: (N_w, 3) predicted water positions + water_true: (N_w, 3) ground truth water positions + title: Plot title string + xlim: Optional (min, max) x-axis limits in Angstroms + ylim: Optional (min, max) y-axis limits in Angstroms + zlim: Optional (min, max) z-axis limits in Angstroms + """ ax.clear() if protein_pos.size > 0: ax.scatter( - protein_pos[:, 0], protein_pos[:, 1], protein_pos[:, 2], - c='gray', alpha=0.3, s=8, label='Protein' + protein_pos[:, 0], + protein_pos[:, 1], + protein_pos[:, 2], + c="gray", + alpha=0.3, + s=8, + label="Protein", ) if mate_pos is not None and mate_pos.size > 0: ax.scatter( - mate_pos[:, 0], mate_pos[:, 1], mate_pos[:, 2], - c='dimgrey', alpha=0.6, s=10, label='Mates' + mate_pos[:, 0], + mate_pos[:, 1], + mate_pos[:, 2], + c="dimgrey", + alpha=0.6, + s=10, + label="Mates", ) ax.scatter( - water_true[:, 0], water_true[:, 1], water_true[:, 2], - c='red', marker='*', alpha=0.9, s=16, label='True Water' + water_true[:, 0], + water_true[:, 1], + water_true[:, 2], + c="red", + marker="*", + alpha=0.9, + s=16, + label="True Water", ) ax.scatter( - water_pred[:, 0], water_pred[:, 1], water_pred[:, 2], - c='blue', alpha=0.9, s=14, label='Predicted Water' + water_pred[:, 0], + water_pred[:, 1], + water_pred[:, 2], + c="blue", + alpha=0.9, + s=14, + label="Predicted Water", ) - ax.set_xlabel('X (Å)') - ax.set_ylabel('Y (Å)') - ax.set_zlabel('Z (Å)') + ax.set_xlabel("X (Å)") + ax.set_ylabel("Y (Å)") + ax.set_zlabel("Z (Å)") ax.set_title(title) - ax.legend(loc='upper right') + ax.legend(loc="upper right") if xlim is not None: ax.set_xlim(xlim) @@ -299,7 +355,6 @@ 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, @@ -343,15 +398,22 @@ def create_trajectory_gif( for i in frame_indices: water_pred = trajectory[i] fig = plt.figure(figsize=(10, 8)) - ax = fig.add_subplot(111, projection='3d') + ax = fig.add_subplot(111, projection="3d") - frame_title = f"{title} Step {i}/{len(trajectory)-1}" + frame_title = f"{title} Step {i}/{len(trajectory) - 1}" if pdb_id: frame_title = f"{pdb_id} | {frame_title}" plot_3d_frame( - ax, protein_pos, None, water_pred, water_true, - title=frame_title, xlim=xlim, ylim=ylim, zlim=zlim + ax, + protein_pos, + None, + water_pred, + water_true, + title=frame_title, + xlim=xlim, + ylim=ylim, + zlim=zlim, ) fig.canvas.draw() @@ -369,18 +431,25 @@ def create_trajectory_gif( save_all=True, append_images=all_frames[1:], duration=1000 // fps, - loop=0 + loop=0, ) - def save_protein_plot( pred_ca: torch.Tensor, true_ca: torch.Tensor, step: int, save_dir: str, ): - """Align and plot CA traces with Kabsch alignment.""" - + """ + Save 3D plot of predicted vs true CA traces with Kabsch alignment. + + Args: + pred_ca: (N_res, 3) predicted CA coordinates + true_ca: (N_res, 3) ground truth CA coordinates + step: Step number for filename + save_dir: Directory to save plot image + """ + P = pred_ca.detach().float().cpu().numpy() Q = true_ca.detach().float().cpu().numpy() @@ -398,9 +467,24 @@ def save_protein_plot( P_aligned = np.dot(P, R) fig = plt.figure(figsize=(8, 8)) - ax = fig.add_subplot(111, projection='3d') - ax.plot(Q[:, 0], Q[:, 1], Q[:, 2], color='black', linewidth=2, label='Ground Truth', alpha=0.6) - ax.plot(P_aligned[:, 0], P_aligned[:, 1], P_aligned[:, 2], color='red', linewidth=2, label='Prediction') + ax = fig.add_subplot(111, projection="3d") + ax.plot( + Q[:, 0], + Q[:, 1], + Q[:, 2], + color="black", + linewidth=2, + label="Ground Truth", + alpha=0.6, + ) + ax.plot( + P_aligned[:, 0], + P_aligned[:, 1], + P_aligned[:, 2], + color="red", + linewidth=2, + label="Prediction", + ) ax.legend() ax.set_title(f"Step {step}") plt.savefig(f"{save_dir}/step_{step}.png") diff --git a/tests/test_utils.py b/tests/test_utils.py index fafbe22..1691878 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -168,7 +168,7 @@ def test_coordinates_preserved(self): @pytest.mark.unit class TestOTCoupling: - """Tests for Hungarian matching in flow matching.""" + """Tests for OT coupling in flow matching.""" def test_output_shapes(self): """Output shapes should match input shapes.""" @@ -268,7 +268,6 @@ 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."""