Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 10 additions & 8 deletions src/constants.py
Original file line number Diff line number Diff line change
@@ -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]
172 changes: 128 additions & 44 deletions src/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
):
Comment on lines +28 to +31

Copilot AI Feb 26, 2026

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

setup_logging_for_tqdm now has new behavior (configurable log level + optional file sink/dir creation) but there are no unit tests covering these paths. Since tests/test_utils.py already exercises other utilities, consider adding a small test that calls setup_logging_for_tqdm(level=..., log_file=tmp_path/...) and asserts it doesn’t raise and creates/writes the log file.

Copilot uses AI. Check for mistakes.
"""
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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

keep all imports together at the top. Running ruff format and/or ruff check --fix can help do this automatically.


from loguru import logger

logger.remove()
logger.add(
lambda msg: tqdm.write(msg, end=""),
level=level.upper(),
colorize=True,
format="<level>{level: <8}</level> | <level>{message}</level>",
)
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:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This really applies to the code above: keep your code organized and put constants at the top. Generally it should go imports, then constants, then methods, then classes. It just makes for easier reading.

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

Expand All @@ -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,
Expand All @@ -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

Expand All @@ -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,
Expand Down Expand Up @@ -160,7 +184,6 @@ def recall_precision(

return recall, precision


@torch.no_grad()
def compute_rmsd(
pred: torch.Tensor | np.ndarray,
Expand Down Expand Up @@ -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,
Expand All @@ -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)

Expand All @@ -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,
Expand All @@ -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)
Expand All @@ -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,
Expand Down Expand Up @@ -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()
Expand All @@ -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()

Expand All @@ -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")
Expand Down
3 changes: 1 addition & 2 deletions tests/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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."""
Expand Down