added ReBind QM9 baseline; organometallic models running - #1
Conversation
sid-betalol
left a comment
There was a problem hiding this comment.
Thanks for the great work, @jwtoney! This will need some changes, though, for some minor issues.
| "ignore::UserWarning", | ||
| ] | ||
|
|
||
| [tool.uv.sources] |
There was a problem hiding this comment.
torch is globally pinned to the pytorch-cu124. On macOS arm64, uv run pytest -q fails before pytest starts:
error: Distribution `torch==2.6.0+cu124 ...` can't be installed because it doesn't have a source distribution or wheel for the current platform
I suggest making the default dependency set CPU/platform-compatible, and moving CUDA Torch into an optional group or environment-specific install path. For example:
- keep normal
torchin base dependencies souv sync --devworks on macOS/Linux CPU; - add a CUDA extra or separate docs command for GPU training environments;
There was a problem hiding this comment.
Sounds good, made these changes
| ``` | ||
|
|
||
| ```bash | ||
| sbatch scripts/train.slurm configs/qm9.yaml |
There was a problem hiding this comment.
scripts/train.slurm was not pushed. Do you mean scripts/train.sh?
There was a problem hiding this comment.
Yes, thanks - I renamed the script but didn't edit the README. I've done that now
| padding_mask=node_mask, | ||
| compute_loss=True, | ||
| ) | ||
| loss_cache, conformer_cache = cache_out["loss"], cache_out["conformer_hat"] |
There was a problem hiding this comment.
Maybe I am misunderstanding something for this comment, so please clarify.
The conformer head output is stored in conformer_cache:
loss_cache, conformer_cache = cache_out["loss"], cache_out["conformer_hat"]but later the code sets:
inputs["pred_conformation"] = node_embeddingand passes that to the residual head as conformer_base:
conformer_base=inputs["pred_conformation"],That looks wrong: node_embedding has shape (batch, atoms, d_model), while a conformer base should be coordinates with shape (batch, atoms, 3)? The natural value here appears to be conformer_cache, not node_embedding. If this path is exercised, it can either crash on a shape mismatch or train against the wrong tensor.
If what I am saying is correct, then the fix would be to change inputs["pred_conformation"] to use conformer_cache, then add/strengthen a forward-pass test that exercises the patched forward with a real collated batch and asserts finite loss/output shapes.
There was a problem hiding this comment.
I believe this is implemented correctly, I took it directly from their github and it seems to be the faithful implementation of section 4.3 of the paper. I did notice Laplacian positional encoding was missing at one place so I added that back
|
|
||
| import pytest | ||
|
|
||
| QM9_PATH = Path("/home/gridsan/jtoney/ElemNet/benchmarking/datasets/QM9-full.csv") |
There was a problem hiding this comment.
The test fixtures point at absolute cluster paths:
QM9_PATH = Path("/home/gridsan/jtoney/ElemNet/benchmarking/datasets/QM9-full.csv")
TMQMG_PATH = Path("/home/gridsan/jtoney/ElemNet/benchmarking/datasets/tmQMg-full.csv")
BOSTMC_PATH = Path("/home/gridsan/jtoney/BOSTMC/datasets/BOSTMC-low-spin.csv")On any contributor machine or GitHub runner without those files, the dataset and forward tests will skip. That means CI will not actually validate the new featurization, MOL2 parser, LJ patch, or training path.
I suggest committing tiny synthetic fixture CSVs under tests/fixtures/ that cover:
- one QM9-style XYZ row;
- one MOL2+XYZ organometallic row with a transition metal;
- at least one ring/aromatic case if ring flags matter.
We should use private full datasets only for optional/integration tests, not core CI tests.
There was a problem hiding this comment.
Added tests/fixtures/qm9_mini.csv, tmqmg_mini.csv, and bostmc_mini.csv for CI tests
| @@ -0,0 +1,123 @@ | |||
| """Convert raw XYZ / MOL2 blocks into ReBind-format graph dicts. | |||
|
|
|||
There was a problem hiding this comment.
This module touches external/ReBIND as soon as it is imported. That means a normal import path like:
from step_up.data.csv_dataset import CSVMoleculeDatasetcan try to load vendored ReBIND code before any test body runs. This is fragile because pytest imports test modules during collection. If the repo was cloned without --recursive, external/ReBIND may exist as an empty submodule directory, but the actual files such as external/ReBIND/models and external/ReBIND/data/utils.py are missing. In that case, pytest can fail during collection with an import error instead of reaching a test fixture that could skip with a useful message.
There is also an inconsistency between the two modules: featurize.py checks for a concrete file (data/utils.py) and raises a helpful FileNotFoundError, while rebind.py only checks whether the submodule root path exists. An empty submodule directory passes that check, then fails later with a less helpful ModuleNotFoundError.
I suggest we make rebind.py check for concrete submodule contents, e.g. external/ReBIND/models or external/ReBIND/data/utils.py, and raise the same actionable error: Run: git submodule update --init --recursive. Also ensure CI and docs use recursive submodule checkout. Longer term, defer vendored ReBIND imports until build_rebind() or the specific featurization function is called, so importing dataset utilities does not require the model submodule to be initialized.
There was a problem hiding this comment.
Okay refactored this, featurize.py checks for data/utils.py and rebind.py checks for models/rebind/modeling_rebind.py, and both raise the same helpful error messages
| # - `charge` and `spinmult` columns are currently IGNORED. | ||
| # - Recommendation for first publishable run: pre-filter to singlets-only (spinmult == 1) for the cleanest comparison to tmQMg. The CSVMoleculeDataset does not yet expose a filter knob, add one when this config is first run. | ||
| # - Follow-up: condition the model on (charge, spinmult) as a global feature. | ||
| dataset_path: /home/gridsan/jtoney/BOSTMC/datasets/BOSTMC-low-spin.csv |
There was a problem hiding this comment.
The config comment says charge and spinmult are ignored and recommends pre-filtering to singlets before a publishable run, but the config still points at BOSTMC-low-spin.csv with no filter support:
dataset_path: /home/gridsan/jtoney/BOSTMC/datasets/BOSTMC-low-spin.csvFor a benchmark, mixing singlets and doublets without conditioning on spin/charge can make comparisons ambiguous and hard to reproduce.
It'd be good to either add filtering to CSVMoleculeDataset / config and make this config singlet-only, or rename/mark the BOSTMC full config as experimental until spin/charge conditioning exists.
There was a problem hiding this comment.
Added filter_column/filter_value to the dataset and config
|
|
||
| with open(out_dir / "history.json", "w") as f: | ||
| json.dump(history, f, indent=2) | ||
| del test_set # held out; first round does not evaluate on test |
There was a problem hiding this comment.
The code creates a test split but then discards it:
train_set, val_set, test_set = random_split(...)
...
del test_set # held out; first round does not evaluate on testDo you think we should save test metrics also?
After training, we could load/use the best checkpoint and evaluate on test_set, then write test_metrics.json. It might be worth setting up a wandb for this project @luispintoc?
There was a problem hiding this comment.
Yes, after training the run's best-val checkpoint is evaluated on the test split and saved to test_metrics.json with its epoch
| zs = atomic_numbers.to(torch.long) | ||
| per_z: dict[int, float] = {} | ||
| for z in torch.unique(zs): | ||
| if int(z.item()) == 0: |
There was a problem hiding this comment.
The function says atomic_numbers should contain true atomic numbers and uses 0 as padding:
if int(z.item()) == 0:
continueBut the graph data stores node_type as Z - 1, where hydrogen is 0. The test also passes [5, 6, 7] while commenting that these are C/N/O, which are actually Z - 1 indices, not atomic numbers. If callers pass node_type directly, hydrogen is silently dropped, and all element keys are off by one.
We should choose one representation and enforce it. Either:
- rename the argument to
atomic_number_indices, documentZ - 1, and use a separate padding mask instead of skipping0; or - require true atomic numbers and explicitly convert
node_type + 1before calling.
Add a test that includes hydrogen to prevent this regression.
There was a problem hiding this comment.
I used true atomic numbers plus node_mask for padding, so hydrogens are no longer dropped. An error is raised if any atom has Z < 1
| self._df = _read_csv_subset(self.path, cols, nrows=subset_size) | ||
|
|
||
| if validate: | ||
| self._valid_indices = self._validate_rows() |
There was a problem hiding this comment.
CSVMoleculeDataset._validate_rows catches every exception, records the class name, and keeps going. If every row fails, len(dataset) becomes 0. train.py then uses max(len(train_loader), 1) and can continue into a run with empty loaders, meaningless metrics, and no useful checkpoint.
We should fail fast when validation keeps zero rows, and consider adding a configurable maximum drop rate for full runs.
There was a problem hiding this comment.
Validation raises if it keeps zero rows, max_drop_fraction caps the drop rate, and train.py rejects empty train or val splits
|
|
||
| The CSV is loaded into a pandas DataFrame at construction time (the | ||
| ``subset_size`` knob keeps memory bounded for smoke runs). When | ||
| ``validate=True`` (the default), the dataset then walks every row, calls |
There was a problem hiding this comment.
Construction with validate=True featurizes every row once to find valid indices. __getitem__ featurizes the same row again during training.
We can cache validated graph dicts for small/subset runs, or store validation results plus a cheaper row-level validity marker? At minimum, we should expose a config knob to disable validation for already-clean production datasets.
There was a problem hiding this comment.
Added these; cache_dataset keeps the featurized graphs in memory. validate_dataset false skips the upfront pass for datasets known to be clean
a3bda28 to
413b4a4
Compare
|
I pushed the script that contains all LJ parameters. I just realized that you've added rebind as a submodule, maybe you might want to fork it under lemat and we can modify it there, otherwise you won't be able to overwrite the LJ parameters |
Second benchmark model plus the corrections that came out of running the first one end to end. GTMGC (ICLR 2024) joins ReBind behind `model:` in the config. It shares ReBind's code base, so both wrappers now sit on `models/common.py` (the out-of-place Laplacian addition and the charge/spin conditioning) and the vendored tree loads under a private module name to keep its `models` package from colliding with ReBind's. `from_pretrained` breaks on transformers 5.8.1, so released weights load through an explicit config + state-dict path. `scripts/tokenize_molebert.py` precomputes the Mole-BERT atom ids that its `atom_tokenized_ids` embedding needs. `eval/conformer_eval.py` implements ReBind's published protocol: D-MAE and D-RMSE pooled over every pairwise distance, C-RMSD via RDKit GetBestRMS with hydrogens removed. Both vendored models Kabsch-align inside the prediction head over the *padded* batch tensor, so padding zeros drag the centroid and rotation off; trusting that alignment inflated RMSD roughly threefold (BOSTMC read 8.17 where it is 2.53). `kabsch_rmsd` now aligns per molecule over real atoms only. ReBind hardcodes Lennard-Jones parameters up to Kr and KeyErrors above it, which we had been papering over with a flat sigma/epsilon for every heavier element. `models/lj_params.py` carries the full UFF table through Lr, from the values Luis Pinto compiled in rebind_utils.py; it is bit-identical to ReBind's over Z=1..36, so QM9 is untouched. This is not cosmetic: 65% of tmQMg and 46% of BOSTMC low-spin structures contain an element past Kr, usually the metal centre itself. Also here: stable hash-based splits and published split/ID-list loading, charge and spin conditioning for the organometallic sets with the ablation script that measures whether a model uses it, the QM9 preparation script for ReBind's published split, per-dataset configs, and the first round of results in the README. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
f7ab2f9 to
c214a00
Compare
Its unique content was the extended Lennard-Jones table, which is now in src/step_up/models/lj_params.py and patched into the vendored ReBind at runtime, so no fork of ReBind is needed. The rest of the file was a copy of external/ReBIND/models/modules/utils.py with the feature vocabulary changed (numH, bond direction, bond stereo and conjugation dropped, DATIVE added), which would not have been safe to use as-is: ReBind's embedding layers zip those vocab dimensions against the 9 atom / 4 bond feature columns our featurizer emits, so the shorter lists would silently drop features. Keeping the paper's vocabulary. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
The MOL2 rows build no RDKit molecule, so `Chem.RemoveHs` never ran on them and the organometallic RMSDs were being reported over every atom while the QM9 C-RMSD excluded hydrogen. Hydrogen is a quarter to a half of the atoms in these complexes and the hardest to place, so that is not a comparison. The aligned-coordinate path now drops hydrogens by atom type, which needs no RDKit. D-MAE and D-RMSE still pool over every atom, as ReBind's evaluate.py does. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
|
@luispintoc Thanks! I added a signature so I think you need to reapprove it and then we can squash and merge? |
Pairs every benchmark run with a variant the model never sees hydrogen in, so "bad at hydrogen" can be told apart from "bad at geometry". This is a different model, not the eval-time --keep-hs flag: hydrogens are gone before featurization, so they are neither embedded, attended over, nor predicted. Both featurization paths support it. The RDKit path (QM9) uses Chem.RemoveHs. The MOL2 path (tmQMg, BOSTMC) has no RDKit molecule to lean on, so the parser filters and reindexes atoms and bonds itself. Either way each heavy atom keeps its hydrogen count in the numH feature, so the graph still knows how many hydrogens there were; only their positions are gone. `rdkit_mol()` returns the stripped molecule too, or C-RMSD would write predicted coordinates onto the wrong atoms. QM9 without hydrogens contains one-atom molecules -- methane, ammonia -- which is where a degree-zero graph could divide by zero in the Laplacian positional encoding. It does not; there is a test pinning it. Three configs staged: qm9_rebind_noh, tmqmg_complete_noh, bostmc_noh. They are labelled as experiments rather than reproductions, since ReBind's published setup keeps hydrogen explicit. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
One row per test molecule: id, charge, spin_multiplicity, n_atoms, elements, xyz_true, xyz_pred, and that molecule's own d_mae / d_rmse / rmsd. Both geometries are XYZ blocks with atoms in the same order, so they overlay directly, and scripts/eval_and_dump.sh produces the pooled metrics and the CSV from one checkpoint in one job so the two always describe the same state. xyz_pred is Kabsch-aligned onto xyz_true per molecule. The models align inside their prediction head over the padded batch tensor, so raw output sits in a frame corrupted by padding zeros. Only a rotation and translation are applied; a test asserts the prediction's internal distance matrix survives untouched. charge and spin_multiplicity are left blank when the source CSV has no such column, rather than filled with 0 and 1. tmQMg has no spin column at all, and the conditioning collator already defaults that input to zero unpaired electrons -- a default that happens to be consistent with the data (no complex has an odd electron count, so none is forced open-shell) but is not read from it. The tmQMg configs now say that instead of asserting every complex is a closed-shell singlet, and the dump does not repeat the default as if it were data. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
scripts/compare_runs.py reads each run's metrics JSON and its structure dump, normalizes D-MAE by that test set's own mean pairwise distance -- without which a 64-atom complex and a 9-atom skeleton share a column that means nothing -- and recomputes D-MAE from the dumped geometry as a cross-check. All six runs agree with their metrics exactly. Two results worth the README space. Dropping hydrogen does not produce a better heavy-atom structure. It improves D-MAE a lot (QM9 7.5% to 1.7% relative) but that is mostly the hardest atoms leaving the average: on the identical heavy atoms under the identical metric the model trained *with* hydrogen wins, tmQMg RMSD 1.810 against 1.823 and BOSTMC 2.024 against 2.081. ReBind's LJ rewiring is dead code in the released implementation, so every number here was produced without the mechanism the paper is named for. `Collator.__call__` samples `keys` from the input graph dict while `_transform` attaches `num_near_edges` afterwards, so the padding guard never fires and the tensor stays all zeros; with k=0 the top-k retains nothing and both rewiring channels reach the decoder as zero matrices. Upstream's own mol_to_graph_dict has the same keys, so this is not our featurization. Measured on a tmQMg batch with 4d/5d metals: 13,380 candidate pairs, 0 retained, and perturbing the LJ table -- the old flat values, epsilon x1000, sigma x2 -- moves the predicted coordinates by exactly zero. That explains why the LJ-corrected tmQMg rerun scored 0.8918838762 against the pre-fix 0.8918838298, identical to 7 significant figures: the two trainings differed in nothing that reaches the loss. It also means our QM9 reproduction beats the published C-RMSD while running the paper's own ablation. Left documented rather than patched, since fixing the propagation produces a model unlike the released one. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
luispintoc
left a comment
There was a problem hiding this comment.
Looks good, feel free to merge
Added baselines to test ReBind (https://arxiv.org/abs/2410.14696) on QM9, tmQMg, and BOS-TMC datasets. QM9 has completed, tmQMg and BOS-TMC are still running. Using hyperparameters and epochs from original ReBind implementation.