Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
96 commits
Select commit Hold shift + click to select a range
08980dc
chore: clean up .gitignore by removing unnecessary entries
sundusaijaz Feb 21, 2026
f8879c0
feat: add new AtomsDataModuleV2 and StatsAtomrefProvider for enhanced…
sundusaijaz Feb 22, 2026
b043e7e
style: improve code formatting and readability in multiple files
sundusaijaz Feb 22, 2026
6e67e01
refactor: simplify AtomsDataModuleV2 by removing unused parameters an…
sundusaijaz Feb 22, 2026
f7dd667
docs: update AtomsDataModuleV2 docstring for clarity by removing redu…
sundusaijaz Feb 22, 2026
8b5a014
test: add pytests for AtomsDataset and AtomsDataModuleV2 functionality
sundusaijaz Feb 22, 2026
215587e
refactor: update QM9 and StatsAtomrefProvider docstrings for clarity …
sundusaijaz Feb 22, 2026
db21004
feat: refactor AtomsDataModuleV2
sundusaijaz Mar 2, 2026
01fb1dd
refactor: update StatsAtomrefProvider to use BaseAtomsData and simpli…
sundusaijaz Mar 2, 2026
6d897a9
refactor: update calculate_stats and estimate_atomrefs to use BaseAto…
sundusaijaz Mar 2, 2026
7af99a8
refactor: simplify initialization in StatsAtomrefProvider
sundusaijaz Mar 2, 2026
e592743
refactor: update Transform class
sundusaijaz Mar 2, 2026
ba8d46f
refactor: QM9 class by removing unused parameters and simplifying doc…
sundusaijaz Mar 2, 2026
ee24a7e
refactor: enhance AtomsDataModuleV2 and QM9 class by simplifying init…
sundusaijaz Mar 3, 2026
8a90055
fix: black format
sundusaijaz Mar 3, 2026
ee32d85
refactor: update custom and qm9 config files
sundusaijaz Mar 4, 2026
ecaca86
refactor: improve model testing and checkpoint handling in cli
sundusaijaz Mar 4, 2026
1339fb6
refactor: merged ASEAtomsData class and BaseAtomsData
sundusaijaz Mar 4, 2026
cc3cb3c
refactor: update data handling
sundusaijaz Mar 7, 2026
b48ca02
refactor: simplify ASEAtomsData by removing unused methods and proper…
sundusaijaz Mar 7, 2026
bfc2c97
refactor: update dataset method signatures to use ASEAtomsData
sundusaijaz Mar 7, 2026
860fbee
refactor: update references from BaseAtomsData to ASEAtomsData in dat…
sundusaijaz Mar 7, 2026
c83501a
refactor: update checkpoint loading in training process and adjust da…
sundusaijaz Mar 7, 2026
3fa44d3
refactor: clean up code formatting and remove legacy QM9 dataset file
sundusaijaz Mar 8, 2026
7406cef
refactor: update rMD17 dataset class
sundusaijaz Mar 8, 2026
091a0e5
refactor: remove irrelevant refactor pytest
sundusaijaz Mar 8, 2026
2d92e05
refactor: update md17, md22, qm7x, rmd17
sundusaijaz Mar 8, 2026
61740d9
refactor: update dataset classes mp, ani1, iso17
sundusaijaz Mar 8, 2026
c242648
refactor: fix format error in MaterialsProject
sundusaijaz Mar 8, 2026
fb2bd3d
refactor: remove format parameter in QM7X dataset loading
sundusaijaz Mar 8, 2026
1fb9d97
refactor: remove legacy atoms_legacy.py file and streamline dataset l…
sundusaijaz Mar 12, 2026
bc8cb02
refactor: change ASEAtomsData class with additional transform options…
sundusaijaz Mar 12, 2026
4efe59b
refactor: remove format parameter from all dataset classes
sundusaijaz Mar 12, 2026
de23e5d
refactor: simplify format handling in AtomsDataModule (old)
sundusaijaz Mar 12, 2026
dadbf1b
refactor: removed dict in ASEAtomsData and simplify download method i…
sundusaijaz Mar 12, 2026
f788a91
refactor: add deprecation warnings for legacy datamodule methods in a…
sundusaijaz Mar 12, 2026
70cb3e1
refactor: add docstring and deprecation warnings for legacy argument…
sundusaijaz Mar 12, 2026
3bbc31d
refactor: replace property_unit_dict with _native_property_units meth…
sundusaijaz Mar 12, 2026
ed73284
refactor: enhance QM9 dataset with train/val/test transform options
sundusaijaz Mar 12, 2026
e627a69
refactor: add train/val/test transform options and docstring across m…
sundusaijaz Mar 12, 2026
607eab1
refactor: update docstrings in atomistic transforms
sundusaijaz Mar 12, 2026
2a6ce76
refactor: add docstrings in ASEAtomsData
sundusaijaz Mar 12, 2026
f0dd4c5
refactor: add docstrings for calculate_stats() and estimate_atomrefs()
sundusaijaz Mar 12, 2026
4ca3648
refactor: simplify transform initialization in ASEAtomsData and updat…
sundusaijaz Mar 12, 2026
93d9ce8
refactor: restructure configs for all datasets
sundusaijaz Mar 12, 2026
99ba778
refactor: update omdb to support datamodulev2
sundusaijaz Mar 12, 2026
5d44d74
refactor: streamline transform assignment and initialization in Atoms…
sundusaijaz Mar 14, 2026
43e605d
refactor: update pytest test_stats to accept data and batch parameter…
sundusaijaz Mar 14, 2026
db0a96e
refactor: improve ANI1 dataset loading and validation
sundusaijaz Mar 15, 2026
074bee9
refactor: enhance QM9 dataset loading
sundusaijaz Mar 15, 2026
d4ca4b0
refactor: simplify download method in ISO17 dataset
sundusaijaz Mar 15, 2026
985032e
refactor: enhance MaterialsProject and update docstring
sundusaijaz Mar 15, 2026
6bf81d3
refactor: optimize GDMLDataset download method in md17
sundusaijaz Mar 15, 2026
6f30249
refactor: enhance QM7X dataset docstring and improve download methods
sundusaijaz Mar 15, 2026
b5b8969
refactor: improve rMD17 dataset loading and metadata handling
sundusaijaz Mar 15, 2026
799d5a8
refactor: enhance ASEAtomsData and QM9 dataset handling
sundusaijaz Mar 17, 2026
6ee9699
refactor: streamline metadata _check_db() and dataset creation in ASE…
sundusaijaz Mar 19, 2026
86df4d8
refactor: enhance ASEAtomsData split transform and db creation
sundusaijaz Mar 19, 2026
1690b19
refactor: streamline ANI1 and QM9 dataset handling and download methods
sundusaijaz Mar 19, 2026
3962eb8
refactor: update ISO17 dataset download method and improve property u…
sundusaijaz Mar 19, 2026
66ffe81
refactor: improve MaterialsProject API key validation and simplify do…
sundusaijaz Mar 19, 2026
f961526
refactor: streamline GDMLDataset methods and enhance metadata handlin…
sundusaijaz Mar 23, 2026
716b34e
refactor: add type hint to _check_db() method in ASEAtomsData
sundusaijaz Mar 23, 2026
2b51b1c
refactor: simplify download method in omdb
sundusaijaz Mar 23, 2026
ca88623
refactor: simplify QM7X download method and enhance metadata handling
sundusaijaz Mar 23, 2026
6211431
refactor: remove unused imports and streamline rMD17 dataset methods
sundusaijaz Mar 23, 2026
cf4c406
refactor: remove unused split_id from rMD17 dataset configuration
sundusaijaz Mar 27, 2026
252627e
refactor: adjust train and test split calculations in rMD17 dataset
sundusaijaz Mar 27, 2026
1969fbf
refactor: remove AtomsDataModuleV2 and consolidate functionality into…
sundusaijaz Mar 31, 2026
f966c1b
refactor: update AtomsDataModule deprecation warning and configuratio…
sundusaijaz Mar 31, 2026
a47065d
refactor: remove import of datamodule_v2
sundusaijaz Mar 31, 2026
5fa894b
refactor: improve pytest data handling
sundusaijaz Mar 31, 2026
b244d9f
refactor: move transform from AtomsDataModule to ASEAtomsData
sundusaijaz Apr 10, 2026
86f812c
refactor: simplify transform initialization in AtomsDataModule and AS…
sundusaijaz Apr 10, 2026
cf94e59
refactor: rename atomrefs to _atomrefs rMD17 dataset
sundusaijaz Apr 10, 2026
e2e4ecf
reformatted
sundusaijaz May 12, 2026
73dd800
fix: address dataset refactor review findings
stefaanhessmann Jul 22, 2026
25b839a
Merge remote sa/dataset_refactor into review-fixed state
stefaanhessmann Jul 22, 2026
211061d
style: apply black formatting
stefaanhessmann Jul 22, 2026
9ca9186
style: reformat with pinned black 24.4.2
stefaanhessmann Jul 22, 2026
e53a76f
test: port end-to-end statistics regression checks into the suite
stefaanhessmann Jul 22, 2026
8cfda2a
feat: teach stats functions to take explicit indices
stefaanhessmann Jul 22, 2026
664b15b
refactor: construct StatsAtomrefProvider from base dataset and train …
stefaanhessmann Jul 22, 2026
3f30c66
feat: add strict mode to the atomref query
stefaanhessmann Jul 22, 2026
72a3255
refactor: unify transform initialization to a single stats-source arg…
stefaanhessmann Jul 22, 2026
8886457
refactor!: delete the Transform.datamodule() hook
stefaanhessmann Jul 22, 2026
a05d6cc
feat: fingerprint the train partition
stefaanhessmann Jul 22, 2026
e29eee5
feat: persist training statistics to disk next to the split file
stefaanhessmann Jul 22, 2026
a480ed8
feat: persist estimated atomrefs alongside statistics
stefaanhessmann Jul 22, 2026
ae43854
docs: update statistics documentation
stefaanhessmann Jul 22, 2026
a49f095
style: apply pinned black formatting
stefaanhessmann Jul 22, 2026
16e9374
fix: compute persisted stats under the splitting lock
stefaanhessmann Jul 22, 2026
a83fddb
fix: treat unreadable stats files as empty
stefaanhessmann Jul 22, 2026
e4eac8e
fix: normalize stats-file paths to the .npz extension np.savez writes
stefaanhessmann Jul 22, 2026
8e5fa6e
test: invalidate persisted stats across an actual split regeneration
stefaanhessmann Jul 22, 2026
3ec54af
Merge pull request #786 from atomistic-machine-learning/sa/stats_refa…
sundusaijaz Jul 23, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -128,4 +128,4 @@ interfaces/lammps/examples/*/*.dat
interfaces/lammps/examples/*/deployed_model

