Skip to content

added ReBind QM9 baseline; organometallic models running - #1

Merged
jwtoney merged 12 commits into
mainfrom
rebind_baseline
Oct 2, 2026
Merged

jwtoney merged 12 commits into
mainfrom
rebind_baseline

Conversation

@jwtoney

@jwtoney jwtoney commented May 25, 2026

Copy link
Copy Markdown
Collaborator

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.

@jwtoney
jwtoney requested a review from sid-betalol May 25, 2026 00:34

@sid-betalol sid-betalol left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Thanks for the great work, @jwtoney! This will need some changes, though, for some minor issues.

Comment thread pyproject.toml
"ignore::UserWarning",
]

[tool.uv.sources]

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

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 torch in base dependencies so uv sync --dev works on macOS/Linux CPU;
  • add a CUDA extra or separate docs command for GPU training environments;

@jwtoney jwtoney Jun 8, 2026 •

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Sounds good, made these changes

Comment thread README.md Outdated
```

```bash
sbatch scripts/train.slurm configs/qm9.yaml

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

scripts/train.slurm was not pushed. Do you mean scripts/train.sh?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

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_embedding

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

@jwtoney jwtoney Jun 8, 2026 •

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

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

Comment thread tests/conftest.py Outdated

import pytest

QM9_PATH = Path("/home/gridsan/jtoney/ElemNet/benchmarking/datasets/QM9-full.csv")

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

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.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

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.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

This module touches external/ReBIND as soon as it is imported. That means a normal import path like:

from step_up.data.csv_dataset import CSVMoleculeDataset

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

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

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

Comment thread configs/bostmc.yaml Outdated
# - `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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

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

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

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Added filter_column/filter_value to the dataset and config

Comment thread src/step_up/train.py Outdated

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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

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 test

Do 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?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Yes, after training the run's best-val checkpoint is evaluated on the test split and saved to test_metrics.json with its epoch

Comment thread src/step_up/eval/metrics.py Outdated
zs = atomic_numbers.to(torch.long)
per_z: dict[int, float] = {}
for z in torch.unique(zs):
if int(z.item()) == 0:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

The function says atomic_numbers should contain true atomic numbers and uses 0 as padding:

if int(z.item()) == 0:
    continue

But 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, document Z - 1, and use a separate padding mask instead of skipping 0; or
  • require true atomic numbers and explicitly convert node_type + 1 before calling.

Add a test that includes hydrogen to prevent this regression.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

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

Comment thread src/step_up/data/csv_dataset.py Outdated
self._df = _read_csv_subset(self.path, cols, nrows=subset_size)

if validate:
self._valid_indices = self._validate_rows()

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

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.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Validation raises if it keeps zero rows, max_drop_fraction caps the drop rate, and train.py rejects empty train or val splits

Comment thread src/step_up/data/csv_dataset.py Outdated

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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

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.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Added these; cache_dataset keeps the featurized graphs in memory. validate_dataset false skips the upfront pass for datasets known to be clean

luispintoc
luispintoc previously approved these changes Sep 21, 2026

@luispintoc luispintoc left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

LGTM

@luispintoc

luispintoc commented Sep 21, 2026 •

Copy link
Copy Markdown
Collaborator

GitHub still blocks the merge for two reasons @jwtoney

Unsigned commits. 413b4a4 is verified, but these two are not:

  • 46cc29d — addressed requested changes
  • 675a643 — Remove personal Slurm wrapper jobscript.sh

You can merge after solving it

@luispintoc

Copy link
Copy Markdown
Collaborator

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

@luispintoc
luispintoc self-requested a review September 25, 2026 08:49
jwtoney and others added 6 commits September 27, 2026 10:54
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>
jwtoney and others added 2 commits September 27, 2026 11:51
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>
@jwtoney

jwtoney commented Sep 28, 2026

Copy link
Copy Markdown
Collaborator Author

@luispintoc Thanks! I added a signature so I think you need to reapprove it and then we can squash and merge?

jwtoney and others added 3 commits September 28, 2026 12:54
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 luispintoc left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Looks good, feel free to merge

@jwtoney
jwtoney merged commit 81df1eb into main Oct 2, 2026
2 checks passed
@jwtoney
jwtoney deleted the rebind_baseline branch October 2, 2026 14:32
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants