From 708458cde2fc25e782159347cb6a734603be3d29 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Tue, 3 Mar 2026 16:25:34 -0600 Subject: [PATCH 1/3] refactoring training and inference script to match source code changes --- scripts/inference.py | 185 ++++++-- scripts/train.py | 1080 +++++++++++++++++++++++++++++++++--------- 2 files changed, 995 insertions(+), 270 deletions(-) diff --git a/scripts/inference.py b/scripts/inference.py index 1c5abe5..6f83519 100644 --- a/scripts/inference.py +++ b/scripts/inference.py @@ -27,6 +27,7 @@ from loguru import logger from tqdm import tqdm +from src.constants import NUM_RBF from src.dataset import ProteinWaterDataset from src.encoder_base import build_encoder from src.flow import FlowMatcher, FlowWaterGVP @@ -43,6 +44,21 @@ def parse_args(): + """ + Parse command-line arguments for inference configuration. + + Returns: + argparse.Namespace with inference parameters and paths + """ + # TODO: Add support for loading configuration from YAML/JSON config files. + # This would simplify inference invocation and ensure reproducibility. + # Example: --config config.yaml would load inference settings. + + # TODO: Remove hardcoded default paths. These should be required arguments + # or loaded from environment variables / config files for portability. + # Current hardcoded paths: + # - processed_dir: /home/srivasv/flow_cache/ + # - base_pdb_dir: /sb/wankowicz_lab/data/srivasv/pdb_redo_data p = argparse.ArgumentParser(description="Run WaterFlow inference on PDB files") p.add_argument( @@ -82,6 +98,13 @@ def parse_args(): action="store_true", help="Include symmetry mate atoms as protein nodes", ) + p.add_argument( + "--geometry_cache", + type=str, + default=None, + help="Geometry cache name to use (e.g., 'geometry' or 'geometry_unfiltered'). " + "Overrides the model's config if specified. Use this to evaluate against a specific ground truth.", + ) # checkpoint arguments p.add_argument( @@ -142,7 +165,7 @@ def parse_args(): type=float, default=None, help="Sample num_residues * water_ratio waters instead of using ground truth count. " - "E.g., --water_ratio 0.5 samples 50 waters for a 100-residue protein.", + "E.g., --water_ratio 0.5 samples 50 waters for a 100-residue protein.", ) args = p.parse_args() @@ -151,7 +174,18 @@ def parse_args(): def load_config(run_dir: Path) -> dict: - """Load training config from run directory.""" + """ + Load training configuration from run directory. + + Args: + run_dir: Path to training run directory containing config.json + + Returns: + Dict with training configuration parameters + + Raises: + FileNotFoundError: If config.json doesn't exist in run_dir + """ config_path = run_dir / "config.json" if not config_path.exists(): raise FileNotFoundError(f"Config file not found: {config_path}") @@ -163,39 +197,78 @@ def load_config(run_dir: Path) -> dict: def build_model_from_config(config: dict, device: torch.device) -> nn.Module: - """Build model architecture from config using registry-based encoder construction.""" - # Build encoder config for registry - # Example GVP config: - # {'encoder_type': 'gvp', 'hidden_s': 256, 'hidden_v': 64, 'node_scalar_in': 16} - # Example SLAE config: - # {'encoder_type': 'slae', 'hidden_s': 256, 'hidden_v': 64, 'slae_dim': 128, - # 'encoder_ckpt': '/path/to/slae_checkpoint.pt'} - encoder_config = { - 'encoder_type': 'slae' if config.get("use_slae", False) else 'gvp', - 'hidden_s': config.get("hidden_s", 256), - 'hidden_v': config.get("hidden_v", 64), - 'node_scalar_in': config.get("node_scalar_in", 16), - 'freeze_encoder': config.get("freeze_encoder", False), - 'slae_dim': config.get("slae_dim", 128), - 'encoder_ckpt': config.get("encoder_ckpt"), - } + """ + Build model architecture from training configuration. + + Uses registry-based encoder construction to instantiate the correct + encoder type (GVP, SLAE, or ESM) based on config. + + Args: + config: Training configuration dict with model hyperparameters. + Expected keys include: + - encoder_type: "gvp", "slae", or "esm" + - hidden_s, hidden_v: Hidden dimensions for scalars/vectors + - flow_layers: Number of flow layers + - For SLAE: slae_dim (default 128) + - For ESM: esm_dim (default 1536) + device: Device to place model on + + Returns: + FlowWaterGVP model instance + """ + # Use resolved_encoder_config if available (from training), otherwise build from config + resolved = config.get("resolved_encoder_config") + if resolved: + encoder_config = resolved.copy() + else: + encoder_type = config.get("encoder_type", "gvp") + encoder_config = { + "encoder_type": encoder_type, + "hidden_s": config.get("hidden_s") or 256, + "hidden_v": config.get("hidden_v") or 64, + "node_scalar_in": config.get("node_scalar_in") or 16, + "freeze_encoder": config.get("freeze_encoder", False), + "encoder_ckpt": config.get("encoder_ckpt"), + } + + # Add encoder-specific dimension (use 'or' to handle None values) + if encoder_type == "slae": + encoder_config["slae_dim"] = config.get("slae_dim") or 128 + elif encoder_type == "esm": + encoder_config["esm_dim"] = config.get("esm_dim") or 1536 encoder = build_encoder(encoder_config, device) model = FlowWaterGVP( encoder=encoder, - hidden_dims=(config.get("hidden_s", 256), config.get("hidden_v", 64)), - edge_scalar_dim=32, - layers=config.get("flow_layers", 5), - k_pw=config.get("k_pw", 24), - k_ww=config.get("k_ww", 24), + hidden_dims=(config.get("hidden_s") or 256, config.get("hidden_v") or 64), + edge_scalar_dim=config.get("edge_scalar_dim") or NUM_RBF, + layers=config.get("flow_layers") or 3, + drop_rate=config.get("drop_rate", 0.1), + n_message_gvps=config.get("n_message_gvps", 2), + n_update_gvps=config.get("n_update_gvps", 2), + k_pw=config.get("k_pw") or 16, + k_ww=config.get("k_ww") or 16, ).to(device) return model def load_checkpoint(model: nn.Module, checkpoint_path: Path, device: torch.device): - """Load model weights from checkpoint.""" + """ + Load model weights from checkpoint file. + + Args: + model: FlowWaterGVP model instance to load weights into + checkpoint_path: Path to checkpoint .pt file + device: Device to map checkpoint tensors to + + Returns: + Epoch number from checkpoint, or None if not stored + + Raises: + FileNotFoundError: If checkpoint file doesn't exist + """ if not checkpoint_path.exists(): raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}") @@ -253,13 +326,15 @@ def run_inference_batch( # build result dicts similar to rk4 results = [] for graph, water_pred in zip(graphs, water_preds): - results.append({ - "protein_pos": graph["protein"].pos.numpy(), - "water_true": graph["water"].pos.numpy(), - "water_pred": water_pred, - "trajectory": None, - "pdb_id": getattr(graph, 'pdb_id', None), - }) + results.append( + { + "protein_pos": graph["protein"].pos.numpy(), + "water_true": graph["water"].pos.numpy(), + "water_pred": water_pred, + "trajectory": None, + "pdb_id": getattr(graph, "pdb_id", None), + } + ) return results @@ -270,7 +345,15 @@ def save_plot( output_path: Path, metrics: dict, ): - """Save 3D visualization plot.""" + """ + Save 3D visualization plot of water prediction results. + + Args: + result: Dict with 'protein_pos', 'water_pred', 'water_true' arrays + pdb_id: PDB identifier for title + output_path: Path to save PNG image + metrics: Dict with 'rmsd', 'precision', 'recall', 'f1' for title + """ fig = plt.figure(figsize=(12, 10)) ax = fig.add_subplot(111, projection="3d") @@ -294,6 +377,8 @@ def save_plot( def main(): + """Run inference pipeline on a list of PDB structures.""" + setup_logging_for_tqdm() args = parse_args() # setup paths @@ -310,7 +395,7 @@ def main(): logger.info(f"Using device: {device}") # load config and build model - logger.info(f"\nLoading model from: {run_dir}") + logger.info(f"Loading model from: {run_dir}") config = load_config(run_dir) model = build_model_from_config(config, device) @@ -326,28 +411,39 @@ def main(): ) # Load dataset - logger.info(f"\nLoading PDBs from: {args.pdb_list}") + logger.info(f"Loading PDBs from: {args.pdb_list}") # Determine include_mates from args or config include_mates = args.include_mates or config.get("include_mates", False) + encoder_type = config.get("encoder_type", "gvp") + + # Use --geometry_cache if provided, otherwise use config's geometry_cache_name + geometry_cache_name = args.geometry_cache or config.get( + "geometry_cache_name", "geometry" + ) dataset = ProteinWaterDataset( pdb_list_file=args.pdb_list, processed_dir=args.processed_dir, base_pdb_dir=args.base_pdb_dir, + encoder_type=encoder_type, include_mates=include_mates, + geometry_cache_name=geometry_cache_name, preprocess=True, ) logger.info(f"Found {len(dataset)} PDB entries") + logger.info(f"Using geometry cache: {geometry_cache_name}") # run inference - logger.info(f"\nRunning inference with method={args.method}, steps={args.num_steps}") + logger.info(f"Running inference with method={args.method}, steps={args.num_steps}") logger.info(f"Self-conditioning: {args.use_sc}") logger.info(f"Threshold for metrics: {args.threshold}Å") logger.info(f"Batch size: {args.batch_size}") if args.water_ratio is not None: - logger.info(f"Water ratio: {args.water_ratio} (sampling num_residues × {args.water_ratio} waters)") + logger.info( + f"Water ratio: {args.water_ratio} (sampling num_residues × {args.water_ratio} waters)" + ) else: logger.info("Water ratio: None (using ground truth water count)") logger.info("-" * 60) @@ -389,8 +485,8 @@ def main(): # process each result in the batch for result in batch_results: pdb_id = result.get("pdb_id", f"unknown_{len(all_metrics)}") - water_true = result["water_true"] water_pred = result["water_pred"] + water_true = result["water_true"] # compute metrics metrics = compute_placement_metrics( @@ -438,8 +534,12 @@ def main(): "avg_recall": float(np.mean([m["recall"] for m in all_metrics])), "avg_f1": float(np.mean([m["f1"] for m in all_metrics])), "avg_auc_pr": float(np.mean([m["auc_pr"] for m in all_metrics])), - "avg_n_waters_true": float(np.mean([m["n_waters_true"] for m in all_metrics])), - "avg_n_waters_pred": float(np.mean([m["n_waters_pred"] for m in all_metrics])), + "avg_n_waters_true": float( + np.mean([m["n_waters_true"] for m in all_metrics]) + ), + "avg_n_waters_pred": float( + np.mean([m["n_waters_pred"] for m in all_metrics]) + ), } logger.info("\n" + "=" * 60) @@ -471,21 +571,22 @@ def main(): "threshold": args.threshold, "include_mates": include_mates, "water_ratio": args.water_ratio, + "geometry_cache": geometry_cache_name, }, }, f, indent=2, ) - logger.info(f"\nMetrics saved to: {metrics_path}") + logger.info(f"Metrics saved to: {metrics_path}") else: - logger.info("\nNo valid samples processed.") + logger.warning("No valid samples processed.") logger.info(f"Plots saved to: {output_dir / 'plots'}") if args.save_gifs: logger.info(f"GIFs saved to: {output_dir / 'gifs'}") - logger.info("\nInference complete.") + logger.info("Inference complete.") if __name__ == "__main__": diff --git a/scripts/train.py b/scripts/train.py index 802c9d5..9f97e08 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -1,22 +1,42 @@ -# train.py +""" +Training pipeline for WaterFlow model. + +This module provides the main training script for the WaterFlow water placement +model. It handles: +- Dataset loading and preprocessing with configurable quality filters +- Model construction with pluggable encoders (GVP, SLAE, ESM) +- Training loop with gradient accumulation and warmup scheduling +- Validation and evaluation with RK4 trajectory integration +- Checkpointing and W&B logging + +Usage: + python scripts/train.py \\ + --train_list /path/to/train.txt \\ + --val_list /path/to/val.txt \\ + --encoder_type gvp \\ + --epochs 200 \\ + --batch_size 4 +""" import argparse -import os +import json from datetime import datetime from pathlib import Path import matplotlib.pyplot as plt import numpy as np import torch -import torch.nn as nn import wandb from loguru import logger from torch.optim import AdamW -from torch.optim.lr_scheduler import CosineAnnealingLR +from torch.optim.lr_scheduler import CosineAnnealingLR, StepLR, LinearLR from tqdm import tqdm -from src.dataset import get_dataloader +from torch.utils.data import DataLoader +from torch_geometric.data import HeteroData + from src.encoder_base import build_encoder +from src.dataset import get_dataloader from src.flow import FlowMatcher, FlowWaterGVP from src.utils import ( compute_placement_metrics, @@ -26,52 +46,256 @@ setup_logging_for_tqdm, ) -# Configure logging to work with tqdm progress bars -setup_logging_for_tqdm() - -def generate_run_name(args): +def generate_run_name(args: argparse.Namespace) -> str: """Generate a run name from timestamp and key parameters.""" timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") - encoder_type = "slae" if args.use_slae else "gvp" layers = f"L{args.flow_layers}" hidden = f"h{args.hidden_s}" - name = f"{timestamp}_{encoder_type}_{layers}_{hidden}" + name = f"{timestamp}_{args.encoder_type}_{layers}_{hidden}" return name def parse_args(): + """ + Parse command-line arguments for training configuration. + + Returns: + argparse.Namespace with all training hyperparameters and paths + """ + # TODO: Add support for loading configuration from YAML/JSON config files. + # This would allow users to save and share training configurations easily. + # Example: --config config.yaml would load all arguments from the file, + # with CLI args taking precedence for overrides. + + # TODO: Remove hardcoded default paths. These should be required arguments + # or loaded from environment variables / config files for portability. + # Current hardcoded paths: + # - processed_dir: /home/srivasv/flow_cache/ + # - base_pdb_dir: /sb/wankowicz_lab/data/srivasv/pdb_redo_data + # - edia_dir: /sb/wankowicz_lab/data/srivasv/edia_results + # - save_dir: /home/srivasv/flow_checkpoints + # - wandb_dir: /home/srivasv/wandb_logs p = argparse.ArgumentParser() # data p.add_argument("--train_list", type=str, required=True) p.add_argument("--val_list", type=str, required=True) - p.add_argument("--processed_dir", type=str, default="/home/srivasv/flow_cache/") - p.add_argument("--base_pdb_dir", type=str, default="/sb/wankowicz_lab/data/srivasv/pdb_redo_data") - p.add_argument("--include_mates", action="store_true", help="Include symmetry mate atoms as protein nodes") - p.add_argument("--duplicate_single_sample", type=int, default=1, - help="If training on single sample, duplicate it N times for more gradient updates per epoch") + p.add_argument( + "--processed_dir", + type=str, + default="/home/srivasv/flow_cache/", + help=( + "Cache root. Geometry caches are expected in /geometry, " + "embeddings in /." + ), + ) + p.add_argument( + "--base_pdb_dir", + type=str, + default="/sb/wankowicz_lab/data/srivasv/pdb_redo_data", + ) + p.add_argument( + "--geometry_cache_name", + type=str, + default="geometry", + help="Base name for geometry cache directory (e.g., 'geometry' -> geometry/ or geometry_unfiltered/)", + ) + p.add_argument( + "--include_mates", + action="store_true", + help="Include symmetry mate atoms as protein nodes", + ) + p.add_argument( + "--duplicate_single_sample", + type=int, + default=1, + help="If training on single sample, duplicate it N times for more gradient updates per epoch", + ) + p.add_argument( + "--edia_dir", + type=str, + default="/sb/wankowicz_lab/data/srivasv/edia_results", + help=( + "Water filter: EDIA root directory " + "({edia_dir}/{pdb_id}/{pdb_id}_residue_stats.csv)." + ), + ) + + # dataset quality checks (always on) + p.add_argument( + "--max_com_dist", + type=float, + default=25.0, + help="Quality: max allowed protein-water center-of-mass distance (Angstroms).", + ) + p.add_argument( + "--max_clash_fraction", + type=float, + default=0.05, + help="Quality: max allowed fraction of waters clashing with protein.", + ) + p.add_argument( + "--clash_dist", + type=float, + default=2.0, + help="Quality: distance threshold for defining a water-protein clash (Angstroms).", + ) + p.add_argument( + "--interface_dist_threshold", + type=float, + default=4.0, + help="Quality: max inter-chain interface distance to treat chains as interacting (Angstroms).", + ) + p.add_argument( + "--min_water_residue_ratio", + type=float, + default=0.6, + help="Quality: minimum waters/residue ratio required per structure.", + ) + + # per-water filtering (toggleable) + p.add_argument( + "--max_protein_dist", + type=float, + default=5.0, + help="Water filter: remove waters farther than this from nearest protein atom (Angstroms).", + ) + p.add_argument( + "--min_edia", + type=float, + default=0.4, + help="Water filter: remove waters with EDIA below this threshold.", + ) + p.add_argument( + "--max_bfactor_zscore", + type=float, + default=1.5, + help="Water filter: remove waters with normalized B-factor above this threshold.", + ) + p.add_argument( + "--no_filter_by_distance", + dest="filter_by_distance", + action="store_false", + help="Disable distance-from-protein water filtering (ignores --max_protein_dist).", + ) + p.add_argument( + "--no_filter_by_edia", + dest="filter_by_edia", + action="store_false", + help="Disable EDIA-based water filtering (ignores --min_edia).", + ) + p.add_argument( + "--no_filter_by_bfactor", + dest="filter_by_bfactor", + action="store_false", + help="Disable B-factor-based water filtering (ignores --max_bfactor_zscore).", + ) + p.set_defaults(filter_by_distance=True, filter_by_edia=True, filter_by_bfactor=True) # model + p.add_argument( + "--encoder_type", type=str, default="gvp", choices=["gvp", "slae", "esm"] + ) p.add_argument("--encoder_ckpt", type=str, default=None) p.add_argument("--freeze_encoder", action="store_true") p.add_argument("--hidden_s", type=int, default=256) p.add_argument("--hidden_v", type=int, default=64) p.add_argument("--flow_layers", type=int, default=3) + p.add_argument( + "--n_message_gvps", + type=int, + default=2, + help="Number of GVPs in message function per edge type (default: 2)", + ) + p.add_argument( + "--n_update_gvps", + type=int, + default=2, + help="Number of GVPs in node update function (default: 2)", + ) + p.add_argument( + "--drop_rate", + type=float, + default=0.1, + help="Dropout rate for GVP layers (default: 0.1)", + ) p.add_argument("--k_pw", type=int, default=16) p.add_argument("--k_ww", type=int, default=16) - # SLAE encoder options - p.add_argument("--use_slae", action="store_true", help="Use SLAE encoder instead of GVP") - p.add_argument("--slae_dim", type=int, default=128, help="SLAE embedding dimension") + # optional encoder-specific overrides + p.add_argument( + "--slae_dim", + type=int, + default=None, + help="Optional SLAE embedding dimension override", + ) + p.add_argument( + "--esm_dim", + type=int, + default=None, + help="Optional ESM embedding dimension override", + ) # training - p.add_argument("--epochs", type=int, default=100) + p.add_argument("--epochs", type=int, default=200) p.add_argument("--batch_size", type=int, default=4) - p.add_argument("--lr", type=float, default=5e-4) - p.add_argument("--weight_decay", type=float, default=1e-3) + p.add_argument("--lr", type=float, default=1e-3) + p.add_argument("--weight_decay", type=float, default=1e-4) p.add_argument("--grad_clip", type=float, default=1.0) - p.add_argument("--num_workers", type=int, default=4) + p.add_argument( + "--grad_accum_steps", + type=int, + default=1, + help="Number of gradient accumulation steps", + ) + p.add_argument("--num_workers", type=int, default=8) + p.add_argument( + "--prefetch_factor", + type=int, + default=4, + help="Number of batches to prefetch per worker", + ) + p.add_argument( + "--pin_memory", + action="store_true", + default=True, + help="Pin memory for faster CPU-GPU transfer", + ) + p.add_argument( + "--no_pin_memory", + dest="pin_memory", + action="store_false", + help="Disable pin_memory", + ) + p.add_argument( + "--persistent_workers", + action="store_true", + default=True, + help="Keep workers alive between epochs", + ) + p.add_argument( + "--no_persistent_workers", + dest="persistent_workers", + action="store_false", + help="Disable persistent_workers", + ) + + # scheduler + p.add_argument( + "--scheduler", type=str, default="cosine", choices=["cosine", "step", "none"] + ) + p.add_argument("--warmup_steps", type=int, default=0, help="Linear warmup steps") + p.add_argument( + "--eta_min_factor", + type=float, + default=0.001, + help="eta_min = lr * eta_min_factor", + ) + p.add_argument( + "--step_size", type=int, default=50, help="StepLR step size (epochs)" + ) + p.add_argument("--step_gamma", type=float, default=0.5, help="StepLR gamma") # flow matching p.add_argument("--use_self_cond", action="store_true") @@ -83,46 +307,254 @@ def parse_args(): # checkpointing p.add_argument("--save_dir", type=str, default="/home/srivasv/flow_checkpoints") - p.add_argument("--run_name", type=str, default=None, help="Name for this run (auto-generated if not provided)") - p.add_argument("--save_every", type=int, default=100) + p.add_argument( + "--run_name", + type=str, + default=None, + help="Name for this run (auto-generated if not provided)", + ) + p.add_argument("--save_every", type=int, default=10) p.add_argument("--eval_every", type=int, default=5) p.add_argument("--n_eval_samples", type=int, default=3) p.add_argument("--rk4_steps", type=int, default=100) - p.add_argument("--save_gifs", action="store_true", help="Save trajectory GIFs during eval") + p.add_argument( + "--save_gifs", action="store_true", help="Save trajectory GIFs during eval" + ) - # wandb + # logging / wandb + p.add_argument("--log_level", type=str, default="INFO") + p.add_argument("--log_file", type=str, default=None) p.add_argument("--wandb_project", type=str, default="water-flow") - p.add_argument("--wandb_run", type=str, default=None) p.add_argument("--wandb_dir", type=str, default="/home/srivasv/wandb_logs") p.add_argument("--device", type=str, default="cuda") return p.parse_args() -def build_model(args, device, node_scalar_in=16): - """Build encoder and flow model using registry-based encoder construction.""" - encoder_type = 'slae' if args.use_slae else 'gvp' - logger.info(f"Building model with {encoder_type.upper()} encoder") +def _extract_quality_config(args: argparse.Namespace) -> dict: + """Extract dataset quality check parameters (always active in preprocessing).""" + return { + "max_com_dist": args.max_com_dist, + "max_clash_fraction": args.max_clash_fraction, + "clash_dist": args.clash_dist, + "interface_dist_threshold": args.interface_dist_threshold, + "min_water_residue_ratio": args.min_water_residue_ratio, + } + + +def _extract_water_filter_config(args: argparse.Namespace) -> dict: + """Extract per-water filtering parameters (toggleable).""" + return { + "edia_dir": args.edia_dir, + "max_protein_dist": args.max_protein_dist, + "min_edia": args.min_edia, + "max_bfactor_zscore": args.max_bfactor_zscore, + "filter_by_distance": args.filter_by_distance, + "filter_by_edia": args.filter_by_edia, + "filter_by_bfactor": args.filter_by_bfactor, + } + + +def _build_dataset_config(args: argparse.Namespace) -> tuple[dict, dict, dict]: + """ + Build grouped dataset configuration from command-line arguments. + + Args: + args: Parsed command-line arguments + + Returns: + Tuple of (dataset_kwargs, quality_kwargs, water_filter_kwargs): + - dataset_kwargs: Merged dict for DataLoader creation + - quality_kwargs: Structure-level quality check parameters + - water_filter_kwargs: Per-water filtering parameters + """ + quality_kwargs = _extract_quality_config(args) + water_filter_kwargs = _extract_water_filter_config(args) + dataset_kwargs = { + "encoder_type": args.encoder_type, + "base_pdb_dir": args.base_pdb_dir, + "geometry_cache_name": args.geometry_cache_name, + "include_mates": args.include_mates, + **quality_kwargs, + **water_filter_kwargs, + } + return dataset_kwargs, quality_kwargs, water_filter_kwargs + + +def _ignored_water_filter_thresholds(args) -> list[str]: + """ + Identify water filter thresholds that are disabled. + + Args: + args: Parsed command-line arguments with filter_by_* flags + + Returns: + List of threshold parameter names that are disabled (e.g., ['min_edia']) + """ + ignored = [] + if not args.filter_by_distance: + ignored.append("max_protein_dist") + if not args.filter_by_edia: + ignored.append("min_edia") + if not args.filter_by_bfactor: + ignored.append("max_bfactor_zscore") + return ignored + + +def _log_dataset_filter_config(args, quality_kwargs: dict): + """ + Log dataset quality check and water filter configuration. + + Args: + args: Parsed command-line arguments with filter settings + quality_kwargs: Structure-level quality check parameters to log + """ + active_filters = { + "distance": args.filter_by_distance, + "edia": args.filter_by_edia, + "bfactor": args.filter_by_bfactor, + } + logger.info(f"Dataset quality checks (always on): {quality_kwargs}") + logger.info(f"Water filters (toggleable): {active_filters}") - if args.use_slae: - logger.info(f" SLAE dim: {args.slae_dim}") + ignored = _ignored_water_filter_thresholds(args) + if ignored: + logger.info(f"Ignored water-filter thresholds (disabled): {ignored}") + + if args.filter_by_edia and args.edia_dir is None: + logger.info( + "EDIA filter enabled but --edia_dir is not set; EDIA filtering will be skipped." + ) + + +def _required_embedding_field(encoder_type: str) -> str | None: + """ + Get the required embedding field name for a given encoder type. + Args: + encoder_type: Encoder identifier ('gvp', 'slae', or 'esm') + + Returns: + Field name string (e.g., 'slae_embedding') or None if encoder doesn't need embeddings + """ + if encoder_type == "slae": + return "slae_embedding" + if encoder_type == "esm": + return "esm_embedding" + return None + + +def _resolve_embedding_dim( + sample_data, + encoder_type: str, + override_dim: int | None, +) -> int | None: + """ + Infer or validate embedding dimension from sample data. + + Args: + sample_data: HeteroData sample from the dataset + encoder_type: Encoder identifier ('gvp', 'slae', or 'esm') + override_dim: User-specified dimension override, or None to infer + + Returns: + Embedding dimension, or None if encoder doesn't use embeddings + + Raises: + ValueError: If required embedding field is missing or dimension mismatch + """ + field = _required_embedding_field(encoder_type) + if field is None: + return None + if field not in sample_data["protein"]: + raise ValueError( + f"Selected encoder '{encoder_type}' requires protein.{field}, " + f"but it is missing from dataset samples. " + f"Expected cache at {field.split('_')[0]}/.pt under --processed_dir." + ) + + inferred_dim = int(sample_data["protein"][field].shape[-1]) + if override_dim is not None and int(override_dim) != inferred_dim: + raise ValueError( + f"{encoder_type} dim override mismatch: override={override_dim}, " + f"inferred={inferred_dim} from sample data" + ) + return inferred_dim if override_dim is None else int(override_dim) + + +def resolve_encoder_config(args, sample_data, node_scalar_in: int): + """ + Build a registry-friendly encoder config with inferred dimensions. + + Args: + args: Parsed command-line arguments containing encoder settings + sample_data: HeteroData sample used to infer embedding dimensions + node_scalar_in: Number of input scalar features per node + + Returns: + dict: Encoder configuration ready for build_encoder(), e.g.: + - GVP: {"encoder_type": "gvp", "hidden_s": 256, "hidden_v": 64, ...} + - SLAE: {"encoder_type": "slae", "slae_dim": 128, ...} + - ESM: {"encoder_type": "esm", "esm_dim": 1536, ...} + """ encoder_config = { - 'encoder_type': encoder_type, - 'hidden_s': args.hidden_s, - 'hidden_v': args.hidden_v, - 'node_scalar_in': node_scalar_in, - 'freeze_encoder': args.freeze_encoder, - 'slae_dim': args.slae_dim, - 'encoder_ckpt': args.encoder_ckpt, + "encoder_type": args.encoder_type, + "hidden_s": args.hidden_s, + "hidden_v": args.hidden_v, + "node_scalar_in": node_scalar_in, + "freeze_encoder": args.freeze_encoder, + "encoder_ckpt": args.encoder_ckpt, } + if args.encoder_type == "slae": + encoder_config["slae_dim"] = _resolve_embedding_dim( + sample_data, "slae", args.slae_dim + ) + elif args.encoder_type == "esm": + encoder_config["esm_dim"] = _resolve_embedding_dim( + sample_data, "esm", args.esm_dim + ) + + return encoder_config + + +def log_encoder_sample_stats(sample_data: HeteroData, encoder_type: str) -> None: + """Log summary statistics for the selected encoder input features.""" + field = _required_embedding_field(encoder_type) + if field is None: + return + emb = sample_data["protein"][field] + logger.info( + f"{field} shape={tuple(emb.shape)} " + f"mean={emb.mean():.4f} std={emb.std():.4f} min={emb.min():.4f} max={emb.max():.4f}" + ) + + +def build_model( + args: argparse.Namespace, device: torch.device, encoder_config: dict +) -> FlowWaterGVP: + """ + Build encoder and flow model using registry-based encoder construction. + + Args: + args: Parsed command-line arguments with model hyperparameters + device: Torch device to place the model on + encoder_config: Registry-friendly config from resolve_encoder_config() + + Returns: + FlowWaterGVP: Initialized model with the specified encoder + """ + logger.info(f"Building model with {args.encoder_type.upper()} encoder") + logger.info(f"Resolved encoder config: {encoder_config}") + encoder = build_encoder(encoder_config, device) model = FlowWaterGVP( encoder=encoder, hidden_dims=(args.hidden_s, args.hidden_v), - edge_scalar_dim=32, layers=args.flow_layers, + n_message_gvps=args.n_message_gvps, + n_update_gvps=args.n_update_gvps, + drop_rate=args.drop_rate, k_pw=args.k_pw, k_ww=args.k_ww, ).to(device) @@ -130,7 +562,9 @@ def build_model(args, device, node_scalar_in=16): return model -def run_eval_sampling(flow_matcher, val_loader, args, epoch, device, global_step, eval_indices, run_dir): +def run_eval_sampling( + flow_matcher, val_loader, args, epoch, device, global_step, eval_indices, run_dir +): """Run RK4 integration on fixed eval samples and log results. Args: @@ -142,7 +576,7 @@ def run_eval_sampling(flow_matcher, val_loader, args, epoch, device, global_step for i, idx in enumerate(eval_indices): graph = val_loader.dataset[idx] - if graph['water'].num_nodes == 0: + if graph["water"].num_nodes == 0: continue out = flow_matcher.rk4_integrate( @@ -155,31 +589,31 @@ def run_eval_sampling(flow_matcher, val_loader, args, epoch, device, global_step # compute metrics final_metrics = compute_placement_metrics( - pred=out['water_pred'], - true=out['water_true'], - threshold=1.0 + pred=out["water_pred"], true=out["water_true"], threshold=1.0 ) - final_rmsd = compute_rmsd(out['water_pred'], out['water_true']) + final_rmsd = compute_rmsd(out["water_pred"], out["water_true"]) - results.append({ - 'rmsd': final_rmsd, - 'precision': final_metrics['precision'], - 'recall': final_metrics['recall'], - 'f1': final_metrics['f1'], - 'auc_pr': final_metrics['auc_pr'] - }) + results.append( + { + "rmsd": final_rmsd, + "precision": final_metrics["precision"], + "recall": final_metrics["recall"], + "f1": final_metrics["f1"], + "auc_pr": final_metrics["auc_pr"], + } + ) # plot final frame fig = plt.figure(figsize=(10, 8)) - ax = fig.add_subplot(111, projection='3d') + ax = fig.add_subplot(111, projection="3d") plot_3d_frame( ax, - out['protein_pos'], + out["protein_pos"], None, - out['water_pred'], - out['water_true'], - title=f"Epoch {epoch} Sample {i} | RMSD={final_rmsd:.2f}Å | F1={final_metrics['f1']:.3f}" + out["water_pred"], + out["water_true"], + title=f"Epoch {epoch} Sample {i} | RMSD={final_rmsd:.2f}A | F1={final_metrics['f1']:.3f}", ) plot_path = run_dir / "plots" / f"epoch{epoch}_sample{i}.png" @@ -188,97 +622,192 @@ def run_eval_sampling(flow_matcher, val_loader, args, epoch, device, global_step plt.close() # save GIF if requested - if args.save_gifs and 'trajectory' in out: + if args.save_gifs and "trajectory" in out: gif_path = run_dir / "gifs" / f"epoch{epoch}_sample{i}.gif" gif_path.parent.mkdir(parents=True, exist_ok=True) create_trajectory_gif( - trajectory=out['trajectory'], - protein_pos=out['protein_pos'], - water_true=out['water_true'], + trajectory=out["trajectory"], + protein_pos=out["protein_pos"], + water_true=out["water_true"], save_path=str(gif_path), title=f"Epoch {epoch} Sample {i}", fps=10, - pdb_id=graph.pdb_id + pdb_id=graph.pdb_id, ) if results: avg_metrics = { - "eval/avg_rmsd": np.mean([r['rmsd'] for r in results]), - "eval/avg_precision": np.mean([r['precision'] for r in results]), - "eval/avg_recall": np.mean([r['recall'] for r in results]), - "eval/avg_f1": np.mean([r['f1'] for r in results]), - "eval/avg_auc_pr": np.mean([r['auc_pr'] for r in results]), + "eval/avg_rmsd": np.mean([r["rmsd"] for r in results]), + "eval/avg_precision": np.mean([r["precision"] for r in results]), + "eval/avg_recall": np.mean([r["recall"] for r in results]), + "eval/avg_f1": np.mean([r["f1"] for r in results]), + "eval/avg_auc_pr": np.mean([r["auc_pr"] for r in results]), } wandb.log(avg_metrics, step=global_step) return avg_metrics return {} -def train_epoch(flow_matcher, train_loader, optimizer, args, epoch): - """Single training epoch.""" +def train_epoch( + flow_matcher: FlowMatcher, + train_loader: DataLoader, + optimizer: AdamW, + warmup_scheduler, + args: argparse.Namespace, + epoch: int, + optimizer_step_count: int, +) -> tuple[dict[str, float], int, int]: + """Single training epoch with gradient accumulation and warmup support.""" flow_matcher.model.train() total_loss, total_rmsd = 0.0, 0.0 + skipped_batches = 0 + processed_batches = 0 + + optimizer.zero_grad(set_to_none=True) pbar = tqdm(train_loader, desc=f"Epoch {epoch} [Train]") for step, batch in enumerate(pbar): batch = batch.to(args.device) - if batch['water'].num_nodes == 0: + if batch["water"].num_nodes == 0: + skipped_batches += 1 continue metrics = flow_matcher.training_step( - batch, optimizer, - grad_clip=args.grad_clip, + batch, use_self_conditioning=args.use_self_cond, + accumulation_steps=args.grad_accum_steps, ) - # print per-sample losses if batch loss exceeded 100.0 - if metrics['per_sample_info'] is not None: - per_sample_losses = metrics['per_sample_info']['losses'].cpu() - num_graphs = metrics['per_sample_info']['num_graphs'] - - if hasattr(batch, 'pdb_id'): - # pdb_id might be a list when batched - pdb_ids = batch.pdb_id if isinstance(batch.pdb_id, list) else [batch.pdb_id] - logger.info(f"\n{'='*60}") - logger.warning(f"WARNING: Batch loss {metrics['loss']:.2f} exceeded 100.0!") - logger.info(f"Per-sample losses ({num_graphs} samples):") + if metrics["per_sample_info"] is not None: + per_sample_losses = metrics["per_sample_info"]["losses"].cpu() + num_graphs = metrics["per_sample_info"]["num_graphs"] + + if hasattr(batch, "pdb_id"): + pdb_ids = ( + batch.pdb_id if isinstance(batch.pdb_id, list) else [batch.pdb_id] + ) + logger.warning("=" * 60) + logger.warning(f"Batch loss {metrics['loss']:.2f} exceeded 100.0!") + logger.warning(f"Per-sample losses ({num_graphs} samples):") for i in range(num_graphs): - pdb_id = pdb_ids[i] if i < len(pdb_ids) else 'unknown' + pdb_id = pdb_ids[i] if i < len(pdb_ids) else "unknown" sample_loss = per_sample_losses[i].item() - logger.info(f" [{i}] {pdb_id}: {sample_loss:.2f}") - logger.info(f"{'='*60}") - - total_loss += metrics['loss'] - total_rmsd += metrics['rmsd'] - pbar.set_postfix(loss=f"{metrics['loss']:.4f}", rmsd=f"{metrics['rmsd']:.2f}") + logger.warning(f"[{i}] {pdb_id}: {sample_loss:.2f}") + logger.warning("=" * 60) + + processed_batches += 1 + total_loss += metrics["loss"] + total_rmsd += metrics["rmsd"] + + # Step optimizer every grad_accum_steps + if (step + 1) % args.grad_accum_steps == 0: + if args.grad_clip > 0: + torch.nn.utils.clip_grad_norm_( + [p for p in flow_matcher.model.parameters() if p.requires_grad], + max_norm=args.grad_clip, + ) + optimizer.step() + optimizer.zero_grad(set_to_none=True) + optimizer_step_count += 1 + + # Step warmup scheduler per optimizer step + if ( + warmup_scheduler is not None + and optimizer_step_count <= args.warmup_steps + ): + warmup_scheduler.step() + + current_lr = optimizer.param_groups[0]["lr"] + pbar.set_postfix( + loss=f"{metrics['loss']:.4f}", + rmsd=f"{metrics['rmsd']:.2f}", + lr=f"{current_lr:.2e}", + ) global_step = (epoch - 1) * len(train_loader) + step - wandb.log({ - "train/iter_loss": metrics['loss'], - "train/iter_rmsd": metrics['rmsd'], - }, step=global_step) + wandb.log( + { + "train/iter_loss": metrics["loss"], + "train/iter_rmsd": metrics["rmsd"], + "lr": current_lr, + }, + step=global_step, + ) - n = len(train_loader) + # Handle remaining gradients at end of epoch + if (step + 1) % args.grad_accum_steps != 0: + if args.grad_clip > 0: + torch.nn.utils.clip_grad_norm_( + [p for p in flow_matcher.model.parameters() if p.requires_grad], + max_norm=args.grad_clip, + ) + optimizer.step() + optimizer.zero_grad(set_to_none=True) + optimizer_step_count += 1 + if warmup_scheduler is not None and optimizer_step_count <= args.warmup_steps: + warmup_scheduler.step() final_global_step = (epoch - 1) * len(train_loader) + len(train_loader) - 1 - return {'train/epoch_loss': total_loss / n, 'train/epoch_rmsd': total_rmsd / n}, final_global_step + + if processed_batches == 0: + logger.warning( + f"Epoch {epoch}: skipped all {skipped_batches} train batches (no waters)." + ) + return ( + {"train/epoch_loss": float("inf"), "train/epoch_rmsd": float("inf")}, + final_global_step, + optimizer_step_count, + ) + + logger.info( + f"Epoch {epoch} [Train] processed_batches={processed_batches}, skipped_batches={skipped_batches}" + ) + return ( + { + "train/epoch_loss": total_loss / processed_batches, + "train/epoch_rmsd": total_rmsd / processed_batches, + }, + final_global_step, + optimizer_step_count, + ) + @torch.no_grad() -def val_epoch(flow_matcher, val_loader, args, epoch): +def val_epoch( + flow_matcher: FlowMatcher, + val_loader: DataLoader, + args: argparse.Namespace, + epoch: int, +) -> dict[str, float]: """Single validation epoch.""" flow_matcher.model.eval() total_loss, total_rmsd = 0.0, 0.0 - + skipped_batches = 0 + processed_batches = 0 + for batch in tqdm(val_loader, desc=f"Epoch {epoch} [Val]"): batch = batch.to(args.device) - if batch['water'].num_nodes == 0: + if batch["water"].num_nodes == 0: + skipped_batches += 1 continue metrics = flow_matcher.validation_step(batch) - total_loss += metrics['loss'] - total_rmsd += metrics['rmsd'] - - n = len(val_loader) - return {'val/loss': total_loss / n, 'val/rmsd': total_rmsd / n} + processed_batches += 1 + total_loss += metrics["loss"] + total_rmsd += metrics["rmsd"] + + if processed_batches == 0: + logger.warning( + f"Epoch {epoch}: skipped all {skipped_batches} val batches (no waters)." + ) + return {"val/loss": float("inf"), "val/rmsd": float("inf")} + + logger.info( + f"Epoch {epoch} [Val] processed_batches={processed_batches}, skipped_batches={skipped_batches}" + ) + return { + "val/loss": total_loss / processed_batches, + "val/rmsd": total_rmsd / processed_batches, + } def count_parameters(model): @@ -288,41 +817,85 @@ def count_parameters(model): return trainable, total -def check_slae_embeddings(loader, device): - """Check if SLAE embeddings are present and compute statistics.""" - logger.info("\nChecking SLAE embeddings in dataset...") - batch = next(iter(loader)) - batch = batch.to(device) +def save_checkpoint( + model, + optimizer, + warmup_scheduler, + main_scheduler, + epoch, + optimizer_step_count, + path, + best=False, +): + """ + Save model checkpoint with optimizer and scheduler states. + + Args: + model: FlowWaterGVP model instance + optimizer: AdamW optimizer instance + warmup_scheduler: LinearLR warmup scheduler, or None + main_scheduler: Main LR scheduler (CosineAnnealingLR or StepLR), or None + epoch: Current epoch number + optimizer_step_count: Total number of optimizer steps taken + path: Path object for checkpoint file destination + best: If True, log as best checkpoint + """ + path.parent.mkdir(parents=True, exist_ok=True) + torch.save( + { + "epoch": epoch, + "optimizer_step_count": optimizer_step_count, + "model_state_dict": model.state_dict(), + "optimizer_state_dict": optimizer.state_dict(), + "warmup_scheduler_state_dict": warmup_scheduler.state_dict() + if warmup_scheduler + else None, + "main_scheduler_state_dict": main_scheduler.state_dict() + if main_scheduler + else None, + }, + path, + ) + logger.info(f"{'Best ' if best else ''}Checkpoint saved: {path}") - if 'slae_embedding' not in batch['protein']: - logger.warning(" WARNING: No SLAE embeddings found in data!") - logger.info(" Please run scripts/precompute_slae_embeddings.py first") - return False - emb = batch['protein'].slae_embedding - logger.info(f" SLAE embedding shape: {emb.shape}") - logger.info(f" SLAE embedding stats: mean={emb.mean():.4f}, std={emb.std():.4f}, min={emb.min():.4f}, max={emb.max():.4f}") +def build_scheduler(optimizer, args): + """ + Build warmup and main learning rate schedulers. - # check if embeddings are all zeros or constant - if emb.std() < 1e-6: - logger.warning(" WARNING: SLAE embeddings appear to be constant/zero!") - return False + Supports hybrid stepping: warmup scheduler steps per optimizer step, + main scheduler steps per epoch after warmup completes. - return True + Args: + optimizer: AdamW optimizer instance + args: Parsed arguments with scheduler configuration -def save_checkpoint(model, optimizer, scheduler, epoch, path, best=False): - """Save model checkpoint.""" - path.parent.mkdir(parents=True, exist_ok=True) - torch.save({ - 'epoch': epoch, - 'model_state_dict': model.state_dict(), - 'optimizer_state_dict': optimizer.state_dict(), - 'scheduler_state_dict': scheduler.state_dict() if scheduler else None, - }, path) - logger.info(f"{'Best ' if best else ''}Checkpoint saved: {path}") + Returns: + Tuple of (warmup_scheduler, main_scheduler), either may be None + """ + # Warmup scheduler (stepped per optimizer step) + warmup_scheduler = None + if args.warmup_steps > 0: + warmup_scheduler = LinearLR( + optimizer, start_factor=1e-8, end_factor=1.0, total_iters=args.warmup_steps + ) + + # Main scheduler (stepped per epoch, after warmup) + main_scheduler = None + if args.scheduler == "cosine": + main_scheduler = CosineAnnealingLR( + optimizer, T_max=args.epochs, eta_min=args.lr * args.eta_min_factor + ) + elif args.scheduler == "step": + main_scheduler = StepLR( + optimizer, step_size=args.step_size, gamma=args.step_gamma + ) + + return warmup_scheduler, main_scheduler def main(): + """Run the full training pipeline.""" args = parse_args() device = torch.device(args.device if torch.cuda.is_available() else "cpu") @@ -335,103 +908,117 @@ def main(): (run_dir / "plots").mkdir(exist_ok=True) (run_dir / "gifs").mkdir(exist_ok=True) - logger.info(f"\n{'='*60}") + log_file = Path(args.log_file) if args.log_file else run_dir / "train.log" + setup_logging_for_tqdm(level=args.log_level, log_file=str(log_file)) + + logger.info("=" * 60) logger.info(f"Run name: {args.run_name}") logger.info(f"Run directory: {run_dir}") - logger.info(f"{'='*60}\n") + logger.info(f"Log file: {log_file}") + logger.info("=" * 60) + + # data loaders + dataset_kwargs, quality_kwargs, _ = _build_dataset_config(args) + _log_dataset_filter_config(args, quality_kwargs) - import json - config_file = run_dir / "config.json" - with open(config_file, 'w') as f: - json.dump(vars(args), f, indent=2) - logger.info(f"Configuration saved to: {config_file}\n") - - # dataloaders train_loader = get_dataloader( - args.train_list, args.processed_dir, - batch_size=args.batch_size, shuffle=True, + pdb_list_file=args.train_list, + processed_dir=args.processed_dir, + batch_size=args.batch_size, + shuffle=True, num_workers=args.num_workers, - base_pdb_dir=args.base_pdb_dir, - include_mates=args.include_mates, + pin_memory=args.pin_memory, + prefetch_factor=args.prefetch_factor, + persistent_workers=args.persistent_workers, duplicate_single_sample=args.duplicate_single_sample, + **dataset_kwargs, ) + val_loader = get_dataloader( - args.val_list, args.processed_dir, - batch_size=args.batch_size, shuffle=False, + pdb_list_file=args.val_list, + processed_dir=args.processed_dir, + batch_size=args.batch_size, + shuffle=False, num_workers=args.num_workers, - base_pdb_dir=args.base_pdb_dir, - include_mates=args.include_mates, - duplicate_single_sample=1, # Don't duplicate for validation + pin_memory=args.pin_memory, + prefetch_factor=args.prefetch_factor, + persistent_workers=args.persistent_workers, + duplicate_single_sample=args.duplicate_single_sample, + **dataset_kwargs, ) - # sample fixed eval indices (same proteins evaluated every epoch) - np.random.seed(42) + # sample fixed eval indices + np.random.seed(42) eval_indices = np.random.choice( len(val_loader.dataset), min(args.n_eval_samples, len(val_loader.dataset)), - replace=False + replace=False, ).tolist() - # save eval indices for reproducibility eval_indices_file = run_dir / "eval_indices.txt" - with open(eval_indices_file, 'w') as f: + with open(eval_indices_file, "w") as f: f.write("# Fixed evaluation sample indices\n") for idx in eval_indices: graph = val_loader.dataset[idx] - pdb_id = getattr(graph, 'pdb_id', 'unknown') + pdb_id = getattr(graph, "pdb_id", "unknown") f.write(f"{idx}\t{pdb_id}\n") logger.info(f"Fixed eval indices saved to: {eval_indices_file}") - logger.info(f"Evaluating on {len(eval_indices)} proteins at each eval epoch\n") + logger.info(f"Evaluating on {len(eval_indices)} proteins at each eval epoch") + + # detect input dimension and resolve encoder configuration from sample data + sample_data = train_loader.dataset[0] + node_scalar_in = int(sample_data["protein"].x.shape[-1]) + logger.info(f"Detected protein input dimension: {node_scalar_in}") + + log_encoder_sample_stats(sample_data, args.encoder_type) + encoder_config = resolve_encoder_config( + args, sample_data, node_scalar_in=node_scalar_in + ) + + config_dict = vars(args).copy() + config_dict["active_water_filters"] = { + "distance": args.filter_by_distance, + "edia": args.filter_by_edia, + "bfactor": args.filter_by_bfactor, + } + config_dict["ignored_water_filter_thresholds"] = _ignored_water_filter_thresholds( + args + ) + config_dict["node_scalar_in"] = node_scalar_in + config_dict["resolved_encoder_config"] = encoder_config + config_file = run_dir / "config.json" + with open(config_file, "w") as f: + json.dump(config_dict, f, indent=2) + logger.info(f"Configuration saved to: {config_file}") wandb.init( project=args.wandb_project, dir=args.wandb_dir, - name=args.wandb_run, - config=vars(args), + name=args.run_name, + config=config_dict, ) - # detect input dimension from data (16 for element_onehot, 37 for atom37 format) - sample_data = train_loader.dataset[0] - node_scalar_in = sample_data['protein'].x.shape[-1] - logger.info(f"Detected protein input dimension: {node_scalar_in}") - - model = build_model(args, device, node_scalar_in=node_scalar_in) + model = build_model(args, device, encoder_config=encoder_config) trainable_params, total_params = count_parameters(model) - logger.info(f"\nModel statistics:") - logger.info(f" Trainable parameters: {trainable_params:,}") - logger.info(f" Total parameters: {total_params:,}") - - # check SLAE embeddings if using SLAE mode - if args.use_slae: - embeddings_ok = check_slae_embeddings(train_loader, device) - if not embeddings_ok: - logger.error("\nERROR: SLAE embeddings are missing or invalid!") - logger.info("Please run: python scripts/precompute_slae_embeddings.py \\") - logger.info(f" --train_list {args.train_list} \\") - logger.info(f" --val_list {args.val_list} \\") - logger.info(f" --processed_dir {args.processed_dir}") - return - - # test forward pass to check if adapter is working - logger.info("\nTesting forward pass with SLAE...") + logger.info("Model statistics:") + logger.info(f"Trainable parameters: {trainable_params:,}") + logger.info(f"Total parameters: {total_params:,}") + + # quick forward pass sanity check for embedding-based encoders + if args.encoder_type in {"slae", "esm"}: + logger.info(f"Testing forward pass with {args.encoder_type.upper()}...") model.eval() batch = next(iter(train_loader)).to(device) with torch.no_grad(): - # determine number of graphs in batch - num_graphs = int(batch['protein'].batch.max().item()) + 1 + num_graphs = int(batch["protein"].batch.max().item()) + 1 t = torch.zeros(num_graphs, device=device) - try: - v_out = model(batch, t) - logger.info(f" Forward pass successful! Output shape: {v_out.shape}") - logger.info(f" Output stats: mean={v_out.mean():.4f}, std={v_out.std():.4f}") - if v_out.std() < 1e-6: - logger.warning(" WARNING: Model output is constant! This indicates a problem.") - except Exception as e: - logger.error(f" ERROR in forward pass: {e}") - return + v_out = model(batch, t) + logger.info(f"Forward pass successful! Output shape: {v_out.shape}") + logger.info(f"Output stats: mean={v_out.mean():.4f}, std={v_out.std():.4f}") + if v_out.std() < 1e-6: + logger.warning("Model output is constant! This indicates a problem.") model.train() - - # flow matcher + flow_matcher = FlowMatcher( model=model, p_self_cond=args.p_self_cond, @@ -440,58 +1027,95 @@ def main(): t_distort=args.t_distort, sigma_distort=args.sigma_distort, ) - - # optimizer & scheduler + optimizer = AdamW( [p for p in model.parameters() if p.requires_grad], - lr=args.lr, weight_decay=args.weight_decay + lr=args.lr, + weight_decay=args.weight_decay, ) - scheduler = CosineAnnealingLR(optimizer, T_max=args.epochs, eta_min=args.lr * 0.01) - - best_val_loss = float('inf') + warmup_scheduler, main_scheduler = build_scheduler(optimizer, args) + + best_val_loss = float("inf") + optimizer_step_count = 0 for epoch in range(1, args.epochs + 1): - - # train - train_metrics, global_step = train_epoch(flow_matcher, train_loader, optimizer, args, epoch) + train_metrics, global_step, optimizer_step_count = train_epoch( + flow_matcher, + train_loader, + optimizer, + warmup_scheduler, + args, + epoch, + optimizer_step_count, + ) + # Log epoch-level metrics with epoch number for per-epoch tracking + train_metrics["epoch"] = epoch wandb.log(train_metrics, step=global_step) - # val val_metrics = val_epoch(flow_matcher, val_loader, args, epoch) + val_metrics["epoch"] = epoch wandb.log(val_metrics, step=global_step) - wandb.log({"lr": scheduler.get_last_lr()[0]}, step=global_step) - scheduler.step() + # Step main scheduler per epoch (after warmup completes) + if main_scheduler is not None and optimizer_step_count >= args.warmup_steps: + main_scheduler.step() - logger.info(f"Epoch {epoch}: train_loss={train_metrics['train/epoch_loss']:.4f}, " - f"val_loss={val_metrics['val/loss']:.4f}, val_rmsd={val_metrics['val/rmsd']:.2f}") + logger.info( + f"Epoch {epoch}: train_loss={train_metrics['train/epoch_loss']:.4f}, " + f"val_loss={val_metrics['val/loss']:.4f}, val_rmsd={val_metrics['val/rmsd']:.2f}" + ) - # save best - if val_metrics['val/loss'] < best_val_loss: - best_val_loss = val_metrics['val/loss'] - save_checkpoint(model, optimizer, scheduler, epoch, - run_dir / "checkpoints" / "best.pt", best=True) + if val_metrics["val/loss"] < best_val_loss: + best_val_loss = val_metrics["val/loss"] + save_checkpoint( + model, + optimizer, + warmup_scheduler, + main_scheduler, + epoch, + optimizer_step_count, + run_dir / "checkpoints" / "best.pt", + best=True, + ) - # periodic save if epoch % args.save_every == 0: - save_checkpoint(model, optimizer, scheduler, epoch, - run_dir / "checkpoints" / f"epoch_{epoch}.pt") + save_checkpoint( + model, + optimizer, + warmup_scheduler, + main_scheduler, + epoch, + optimizer_step_count, + run_dir / "checkpoints" / f"epoch_{epoch}.pt", + ) - # eval sampling if epoch % args.eval_every == 0: eval_metrics = run_eval_sampling( - flow_matcher, val_loader, args, epoch, device, global_step, eval_indices, run_dir + flow_matcher, + val_loader, + args, + epoch, + device, + global_step, + eval_indices, + run_dir, ) if eval_metrics: - logger.info(f" Eval: RMSD={eval_metrics['eval/avg_rmsd']:.2f}Å, " - f"Precision={eval_metrics['eval/avg_precision']:.2%}, " - f"Recall={eval_metrics['eval/avg_recall']:.2%}, " - f"F1={eval_metrics['eval/avg_f1']:.3f}, " - f"AUC-PR={eval_metrics['eval/avg_auc_pr']:.3f}") - + logger.info( + f"Eval: RMSD={eval_metrics['eval/avg_rmsd']:.2f}A, " + f"Precision={eval_metrics['eval/avg_precision']:.2%}, " + f"Recall={eval_metrics['eval/avg_recall']:.2%}, " + f"F1={eval_metrics['eval/avg_f1']:.3f}, " + f"AUC-PR={eval_metrics['eval/avg_auc_pr']:.3f}" + ) + wandb.finish() logger.info("Training complete.") if __name__ == "__main__": - main() + try: + main() + except Exception: + logger.exception("Training failed with an unhandled exception.") + raise From a1ae9369e006ba23e30d12fdf10e7e24cb6a34e9 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Wed, 4 Mar 2026 16:08:44 -0600 Subject: [PATCH 2/3] addressing changes in inference script --- scripts/inference.py | 122 +++++++++++++++++++++++++++---------------- 1 file changed, 76 insertions(+), 46 deletions(-) diff --git a/scripts/inference.py b/scripts/inference.py index 6f83519..be3f466 100644 --- a/scripts/inference.py +++ b/scripts/inference.py @@ -85,7 +85,9 @@ def parse_args(): "--processed_dir", type=str, default="/home/srivasv/flow_cache/", - help="Directory for cached preprocessed .pt files", + help="Parent directory containing all cached preprocessed data. This directory holds " + "protein graphs, embeddings, and geometry subdirectories (e.g., processed_dir/geometry/). " + "Each PDB's preprocessed data is stored in subdirectories organized by cache type.", ) p.add_argument( "--base_pdb_dir", @@ -102,8 +104,9 @@ def parse_args(): "--geometry_cache", type=str, default=None, - help="Geometry cache name to use (e.g., 'geometry' or 'geometry_unfiltered'). " - "Overrides the model's config if specified. Use this to evaluate against a specific ground truth.", + help="Subdirectory name within processed_dir specifying which water coordinate set to use. " + "Options include 'geometry' (filtered waters meeting quality criteria) or " + "'geometry_unfiltered' (all crystallographic waters). Overrides the model's config if specified.", ) # checkpoint arguments @@ -168,6 +171,14 @@ def parse_args(): "E.g., --water_ratio 0.5 samples 50 waters for a 100-residue protein.", ) + p.add_argument( + "--skip_metrics", + action="store_true", + help="Skip metrics computation (precision, recall, RMSD) that require ground truth. " + "Use when running inference on structures without ground truth waters. " + "Automatically enabled when --water_ratio is specified.", + ) + args = p.parse_args() return args @@ -316,25 +327,13 @@ def run_inference_batch( water_ratio=water_ratio, ) else: # euler - water_preds = flow_matcher.euler_integrate( + results = flow_matcher.euler_integrate( graphs, num_steps=num_steps, use_sc=use_sc, device=device, water_ratio=water_ratio, ) - # build result dicts similar to rk4 - results = [] - for graph, water_pred in zip(graphs, water_preds): - results.append( - { - "protein_pos": graph["protein"].pos.numpy(), - "water_true": graph["water"].pos.numpy(), - "water_pred": water_pred, - "trajectory": None, - "pdb_id": getattr(graph, "pdb_id", None), - } - ) return results @@ -343,7 +342,7 @@ def save_plot( result: dict, pdb_id: str, output_path: Path, - metrics: dict, + metrics: dict | None, ): """ Save 3D visualization plot of water prediction results. @@ -352,22 +351,31 @@ def save_plot( result: Dict with 'protein_pos', 'water_pred', 'water_true' arrays pdb_id: PDB identifier for title output_path: Path to save PNG image - metrics: Dict with 'rmsd', 'precision', 'recall', 'f1' for title + metrics: Dict with 'rmsd', 'precision', 'recall', 'f1' for title, or None if no ground truth """ fig = plt.figure(figsize=(12, 10)) ax = fig.add_subplot(111, projection="3d") - title = ( - f"{pdb_id} | RMSD={metrics['rmsd']:.2f}Å | " - f"P={metrics['precision']:.2%} R={metrics['recall']:.2%} F1={metrics['f1']:.3f}" - ) + if metrics is not None: + title = ( + f"{pdb_id} | RMSD={metrics['rmsd']:.2f}Å | " + f"P={metrics['precision']:.2%} R={metrics['recall']:.2%} F1={metrics['f1']:.3f}" + ) + else: + n_pred = result["water_pred"].shape[0] + title = f"{pdb_id} | {n_pred} waters predicted (no ground truth)" + + # water_true may be None or empty when no ground truth is available + water_true = result.get("water_true") + if water_true is not None and water_true.shape[0] == 0: + water_true = None plot_3d_frame( ax, result["protein_pos"], None, # no separate mate positions result["water_pred"], - result["water_true"], + water_true, title=title, ) @@ -440,22 +448,32 @@ def main(): logger.info(f"Self-conditioning: {args.use_sc}") logger.info(f"Threshold for metrics: {args.threshold}Å") logger.info(f"Batch size: {args.batch_size}") + + # Determine if metrics should be skipped + # Skip metrics when explicitly requested or when using water_ratio (no ground truth count) + skip_metrics = args.skip_metrics or args.water_ratio is not None + if args.water_ratio is not None: logger.info( f"Water ratio: {args.water_ratio} (sampling num_residues × {args.water_ratio} waters)" ) else: logger.info("Water ratio: None (using ground truth water count)") + if skip_metrics: + logger.info("Metrics computation: DISABLED (no ground truth comparison)") + else: + logger.info("Metrics computation: ENABLED") logger.info("-" * 60) all_metrics = [] - # collect valid graphs (those with waters for ground truth comparison) + # collect graphs for inference valid_graphs = [] skipped_pdbs = [] for idx in range(len(dataset)): graph = dataset[idx] - if graph["water"].num_nodes == 0: + # Only skip zero-water PDBs when we need ground truth for metrics + if graph["water"].num_nodes == 0 and not skip_metrics: skipped_pdbs.append(graph.pdb_id) else: valid_graphs.append(graph) @@ -487,19 +505,23 @@ def main(): pdb_id = result.get("pdb_id", f"unknown_{len(all_metrics)}") water_pred = result["water_pred"] water_true = result["water_true"] - - # compute metrics - metrics = compute_placement_metrics( - pred=water_pred, - true=water_true, - threshold=args.threshold, - ) - metrics["rmsd"] = compute_rmsd(water_pred, water_true) - metrics["pdb_id"] = pdb_id - metrics["n_waters_true"] = water_true.shape[0] - metrics["n_waters_pred"] = water_pred.shape[0] - - all_metrics.append(metrics) + has_ground_truth = water_true is not None and water_true.shape[0] > 0 + + # compute metrics only when ground truth is available and metrics are enabled + if not skip_metrics and has_ground_truth: + metrics = compute_placement_metrics( + pred=water_pred, + true=water_true, + threshold=args.threshold, + ) + metrics["rmsd"] = compute_rmsd(water_pred, water_true) + metrics["pdb_id"] = pdb_id + metrics["n_waters_true"] = water_true.shape[0] + metrics["n_waters_pred"] = water_pred.shape[0] + all_metrics.append(metrics) + else: + # no ground truth comparison - just store prediction info + metrics = None plot_path = output_dir / "plots" / f"{pdb_id}.png" save_plot(result, pdb_id, plot_path, metrics) @@ -507,22 +529,27 @@ def main(): # save GIF if requested and trajectory available if args.save_gifs and result.get("trajectory") is not None: gif_path = output_dir / "gifs" / f"{pdb_id}.gif" + # Use water_true only if available + gif_water_true = water_true if has_ground_truth else None create_trajectory_gif( trajectory=result["trajectory"], protein_pos=result["protein_pos"], - water_true=water_true, + water_true=gif_water_true, save_path=str(gif_path), title="", fps=10, pdb_id=pdb_id, ) - # print per-sample metrics - tqdm.write( - f" {pdb_id}: RMSD={metrics['rmsd']:.2f}Å | " - f"P={metrics['precision']:.2%} R={metrics['recall']:.2%} " - f"F1={metrics['f1']:.3f} AUC-PR={metrics['auc_pr']:.3f}" - ) + # print per-sample info + if metrics is not None: + tqdm.write( + f" {pdb_id}: RMSD={metrics['rmsd']:.2f}Å | " + f"P={metrics['precision']:.2%} R={metrics['recall']:.2%} " + f"F1={metrics['f1']:.3f} AUC-PR={metrics['auc_pr']:.3f}" + ) + else: + tqdm.write(f" {pdb_id}: {water_pred.shape[0]} waters predicted") # compute and save summary metrics if all_metrics: @@ -580,7 +607,10 @@ def main(): logger.info(f"Metrics saved to: {metrics_path}") else: - logger.warning("No valid samples processed.") + if skip_metrics: + logger.info(f"Processed {len(valid_graphs)} samples (metrics disabled)") + else: + logger.warning("No valid samples processed.") logger.info(f"Plots saved to: {output_dir / 'plots'}") if args.save_gifs: From 290795284ce6c2e85a01589f333ea39e9f559956 Mon Sep 17 00:00:00 2001 From: Vratin Srivastava Date: Tue, 10 Mar 2026 09:48:38 -0500 Subject: [PATCH 3/3] addressing PR comments for train.py --- scripts/train.py | 22 +++++++--------------- 1 file changed, 7 insertions(+), 15 deletions(-) diff --git a/scripts/train.py b/scripts/train.py index 9f97e08..84e3a46 100644 --- a/scripts/train.py +++ b/scripts/train.py @@ -259,27 +259,13 @@ def parse_args(): p.add_argument( "--pin_memory", action="store_true", - default=True, help="Pin memory for faster CPU-GPU transfer", ) - p.add_argument( - "--no_pin_memory", - dest="pin_memory", - action="store_false", - help="Disable pin_memory", - ) p.add_argument( "--persistent_workers", action="store_true", - default=True, help="Keep workers alive between epochs", ) - p.add_argument( - "--no_persistent_workers", - dest="persistent_workers", - action="store_false", - help="Disable persistent_workers", - ) # scheduler p.add_argument( @@ -320,6 +306,12 @@ def parse_args(): p.add_argument( "--save_gifs", action="store_true", help="Save trajectory GIFs during eval" ) + p.add_argument( + "--threshold", + type=float, + default=1.0, + help="Distance threshold in Angstroms for precision/recall (default: 1.0)", + ) # logging / wandb p.add_argument("--log_level", type=str, default="INFO") @@ -589,7 +581,7 @@ def run_eval_sampling( # compute metrics final_metrics = compute_placement_metrics( - pred=out["water_pred"], true=out["water_true"], threshold=1.0 + pred=out["water_pred"], true=out["water_true"], threshold=args.threshold ) final_rmsd = compute_rmsd(out["water_pred"], out["water_true"])