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 @@ [![DOI](https://zenodo.org/badge/1023802589.svg)](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)