diff --git a/.gitignore b/.gitignore index a2b3344..70bf226 100644 --- a/.gitignore +++ b/.gitignore @@ -70,6 +70,7 @@ training/data/rcsb_pdbs training/logs/ training/checkpoints_direct_feat/ training/checkpoints_delta_feat/ +training/figures/ training/models/*.pt training/models/*.npz training/models/*.json diff --git a/CITATION.cff b/CITATION.cff index 9cb6fcc..e2d91fe 100644 --- a/CITATION.cff +++ b/CITATION.cff @@ -7,8 +7,8 @@ authors: given-names: "Samuel" email: "samuels.lobo@gmail.com" doi: 10.5281/zenodo.19744336 -version: "0.1.2" -date-released: 2026-04-24 +version: "0.1.3" +date-released: 2026-05-28 url: "https://github.com/samlobe/FastHydroMap" repository-code: "https://github.com/samlobe/FastHydroMap" license: "MIT" diff --git a/README.md b/README.md index b00dc8b..b11778d 100644 --- a/README.md +++ b/README.md @@ -5,6 +5,7 @@ [](https://doi.org/10.5281/zenodo.19744336) FastHydroMap predicts per-residue dewetting free energies (`Fdewet`) from protein structures and trajectories. +It can also predict water structuring (`PC1`, `PC2`, `PC3` of the [water triplet angle distribution](https://doi.org/10.1021/acs.jpcb.3c00826)).
/dev/null
fasthydromap install-torch --variant cpu
fasthydromap predict examples/1A1U.pdb -o "${OUTDIR}/1A1U_fdewet"
+fasthydromap predict examples/1A1U.pdb --quantity pc1 -o "${OUTDIR}/1A1U_pc1"
diff --git a/scripts/test_wheel_install.sh b/scripts/test_wheel_install.sh
index 22a9192..da0175c 100755
--- a/scripts/test_wheel_install.sh
+++ b/scripts/test_wheel_install.sh
@@ -21,9 +21,12 @@ fasthydromap install-torch --dry-run
fasthydromap install-torch --variant cpu
OUTROOT="${TEST_ENV}/smoke_1A1U"
fasthydromap predict "${ROOT_DIR}/examples/1A1U.pdb" -o "${OUTROOT}"
+fasthydromap predict "${ROOT_DIR}/examples/1A1U.pdb" --quantity pc1 -o "${TEST_ENV}/smoke_1A1U_pc1"
test -f "${OUTROOT}.csv"
test -f "${OUTROOT}.pdb"
+test -f "${TEST_ENV}/smoke_1A1U_pc1.csv"
+test -f "${TEST_ENV}/smoke_1A1U_pc1.pdb"
echo
echo "Wheel install smoke test passed in ${TEST_ENV}"
diff --git a/src/FastHydroMap/cli.py b/src/FastHydroMap/cli.py
index 1f5ad5f..b01b92f 100644
--- a/src/FastHydroMap/cli.py
+++ b/src/FastHydroMap/cli.py
@@ -9,13 +9,21 @@
from .io.pdb import write_bfactor
from .install_torch import install_torch, torch_install_command
+QUANTITIES = ("fdewet", "pc1", "pc2", "pc3")
+QUANTITY_LABELS = {
+ "fdewet": "Fdewet",
+ "pc1": "PC1",
+ "pc2": "PC2",
+ "pc3": "PC3",
+}
+
# ---------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------
def _build_parser() -> argparse.ArgumentParser:
ap = argparse.ArgumentParser(prog="fasthydromap",
- description="FastHydroMap – infer per-residue Fdewet")
+ description="FastHydroMap – infer per-residue Fdewet or PC maps")
sub = ap.add_subparsers(dest="cmd", required=True)
# -------- predict -------------------------------------------------
@@ -31,6 +39,12 @@ def _build_parser() -> argparse.ArgumentParser:
action="store_true",
help="for single-structure predictions, include intrinsic and context columns",
)
+ p.add_argument(
+ "--quantity",
+ choices=QUANTITIES,
+ default="fdewet",
+ help="per-residue quantity to predict (default: fdewet)",
+ )
# -------- predict-trajectory --------------------------------------
pt = sub.add_parser(
@@ -51,6 +65,12 @@ def _build_parser() -> argparse.ArgumentParser:
action="store_true",
help="write intrinsic, context, and per-frame summary CSVs in addition to total",
)
+ pt.add_argument(
+ "--quantity",
+ choices=QUANTITIES,
+ default="fdewet",
+ help="per-residue quantity to predict (default: fdewet)",
+ )
# -------- install-torch ------------------------------------------
it = sub.add_parser(
@@ -129,7 +149,9 @@ def main() -> None:
return
FdewetPredictor = _load_predictor_or_exit()
- predictor = FdewetPredictor()
+ predictor = FdewetPredictor(quantity=args.quantity)
+ quantity_label = QUANTITY_LABELS[args.quantity]
+ quantity_suffix = args.quantity
if args.cmd == "predict":
if args.parts and args.dcd is not None:
@@ -161,7 +183,7 @@ def main() -> None:
# -------------------------------------------------------------
outroot = args.outroot or args.pdb.with_suffix("")
if args.outroot is None:
- outroot = outroot.with_name(outroot.name + "_fdewet")
+ outroot = outroot.with_name(outroot.name + f"_{quantity_suffix}")
csv_path = Path(f"{outroot}.csv")
pdb_path = Path(f"{outroot}.pdb")
@@ -170,11 +192,11 @@ def main() -> None:
if args.dcd is None:
data = {
"residue": [str(r) for r in res_ids],
- "Fdewet": scores.round(2),
+ quantity_label: scores.round(2),
}
if args.parts:
- data["Fdewet_intrinsic"] = parts["intrinsic"].round(2)
- data["Fdewet_context"] = parts["context"].round(2)
+ data[f"{quantity_label}_intrinsic"] = parts["intrinsic"].round(2)
+ data[f"{quantity_label}_context"] = parts["context"].round(2)
df = pd.DataFrame(data)
else:
col_names = [str(r) for r in res_ids] # columns = residue numbers
@@ -203,7 +225,7 @@ def main() -> None:
outroot = args.outroot or args.dcd.with_suffix("")
if args.outroot is None:
- outroot = outroot.with_name(outroot.name + "_fdewet_traj")
+ outroot = outroot.with_name(outroot.name + f"_{quantity_suffix}_traj")
total_csv = Path(f"{outroot}_total.csv")
intrinsic_csv = Path(f"{outroot}_intrinsic.csv")
diff --git a/src/FastHydroMap/predictors/fdewet.py b/src/FastHydroMap/predictors/fdewet.py
index 895136e..0d9711b 100644
--- a/src/FastHydroMap/predictors/fdewet.py
+++ b/src/FastHydroMap/predictors/fdewet.py
@@ -24,6 +24,52 @@
from ..utils.atom_names import backbone_alias_priority, canonical_backbone_atom_name
WEIGHT_DIR = Path(__file__).parents[1] / "weights"
+QUANTITY_SPECS = {
+ "fdewet": {
+ "label": "Fdewet",
+ "weight": WEIGHT_DIR / "mpnn_latest.pt",
+ "k_nn": 12,
+ "n_rbf": 3,
+ "rbf_min": 2.0,
+ "rbf_max": 14.0,
+ "rbf_sigma": 4.0,
+ },
+ "pc1": {
+ "label": "PC1",
+ "weight": WEIGHT_DIR / "mpnn_pc1_latest.pt",
+ "k_nn": 12,
+ "n_rbf": 3,
+ "rbf_min": 2.0,
+ "rbf_max": 14.0,
+ "rbf_sigma": 4.0,
+ },
+ "pc2": {
+ "label": "PC2",
+ "weight": WEIGHT_DIR / "mpnn_pc2_latest.pt",
+ "k_nn": 12,
+ "n_rbf": 3,
+ "rbf_min": 2.0,
+ "rbf_max": 14.0,
+ "rbf_sigma": 4.0,
+ },
+ "pc3": {
+ "label": "PC3",
+ "weight": WEIGHT_DIR / "mpnn_pc3_latest.pt",
+ "k_nn": 12,
+ "n_rbf": 3,
+ "rbf_min": 2.0,
+ "rbf_max": 14.0,
+ "rbf_sigma": 4.0,
+ },
+}
+
+
+def normalize_quantity(quantity: str) -> str:
+ q = quantity.lower()
+ if q in QUANTITY_SPECS:
+ return q
+ raise ValueError(f"unknown FastHydroMap quantity {quantity!r}")
+
def _display_residue_label(
chain_id: str,
@@ -42,15 +88,26 @@ def _display_residue_label(
class FdewetPredictor:
def __init__(
self,
+ quantity: str = "fdewet",
k_nn: int = 12,
n_rbf: int = 3,
rbf_min: float = 2.0,
rbf_max: float = 14.0,
rbf_sigma: float = 4.0,
- mpnn_pt: Path = WEIGHT_DIR / "mpnn_latest.pt",
+ mpnn_pt: Path | None = None,
sasa_stats_npz: Path = WEIGHT_DIR / "sasa_feature_stats.npz",
device: str | torch.device | None = None,
):
+ self.quantity = normalize_quantity(quantity)
+ spec = QUANTITY_SPECS[self.quantity]
+ if mpnn_pt is None:
+ mpnn_pt = spec["weight"]
+ k_nn = int(spec["k_nn"])
+ n_rbf = int(spec["n_rbf"])
+ rbf_min = float(spec["rbf_min"])
+ rbf_max = float(spec["rbf_max"])
+ rbf_sigma = float(spec["rbf_sigma"])
+ self.output_label = str(spec["label"])
self.k = k_nn
self.n_rbf = n_rbf
self.rbf_min = rbf_min
diff --git a/src/FastHydroMap/weights/mpnn_pc1_latest.pt b/src/FastHydroMap/weights/mpnn_pc1_latest.pt
new file mode 100644
index 0000000..166c212
Binary files /dev/null and b/src/FastHydroMap/weights/mpnn_pc1_latest.pt differ
diff --git a/src/FastHydroMap/weights/mpnn_pc2_latest.pt b/src/FastHydroMap/weights/mpnn_pc2_latest.pt
new file mode 100644
index 0000000..a088b46
Binary files /dev/null and b/src/FastHydroMap/weights/mpnn_pc2_latest.pt differ
diff --git a/src/FastHydroMap/weights/mpnn_pc3_latest.pt b/src/FastHydroMap/weights/mpnn_pc3_latest.pt
new file mode 100644
index 0000000..1a427c5
Binary files /dev/null and b/src/FastHydroMap/weights/mpnn_pc3_latest.pt differ
diff --git a/tests/test_predictor_regression.py b/tests/test_predictor_regression.py
index 98d0fe1..5667ea2 100644
--- a/tests/test_predictor_regression.py
+++ b/tests/test_predictor_regression.py
@@ -145,6 +145,22 @@ def test_cli_predict_writes_outputs_for_path_outroot(monkeypatch, tmp_path):
assert np.max(np.abs(out_scores - ref_scores)) < 0.5
+def test_cli_predict_can_write_pc_quantity(monkeypatch, tmp_path):
+ outroot = tmp_path / "1A1U_pc1"
+ monkeypatch.setattr(
+ sys,
+ "argv",
+ ["fasthydromap", "predict", str(PDB_PATH), "-o", str(outroot), "--quantity", "pc1"],
+ )
+
+ main()
+
+ df = pd.read_csv(Path(f"{outroot}.csv"))
+ assert list(df.columns) == ["residue", "PC1"]
+ assert len(df) > 0
+ assert np.isfinite(df["PC1"].to_numpy(np.float32)).all()
+
+
def test_cli_predict_parts_writes_single_structure_decomposition(monkeypatch, tmp_path):
outroot = tmp_path / "1A1U_fdewet_parts"
monkeypatch.setattr(
diff --git a/training/03_train_mpnn.py b/training/03_train_mpnn.py
new file mode 100644
index 0000000..76256dd
--- /dev/null
+++ b/training/03_train_mpnn.py
@@ -0,0 +1,394 @@
+#!/usr/bin/env python3
+"""Train FastHydroMap direct MPNN regressors."""
+
+from __future__ import annotations
+
+import argparse
+import json
+import shutil
+from pathlib import Path
+
+import torch
+from torch_geometric.loader import DataLoader
+
+from train_mpnn_common import (
+ DEVICE,
+ GraphDSDirect,
+ SUPPORTED_MASK_SOURCES,
+ SUPPORTED_TARGETS,
+ TrainConfig,
+ compute_winsor_bounds,
+ edge_dim_from_cfg,
+ evaluate,
+ graph_cache_paths,
+ load_meta_and_splits,
+ make_model,
+ make_optimizer,
+ masked_mean_target,
+ set_seed,
+ split_indices,
+ train_one_epoch,
+)
+
+ROOT = Path(__file__).resolve().parent
+CKPT_DIR = ROOT / "checkpoints_direct_feat"
+MODEL_DIR = ROOT / "models"
+PKG_WEIGHT_DIR = ROOT.parent / "src" / "FastHydroMap" / "weights"
+CKPT_DIR.mkdir(exist_ok=True)
+MODEL_DIR.mkdir(exist_ok=True)
+PKG_WEIGHT_DIR.mkdir(exist_ok=True)
+
+TARGET_PACKAGE_WEIGHTS = {
+ "Fdewet_pred": "mpnn_latest.pt",
+ "PC1": "mpnn_pc1_latest.pt",
+ "PC2": "mpnn_pc2_latest.pt",
+ "PC3": "mpnn_pc3_latest.pt",
+}
+
+
+def _tag_float(x: float) -> str:
+ if float(x).is_integer():
+ return str(int(x))
+ return str(x).replace("-", "m").replace(".", "p")
+
+
+def model_tag(cfg: TrainConfig) -> str:
+ effective_mask = cfg.mask_source
+ if effective_mask == "auto":
+ effective_mask = "fdewet" if cfg.target == "Fdewet_pred" else "trusted"
+
+ graph = (
+ f"k{cfg.k_nn}_rbf{cfg.n_rbf}_"
+ f"r{_tag_float(cfg.rbf_min)}to{_tag_float(cfg.rbf_max)}_"
+ f"s{_tag_float(cfg.rbf_sigma)}"
+ )
+ model = f"{graph}_h{cfg.hidden}_d{cfg.depth}_head{cfg.head_hidden}"
+
+ if cfg.target == "Fdewet_pred" and effective_mask == "fdewet":
+ tag = model
+ else:
+ tag = f"{cfg.target.lower()}_{effective_mask}_{model}"
+
+ if cfg.winsor_lower is not None:
+ tag = f"{tag}_winsor_p{int(cfg.winsor_lower * 100)}_p{int(cfg.winsor_upper * 100)}"
+ return tag
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser(description=__doc__)
+ p.add_argument("--stage", choices=("val", "prod"), default="val")
+ p.add_argument("--seed", type=int, default=48)
+ p.add_argument("--target", choices=SUPPORTED_TARGETS, default="Fdewet_pred")
+ p.add_argument("--mask-source", choices=SUPPORTED_MASK_SOURCES, default="auto")
+ p.add_argument("--k-nn", type=int, default=TrainConfig.k_nn)
+ p.add_argument("--n-rbf", type=int, default=TrainConfig.n_rbf)
+ p.add_argument("--rbf-min", type=float, default=TrainConfig.rbf_min)
+ p.add_argument("--rbf-max", type=float, default=TrainConfig.rbf_max)
+ p.add_argument("--rbf-sigma", type=float, default=TrainConfig.rbf_sigma)
+ p.add_argument("--hidden", type=int, default=TrainConfig.hidden)
+ p.add_argument("--depth", type=int, default=TrainConfig.depth)
+ p.add_argument("--head-hidden", type=int, default=TrainConfig.head_hidden)
+ p.add_argument("--dropout", type=float, default=TrainConfig.dropout)
+ p.add_argument("--edge-drop", type=float, default=TrainConfig.edge_drop)
+ p.add_argument("--weight-decay", type=float, default=TrainConfig.weight_decay)
+ p.add_argument("--winsor-lower", type=float, default=None)
+ p.add_argument("--winsor-upper", type=float, default=None)
+ p.add_argument(
+ "--epochs",
+ type=int,
+ default=None,
+ help="production epochs; if omitted, use best_epoch from the validation summary",
+ )
+ p.add_argument("--report-test", action="store_true", help="evaluate the held-out test split")
+ p.add_argument(
+ "--copy-to-package",
+ action="store_true",
+ help="for production training, copy the resulting weight to src/FastHydroMap/weights",
+ )
+ return p.parse_args()
+
+
+def config_from_args(args: argparse.Namespace) -> TrainConfig:
+ return TrainConfig(
+ seed=args.seed,
+ target=args.target,
+ mask_source=args.mask_source,
+ k_nn=args.k_nn,
+ n_rbf=args.n_rbf,
+ rbf_min=args.rbf_min,
+ rbf_max=args.rbf_max,
+ rbf_sigma=args.rbf_sigma,
+ hidden=args.hidden,
+ depth=args.depth,
+ head_hidden=args.head_hidden,
+ dropout=args.dropout,
+ edge_drop=args.edge_drop,
+ weight_decay=args.weight_decay,
+ winsor_lower=args.winsor_lower,
+ winsor_upper=args.winsor_upper,
+ )
+
+
+def config_summary(cfg: TrainConfig) -> dict[str, float | int | str | None]:
+ return {
+ "k_nn": cfg.k_nn,
+ "n_rbf": cfg.n_rbf,
+ "rbf_min": cfg.rbf_min,
+ "rbf_max": cfg.rbf_max,
+ "rbf_sigma": cfg.rbf_sigma,
+ "hidden": cfg.hidden,
+ "depth": cfg.depth,
+ "head_hidden": cfg.head_hidden,
+ "dropout": cfg.dropout,
+ "edge_drop": cfg.edge_drop,
+ "weight_decay": cfg.weight_decay,
+ "winsor_lower": cfg.winsor_lower,
+ "winsor_upper": cfg.winsor_upper,
+ }
+
+
+def load_training_inputs(cfg: TrainConfig):
+ graph_pt, pid_pt = graph_cache_paths(cfg)
+ if not graph_pt.exists() or not pid_pt.exists():
+ raise FileNotFoundError(
+ f"missing graph cache: {graph_pt.name} / {pid_pt.name}. "
+ "Build it first with 02_build_mpnn_graphs.py using matching graph settings."
+ )
+ meta, splits = load_meta_and_splits()
+ return graph_pt, pid_pt, meta, splits
+
+
+def make_dataset_and_loaders(
+ cfg: TrainConfig,
+ split_ids: dict[str, list[str]],
+ *,
+ clip_from_split: str,
+ shuffle_split: str,
+):
+ graph_pt, pid_pt, meta, _ = load_training_inputs(cfg)
+ split_ix, split_ids_kept = split_indices(pid_pt, split_ids)
+ target_clip_bounds = compute_winsor_bounds(
+ meta,
+ split_ids_kept[clip_from_split],
+ cfg.target,
+ cfg.mask_source,
+ cfg.winsor_lower,
+ cfg.winsor_upper,
+ )
+ mu = masked_mean_target(
+ meta,
+ split_ids_kept[clip_from_split],
+ cfg.target,
+ cfg.mask_source,
+ target_clip_bounds,
+ )
+ ds = GraphDSDirect(
+ graph_pt,
+ meta,
+ target=cfg.target,
+ mask_source=cfg.mask_source,
+ target_clip_bounds=target_clip_bounds,
+ )
+ loaders = {}
+ for name, indices in split_ix.items():
+ loaders[name] = DataLoader(
+ ds[indices],
+ cfg.batch_size,
+ shuffle=(name == shuffle_split),
+ num_workers=4,
+ )
+ return loaders, graph_pt, mu, target_clip_bounds
+
+
+def base_summary(
+ cfg: TrainConfig,
+ *,
+ stage: str,
+ tag: str,
+ graph_pt: Path,
+ mu: float,
+ target_clip_bounds: tuple[float, float] | None,
+) -> dict:
+ return {
+ "stage": stage,
+ "seed": cfg.seed,
+ "target": cfg.target,
+ "mask_source": cfg.mask_source,
+ "tag": tag,
+ "mu": float(mu),
+ "edge_dim": edge_dim_from_cfg(cfg),
+ "target_clip_bounds": list(target_clip_bounds) if target_clip_bounds is not None else None,
+ "config": config_summary(cfg),
+ "graph_pt": str(graph_pt),
+ }
+
+
+def run_validation(cfg: TrainConfig, *, report_test: bool) -> None:
+ _, _, _, splits = load_training_inputs(cfg)
+ split_ids = {"train": splits["train"], "val": splits["val"], "test": splits["test"]}
+ loaders, graph_pt, mu, target_clip_bounds = make_dataset_and_loaders(
+ cfg,
+ split_ids,
+ clip_from_split="train",
+ shuffle_split="train",
+ )
+
+ tag = model_tag(cfg)
+ best_pt = CKPT_DIR / f"best_{tag}_seed{cfg.seed}.pt"
+ summary_json = MODEL_DIR / f"03_train_mpnn_val_{tag}_seed{cfg.seed}.json"
+
+ model = make_model(cfg, mu=mu, edge_dim=edge_dim_from_cfg(cfg))
+ opt = make_optimizer(model, cfg)
+
+ best_val = float("inf")
+ best_epoch = 0
+ wait = 0
+ for epoch in range(1, cfg.epochs + 1):
+ train_one_epoch(model, loaders["train"], opt, cfg)
+ val_rmse, val_r, val_rho, _ = evaluate(model, loaders["val"])
+ print(f"[val] epoch {epoch:02d} rmse {val_rmse:.4f} | r {val_r:.3f} | rho {val_rho:.3f}", flush=True)
+ if val_rmse + 1e-4 < best_val:
+ best_val = val_rmse
+ best_epoch = epoch
+ wait = 0
+ torch.save(model.state_dict(), best_pt)
+ print(f"[val] new best -> {best_pt.name}", flush=True)
+ else:
+ wait += 1
+ if wait >= cfg.patience:
+ print("[val] early stop", flush=True)
+ break
+
+ model.load_state_dict(torch.load(best_pt, map_location=DEVICE))
+ val_rmse, val_r, val_rho, val_stats = evaluate(model, loaders["val"])
+ summary = base_summary(
+ cfg,
+ stage="val",
+ tag=tag,
+ graph_pt=graph_pt,
+ mu=mu,
+ target_clip_bounds=target_clip_bounds,
+ )
+ summary.update(
+ {
+ "best_epoch": best_epoch,
+ "best_val_rmse": float(best_val),
+ "val_rmse_reloaded": float(val_rmse),
+ "val_r": float(val_r),
+ "val_rho": float(val_rho),
+ "best_ckpt": str(best_pt),
+ "val_stats": {k: float(v) for k, v in val_stats.items()},
+ }
+ )
+
+ if report_test:
+ test_rmse, test_r, test_rho, test_stats = evaluate(model, loaders["test"])
+ summary.update(
+ {
+ "test_rmse": float(test_rmse),
+ "test_r": float(test_r),
+ "test_rho": float(test_rho),
+ "test_stats": {k: float(v) for k, v in test_stats.items()},
+ }
+ )
+ print(f"[test] rmse {test_rmse:.4f} | r {test_r:.3f} | rho {test_rho:.3f}", flush=True)
+
+ summary_json.write_text(json.dumps(summary, indent=2) + "\n")
+ print(f"[done] best val rmse {best_val:.4f} at epoch {best_epoch}", flush=True)
+ print(f"[done] summary -> {summary_json}", flush=True)
+
+
+def run_production(
+ cfg: TrainConfig,
+ *,
+ epochs: int | None,
+ report_test: bool,
+ copy_to_package: bool,
+) -> None:
+ _, _, _, splits = load_training_inputs(cfg)
+ split_ids = {
+ "trainval": [*splits["train"], *splits["val"]],
+ "test": splits["test"],
+ }
+ loaders, graph_pt, mu, target_clip_bounds = make_dataset_and_loaders(
+ cfg,
+ split_ids,
+ clip_from_split="trainval",
+ shuffle_split="trainval",
+ )
+
+ tag = model_tag(cfg)
+ if epochs is None:
+ val_summary = MODEL_DIR / f"03_train_mpnn_val_{tag}_seed{cfg.seed}.json"
+ if not val_summary.exists():
+ raise FileNotFoundError(f"missing validation summary for default epochs: {val_summary}")
+ epochs = int(json.loads(val_summary.read_text())["best_epoch"])
+
+ model = make_model(cfg, mu=mu, edge_dim=edge_dim_from_cfg(cfg))
+ opt = make_optimizer(model, cfg)
+ local_pt = MODEL_DIR / f"mpnn_direct_prod_{tag}.pt"
+ summary_json = MODEL_DIR / f"03_train_mpnn_prod_{tag}_seed{cfg.seed}.json"
+
+ for epoch in range(1, epochs + 1):
+ train_one_epoch(model, loaders["trainval"], opt, cfg)
+ print(f"[prod] epoch {epoch:02d}/{epochs}", flush=True)
+
+ torch.save(model.state_dict(), local_pt)
+
+ package_weight = None
+ if copy_to_package:
+ package_weight = PKG_WEIGHT_DIR / TARGET_PACKAGE_WEIGHTS[cfg.target]
+ shutil.copy2(local_pt, package_weight)
+
+ summary = base_summary(
+ cfg,
+ stage="prod",
+ tag=tag,
+ graph_pt=graph_pt,
+ mu=mu,
+ target_clip_bounds=target_clip_bounds,
+ )
+ summary.update(
+ {
+ "epochs": epochs,
+ "local_weight": str(local_pt),
+ "package_weight": str(package_weight) if package_weight is not None else None,
+ }
+ )
+
+ if report_test:
+ test_rmse, test_r, test_rho, test_stats = evaluate(model, loaders["test"])
+ summary.update(
+ {
+ "test_rmse": float(test_rmse),
+ "test_r": float(test_r),
+ "test_rho": float(test_rho),
+ "test_stats": {k: float(v) for k, v in test_stats.items()},
+ }
+ )
+ print(f"[test] rmse {test_rmse:.4f} | r {test_r:.3f} | rho {test_rho:.3f}", flush=True)
+
+ summary_json.write_text(json.dumps(summary, indent=2) + "\n")
+ print(f"[done] local weight -> {local_pt}", flush=True)
+ if package_weight is not None:
+ print(f"[done] package weight -> {package_weight}", flush=True)
+ print(f"[done] summary -> {summary_json}", flush=True)
+
+
+def main() -> None:
+ args = parse_args()
+ cfg = config_from_args(args)
+ set_seed(cfg.seed)
+ if args.stage == "val":
+ run_validation(cfg, report_test=args.report_test)
+ else:
+ run_production(
+ cfg,
+ epochs=args.epochs,
+ report_test=args.report_test,
+ copy_to_package=args.copy_to_package,
+ )
+
+
+if __name__ == "__main__":
+ main()
diff --git a/training/03_train_mpnn_val.py b/training/03_train_mpnn_val.py
deleted file mode 100644
index ca7641e..0000000
--- a/training/03_train_mpnn_val.py
+++ /dev/null
@@ -1,128 +0,0 @@
-#!/usr/bin/env python3
-"""
-Train FastHydroMap direct MPNN on train split, early-stop on val split.
-"""
-
-from __future__ import annotations
-
-import argparse
-import json
-from pathlib import Path
-
-import torch
-from torch_geometric.loader import DataLoader
-
-from train_mpnn_common import (
- DEVICE,
- GraphDSDirect,
- TrainConfig,
- edge_dim_from_cfg,
- evaluate,
- graph_cache_paths,
- load_meta_and_splits,
- make_model,
- make_optimizer,
- masked_mean_target,
- set_seed,
- split_indices,
- train_one_epoch,
-)
-
-ROOT = Path(__file__).resolve().parent
-CKPT_DIR = ROOT / "checkpoints_direct_feat"
-MODEL_DIR = ROOT / "models"
-CKPT_DIR.mkdir(exist_ok=True)
-MODEL_DIR.mkdir(exist_ok=True)
-
-
-def parse_args():
- p = argparse.ArgumentParser()
- p.add_argument("--seed", type=int, default=48)
- p.add_argument("--report-test", action="store_true", help="evaluate held-out test split after training")
- return p.parse_args()
-
-
-def main():
- args = parse_args()
- cfg = TrainConfig(seed=args.seed)
- set_seed(cfg.seed)
-
- graph_pt, pid_pt = graph_cache_paths(cfg)
- if not graph_pt.exists() or not pid_pt.exists():
- raise FileNotFoundError(
- f"missing graph cache: {graph_pt.name} / {pid_pt.name}. "
- "build it first with 02_build_mpnn_graphs.py using matching settings."
- )
-
- meta, splits = load_meta_and_splits()
- split_ids = {"train": splits["train"], "val": splits["val"], "test": splits["test"]}
- split_ix, split_ids_kept = split_indices(pid_pt, split_ids)
- mu = masked_mean_target(meta, split_ids_kept["train"])
- edge_dim = edge_dim_from_cfg(cfg)
-
- ds = GraphDSDirect(graph_pt, meta)
- trL = DataLoader(ds[split_ix["train"]], cfg.batch_size, shuffle=True, num_workers=4)
- vaL = DataLoader(ds[split_ix["val"]], cfg.batch_size, shuffle=False, num_workers=4)
- teL = DataLoader(ds[split_ix["test"]], cfg.batch_size, shuffle=False, num_workers=4)
-
- model = make_model(cfg, mu=mu, edge_dim=edge_dim)
- opt = make_optimizer(model, cfg)
-
- tag = "k12_rbf3_r2to14_s4_h24_d2_head20"
- best_pt = CKPT_DIR / f"best_{tag}_seed{cfg.seed}.pt"
- summary_json = MODEL_DIR / f"03_train_mpnn_val_{tag}_seed{cfg.seed}.json"
-
- best_val = float("inf")
- best_epoch = 0
- wait = 0
- for epoch in range(1, cfg.epochs + 1):
- train_one_epoch(model, trL, opt, cfg)
- val_rmse, val_r, val_rho, _ = evaluate(model, vaL)
- print(f"[val] epoch {epoch:02d} rmse {val_rmse:.4f} | r {val_r:.3f} | rho {val_rho:.3f}", flush=True)
- if val_rmse + 1e-4 < best_val:
- best_val = val_rmse
- best_epoch = epoch
- wait = 0
- torch.save(model.state_dict(), best_pt)
- print(f"[val] new best -> {best_pt.name}", flush=True)
- else:
- wait += 1
- if wait >= cfg.patience:
- print("[val] early stop", flush=True)
- break
-
- model.load_state_dict(torch.load(best_pt, map_location=DEVICE))
- val_rmse, val_r, val_rho, val_stats = evaluate(model, vaL)
- summary = {
- "seed": cfg.seed,
- "best_epoch": best_epoch,
- "best_val_rmse": float(best_val),
- "val_rmse_reloaded": float(val_rmse),
- "val_r": float(val_r),
- "val_rho": float(val_rho),
- "mu": float(mu),
- "edge_dim": edge_dim,
- "graph_pt": str(graph_pt),
- "best_ckpt": str(best_pt),
- "val_stats": {k: float(v) for k, v in val_stats.items()},
- }
-
- if args.report_test:
- test_rmse, test_r, test_rho, test_stats = evaluate(model, teL)
- summary.update(
- {
- "test_rmse": float(test_rmse),
- "test_r": float(test_r),
- "test_rho": float(test_rho),
- "test_stats": {k: float(v) for k, v in test_stats.items()},
- }
- )
- print(f"[test] rmse {test_rmse:.4f} | r {test_r:.3f} | rho {test_rho:.3f}", flush=True)
-
- summary_json.write_text(json.dumps(summary, indent=2) + "\n")
- print(f"[done] best val rmse {best_val:.4f} at epoch {best_epoch}", flush=True)
- print(f"[done] summary -> {summary_json}", flush=True)
-
-
-if __name__ == "__main__":
- main()
diff --git a/training/04_train_mpnn_prod.py b/training/04_train_mpnn_prod.py
deleted file mode 100644
index 4cc0a76..0000000
--- a/training/04_train_mpnn_prod.py
+++ /dev/null
@@ -1,122 +0,0 @@
-#!/usr/bin/env python3
-"""
-Train FastHydroMap direct MPNN on train+val for production weights.
-"""
-
-from __future__ import annotations
-
-import argparse
-import json
-import shutil
-from pathlib import Path
-
-import torch
-from torch_geometric.loader import DataLoader
-
-from train_mpnn_common import (
- GraphDSDirect,
- TrainConfig,
- edge_dim_from_cfg,
- evaluate,
- graph_cache_paths,
- load_meta_and_splits,
- make_model,
- make_optimizer,
- masked_mean_target,
- set_seed,
- split_indices,
- train_one_epoch,
-)
-
-ROOT = Path(__file__).resolve().parent
-MODEL_DIR = ROOT / "models"
-PKG_WEIGHT_DIR = ROOT.parent / "src" / "FastHydroMap" / "weights"
-MODEL_DIR.mkdir(exist_ok=True)
-PKG_WEIGHT_DIR.mkdir(exist_ok=True)
-VAL_SUMMARY = MODEL_DIR / "03_train_mpnn_val_k12_rbf3_r2to14_s4_h24_d2_head20_seed48.json"
-
-
-def parse_args():
- p = argparse.ArgumentParser()
- p.add_argument("--seed", type=int, default=48)
- p.add_argument("--epochs", type=int, default=None, help="if omitted, use best_epoch from val summary")
- p.add_argument("--report-test", action="store_true", help="evaluate held-out test split at end")
- return p.parse_args()
-
-
-def main():
- args = parse_args()
- cfg = TrainConfig(seed=args.seed)
- set_seed(cfg.seed)
-
- graph_pt, pid_pt = graph_cache_paths(cfg)
- if not graph_pt.exists() or not pid_pt.exists():
- raise FileNotFoundError(
- f"missing graph cache: {graph_pt.name} / {pid_pt.name}. "
- "build it first with 02_build_mpnn_graphs.py using matching settings."
- )
-
- epochs = args.epochs
- if epochs is None:
- if not VAL_SUMMARY.exists():
- raise FileNotFoundError(f"missing val summary for default epochs: {VAL_SUMMARY}")
- epochs = int(json.loads(VAL_SUMMARY.read_text())["best_epoch"])
-
- meta, splits = load_meta_and_splits()
- split_ids = {
- "trainval": [*splits["train"], *splits["val"]],
- "test": splits["test"],
- }
- split_ix, split_ids_kept = split_indices(pid_pt, split_ids)
- mu = masked_mean_target(meta, split_ids_kept["trainval"])
- edge_dim = edge_dim_from_cfg(cfg)
-
- ds = GraphDSDirect(graph_pt, meta)
- trL = DataLoader(ds[split_ix["trainval"]], cfg.batch_size, shuffle=True, num_workers=4)
- teL = DataLoader(ds[split_ix["test"]], cfg.batch_size, shuffle=False, num_workers=4)
-
- model = make_model(cfg, mu=mu, edge_dim=edge_dim)
- opt = make_optimizer(model, cfg)
-
- tag = "k12_rbf3_r2to14_s4_h24_d2_head20"
- local_pt = MODEL_DIR / f"mpnn_direct_prod_{tag}.pt"
- pkg_pt = PKG_WEIGHT_DIR / f"mpnn_direct_prod_{tag}.pt"
- summary_json = MODEL_DIR / f"04_train_mpnn_prod_{tag}_seed{cfg.seed}.json"
-
- for epoch in range(1, epochs + 1):
- train_one_epoch(model, trL, opt, cfg)
- print(f"[prod] epoch {epoch:02d}/{epochs}", flush=True)
-
- torch.save(model.state_dict(), local_pt)
- shutil.copy2(local_pt, pkg_pt)
-
- summary = {
- "seed": cfg.seed,
- "epochs": epochs,
- "mu": float(mu),
- "edge_dim": edge_dim,
- "local_weight": str(local_pt),
- "package_weight": str(pkg_pt),
- "graph_pt": str(graph_pt),
- }
-
- if args.report_test:
- test_rmse, test_r, test_rho, test_stats = evaluate(model, teL)
- summary.update(
- {
- "test_rmse": float(test_rmse),
- "test_r": float(test_r),
- "test_rho": float(test_rho),
- "test_stats": {k: float(v) for k, v in test_stats.items()},
- }
- )
- print(f"[test] rmse {test_rmse:.4f} | r {test_r:.3f} | rho {test_rho:.3f}", flush=True)
-
- summary_json.write_text(json.dumps(summary, indent=2) + "\n")
- print(f"[done] local weight -> {local_pt}", flush=True)
- print(f"[done] package weight -> {pkg_pt}", flush=True)
- print(f"[done] summary -> {summary_json}", flush=True)
-
-
-if __name__ == "__main__":
- main()
diff --git a/training/README.md b/training/README.md
index 7f9c9ad..18f0707 100644
--- a/training/README.md
+++ b/training/README.md
@@ -1,7 +1,8 @@
# FastHydroMap Training
This directory contains the minimal training pipeline used to build the
-FastHydroMap direct MPNN weights from residue-level Fdewet targets.
+FastHydroMap direct MPNN weights from residue-level `Fdewet`, `PC1`, `PC2`, or
+`PC3` targets.
The large graph tensors and checkpoints are generated artifacts and are not
stored in Git. They can be rebuilt from the scripts, CSV metadata, and source
@@ -15,9 +16,7 @@ PDB structures.
statistics from the source PDB structures.
- `02_build_mpnn_graphs.py`: converts the CSV and PDB structures into cached
PyTorch Geometric graph tensors.
-- `03_train_mpnn_val.py`: trains on the training split and early-stops on the
- validation split.
-- `04_train_mpnn_prod.py`: retrains on train+validation for production weights.
+- `03_train_mpnn.py`: trains validation-stage or production MPNN weights.
- `train_mpnn_common.py`: shared dataset, model, optimizer, and evaluation
helpers.
- `residue_keys.py`: stable residue identifiers for chains and insertion codes.
@@ -38,19 +37,49 @@ Download the source PDB files to `training/data/rcsb_pdbs/`.
- `head_hidden=20`
- trust mask: `avg_n_waters > 7.0` and `3.8 <= Fdewet_pred <= 8.7`
-## Reproduce Training
+## Reproduce Fdewet Training
From the repository root:
```bash
python training/02_build_mpnn_graphs.py --k 12 --n-rbf 3 --rbf-min 2.0 --rbf-max 14.0 --rbf-sigma 4.0
-python training/03_train_mpnn_val.py --seed 48 --report-test
-python training/04_train_mpnn_prod.py --seed 48
+python training/03_train_mpnn.py --stage val --seed 48 --report-test
+python training/03_train_mpnn.py --stage prod --seed 48 --copy-to-package
```
If you want to skip the validation-stage JSON lookup, pass the production epoch
count explicitly:
```bash
-python training/04_train_mpnn_prod.py --seed 48 --epochs 22
+python training/03_train_mpnn.py --stage prod --seed 48 --epochs 22 --copy-to-package
+```
+
+## PC Target Training
+
+`data/all_residue_results.csv` includes `PC1`, `PC2`, and `PC3`. Use
+`--target` to train the same architecture against a PC target:
+
+```bash
+python training/03_train_mpnn.py --stage val --target PC1 --report-test
+python training/03_train_mpnn.py --stage prod --target PC1 --copy-to-package
+```
+
+For PC targets, `--mask-source auto` uses the CSV `trusted` column. The released
+PC weights were trained with target winsorization at the 2nd and 98th
+percentiles of the trusted fitting split:
+
+```bash
+python training/03_train_mpnn.py --stage val --target PC1 --winsor-lower 0.02 --winsor-upper 0.98 --report-test
+python training/03_train_mpnn.py --stage prod --target PC1 --winsor-lower 0.02 --winsor-upper 0.98 --copy-to-package
+```
+
+Repeat with `--target PC2` or `--target PC3` for the other PC regressors. The
+production stage writes a target-specific local weight under `training/models/`.
+When `--copy-to-package` is supplied, the package weight is updated at:
+
+```bash
+src/FastHydroMap/weights/mpnn_latest.pt # Fdewet_pred
+src/FastHydroMap/weights/mpnn_pc1_latest.pt # PC1
+src/FastHydroMap/weights/mpnn_pc2_latest.pt # PC2
+src/FastHydroMap/weights/mpnn_pc3_latest.pt # PC3
```
diff --git a/training/train_mpnn_common.py b/training/train_mpnn_common.py
index 396b5bb..91453f6 100644
--- a/training/train_mpnn_common.py
+++ b/training/train_mpnn_common.py
@@ -25,6 +25,8 @@
META_CSV = ROOT / "all_residue_results.csv"
SPLIT_YML = ROOT / "splits.yaml"
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
+SUPPORTED_TARGETS = ("Fdewet_pred", "PC1", "PC2", "PC3")
+SUPPORTED_MASK_SOURCES = ("auto", "fdewet", "trusted", "all")
@dataclass(frozen=True)
@@ -48,6 +50,20 @@ class TrainConfig:
clip: float = 2.0
factor: float = 1.0
seed: int = 48
+ target: str = "Fdewet_pred"
+ mask_source: str = "auto"
+ winsor_lower: float | None = None
+ winsor_upper: float | None = None
+
+ def __post_init__(self):
+ if self.target not in SUPPORTED_TARGETS:
+ raise ValueError(f"target must be one of {SUPPORTED_TARGETS}; got {self.target!r}")
+ if self.mask_source not in SUPPORTED_MASK_SOURCES:
+ raise ValueError(f"mask_source must be one of {SUPPORTED_MASK_SOURCES}; got {self.mask_source!r}")
+ if (self.winsor_lower is None) != (self.winsor_upper is None):
+ raise ValueError("winsor_lower and winsor_upper must be set together")
+ if self.winsor_lower is not None and not (0.0 <= self.winsor_lower < self.winsor_upper <= 1.0):
+ raise ValueError("winsor bounds must satisfy 0 <= lower < upper <= 1")
def _flt_tag(x: float) -> str:
@@ -74,23 +90,88 @@ def set_seed(seed: int) -> None:
np.random.seed(seed)
-def masked_mean_target(meta: pd.DataFrame, pdb_ids: list[str]) -> float:
+def _target_mask(rows: pd.DataFrame, target: str, mask_source: str) -> np.ndarray:
+ if target not in rows.columns:
+ raise ValueError(f"target column {target!r} not found in metadata")
+
+ if mask_source == "auto":
+ mask_source = "fdewet" if target == "Fdewet_pred" else "trusted"
+
+ if mask_source == "fdewet":
+ mask = (
+ (rows.avg_n_waters.values > 7.0)
+ & (rows.Fdewet_pred.values >= 3.8)
+ & (rows.Fdewet_pred.values <= 8.7)
+ )
+ elif mask_source == "trusted":
+ if "trusted" not in rows.columns:
+ raise ValueError("trusted mask requested, but metadata has no 'trusted' column")
+ mask = rows.trusted.astype(bool).values
+ elif mask_source == "all":
+ mask = np.ones(len(rows), dtype=np.bool_)
+ else:
+ raise ValueError(f"unknown mask source: {mask_source}")
+
+ return mask & pd.notna(rows[target]).values
+
+
+def compute_winsor_bounds(
+ meta: pd.DataFrame,
+ pdb_ids: list[str],
+ target: str,
+ mask_source: str,
+ lower: float | None,
+ upper: float | None,
+) -> tuple[float, float] | None:
+ if lower is None or upper is None:
+ return None
subset = meta[meta["pdb_id"].isin(pdb_ids)]
- trusted = (
- (subset.avg_n_waters.values > 7.0)
- & (subset.Fdewet_pred.values >= 3.8)
- & (subset.Fdewet_pred.values <= 8.7)
- )
- return float(subset.loc[trusted, "Fdewet_pred"].mean())
+ mask = _target_mask(subset, target, mask_source)
+ vals = subset.loc[mask, target].astype(float)
+ return float(vals.quantile(lower)), float(vals.quantile(upper))
+
+
+def apply_target_clip(values: np.ndarray, bounds: tuple[float, float] | None) -> np.ndarray:
+ if bounds is None:
+ return values
+ lo, hi = bounds
+ return np.clip(values, lo, hi)
+
+
+def masked_mean_target(
+ meta: pd.DataFrame,
+ pdb_ids: list[str],
+ target: str,
+ mask_source: str,
+ target_clip_bounds: tuple[float, float] | None = None,
+) -> float:
+ subset = meta[meta["pdb_id"].isin(pdb_ids)]
+ mask = _target_mask(subset, target, mask_source)
+ vals = subset.loc[mask, target].values.astype(np.float32)
+ vals = apply_target_clip(vals, target_clip_bounds)
+ return float(vals.mean())
class GraphDSDirect(InMemoryDataset):
- def __init__(self, graphs_pt: Path, meta_df: pd.DataFrame):
+ def __init__(
+ self,
+ graphs_pt: Path,
+ meta_df: pd.DataFrame,
+ target: str = "Fdewet_pred",
+ mask_source: str = "auto",
+ target_clip_bounds: tuple[float, float] | None = None,
+ ):
super().__init__("")
self.data, self.slices = torch.load(graphs_pt)
- self._augment_with_meta(meta_df)
-
- def _augment_with_meta(self, meta: pd.DataFrame):
+ self._augment_with_meta(meta_df, target, mask_source, target_clip_bounds)
+
+ def _augment_with_meta(
+ self,
+ meta: pd.DataFrame,
+ target: str,
+ mask_source: str,
+ target_clip_bounds: tuple[float, float] | None,
+ ):
meta = ensure_residue_key_columns(meta)
mi_uid = meta.set_index(["pdb_id", "res_uid"])
mi_resid = meta.set_index(["pdb_id", "resid"])
@@ -105,15 +186,17 @@ def _augment_with_meta(self, meta: pd.DataFrame):
if rows.isnull().any().any():
raise ValueError(f"missing metadata for graph {g.pdb_id}")
- target = rows.Fdewet_pred.values.astype(np.float32)
- trusted = (
- (rows.avg_n_waters.values > 7.0)
- & (rows.Fdewet_pred.values >= 3.8)
- & (rows.Fdewet_pred.values <= 8.7)
- ).astype(np.bool_)
+ raw_target_values = rows[target].values.astype(np.float32)
+ target_values = apply_target_clip(raw_target_values, target_clip_bounds)
+ trusted = _target_mask(rows, target, mask_source).astype(np.bool_)
- g.target = torch.from_numpy(target)
+ g.target = torch.from_numpy(target_values)
+ g.raw_target = torch.from_numpy(raw_target_values)
g.mask = torch.from_numpy(trusted)
+ g.target_name = target
+ g.mask_source = mask_source
+ if target_clip_bounds is not None:
+ g.target_clip_bounds = target_clip_bounds
new_graphs.append(g)
self.data, self.slices = InMemoryDataset.collate(new_graphs)