# batchwise optimizer examples
examples/howtos/howto_batchwise_relaxations_outputs/*
examples/howtos/howto_batchwise_relaxations_outputs/*
Comment thread
sundusaijaz marked this conversation as resolved.
17 changes: 3 additions & 14 deletions docs/api/data.rst
Original file line number Diff line number Diff line change
Expand Up @@ -10,24 +10,11 @@ Atoms data
:recursive:
:template: classtemplate.rst

BaseAtomsData
ASEAtomsData
DownloadableASEAtomsData
AtomsLoader
resolve_format
AtomsDataFormat
StratifiedSampler


Creation
--------
.. autosummary::
:toctree: generated
:nosignatures:
:template: classtemplate.rst

create_dataset
load_dataset

Data modules
------------

Expand All @@ -47,5 +34,7 @@ Statistics
:template: classtemplate.rst

calculate_stats
estimate_atomrefs
StatsAtomrefProvider
NumberOfAtomsCriterion
PropertyCriterion
1 change: 0 additions & 1 deletion docs/api/datasets.rst
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,6 @@ Molecules
ISO17
MD22
QM7X
TMQM

Materials
---------
Expand Down
7 changes: 5 additions & 2 deletions docs/userguide/overview.rst
Original file line number Diff line number Diff line change
Expand Up @@ -31,8 +31,11 @@ and calculation of neighbor lists.

Furthermore, we support PyTorch Lightning datamodules through :class:`AtomsDataModule`,
which combines :class:`ASEAtomsData` with code for preparation, setup and partitioning
into train/validation/test splits. We provide specific implementations of
:class:`AtomsDataModule` for several benchmark datasets.
into train/validation/test splits. The datamodule also owns the training statistics
(per-property mean/std and atom references): it exposes them via ``get_stats`` and
``get_atomrefs``, initializes all transforms that require them, and persists computed
values next to the split file so reruns read instead of recompute. We provide specific
implementations of :class:`AtomsDataModule` for several benchmark datasets.


Model
Expand Down
85 changes: 47 additions & 38 deletions examples/howtos/howto_ensemble_calculation.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -42,7 +42,11 @@
"from ase.md.velocitydistribution import MaxwellBoltzmannDistribution\n",
"from ase.md.langevin import Langevin\n",
"\n",
"from schnetpack.interfaces.ase_interface import SpkEnsembleCalculator, AbsoluteUncertainty, RelativeUncertainty\n",
"from schnetpack.interfaces.ase_interface import (\n",
" SpkEnsembleCalculator,\n",
" AbsoluteUncertainty,\n",
" RelativeUncertainty,\n",
")\n",
"import schnetpack.transform as trn\n",
"from schnetpack.datasets import MD17\n",
"import torch\n",
Expand Down Expand Up @@ -78,11 +82,13 @@
},
"outputs": [],
"source": [
"model_path_list = ['../trained_models/rmd17_ethanol/painn_1/best_model',\n",
" '../trained_models/rmd17_ethanol/painn_2/best_model',\n",
" '../trained_models/rmd17_ethanol/painn_3/best_model',\n",
" '../trained_models/rmd17_ethanol/painn_4/best_model',\n",
" '../trained_models/rmd17_ethanol/painn_5/best_model']"
"model_path_list = [\n",
" \"../trained_models/rmd17_ethanol/painn_1/best_model\",\n",
" \"../trained_models/rmd17_ethanol/painn_2/best_model\",\n",
" \"../trained_models/rmd17_ethanol/painn_3/best_model\",\n",
" \"../trained_models/rmd17_ethanol/painn_4/best_model\",\n",
" \"../trained_models/rmd17_ethanol/painn_5/best_model\",\n",
"]"
]
},
{
Expand Down Expand Up @@ -113,7 +119,7 @@
},
"outputs": [],
"source": [
"uncertainty_abs = AbsoluteUncertainty(energy_weight=0.5,force_weight=1.0)\n",
"uncertainty_abs = AbsoluteUncertainty(energy_weight=0.5, force_weight=1.0)\n",
"uncertainty_rel = RelativeUncertainty(energy_weight=1.0, force_weight=2.0)\n",
"\n",
"uncertainty = [uncertainty_abs, uncertainty_rel]\n",
Expand All @@ -125,7 +131,8 @@
" force_key=MD17.forces,\n",
" energy_unit=\"kcal/mol\",\n",
" position_unit=\"Ang\",\n",
" uncertainty_fn=uncertainty)"
" uncertainty_fn=uncertainty,\n",
")"
]
},
{
Expand All @@ -146,8 +153,8 @@
},
"outputs": [],
"source": [
"#load data into atoms object\n",
"atoms = read('../../tests/testdata/md_ethanol.xyz', index=0)\n",
"# load data into atoms object\n",
"atoms = read(\"../../tests/testdata/md_ethanol.xyz\", index=0)\n",
"# specify atoms calculator\n",
"atoms.calc = ensemble_calculator"
]
Expand Down Expand Up @@ -290,22 +297,22 @@
"fig, ax1 = plt.subplots(figsize=(8, 6))\n",
"\n",
"# Plot absolute uncertainty on left y-axis\n",
"ax1.plot(steps, abs_vals, label=\"Absolute Uncertainty\", marker='o', color='tab:blue')\n",
"ax1.plot(steps, abs_vals, label=\"Absolute Uncertainty\", marker=\"o\", color=\"tab:blue\")\n",
"ax1.set_xlabel(\"Optimization Step\")\n",
"ax1.set_ylabel(\"Absolute Uncertainty\", color='tab:blue')\n",
"ax1.tick_params(axis='y', labelcolor='tab:blue')\n",
"ax1.set_ylabel(\"Absolute Uncertainty\", color=\"tab:blue\")\n",
"ax1.tick_params(axis=\"y\", labelcolor=\"tab:blue\")\n",
"ax1.grid(True)\n",
"\n",
"# Create second y-axis for relative uncertainty\n",
"ax2 = ax1.twinx()\n",
"ax2.plot(steps, rel_vals, label=\"Relative Uncertainty\", marker='x', color='tab:red')\n",
"ax2.set_ylabel(\"Relative Uncertainty\", color='tab:red')\n",
"ax2.tick_params(axis='y', labelcolor='tab:red')\n",
"ax2.plot(steps, rel_vals, label=\"Relative Uncertainty\", marker=\"x\", color=\"tab:red\")\n",
"ax2.set_ylabel(\"Relative Uncertainty\", color=\"tab:red\")\n",
"ax2.tick_params(axis=\"y\", labelcolor=\"tab:red\")\n",
"\n",
"# Title and layout\n",
"plt.title(\"Uncertainty during Optimization\")\n",
"fig.tight_layout()\n",
"plt.show()\n"
"plt.show()"
]
},
{
Expand Down Expand Up @@ -344,7 +351,7 @@
},
"outputs": [],
"source": [
"uncertainty_abs = AbsoluteUncertainty(energy_weight=0.5,force_weight=1.0)\n",
"uncertainty_abs = AbsoluteUncertainty(energy_weight=0.5, force_weight=1.0)\n",
"\n",
"abs_ensemble_calculator = SpkEnsembleCalculator(\n",
" models=model_path_list,\n",
Expand All @@ -353,7 +360,8 @@
" force_key=MD17.forces,\n",
" energy_unit=\"kcal/mol\",\n",
" position_unit=\"Ang\",\n",
" uncertainty_fn=uncertainty_abs)"
" uncertainty_fn=uncertainty_abs,\n",
")"
]
},
{
Expand All @@ -367,13 +375,13 @@
},
"outputs": [],
"source": [
"target_temperatures = [_ for _ in range(50, 800, 100)] \n",
"n_steps = 1000 \n",
"sampling_interval = 10 \n",
"step_size = 0.5 \n",
"target_temperatures = [_ for _ in range(50, 800, 100)]\n",
"n_steps = 1000\n",
"sampling_interval = 10\n",
"step_size = 0.5\n",
"\n",
"# setting up initial atoms\n",
"atoms = read('../../tests/testdata/md_ethanol.xyz', index=0)\n",
"atoms = read(\"../../tests/testdata/md_ethanol.xyz\", index=0)\n",
"atoms.calc = abs_ensemble_calculator\n",
"\n",
"MaxwellBoltzmannDistribution(atoms, temperature_K=target_temperatures[0])\n",
Expand All @@ -385,19 +393,19 @@
"for target_temperature in target_temperatures:\n",
" print(f\"Temp: {target_temperature:.2f} K\")\n",
" for step in tqdm(range(n_steps // sampling_interval)):\n",
" \n",
"\n",
" dyn = Langevin(\n",
" atoms, \n",
" timestep=step_size * units.fs, \n",
" atoms,\n",
" timestep=step_size * units.fs,\n",
" temperature_K=target_temperature,\n",
" friction=0.01 / units.fs\n",
" friction=0.01 / units.fs,\n",
" )\n",
" \n",
"\n",
" dyn.run(sampling_interval)\n",
" \n",
"\n",
" temp.append(atoms.get_temperature())\n",
" uncertainties.append(abs_ensemble_calculator.get_uncertainty(atoms))\n",
" \n",
"\n",
" ats_traj.append(atoms.copy())"
]
},
Expand All @@ -409,22 +417,22 @@
"source": [
"fig, ax1 = plt.subplots(figsize=(8, 6))\n",
"\n",
"ax1.plot(uncertainties, marker='o', color='blue', label='Uncertainty')\n",
"ax1.plot(uncertainties, marker=\"o\", color=\"blue\", label=\"Uncertainty\")\n",
"ax1.set_xlabel(\"MD Step\")\n",
"ax1.set_ylabel(\"Uncertainty\", color='blue')\n",
"ax1.tick_params(axis='y', labelcolor='blue')\n",
"ax1.set_ylabel(\"Uncertainty\", color=\"blue\")\n",
"ax1.tick_params(axis=\"y\", labelcolor=\"blue\")\n",
"\n",
"ax2 = ax1.twinx()\n",
"ax2.plot(temp, marker='x', color='red', label='Temperature')\n",
"ax2.set_ylabel(\"Temperature (K)\", color='red')\n",
"ax2.tick_params(axis='y', labelcolor='red')\n",
"ax2.plot(temp, marker=\"x\", color=\"red\", label=\"Temperature\")\n",
"ax2.set_ylabel(\"Temperature (K)\", color=\"red\")\n",
"ax2.tick_params(axis=\"y\", labelcolor=\"red\")\n",
"\n",
"plt.title(\"Molecular Dynamics: Uncertainty and Temperature Profile\")\n",
"ax1.grid(True)\n",
"\n",
"lines_1, labels_1 = ax1.get_legend_handles_labels()\n",
"lines_2, labels_2 = ax2.get_legend_handles_labels()\n",
"ax1.legend(lines_1 + lines_2, labels_1 + labels_2, loc='upper right')\n",
"ax1.legend(lines_1 + lines_2, labels_1 + labels_2, loc=\"upper right\")\n",
"\n",
"plt.tight_layout()\n",
"plt.show()"
Expand All @@ -444,6 +452,7 @@
"outputs": [],
"source": [
"from ase.visualize import view\n",
"\n",
"view(ats_traj)"
]
}
Expand Down
4 changes: 1 addition & 3 deletions examples/tutorials/tutorial_02_qm9.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -377,9 +377,7 @@
"from ase import Atoms\n",
"from schnetpack.utils.compatibility import load_model\n",
"\n",
"best_model = load_model(\n",
" os.path.join(qm9tut, \"best_inference_model\"), device=\"cpu\"\n",
")"
"best_model = load_model(os.path.join(qm9tut, \"best_inference_model\"), device=\"cpu\")"
]
},
{
Expand Down
1 change: 0 additions & 1 deletion src/schnetpack/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,5 +16,4 @@
from schnetpack.task import *
from schnetpack import md


__version__ = "2.2.0"
23 changes: 13 additions & 10 deletions src/schnetpack/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,11 +17,10 @@
import schnetpack as spk
from schnetpack.utils import str2class
from schnetpack.utils.script import log_hyperparameters, print_config
from schnetpack.data import BaseAtomsData, AtomsLoader
from schnetpack.data import ASEAtomsData, AtomsLoader
from schnetpack.train import PredictionWriter
from schnetpack import properties
from schnetpack.utils import load_model

from schnetpack.utils import load_model, load_task_from_checkpoint

log = logging.getLogger(__name__)

Expand Down Expand Up @@ -176,16 +175,20 @@ def train(config: DictConfig):
log.info("Starting training.")
trainer.fit(model=task, datamodule=datamodule, ckpt_path=config.run.ckpt_path)

# Evaluate model on test set after training
log.info("Starting testing.")
trainer.test(model=task, datamodule=datamodule, ckpt_path="best")

# Store best model
# Load the best checkpoint through the compatibility helper (it handles
# `weights_only` across PL versions) and test that task directly, instead
# of having Lightning re-load the checkpoint internally via
# ckpt_path="best", whose weights_only handling is version-dependent.
best_path = trainer.checkpoint_callback.best_model_path
log.info(f"Best checkpoint path:\n{best_path}")
best_task = load_task_from_checkpoint(type(task), best_path)

# Evaluate best model on test set after training
log.info("Starting testing.")
trainer.test(model=best_task, datamodule=datamodule)

# Store best model
log.info(f"Store best model")
best_task = type(task).load_from_checkpoint(best_path)
torch.save(best_task, config.globals.model_path + ".task")

best_task.save_model(config.globals.model_path, do_postprocessing=True)
Expand All @@ -195,7 +198,7 @@ def train(config: DictConfig):
@hydra.main(config_path="configs", config_name="predict", version_base="1.2")
def predict(config: DictConfig):
log.info(f"Load data from `{config.data.datapath}`")
dataset: BaseAtomsData = hydra.utils.instantiate(config.data)
dataset: ASEAtomsData = hydra.utils.instantiate(config.data)
loader = AtomsLoader(dataset, batch_size=config.batch_size, num_workers=8)

model = load_model("best_model")
Expand Down
18 changes: 9 additions & 9 deletions src/schnetpack/configs/data/ani1.yaml
Original file line number Diff line number Diff line change
@@ -1,16 +1,16 @@
# @package data
defaults:
- custom

_target_: schnetpack.datasets.ANI1
dataset:
_target_: schnetpack.datasets.ANI1
datapath: ${run.data_dir}/ani1.db # data_dir is specified in train.yaml
num_heavy_atoms: 8
high_energies: false
distance_unit: Ang
property_units:
energy: eV

datapath: ${run.data_dir}/ani1.db # data_dir is specified in train.yaml
batch_size: 32
num_train: 10000000
num_val: 100000
num_heavy_atoms: 8
high_energies: False

# convert to typically used units
distance_unit: Ang
property_units:
energy: eV
24 changes: 19 additions & 5 deletions src/schnetpack/configs/data/custom.yaml
Original file line number Diff line number Diff line change
@@ -1,12 +1,26 @@
# @package data
_target_: schnetpack.data.AtomsDataModule

datapath: ???
data_workdir: null
dataset:
_target_: schnetpack.data.ASEAtomsData
datapath: ???
load_properties: null
# null: keep the dataset's native units — setting a unit here converts on
# load.
distance_unit: null
property_units: null
transforms: null
train_transforms: null
val_transforms: null
test_transforms: null

batch_size: 10
num_train: ???
num_val: ???
num_test: null
split_file: ${run.data_dir}/split.npz
splitting: null
num_workers: 8
num_val_workers: null
num_test_workers: null
train_sampler_cls: null
train_sampler_cls: null
train_sampler_args: {}
pin_memory: false
Loading
Loading