diff --git a/.gitattributes b/.gitattributes index a6344aac8c09253b3b630fb776ae94478aa0275b..0b52315de74879860562ab84b6e4584b4ab87d91 100644 --- a/.gitattributes +++ b/.gitattributes @@ -1,35 +1,49 @@ *.7z filter=lfs diff=lfs merge=lfs -text *.arrow filter=lfs diff=lfs merge=lfs -text *.bin filter=lfs diff=lfs merge=lfs -text +*.bin.* filter=lfs diff=lfs merge=lfs -text *.bz2 filter=lfs diff=lfs merge=lfs -text -*.ckpt filter=lfs diff=lfs merge=lfs -text *.ftz filter=lfs diff=lfs merge=lfs -text *.gz filter=lfs diff=lfs merge=lfs -text *.h5 filter=lfs diff=lfs merge=lfs -text *.joblib filter=lfs diff=lfs merge=lfs -text *.lfs.* filter=lfs diff=lfs merge=lfs -text -*.mlmodel filter=lfs diff=lfs merge=lfs -text *.model filter=lfs diff=lfs merge=lfs -text *.msgpack filter=lfs diff=lfs merge=lfs -text -*.npy filter=lfs diff=lfs merge=lfs -text -*.npz filter=lfs diff=lfs merge=lfs -text *.onnx filter=lfs diff=lfs merge=lfs -text *.ot filter=lfs diff=lfs merge=lfs -text *.parquet filter=lfs diff=lfs merge=lfs -text *.pb filter=lfs diff=lfs merge=lfs -text -*.pickle filter=lfs diff=lfs merge=lfs -text -*.pkl filter=lfs diff=lfs merge=lfs -text *.pt filter=lfs diff=lfs merge=lfs -text *.pth filter=lfs diff=lfs merge=lfs -text *.rar filter=lfs diff=lfs merge=lfs -text -*.safetensors filter=lfs diff=lfs merge=lfs -text saved_model/**/* filter=lfs diff=lfs merge=lfs -text *.tar.* filter=lfs diff=lfs merge=lfs -text -*.tar filter=lfs diff=lfs merge=lfs -text *.tflite filter=lfs diff=lfs merge=lfs -text *.tgz filter=lfs diff=lfs merge=lfs -text -*.wasm filter=lfs diff=lfs merge=lfs -text *.xz filter=lfs diff=lfs merge=lfs -text *.zip filter=lfs diff=lfs merge=lfs -text +*.zstandard filter=lfs diff=lfs merge=lfs -text +*.tfevents* filter=lfs diff=lfs merge=lfs -text +*.db* filter=lfs diff=lfs merge=lfs -text +*.ark* filter=lfs diff=lfs merge=lfs -text +**/*ckpt*data* filter=lfs diff=lfs merge=lfs -text +**/*ckpt*.meta filter=lfs diff=lfs merge=lfs -text +**/*ckpt*.index filter=lfs diff=lfs merge=lfs -text +*.safetensors filter=lfs diff=lfs merge=lfs -text +*.ckpt filter=lfs diff=lfs merge=lfs -text +*.gguf* filter=lfs diff=lfs merge=lfs -text +*.ggml filter=lfs diff=lfs merge=lfs -text +*.llamafile* filter=lfs diff=lfs merge=lfs -text +*.pt2 filter=lfs diff=lfs merge=lfs -text +*.mlmodel filter=lfs diff=lfs merge=lfs -text +*.npy filter=lfs diff=lfs merge=lfs -text +*.npz filter=lfs diff=lfs merge=lfs -text +*.pickle filter=lfs diff=lfs merge=lfs -text +*.pkl filter=lfs diff=lfs merge=lfs -text +*.tar filter=lfs diff=lfs merge=lfs -text +*.wasm filter=lfs diff=lfs merge=lfs -text *.zst filter=lfs diff=lfs merge=lfs -text *tfevents* filter=lfs diff=lfs merge=lfs -text +weight/NequIP-OAM-L-0.1.nequip.pth filter=lfs diff=lfs merge=lfs -text +weight/NequIP-OAM-L-0.1.nequip.zip filter=lfs diff=lfs merge=lfs -text diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000000000000000000000000000000000000..82699b13247f9e7114768dd2dfe83823832d4374 --- /dev/null +++ b/LICENSE @@ -0,0 +1,203 @@ +Copyright 2025 Onescience Authors. All rights reserved. + + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION + + 1. Definitions. + + "License" shall mean the terms and conditions for use, reproduction, + and distribution as defined by Sections 1 through 9 of this document. + + "Licensor" shall mean the copyright owner or entity authorized by + the copyright owner that is granting the License. + + "Legal Entity" shall mean the union of the acting entity and all + other entities that control, are controlled by, or are under common + control with that entity. For the purposes of this definition, + "control" means (i) the power, direct or indirect, to cause the + direction or management of such entity, whether by contract or + otherwise, or (ii) ownership of fifty percent (50%) or more of the + outstanding shares, or (iii) beneficial ownership of such entity. + + "You" (or "Your") shall mean an individual or Legal Entity + exercising permissions granted by this License. + + "Source" form shall mean the preferred form for making modifications, + including but not limited to software source code, documentation + source, and configuration files. + + "Object" form shall mean any form resulting from mechanical + transformation or translation of a Source form, including but + not limited to compiled object code, generated documentation, + and conversions to other media types. + + "Work" shall mean the work of authorship, whether in Source or + Object form, made available under the License, as indicated by a + copyright notice that is included in or attached to the work + (an example is provided in the Appendix below). + + "Derivative Works" shall mean any work, whether in Source or Object + form, that is based on (or derived from) the Work and for which the + editorial revisions, annotations, elaborations, or other modifications + represent, as a whole, an original work of authorship. For the purposes + of this License, Derivative Works shall not include works that remain + separable from, or merely link (or bind by name) to the interfaces of, + the Work and Derivative Works thereof. + + "Contribution" shall mean any work of authorship, including + the original version of the Work and any modifications or additions + to that Work or Derivative Works thereof, that is intentionally + submitted to Licensor for inclusion in the Work by the copyright owner + or by an individual or Legal Entity authorized to submit on behalf of + the copyright owner. For the purposes of this definition, "submitted" + means any form of electronic, verbal, or written communication sent + to the Licensor or its representatives, including but not limited to + communication on electronic mailing lists, source code control systems, + and issue tracking systems that are managed by, or on behalf of, the + Licensor for the purpose of discussing and improving the Work, but + excluding communication that is conspicuously marked or otherwise + designated in writing by the copyright owner as "Not a Contribution." + + "Contributor" shall mean Licensor and any individual or Legal Entity + on behalf of whom a Contribution has been received by Licensor and + subsequently incorporated within the Work. + + 2. Grant of Copyright License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + copyright license to reproduce, prepare Derivative Works of, + publicly display, publicly perform, sublicense, and distribute the + Work and such Derivative Works in Source or Object form. + + 3. Grant of Patent License. Subject to the terms and conditions of + this License, each Contributor hereby grants to You a perpetual, + worldwide, non-exclusive, no-charge, royalty-free, irrevocable + (except as stated in this section) patent license to make, have made, + use, offer to sell, sell, import, and otherwise transfer the Work, + where such license applies only to those patent claims licensable + by such Contributor that are necessarily infringed by their + Contribution(s) alone or by combination of their Contribution(s) + with the Work to which such Contribution(s) was submitted. If You + institute patent litigation against any entity (including a + cross-claim or counterclaim in a lawsuit) alleging that the Work + or a Contribution incorporated within the Work constitutes direct + or contributory patent infringement, then any patent licenses + granted to You under this License for that Work shall terminate + as of the date such litigation is filed. + + 4. Redistribution. You may reproduce and distribute copies of the + Work or Derivative Works thereof in any medium, with or without + modifications, and in Source or Object form, provided that You + meet the following conditions: + + (a) You must give any other recipients of the Work or + Derivative Works a copy of this License; and + + (b) You must cause any modified files to carry prominent notices + stating that You changed the files; and + + (c) You must retain, in the Source form of any Derivative Works + that You distribute, all copyright, patent, trademark, and + attribution notices from the Source form of the Work, + excluding those notices that do not pertain to any part of + the Derivative Works; and + + (d) If the Work includes a "NOTICE" text file as part of its + distribution, then any Derivative Works that You distribute must + include a readable copy of the attribution notices contained + within such NOTICE file, excluding those notices that do not + pertain to any part of the Derivative Works, in at least one + of the following places: within a NOTICE text file distributed + as part of the Derivative Works; within the Source form or + documentation, if provided along with the Derivative Works; or, + within a display generated by the Derivative Works, if and + wherever such third-party notices normally appear. The contents + of the NOTICE file are for informational purposes only and + do not modify the License. You may add Your own attribution + notices within Derivative Works that You distribute, alongside + or as an addendum to the NOTICE text from the Work, provided + that such additional attribution notices cannot be construed + as modifying the License. + + You may add Your own copyright statement to Your modifications and + may provide additional or different license terms and conditions + for use, reproduction, or distribution of Your modifications, or + for any such Derivative Works as a whole, provided Your use, + reproduction, and distribution of the Work otherwise complies with + the conditions stated in this License. + + 5. Submission of Contributions. Unless You explicitly state otherwise, + any Contribution intentionally submitted for inclusion in the Work + by You to the Licensor shall be under the terms and conditions of + this License, without any additional terms or conditions. + Notwithstanding the above, nothing herein shall supersede or modify + the terms of any separate license agreement you may have executed + with Licensor regarding such Contributions. + + 6. Trademarks. This License does not grant permission to use the trade + names, trademarks, service marks, or product names of the Licensor, + except as required for reasonable and customary use in describing the + origin of the Work and reproducing the content of the NOTICE file. + + 7. Disclaimer of Warranty. Unless required by applicable law or + agreed to in writing, Licensor provides the Work (and each + Contributor provides its Contributions) on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or + implied, including, without limitation, any warranties or conditions + of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A + PARTICULAR PURPOSE. You are solely responsible for determining the + appropriateness of using or redistributing the Work and assume any + risks associated with Your exercise of permissions under this License. + + 8. Limitation of Liability. In no event and under no legal theory, + whether in tort (including negligence), contract, or otherwise, + unless required by applicable law (such as deliberate and grossly + negligent acts) or agreed to in writing, shall any Contributor be + liable to You for damages, including any direct, indirect, special, + incidental, or consequential damages of any character arising as a + result of this License or out of the use or inability to use the + Work (including but not limited to damages for loss of goodwill, + work stoppage, computer failure or malfunction, or any and all + other commercial damages or losses), even if such Contributor + has been advised of the possibility of such damages. + + 9. Accepting Warranty or Additional Liability. While redistributing + the Work or Derivative Works thereof, You may choose to offer, + and charge a fee for, acceptance of support, warranty, indemnity, + or other liability obligations and/or rights consistent with this + License. However, in accepting such obligations, You may act only + on Your own behalf and on Your sole responsibility, not on behalf + of any other Contributor, and only if You agree to indemnify, + defend, and hold each Contributor harmless for any liability + incurred by, or claims asserted against, such Contributor by reason + of your accepting any such warranty or additional liability. + + END OF TERMS AND CONDITIONS + + APPENDIX: How to apply the Apache License to your work. + + To apply the Apache License to your work, attach the following + boilerplate notice, with the fields enclosed by brackets "[]" + replaced with your own identifying information. (Don't include + the brackets!) The text should be enclosed in the appropriate + comment syntax for the file format. We also recommend that a + file or class name and description of purpose be included on the + same "printed page" as the copyright notice for easier + identification within third-party archives. + + Copyright 2025 Onescience Authors. + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/README.md b/README.md new file mode 100644 index 0000000000000000000000000000000000000000..2cceb94d1fe34a1db13d32824f9c737ed92b493b --- /dev/null +++ b/README.md @@ -0,0 +1,221 @@ +--- +license: apache-2.0 +tasks: + - materials-simulation + - molecular-dynamics + - energy-prediction + - force-prediction +frameworks: + - pytorch +language: + - en +tags: + - OneScience + - NequIP + - machine-learning-potential + - molecular-simulation + - materials-computing + - graph-neural-network + - equivariant-neural-network + - training + - fine-tuning + - inference +datasets: + - OneScience-Group/FCC_Cu +--- +

+ + NequIP + +

+ +# Model Introduction + +NequIP is a machine-learning interatomic potential (MLIP) for molecular and materials systems. Built on an E(3)-equivariant graph neural network, it predicts the energies and forces of atomic structures. + +Reference implementation: https://github.com/mir-group/nequip + +# Model Description + +This repository provides the OneScience-integrated NequIP model code, OAM-L model weights, and runnable examples for training, fine-tuning, and inference. The `model/` directory corresponds only to `src/onescience/models/nequip/` in the main OneScience repository; training utilities, data-processing tools, and other shared modules are provided by the installed OneScience package. + +The included OAM-L weights are: + +| File | Purpose | +| --- | --- | +| `weight/NequIP-OAM-L-0.1.nequip.pth` | Compiled model for ASE single-point energy, atomic force, and stress inference | +| `weight/NequIP-OAM-L-0.1.nequip.zip` | NequIP package for OAM-L fine-tuning and checkpoint inference | + +# Use Cases + +| Use case | Description | +| :---: | :--- | +| Interatomic-potential training | Train a NequIP model using the example configurations and ASE extxyz data | +| Pretrained-model fine-tuning | Fine-tune the OAM-L package using data labeled with energy and forces | +| Single-point energy and force inference | Predict the energy, atomic forces, and stress of a structure with a compiled model or fine-tuned checkpoint | +| Structure relaxation | Optimize atomic positions with ASE | +| Energy-volume curve | Scan the volume of a periodic crystal and calculate the corresponding energy | +| Slurm/DCU training | Submit single-device or multi-device jobs using the included configurations and launch scripts | + +# Usage + +## 1. Using OneCode + +Try intelligent, one-click AI4S programming in the OneCode online environment: + +[Try intelligent, one-click AI4S programming](https://web-2069360198568017922-iaaj.ksai.scnet.cn:58043/home) + +## 2. Manual Installation and Usage + +**Hardware requirements** + +- A GPU or DCU is recommended for training. +- A CPU can be used for import checks and small-configuration connectivity tests, but full training will be slow. +- DCU users must install DTK in advance. DTK 25.04.2 or later, or the OneScience-recommended version matching the current cluster, is recommended. + +### Download the Model Package + +```bash +hf download --model OneScience-Group/NequIP --local-dir ./NequIP +cd NequIP +``` + +### Install the Runtime Environment + +**DCU environment** + +```bash +# Activate DTK and Conda first +conda create -n onescience311 python=3.11 -y +conda activate onescience311 +# uv installation is also supported +pip install onescience[matchem-dcu] -i http://mirrors.onescience.ai:3141/pypi/simple/ --trusted-host mirrors.onescience.ai +``` + +**GPU environment** + +```bash +# Activate Conda first +conda create -n onescience311 python=3.11 -y libstdcxx-ng=12 libgcc-ng=12 gcc_linux-64=12 gxx_linux-64=12 +conda activate onescience311 +# uv installation is also supported +pip install onescience[matchem-gpu] -i http://mirrors.onescience.ai:3141/pypi/simple/ --trusted-host mirrors.onescience.ai +``` + +### Training Data + +Training data is not bundled with this repository. Using the introductory FCC Cu dataset as an example, download it from Hugging Face to `data/` in the repository root: + +```bash +hf download --dataset OneScience-Group/FCC_Cu --local-dir ./data +``` + +After downloading, the raw data is located at `data/data/FCC_Cu/raw/fcu.xyz`. The dataset contains 6,855 structures, each with 52 atoms of C, H, O, and Cu. It uses an ASE-readable extxyz format and includes periodic cells together with energy and force labels. For production training or fine-tuning, use data consistent with the target system, label definitions, and units. + +The training scripts read models and data from shared directories. Set the paths for your cluster before training or fine-tuning: + +```bash +export ONESCIENCE_MODELS_DIR=/path/to/onescience-models +export ONESCIENCE_DATASETS_DIR=/path/to/onescience-datasets +``` + +To use the OAM-L weights included in this repository, copy them into the shared model directory: + +```bash +mkdir -p "$ONESCIENCE_MODELS_DIR/NequIP" +cp weight/NequIP-OAM-L-0.1.nequip.pth "$ONESCIENCE_MODELS_DIR/NequIP/" +cp weight/NequIP-OAM-L-0.1.nequip.zip "$ONESCIENCE_MODELS_DIR/NequIP/" +``` + +### Training + +Generate minimal smoke-test data and run local training: + +```bash +python demo/prepare_smoke_data.py +bash demo/run.sh --config configs/tutorial_smoke.yaml +``` + +Download the official FCU tutorial data and submit a training job: + +```bash +python demo/download_tutorial_data.py +bash demo/run.sh --config configs/tutorial_fcu.yaml --submit +``` + +The eight-DCU configurations run locally or submit to Slurm automatically, depending on the currently available resources: + +```bash +bash demo/run.sh --config configs/tutorial_smoke_8dcu.yaml +bash demo/run.sh --config configs/tutorial_fcu_8dcu.yaml +``` + +Outputs are written to `outputs/` by default. The actual wait time for a training job depends on the cluster queue and available resources. + +### Model Weights + +This repository includes the OAM-L trained weights: + +```text +e83a1d656f8b19b55d2f05708c83e054612f713e9a1b06266aa010db58e56517 weight/NequIP-OAM-L-0.1.nequip.pth +5d01a4fab228abb3cdb6ace0033f93993729956bca6a42234a2a8816825b9a0f weight/NequIP-OAM-L-0.1.nequip.zip +``` + +### Fine-Tuning + +Validate the OAM-L fine-tuning workflow with generated smoke-test data: + +```bash +python demo/prepare_smoke_data.py +bash demo/run.sh --config configs/oam_l_finetune_smoke.yaml --submit +``` + +Use the production fine-tuning configuration: + +```bash +bash demo/run.sh --config configs/oam_l_finetune.yaml --submit +``` + +Provide production fine-tuning data through `ONESCIENCE_DATASETS_DIR` in an ASE-readable extxyz format. Every frame must contain at least `energy` and `forces`; element types, units, and label definitions must be consistent with the OAM-L package and configuration. + +### Inference + +Use the compiled model for single-point energy, atomic force, and stress prediction: + +```bash +python single_point.py --compiled-model weight/NequIP-OAM-L-0.1.nequip.pth +python single_point.py \ + --compiled-model weight/NequIP-OAM-L-0.1.nequip.pth \ + --input structure.cif \ + --output outputs/single_point.json +``` + +Calculate an energy-volume curve and perform structure relaxation: + +```bash +python energy_volume.py +python structure_relaxation.py --fmax 0.05 --steps 100 --output-dir outputs/oam_l_relax +``` + +Run inference with a checkpoint produced by fine-tuning: + +```bash +python single_point.py \ + --checkpoint outputs//checkpoints/best.ckpt \ + --package weight/NequIP-OAM-L-0.1.nequip.zip \ + --output outputs//single_point.json +``` + +# Official OneScience Resources + +| Platform | OneScience Main Repository | Skills Repository | +| --- | --- | --- | +| Gitee | https://gitee.com/onescience-ai/onescience | https://gitee.com/onescience-ai/oneskills | +| GitHub | https://github.com/onescience-ai/OneScience | https://github.com/onescience-ai/oneskills | + +# Citation and License + +- The NequIP-related code comes from the OneScience MatChem integration and refers to the upstream NequIP project (https://github.com/mir-group/nequip). The OneScience integration code follows the [Apache License 2.0](https://www.apache.org/licenses/LICENSE-2.0) used by the main repository. +- If you use NequIP or OAM-L training results in research, please cite the original NequIP paper, the relevant OneScience projects, and the datasets used. +- Redistribution rights for the OAM-L model weights are governed by the original OneScience/OAM-L release terms. Confirm the applicable rights and restrictions before use. +- The FCC Cu dataset is published separately at [OneScience-Group/FCC_Cu](https://huggingface.co/datasets/OneScience-Group/FCC_Cu). Its license and provenance are documented on the dataset card and by its upstream source. diff --git a/demo/_parse_config.py b/demo/_parse_config.py new file mode 100644 index 0000000000000000000000000000000000000000..5d29c5d0157874f97215c160bf1065b27f94bce2 --- /dev/null +++ b/demo/_parse_config.py @@ -0,0 +1,115 @@ +#!/usr/bin/env python3 +"""Extract NequIP demo launch metadata and its Hydra training config.""" + +from __future__ import annotations + +import re +import shlex +import sys +from pathlib import Path + +import yaml + + +META_KEYS = {"name", "launch", "slurm", "env"} + + +def _assignment(name: str, value) -> None: + print(f"{name}={shlex.quote(str(value))}") + + +def _positive_int(value, field: str) -> int: + try: + parsed = int(value) + except (TypeError, ValueError) as error: + raise ValueError(f"{field} must be a positive integer") from error + if parsed < 1: + raise ValueError(f"{field} must be a positive integer") + return parsed + + +def _config(path: str) -> dict: + text = Path(path).read_text(encoding="utf-8") + text = re.sub( + r"\$\{demo_dir:([^}]+)\}", + lambda match: str(Path(__file__).parent.resolve() / match.group(1)), + text, + ) + return yaml.safe_load(text) or {} + + +def _print_launch(cfg: dict) -> None: + launch = cfg.get("launch", {}) or {} + trainer = cfg.get("trainer", {}) or {} + mode = launch.get("mode", "local") + if mode not in {"auto", "local", "submit"}: + raise ValueError("launch.mode must be 'auto', 'local', or 'submit'") + + nodes = _positive_int(launch.get("num_nodes", 1), "launch.num_nodes") + devices = _positive_int(launch.get("num_gpus", 1), "launch.num_gpus") + trainer_nodes = _positive_int(trainer.get("num_nodes", 1), "trainer.num_nodes") + trainer_devices = _positive_int(trainer.get("devices", 1), "trainer.devices") + if nodes != trainer_nodes: + raise ValueError("launch.num_nodes must equal trainer.num_nodes") + if devices != trainer_devices: + raise ValueError("launch.num_gpus must equal trainer.devices") + _assignment("RUN_MODE", mode) + _assignment("NODES", nodes) + _assignment("GPUS_PER_NODE", devices) + _assignment("WORLD_SIZE", nodes * devices) + + +def _print_slurm(cfg: dict) -> None: + slurm = cfg.get("slurm", {}) or {} + _assignment("PARTITION", slurm.get("partition", "hx1hdnormal01")) + _assignment("TIME", slurm.get("time", "01:00:00")) + _assignment( + "CPUS_PER_TASK", + _positive_int(slurm.get("cpus_per_task", 8), "slurm.cpus_per_task"), + ) + _assignment("NODELIST", slurm.get("nodelist", "")) + + +def _print_env(cfg: dict) -> None: + env = cfg.get("env", {}) or {} + for name, value in env.items(): + if not re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*", name): + raise ValueError(f"invalid environment variable name: {name}") + _assignment(f"export {name}", value) + + +def _print_training_config(cfg: dict) -> None: + training = {key: value for key, value in cfg.items() if key not in META_KEYS} + yaml.safe_dump( + training, + sys.stdout, + sort_keys=False, + default_flow_style=False, + allow_unicode=True, + ) + + +def main() -> None: + if len(sys.argv) != 3: + raise SystemExit( + "usage: _parse_config.py " + "" + ) + cfg = _config(sys.argv[1]) + action = sys.argv[2] + actions = { + "name": lambda: print(cfg.get("name", "nequip_run")), + "launch": lambda: _print_launch(cfg), + "slurm": lambda: _print_slurm(cfg), + "env": lambda: _print_env(cfg), + "training-config": lambda: _print_training_config(cfg), + "finetune-config": lambda: _print_training_config(cfg), + } + try: + actions[action]() + except KeyError as error: + raise SystemExit(f"unknown action: {action}") from error + + +if __name__ == "__main__": + main() diff --git a/demo/configs/oam_l_finetune.yaml b/demo/configs/oam_l_finetune.yaml new file mode 100644 index 0000000000000000000000000000000000000000..b0e912ab924f866e359ee086bd23eafe0b1c4047 --- /dev/null +++ b/demo/configs/oam_l_finetune.yaml @@ -0,0 +1,104 @@ +# Production-shaped fine-tuning template for the official NequIP OAM-L model. +# NequIP does not prescribe a universal fine-tuning dataset or epoch count. +# Supply consistent energy/force reference calculations at the path below. +run: [train, test] + +package_path: ${oc.env:ONESCIENCE_MODELS_DIR}/NequIP/NequIP-OAM-L-0.1.nequip.zip +model_type_names: ${type_names_from_package:${package_path}} +cutoff_radius: ${cutoff_radius_from_package:${package_path}} +monitored_metric: val0_epoch/weighted_sum + +data: + _target_: onescience.datapipes.materials.nequip.datamodule.ASEDataModule + seed: 456 + split_dataset: + file_path: ${oc.env:ONESCIENCE_DATASETS_DIR}/matchem/NequIP/oam_l_finetune.xyz + train: 0.8 + val: 0.1 + test: 0.1 + transforms: + - _target_: onescience.datapipes.materials.nequip.transforms.ChemicalSpeciesToAtomTypeMapper + model_type_names: ${model_type_names} + - _target_: onescience.datapipes.materials.nequip.transforms.NeighborListTransform + r_max: ${cutoff_radius} + train_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 1 + num_workers: 0 + shuffle: true + val_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 4 + num_workers: 0 + test_dataloader: ${data.val_dataloader} + +trainer: + _target_: lightning.Trainer + accelerator: gpu + devices: 1 + num_nodes: 1 + max_epochs: 100 + enable_checkpointing: true + logger: + _target_: lightning.pytorch.loggers.CSVLogger + save_dir: ${hydra:runtime.output_dir} + name: metrics + version: 0 + enable_progress_bar: true + log_every_n_steps: 10 + callbacks: + - _target_: onescience.utils.nequip.train.callbacks.PlainTextMetricsLogger + - _target_: lightning.pytorch.callbacks.EarlyStopping + monitor: ${monitored_metric} + min_delta: 1e-4 + patience: 10 + - _target_: lightning.pytorch.callbacks.ModelCheckpoint + monitor: ${monitored_metric} + dirpath: ${hydra:runtime.output_dir}/checkpoints + filename: best + save_last: true + +training_module: + _target_: onescience.utils.nequip.train.EMALightningModule + ema_decay: 0.999 + loss: + _target_: onescience.utils.nequip.train.EnergyForceLoss + per_atom_energy: true + coeffs: + total_energy: 1.0 + forces: 1.0 + train_metrics: + _target_: onescience.utils.nequip.train.EnergyForceMetrics + coeffs: + total_energy_mae: 1.0 + forces_mae: 1.0 + val_metrics: ${training_module.train_metrics} + test_metrics: ${training_module.train_metrics} + optimizer: + _target_: torch.optim.Adam + lr: 1.0e-5 + lr_scheduler: + scheduler: + _target_: torch.optim.lr_scheduler.ReduceLROnPlateau + factor: 0.5 + patience: 5 + min_lr: 1.0e-7 + monitor: ${monitored_metric} + interval: epoch + frequency: 1 + model: + _target_: onescience.models.nequip.model.ModelFromPackage + package_path: ${package_path} + +name: nequip_oam_l_finetune +launch: + mode: local + num_nodes: 1 + num_gpus: 1 +slurm: + partition: hx1hdnormal01 + nodelist: "" + time: "1-00:00:00" + cpus_per_task: 8 +env: + OMP_NUM_THREADS: 8 diff --git a/demo/configs/oam_l_finetune_smoke.yaml b/demo/configs/oam_l_finetune_smoke.yaml new file mode 100644 index 0000000000000000000000000000000000000000..95163b575a3fa8319aa9c2f82912c0a176a317f9 --- /dev/null +++ b/demo/configs/oam_l_finetune_smoke.yaml @@ -0,0 +1,95 @@ +# Fine-tuning smoke test for the official NequIP OAM-L package. +# The bundled Cu data only validates the workflow; replace it with consistent +# reference calculations and remove the batch limits for a scientific run. +run: [train, test] + +package_path: ${oc.env:ONESCIENCE_MODELS_DIR}/NequIP/NequIP-OAM-L-0.1.nequip.zip +model_type_names: ${type_names_from_package:${package_path}} +cutoff_radius: ${cutoff_radius_from_package:${package_path}} +monitored_metric: val0_epoch/weighted_sum + +data: + _target_: onescience.datapipes.materials.nequip.datamodule.ASEDataModule + seed: 456 + split_dataset: + file_path: ${demo_dir:reference_data/smoke.xyz} + train: 0.75 + val: 0.125 + test: 0.125 + transforms: + - _target_: onescience.datapipes.materials.nequip.transforms.ChemicalSpeciesToAtomTypeMapper + model_type_names: ${model_type_names} + - _target_: onescience.datapipes.materials.nequip.transforms.NeighborListTransform + r_max: ${cutoff_radius} + train_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 1 + num_workers: 0 + shuffle: true + val_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 1 + num_workers: 0 + test_dataloader: ${data.val_dataloader} + +trainer: + _target_: lightning.Trainer + accelerator: gpu + devices: 1 + num_nodes: 1 + max_epochs: 1 + limit_train_batches: 1 + limit_val_batches: 1 + limit_test_batches: 1 + num_sanity_val_steps: 0 + enable_checkpointing: true + logger: + _target_: lightning.pytorch.loggers.CSVLogger + save_dir: ${hydra:runtime.output_dir} + name: metrics + version: 0 + enable_progress_bar: true + log_every_n_steps: 1 + callbacks: + - _target_: onescience.utils.nequip.train.callbacks.PlainTextMetricsLogger + - _target_: lightning.pytorch.callbacks.ModelCheckpoint + monitor: ${monitored_metric} + dirpath: ${hydra:runtime.output_dir}/checkpoints + filename: best + save_last: true + +training_module: + _target_: onescience.utils.nequip.train.EMALightningModule + ema_decay: 0.999 + loss: + _target_: onescience.utils.nequip.train.EnergyForceLoss + per_atom_energy: true + coeffs: + total_energy: 1.0 + forces: 1.0 + train_metrics: + _target_: onescience.utils.nequip.train.EnergyForceMetrics + coeffs: + total_energy_mae: 1.0 + forces_mae: 1.0 + val_metrics: ${training_module.train_metrics} + test_metrics: ${training_module.train_metrics} + optimizer: + _target_: torch.optim.Adam + lr: 1.0e-5 + model: + _target_: onescience.models.nequip.model.ModelFromPackage + package_path: ${package_path} + +name: nequip_oam_l_finetune_smoke +launch: + mode: local + num_nodes: 1 + num_gpus: 1 +slurm: + partition: hx1hdnormal01 + nodelist: "" + time: "00:30:00" + cpus_per_task: 8 +env: + OMP_NUM_THREADS: 8 diff --git a/demo/configs/tutorial_fcu.yaml b/demo/configs/tutorial_fcu.yaml new file mode 100644 index 0000000000000000000000000000000000000000..7a03b929f399d18a74d7c959130a54f93b3f8f2e --- /dev/null +++ b/demo/configs/tutorial_fcu.yaml @@ -0,0 +1,130 @@ +# NequIP 0.19 tutorial reproduction using the official fcu.xyz dataset. +# The validated smoke schedule uses two complete epochs. Set max_epochs to 1000 +# to match the upstream tutorial's full training schedule. +run: [train, test] + +cutoff_radius: 5.0 +num_layers: 4 +l_max: 1 +num_features: 32 +model_type_names: [C, H, O, Cu] +chemical_species: ${model_type_names} +monitored_metric: val0_epoch/weighted_sum + +data: + _target_: onescience.datapipes.materials.nequip.datamodule.ASEDataModule + seed: 456 + split_dataset: + file_path: ${oc.env:ONESCIENCE_DATASETS_DIR}/matchem/NequIP/fcu.xyz + train: 0.8 + val: 0.1 + test: 0.1 + transforms: + - _target_: onescience.datapipes.materials.nequip.transforms.ChemicalSpeciesToAtomTypeMapper + model_type_names: ${model_type_names} + - _target_: onescience.datapipes.materials.nequip.transforms.NeighborListTransform + r_max: ${cutoff_radius} + train_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 5 + num_workers: 0 + shuffle: true + val_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 10 + num_workers: 0 + test_dataloader: ${data.val_dataloader} + stats_manager: + _target_: onescience.datapipes.materials.nequip.CommonDataStatisticsManager + dataloader_kwargs: + batch_size: 10 + type_names: ${model_type_names} + +trainer: + _target_: lightning.Trainer + accelerator: gpu + devices: 1 + num_nodes: 1 + enable_checkpointing: true + max_epochs: 2 + log_every_n_steps: 1 + logger: false + enable_progress_bar: false + callbacks: + - _target_: lightning.pytorch.callbacks.EarlyStopping + monitor: ${monitored_metric} + min_delta: 1e-3 + patience: 20 + - _target_: lightning.pytorch.callbacks.ModelCheckpoint + monitor: ${monitored_metric} + dirpath: ${hydra:runtime.output_dir}/checkpoints + filename: best + save_last: true + +training_module: + _target_: onescience.utils.nequip.train.EMALightningModule + ema_decay: 0.999 + loss: + _target_: onescience.utils.nequip.train.EnergyForceLoss + per_atom_energy: true + coeffs: + total_energy: 1.0 + forces: 1.0 + val_metrics: + _target_: onescience.utils.nequip.train.EnergyForceMetrics + coeffs: + total_energy_mae: 1.0 + forces_mae: 1.0 + train_metrics: ${training_module.val_metrics} + test_metrics: ${training_module.val_metrics} + optimizer: + _target_: torch.optim.Adam + lr: 0.01 + lr_scheduler: + scheduler: + _target_: torch.optim.lr_scheduler.ReduceLROnPlateau + factor: 0.6 + patience: 5 + threshold: 0.2 + min_lr: 1e-6 + monitor: ${monitored_metric} + interval: epoch + frequency: 1 + model: + _target_: onescience.models.nequip.model.NequIPGNNModel + compile_mode: eager + seed: 456 + model_dtype: float32 + type_names: ${model_type_names} + r_max: ${cutoff_radius} + num_bessels: 8 + bessel_trainable: false + polynomial_cutoff_p: 6 + num_layers: ${num_layers} + l_max: ${l_max} + parity: true + num_features: ${num_features} + radial_mlp_depth: 2 + radial_mlp_width: 64 + avg_num_neighbors: ${training_data_stats:num_neighbors_mean} + per_type_energy_scales: ${training_data_stats:per_type_forces_rms} + per_type_energy_shifts: ${training_data_stats:per_atom_energy_mean} + per_type_energy_scales_trainable: false + per_type_energy_shifts_trainable: false + pair_potential: + _target_: onescience.models.nequip.nn.pair_potential.ZBL + units: metal + chemical_species: ${chemical_species} + +name: nequip_fcu_tutorial +launch: + mode: local + num_nodes: 1 + num_gpus: 1 +slurm: + partition: hx1hdnormal01 + nodelist: a01r1n02 + time: "00:30:00" + cpus_per_task: 8 +env: + OMP_NUM_THREADS: 8 diff --git a/demo/configs/tutorial_fcu_8dcu.yaml b/demo/configs/tutorial_fcu_8dcu.yaml new file mode 100644 index 0000000000000000000000000000000000000000..e0a640036365abe33ef3f4d6f44b985bc90e9434 --- /dev/null +++ b/demo/configs/tutorial_fcu_8dcu.yaml @@ -0,0 +1,140 @@ +# Eight-DCU DDP training plan for the official NequIP fcu.xyz tutorial. +# DataLoader batch size is per rank: batch_size=5 gives a global batch of 40. +# The learning rate remains conservative at the official 0.01; tune it for +# production based on validation behavior rather than scaling it blindly. +run: [train, test] + +cutoff_radius: 5.0 +num_layers: 4 +l_max: 1 +num_features: 32 +model_type_names: [C, H, O, Cu] +chemical_species: ${model_type_names} +monitored_metric: val0_epoch/weighted_sum + +data: + _target_: onescience.datapipes.materials.nequip.datamodule.ASEDataModule + seed: 456 + split_dataset: + file_path: ${oc.env:ONESCIENCE_DATASETS_DIR}/matchem/NequIP/fcu.xyz + train: 0.8 + val: 0.1 + test: 0.1 + transforms: + - _target_: onescience.datapipes.materials.nequip.transforms.ChemicalSpeciesToAtomTypeMapper + model_type_names: ${model_type_names} + - _target_: onescience.datapipes.materials.nequip.transforms.NeighborListTransform + r_max: ${cutoff_radius} + train_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 5 + num_workers: 0 + shuffle: true + val_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 10 + num_workers: 0 + test_dataloader: ${data.val_dataloader} + stats_manager: + _target_: onescience.datapipes.materials.nequip.CommonDataStatisticsManager + dataloader_kwargs: + batch_size: 10 + type_names: ${model_type_names} + +trainer: + _target_: lightning.Trainer + accelerator: gpu + devices: 8 + num_nodes: 1 + strategy: + _target_: lightning.pytorch.strategies.DDPStrategy + enable_checkpointing: true + max_epochs: 1000 + max_time: "03:00:00:00" + log_every_n_steps: 20 + logger: + _target_: lightning.pytorch.loggers.CSVLogger + save_dir: ${hydra:runtime.output_dir} + name: metrics + version: 0 + enable_progress_bar: true + callbacks: + - _target_: onescience.utils.nequip.train.callbacks.PlainTextMetricsLogger + - _target_: lightning.pytorch.callbacks.EarlyStopping + monitor: ${monitored_metric} + min_delta: 1e-3 + patience: 20 + - _target_: lightning.pytorch.callbacks.ModelCheckpoint + monitor: ${monitored_metric} + dirpath: ${hydra:runtime.output_dir}/checkpoints + filename: best + save_last: true + +training_module: + _target_: onescience.utils.nequip.train.EMALightningModule + ema_decay: 0.999 + loss: + _target_: onescience.utils.nequip.train.EnergyForceLoss + per_atom_energy: true + coeffs: + total_energy: 1.0 + forces: 1.0 + val_metrics: + _target_: onescience.utils.nequip.train.EnergyForceMetrics + coeffs: + total_energy_mae: 1.0 + forces_mae: 1.0 + train_metrics: ${training_module.val_metrics} + test_metrics: ${training_module.val_metrics} + optimizer: + _target_: torch.optim.Adam + lr: 0.01 + lr_scheduler: + scheduler: + _target_: torch.optim.lr_scheduler.ReduceLROnPlateau + factor: 0.6 + patience: 5 + threshold: 0.2 + min_lr: 1.0e-6 + monitor: ${monitored_metric} + interval: epoch + frequency: 1 + model: + _target_: onescience.models.nequip.model.NequIPGNNModel + compile_mode: eager + seed: 456 + model_dtype: float32 + type_names: ${model_type_names} + r_max: ${cutoff_radius} + num_bessels: 8 + bessel_trainable: false + polynomial_cutoff_p: 6 + num_layers: ${num_layers} + l_max: ${l_max} + parity: true + num_features: ${num_features} + radial_mlp_depth: 2 + radial_mlp_width: 64 + avg_num_neighbors: ${training_data_stats:num_neighbors_mean} + per_type_energy_scales: ${training_data_stats:per_type_forces_rms} + per_type_energy_shifts: ${training_data_stats:per_atom_energy_mean} + per_type_energy_scales_trainable: false + per_type_energy_shifts_trainable: false + pair_potential: + _target_: onescience.models.nequip.nn.pair_potential.ZBL + units: metal + chemical_species: ${chemical_species} + +name: nequip_fcu_tutorial_8dcu +launch: + mode: auto + num_nodes: 1 + num_gpus: 8 +slurm: + partition: hx1hdnormal01 + nodelist: "" + time: "3-00:00:00" + cpus_per_task: 8 +env: + OMP_NUM_THREADS: 8 + NCCL_DEBUG: "WARN" diff --git a/demo/configs/tutorial_fcu_full.yaml b/demo/configs/tutorial_fcu_full.yaml new file mode 100644 index 0000000000000000000000000000000000000000..e6f8460785b91092a1c90367e6ce3a5eef2bd9c8 --- /dev/null +++ b/demo/configs/tutorial_fcu_full.yaml @@ -0,0 +1,131 @@ +# Full OneScience run of the official NequIP fcu.xyz tutorial configuration. +# The model, data split, losses, and 1000-epoch schedule follow upstream. Eager +# mode and local logging are retained for the validated PyTorch 2.5 DTK stack. +run: [train, test] + +cutoff_radius: 5.0 +num_layers: 4 +l_max: 1 +num_features: 32 +model_type_names: [C, H, O, Cu] +chemical_species: ${model_type_names} +monitored_metric: val0_epoch/weighted_sum + +data: + _target_: onescience.datapipes.materials.nequip.datamodule.ASEDataModule + seed: 456 + split_dataset: + file_path: ${oc.env:ONESCIENCE_DATASETS_DIR}/matchem/NequIP/fcu.xyz + train: 0.8 + val: 0.1 + test: 0.1 + transforms: + - _target_: onescience.datapipes.materials.nequip.transforms.ChemicalSpeciesToAtomTypeMapper + model_type_names: ${model_type_names} + - _target_: onescience.datapipes.materials.nequip.transforms.NeighborListTransform + r_max: ${cutoff_radius} + train_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 5 + num_workers: 0 + shuffle: true + val_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 10 + num_workers: 0 + test_dataloader: ${data.val_dataloader} + stats_manager: + _target_: onescience.datapipes.materials.nequip.CommonDataStatisticsManager + dataloader_kwargs: + batch_size: 10 + type_names: ${model_type_names} + +trainer: + _target_: lightning.Trainer + accelerator: gpu + devices: 1 + num_nodes: 1 + enable_checkpointing: true + max_epochs: 1000 + max_time: "03:00:00:00" + log_every_n_steps: 20 + logger: false + enable_progress_bar: false + callbacks: + - _target_: lightning.pytorch.callbacks.EarlyStopping + monitor: ${monitored_metric} + min_delta: 1e-3 + patience: 20 + - _target_: lightning.pytorch.callbacks.ModelCheckpoint + monitor: ${monitored_metric} + dirpath: ${hydra:runtime.output_dir}/checkpoints + filename: best + save_last: true + +training_module: + _target_: onescience.utils.nequip.train.EMALightningModule + ema_decay: 0.999 + loss: + _target_: onescience.utils.nequip.train.EnergyForceLoss + per_atom_energy: true + coeffs: + total_energy: 1.0 + forces: 1.0 + val_metrics: + _target_: onescience.utils.nequip.train.EnergyForceMetrics + coeffs: + total_energy_mae: 1.0 + forces_mae: 1.0 + train_metrics: ${training_module.val_metrics} + test_metrics: ${training_module.val_metrics} + optimizer: + _target_: torch.optim.Adam + lr: 0.01 + lr_scheduler: + scheduler: + _target_: torch.optim.lr_scheduler.ReduceLROnPlateau + factor: 0.6 + patience: 5 + threshold: 0.2 + min_lr: 1.0e-6 + monitor: ${monitored_metric} + interval: epoch + frequency: 1 + model: + _target_: onescience.models.nequip.model.NequIPGNNModel + compile_mode: eager + seed: 456 + model_dtype: float32 + type_names: ${model_type_names} + r_max: ${cutoff_radius} + num_bessels: 8 + bessel_trainable: false + polynomial_cutoff_p: 6 + num_layers: ${num_layers} + l_max: ${l_max} + parity: true + num_features: ${num_features} + radial_mlp_depth: 2 + radial_mlp_width: 64 + avg_num_neighbors: ${training_data_stats:num_neighbors_mean} + per_type_energy_scales: ${training_data_stats:per_type_forces_rms} + per_type_energy_shifts: ${training_data_stats:per_atom_energy_mean} + per_type_energy_scales_trainable: false + per_type_energy_shifts_trainable: false + pair_potential: + _target_: onescience.models.nequip.nn.pair_potential.ZBL + units: metal + chemical_species: ${chemical_species} + +name: nequip_fcu_tutorial_full +launch: + mode: local + num_nodes: 1 + num_gpus: 1 +slurm: + partition: hx1hdnormal01 + nodelist: a01r1n02 + time: "3-00:00:00" + cpus_per_task: 8 +env: + OMP_NUM_THREADS: 8 diff --git a/demo/configs/tutorial_smoke.yaml b/demo/configs/tutorial_smoke.yaml new file mode 100644 index 0000000000000000000000000000000000000000..e5e5321655c6439f0bb23e5dc6a551334f822422 --- /dev/null +++ b/demo/configs/tutorial_smoke.yaml @@ -0,0 +1,146 @@ +# yamllint disable rule:line-length +# Smoke-test config for NequIP on OneScience. +# This config uses a tiny synthetic Cu dataset and a small model to verify the +# training pipeline (data loading, model build, forward, backward, checkpoint). + +run: [train, test] + +cutoff_radius: 4.0 + +num_layers: 2 +l_max: 1 +num_features: 8 + +model_type_names: [Cu] +chemical_species: ${model_type_names} + +monitored_metric: val0_epoch/weighted_sum + +# ============ +# DATA +# ============ +data: + _target_: onescience.datapipes.materials.nequip.datamodule.ASEDataModule + seed: 456 + + split_dataset: + file_path: ${demo_dir:reference_data/smoke.xyz} + train: 0.75 + val: 0.125 + test: 0.125 + + transforms: + - _target_: onescience.datapipes.materials.nequip.transforms.ChemicalSpeciesToAtomTypeMapper + model_type_names: ${model_type_names} + - _target_: onescience.datapipes.materials.nequip.transforms.NeighborListTransform + r_max: ${cutoff_radius} + + train_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 2 + num_workers: 0 + shuffle: true + val_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 2 + num_workers: 0 + test_dataloader: ${data.val_dataloader} + + stats_manager: + _target_: onescience.datapipes.materials.nequip.CommonDataStatisticsManager + dataloader_kwargs: + batch_size: 2 + type_names: ${model_type_names} + +# ============= +# TRAINER +# ============= +trainer: + _target_: lightning.Trainer + accelerator: gpu + devices: 1 + num_nodes: 1 + enable_checkpointing: true + max_epochs: 2 + log_every_n_steps: 1 + logger: false + enable_progress_bar: true + + callbacks: + - _target_: lightning.pytorch.callbacks.ModelCheckpoint + monitor: ${monitored_metric} + dirpath: ${hydra:runtime.output_dir}/checkpoints + filename: best + save_last: true + +# ===================== +# TRAINING MODULE +# ===================== +training_module: + _target_: onescience.utils.nequip.train.EMALightningModule + ema_decay: 0.999 + + loss: + _target_: onescience.utils.nequip.train.EnergyForceLoss + per_atom_energy: true + coeffs: + total_energy: 1.0 + forces: 1.0 + + val_metrics: + _target_: onescience.utils.nequip.train.EnergyForceMetrics + coeffs: + total_energy_mae: 1.0 + forces_mae: 1.0 + train_metrics: ${training_module.val_metrics} + test_metrics: ${training_module.val_metrics} + + optimizer: + _target_: torch.optim.Adam + lr: 0.01 + + lr_scheduler: + scheduler: + _target_: torch.optim.lr_scheduler.ReduceLROnPlateau + factor: 0.6 + patience: 5 + threshold: 0.2 + min_lr: 1e-6 + monitor: ${monitored_metric} + interval: epoch + frequency: 1 + + model: + _target_: onescience.models.nequip.model.NequIPGNNModel + seed: 456 + model_dtype: float32 + type_names: ${model_type_names} + r_max: ${cutoff_radius} + num_bessels: 4 + bessel_trainable: false + polynomial_cutoff_p: 6 + num_layers: ${num_layers} + l_max: ${l_max} + parity: false + num_features: ${num_features} + radial_mlp_depth: 1 + radial_mlp_width: 16 + avg_num_neighbors: ${training_data_stats:num_neighbors_mean} + per_type_energy_scales: ${training_data_stats:per_type_forces_rms} + per_type_energy_shifts: ${training_data_stats:per_atom_energy_mean} + per_type_energy_scales_trainable: false + per_type_energy_shifts_trainable: false + +# Slurm / launch metadata used by demo/run.sh +name: nequip_smoke +launch: + mode: local + num_nodes: 1 + num_gpus: 1 +slurm: + partition: hx1hdnormal01 + nodelist: a01r1n02 + time: "00:10:00" + cpus_per_task: 8 +env: + OMP_NUM_THREADS: 1 diff --git a/demo/configs/tutorial_smoke_8dcu.yaml b/demo/configs/tutorial_smoke_8dcu.yaml new file mode 100644 index 0000000000000000000000000000000000000000..25e5cf0492656071d5c0a1a4f4315ff50367cf99 --- /dev/null +++ b/demo/configs/tutorial_smoke_8dcu.yaml @@ -0,0 +1,133 @@ +# DDP smoke test for the vendored NequIP trainer on one node with eight DCUs. +# Batch size is per rank, so batch_size=1 gives a global batch size of 8. +run: [train, test] + +cutoff_radius: 4.0 +num_layers: 2 +l_max: 1 +num_features: 8 +model_type_names: [Cu] +chemical_species: ${model_type_names} +monitored_metric: val0_epoch/weighted_sum + +data: + _target_: onescience.datapipes.materials.nequip.datamodule.ASEDataModule + seed: 456 + split_dataset: + file_path: ${demo_dir:reference_data/smoke.xyz} + train: 0.75 + val: 0.125 + test: 0.125 + transforms: + - _target_: onescience.datapipes.materials.nequip.transforms.ChemicalSpeciesToAtomTypeMapper + model_type_names: ${model_type_names} + - _target_: onescience.datapipes.materials.nequip.transforms.NeighborListTransform + r_max: ${cutoff_radius} + train_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 1 + num_workers: 0 + shuffle: true + val_dataloader: + _target_: torch.utils.data.DataLoader + batch_size: 1 + num_workers: 0 + test_dataloader: ${data.val_dataloader} + stats_manager: + _target_: onescience.datapipes.materials.nequip.CommonDataStatisticsManager + dataloader_kwargs: + batch_size: 2 + type_names: ${model_type_names} + +trainer: + _target_: lightning.Trainer + accelerator: gpu + devices: 8 + num_nodes: 1 + strategy: + _target_: lightning.pytorch.strategies.DDPStrategy + enable_checkpointing: true + max_epochs: 1 + limit_train_batches: 1 + limit_val_batches: 1 + limit_test_batches: 1 + num_sanity_val_steps: 0 + log_every_n_steps: 1 + logger: + _target_: lightning.pytorch.loggers.CSVLogger + save_dir: ${hydra:runtime.output_dir} + name: metrics + version: 0 + enable_progress_bar: true + callbacks: + - _target_: onescience.utils.nequip.train.callbacks.PlainTextMetricsLogger + - _target_: lightning.pytorch.callbacks.ModelCheckpoint + monitor: ${monitored_metric} + dirpath: ${hydra:runtime.output_dir}/checkpoints + filename: best + save_last: true + +training_module: + _target_: onescience.utils.nequip.train.EMALightningModule + ema_decay: 0.999 + loss: + _target_: onescience.utils.nequip.train.EnergyForceLoss + per_atom_energy: true + coeffs: + total_energy: 1.0 + forces: 1.0 + val_metrics: + _target_: onescience.utils.nequip.train.EnergyForceMetrics + coeffs: + total_energy_mae: 1.0 + forces_mae: 1.0 + train_metrics: ${training_module.val_metrics} + test_metrics: ${training_module.val_metrics} + optimizer: + _target_: torch.optim.Adam + lr: 0.01 + lr_scheduler: + scheduler: + _target_: torch.optim.lr_scheduler.ReduceLROnPlateau + factor: 0.6 + patience: 5 + threshold: 0.2 + min_lr: 1.0e-6 + monitor: ${monitored_metric} + interval: epoch + frequency: 1 + model: + _target_: onescience.models.nequip.model.NequIPGNNModel + compile_mode: eager + seed: 456 + model_dtype: float32 + type_names: ${model_type_names} + r_max: ${cutoff_radius} + num_bessels: 4 + bessel_trainable: false + polynomial_cutoff_p: 6 + num_layers: ${num_layers} + l_max: ${l_max} + parity: false + num_features: ${num_features} + radial_mlp_depth: 1 + radial_mlp_width: 16 + avg_num_neighbors: ${training_data_stats:num_neighbors_mean} + per_type_energy_scales: ${training_data_stats:per_type_forces_rms} + per_type_energy_shifts: ${training_data_stats:per_atom_energy_mean} + per_type_energy_scales_trainable: false + per_type_energy_shifts_trainable: false + +name: nequip_smoke_8dcu +launch: + mode: auto + num_nodes: 1 + num_gpus: 8 +slurm: + partition: hx1hdnormal01 + nodelist: "" + time: "00:10:00" + cpus_per_task: 8 +env: + OMP_NUM_THREADS: 1 + NCCL_DEBUG: INFO diff --git a/demo/download_tutorial_data.py b/demo/download_tutorial_data.py new file mode 100644 index 0000000000000000000000000000000000000000..ecb72dab9e55c2d0dc175eda56b781a788b830c8 --- /dev/null +++ b/demo/download_tutorial_data.py @@ -0,0 +1,64 @@ +"""Download and verify the official NequIP fcu.xyz tutorial dataset.""" + +from __future__ import annotations + +import argparse +import hashlib +import os +import tempfile +import urllib.request +from pathlib import Path + + +URL = "https://archive.materialscloud.org/records/ycbvx-knj69/files/fcu.xyz?download=1" +SHA256 = "57f00395d6945a3018a873d229fd7fbb7352a44a66f00f3c6e8a36247e0851e5" + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for block in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def default_output() -> Path: + datasets_dir = os.environ.get("ONESCIENCE_DATASETS_DIR") + if not datasets_dir: + raise RuntimeError("ONESCIENCE_DATASETS_DIR is not set; load matchem_env.sh first") + return Path(datasets_dir) / "matchem" / "NequIP" / "fcu.xyz" + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output", type=Path, help="Destination for fcu.xyz") + args = parser.parse_args() + output = (args.output or default_output()).expanduser().resolve() + output.parent.mkdir(parents=True, exist_ok=True) + + if output.is_file() and sha256(output) == SHA256: + print(f"Using verified dataset: {output}") + return + + request = urllib.request.Request(URL, headers={"User-Agent": "OneScience-NequIP/0.19"}) + temporary_path: Path | None = None + try: + with tempfile.NamedTemporaryFile( + prefix="fcu_", suffix=".xyz.part", dir=output.parent, delete=False + ) as temporary: + temporary_path = Path(temporary.name) + with urllib.request.urlopen(request) as response: + while block := response.read(1024 * 1024): + temporary.write(block) + actual = sha256(temporary_path) + if actual != SHA256: + raise RuntimeError(f"fcu.xyz SHA256 mismatch: expected {SHA256}, got {actual}") + temporary_path.replace(output) + finally: + if temporary_path is not None and temporary_path.exists(): + temporary_path.unlink() + print(f"Downloaded verified dataset: {output}") + + +if __name__ == "__main__": + main() diff --git a/demo/prepare_smoke_data.py b/demo/prepare_smoke_data.py new file mode 100644 index 0000000000000000000000000000000000000000..0b74c8aab10617f308bf77f2920f19445aaf991f --- /dev/null +++ b/demo/prepare_smoke_data.py @@ -0,0 +1,35 @@ +"""Generate a tiny extxyz smoke dataset for NequIP demo training.""" + +from __future__ import annotations + +import os +from pathlib import Path + +import numpy as np +from ase import Atoms +from ase.build import bulk +from ase.io import write + + +def main() -> None: + out_dir = Path(__file__).parent / "reference_data" + out_dir.mkdir(parents=True, exist_ok=True) + out_file = out_dir / "smoke.xyz" + + rng = np.random.default_rng(123) + structures = [] + # A few small Cu clusters with random displacements. + base = bulk("Cu", "fcc", a=3.6) * (2, 2, 2) + for i in range(8): + atoms = base.copy() + atoms.positions += rng.normal(scale=0.05, size=atoms.positions.shape) + atoms.info["energy"] = float(-len(atoms) * 3.5 + rng.normal(scale=0.5)) + atoms.arrays["forces"] = rng.normal(scale=0.1, size=atoms.positions.shape) + structures.append(atoms) + + write(out_file, structures, format="extxyz") + print(f"Wrote {len(structures)} structures to {out_file}") + + +if __name__ == "__main__": + main() diff --git a/demo/run.sh b/demo/run.sh new file mode 100644 index 0000000000000000000000000000000000000000..2779acca3e0a5428bb765703c7c016c2c869f9c3 --- /dev/null +++ b/demo/run.sh @@ -0,0 +1,162 @@ +#!/bin/bash +# Run NequIP training locally or submit it to Slurm from one YAML file. +set -euo pipefail + +DEMO_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +NEQUIP_DIR="$(cd "$DEMO_DIR/.." && pwd)" +PARSER="$DEMO_DIR/_parse_config.py" +CONFIG="" +SUBMIT=false + +while [[ $# -gt 0 ]]; do + case "$1" in + --config) CONFIG="$2"; shift 2 ;; + --config=*) CONFIG="${1#*=}"; shift ;; + --submit) SUBMIT=true; shift ;; + -h|--help) + echo "Usage: bash demo/run.sh --config configs/.yaml [--submit]" + echo "launch.mode: auto uses matching resources or submits when needed." + echo "launch.mode: local runs directly; submit always submits to Slurm." + exit 0 + ;; + *) echo "Unknown argument: $1" >&2; exit 2 ;; + esac +done + +[[ -n "$CONFIG" ]] || { echo "Please specify --config configs/.yaml" >&2; exit 2; } +[[ "$CONFIG" = /* ]] || CONFIG="$DEMO_DIR/$CONFIG" +[[ -f "$CONFIG" ]] || { echo "Config not found: $CONFIG" >&2; exit 2; } + +if [[ -z "${CONDA_PREFIX:-}" ]]; then + echo "Activate a OneScience MatChem conda environment before running this script." >&2 + exit 2 +fi +if [[ -z "${ONESCIENCE_MODELS_DIR:-}" || -z "${ONESCIENCE_DATASETS_DIR:-}" ]]; then + echo "Set ONESCIENCE_MODELS_DIR and ONESCIENCE_DATASETS_DIR before running this script." >&2 + exit 2 +fi +export MATCHEM_CONDA_NAME="${MATCHEM_CONDA_NAME:-$(basename "$CONDA_PREFIX")}" + +NAME="$(python3 "$PARSER" "$CONFIG" name)" +eval "$(python3 "$PARSER" "$CONFIG" launch)" +eval "$(python3 "$PARSER" "$CONFIG" slurm)" +ENV_EXPORTS="$(python3 "$PARSER" "$CONFIG" env)" +if [[ "$RUN_MODE" == "submit" ]]; then + SUBMIT=true +fi + +if [[ "$RUN_MODE" == "auto" ]] && ! $SUBMIT; then + IN_SLURM_ALLOCATION=false + AVAILABLE_NODES=1 + if [[ -n "${SLURM_JOB_ID:-}" ]]; then + IN_SLURM_ALLOCATION=true + AVAILABLE_NODES="${SLURM_NNODES:-${SLURM_JOB_NUM_NODES:-1}}" + if ! [[ "$AVAILABLE_NODES" =~ ^[1-9][0-9]*$ ]]; then + echo "Cannot determine allocated nodes from Slurm: $AVAILABLE_NODES" >&2 + exit 2 + fi + fi + + AVAILABLE_GPUS="$( + python3 -c 'import torch; print(torch.cuda.device_count() if torch.cuda.is_available() else 0)' \ + 2>/dev/null || true + )" + if ! [[ "$AVAILABLE_GPUS" =~ ^[0-9]+$ ]]; then + AVAILABLE_GPUS=0 + fi + + RESOURCE_MISMATCH="" + if (( AVAILABLE_NODES < NODES )); then + RESOURCE_MISMATCH="the config requests $NODES nodes but only $AVAILABLE_NODES are available" + elif (( AVAILABLE_GPUS < GPUS_PER_NODE )); then + RESOURCE_MISMATCH="the config requests $GPUS_PER_NODE DCUs per node but only $AVAILABLE_GPUS are visible" + fi + + if [[ -n "$RESOURCE_MISMATCH" ]]; then + if ! command -v sbatch >/dev/null 2>&1; then + echo "Current resources are insufficient: $RESOURCE_MISMATCH, and sbatch is unavailable." >&2 + exit 2 + fi + if $IN_SLURM_ALLOCATION; then + echo "Current Slurm allocation is insufficient: $RESOURCE_MISMATCH. Submitting a new Slurm job." + else + echo "Current resources are insufficient: $RESOURCE_MISMATCH. Submitting to Slurm." + fi + SUBMIT=true + else + echo "Current resources satisfy the config: nodes=$NODES, DCUs/node=$GPUS_PER_NODE." + fi +fi + +TIMESTAMP="$(date +%Y%m%d_%H%M%S)" +OUTPUT_ROOT="${ONESCIENCE_NEQUIP_OUTPUT_ROOT:-$NEQUIP_DIR/outputs}" +OUTPUT_DIR="$OUTPUT_ROOT/${NAME}_${TIMESTAMP}" +mkdir -p "$OUTPUT_DIR/checkpoints" +cp "$CONFIG" "$OUTPUT_DIR/source_config.yaml" +python3 "$PARSER" "$CONFIG" training-config > "$OUTPUT_DIR/config.yaml" + +if $SUBMIT; then + SLURM_SCRIPT="$OUTPUT_DIR/submit.sh" + cat > "$SLURM_SCRIPT" <> "$SLURM_SCRIPT" + fi + cat >> "$SLURM_SCRIPT" < 1 )); then + unset CUDA_VISIBLE_DEVICES HIP_VISIBLE_DEVICES ROCR_VISIBLE_DEVICES +fi +$ENV_EXPORTS +cd "$OUTPUT_DIR" +EOF + if (( WORLD_SIZE > 1 )); then + cat >> "$SLURM_SCRIPT" <> "$SLURM_SCRIPT" + fi + chmod u+x "$SLURM_SCRIPT" + echo "Submitting NequIP job: $SLURM_SCRIPT" + sbatch "$SLURM_SCRIPT" + exit 0 +fi + +eval "$ENV_EXPORTS" +cd "$OUTPUT_DIR" +if (( WORLD_SIZE > 1 )); then + if (( NODES > 1 )); then + if [[ "$RUN_MODE" == "auto" && -n "${SLURM_JOB_ID:-}" ]]; then + exec srun --kill-on-bad-exit=1 \ + --nodes="$NODES" \ + --ntasks="$WORLD_SIZE" \ + --ntasks-per-node="$GPUS_PER_NODE" \ + python "$NEQUIP_DIR/train.py" "hydra.run.dir=$OUTPUT_DIR" + fi + echo "Multi-node NequIP training must be launched through Slurm (--submit)." >&2 + exit 2 + fi + exec torchrun --standalone --nproc_per_node="$GPUS_PER_NODE" \ + "$NEQUIP_DIR/train.py" "hydra.run.dir=$OUTPUT_DIR" +fi +exec python "$NEQUIP_DIR/train.py" "hydra.run.dir=$OUTPUT_DIR" diff --git a/energy_volume.py b/energy_volume.py new file mode 100644 index 0000000000000000000000000000000000000000..ea6a1393fdf14bfff3f6b8312aa0230fb1605196 --- /dev/null +++ b/energy_volume.py @@ -0,0 +1,129 @@ +"""Compute the ASE energy-volume curve from the official NequIP example.""" + +from __future__ import annotations + +import argparse +import json +import os +from pathlib import Path + +import numpy as np +import torch +from ase.build import bulk + +from onescience.utils.nequip.integrations.ase import NequIPCalculator + + +def default_compiled_model() -> str | None: + models_dir = os.environ.get("ONESCIENCE_MODELS_DIR") + if not models_dir: + return None + return str(Path(models_dir) / "NequIP" / "NequIP-OAM-L-0.1.nequip.pth") + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--compiled-model", default=default_compiled_model()) + parser.add_argument("--device", default="cuda") + parser.add_argument("--element", default="Si") + parser.add_argument("--crystal-structure", default="diamond") + parser.add_argument("--lattice-constant", type=float, default=5.43) + parser.add_argument("--supercell", type=int, default=3) + parser.add_argument("--scale-min", type=float, default=0.95) + parser.add_argument("--scale-max", type=float, default=1.05) + parser.add_argument("--num-points", type=int, default=10) + parser.add_argument("--output", default="outputs/energy_volume.json") + parser.add_argument("--plot", default="outputs/energy_volume.png") + args = parser.parse_args() + + if not args.compiled_model: + parser.error("--compiled-model is required when ONESCIENCE_MODELS_DIR is unset") + compiled_model = Path(args.compiled_model).expanduser().resolve() + if not compiled_model.is_file(): + parser.error(f"compiled model not found: {compiled_model}") + if args.num_points < 2: + parser.error("--num-points must be at least 2") + if args.supercell < 1: + parser.error("--supercell must be positive") + + calculator = NequIPCalculator.from_compiled_model( + compile_path=str(compiled_model), + chemical_species_to_atom_type_map={args.element: args.element}, + device=args.device, + ) + + points = [] + for scale in np.linspace(args.scale_min, args.scale_max, args.num_points): + atoms = bulk( + args.element, + crystalstructure=args.crystal_structure, + a=args.lattice_constant * float(scale), + cubic=True, + ) + atoms *= (args.supercell,) * 3 + atoms.calc = calculator + energy = float(atoms.get_potential_energy()) + forces = atoms.get_forces() + points.append( + { + "scale": float(scale), + "volume_angstrom3": float(atoms.get_volume()), + "energy_ev": energy, + "energy_ev_per_atom": energy / len(atoms), + "max_force_ev_per_angstrom": float( + np.linalg.norm(forces, axis=1).max() + ), + } + ) + + energies = np.asarray([point["energy_ev"] for point in points]) + volumes = np.asarray([point["volume_angstrom3"] for point in points]) + minimum_index = int(np.argmin(energies)) + result = { + "compiled_model": str(compiled_model), + "device": args.device, + "device_name": torch.cuda.get_device_name(0) + if args.device.startswith("cuda") and torch.cuda.is_available() + else "cpu", + "element": args.element, + "crystal_structure": args.crystal_structure, + "base_lattice_constant_angstrom": args.lattice_constant, + "supercell": [args.supercell] * 3, + "num_atoms": len(atoms), + "points": points, + "sampled_minimum": points[minimum_index], + } + + output_path = Path(args.output).expanduser().resolve() + output_path.parent.mkdir(parents=True, exist_ok=True) + output_path.write_text(json.dumps(result, indent=2) + "\n", encoding="utf-8") + + if args.plot: + import matplotlib + + matplotlib.use("Agg") + import matplotlib.pyplot as plt + + plot_path = Path(args.plot).expanduser().resolve() + plot_path.parent.mkdir(parents=True, exist_ok=True) + plt.figure(figsize=(8, 6)) + plt.plot(volumes, energies, marker="o", label="E-V Curve") + plt.xlabel("Volume (Angstrom^3)", fontsize=14) + plt.ylabel("Energy (eV)", fontsize=14) + plt.title(f"Energy-Volume Curve for Cubic {args.element}", fontsize=16) + plt.legend(fontsize=12) + plt.grid() + plt.tight_layout() + plt.savefig(plot_path, dpi=160) + plt.close() + result["plot"] = str(plot_path) + output_path.write_text(json.dumps(result, indent=2) + "\n", encoding="utf-8") + + print(f"points: {len(points)}") + print(f"atoms per point: {result['num_atoms']}") + print(f"sampled minimum: {result['sampled_minimum']}") + print(f"result: {output_path}") + + +if __name__ == "__main__": + main() diff --git a/model/__init__.py b/model/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..d861da73337a56bdb28c5266bc720ab22a05e1ea --- /dev/null +++ b/model/__init__.py @@ -0,0 +1,41 @@ +from ._version import __version__ # noqa: F401 + +import packaging.version + +import torch + +# Load all installed nequip extension packages +# This allows installed extensions to register themselves in +# the nequip infrastructure with calls like `register_fields` + +# see https://packaging.python.org/en/guides/creating-and-discovering-plugins/#using-package-metadata +# we use "try ... except ..." to avoid importing sys.version_info +try: + # python >= 3.10 + from importlib.metadata import entry_points + + _DISCOVERED_NEQUIP_EXTENSION = entry_points(group="nequip.extension") +except (ImportError, TypeError): + # python < 3.10 + from importlib_metadata import entry_points + + _DISCOVERED_NEQUIP_EXTENSION = entry_points(group="nequip.extension") + +from onescience.utils.nequip.internal.resolvers import _register_default_resolvers +from onescience.utils.nequip.internal.versions.version_utils import get_version_safe + + +# torch version checks +torch_version = packaging.version.parse(get_version_safe(torch.__name__).split("+")[0]) + +# only allow 2.2.* or higher, required for `lightning` and `torchmetrics` compatibility +assert torch_version >= packaging.version.parse("2.2"), ( + f"NequIP supports 2.2.* or later, but {torch_version} found" +) + +for ep in _DISCOVERED_NEQUIP_EXTENSION: + if ep.name == "init_always": + ep.load() + +# register OmegaConf resolvers +_register_default_resolvers() diff --git a/model/_version.py b/model/_version.py new file mode 100644 index 0000000000000000000000000000000000000000..4b862b379f8143421f295944db52b786ec63a9e4 --- /dev/null +++ b/model/_version.py @@ -0,0 +1,5 @@ +# nequip package version file +# See Python packaging guide +# https://packaging.python.org/guides/single-sourcing-package-version/ + +__version__ = "0.19.0" diff --git a/model/model/__init__.py b/model/model/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..ae3530276331b3436558bddb0a8633712197cba4 --- /dev/null +++ b/model/model/__init__.py @@ -0,0 +1,29 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +from .utils import model_builder, override_model_compile_mode +from .modify_utils import modify +from .saved_models import ( + ModelFromCheckpoint, + ModelFromPackage, + ModelTypeNamesFromPackage, +) +from .nequip_models import ( + NequIPGNNModel, + PresetNequIPGNNModel, + FullNequIPGNNModel, +) +from .pair_potential import ZBLPairPotential +from .param_groups import MuonParamGroups + +__all__ = [ + "model_builder", + "override_model_compile_mode", + "modify", + "ModelFromCheckpoint", + "ModelFromPackage", + "ModelTypeNamesFromPackage", + "NequIPGNNModel", + "PresetNequIPGNNModel", + "FullNequIPGNNModel", + "ZBLPairPotential", + "MuonParamGroups", +] diff --git a/model/model/energy_modules.py b/model/model/energy_modules.py new file mode 100644 index 0000000000000000000000000000000000000000..1e46740fa7132b4254dc627a825e93f1e423cc59 --- /dev/null +++ b/model/model/energy_modules.py @@ -0,0 +1,35 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +from typing import Dict, Optional, Sequence + +from hydra.utils import instantiate + +from onescience.datapipes.materials.nequip import AtomicDataDict +from onescience.models.nequip.nn import AtomwiseReduce, SequentialGraphNetwork + + +def _append_energy_modules( + model: SequentialGraphNetwork, + type_names: Sequence[str], + pair_potential: Optional[Dict] = None, +): + # === pair potentials === + prev_irreps_out = model.irreps_out + if pair_potential is not None: + pair_potential = instantiate( + pair_potential, + type_names=type_names, + irreps_in=prev_irreps_out, + ) + prev_irreps_out = pair_potential.irreps_out + model.append("pair_potential", pair_potential) + + # === sum to total energy === + # perform sum after applying `pair_potential` + total_energy_sum = AtomwiseReduce( + irreps_in=prev_irreps_out, + reduce="sum", + field=AtomicDataDict.PER_ATOM_ENERGY_KEY, + out_field=AtomicDataDict.TOTAL_ENERGY_KEY, + ) + model.append("total_energy_sum", total_energy_sum) + return model diff --git a/model/model/inference_models/__init__.py b/model/model/inference_models/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..e72464c9e55205c6aa4ecc79eb489f6d08951154 --- /dev/null +++ b/model/model/inference_models/__init__.py @@ -0,0 +1,7 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. + +from .torchscript import load_torchscript_model +from .aotinductor import load_aotinductor_model +from .compiled import load_compiled_model + +__all__ = ["load_torchscript_model", "load_aotinductor_model", "load_compiled_model"] diff --git a/model/model/inference_models/aotinductor.py b/model/model/inference_models/aotinductor.py new file mode 100644 index 0000000000000000000000000000000000000000..19f69a34edb3b74414bb8b46eefe34283340f685 --- /dev/null +++ b/model/model/inference_models/aotinductor.py @@ -0,0 +1,128 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch +from typing import Union, Tuple, List, Optional + +from onescience.models.nequip.nn import graph_model +from onescience.utils.nequip.internal.versions import check_pt2_compile_compatibility +from onescience.utils.nequip.internal.aoti_metadata import ( + NEQUIP_AOTI_INPUTS_KEY, + NEQUIP_AOTI_OUTPUTS_KEY, + parse_aoti_keys, + import_custom_ops_libs, +) +from onescience.models.nequip.nn.compile import DictInputOutputWrapper + + +def _resolve_aot_keys( + provided_keys: Optional[List[str]], + metadata: dict, + metadata_key: str, + kind: str, + compile_path: str, +) -> List[str]: + """ + As of NequIP v0.17.0, we include the input and output fields in the AOTI artefact's metadata. + Previously, we always pass it from outside, but that's brittle. + In principle, now we can always use the metadata from the AOTI artefact to inform the input and output fields, + but there might be failure modes where a `--target batch` model is used for the ASE intergation or a `--target ase` model is used for the torchsim integration. + So it's safer for those integrations to also provide the input and output keys for what they expect. + This function will check for their consistency for safety. + """ + # for backwards compatibility since previous AOTI models don't store the key + metadata_entry = metadata.get(metadata_key, None) + + if provided_keys is None: + if metadata_entry is None: + raise ValueError( + f"{kind}_keys are required for `{compile_path}` because this AOTI artifact does not store `{kind}` keys metadata. " + "Please pass them explicitly or recompile with a newer `nequip-compile`." + ) + else: + return parse_aoti_keys(metadata_entry) + else: + provided_keys = list(provided_keys) + if metadata_entry is None: + return provided_keys + else: + # both probided, so we check their consistency + metadata_keys = parse_aoti_keys(metadata_entry) + if provided_keys != metadata_keys: + raise ValueError( + f"Provided {kind} keys do not match metadata for `{compile_path}`.\n" + f"provided={provided_keys}\n" + f"metadata={metadata_keys}" + ) + return provided_keys + + +def load_aotinductor_model( + compile_path: str, + device: Union[str, torch.device], + input_keys: Optional[List[str]] = None, + output_keys: Optional[List[str]] = None, +) -> Tuple[torch.nn.Module, dict]: + """Load an AOTInductor model from a .nequip.pt2 file. + + Args: + compile_path: path to compiled model file ending with .nequip.pt2 + device: the device to use + input_keys: optional list of expected input field names for DictInputOutputWrapper + output_keys: optional list of expected output field names for DictInputOutputWrapper + + Returns: + tuple of (wrapped_model, processed_metadata) + """ + # sanity checks + check_pt2_compile_compatibility() + + # import any required custom ops libraries before the C++ loader runs + import_custom_ops_libs(compile_path) + + # load compiled model + compiled_model = torch._inductor.aoti_load_package(compile_path) + + # get and process metadata + metadata = compiled_model.get_metadata() + + input_keys = _resolve_aot_keys( + provided_keys=input_keys, + metadata=metadata, + metadata_key=NEQUIP_AOTI_INPUTS_KEY, + kind="input", + compile_path=compile_path, + ) + output_keys = _resolve_aot_keys( + provided_keys=output_keys, + metadata=metadata, + metadata_key=NEQUIP_AOTI_OUTPUTS_KEY, + kind="output", + compile_path=compile_path, + ) + + model = DictInputOutputWrapper(compiled_model, input_keys, output_keys) + + # check device compatibility + compile_device = metadata["AOTI_DEVICE_KEY"] + if torch.device(compile_device) != torch.device(device): + raise RuntimeError( + f"`{compile_path}` was compiled for `{compile_device}` and won't work with device={device}, use device={compile_device} instead." + ) + + # process standard metadata + metadata[graph_model.R_MAX_KEY] = float(metadata[graph_model.R_MAX_KEY]) + metadata[graph_model.TYPE_NAMES_KEY] = metadata[graph_model.TYPE_NAMES_KEY].split( + " " + ) + + # process per-edge-type cutoffs if present + if graph_model.PER_EDGE_TYPE_CUTOFF_KEY in metadata: + from onescience.models.nequip.nn.embedding.utils import cutoff_str_to_fulldict + + cutoff_str = metadata[graph_model.PER_EDGE_TYPE_CUTOFF_KEY] + metadata[graph_model.PER_EDGE_TYPE_CUTOFF_KEY] = cutoff_str_to_fulldict( + cutoff_str, metadata[graph_model.TYPE_NAMES_KEY] + ) + else: + metadata[graph_model.PER_EDGE_TYPE_CUTOFF_KEY] = None + + return model, metadata diff --git a/model/model/inference_models/compiled.py b/model/model/inference_models/compiled.py new file mode 100644 index 0000000000000000000000000000000000000000..731e6deecc913aeee2e0a4baa4d5898fdf8bfed4 --- /dev/null +++ b/model/model/inference_models/compiled.py @@ -0,0 +1,60 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README +# at the root for information on using it. +import torch + +from pathlib import Path +from typing import Union, Tuple, List, Optional + +from .torchscript import load_torchscript_model +from .aotinductor import load_aotinductor_model +from onescience.utils.nequip.internal.global_state import TF32_KEY, set_global_state + + +def load_compiled_model( + compile_path: str, + device: Union[str, torch.device], + input_keys: Optional[List[str]] = None, + output_keys: Optional[List[str]] = None, +) -> Tuple[torch.nn.Module, dict]: + """Load a compiled model from either TorchScript or AOTInductor format. + + This function can load compiled models created with ``nequip-compile``: + + - **TorchScript models** (``.nequip.pth``): legacy compiled format + - **AOT Inductor models** (``.nequip.pt2``): modern compiled format with better performance + + Args: + compile_path: path to compiled model file (``.nequip.pth`` or ``.nequip.pt2``) + device: the device to use + input_keys: optional input field names for AOTInductor models (for ``.nequip.pt2``) + output_keys: optional output field names for AOTInductor models (for ``.nequip.pt2``) + + Returns: + tuple: ``(model, metadata)`` with model prepared for inference + """ + compile_fname = Path(compile_path).name + + if compile_fname.endswith(".nequip.pth"): + model, metadata = load_torchscript_model(compile_path, device) + elif compile_fname.endswith(".nequip.pt2"): + model, metadata = load_aotinductor_model( + compile_path, device, input_keys, output_keys + ) + else: + raise ValueError( + f"Unknown file type: {compile_fname} " + f"(expected `*.nequip.pth` or `*.nequip.pt2`)" + ) + + # set global state from metadata + set_global_state( + **{ + TF32_KEY: bool(int(metadata[TF32_KEY])), + } + ) + + # prepare model for inference + model = model.to(device) + model.eval() + + return model, metadata diff --git a/model/model/inference_models/torchscript.py b/model/model/inference_models/torchscript.py new file mode 100644 index 0000000000000000000000000000000000000000..48182c8fa1cc6dcf6e7f0c7e8ab90b3176c777ea --- /dev/null +++ b/model/model/inference_models/torchscript.py @@ -0,0 +1,73 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README +# at the root for information on using it. +import torch +from e3nn.util.jit import script + +from onescience.models.nequip.nn import graph_model +from onescience.utils.nequip.internal.compile import prepare_model_for_compile +from onescience.utils.nequip.internal.global_state import TF32_KEY + +from typing import Union, Tuple + + +def load_torchscript_model( + compile_path: str, + device: Union[str, torch.device] = "cpu", +) -> Tuple[torch.nn.Module, dict]: + """Load a torchscript model from a .nequip.pth file. + + Args: + compile_path (str): path to compiled model file ending with .nequip.pth + device (Union[str, torch.device]): the device to use + """ + # load model with metadata + metadata = { + graph_model.R_MAX_KEY: None, + graph_model.TYPE_NAMES_KEY: None, + graph_model.PER_EDGE_TYPE_CUTOFF_KEY: None, + TF32_KEY: None, + } + model = torch.jit.load(compile_path, _extra_files=metadata, map_location=device) + model = torch.jit.freeze(model) + + # process metadata + metadata[graph_model.R_MAX_KEY] = float(metadata[graph_model.R_MAX_KEY]) + metadata[graph_model.TYPE_NAMES_KEY] = ( + metadata[graph_model.TYPE_NAMES_KEY].decode("utf-8").split(" ") + ) + + # process per-edge-type cutoffs if present + if metadata[graph_model.PER_EDGE_TYPE_CUTOFF_KEY] is not None: + from onescience.models.nequip.nn.embedding.utils import cutoff_str_to_fulldict + + cutoff_str = metadata[graph_model.PER_EDGE_TYPE_CUTOFF_KEY].decode("utf-8") + metadata[graph_model.PER_EDGE_TYPE_CUTOFF_KEY] = cutoff_str_to_fulldict( + cutoff_str, metadata[graph_model.TYPE_NAMES_KEY] + ) + + return model, metadata + + +def save_torchscript_model( + model: torch.nn.Module, + metadata: dict, + output_path: str, + device: Union[str, torch.device], +) -> None: + """Save a model as a torchscript .nequip.pth file. + + Args: + model: model to save + metadata: metadata dictionary to save with the model + output_path: path to save the compiled model + device: device to prepare model on + """ + # encode metadata for torchscript + encoded_metadata = {k: str(v).encode("ascii") for k, v in metadata.items()} + + # prepare and script model + model = prepare_model_for_compile(model, device) + script_model = script(model) + + # save with metadata + torch.jit.save(script_model, output_path, _extra_files=encoded_metadata) diff --git a/model/model/modify_utils.py b/model/model/modify_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..0429f798ec9c0cf3163d1199e324ca259d2fd40f --- /dev/null +++ b/model/model/modify_utils.py @@ -0,0 +1,131 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch + +from onescience.models.nequip.nn.model_modifier_utils import ( + is_model_modifier, + is_persistent_model_modifier, +) + +import inspect +import contextvars +import contextlib +from hydra.utils import get_method +from typing import Dict, List, Union, Any, Optional + +_ONLY_APPLY_PERSISTENT = contextvars.ContextVar("_ONLY_APPLY_PERSISTENT", default=False) + + +@contextlib.contextmanager +def only_apply_persistent_modifiers(persistent_only: bool): + """ + Used during `nequip-package` to only apply persistent modifiers. + """ + global _ONLY_APPLY_PERSISTENT + init_state = _ONLY_APPLY_PERSISTENT.get() + assert not init_state, ( + "this error implies that the `only_apply_persistent_modifiers` context manager is being nested, which is unexpected behavior" + ) + _ONLY_APPLY_PERSISTENT.set(persistent_only) + try: + yield + finally: + _ONLY_APPLY_PERSISTENT.set(init_state) + + +def get_all_modifiers( + module: torch.nn.Module, _all_modifiers: Optional[Dict[str, callable]] = None +) -> Dict[str, callable]: + """ + Find all model modifiers available in a model. + + Args: + module (torch.nn.Module): The model to collect modifiers from. + + Returns: + Dict[str, callable]: A dictionary mapping modifier names to their functions. + """ + if _all_modifiers is None: + _all_modifiers = {} + + for name, member in inspect.getmembers(module, predicate=inspect.ismethod): + if is_model_modifier(member): + if name in _all_modifiers: + # confirm (indirectly) that these are @classmethods (bound instance methods will not be equal) + # this ensures that having a globally unique name for each modifier does not hide differences between different copies of the same modifier hiding in a single module tree + assert _all_modifiers[name] == member, ( + f"Found at least two non-unique modifiers with same name `{name}`: {_all_modifiers[name]!r} and {member!r}" + ) + _all_modifiers[name] = member + + for _, child in module.named_children(): + get_all_modifiers(child, _all_modifiers=_all_modifiers) + + return _all_modifiers + + +def modify( + model: Union[Dict[str, torch.nn.Module], torch.nn.Module], + modifiers: Union[List[Dict[str, Any]], Dict[str, List[Dict[str, Any]]]], +) -> Union[Dict[str, torch.nn.Module], torch.nn.Module]: + """Applies a sequence of model modifier functions to a model. + + The modifiers will be applied in the specified order. Whether the order of modifiers matters depends on the specific modifiers used. + + Args: + model (Union[Dict[str, torch.nn.Module], torch.nn.Module]): The model(s) to modify. + modifiers (Union[List[Dict[str, Any]], Dict[str, List[Dict[str, Any]]]]): A list of modifier configurations (if ``model`` is a single model) or a dictionary mapping model names to lists of modifier configurations (if ``model`` is a dictionary). + Each modifier configuration is a dictionary. The dictionary must contain a key "modifier" that specifies the name of the modifier function to apply as a string. All other keys in the dictionary are passed as keyword arguments to the modifier function. + + Returns: + Union[Dict[str, torch.nn.Module], torch.nn.Module]: The modified model(s). + """ + # check persistence + global _ONLY_APPLY_PERSISTENT + persistent_only: bool = _ONLY_APPLY_PERSISTENT.get() + + # build inner model if not already built + if not isinstance(model, torch.nn.Module): + # don't use `hydra.utils.instantiate` because it may lead to a hydra dependency during packaging + model = model.copy() + model_fn = get_method(model.pop("_target_")) + model = model_fn(**model) + + def _apply_modifier( + avail_modifiers: Dict[str, callable], + modifier_cfg: Dict[str, Any], + this_model: torch.nn.Module, + ) -> None: + modifier_cfg = modifier_cfg.copy() + modifier_name = modifier_cfg.pop("modifier") + if modifier_name not in avail_modifiers.keys(): + avail_names = list(avail_modifiers.keys()) + raise RuntimeError( + f"`{modifier_name}` is not a registered model modifier. The following are registered model modifiers: {avail_names}" + ) + modifier_fn = avail_modifiers[modifier_name] + is_persistent = is_persistent_model_modifier(modifier_fn) + # only skip if doing `persistent_only` and modifier is non-persistent, otherwise always apply + if not (persistent_only and not is_persistent): + this_model = modifier_fn(this_model, **modifier_cfg) + + if isinstance(model, torch.nn.ModuleDict): + # because `model` is actually a `ModuleDict`, we make the modifiers flexible while keeping a simple default for the more common single-model use case + # a single list of modifiers is given, we assume it'll be uniformly applied to everything + if isinstance(modifiers, list): + modifiers = {model_name: modifiers.copy() for model_name in model.keys()} + # ^ the above allows us to use a common loop over individual sub-models and apply the relevant model-specific modifiers + + for model_name, submodel in model.items(): + avail_modifiers: Dict[str, callable] = get_all_modifiers(submodel) + for modifier in modifiers[model_name]: + _apply_modifier(avail_modifiers, modifier, submodel) + + elif isinstance(model, torch.nn.Module): + assert isinstance(modifiers, list) + avail_modifiers: Dict[str, callable] = get_all_modifiers(model) + for modifier in modifiers: + _apply_modifier(avail_modifiers, modifier, model) + else: + raise RuntimeError("Unrecognized model object found.") + + return model diff --git a/model/model/nequip_models.py b/model/model/nequip_models.py new file mode 100644 index 0000000000000000000000000000000000000000..7c88aa8c4094e005fd183092728bff7addc83031 --- /dev/null +++ b/model/model/nequip_models.py @@ -0,0 +1,399 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import math +from e3nn import o3 + +from onescience.datapipes.materials.nequip import AtomicDataDict + +from onescience.models.nequip.nn import ( + GraphModel, + SequentialGraphNetwork, + ScalarMLP, + PerTypeScaleShift, + ConvNetLayer, + ForceStressOutput, + ApplyFactor, +) +from onescience.models.nequip.nn.embedding import ( + NodeTypeEmbed, + PolynomialCutoff, + EdgeLengthNormalizer, + BesselEdgeLengthEncoding, + SphericalHarmonicEdgeAttrs, +) + +from .utils import model_builder +from .energy_modules import _append_energy_modules +import warnings +from typing import Sequence, Optional, List, Dict, Union, Callable + + +_NEQUIP_GNN_PRESETS = { + "S": { + "num_layers": 2, + "l_max": 1, + "num_features": [128, 64], + }, + "M": { + "num_layers": 4, + "l_max": 2, + "num_features": [128, 64, 32], + }, + "L": { + "num_layers": 6, + "l_max": 3, + "num_features": [128, 64, 32, 32], + }, + "XL": { + "num_layers": 6, + "l_max": 4, + "num_features": [320, 96, 64, 32, 32], + }, +} + +_NEQUIP_GNN_STANDARD_PRESET = { + "parity": False, + "type_embed_num_features": 32, + "radial_mlp_depth": 1, + "radial_mlp_width": 128, +} + + +def _format_nequip_gnn_preset_docstring() -> str: + shared_defaults = "\n".join( + [ + f" - ``{key}``: ``{value!r}``" + for key, value in _NEQUIP_GNN_STANDARD_PRESET.items() + ] + ) + preset_defaults = "\n".join( + [ + f" - ``{preset}``: ``{defaults!r}``" + for preset, defaults in _NEQUIP_GNN_PRESETS.items() + ] + ) + return f"""Build :func:`NequIPGNNModel` from a named architecture preset. + + This is a wrapper of :func:`NequIPGNNModel` that injects preset hyperparameters based on model sizes of the NequIP foundation potentials. + All arguments are the same as :func:`NequIPGNNModel`, except this builder also requires ``preset`` and applies preset defaults before ``**kwargs``. + For full argument documentation, see :func:`NequIPGNNModel`. + Users can override the preset defaults by providing arguments for the fields to be overriden. + + Preset argument: + preset (str): one of {", ".join([f"``{name}``" for name in _NEQUIP_GNN_PRESETS.keys()])} + + Override order: + 1. shared defaults + 2. per-preset defaults + 3. explicit ``**kwargs`` (highest priority) + + Shared defaults: +{shared_defaults} + + Per-preset defaults: +{preset_defaults} + """ + + +@model_builder +def PresetNequIPGNNModel( + preset: str, + **kwargs, +) -> GraphModel: + preset = preset.upper() + assert preset in _NEQUIP_GNN_PRESETS, ( + f"`preset` must be one of {list(_NEQUIP_GNN_PRESETS.keys())}, but found `{preset}`" + ) + model_kwargs = { + **_NEQUIP_GNN_STANDARD_PRESET, + **_NEQUIP_GNN_PRESETS[preset], + } + # explicit kwargs override standard and preset defaults + model_kwargs.update(kwargs) + + return NequIPGNNModel(**model_kwargs) + + +@model_builder +def NequIPGNNModel( + num_layers: int = 4, + l_max: int = 1, + parity: bool = True, + num_features: Union[int, List[int]] = 32, + type_embed_num_features: Optional[int] = None, + radial_mlp_depth: int = 1, + radial_mlp_width: int = 128, + **kwargs, +) -> GraphModel: + """NequIP GNN model that can predict energies only or energies with forces/stresses. + + Args: + seed (int): seed for reproducibility + model_dtype (str): ``float32`` or ``float64`` + r_max (float): cutoff radius + per_edge_type_cutoff (Dict): one can optionally specify cutoffs for each edge type [must be smaller than ``r_max``] (default ``None``) + type_names (Sequence[str]): list of atom type names + num_layers (int): number of interaction blocks, we find 3-5 to work best (default ``4``) + l_max (int): the maximum rotation order for the network's features, ``1`` is a good default, ``2`` is more accurate but slower (default ``1``) + parity (bool): whether to include features with odd mirror parity -- often turning parity off gives equally good results but faster networks, so it's worth testing (default ``True``) + num_features (int/List[int]): multiplicity of the features, smaller is faster (default ``32``); it is also possible to provide the multiplicity for each irrep, e.g. for ``l_max=2`` and ``parity=False``, ``num_features=[5, 2, 7]`` refers to ``5x0e``, ``2x1o`` and ``7x2e`` features + type_embed_num_features (int): number of features for the type embedding layer; if not provided, defaults to ``num_features[0]`` (default ``None``) + radial_mlp_depth (int): number of radial layers, usually 1-3 works best, smaller is faster (default ``1``) + radial_mlp_width (int): number of hidden neurons in radial function, smaller is faster (default ``128``) + readout_mlp_hidden_layers_depth (int): number of hidden layers in the readout MLP (default ``0``) + readout_mlp_hidden_layers_width (int): width of hidden layers in the readout MLP (default 0e contribution of ``num_features``) + readout_mlp_nonlinearity (str): ``silu``, ``mish``, ``gelu``, or ``None`` (default ``silu``) + num_bessels (int): number of Bessel basis functions (default ``8``) + bessel_trainable (bool): whether the Bessel roots are trainable (default ``False``) + polynomial_cutoff_p (int): p-exponent used in polynomial cutoff function, smaller p corresponds to stronger decay with distance (default ``6``) + avg_num_neighbors (float/Dict[str, float]): used to normalize edge sums for better numerics (default ``None``) + per_type_energy_scales (float/List[float]): per-atom energy scales, which could be derived from the force RMS of the data (default ``None``) + per_type_energy_shifts (float/List[float]): per-atom energy shifts, which should generally be isolated atom reference energies or estimated from average per-atom energies of the data (default ``None``) + per_type_energy_scales_trainable (bool): whether the per-atom energy scales are trainable (default ``False``) + per_type_energy_shifts_trainable (bool): whether the per-atom energy shifts are trainable (default ``False``) + pair_potential (torch.nn.Module): additional pair potential term, e.g. :class:`~nequip.nn.pair_potential.ZBL` (default ``None``) + do_derivatives (bool): whether to compute forces and stresses via autograd (default ``True``) + """ + # === sanity checks and warnings === + assert num_layers > 0, ( + f"at least one convnet layer required, but found `num_layers={num_layers}`" + ) + + # === spherical harmonics === + irreps_edge_sh = repr(o3.Irreps.spherical_harmonics(lmax=l_max)) + + # === handle `num_features` === + if isinstance(num_features, int): + num_features = [num_features] * (l_max + 1) + assert len(num_features) == l_max + 1, ( + f"`num_features` should be of length `l_max + 1` ({l_max + 1}), but found `num_features={num_features}` with {len(num_features)} entries." + ) + + # === type embedding === + type_embed_num_features = ( + type_embed_num_features + if type_embed_num_features is not None + else num_features[0] + ) + + # === convnet === + # convert a single set of parameters uniformly for every layer + feature_irreps_hidden = repr( + o3.Irreps( + [ + (num_features[l], (l, p)) + for l in range(l_max + 1) + for p in ( + (1, -1) if parity else ((1,) if l % 2 == 0 else (-1,)) + ) # p = 1 for even l, -1 for odd l, with parity = False + ] + ) + ) + feature_irreps_hidden_list = [feature_irreps_hidden] * (num_layers - 1) + radial_mlp_depth_list = [radial_mlp_depth] * num_layers + radial_mlp_width_list = [radial_mlp_width] * num_layers + + # === post convnets === + feature_irreps_hidden_list += [repr(o3.Irreps([(num_features[0], (0, 1))]))] + + # === build model === + model = FullNequIPGNNModel( + irreps_edge_sh=irreps_edge_sh, + type_embed_num_features=type_embed_num_features, + feature_irreps_hidden=feature_irreps_hidden_list, + radial_mlp_depth=radial_mlp_depth_list, + radial_mlp_width=radial_mlp_width_list, + **kwargs, + ) + return model + + +PresetNequIPGNNModel.__doc__ = _format_nequip_gnn_preset_docstring() + + +@model_builder +def FullNequIPGNNModel( + r_max: float, + type_names: Sequence[str], + # convnet params + radial_mlp_depth: Sequence[int], + radial_mlp_width: Sequence[int], + feature_irreps_hidden: Sequence[Union[str, o3.Irreps]], + # irreps and dims + irreps_edge_sh: Union[int, str, o3.Irreps], + type_embed_num_features: int, + categorical_graph_field_embed: Optional[List[Dict[str, int]]] = None, + # readout + readout_mlp_hidden_layers_depth: int = 0, + readout_mlp_hidden_layers_width: Optional[int] = None, + readout_mlp_nonlinearity: Optional[str] = "silu", + # edge length encoding + per_edge_type_cutoff: Optional[Dict[str, Union[float, Dict[str, float]]]] = None, + num_bessels: int = 8, + bessel_trainable: bool = False, + polynomial_cutoff_p: int = 6, + # edge sum normalization + avg_num_neighbors: Optional[Union[float, Dict[str, float]]] = None, + # per atom energy params + per_type_energy_scales: Optional[Union[float, Sequence[float]]] = None, + per_type_energy_shifts: Optional[Union[float, Sequence[float]]] = None, + per_type_energy_scales_trainable: Optional[bool] = False, + per_type_energy_shifts_trainable: Optional[bool] = False, + pair_potential: Optional[Dict] = None, + # derivatives + do_derivatives: bool = True, + # developmental params + convnet_sc: bool = True, + learnable_shift: bool = False, + # == things that generally shouldn't be changed == + # convnet + convnet_resnet: bool = False, + convnet_nonlinearity_type: str = "gate", + convnet_nonlinearity_scalars: Dict[int, Callable] = {"e": "silu", "o": "tanh"}, + convnet_nonlinearity_gates: Dict[int, Callable] = {"e": "silu", "o": "tanh"}, +) -> GraphModel: + """NequIP GNN model that predicts energies based on a more extensive set of arguments.""" + # === sanity checks and warnings === + assert all(tn.isalnum() for tn in type_names), ( + "`type_names` must contain only alphanumeric characters" + ) + + # learnable_shift requires skip connections to be enabled + assert not learnable_shift or (convnet_sc or convnet_resnet), ( + "`learnable_shift=True` requires at least one of `convnet_sc` or `convnet_resnet` to be True" + ) + + # require every convnet layer to be specified explicitly in a list + # infer num_layers from the list size + assert ( + len(radial_mlp_depth) == len(radial_mlp_width) == len(feature_irreps_hidden) + ), ( + f"radial_mlp_depth: {radial_mlp_depth}, radial_mlp_width: {radial_mlp_width}, feature_irreps_hidden: {feature_irreps_hidden} should all have the same length" + ) + num_layers = len(radial_mlp_depth) + + # assert that last convnet produces only scalars + assert all([l == 0 for l in o3.Irreps(feature_irreps_hidden[-1]).ls]), ( + f"last convnet layer output must only contain scalars but found {feature_irreps_hidden[-1]}" + ) + + if per_type_energy_scales is None: + warnings.warn( + "Found `per_type_energy_scales=None` -- it is recommended to set `per_type_energy_scales` for better numerics during training." + ) + if per_type_energy_shifts is None: + warnings.warn( + "Found `per_type_energy_shifts=None` -- it is HIGHLY recommended to set `per_type_energy_shifts` as it determines the per-atom energies approaching the isolated atom regime." + ) + + # === encode and embed features === + # == node scalar embedding == + # NOTE: node embed is done first in case we need to pass in categorical graph fields as inputs + # see how `irreps_in` is registered in the `NodeTypeEmbed` class + type_embed = NodeTypeEmbed( + type_names=type_names, + num_features=type_embed_num_features, + categorical_graph_field_embed=categorical_graph_field_embed, + ) + + # == edge tensor embedding == + spharm = SphericalHarmonicEdgeAttrs( + irreps_edge_sh=irreps_edge_sh, + irreps_in=type_embed.irreps_out, + ) + # == edge scalar embedding == + edge_norm = EdgeLengthNormalizer( + r_max=r_max, + type_names=type_names, + per_edge_type_cutoff=per_edge_type_cutoff, + irreps_in=spharm.irreps_out, + ) + bessel_encode = BesselEdgeLengthEncoding( + num_bessels=num_bessels, + trainable=bessel_trainable, + cutoff=PolynomialCutoff(polynomial_cutoff_p), + edge_invariant_field=AtomicDataDict.EDGE_EMBEDDING_KEY, + irreps_in=edge_norm.irreps_out, + ) + # for backwards compatibility of NequIP's bessel encoding + factor = ApplyFactor( + in_field=AtomicDataDict.EDGE_EMBEDDING_KEY, + factor=(2 * math.pi) / (r_max * r_max), + irreps_in=bessel_encode.irreps_out, + ) + + modules = { + "type_embed": type_embed, + "spharm": spharm, + "edge_norm": edge_norm, + "bessel_encode": bessel_encode, + "factor": factor, + } + prev_irreps_out = factor.irreps_out + + # === convnet layers === + for layer_i in range(num_layers): + current_convnet = ConvNetLayer( + irreps_in=prev_irreps_out, + feature_irreps_hidden=feature_irreps_hidden[layer_i], + convolution_kwargs={ + "radial_mlp_depth": radial_mlp_depth[layer_i], + "radial_mlp_width": radial_mlp_width[layer_i], + # to ensure isolated atom limit + "use_sc": convnet_sc + if learnable_shift + else (layer_i != 0) and convnet_sc, + "is_first_layer": layer_i == 0, + # normalization parameters + "avg_num_neighbors": avg_num_neighbors, + "type_names": type_names, + }, + resnet=convnet_resnet + if learnable_shift + else (layer_i != 0) and convnet_resnet, + nonlinearity_type=convnet_nonlinearity_type, + nonlinearity_scalars=convnet_nonlinearity_scalars, + nonlinearity_gates=convnet_nonlinearity_gates, + ) + prev_irreps_out = current_convnet.irreps_out + modules.update({f"layer{layer_i}_convnet": current_convnet}) + + # === readout === + if readout_mlp_hidden_layers_width is None: + readout_mlp_hidden_layers_width = o3.Irreps(feature_irreps_hidden[-1]).dim + per_atom_energy_readout = ScalarMLP( + output_dim=1, + hidden_layers_depth=readout_mlp_hidden_layers_depth, + hidden_layers_width=readout_mlp_hidden_layers_width, + nonlinearity=readout_mlp_nonlinearity, + bias=False, + forward_weight_init=True, + field=AtomicDataDict.NODE_FEATURES_KEY, + out_field=AtomicDataDict.PER_ATOM_ENERGY_KEY, + irreps_in=prev_irreps_out, + ) + + per_type_energy_scale_shift = PerTypeScaleShift( + type_names=type_names, + field=AtomicDataDict.PER_ATOM_ENERGY_KEY, + out_field=AtomicDataDict.PER_ATOM_ENERGY_KEY, + scales=per_type_energy_scales, + shifts=per_type_energy_shifts, + scales_trainable=per_type_energy_scales_trainable, + shifts_trainable=per_type_energy_shifts_trainable, + irreps_in=per_atom_energy_readout.irreps_out, + ) + + modules.update( + { + "per_atom_energy_readout": per_atom_energy_readout, + "per_type_energy_scale_shift": per_type_energy_scale_shift, + } + ) + + energy_model = SequentialGraphNetwork(modules) + energy_model = _append_energy_modules( + model=energy_model, + type_names=type_names, + pair_potential=pair_potential, + ) + return ForceStressOutput(energy_model, do_derivatives) diff --git a/model/model/pair_potential.py b/model/model/pair_potential.py new file mode 100644 index 0000000000000000000000000000000000000000..9331cad879db1f20bc063815b55c2edb9fd9d0b2 --- /dev/null +++ b/model/model/pair_potential.py @@ -0,0 +1,50 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +from onescience.models.nequip.nn import SequentialGraphNetwork, AtomwiseReduce, ForceStressOutput +from onescience.models.nequip.nn.embedding import EdgeLengthNormalizer +from onescience.datapipes.materials.nequip import AtomicDataDict +from onescience.models.nequip.nn.pair_potential import ZBL +from .utils import model_builder + +from typing import Optional, Dict, Union, Sequence + + +@model_builder +def ZBLPairPotential( + r_max: float, + type_names: Sequence[str], + chemical_species: Sequence[str], + units: str, + polynomial_cutoff_p: int = 6, + per_edge_type_cutoff: Optional[Dict[str, Union[float, Dict[str, float]]]] = None, +): + """ + Model builder for a force field containing only a ZBL pair potential term, mainly for internal testing purposes. + """ + edge_norm = EdgeLengthNormalizer( + r_max=r_max, + type_names=type_names, + per_edge_type_cutoff=per_edge_type_cutoff, + ) + zbl_module = ZBL( + type_names=type_names, + chemical_species=chemical_species, + units=units, + polynomial_cutoff_p=polynomial_cutoff_p, + irreps_in=edge_norm.irreps_out, + ) + energy_sum = AtomwiseReduce( + reduce="sum", + field=AtomicDataDict.PER_ATOM_ENERGY_KEY, + out_field=AtomicDataDict.TOTAL_ENERGY_KEY, + irreps_in=zbl_module.irreps_out, + ) + energy_model = SequentialGraphNetwork( + { + "edge_norm": edge_norm, + "pair_potential": zbl_module, + "total_energy_sum": energy_sum, + } + ) + model = ForceStressOutput(func=energy_model) + + return model diff --git a/model/model/param_groups.py b/model/model/param_groups.py new file mode 100644 index 0000000000000000000000000000000000000000..263c8ad4bfdd76f971ce2268909e2e4471306127 --- /dev/null +++ b/model/model/param_groups.py @@ -0,0 +1,97 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch + + +def _normalize_weight_index_slices(weight_index_slices): + normalized = [] + for entry in weight_index_slices: + index_slice = getattr(entry, "slice_1D", None) + shape_2d = getattr(entry, "shape_2D", None) + if index_slice is None or shape_2d is None: + index_slice, shape_2d = entry + if isinstance(index_slice, slice): + index_slice = (index_slice.start, index_slice.stop, index_slice.step) + else: + index_slice = tuple(index_slice) + assert len(index_slice) == 3 + shape_2d = tuple(shape_2d) + assert len(shape_2d) == 2 + normalized.append((index_slice, shape_2d)) + return normalized + + +def MuonParamGroups( + model: torch.nn.Module, + muon: dict, + adam: dict, +): + """ + Build optimizer parameter groups, splitting parameters between a Muon-based optimizer + and Adam (or Adam-like) optimizer. + + Assigned to Adam group: + - Any parameter whose name does **not** contain the substring ``"layer"``. + - Any parameter not matching the Muon-specific rules below. + + Assigned to Muon group: + - Edge MLP weights: parameters whose name contains ``"edge_mlp"`` and that are + 2D tensors (i.e., matrix weights). + - e3nn convolution linear weights: parameters whose name contains ``"conv.linear"``. + + For e3nn ``Linear`` layers, the returned Muon parameter group includes an + ``e3nn_reshaping`` dictionary mapping the index of the parameter within the Muon + group to the module's ``weight_index_slices``. This metadata will be used by the + to reshape or operate on corresponding matrix weights. + + Args: + model (torch.nn.Module): The model to optimize. + muon (dict): Muon config parameters. + adam (dict): Adam config parameters. + + """ + muon_weights = [] + adam_weights = [] + + e3nn_reshaping = {} + + modules = dict(model.named_modules()) + + for name, param in model.named_parameters(): + # Assumes all input and output layers are + # not called layers. + if "layer" not in name: + adam_weights.append(param) + continue + + # First, all edge_mlps should be muon + if "edge_mlp" in name and param.ndim == 2: + muon_weights.append(param) + continue + + if "conv.linear" in name: + # e3nn conv layers. + + # Find the e3nn Linear module this represents + module_name, _, _ = name.rpartition(".") + module = modules[module_name] + + # use Muon only when reshape metadata is available + weight_index_slices = getattr(module, "weight_index_slices", None) + if weight_index_slices is None: + adam_weights.append(param) + continue + + # store plain tuples to keep optimizer state picklable + index = len(muon_weights) + e3nn_reshaping[index] = _normalize_weight_index_slices(weight_index_slices) + muon_weights.append(param) + continue + + adam_weights.append(param) + + param_groups = [ + dict(params=muon_weights, use_muon=True, e3nn_reshaping=e3nn_reshaping, **muon), + dict(params=adam_weights, use_muon=False, **adam), + ] + + return param_groups diff --git a/model/model/saved_models/__init__.py b/model/model/saved_models/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..97930c9250e4dac5accb9f41e6fb0fb2abec85fb --- /dev/null +++ b/model/model/saved_models/__init__.py @@ -0,0 +1,12 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. + +from .checkpoint import ModelFromCheckpoint +from .package import ModelFromPackage, ModelTypeNamesFromPackage +from .load_utils import load_saved_model + +__all__ = [ + "ModelFromCheckpoint", + "ModelFromPackage", + "ModelTypeNamesFromPackage", + "load_saved_model", +] diff --git a/model/model/saved_models/_utils.py b/model/model/saved_models/_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..0a5bb39d96bf26b6cc91d09532b14530fb515615 --- /dev/null +++ b/model/model/saved_models/_utils.py @@ -0,0 +1,33 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +""" +Shared utilities for loading models from saved formats (checkpoints and packages). +""" + +import os +from typing import List + +from onescience.models.nequip.model.utils import _COMPILE_MODE_OPTIONS + + +def _check_compile_mode(compile_mode: str, client: str, exclude_keys: List[str] = []): + """Helper function for checking input arguments.""" + allowed_options = [ + mode for mode in _COMPILE_MODE_OPTIONS if mode not in exclude_keys + ] + assert compile_mode in allowed_options, ( + f"`compile_mode={compile_mode}` is not recognized for `{client}`, only the following are supported: {allowed_options}" + ) + + +def _check_file_exists(file_path: str, file_type: str): + """Check if a checkpoint or package file exists.""" + if not os.path.isfile(file_path): + assert file_type in ("checkpoint", "package") + client = ( + "`ModelFromCheckpoint`" + if file_type == "checkpoint" + else "`ModelFromPackage`" + ) + raise RuntimeError( + f"{file_type} file provided at `{file_path}` is not found. NOTE: Any process that loads a checkpoint produced from training runs based on {client} will look for the original {file_type} file at the location specified during training. It is also recommended to use full paths (instead or relative paths) to avoid potential errors." + ) diff --git a/model/model/saved_models/checkpoint.py b/model/model/saved_models/checkpoint.py new file mode 100644 index 0000000000000000000000000000000000000000..871efc4d7b7b8ad1e45461ba6692ab1c669338a0 --- /dev/null +++ b/model/model/saved_models/checkpoint.py @@ -0,0 +1,148 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +""" +Functions for loading models from checkpoint files. +""" + +import torch +import hydra +import warnings + +from onescience.models.nequip.model.utils import ( + override_model_compile_mode, + _EAGER_MODEL_KEY, +) +from onescience.datapipes.materials.nequip import AtomicDataDict +from onescience.datapipes.materials.nequip.transforms import NonPeriodicCellTransform + + +from onescience.utils.nequip.internal.global_dtype import _GLOBAL_DTYPE +from onescience.utils.nequip.internal.logger import RankedLogger + +from ._utils import _check_compile_mode, _check_file_exists + +# === setup logging === +logger = RankedLogger(__name__, rank_zero_only=True) + + +def ModelFromCheckpoint(checkpoint_path: str, compile_mode: str = _EAGER_MODEL_KEY): + """Builds model from a NequIP framework checkpoint file. + + This function can be used in the config file as follows. + + .. code-block:: yaml + + model: + _target_: onescience.models.nequip.model.ModelFromCheckpoint + checkpoint_path: path/to/ckpt + compile_mode: eager/compile + + .. warning:: + DO NOT CHANGE the directory structure or location of the checkpoint file if this model loader is used for training. Any process that loads a checkpoint produced from training runs originating from a package file will look for the original package file at the location specified during training. It is also recommended to use full paths (instead or relative paths) to avoid potential errors. + + Args: + checkpoint_path (str): path to a ``nequip`` framework checkpoint file + compile_mode (str): ``eager`` or ``compile`` allowed for training + """ + # === sanity checks === + _check_file_exists(file_path=checkpoint_path, file_type="checkpoint") + _check_compile_mode(compile_mode, "ModelFromCheckpoint") + logger.info(f"Loading model from checkpoint file: {checkpoint_path} ...") + + # === load checkpoint and extract info === + checkpoint = torch.load( + checkpoint_path, + map_location="cpu", + weights_only=False, + ) + + # === versions === + ckpt_versions = checkpoint["hyper_parameters"]["info_dict"]["versions"] + from onescience.utils.nequip.internal import get_current_code_versions + + session_versions = get_current_code_versions(verbose=False) + + for code, session_version in session_versions.items(): + if code in ckpt_versions: + ckpt_version = ckpt_versions[code] + # sanity check that versions for current build matches versions from ckpt + if ckpt_version != session_version: + warnings.warn( + f"`{code}` versions differ between the checkpoint file ({ckpt_version}) and the current run ({session_version}) -- `ModelFromCheckpoint` will be built with the current run's versions, but please check that this decision is as intended." + ) + + # === load model via lightning module === + # Rewrite legacy upstream ``nequip.`` targets to the OneScience namespace. + from onescience.utils.nequip.internal.compat import rewrite_nequip_targets + + compatible_hyper_parameters = rewrite_nequip_targets( + checkpoint["hyper_parameters"] + ) + info_dict = compatible_hyper_parameters["info_dict"] + training_module = hydra.utils.get_class(info_dict["training_module"]["_target_"]) + # ensure that model is built with correct `compile_mode` + with override_model_compile_mode(compile_mode): + lightning_module = training_module.load_from_checkpoint( + checkpoint_path, + weights_only=False, + **compatible_hyper_parameters, + ) + + model = lightning_module.evaluation_model + return model + + +def data_dict_from_checkpoint(ckpt_path: str) -> AtomicDataDict.Type: + from onescience.utils.nequip.internal.dtype import torch_default_dtype + + with torch_default_dtype(_GLOBAL_DTYPE): + # === get data from checkpoint === + checkpoint = torch.load( + ckpt_path, + map_location="cpu", + weights_only=False, + ) + from onescience.utils.nequip.internal.compat import rewrite_nequip_targets + + data_config = rewrite_nequip_targets( + checkpoint["hyper_parameters"]["info_dict"]["data"].copy() + ) + if "train_dataloader" not in data_config: + data_config["train_dataloader"] = { + "_target_": "torch.utils.data.DataLoader" + } + data_config["train_dataloader"]["batch_size"] = 1 + datamodule = hydra.utils.instantiate(data_config, _recursive_=False) + # TODO: better way of doing this? + # instantiate the datamodule, dataset, and get train dataloader + try: + datamodule.prepare_data() + # instantiate train dataset + datamodule.setup(stage="fit") + dloader = datamodule.train_dataloader() + for data in dloader: + if AtomicDataDict.num_nodes(data) > 3: + break + finally: + datamodule.teardown(stage="fit") + + # === sanitize data === + if AtomicDataDict.CELL_KEY not in data: + # try to construct sensible cell for nonperiodic system + transform = NonPeriodicCellTransform(padding=10.0, override_cell=False) + data = transform(data) + + # if still no cell (transform was no-op), create a large cell + if AtomicDataDict.CELL_KEY not in data: + data[AtomicDataDict.CELL_KEY] = 1e5 * torch.eye( + 3, + dtype=_GLOBAL_DTYPE, + device=data[AtomicDataDict.POSITIONS_KEY].device, + ).unsqueeze(0) + + data[AtomicDataDict.EDGE_CELL_SHIFT_KEY] = torch.zeros( + (AtomicDataDict.num_edges(data), 3), + dtype=_GLOBAL_DTYPE, + device=data[AtomicDataDict.POSITIONS_KEY].device, + ) + + return data diff --git a/model/model/saved_models/load_utils.py b/model/model/saved_models/load_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..d3fc0c1b521ff0a8a099929c3947a859e29cfc38 --- /dev/null +++ b/model/model/saved_models/load_utils.py @@ -0,0 +1,150 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. + +import contextlib +import pathlib +import requests +from tqdm.auto import tqdm + +from onescience.models.nequip.model.utils import _EAGER_MODEL_KEY +from onescience.models.nequip.model.saved_models import ModelFromPackage, ModelFromCheckpoint +from onescience.models.nequip.model.modify_utils import only_apply_persistent_modifiers +from onescience.utils.nequip.train.lightning import _SOLE_MODEL_KEY +from onescience.utils.nequip.internal import model_repository +from onescience.utils.nequip.internal.logger import RankedLogger +from onescience.utils.nequip.internal.model_cache import get_cached_model, cache_model + +logger = RankedLogger(__name__, rank_zero_only=True) + + +@contextlib.contextmanager +def _get_model_file_path(input_path): + """Context manager that provides a file path for both local and nequip.net models. + + For local files: yields the input path directly + For nequip.net downloads: uses cache if available, otherwise downloads and caches + (default cache location: ``~/.nequip/model_cache``, configurable via ``NEQUIP_CACHE_DIR``) + + Args: + input_path: path to the model checkpoint or package file, or nequip.net model ID + (format: ``nequip.net:group-name/model-name:version``) + + Yields: + pathlib.Path: Path to the model file (either original or cached) + """ + is_nequip_net_download: bool = str(input_path).startswith("nequip.net:") + + if is_nequip_net_download: + # get model ID + model_id = str(input_path)[len("nequip.net:") :] + logger.info(f"Fetching {model_id} from onescience.models.nequip.net...") + + # get download URL + with model_repository.NequIPNetAPIClient() as client: + model_info = client.get_model_download_info(model_id) + + if model_info.newer_version_id is not None: + logger.info( + f"Model {model_id} has a newer version available: {model_info.newer_version_id}" + ) + + download_url = model_info.artifact.download_url + + # check cache first + cached_path = get_cached_model(model_id, download_url) + if cached_path is not None: + yield cached_path + return + + # cache miss: download and cache + def download_fn(target_path: pathlib.Path): + response = requests.get(download_url, stream=True) + response.raise_for_status() + + total_size = int(response.headers.get("content-length", 0)) + + with open(target_path, "wb") as f: + with tqdm( + total=total_size, + unit="B", + unit_scale=True, + desc=f"Downloading from {model_info.artifact.host_name}", + ) as pbar: + for chunk in response.iter_content(chunk_size=65536): + if chunk: + f.write(chunk) + pbar.update(len(chunk)) + + # download and cache (cache_model will skip caching if NEQUIP_NO_CACHE is set) + cached_path = cache_model(model_id, download_url, download_fn) + logger.info("Download complete, loading model...") + yield cached_path + else: + logger.info(f"Loading model from {input_path} ...") + yield pathlib.Path(input_path) + + +def load_saved_model( + input_path, + compile_mode: str = _EAGER_MODEL_KEY, + model_key: str = _SOLE_MODEL_KEY, + return_data_dict: bool = False, +): + """Load a saved model from checkpoint, package, or nequip.net. + + This function can load models from: + + - **Checkpoint files** (``.ckpt``): saved during training runs + - **Package files** (``.nequip.zip``): created with ``nequip-package`` + - **nequip.net models**: using model ID format ``nequip.net:group-name/model-name:version`` from `nequip.net `__ + + Args: + input_path: path to the model checkpoint or package file, or nequip.net model ID + (format: ``nequip.net:group-name/model-name:version``) + compile_mode (str): compile mode for the model (default: ``"eager"``) + model_key (str): key to select the model from ModuleDict (default: ``"sole_model"``) + return_data_dict (bool): if ``True``, also return the data dict for compilation (default: ``False``) + + Returns: + torch.nn.Module or tuple: the loaded model, or ``(model, data)`` tuple if ``return_data_dict=True`` + """ + + with _get_model_file_path(input_path) as actual_input_path: + # check if the resolved file exists + if not actual_input_path.exists(): + raise ValueError( + f"Model file does not exist: {input_path} (resolved to: {actual_input_path})" + ) + + # use package load path if extension matches, otherwise assume checkpoint file + use_ckpt = not str(actual_input_path).endswith(".nequip.zip") + + # load model + if use_ckpt: + # we only apply persistent modifiers when building from checkpoint + # i.e. acceleration modifiers won't be applied, and have to be specified during compile time + with only_apply_persistent_modifiers(persistent_only=True): + model = ModelFromCheckpoint( + actual_input_path, compile_mode=compile_mode + ) + else: + # packaged models will never have non-persistent modifiers built in + model = ModelFromPackage(actual_input_path, compile_mode=compile_mode) + + if model_key is not None: + model = model[model_key] + # ^ `ModuleDict` of `GraphModel` is loaded, we then select the desired `GraphModel` (`model_key` defaults to work for single model case) + # otherwise, return the `ModuleDict` + + # load data dict if requested + if return_data_dict: + from onescience.models.nequip.model.saved_models.checkpoint import data_dict_from_checkpoint + from onescience.models.nequip.model.saved_models.package import data_dict_from_package + + if use_ckpt: + data = data_dict_from_checkpoint(str(actual_input_path)) + else: + data = data_dict_from_package(str(actual_input_path)) + + return model, data + else: + return model diff --git a/model/model/saved_models/package.py b/model/model/saved_models/package.py new file mode 100644 index 0000000000000000000000000000000000000000..3b32d13242f9f8ee010455717cdefaab60fba531 --- /dev/null +++ b/model/model/saved_models/package.py @@ -0,0 +1,190 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +""" +Functions for loading models from package files. +""" + +import torch +import yaml +import warnings +import contextlib +import io +from typing import Dict, Any + +from onescience.datapipes.materials.nequip import AtomicDataDict +from onescience.models.nequip.model.utils import ( + get_current_compile_mode, + _EAGER_MODEL_KEY, +) +from onescience.utils.nequip.cli._workflow_utils import get_workflow_state +from onescience.utils.nequip.internal.logger import RankedLogger + +from ._utils import _check_compile_mode, _check_file_exists +from onescience.utils.nequip.internal.asserts import assert_package_extension + +# === setup logging === +logger = RankedLogger(__name__, rank_zero_only=True) + + +@contextlib.contextmanager +def _cpu_deserialize_if_no_cuda(): + """Force CUDA-saved storages inside packaged models to load on CPU when CUDA is unavailable.""" + if torch.cuda.is_available(): + yield + return + + orig = torch.storage._load_from_bytes + + def _load_from_bytes_cpu(b): + return torch.load(io.BytesIO(b), map_location="cpu", weights_only=False) + + torch.storage._load_from_bytes = _load_from_bytes_cpu + try: + yield + finally: + torch.storage._load_from_bytes = orig + + +# === package importer utilities === +# most of the complexity for `ModelFromPackage` is due to the need to keep track of the `Importer` if we ever repackage +# see `nequip/scripts/package.py` to get the full picture of how they interact +# we expect the following variable to only be used during `nequip-package` + +_PACKAGE_TIME_SHARED_IMPORTER = None + + +def _get_shared_importer(): + global _PACKAGE_TIME_SHARED_IMPORTER + return _PACKAGE_TIME_SHARED_IMPORTER + + +def _get_package_metadata(imp) -> Dict[str, Any]: + """Load packaged model metadata from an existing PackageImporter.""" + pkg_metadata: Dict[str, Any] = yaml.safe_load( + imp.load_text(package="model", resource="package_metadata.txt") + ) + assert int(pkg_metadata["package_version_id"]) > 0 + # ^ extra sanity check since saving metadata in txt files was implemented in packaging version 1 + + return pkg_metadata + + +# === warning management === + + +@contextlib.contextmanager +def _suppress_package_importer_exporter_warnings(): + # Ideally this ceases to exist or becomes a no-op in future versions of PyTorch + with warnings.catch_warnings(): + # suppress torch.package TypedStorage warning + warnings.filterwarnings( + "ignore", + message="TypedStorage is deprecated.*", + category=UserWarning, + module=r"torch\.package\.(package_exporter|package_importer)", + ) + yield + + +# === loading models from package files === + + +def ModelFromPackage(package_path: str, compile_mode: str = _EAGER_MODEL_KEY): + """Builds model from a NequIP framework packaged zip file constructed with ``nequip-package``. + + This function can be used in the config file as follows. + + .. code-block:: yaml + + model: + _target_: onescience.models.nequip.model.ModelFromPackage + package_path: path/to/pkg + compile_mode: eager/compile + + .. warning:: + DO NOT CHANGE the directory structure or location of the package file if this model loader is used for training. Any process that loads a checkpoint produced from training runs originating from a package file will look for the original package file at the location specified during training. It is also recommended to use full paths (instead or relative paths) to avoid potential errors. + + Args: + package_path (str): path to NequIP framework packaged model with the ``.nequip.zip`` extension (an error will be thrown if the file has a different extension) + compile_mode (str): ``eager`` or ``compile`` allowed for training + """ + # === sanity checks === + _check_file_exists(file_path=package_path, file_type="package") + assert_package_extension(package_path) + + # === account for checkpoint loading === + # if `ModelFromPackage` is used by itself, `override=False` and the input `compile_mode` argument is used + # if this function is called at the end of checkpoint loading via `ModelFromCheckpoint`, `override=True` and the overriden `compile_mode` takes precedence + cm, override = get_current_compile_mode(return_override=True) + compile_mode = cm if override else compile_mode + + # === sanity check compile modes === + workflow_state = get_workflow_state() + _check_compile_mode(compile_mode, "ModelFromPackage") + + # === load model === + logger.info(f"Loading model from package file: {package_path} ...") + with _suppress_package_importer_exporter_warnings(): + # during `nequip-package`, we need to use the same importer for all the models for successful repackaging + # see https://pytorch.org/docs/stable/package.html#re-export-an-imported-object + if workflow_state == "package": + global _PACKAGE_TIME_SHARED_IMPORTER + imp = _PACKAGE_TIME_SHARED_IMPORTER + # we load the importer from `package_path` for the first time + if imp is None: + imp = torch.package.PackageImporter(package_path) + _PACKAGE_TIME_SHARED_IMPORTER = imp + # if it's not `None`, it means we've previously loaded a model during `nequip-package` and should keep using the same importer + else: + # if not doing `nequip-package`, we just load a new importer every time `ModelFromPackage` is called + imp = torch.package.PackageImporter(package_path) + + # do sanity checking with available models + pkg_metadata = _get_package_metadata(imp) + available_models = pkg_metadata["available_models"] + # throw warning if desired `compile_mode` is not available, and default to eager + if compile_mode not in available_models: + warnings.warn( + f"Requested `{compile_mode}` model is not present in the package file ({package_path}). `nequip-{workflow_state}` task will default to using the `{_EAGER_MODEL_KEY}` model." + ) + compile_mode = _EAGER_MODEL_KEY + + with _cpu_deserialize_if_no_cuda(): + model = imp.load_pickle( + package="model", + resource=f"{compile_mode}_model.pkl", + map_location="cpu", + ) + + # NOTE: model returned is not a GraphModel object tied to the `nequip` in current Python env, but a GraphModel object from the packaged zip file + return model + + +def data_dict_from_package(package_path: str) -> AtomicDataDict.Type: + """Load example data from a .nequip.zip package file.""" + with _suppress_package_importer_exporter_warnings(): + imp = torch.package.PackageImporter(package_path) + with _cpu_deserialize_if_no_cuda(): + data = imp.load_pickle(package="model", resource="example_data.pkl") + return data + + +def ModelTypeNamesFromPackage(package_path: str): + """Extract model type names from a packaged model file. + + Useful for setting up type mappers when fine-tuning models or when you need to know what atom types a model was trained on. + + Args: + package_path (str): path to packaged model file + """ + from typing import List + + _check_file_exists(file_path=package_path, file_type="package") + + with _suppress_package_importer_exporter_warnings(): + imp = torch.package.PackageImporter(package_path) + pkg_metadata = _get_package_metadata(imp) + + atom_types_dict = pkg_metadata["atom_types"] + # convert dict {idx: name} to list [name, ...] + type_names: List[str] = [atom_types_dict[i] for i in range(len(atom_types_dict))] + return type_names diff --git a/model/model/utils.py b/model/model/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..551ab9004489b86a73f7d4d5c58e4864c5743b94 --- /dev/null +++ b/model/model/utils.py @@ -0,0 +1,230 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch +from lightning.pytorch.utilities.seed import isolate_rng + +from onescience.models.nequip.nn.graph_model import GraphModel +from onescience.models.nequip.nn.compile import CompileGraphModel +from onescience.utils.nequip.internal import ( + dtype_from_name, + torch_default_dtype, + conditional_torchscript_mode, +) +from onescience.utils.nequip.internal.global_state import ( + global_state_initialized, + get_latest_global_state, + TF32_KEY, +) + +import functools +import contextvars +import contextlib + +from typing import Optional, Final + +_IS_BUILDING_MODEL = contextvars.ContextVar("_IS_BUILDING_MODEL", default=False) +_CURRENT_MODEL_BUILDER_DEFAULTS = contextvars.ContextVar( + "_CURRENT_MODEL_BUILDER_DEFAULTS", + default=None, +) + +# the following is the set of model build types for specific purposes +_EAGER_MODEL_KEY = "eager" +_TRAIN_TIME_COMPILE_KEY: Final[str] = "compile" + +_COMPILE_MODE_OPTIONS = { + _EAGER_MODEL_KEY, + _TRAIN_TIME_COMPILE_KEY, +} + + +_OVERRIDE_COMPILE_MODE = contextvars.ContextVar("_OVERRIDE_COMPILE_MODE", default=False) +_CURRENT_COMPILE_MODE = contextvars.ContextVar( + "_CURRENT_COMPILE_MODE", default=_EAGER_MODEL_KEY +) + + +@contextlib.contextmanager +def override_model_compile_mode(compile_mode: Optional[str]): + """ + Overrides the ``compile_mode`` for model building. + If several of these context managers are nested, the outermost one will be prioritized while the inner ones are ignored. + The intended client is `ModelFromCheckpoint`. + Anybody using this function should be warned that the behavior is designed for loading models from checkpoints and packages correctly. + """ + assert compile_mode in _COMPILE_MODE_OPTIONS + global _OVERRIDE_COMPILE_MODE + global _CURRENT_COMPILE_MODE + init_state = _OVERRIDE_COMPILE_MODE.get() + # in the case of nested overrides, we prioritize the outermost context manager + if init_state: + yield + else: + init_mode = _CURRENT_COMPILE_MODE.get() + _OVERRIDE_COMPILE_MODE.set(True) + _CURRENT_COMPILE_MODE.set(compile_mode) + try: + yield + finally: + _OVERRIDE_COMPILE_MODE.set(init_state) + _CURRENT_COMPILE_MODE.set(init_mode) + + +@contextlib.contextmanager +def fresh_model_builder_context(): + """Temporarily treat nested model-builder calls as fresh top-level builds. + + This is an explicit escape hatch for composing models where an inner builder + should run with full `@model_builder` behavior (dtype/seed/wrapping), + instead of being returned as a raw nested module. + + Required builder args (`seed`, `model_dtype`, `type_names`) are inherited + from the active outer model-builder context when not explicitly provided. + """ + # TODO: decide compile_mode semantics for fresh nested builds: + # should they inherit outer builder compile_mode or use current default/override? + global _IS_BUILDING_MODEL + init_state = _IS_BUILDING_MODEL.get() + _IS_BUILDING_MODEL.set(False) + try: + yield + finally: + _IS_BUILDING_MODEL.set(init_state) + + +def get_current_compile_mode(return_override: bool = False): + # returns tuple of (whether compile mode is overriden, compile mode) + global _CURRENT_COMPILE_MODE + if return_override: + global _OVERRIDE_COMPILE_MODE + return _CURRENT_COMPILE_MODE.get(), _OVERRIDE_COMPILE_MODE.get() + else: + return _CURRENT_COMPILE_MODE.get() + + +def model_builder(func=None, *, wrapper_class=None, compile_wrapper_class=None): + """Decorator for model builder functions in the ``nequip`` ecosystem. + + Handles model building with proper seeding, floating point precision (``float32`` or ``float64``), and wraps the result with ``GraphModel``. Requires ``seed``, ``model_dtype``, and ``type_names`` arguments. + Supports ``eager`` and ``compile`` modes via ``compile_mode``. + + The ``seed``, ``model_dtype``, and ``compile_mode`` arguments are consumed by the decorator and not passed to the decorated function. + + Can be used in two ways: + - @model_builder (uses GraphModel wrapper, backward compatible) + - @model_builder(wrapper_class=CustomGraphModel) (uses custom wrapper) + + Args: + func: The function to decorate (when used without parentheses) + wrapper_class: Custom GraphModel subclass to use for wrapping (default: GraphModel) + compile_wrapper_class: Custom wrapper for compile mode (default: CompileGraphModel) + """ + + # default wrapper classes + if wrapper_class is None: + wrapper_class = GraphModel + if compile_wrapper_class is None: + compile_wrapper_class = CompileGraphModel + + def decorator(f): + @functools.wraps(f) + def wrapper(*args, **kwargs): + # to handle nested model building + global _IS_BUILDING_MODEL + + # to handle compile modes + global _OVERRIDE_COMPILE_MODE + global _CURRENT_COMPILE_MODE + + # this means we're in an inner model, so we shouldn't apply the model builder operations, and just pass the function + if _IS_BUILDING_MODEL.get(): + return f(*args, **kwargs) + + # this means we're in the outer model, and have to apply the model builder operations + _IS_BUILDING_MODEL.set(True) + prev_builder_defaults = _CURRENT_MODEL_BUILDER_DEFAULTS.get() + try: + default_builder_kwargs = _CURRENT_MODEL_BUILDER_DEFAULTS.get() + if default_builder_kwargs is not None: + for key in ("seed", "model_dtype", "type_names"): + if key not in kwargs and key in default_builder_kwargs: + kwargs[key] = default_builder_kwargs[key] + + model_cfg = kwargs.copy() + # === sanity checks === + assert global_state_initialized(), ( + "global state must be initialized before building models" + ) + assert all( + key in kwargs for key in ["seed", "model_dtype", "type_names"] + ), ( + "`seed`, `model_dtype`, and `type_names` are mandatory model arguments." + ) + + if get_latest_global_state().get(TF32_KEY, False): + assert kwargs["model_dtype"] == "float32", ( + "`allow_tf32=True` only works with `model_dtype=float32`" + ) + + # seed and model_dtype are removed from kwargs, so they will NOT get passed to inner models + seed = kwargs.pop("seed") + model_dtype = kwargs.pop("model_dtype") + dtype = dtype_from_name(model_dtype) + inherited_builder_defaults = { + "seed": seed, + "model_dtype": model_dtype, + "type_names": kwargs["type_names"], + } + _CURRENT_MODEL_BUILDER_DEFAULTS.set(inherited_builder_defaults) + + # === compilation options === + # `compile_mode` dictates the optimization path chosen + # users can set this with the `compile_mode` arg to the model builder + # devs can override it with `override_model_compile_mode` + + # always pop because inner models won't need `compile_mode` arg + compile_mode = kwargs.pop("compile_mode", _CURRENT_COMPILE_MODE.get()) + # compile mode overriding logic + if _OVERRIDE_COMPILE_MODE.get(): + compile_mode = _CURRENT_COMPILE_MODE.get() + assert compile_mode in _COMPILE_MODE_OPTIONS, ( + f"`compile_mode` can only be any of {_COMPILE_MODE_OPTIONS}, but `{compile_mode}` found" + ) + + # use custom wrapper class or default + if compile_mode == _TRAIN_TIME_COMPILE_KEY: + # === torch version check === + from onescience.utils.nequip.internal.versions import check_pt2_compile_compatibility + + check_pt2_compile_compatibility() + graph_model_module = compile_wrapper_class + else: + graph_model_module = wrapper_class + + # never script + with conditional_torchscript_mode(False): + # set dtype and seed + with torch_default_dtype(dtype): + with isolate_rng(): + torch.manual_seed(seed) + model = f(*args, **kwargs) + # wrap with GraphModel + graph_model = graph_model_module( + model=model, + model_config=model_cfg, + model_input_fields=model.irreps_in, + ) + return graph_model + finally: + _CURRENT_MODEL_BUILDER_DEFAULTS.set(prev_builder_defaults) + # reset to default in case of failure + _IS_BUILDING_MODEL.set(False) + + return wrapper + + # handle both @model_builder and @model_builder(...) + if func is None: + # called with arguments: @model_builder(wrapper_class=X) + return decorator + else: + # called without arguments: @model_builder + return decorator(func) diff --git a/model/nn/__init__.py b/model/nn/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..36bb5b61d10163097ebeb426c2823ccdf6a37007 --- /dev/null +++ b/model/nn/__init__.py @@ -0,0 +1,45 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +from ._graph_mixin import GraphModuleMixin, SequentialGraphNetwork +from .graph_model import GraphModel +from .atomwise import ( + AtomwiseOperation, + AtomwiseReduce, + AtomwiseLinear, + PerTypeScaleShift, +) +from .nonlinearities import ShiftedSoftplus +from .mlp import ScalarMLP, ScalarMLPFunction +from .interaction_block import InteractionBlock +from .convnetlayer import ConvNetLayer +from .grad_output import PartialForceOutput, ForceStressOutput +from .misc import Concat, ApplyFactor, SaveForOutput +from .utils import scatter, tp_path_exists, with_edge_vectors_, with_edge_type_ +from .model_modifier_utils import model_modifier, replace_submodules +from .norm import AvgNumNeighborsNorm + +__all__ = [ + "GraphModel", + "GraphModuleMixin", + "SequentialGraphNetwork", + "AtomwiseOperation", + "AtomwiseReduce", + "AtomwiseLinear", + "PerTypeScaleShift", + "ShiftedSoftplus", + "ScalarMLP", + "ScalarMLPFunction", + "InteractionBlock", + "PartialForceOutput", + "ForceStressOutput", + "ConvNetLayer", + "Concat", + "ApplyFactor", + "SaveForOutput", + "scatter", + "tp_path_exists", + "with_edge_vectors_", + "with_edge_type_", + "model_modifier", + "replace_submodules", + "AvgNumNeighborsNorm", +] diff --git a/model/nn/_ghost_exchange_base.py b/model/nn/_ghost_exchange_base.py new file mode 100644 index 0000000000000000000000000000000000000000..5b65b8fdf3979d82f2394aa8116e5b0ecdf9e318 --- /dev/null +++ b/model/nn/_ghost_exchange_base.py @@ -0,0 +1,57 @@ +import torch + +from onescience.datapipes.materials.nequip import AtomicDataDict +from ._graph_mixin import GraphModuleMixin +from .model_modifier_utils import replace_submodules, model_modifier + + +class GhostExchangeModule(GraphModuleMixin, torch.nn.Module): + """Base class for ghost atom exchange modules.""" + + def __init__( + self, + field: str = AtomicDataDict.NODE_FEATURES_KEY, + irreps_in={}, + ): + super().__init__() + self.field = field + + self._init_irreps( + irreps_in=irreps_in, + my_irreps_in={field: irreps_in[field]}, + irreps_out={field: irreps_in[field]}, + ) + + def forward( + self, + data: AtomicDataDict.Type, + ghost_included: bool, + ) -> AtomicDataDict.Type: + raise NotImplementedError("Subclasses must implement forward method") + + +class NoOpGhostExchangeModule(GhostExchangeModule): + """Base ghost exchange module that performs a no-op.""" + + def forward( + self, + data: AtomicDataDict.Type, + ghost_included: bool, + ) -> AtomicDataDict.Type: + return data + + @model_modifier(persistent=True, private=True) + @classmethod + def enable_LAMMPSMLIAPGhostExchange(cls, model): + """Enable LAMMPS ML-IAP ghost exchange for inference in LAMMPS ML-IAP.""" + + from ._ghost_exchange_lmp_mliap import LAMMPSMLIAPGhostExchangeModule + + def factory(old): + new = LAMMPSMLIAPGhostExchangeModule( + field=old.field, + irreps_in=old.irreps_in, + ) + return new + + return replace_submodules(model, cls, factory) diff --git a/model/nn/_ghost_exchange_lmp_mliap.py b/model/nn/_ghost_exchange_lmp_mliap.py new file mode 100644 index 0000000000000000000000000000000000000000..7a20f994d200a8fc024572123f51de9f85babdda --- /dev/null +++ b/model/nn/_ghost_exchange_lmp_mliap.py @@ -0,0 +1,64 @@ +import torch + +from onescience.datapipes.materials.nequip import AtomicDataDict +from ._ghost_exchange_base import GhostExchangeModule + + +# NOTE: can't use custom ops https://docs.pytorch.org/tutorials/advanced/python_custom_ops.html#python-custom-ops-tutorial +# because of complications with `lmp_data` type and PyTorch custom ops registration system + + +class LAMMPSMLIAPGhostExchangeOp(torch.autograd.Function): + @staticmethod + def forward(ctx, *args): + node_features, lmp_data = args + original_shape = node_features.shape + node_features_flat = node_features.view(node_features.size(0), -1) + out_flat = torch.empty_like(node_features_flat) + lmp_data.forward_exchange(node_features_flat, out_flat, out_flat.size(-1)) + + # save for backward + ctx.original_shape = original_shape + ctx.lmp_data = lmp_data + + return out_flat.view(original_shape) + + @staticmethod + def backward(ctx, grad_output): + grad_output_flat = grad_output.view(grad_output.size(0), -1) + gout_flat = torch.empty_like(grad_output_flat) + ctx.lmp_data.reverse_exchange(grad_output_flat, gout_flat, gout_flat.size(-1)) + return gout_flat.view(ctx.original_shape), None + + +class LAMMPSMLIAPGhostExchangeModule(GhostExchangeModule): + """LAMMPS ML-IAP ghost atom exchange module.""" + + def forward( + self, data: AtomicDataDict.Type, ghost_included=False + ) -> AtomicDataDict.Type: + assert AtomicDataDict.LMP_MLIAP_DATA_KEY in data, ( + "`LAMMPSMLIAPGhostExchangeModule` shouldn't be used if LAMMPS ML-IAP data is not provided as input." + ) + + node_features = data[self.field] + lmp_data = data[AtomicDataDict.LMP_MLIAP_DATA_KEY] + + if ghost_included: + local_node_features = torch.narrow(node_features, 0, 0, lmp_data.nlocal) + else: + local_node_features = node_features + num_ghost_atoms = lmp_data.ntotal - lmp_data.nlocal + ghost_zeros = torch.zeros( + (num_ghost_atoms,) + node_features.shape[1:], + dtype=node_features.dtype, + device=node_features.device, + ) + + prepared_node_features = torch.cat((local_node_features, ghost_zeros), dim=0) + + # perform LAMMPS exchange + data[self.field] = LAMMPSMLIAPGhostExchangeOp.apply( + prepared_node_features, lmp_data + ) + return data diff --git a/model/nn/_graph_mixin.py b/model/nn/_graph_mixin.py new file mode 100644 index 0000000000000000000000000000000000000000..590a67f35d8169536d7f913357d6731ce26d3fb6 --- /dev/null +++ b/model/nn/_graph_mixin.py @@ -0,0 +1,238 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +from typing import Dict, Any, Sequence, Union, Optional, Final +from collections import OrderedDict + +import torch + +from e3nn.o3._irreps import Irreps + +from onescience.datapipes.materials.nequip import AtomicDataDict + + +class GraphModuleMixin: + r"""Mixin parent class for ``torch.nn.Module``s that act on and return ``AtomicDataDict.Type`` graph data. + + All such classes should call ``_init_irreps`` in their ``__init__`` functions with information on the data fields they expect, require, and produce, as well as their corresponding irreps. + """ + + _is_graph_module_mixin: Final[bool] = True + # ^ to identify `GraphModuleMixin` types from `torch.package`d models (see https://pytorch.org/docs/stable/package.html#torch-package-sharp-edges) + + def _init_irreps( + self, + irreps_in: Optional[Dict[str, Any]] = None, + my_irreps_in: Optional[Dict[str, Any]] = None, + required_irreps_in: Optional[Sequence[str]] = None, + irreps_out: Optional[Dict[str, Any]] = None, + ): + """Setup the expected data fields and their irreps for this graph module. + + ``None`` is a valid irreps in the context for anything that is invariant but not well described by an ``e3nn.o3.Irreps``. An example are edge indexes in a graph, which are invariant but are integers, not ``0e`` scalars. + + Args: + irreps_in (dict): maps names of all input fields from previous modules or + data to their corresponding irreps + my_irreps_in (dict): maps names of fields to the irreps they must have for + this graph module. Will be checked for consistancy with ``irreps_in`` + required_irreps_in: sequence of names of fields that must be present in + ``irreps_in``, but that can have any irreps. + irreps_out (dict): mapping names of fields that are modified/output by + this graph module to their irreps. + """ + # pattern to handle mutable defaults + irreps_in = {} if irreps_in is None else irreps_in + my_irreps_in = {} if my_irreps_in is None else my_irreps_in + required_irreps_in = () if required_irreps_in is None else required_irreps_in + irreps_out = {} if irreps_out is None else irreps_out + + irreps_in = AtomicDataDict._fix_irreps_dict(irreps_in) + + # positions are *always* 1o, and always present + if AtomicDataDict.POSITIONS_KEY in irreps_in: + if irreps_in[AtomicDataDict.POSITIONS_KEY] != Irreps("1x1o"): + raise ValueError( + f"Positions must have irreps 1o, got instead `{irreps_in[AtomicDataDict.POSITIONS_KEY]}`" + ) + irreps_in[AtomicDataDict.POSITIONS_KEY] = Irreps("1o") + + # edges are also always present + if AtomicDataDict.EDGE_INDEX_KEY in irreps_in: + if irreps_in[AtomicDataDict.EDGE_INDEX_KEY] is not None: + raise ValueError( + f"Edge indexes must have irreps None, got instead `{irreps_in[AtomicDataDict.EDGE_INDEX_KEY]}`" + ) + irreps_in[AtomicDataDict.EDGE_INDEX_KEY] = None + + # atom types are also always present + if AtomicDataDict.ATOM_TYPE_KEY in irreps_in: + if irreps_in[AtomicDataDict.ATOM_TYPE_KEY] is not None: + raise ValueError( + f"atom types must have irreps None, got instead `{irreps_in[AtomicDataDict.ATOM_TYPE_KEY]}`" + ) + irreps_in[AtomicDataDict.ATOM_TYPE_KEY] = None + + my_irreps_in = AtomicDataDict._fix_irreps_dict(my_irreps_in) + + irreps_out = AtomicDataDict._fix_irreps_dict(irreps_out) + # Confirm compatibility: + # with my_irreps_in + for k in my_irreps_in: + if k in irreps_in and irreps_in[k] != my_irreps_in[k]: + raise ValueError( + f"The given input irreps {irreps_in[k]} for field '{k}' is incompatible with this configuration {type(self)}; should have been {my_irreps_in[k]}" + ) + # with required_irreps_in + for k in required_irreps_in: + if k not in irreps_in: + raise ValueError( + f"This {type(self)} requires field '{k}' to be in irreps_in" + ) + # Save stuff + self.irreps_in = irreps_in + # The output irreps of any graph module are whatever inputs it has, overwritten with whatever outputs it has. + new_out = irreps_in.copy() + new_out.update(irreps_out) + self.irreps_out = new_out + + def _add_independent_irreps(self, irreps: Dict[str, Any]): + """ + Insert some independent irreps that need to be exposed to the self.irreps_in and self.irreps_out. + The terms that have already appeared in the irreps_in will be removed. + + Args: + irreps (dict): maps names of all new fields + """ + + irreps = { + key: irrep for key, irrep in irreps.items() if key not in self.irreps_in + } + irreps_in = AtomicDataDict._fix_irreps_dict(irreps) + irreps_out = AtomicDataDict._fix_irreps_dict( + {key: irrep for key, irrep in irreps.items() if key not in self.irreps_out} + ) + self.irreps_in.update(irreps_in) + self.irreps_out.update(irreps_out) + + @torch.jit.unused + def _get_metadata_contributions(self) -> Dict[str, str]: + """Override to provide dynamic metadata at compilation time. + + Modules can override this to contribute metadata based on their current state (e.g., learned parameters). + Called by GraphModel during compilation. + All values must be strings. + + Returns: + Dict[str, str]: Metadata key-value pairs. Can override static config values (e.g., per_edge_type_cutoff) or add new keys. + """ + return {} + + +class SequentialGraphNetwork(GraphModuleMixin, torch.nn.Sequential): + r"""A ``torch.nn.Sequential`` of ``GraphModuleMixin``s. + + Args: + modules (list or dict of ``GraphModuleMixin``s): the sequence of graph modules. If a list, the modules will be named ``"module0", "module1", ...``. + """ + + def __init__( + self, + modules: Union[Sequence[GraphModuleMixin], Dict[str, GraphModuleMixin]], + ): + if isinstance(modules, dict): + module_list = list(modules.values()) + else: + module_list = list(modules) + # check in/out irreps compatible + for m1, m2 in zip(module_list, module_list[1:]): + assert AtomicDataDict._irreps_compatible(m1.irreps_out, m2.irreps_in), ( + f"Incompatible irreps_out from {type(m1).__name__} for input to {type(m2).__name__}: {m1.irreps_out} -> {m2.irreps_in}" + ) + self._init_irreps( + irreps_in=module_list[0].irreps_in, + my_irreps_in=module_list[0].irreps_in, + irreps_out=module_list[-1].irreps_out, + ) + # torch.nn.Sequential will name children correctly if passed an OrderedDict + if isinstance(modules, dict): + modules = OrderedDict(modules) + else: + modules = OrderedDict((f"module{i}", m) for i, m in enumerate(module_list)) + super().__init__(modules) + + @torch.jit.unused + def append(self, name: str, module: GraphModuleMixin) -> None: + r"""Append a module to the SequentialGraphNetwork. + + Args: + name (str): the name for the module + module (GraphModuleMixin): the module to append + """ + assert AtomicDataDict._irreps_compatible(self.irreps_out, module.irreps_in) + self.add_module(name, module) + self.irreps_out = dict(module.irreps_out) + return + + @torch.jit.unused + def insert( + self, + name: str, + module: GraphModuleMixin, + after: Optional[str] = None, + before: Optional[str] = None, + ) -> None: + """Insert a module after the module with name ``after``. + + Args: + name: the name of the module to insert + module: the moldule to insert + after: the module to insert after + before: the module to insert before + """ + + if (before is None) is (after is None): + raise ValueError("Only one of before or after argument needs to be defined") + elif before is None: + insert_location = after + else: + insert_location = before + + # This checks names, etc. + self.add_module(name, module) + # Now insert in the right place by overwriting + names = list(self._modules.keys()) + modules = list(self._modules.values()) + idx = names.index(insert_location) + if before is None: + idx += 1 + names.insert(idx, name) + modules.insert(idx, module) + + self._modules = OrderedDict(zip(names, modules)) + + module_list = list(self._modules.values()) + + # sanity check the compatibility + if idx > 0: + assert AtomicDataDict._irreps_compatible( + module_list[idx - 1].irreps_out, module.irreps_in + ) + if len(module_list) > idx: + assert AtomicDataDict._irreps_compatible( + module_list[idx + 1].irreps_in, module.irreps_out + ) + + # insert the new irreps_out to the later modules + for module_id, next_module in enumerate(module_list[idx + 1 :]): + next_module._add_independent_irreps(module.irreps_out) + + # update the final wrapper irreps_out + self.irreps_out = dict(module_list[-1].irreps_out) + + return + + # Copied from https://pytorch.org/docs/stable/_modules/torch/nn/modules/container.html#Sequential + # with type annotations added + def forward(self, input: AtomicDataDict.Type) -> AtomicDataDict.Type: + for module in self: + input = module(input) + return input diff --git a/model/nn/_tp_scatter_base.py b/model/nn/_tp_scatter_base.py new file mode 100644 index 0000000000000000000000000000000000000000..1779386ba27e4368a4a5711c290ae1195c2c5e94 --- /dev/null +++ b/model/nn/_tp_scatter_base.py @@ -0,0 +1,109 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. + +import torch +from e3nn.o3._tensor_product._tensor_product import TensorProduct +from .utils import scatter +from .model_modifier_utils import replace_submodules, model_modifier + + +class TensorProductScatter(torch.nn.Module): + def __init__( + self, + feature_irreps_in, + irreps_edge_attr, + irreps_mid, + instructions, + ) -> None: + super().__init__() + + self.feature_irreps_in = feature_irreps_in + self.irreps_edge_attr = irreps_edge_attr + self.irreps_mid = irreps_mid + self.instructions = instructions + + self.tp = TensorProduct( + feature_irreps_in, + irreps_edge_attr, + irreps_mid, + instructions, + shared_weights=False, + internal_weights=False, + ) + + self.model_dtype = torch.get_default_dtype() + + def forward(self, x, edge_attr, edge_weight, edge_dst, edge_src): + edge_features = self.tp(x[edge_src], edge_attr, edge_weight) + x = scatter(edge_features, edge_dst, dim=0, dim_size=x.size(0)) + return x + + @model_modifier( + persistent=False, + private=False, + unsupported_devices=["cpu"], + supported_compile_modes=["torchscript", "aotinductor"], + ) + @classmethod + def enable_OpenEquivariance(cls, model): + """ + Enable OpenEquivariance tensor product kernel for accelerated NequIP training and inference. + For usage instructions, see https://nequip.readthedocs.io/en/latest/guide/accelerations/openequivariance.html + """ + + from ._tp_scatter_oeq import OpenEquivarianceTensorProductScatter + from onescience.utils.nequip.internal.dtype import torch_default_dtype + from onescience.utils.nequip.internal.versions.torch_versions import _TORCH_GE_2_7 + + if not _TORCH_GE_2_7: + raise RuntimeError("OpenEquivariance requires PyTorch >= 2.7.") + + _TRAIN_TIME_COMPILE: bool = model.is_compile_graph_model + + def factory(old): + with torch_default_dtype(old.model_dtype): + new = OpenEquivarianceTensorProductScatter( + feature_irreps_in=old.feature_irreps_in, + irreps_edge_attr=old.irreps_edge_attr, + irreps_mid=old.irreps_mid, + instructions=old.instructions, + use_opaque=_TRAIN_TIME_COMPILE, + ) + # c.f. https://github.com/mir-group/nequip/issues/572 + # reuse old.tp to preserve e3nn compiled buffers (_tensor_constant*) + # this ensures state dict compatibility whether the modifier is applied or notwa + new.tp = old.tp + return new + + return replace_submodules(model, cls, factory) + + @model_modifier( + persistent=False, + private=False, + unsupported_devices=["cpu"], + supported_compile_modes=["torchscript", "aotinductor"], + ) + @classmethod + def enable_CuEquivariance(cls, model): + """ + [ALPHA SUPPORT] Enable CuEquivariance tensor product kernel for accelerated NequIP inference. + For usage instructions, see https://nequip.readthedocs.io/en/latest/guide/accelerations/cuequivariance.html + """ + + from ._tp_scatter_cueq import CuEquivarianceTensorProductScatter + from onescience.utils.nequip.internal.dtype import torch_default_dtype + + def factory(old): + with torch_default_dtype(old.model_dtype): + new = CuEquivarianceTensorProductScatter( + feature_irreps_in=old.feature_irreps_in, + irreps_edge_attr=old.irreps_edge_attr, + irreps_mid=old.irreps_mid, + instructions=old.instructions, + ) + # c.f. https://github.com/mir-group/nequip/issues/572 + # reuse old.tp to preserve e3nn compiled buffers (_tensor_constant*) + # this ensures state dict compatibility whether the modifier is applied or not + new.tp = old.tp + return new + + return replace_submodules(model, cls, factory) diff --git a/model/nn/_tp_scatter_cueq.py b/model/nn/_tp_scatter_cueq.py new file mode 100644 index 0000000000000000000000000000000000000000..207a34d09939cd1c33d25f1c5fe6dc8a7223ee01 --- /dev/null +++ b/model/nn/_tp_scatter_cueq.py @@ -0,0 +1,122 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. + +from ._tp_scatter_base import TensorProductScatter + + +def nequip_tp_desc( + irreps1, + irreps2, + irreps3, +): + """Construct the NequIP version of channelwise tensor product descriptor. + + subscripts: ``weights[uv],lhs[iu],rhs[jv],output[ku]`` + + Args: + irreps1 (Irreps): Irreps of the first operand. + irreps2 (Irreps): Irreps of the second operand. + irreps3 (Irreps): Irreps of the output to consider. + """ + import cuequivariance as cue + from cuequivariance.group_theory.irreps_array.irrep_utils import into_list_of_irrep + import itertools + + # modified from `channelwise_tensor_product` + # https://github.com/NVIDIA/cuEquivariance/blob/7236768147394a7da6abd7d5209d274704057eed/cuequivariance/cuequivariance/group_theory/descriptors/irreps_tp.py#L149 + + G = irreps1.irrep_class + irreps3_filter = into_list_of_irrep(G, irreps3) + + d = cue.SegmentedTensorProduct.from_subscripts("uv,iu,jv,kuv+ijk") + + for mul, ir in irreps1: + d.add_segment(1, (ir.dim, mul)) + for mul, ir in irreps2: + d.add_segment(2, (ir.dim, mul)) + + irreps3 = [] + for (i1, (mul1, ir1)), (i2, (mul2, ir2)) in itertools.product( + enumerate(irreps1), enumerate(irreps2) + ): + for ir3 in ir1 * ir2: + if ir3 not in irreps3_filter: + continue + + for cg in cue.clebsch_gordan(ir1, ir2, ir3): + d.add_path(None, i1, i2, None, c=cg, dims={"u": mul1, "v": mul2}) + + irreps3.append((mul1 * mul2, ir3)) + + irreps3 = cue.Irreps(G, irreps3) + irreps3, perm, inv = irreps3.sort() + d = d.permute_segments(3, inv) + d = d.normalize_paths_for_operand(-1) + + return cue.EquivariantPolynomial( + [ + cue.IrrepsAndLayout(irreps1.new_scalars(d.operands[0].size), cue.ir_mul), + cue.IrrepsAndLayout(irreps1, cue.ir_mul), + cue.IrrepsAndLayout(irreps2, cue.ir_mul), + ], + [cue.IrrepsAndLayout(irreps3, cue.ir_mul)], + cue.SegmentedPolynomial.eval_last_operand(d), + ) + + +class CuEquivarianceTensorProductScatter(TensorProductScatter): + _nequip_custom_ops_libs = ("cuequivariance_torch",) + + def __init__( + self, + feature_irreps_in, + irreps_edge_attr, + irreps_mid, + instructions, + ) -> None: + super().__init__( + feature_irreps_in=feature_irreps_in, + irreps_edge_attr=irreps_edge_attr, + irreps_mid=irreps_mid, + instructions=instructions, + ) + # ^ we ensure that the base class keeps around a `self.tp` that carries its own set of persistent buffers + # even though `self.tp` is not used, having its (persistent) buffers always around ensures state dict compatibility when adding on or removing this subclass module + + # === CuEq === + + # we do lazy imports of cuequivariance to allow `nequip-package` to pick this file up even if cuequivariance is not installed + # since `nequip-package` ignores files if it errors on loading the file + + import cuequivariance as cue + import cuequivariance_torch as cuet + from cuequivariance.group_theory.experimental.e3nn import O3_e3nn + + self.tp_conv = cuet.SegmentedPolynomial( + nequip_tp_desc( + cue.Irreps(O3_e3nn, feature_irreps_in), + cue.Irreps(O3_e3nn, irreps_edge_attr), + cue.Irreps(O3_e3nn, irreps_mid), + ) + .flatten_coefficient_modes() + .squeeze_modes() + .polynomial, + method="fused_tp", + math_dtype=self.model_dtype, + ) + + self.transpose_feat = cuet.TransposeIrrepsLayout( + feature_irreps_in, source=cue.mul_ir, target=cue.ir_mul + ) + self.transpose_out = cuet.TransposeIrrepsLayout( + irreps_mid, source=cue.ir_mul, target=cue.mul_ir + ) + + def forward(self, x, edge_attr, edge_weight, edge_dst, edge_src): + return self.transpose_out( + self.tp_conv( + [edge_weight, self.transpose_feat(x), edge_attr], + {1: edge_src}, + {0: x}, + {0: edge_dst}, + )[0] + ) diff --git a/model/nn/_tp_scatter_oeq.py b/model/nn/_tp_scatter_oeq.py new file mode 100644 index 0000000000000000000000000000000000000000..ac84136596d9f365e050806aa544ec240455b1e0 --- /dev/null +++ b/model/nn/_tp_scatter_oeq.py @@ -0,0 +1,57 @@ +from ._tp_scatter_base import TensorProductScatter + + +class OpenEquivarianceTensorProductScatter(TensorProductScatter): + _nequip_custom_ops_libs = ("openequivariance",) + + def __init__( + self, + feature_irreps_in, + irreps_edge_attr, + irreps_mid, + instructions, + use_opaque: bool, + ) -> None: + super().__init__( + feature_irreps_in=feature_irreps_in, + irreps_edge_attr=irreps_edge_attr, + irreps_mid=irreps_mid, + instructions=instructions, + ) + # ^ we ensure that the base class keeps around a `self.tp` that carries its own set of persistent buffers + # even though `self.tp` is not used, having its (persistent) buffers always around ensures state dict compatibility when adding on or removing this subclass module + + # === OEQ === + + # we do lazy imports of oeq to allow `nequip-package` to pick this file up even if oeq is not installed + # since `nequip-package` ignores files if it errors on loading the file + + from openequivariance import ( + TensorProductConv, + TPProblem, + torch_to_oeq_dtype, + ) + + tpp = TPProblem( + feature_irreps_in, + irreps_edge_attr, + irreps_mid, + instructions, + irrep_dtype=torch_to_oeq_dtype(self.model_dtype), + weight_dtype=torch_to_oeq_dtype(self.model_dtype), + shared_weights=False, + internal_weights=False, + ) + self.tp_conv = TensorProductConv( + tpp, torch_op=True, deterministic=False, use_opaque=use_opaque + ) + + def forward(self, x, edge_attr, edge_weight, edge_dst, edge_src): + # explicit cast to account for AMP + return self.tp_conv( + x.to(self.model_dtype), + edge_attr.to(self.model_dtype), + edge_weight.to(self.model_dtype), + edge_dst, + edge_src, + ) diff --git a/model/nn/atomwise.py b/model/nn/atomwise.py new file mode 100644 index 0000000000000000000000000000000000000000..a9b603006465e3041252013f76c689b6a0d6dd5f --- /dev/null +++ b/model/nn/atomwise.py @@ -0,0 +1,378 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch +import torch.nn.functional + +from e3nn.o3._linear import Linear + +from onescience.datapipes.materials.nequip import AtomicDataDict +from onescience.datapipes.materials.nequip._key_registry import get_field_type +from ._graph_mixin import GraphModuleMixin +from .utils import scatter +from .model_modifier_utils import model_modifier, replace_submodules +from onescience.utils.nequip.internal.global_dtype import _GLOBAL_DTYPE + +from typing import Optional, List, Dict, Union + + +class AtomwiseOperation(GraphModuleMixin, torch.nn.Module): + def __init__(self, operation, field: str, irreps_in=None): + super().__init__() + self.operation = operation + self.field = field + self._init_irreps( + irreps_in=irreps_in, + my_irreps_in={field: operation.irreps_in}, + irreps_out={field: operation.irreps_out}, + ) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + data[self.field] = self.operation(data[self.field]) + return data + + +class AtomwiseLinear(GraphModuleMixin, torch.nn.Module): + def __init__( + self, + field: str = AtomicDataDict.NODE_FEATURES_KEY, + out_field: Optional[str] = None, + irreps_in=None, + irreps_out=None, + ): + super().__init__() + self.field = field + out_field = out_field if out_field is not None else field + self.out_field = out_field + if irreps_out is None: + irreps_out = irreps_in[field] + + self._init_irreps( + irreps_in=irreps_in, + required_irreps_in=[field], + irreps_out={out_field: irreps_out}, + ) + self.linear = Linear( + irreps_in=self.irreps_in[field], irreps_out=self.irreps_out[out_field] + ) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + data[self.out_field] = self.linear(data[self.field]) + return data + + +class AtomwiseReduce(GraphModuleMixin, torch.nn.Module): + constant: float + + def __init__( + self, + field: str, + out_field: Optional[str] = None, + reduce="sum", + avg_num_atoms=None, + irreps_in={}, + ): + super().__init__() + assert reduce in ("sum", "mean", "normalized_sum") + self.constant = 1.0 + if reduce == "normalized_sum": + assert avg_num_atoms is not None + self.constant = float(avg_num_atoms) ** -0.5 + reduce = "sum" + self.reduce = reduce + self.field = field + self.out_field = f"{reduce}_{field}" if out_field is None else out_field + self._init_irreps( + irreps_in=irreps_in, + irreps_out=( + {self.out_field: irreps_in[self.field]} + if self.field in irreps_in + else {} + ), + ) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + field = data[self.field] + if AtomicDataDict.BATCH_KEY in data: + result = scatter( + field, + data[AtomicDataDict.BATCH_KEY], + dim=0, + dim_size=AtomicDataDict.num_frames(data), + reduce=self.reduce, + ) + else: + # We can significantly simplify and avoid scatters + if self.reduce == "sum": + result = field.sum(dim=0, keepdim=True) + elif self.reduce == "mean": + result = field.mean(dim=0, keepdim=True) + else: + assert False + if self.constant != 1.0: + result = result * self.constant + data[self.out_field] = result + return data + + +class PerTypeScaleShift(GraphModuleMixin, torch.nn.Module): + """Scale and/or shift a predicted per-atom property based on (learnable) per-species/type parameters. + + Note that scaling/shifting is always done casting into the global dtype (``float64``), even if ``model_dtype`` is a lower precision. + + If a single scalar is provided for scales/shifts, a shortcut implementation is used. Otherwise, a more expensive implementation that assigns separate scales/shifts to each atom type is used. + + If scales/shifts are trainable, the more expensive implementation that assigns separate scales/shifts to each atom type is used, even if a single scalar was provided for the initialization. + """ + + field: str + out_field: str + has_scales: bool + has_shifts: bool + scales_trainble: bool + shifts_trainable: bool + + def __init__( + self, + type_names: List[str], + field: str, + out_field: Optional[str] = None, + scales: Optional[Union[float, Dict[str, float]]] = None, + shifts: Optional[Union[float, Dict[str, float]]] = None, + scales_trainable: bool = False, + shifts_trainable: bool = False, + irreps_in={}, + ): + super().__init__() + self.type_names = type_names + self.num_types = len(type_names) + + # === fields and irreps === + self.field = field + self.out_field = field if out_field is None else out_field + assert get_field_type(self.field) == "node" + assert get_field_type(self.out_field) == "node" + + self._init_irreps( + irreps_in=irreps_in, + my_irreps_in={self.field: "0e"}, # input to shift must be a single scalar + irreps_out={self.out_field: irreps_in[self.field]}, + ) + + # === dtype === + self.out_dtype = _GLOBAL_DTYPE + + # === preprocess scales and shifts === + # we only accept single values or dicts + # lists are no longer supported + if isinstance(scales, list) or isinstance(shifts, list): + raise ValueError( + "\n\nLists are no longer supported for per-type energy scales and shifts. Please use dicts that map from the model's `type_names` as keys to the relevant scale or shift values. For example, the following\n\n per_type_energy_shifts: [1, 2, 3]\n\nshould be changed to\n\n per_type_energy_shifts:\n C: 1\n H: 2\n O: 3\n\n" + ) + + # single valued case + if isinstance(scales, float) or isinstance(scales, int): + scales = [scales] + if isinstance(shifts, float) or isinstance(shifts, int): + shifts = [shifts] + + # dict case + if isinstance(scales, dict): + assert set(self.type_names) == set(scales.keys()) + scales = [scales[name] for name in self.type_names] + if isinstance(shifts, dict): + assert set(self.type_names) == set(shifts.keys()) + shifts = [shifts[name] for name in self.type_names] + + # we convert everything to lists at this point for conversion into `torch.Tensor`s + for sc_vars in (scales, shifts): + if sc_vars is not None: + assert isinstance(sc_vars, list) + + # === scales === + self.has_scales = scales is not None + self.scales_trainable = scales_trainable + if self.has_scales: + scales = torch.as_tensor(scales, dtype=self.out_dtype) + if self.scales_trainable and scales.numel() == 1: + # effective no-op if self.num_types == 1 + scales = ( + torch.ones(self.num_types, dtype=scales.dtype, device=scales.device) + * scales + ) + assert scales.shape == (self.num_types,) or scales.numel() == 1, ( + f"Scales expected to have shape ({self.num_types},), but found {scales.shape}" + ) + scales = scales.reshape(-1, 1) + if self.scales_trainable: + self.scales = torch.nn.Parameter(scales) + else: + self.register_buffer("scales", scales) + else: + self.register_buffer("scales", torch.Tensor()) + self.scales_shortcut = self.scales.numel() == 1 + + # === shifts === + self.has_shifts = shifts is not None + self.shifts_trainable = shifts_trainable + if self.has_shifts: + shifts = torch.as_tensor(shifts, dtype=self.out_dtype) + if self.shifts_trainable and shifts.numel() == 1: + # effective no-op if self.num_types == 1 + shifts = ( + torch.ones(self.num_types, dtype=shifts.dtype, device=shifts.device) + * shifts + ) + assert shifts.shape == (self.num_types,) or shifts.numel() == 1, ( + f"Shifts expected to have shape ({self.num_types},), but found {shifts.shape}" + ) + shifts = shifts.reshape(-1, 1) + if self.shifts_trainable: + self.shifts = torch.nn.Parameter(shifts) + else: + self.register_buffer("shifts", shifts) + else: + self.register_buffer("shifts", torch.Tensor()) + self.shifts_shortcut = self.shifts.numel() == 1 + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + """""" + # shortcut if no scales or shifts found (only dtype promotion performed) + if not (self.has_scales or self.has_shifts): + data[self.out_field] = data[self.field].to(self.out_dtype) + return data + + # === set up === + in_field = data[self.field] + types = data[AtomicDataDict.ATOM_TYPE_KEY].view(-1) + # to account for local-ghost truncation in ML-IAP + types = types[: in_field.size(0)] + + if self.has_scales: + if self.scales_shortcut: + scales = self.scales + else: + scales = torch.nn.functional.embedding(types, self.scales) + else: + scales = self.scales # dummy for torchscript + + if self.has_shifts: + if self.shifts_shortcut: + shifts = self.shifts + else: + shifts = torch.nn.functional.embedding(types, self.shifts) + else: + shifts = self.shifts # dummy for torchscript + + # === explicit cast === + in_field = in_field.to(self.out_dtype) + + # === scale/shift === + if self.has_scales and self.has_shifts: + # we can used an FMA for performance + # addcmul computes + # input + tensor1 * tensor2 elementwise + # it will promote to widest dtype, which comes from shifts/scales + in_field = torch.addcmul(shifts, scales, in_field) + else: + # fallback path for mix of enabled shifts and scales + # multiplication / addition promotes dtypes already, so no cast is needed + if self.has_scales: + in_field = scales * in_field + if self.has_shifts: + in_field = shifts + in_field + + data[self.out_field] = in_field + return data + + @model_modifier(persistent=True, private=False) + @classmethod + def modify_PerTypeScaleShift( + cls, + model, + scales: Optional[Union[float, Dict[str, float]]] = None, + shifts: Optional[Union[float, Dict[str, float]]] = None, + scales_trainable: bool = False, + shifts_trainable: bool = False, + ): + """Modify per-type scales and shifts of a model. + + The new ``scales`` and ``shifts`` should be provided as dicts. + The keys must correspond to the ``type_names`` registered in the model being modified, and may not include all the possible ``type_names`` of the original model. + For example, if one uses a pretrained model with 50 atom types, and seeks to only modify 3 per-atom shifts to be consistent with a fine-tuning dataset's DFT settings, one could use + + .. code-block:: yaml + + shifts: + C: 1.23 + H: 0.12 + O: 2.13 + + In this case, the per-type atomic energy shifts of the original model will be used for every other atom type, except for atom types with the new shifts specified. + + For more details on fine-tuning, see https://nequip.readthedocs.io/en/latest/guide/training-techniques/fine_tuning.html + + Args: + scales: the new per-type atomic energy scales + shifts: the new per-type atomic energy shifts (e.g. isolated atom energies of a dataset used for fine-tuning) + scales_trainable (bool): whether the new scales are trainable + shifts_trainable (bool): whether the new shifts are trainable + """ + + def _helper(sc_var, vname, old): + # get original dict values + orig_sc_var = getattr(old, vname).detach().cpu().reshape(-1).tolist() + # handle special case of single-valued shortcut + if len(orig_sc_var) != len(old.type_names): + assert len(orig_sc_var) == 1 + orig_sc_var = orig_sc_var * len(old.type_names) + new_sc_var = {name: val for name, val in zip(old.type_names, orig_sc_var)} + if sc_var is not None: + # preprocess to list if single number + if isinstance(sc_var, float) or isinstance(sc_var, int): + sc_var = {name: sc_var for name in old.type_names} + assert isinstance(sc_var, dict) + assert all(k in old.type_names for k in sc_var.keys()), ( + f"Provided `{vname}` dict keys ({sc_var.keys()}) do not match the expected type names of the model ({old.type_names})." + ) + # update original model's dict with new dict entries + new_sc_var.update(sc_var) + # if no new values provided, we default to the original model's dict entries + return new_sc_var + + def factory(old): + return cls( + type_names=old.type_names, + field=old.field, + out_field=old.out_field, + scales=_helper(scales, "scales", old), + shifts=_helper(shifts, "shifts", old), + scales_trainable=scales_trainable, + shifts_trainable=shifts_trainable, + irreps_in=old.irreps_in, + ) + + return replace_submodules(model, cls, factory) + + def __repr__(self) -> str: + return f"{self.__class__.__name__} \n scales: {_format_type_vals(self.scales.reshape(-1).tolist(), self.type_names)}\n shifts: {_format_type_vals(self.shifts.reshape(-1).tolist(), self.type_names)}" + + +def _format_type_vals( + vals: List[float], type_names: List[str], element_formatter: str = ".6f" +) -> str: + if vals is None or not vals: + return f"[{', '.join(type_names)}: None]" + + if len(vals) == 1: + return (f"[{', '.join(type_names)}: {{:{element_formatter}}}]").format(vals[0]) + elif len(vals) == len(type_names): + return ( + "[" + + ", ".join( + f"{{{i}[0]}}: {{{i}[1]:{element_formatter}}}" for i in range(len(vals)) + ) + + "]" + ).format(*zip(type_names, vals)) + else: + raise ValueError( + f"Don't know how to format vals=`{vals}` for types {type_names} with element_formatter=`{element_formatter}`" + ) diff --git a/model/nn/compile.py b/model/nn/compile.py new file mode 100644 index 0000000000000000000000000000000000000000..883fac140e6a78d7e9f7ef310d9865e0f1a9357f --- /dev/null +++ b/model/nn/compile.py @@ -0,0 +1,236 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch + +from onescience.datapipes.materials.nequip import AtomicDataDict +from .graph_model import GraphModel +from ._graph_mixin import GraphModuleMixin +from onescience.utils.nequip.internal.dtype import ( + test_model_output_similarity_by_dtype, + _pt2_compile_error_message, +) +from onescience.utils.nequip.internal.fx import nequip_make_fx +from onescience.utils.nequip.internal.dtype import dtype_to_name +from typing import Dict, Sequence, List, Optional, Any, Final +from torch.func import functional_call + + +def _list_to_dict( + keys: Sequence[str], args: List[torch.Tensor] +) -> Dict[str, torch.Tensor]: + return {key: arg for key, arg in zip(keys, args)} + + +def _list_from_dict( + keys: Sequence[str], data: Dict[str, torch.Tensor] +) -> List[torch.Tensor]: + return [data[key] for key in keys] + + +class ListInputOutputWrapper(torch.nn.Module): + """ + Wraps a ``torch.nn.Module`` that takes and returns ``Dict[str, torch.Tensor]`` to have it take and return ``Sequence[torch.Tensor]`` for specified input and output fields. + """ + + def __init__( + self, + model: torch.nn.Module, + input_keys: Sequence[str], + output_keys: Sequence[str], + ): + super().__init__() + self.model = model + self.input_keys = list(input_keys) + self.output_keys = list(output_keys) + + def forward(self, *args: torch.Tensor) -> List[torch.Tensor]: + inputs = _list_to_dict(self.input_keys, args) + outputs = self.model(inputs) + return _list_from_dict(self.output_keys, outputs) + + +class DictInputOutputWrapper(torch.nn.Module): + """ + Wraps a model that takes and returns ``Sequence[torch.Tensor]`` to have it take and return ``Dict[str, torch.Tensor]`` for specified input and output fields (i.e. the opposite of ``ListInputOutputWrapper``). + """ + + def __init__(self, model, input_keys: List[str], output_keys: List[str]): + super().__init__() + self.model = model + self.input_keys = input_keys + self.output_keys = output_keys + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + inputs = _list_from_dict(self.input_keys, data) + with torch.inference_mode(): + outputs = self.model(inputs) + return _list_to_dict(self.output_keys, outputs) + + +class ListInputOutputStateDictWrapper(ListInputOutputWrapper): + """Like ``ListInputOutputWrapper``, but also updates the model with state dict entries before each ``forward`` using ``functional_call``.""" + + def __init__( + self, + model: torch.nn.Module, + input_keys: Sequence[str], + output_keys: Sequence[str], + state_dict_keys: Sequence[str], + ): + super().__init__(model, input_keys, output_keys) + self.state_dict_keys = state_dict_keys + + def forward(self, *args: torch.Tensor) -> List[torch.Tensor]: + # won't check that `args` is of the correct length + input_dict = _list_to_dict(self.input_keys, args[: len(self.input_keys)]) + state_dict = _list_to_dict(self.state_dict_keys, args[len(self.input_keys) :]) + # use functional_call to avoid in-place modification + output_dict = functional_call(self.model, state_dict, args=(input_dict,)) + return _list_from_dict(self.output_keys, output_dict) + + +class CompileGraphModel(GraphModel): + """Wrapper that uses ``torch.compile`` to optimize the wrapped module while allowing it to be trained. + + The cache is keyed by input signature (input keys only). + For each input signature, the eager model is run to determine the output keys, and then a compiled model is created for that input/output combination. + The compiled model and output keys are stored together in the cache. + """ + + is_compile_graph_model: Final[bool] = True + # ^ to identify `GraphModel` types from `nequip-package`d models (see https://pytorch.org/docs/stable/package.html#torch-package-sharp-edges) + + def __init__( + self, + model: GraphModuleMixin, + model_config: Optional[Dict[str, str]] = None, + model_input_fields: Dict[str, Any] = {}, + ) -> None: + super().__init__(model, model_config, model_input_fields) + # cache for multiple compiled variants based on input key signatures + # cache structure: {input_signature: (compiled_model, output_fields)} + # NOTE: the cache dict is wrapped in a tuple so that it's not registered and saved in the state dict -- this is necessary to enable `GraphModel` to load `CompileGraphModel` state dicts + # see https://discuss.pytorch.org/t/saving-nn-module-to-parent-nn-module-without-registering-paremeters/132082/6 + self._compiled_cache = ({},) + # weights and buffers should be done lazily because model modification can happen after instantiation + # such that parameters and buffers may change between class instantiation and the lazy compilation in the `forward` + self.weight_names = None + self.buffer_names = None + + def _get_input_signature(self, data: AtomicDataDict.Type) -> tuple: + """Compute a hashable signature for the input keys. + + The unique set of input keys determines a unique set of output keys when run through the model, + so we only need the input keys for the cache lookup signature. + + Uses intersection of data keys and GraphModel inputs, which assumes: + - correctness of irreps registration system + - this particular batch contains all necessary inputs for this variant + """ + input_keys = tuple(sorted(data.keys() & self.model_input_fields)) + return input_keys + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + # short-circuit if one of the batch dims is 1 (0 would be an error) + # this is related to the 0/1 specialization problem + # see https://docs.google.com/document/d/16VPOa3d-Liikf48teAOmxLc92rgvJdfosIy-yoT38Io/edit?fbclid=IwAR3HNwmmexcitV0pbZm_x1a4ykdXZ9th_eJWK-3hBtVgKnrkmemz6Pm5jRQ&tab=t.0#heading=h.ez923tomjvyk + # we just need something that doesn't have a batch dim of 1 to `make_fx` or else it'll shape specialize + # the models compiled for more batch_size > 1 data cannot be used for batch_size=1 data + # (under specific cases related to the `PerTypeScaleShift` module) + # for now we just make sure to always use the eager model when the data has any batch dims of 1 + if ( + AtomicDataDict.num_nodes(data) < 2 + or AtomicDataDict.num_frames(data) < 2 + or AtomicDataDict.num_edges(data) < 2 + ): + # use parent class's forward + return super().forward(data) + + # === get or compile variant for this input signature === + # compilation happens lazily when we encounter a new combination of input keys + input_signature = self._get_input_signature(data) + cache = self._compiled_cache[0] + + if input_signature not in cache: + # get weight names and buffers (only once on first compilation) + if self.weight_names is None: + self.weight_names = [n for n, _ in self.model.named_parameters()] + self.buffer_names = [n for n, _ in self.model.named_buffers()] + + # == get input fields for this variant == + input_fields = list(input_signature) + + # == run eager model to determine output fields == + eager_output = super().forward(data.copy()) + output_fields = tuple(sorted(eager_output.keys())) + del eager_output + + # == preprocess model and make_fx == + model_to_trace = ListInputOutputStateDictWrapper( + model=self.model, + input_keys=input_fields, + output_keys=output_fields, + state_dict_keys=self.weight_names + self.buffer_names, + ) + + weights, buffers = self._get_weights_buffers() + fx_model = nequip_make_fx( + model=model_to_trace, + data=data, + fields=input_fields, + extra_inputs=weights + buffers, + ) + del weights, buffers + + # == compile exported program == + # see https://pytorch.org/tutorials/intermediate/torch_export_tutorial.html#running-the-exported-program + # TODO: compile options + compiled_model = torch.compile( + fx_model, + dynamic=True, + fullgraph=False, + ) + + # store in cache: (compiled_model, output_fields) + cache[input_signature] = (compiled_model, output_fields) + + # run original model and compiled model with data to sanity check + def compiled_forward_for_test(data_test): + return self._compiled_forward( + data_test, compiled_model, input_fields, output_fields + ) + + # only test output fields that are present in data (i.e. labels are present) + test_fields = sorted(set(output_fields) & data.keys()) + test_model_output_similarity_by_dtype( + compiled_forward_for_test, + self.model, + {k: data[k] for k in input_fields}, + dtype_to_name(self.model_dtype), + fields=test_fields, + error_message=_pt2_compile_error_message, + ) + + # === run compiled model for this variant === + compiled_model, output_fields = cache[input_signature] + out_dict = self._compiled_forward( + data, compiled_model, input_signature, output_fields + ) + to_return = data.copy() + to_return.update(out_dict) + return to_return + + def _compiled_forward(self, data, compiled_model, input_fields, output_fields): + # run compiled model with data + weights, buffers = self._get_weights_buffers() + data_list = _list_from_dict(input_fields, data) + out_list = compiled_model(*(data_list + weights + buffers)) + out_dict = _list_to_dict(output_fields, out_list) + return out_dict + + def _get_weights_buffers(self): + # get weights and buffers from trainable model + weight_dict = dict(self.model.named_parameters()) + weights = [weight_dict[name] for name in self.weight_names] + buffer_dict = dict(self.model.named_buffers()) + buffers = [buffer_dict[name] for name in self.buffer_names] + return weights, buffers diff --git a/model/nn/convnetlayer.py b/model/nn/convnetlayer.py new file mode 100644 index 0000000000000000000000000000000000000000..fb854b8768e95354cf48b02464edab88e0374030 --- /dev/null +++ b/model/nn/convnetlayer.py @@ -0,0 +1,170 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch + +from e3nn.o3._irreps import Irreps +from e3nn.nn._gate import Gate +from e3nn.nn._normact import NormActivation + +from onescience.datapipes.materials.nequip import AtomicDataDict +from ._graph_mixin import GraphModuleMixin +from .interaction_block import InteractionBlock +from .nonlinearities import shifted_softplus +from .utils import tp_path_exists + + +from typing import Any, Dict, Optional, Callable + + +acts = { + "abs": torch.abs, + "tanh": torch.tanh, + "ssp": shifted_softplus, + "silu": torch.nn.functional.silu, +} + + +class ConvNetLayer(GraphModuleMixin, torch.nn.Module): + """ + Args: + + """ + + resnet: bool + + def __init__( + self, + irreps_in, + feature_irreps_hidden, + convolution=InteractionBlock, + convolution_kwargs: Optional[Dict[str, Any]] = None, + resnet: bool = False, + nonlinearity_type: str = "gate", + nonlinearity_scalars: Dict[int, Callable] = {"e": "silu", "o": "tanh"}, + nonlinearity_gates: Dict[int, Callable] = {"e": "silu", "o": "tanh"}, + ): + super().__init__() + # initialization + assert nonlinearity_type in ("gate", "norm") + # make the nonlin dicts from parity ints instead of convinience strs + nonlinearity_scalars = { + 1: nonlinearity_scalars["e"], + -1: nonlinearity_scalars["o"], + } + nonlinearity_gates = { + 1: nonlinearity_gates["e"], + -1: nonlinearity_gates["o"], + } + # normalize optional inputs to avoid shared mutable defaults + convolution_kwargs = ( + {} if convolution_kwargs is None else dict(convolution_kwargs) + ) + + self.feature_irreps_hidden = Irreps(feature_irreps_hidden) + self.resnet = resnet + + # We'll set irreps_out later when we know them + self._init_irreps( + irreps_in=irreps_in, + required_irreps_in=[AtomicDataDict.NODE_FEATURES_KEY], + ) + + edge_attr_irreps = self.irreps_in[AtomicDataDict.EDGE_ATTRS_KEY] + irreps_layer_out_prev = self.irreps_in[AtomicDataDict.NODE_FEATURES_KEY] + + irreps_scalars = Irreps( + [ + (mul, ir) + for mul, ir in self.feature_irreps_hidden + if ir.l == 0 + and tp_path_exists(irreps_layer_out_prev, edge_attr_irreps, ir) + ] + ) + + irreps_gated = Irreps( + [ + (mul, ir) + for mul, ir in self.feature_irreps_hidden + if ir.l > 0 + and tp_path_exists(irreps_layer_out_prev, edge_attr_irreps, ir) + ] + ) + + irreps_layer_out = (irreps_scalars + irreps_gated).simplify() + + if nonlinearity_type == "gate": + ir = ( + "0e" + if tp_path_exists(irreps_layer_out_prev, edge_attr_irreps, "0e") + else "0o" + ) + irreps_gates = Irreps([(mul, ir) for mul, _ in irreps_gated]) + + # TO DO, it's not that safe to directly use the + # dictionary + equivariant_nonlin = Gate( + irreps_scalars=irreps_scalars, + act_scalars=[ + acts[nonlinearity_scalars[ir.p]] for _, ir in irreps_scalars + ], + irreps_gates=irreps_gates, + act_gates=[acts[nonlinearity_gates[ir.p]] for _, ir in irreps_gates], + irreps_gated=irreps_gated, + ) + + conv_irreps_out = equivariant_nonlin.irreps_in.simplify() + + else: + conv_irreps_out = irreps_layer_out.simplify() + + equivariant_nonlin = NormActivation( + irreps_in=conv_irreps_out, + # norm is an even scalar, so use nonlinearity_scalars[1] + scalar_nonlinearity=acts[nonlinearity_scalars[1]], + normalize=True, + epsilon=1e-8, + bias=False, + ) + + self.equivariant_nonlin = equivariant_nonlin + + # TODO: partial resnet? + if irreps_layer_out == irreps_layer_out_prev and resnet: + # We are doing resnet updates and can for this layer + self.resnet = True + else: + self.resnet = False + + # TODO: last convolution should go to explicit irreps out + + # override defaults for irreps: + convolution_kwargs.pop("irreps_in", None) + convolution_kwargs.pop("irreps_out", None) + self.conv = convolution( + irreps_in=self.irreps_in, + irreps_out=conv_irreps_out, + **convolution_kwargs, + ) + + # The output features are whatever we got in + # updated with whatever the convolution outputs (which is a full graph module) + self.irreps_out.update(self.conv.irreps_out) + # but with the features updated by the nonlinearity + self.irreps_out[AtomicDataDict.NODE_FEATURES_KEY] = ( + self.equivariant_nonlin.irreps_out + ) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + # save old features for resnet + old_x = data[AtomicDataDict.NODE_FEATURES_KEY] + # run convolution + data = self.conv(data) + # do nonlinearity + data[AtomicDataDict.NODE_FEATURES_KEY] = self.equivariant_nonlin( + data[AtomicDataDict.NODE_FEATURES_KEY] + ) + # do resnet + if self.resnet: + data[AtomicDataDict.NODE_FEATURES_KEY] = ( + old_x + data[AtomicDataDict.NODE_FEATURES_KEY] + ) + return data diff --git a/model/nn/embedding/__init__.py b/model/nn/embedding/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..a2bb6cd3d5ec32a9e07816a580d1723756450f79 --- /dev/null +++ b/model/nn/embedding/__init__.py @@ -0,0 +1,20 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +from .node import NodeTypeEmbed +from .node_tensor import AppendVectorFieldEmbed +from ._edge import ( + EdgeLengthNormalizer, + BesselEdgeLengthEncoding, + SphericalHarmonicEdgeAttrs, + AddRadialCutoffToData, +) +from .cutoffs import PolynomialCutoff + +__all__ = [ + NodeTypeEmbed, + AppendVectorFieldEmbed, + EdgeLengthNormalizer, + BesselEdgeLengthEncoding, + SphericalHarmonicEdgeAttrs, + AddRadialCutoffToData, + PolynomialCutoff, +] diff --git a/model/nn/embedding/_edge.py b/model/nn/embedding/_edge.py new file mode 100644 index 0000000000000000000000000000000000000000..df67277429a82328f0378dd98ea86e436fc34961 --- /dev/null +++ b/model/nn/embedding/_edge.py @@ -0,0 +1,223 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch + +from e3nn.o3._irreps import Irreps +from e3nn.o3._spherical_harmonics import SphericalHarmonics +from e3nn.util.jit import compile_mode + +from onescience.utils.nequip.internal.global_dtype import _GLOBAL_DTYPE +from onescience.utils.nequip.internal.compile import conditional_torchscript_jit +from onescience.datapipes.materials.nequip import AtomicDataDict +from .._graph_mixin import GraphModuleMixin +from ..utils import with_edge_vectors_, with_edge_type_ +from .utils import cutoff_partialdict_to_tensor + +from typing import Optional, List, Dict, Union + + +@compile_mode("script") +class EdgeLengthNormalizer(GraphModuleMixin, torch.nn.Module): + num_types: int + r_max: float + _per_edge_type: bool + + def __init__( + self, + r_max: float, + type_names: List[str], + per_edge_type_cutoff: Optional[ + Dict[str, Union[float, Dict[str, float]]] + ] = None, + # bookkeeping + edge_type_field: str = AtomicDataDict.EDGE_TYPE_KEY, + norm_length_field: str = AtomicDataDict.NORM_LENGTH_KEY, + irreps_in=None, + ): + super().__init__() + + self.r_max = float(r_max) + self.num_types = len(type_names) + self.edge_type_field = edge_type_field + self.norm_length_field = norm_length_field + + self._per_edge_type = False + if per_edge_type_cutoff is not None: + # process per_edge_type_cutoff + self._per_edge_type = True + per_edge_type_cutoff = cutoff_partialdict_to_tensor( + per_edge_type_cutoff, type_names, self.r_max + ) + # compute 1/rmax and flatten for how they're used in forward, i.e. (n_type, n_type) -> (n_type^2,) + rmax_recip = per_edge_type_cutoff.reciprocal().view(-1) + else: + rmax_recip = torch.as_tensor(1.0 / self.r_max, dtype=_GLOBAL_DTYPE) + self.register_buffer("_rmax_recip", rmax_recip) + + irreps_out = {self.norm_length_field: Irreps([(1, (0, 1))])} + if self._per_edge_type: + irreps_out.update({self.edge_type_field: None}) + + self._init_irreps( + irreps_in=irreps_in, + irreps_out=irreps_out, + ) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + # == get lengths with shape (num_edges, 1) == + data = with_edge_vectors_(data, with_lengths=True) + r = data[AtomicDataDict.EDGE_LENGTH_KEY].view(-1, 1) + # == get norm == + rmax_recip = self._rmax_recip + if self._per_edge_type: + # use helper to get edge types + data = with_edge_type_(data, self.edge_type_field) + edge_type = data[self.edge_type_field] + # convert to row-major NxN matrix index with shape (num_edges,) + edge_type_flat = edge_type[0] * self.num_types + edge_type[1] + # (num_type^2,), (num_edges,) -> (num_edges, 1) + rmax_recip = torch.index_select(rmax_recip, 0, edge_type_flat).unsqueeze(-1) + data[self.norm_length_field] = r * rmax_recip + return data + + +@compile_mode("script") +class BesselEdgeLengthEncoding(GraphModuleMixin, torch.nn.Module): + r"""Bessel edge length encoding. + + Args: + num_bessels (int): number of Bessel basis functions + trainable (bool): whether the :math:`n \pi` coefficients are trainable + cutoff (torch.nn.Module): ``torch.nn.Module`` to apply a cutoff function that smoothly goes to zero at the cutoff radius + """ + + def __init__( + self, + cutoff: torch.nn.Module, + num_bessels: int = 8, + trainable: bool = False, + # bookkeeping + edge_invariant_field: str = AtomicDataDict.EDGE_EMBEDDING_KEY, + norm_length_field: str = AtomicDataDict.NORM_LENGTH_KEY, + irreps_in=None, + ): + super().__init__() + # === process inputs === + self.cutoff = conditional_torchscript_jit(cutoff) + self.num_bessels = num_bessels + self.trainable = trainable + self.edge_invariant_field = edge_invariant_field + self.norm_length_field = norm_length_field + + # === bessel weights === + bessel_weights = torch.linspace( + start=1.0, + end=self.num_bessels, + steps=self.num_bessels, + dtype=_GLOBAL_DTYPE, + ).unsqueeze(0) # (1, num_bessel) + if self.trainable: + self.bessel_weights = torch.nn.Parameter(bessel_weights) + else: + self.register_buffer("bessel_weights", bessel_weights) + + self._init_irreps( + irreps_in=irreps_in, + irreps_out={ + self.edge_invariant_field: Irreps([(self.num_bessels, (0, 1))]), + AtomicDataDict.EDGE_CUTOFF_KEY: "0e", + }, + ) + # i.e. `model_dtype` + self._output_dtype = torch.get_default_dtype() + + def extra_repr(self) -> str: + return f"num_bessels={self.num_bessels}" + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + # == Bessel basis == + x = data[self.norm_length_field] # (num_edges, 1) + # (num_edges, 1), (1, num_bessel) -> (num_edges, num_bessel) + bessel = (torch.sinc(x * self.bessel_weights) * self.bessel_weights).to( + self._output_dtype + ) + + # == polynomial cutoff == + cutoff = self.cutoff(x).to(self._output_dtype) + data[AtomicDataDict.EDGE_CUTOFF_KEY] = cutoff + + # == save product == + data[self.edge_invariant_field] = bessel * cutoff + return data + + +@compile_mode("script") +class SphericalHarmonicEdgeAttrs(GraphModuleMixin, torch.nn.Module): + """Construct edge attrs as spherical harmonic projections of edge vectors. + + Parameters follow ``e3nn.o3.spherical_harmonics``. + + Args: + irreps_edge_sh (int, str, or o3.Irreps): if int, will be treated as lmax for o3.Irreps.spherical_harmonics(lmax) + edge_sh_normalization (str): the normalization scheme to use + edge_sh_normalize (bool, default: True): whether to normalize the spherical harmonics + out_field (str, default: AtomicDataDict.EDGE_ATTRS_KEY: data/irreps field + """ + + out_field: str + + def __init__( + self, + irreps_edge_sh: Union[int, str, Irreps], + edge_sh_normalization: str = "component", + edge_sh_normalize: bool = True, + irreps_in=None, + out_field: str = AtomicDataDict.EDGE_ATTRS_KEY, + ): + super().__init__() + self.out_field = out_field + + if isinstance(irreps_edge_sh, int): + self.irreps_edge_sh = Irreps.spherical_harmonics(irreps_edge_sh) + else: + self.irreps_edge_sh = Irreps(irreps_edge_sh) + self._init_irreps( + irreps_in=irreps_in, + irreps_out={out_field: self.irreps_edge_sh}, + ) + self.sh = SphericalHarmonics( + self.irreps_edge_sh, edge_sh_normalize, edge_sh_normalization + ) + # i.e. `model_dtype` + self._output_dtype = torch.get_default_dtype() + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + data = with_edge_vectors_(data, with_lengths=False) + edge_vec = data[AtomicDataDict.EDGE_VECTORS_KEY] + edge_sh = self.sh(edge_vec) + data[self.out_field] = edge_sh.to(self._output_dtype) + return data + + +@compile_mode("script") +class AddRadialCutoffToData(GraphModuleMixin, torch.nn.Module): + def __init__( + self, + cutoff: torch.nn.Module, + norm_length_field: str = AtomicDataDict.NORM_LENGTH_KEY, + irreps_in=None, + ): + super().__init__() + self.cutoff = conditional_torchscript_jit(cutoff) + self.norm_length_field = norm_length_field + self._init_irreps( + irreps_in=irreps_in, irreps_out={AtomicDataDict.EDGE_CUTOFF_KEY: "0e"} + ) + # i.e. `model_dtype` + self._output_dtype = torch.get_default_dtype() + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + if AtomicDataDict.EDGE_CUTOFF_KEY not in data: + x = data[self.norm_length_field] + cutoff = self.cutoff(x).to(self._output_dtype) + data[AtomicDataDict.EDGE_CUTOFF_KEY] = cutoff + return data diff --git a/model/nn/embedding/cutoffs.py b/model/nn/embedding/cutoffs.py new file mode 100644 index 0000000000000000000000000000000000000000..5f6f431cfa1be62a352077b27c4b4eae3a0775b3 --- /dev/null +++ b/model/nn/embedding/cutoffs.py @@ -0,0 +1,27 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch + + +class PolynomialCutoff(torch.nn.Module): + def __init__(self, p: float = 6): + r"""Polynomial cutoff, as proposed in DimeNet: https://arxiv.org/abs/2003.03123 + + Args: + r_max (float): cutoff radius + p (int) : power used in envelope function + """ + super().__init__() + assert p >= 2.0 + self.p = float(p) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Evaluate cutoff function. + + Args: + x (torch.Tensor): input distance + """ + out = 1.0 + out = out - (((self.p + 1.0) * (self.p + 2.0) / 2.0) * torch.pow(x, self.p)) + out = out + (self.p * (self.p + 2.0) * torch.pow(x, self.p + 1.0)) + out = out - ((self.p * (self.p + 1.0) / 2) * torch.pow(x, self.p + 2.0)) + return out * (x < 1.0) diff --git a/model/nn/embedding/node.py b/model/nn/embedding/node.py new file mode 100644 index 0000000000000000000000000000000000000000..f3371090df7f28474a5665e460b80f240dfd4da0 --- /dev/null +++ b/model/nn/embedding/node.py @@ -0,0 +1,175 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +from dataclasses import dataclass +from math import sqrt +import torch + +from e3nn.o3._irreps import Irreps + +from onescience.datapipes.materials.nequip import AtomicDataDict +from onescience.datapipes.materials.nequip._key_registry import _GRAPH_FIELDS +from .._graph_mixin import GraphModuleMixin + +from typing import Optional, Final, List, Dict, Any + + +@dataclass(frozen=True) +class CategoricalGraphFieldEmbedSpec: + field: str + num_features: int + min: int + max: int + init: Optional[str] = None + + @classmethod + def from_dict(cls, field_embed: Dict[str, Any]) -> "CategoricalGraphFieldEmbedSpec": + required_keys: Final[List[str]] = ["field", "num_features", "min", "max"] + missing_keys = [key for key in required_keys if key not in field_embed] + assert len(missing_keys) == 0, ( + f"missing keys {missing_keys} in `categorical_graph_field_embed` entry; required keys are {required_keys}." + ) + return cls( + field=str(field_embed["field"]), + num_features=int(field_embed["num_features"]), + min=int(field_embed["min"]), + max=int(field_embed["max"]), + init=field_embed.get("init", None), + ) + + +class NodeTypeEmbed(GraphModuleMixin, torch.nn.Module): + """Generates node type embeddings. + + Args: + type_names (List[str]): list of type names + num_features (int): embedding dimension + type_embed_init (str): embedding initialization mode for atom type embeddings. + One of ``"uniform"``, ``"zero"``, ``"near_zero"``, or ``None`` (default, keep PyTorch behavior). + set_features (bool): ``node_features`` will be set in addition to ``node_attrs`` if ``True`` (default) + categorical_graph_field_embed: list of dicts, each dict having keys ``field``, ``num_features``, ``min``, ``max``, and optional ``init``. + ``field`` must correspond to a registered graph data field. + The data dict for the field must be populated by an integer quantity that lies between ``min`` and ``max``. + """ + + num_types: int + set_features: bool + type_embed_init: Optional[str] + + def __init__( + self, + type_names: List[str], + num_features: int, + type_embed_init: Optional[str] = None, + set_features: bool = True, + categorical_graph_field_embed: Optional[List[Dict[str, Any]]] = None, + irreps_in: Optional[Dict[str, Any]] = None, + ): + super().__init__() + # normalize optional inputs to avoid shared mutable defaults + irreps_in = {} if irreps_in is None else dict(irreps_in) + # === bookkeeping === + self.num_types = len(type_names) + self.set_features = set_features + self.type_embed_init = type_embed_init + + # === type embedding module === + self.embed_module = torch.nn.Embedding( + num_embeddings=self.num_types, + embedding_dim=num_features, + ) + self._init_embedding(self.embed_module, init=self.type_embed_init) + + # === categorical graph field embedding === + total_features = num_features + self.categorical_graph_field_embed_modules = torch.nn.ModuleDict() + self.categorical_graph_field_embed_shifts = {} + self.do_categorical_graph_field_embed = False + if categorical_graph_field_embed is not None: + self.do_categorical_graph_field_embed = True + for field_embed_dict in categorical_graph_field_embed: + field_embed = CategoricalGraphFieldEmbedSpec.from_dict(field_embed_dict) + assert field_embed.field in _GRAPH_FIELDS, ( + f"`{field_embed.field}` is not a graph field, only graph fields should be provided to `categorical_graph_field_embed`." + ) + assert field_embed.max >= field_embed.min, ( + f"`max` must be >= `min` for field `{field_embed.field}`." + ) + field_init = field_embed.init + + # == important inits == + embed_module = torch.nn.Embedding( + num_embeddings=field_embed.max - field_embed.min + 1, + embedding_dim=field_embed.num_features, + ) + self._init_embedding(embed_module, init=field_init) + self.categorical_graph_field_embed_modules.update( + {field_embed.field: embed_module} + ) + self.categorical_graph_field_embed_shifts.update( + {field_embed.field: field_embed.min} + ) + # ^ we subtract this quantity to make sure the smallest index is 0 + + # == bookkeeping == + total_features += field_embed.num_features + + # register `irreps_in` if not already done + # needed to ensure that the field is propagated into the model + if field_embed.field not in irreps_in: + # categorical, so no irreps + irreps_in[field_embed.field] = None + + irreps_out = {AtomicDataDict.NODE_ATTRS_KEY: Irreps([(total_features, (0, 1))])} + if self.set_features: + irreps_out[AtomicDataDict.NODE_FEATURES_KEY] = irreps_out[ + AtomicDataDict.NODE_ATTRS_KEY + ] + self._init_irreps(irreps_in=irreps_in, irreps_out=irreps_out) + + @staticmethod + def _init_embedding( + module: torch.nn.Embedding, + init: Optional[str], + ) -> None: + if init is None: + return + if init == "uniform": + torch.nn.init.uniform_(module.weight, -sqrt(3.0), sqrt(3.0)) + elif init == "zero": + torch.nn.init.zeros_(module.weight) + elif init == "near_zero": + torch.nn.init.normal_(module.weight, mean=0.0, std=1e-5) + else: + raise ValueError( + f"unsupported embedding init mode `{init}`. supported modes: ('uniform', 'zero', 'near_zero') or None" + ) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + # (num_atoms, 1) -> (num_atoms, num_type_features) + atom_types = data[AtomicDataDict.ATOM_TYPE_KEY].view(-1) + embedding = self.embed_module(atom_types) + + # handle categorical graph field embeddings + if self.do_categorical_graph_field_embed: + embeddings = [embedding] + for field, module in self.categorical_graph_field_embed_modules.items(): + # (num_graph, 1) -> (num_atoms, 1) + if AtomicDataDict.BATCH_KEY in data: + categorical_graph_field = torch.index_select( + data[field].view(-1), 0, data[AtomicDataDict.BATCH_KEY].view(-1) + ) + else: + categorical_graph_field = ( + data[field].view(-1).expand((atom_types.size(0),)) + ) + # (num_atoms,) -> (num_atoms, num_extra_features) + categorical_graph_field_embedding = module( + categorical_graph_field + - self.categorical_graph_field_embed_shifts[field] + ) + embeddings.append(categorical_graph_field_embedding) + embedding = torch.cat(embeddings, dim=1) + + data[AtomicDataDict.NODE_ATTRS_KEY] = embedding + if self.set_features: + data[AtomicDataDict.NODE_FEATURES_KEY] = embedding + return data diff --git a/model/nn/embedding/node_tensor.py b/model/nn/embedding/node_tensor.py new file mode 100644 index 0000000000000000000000000000000000000000..a6f66fda13b5f8beeae691b7a26cb33d8384e673 --- /dev/null +++ b/model/nn/embedding/node_tensor.py @@ -0,0 +1,171 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +from typing import Any, Dict, List, Optional + +import torch + +from e3nn.o3._irreps import Irreps +from e3nn.o3._spherical_harmonics import SphericalHarmonics + +from onescience.datapipes.materials.nequip import AtomicDataDict +from onescience.datapipes.materials.nequip._key_registry import get_field_type +from .._graph_mixin import GraphModuleMixin + + +class AppendVectorFieldEmbed(GraphModuleMixin, torch.nn.Module): + """Append embedded node or graph vector fields to node features. + + Each field is embedded via solid harmonics up to ``l_max``. + The parity of the input vector must be specified per field: ``+1`` for axial vectors + (pseudovectors, e.g. spin, magnetic field) and ``-1`` for polar vectors (e.g. electric field). + + Args: + vector_fields: dict mapping field name to its vector parity (+1 or -1). + l_max: maximum l for the solid harmonic embedding of each field. + append_to_node_attrs: if True, keep ``node_attrs`` equal to appended ``node_features``. + irreps_in: input irreps dictionary passed to ``GraphModuleMixin``. + """ + + def __init__( + self, + vector_fields: Dict[str, int], + l_max: int, + append_to_node_attrs: bool = True, + irreps_in: Optional[Dict[str, Any]] = None, + ): + super().__init__() + + irreps_in = {} if irreps_in is None else dict(irreps_in) + self.append_to_node_attrs = append_to_node_attrs + + assert AtomicDataDict.NODE_FEATURES_KEY in irreps_in, ( + f"`{AtomicDataDict.NODE_FEATURES_KEY}` must be present in `irreps_in`" + ) + if self.append_to_node_attrs: + assert AtomicDataDict.NODE_ATTRS_KEY in irreps_in, ( + f"`{AtomicDataDict.NODE_ATTRS_KEY}` must be present in `irreps_in` when `append_to_node_attrs=True`" + ) + + assert len(vector_fields) > 0, "`vector_fields` cannot be empty" + assert all(p in (1, -1) for p in vector_fields.values()), ( + "all parity values in `vector_fields` must be +1 (axial) or -1 (polar)" + ) + + # preserve insertion order for consistent forward indexing + self.vector_fields: List[str] = list(vector_fields.keys()) + self.field_kinds: Dict[str, str] = self._validate_fields(self.vector_fields) + + # per-field SH modules; e3nn infers irreps_in ("1e" or "1o") from the output irreps + sh_modules = [] + extra_irreps = Irreps() + for field, parity in vector_fields.items(): + required_irreps = Irreps("1e" if parity == 1 else "1o") + if field in irreps_in: + assert irreps_in[field] == required_irreps, ( + f"`{field}` must have irreps {required_irreps} for parity {parity:+d}, " + f"but got {irreps_in[field]}" + ) + else: + irreps_in[field] = required_irreps + + # degree-l SH of a parity-p vector transforms as (l, p**l): + # axial (p=+1): all even — 0e, 1e, 2e, ... + # polar (p=-1): alternating — 0e, 1o, 2e, ... + # e3nn validates this and auto-infers irreps_in from these labels + field_sh_irreps = Irreps([(1, (l, parity**l)) for l in range(l_max + 1)]) + # don't normalize SH for field vectors; this gives solid harmonics + sh_modules.append( + SphericalHarmonics( + field_sh_irreps, normalize=False, normalization="component" + ) + ) + extra_irreps += field_sh_irreps + + self.sh_modules = torch.nn.ModuleList(sh_modules) + + irreps_out = { + AtomicDataDict.NODE_FEATURES_KEY: ( + irreps_in[AtomicDataDict.NODE_FEATURES_KEY] + extra_irreps + ) + } + if self.append_to_node_attrs: + irreps_out[AtomicDataDict.NODE_ATTRS_KEY] = ( + irreps_in[AtomicDataDict.NODE_ATTRS_KEY] + extra_irreps + ) + required_irreps_in = [AtomicDataDict.NODE_FEATURES_KEY] + if self.append_to_node_attrs: + required_irreps_in.append(AtomicDataDict.NODE_ATTRS_KEY) + required_irreps_in.extend(self.vector_fields) + + self._init_irreps( + irreps_in=irreps_in, + required_irreps_in=required_irreps_in, + irreps_out=irreps_out, + ) + + self.model_dtype = torch.get_default_dtype() + + def __repr__(self) -> str: + lines = [f"{self.__class__.__name__}("] + for field, sh in zip(self.vector_fields, self.sh_modules): + lines.append(f" {field}: {sh.irreps_in} -> {sh.irreps_out},") + lines.append( + f" node_features: {self.irreps_in[AtomicDataDict.NODE_FEATURES_KEY]}" + f" -> {self.irreps_out[AtomicDataDict.NODE_FEATURES_KEY]}" + ) + lines.append(")") + return "\n".join(lines) + + @staticmethod + def _validate_fields(vector_fields: List[str]) -> Dict[str, str]: + assert len(vector_fields) > 0, "`vector_fields` cannot be empty" + field_kinds = {} + for field in vector_fields: + field_kind = get_field_type(field, error_on_unregistered=True) + assert field_kind in ("graph", "node"), ( + f"`{field}` has field type `{field_kind}` but only graph/node fields can be appended" + ) + field_kinds[field] = field_kind + return field_kinds + + def _field_to_per_node( + self, + data: AtomicDataDict.Type, + field: str, + num_nodes: int, + ) -> torch.Tensor: + value = data[field].view(-1, 3) + field_kind = self.field_kinds[field] + # short-circuit of node case + if field_kind == "node": + return value + + # (num_graph, 3) -> (num_nodes, 3) + if AtomicDataDict.BATCH_KEY in data: + batch = data[AtomicDataDict.BATCH_KEY].view(-1) + return torch.index_select(value, 0, batch) + # unbatched case -> all nodes get same value + return value.expand(num_nodes, 3) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + node_features = data[AtomicDataDict.NODE_FEATURES_KEY] + + embedded_fields = [] + for i, sh in enumerate(self.sh_modules): + per_node_vector = self._field_to_per_node( + data=data, + field=self.vector_fields[i], + num_nodes=node_features.size(0), + ) + embedded_fields.append(sh(per_node_vector).to(dtype=self.model_dtype)) + + # build the concatenation input list explicitly to satisfy TorchScript + cat_inputs = [node_features] + for embedded in embedded_fields: + cat_inputs.append(embedded) + node_features = torch.cat(cat_inputs, dim=1) + data[AtomicDataDict.NODE_FEATURES_KEY] = node_features + + if self.append_to_node_attrs: + data[AtomicDataDict.NODE_ATTRS_KEY] = node_features + + return data diff --git a/model/nn/embedding/utils.py b/model/nn/embedding/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..7890f1e3a1a7c1824903f447fce576c76d438350 --- /dev/null +++ b/model/nn/embedding/utils.py @@ -0,0 +1,150 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. + +from typing import List, Dict, Union +import torch + +from onescience.utils.nequip.internal.global_dtype import _GLOBAL_DTYPE + + +# conversion flow: partial_dict -> full_dict -> tensor -> str +# | +# v +# full_dict + + +def cutoff_partialdict_to_fulldict( + partial_dict: Dict[str, Union[float, Dict[str, float]]], + type_names: List[str], + r_max: float, +) -> Dict[str, Dict[str, float]]: + """Convert partial cutoff dict to full dict with all entries. + + Fills missing entries with ``r_max``. + + Args: + partial_dict: partial specification from config, + e.g. ``{"H": 2.0, "C": {"H": 4.0, "C": 3.5}}`` + type_names: list of atom type names + r_max: global cutoff radius (default for missing entries) + + Returns: + full dict with all source -> target pairs specified, + e.g. ``{"H": {"H": 2.0, "C": 2.0}, "C": {"H": 4.0, "C": 3.5}}`` + """ + full_dict = {} + for source_type in type_names: + full_dict[source_type] = {} + if source_type in partial_dict: + entry = partial_dict[source_type] + if isinstance(entry, float): + # uniform cutoff for this source type + for target_type in type_names: + full_dict[source_type][target_type] = entry + else: + # per-target specification + for target_type in type_names: + if target_type in entry: + full_dict[source_type][target_type] = entry[target_type] + else: + # missing target defaults to r_max + full_dict[source_type][target_type] = r_max + else: + # missing source defaults to r_max for all targets + for target_type in type_names: + full_dict[source_type][target_type] = r_max + + return full_dict + + +def cutoff_fulldict_to_tensor( + full_dict: Dict[str, Dict[str, float]], + type_names: List[str], +) -> torch.Tensor: + """Convert full cutoff dict to tensor. + + Args: + full_dict: full specification with all source -> target pairs + type_names: list of atom type names + + Returns: + tensor of shape ``(num_types, num_types)`` with per-edge-type cutoffs + """ + num_types = len(type_names) + cutoff_list = [] + for source_type in type_names: + row = [] + for target_type in type_names: + row.append(full_dict[source_type][target_type]) + cutoff_list.append(row) + + cutoff_tensor = torch.as_tensor(cutoff_list, dtype=_GLOBAL_DTYPE).contiguous() + assert cutoff_tensor.shape == (num_types, num_types) + assert torch.all(cutoff_tensor > 0) + return cutoff_tensor + + +def cutoff_tensor_to_str(cutoff_tensor: torch.Tensor) -> str: + """Convert tensor to metadata string format. + + Args: + cutoff_tensor: cutoff values as tensor (any shape, will be flattened) + + Returns: + space-separated string of cutoff values in row-major order + """ + return " ".join(str(r.item()) for r in cutoff_tensor.reshape(-1)) + + +def cutoff_str_to_fulldict( + cutoff_str: str, + type_names: List[str], +) -> Dict[str, Dict[str, float]]: + """Convert metadata string to full dict format. + + Args: + cutoff_str: space-separated string of cutoff values + type_names: list of atom type names + + Returns: + full dict with all source -> target pairs specified + """ + if cutoff_str in ("", None): + return None + + cutoff_values = [float(x) for x in cutoff_str.split()] + num_types = len(type_names) + + assert len(cutoff_values) == num_types * num_types, ( + f"Expected {num_types * num_types} cutoff values, got {len(cutoff_values)}" + ) + + full_dict = {} + for i, source_type in enumerate(type_names): + full_dict[source_type] = {} + for j, target_type in enumerate(type_names): + full_dict[source_type][target_type] = cutoff_values[i * num_types + j] + + return full_dict + + +def cutoff_partialdict_to_tensor( + partial_dict: Dict[str, Union[float, Dict[str, float]]], + type_names: List[str], + r_max: float, +) -> torch.Tensor: + """Composes ``cutoff_partialdict_to_fulldict`` and ``cutoff_fulldict_to_tensor``.""" + full_dict = cutoff_partialdict_to_fulldict(partial_dict, type_names, r_max) + cutoff_tensor = cutoff_fulldict_to_tensor(full_dict, type_names) + assert torch.all(cutoff_tensor <= r_max) + return cutoff_tensor + + +def cutoff_partialdict_to_str( + partial_dict: Dict[str, Union[float, Dict[str, float]]], + type_names: List[str], + r_max: float, +) -> str: + """Composes ``cutoff_partialdict_to_fulldict``, ``cutoff_fulldict_to_tensor``, and ``cutoff_tensor_to_str``.""" + full_dict = cutoff_partialdict_to_fulldict(partial_dict, type_names, r_max) + tensor = cutoff_fulldict_to_tensor(full_dict, type_names) + return cutoff_tensor_to_str(tensor) diff --git a/model/nn/grad_output.py b/model/nn/grad_output.py new file mode 100644 index 0000000000000000000000000000000000000000..7b3775105573250fced527d66da97750a980bdd0 --- /dev/null +++ b/model/nn/grad_output.py @@ -0,0 +1,320 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. + +import torch + +from e3nn.o3._irreps import Irreps +from e3nn.util.jit import compile_mode + +from onescience.datapipes.materials.nequip import AtomicDataDict +from ._graph_mixin import GraphModuleMixin +from .model_modifier_utils import model_modifier, replace_submodules + + +@compile_mode("unsupported") +class PartialForceOutput(GraphModuleMixin, torch.nn.Module): + r"""Generate partial and total forces from an energy model. + + Args: + func: the energy model + vectorize: the vectorize option to ``torch.autograd.functional.jacobian``, + false by default since it doesn't work well. + """ + + vectorize: bool + + def __init__( + self, + func: GraphModuleMixin, + vectorize: bool = False, + vectorize_warnings: bool = False, + ): + super().__init__() + self.func = func + self.vectorize = vectorize + if vectorize_warnings: + # See https://pytorch.org/docs/stable/generated/torch.autograd.functional.jacobian.html + torch._C._debug_only_display_vmap_fallback_warnings(True) + + # check and init irreps + self._init_irreps( + irreps_in=func.irreps_in, + my_irreps_in={AtomicDataDict.PER_ATOM_ENERGY_KEY: Irreps("0e")}, + irreps_out=func.irreps_out, + ) + self.irreps_out[AtomicDataDict.PARTIAL_FORCE_KEY] = Irreps("1o") + self.irreps_out[AtomicDataDict.FORCE_KEY] = Irreps("1o") + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + data = data.copy() + out_data = {} + + def wrapper(pos: torch.Tensor) -> torch.Tensor: + """Wrapper from pos to atomic energy""" + nonlocal data, out_data + data[AtomicDataDict.POSITIONS_KEY] = pos + out_data = self.func(data) + return out_data[AtomicDataDict.PER_ATOM_ENERGY_KEY].squeeze(-1) + + pos = data[AtomicDataDict.POSITIONS_KEY] + + partial_forces = torch.autograd.functional.jacobian( + func=wrapper, + inputs=pos, + create_graph=self.training, # needed to allow gradients of this output during training + vectorize=self.vectorize, + ) + partial_forces = partial_forces.negative() + # output is [n_at, n_at, 3] + + out_data[AtomicDataDict.PARTIAL_FORCE_KEY] = partial_forces + out_data[AtomicDataDict.FORCE_KEY] = partial_forces.sum(dim=0) + + return out_data + + +@compile_mode("script") +class ForceStressOutput(GraphModuleMixin, torch.nn.Module): + r"""Compute forces (and stress if cell is provided) using autograd of an energy model. + + See: + Knuth et. al. Comput. Phys. Commun 190, 33-50, 2015 + https://pure.mpg.de/rest/items/item_2085135_9/component/file_2156800/content + + Args: + func: the energy model to wrap + """ + + do_derivatives: bool + + def __init__(self, func: GraphModuleMixin, do_derivatives: bool = True): + super().__init__() + self.func = func + self.do_derivatives = do_derivatives + + # check and init irreps + self._init_irreps( + irreps_in=self.func.irreps_in.copy(), + irreps_out=self.func.irreps_out.copy(), + ) + self.irreps_out[AtomicDataDict.FORCE_KEY] = "1o" + self.irreps_out[AtomicDataDict.STRESS_KEY] = "1o" + self.irreps_out[AtomicDataDict.VIRIAL_KEY] = "1o" + self.irreps_out[AtomicDataDict.EDGE_FORCE_KEY] = "1o" + + # for torchscript compat + self.register_buffer("_empty", torch.Tensor()) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + # short-circuit + if not self.do_derivatives: + return self.func(data) + + # === LOGIC BRANCHING NOTES === + # if edge vectors not present, we assume that positions are present + # and proceed with the usual procedure to compute forces, virials, stress + # else, we compute edge forces + + # NOTE: if edge vectors are not present, we assume that it is for non-batched inference with no cell + # at the point of making this change, it is specifically for LAMMPS-MLIAP compatibility + if AtomicDataDict.EDGE_VECTORS_KEY not in data: + if AtomicDataDict.BATCH_KEY in data: + batch = data[AtomicDataDict.BATCH_KEY] + num_batch: int = AtomicDataDict.num_frames(data) + else: + # Special case for efficiency + batch = self._empty + num_batch: int = 1 + + pos = data[AtomicDataDict.POSITIONS_KEY] + has_cell: bool = AtomicDataDict.CELL_KEY in data + + if has_cell: + orig_cell = data[AtomicDataDict.CELL_KEY] + # Make the cell per-batch + cell = orig_cell.view(-1, 3, 3).expand(num_batch, 3, 3) + data[AtomicDataDict.CELL_KEY] = cell + else: + # torchscript + orig_cell = self._empty + cell = self._empty + # Add the displacements + # the GradientOutput will make them require grad + # See SchNetPack code: + # https://github.com/atomistic-machine-learning/schnetpack/blob/master/src/schnetpack/atomistic/model.py#L45 + # SchNetPack issue: + # https://github.com/atomistic-machine-learning/schnetpack/issues/165 + # Paper they worked from: + # Knuth et. al. Comput. Phys. Commun 190, 33-50, 2015 + # https://pure.mpg.de/rest/items/item_2085135_9/component/file_2156800/content + + if num_batch > 1: + displacement = torch.zeros( + (num_batch, 3, 3), + dtype=pos.dtype, + device=pos.device, + ) + else: + displacement = torch.zeros( + (3, 3), + dtype=pos.dtype, + device=pos.device, + ) + displacement.requires_grad_(True) + data["_displacement"] = displacement + # in the above paper, the infinitesimal distortion is *symmetric* + # so we symmetrize the displacement before applying it to + # the positions/cell + # This is not strictly necessary (reasoning thanks to Mario): + # the displacement's asymmetric 1o term corresponds to an + # infinitesimal rotation, which should not affect the final + # output (invariance). + # That said, due to numerical error, this will never be + # exactly true. So, we symmetrize the deformation to + # take advantage of this understanding and not rely on + # the invariance here: + symmetric_displacement = 0.5 * ( + displacement + displacement.transpose(-1, -2) + ) + did_pos_req_grad: bool = pos.requires_grad + pos.requires_grad_(True) + if num_batch > 1: + # bmm is natom in batch + # batched [natom, 1, 3] @ [natom, 3, 3] -> [natom, 1, 3] -> [natom, 3] + data[AtomicDataDict.POSITIONS_KEY] = pos + torch.bmm( + pos.unsqueeze(-2), + torch.index_select(symmetric_displacement, 0, batch), + ).squeeze(-2) + else: + # (num_atoms, 3), (3, 3) -> (num_atoms, 3) + data[AtomicDataDict.POSITIONS_KEY] = pos + torch.sum( + pos.view(-1, 3, 1) * symmetric_displacement, 1 + ) + # assert torch.equal(pos, data[AtomicDataDict.POSITIONS_KEY]) + # we only displace the cell if we have one: + if has_cell: + # bmm is num_batch in batch + # here we apply the distortion to the cell as well + # this is critical also for the correctness + # if we didn't symmetrize the distortion, since without this + # there would then be an infinitesimal rotation of the positions + # but not cell, and it thus wouldn't be global and have + # no effect due to equivariance/invariance. + if num_batch > 1: + # [n_batch, 3, 3] @ [n_batch, 3, 3] + data[AtomicDataDict.CELL_KEY] = cell + torch.bmm( + cell, symmetric_displacement + ) + else: + # [3, 3] @ [3, 3] --- enforced to these shapes + data[AtomicDataDict.CELL_KEY] = ( + cell.view(3, 3) + + torch.sum(cell.view(3, 3, 1) * symmetric_displacement, 1) + ).view(1, 3, 3) + + # Call model and get gradients + data = self.func(data) + + grads = torch.autograd.grad( + [data[AtomicDataDict.TOTAL_ENERGY_KEY].sum()], + [pos, data["_displacement"]], + create_graph=self.training, # needed to allow gradients of this output during training + ) + + # Put negative sign on forces + forces = grads[0] + if forces is None: + # condition needed to unwrap optional for torchscript + assert False, "failed to compute forces autograd" + forces = torch.neg(forces) + data[AtomicDataDict.FORCE_KEY] = forces + + # Store virial + virial = grads[1] + if virial is None: + # condition needed to unwrap optional for torchscript + assert False, "failed to compute virial autograd" + virial = virial.view(num_batch, 3, 3) + + # we only compute the stress (1/V * virial) if we have a cell whose volume we can compute + if has_cell: + # ^ can only scale by cell volume if we have one...: + # Rescale stress tensor + # See https://github.com/atomistic-machine-learning/schnetpack/blob/master/src/schnetpack/atomistic/output_modules.py#L180 + # See also https://en.wikipedia.org/wiki/Triple_product + # See also https://gitlab.com/ase/ase/-/blob/master/ase/cell.py, + # which uses np.abs(np.linalg.det(cell)) + # First dim is batch, second is vec, third is xyz + # Note the .abs(), since volume should always be positive + # det is equal to a dot (b cross c) + volume = torch.linalg.det(cell).abs().unsqueeze(-1) + + # NOTE: to support batching periodic and non-periodic structures together, + # the data processing stage is responsible for ensuring that: + # 1. non-periodic systems have a finite dummy cell to prevent infs in the division below + # 2. stress labels for non-periodic systems are NaN and handled with `ignore_nan` in loss and metrics + + stress = virial / volume.view(num_batch, 1, 1) + data[AtomicDataDict.CELL_KEY] = orig_cell + else: + stress = self._empty # torchscript + data[AtomicDataDict.STRESS_KEY] = stress + + # see discussion in https://github.com/libAtoms/QUIP/issues/227 about sign convention + # (and conventions docs page) + # they say the standard convention is virial = -stress x volume + # looking above this means that we need to pick up another negative sign for the virial + # to fit this equation with the stress computed above + virial = torch.neg(virial) + data[AtomicDataDict.VIRIAL_KEY] = virial + + # Remove helper + del data["_displacement"] + if not did_pos_req_grad: + # don't give later modules one that does + pos.requires_grad_(False) + + else: + # we differentiate wrt EDGE_VECTORS_KEY directly in this branch + # NOTE: we only consider the case of non-batched inference, without a cell + # so no batching, no training considerations, no cell + + # make `edge_vectors` requires grad + edge_vectors = data[AtomicDataDict.EDGE_VECTORS_KEY] + edge_vectors.requires_grad_(True) + data[AtomicDataDict.EDGE_VECTORS_KEY] = edge_vectors + + # do energy model forward and backward + data = self.func(data) + edge_forces = torch.autograd.grad( + [data[AtomicDataDict.TOTAL_ENERGY_KEY].sum()], + [edge_vectors], + # no training arg because we only consider inference + )[0] + # assert needed for TorchScript + assert edge_forces is not None + # NOTE: there shouldn't be a sign flip to match LAMMPS convention + data[AtomicDataDict.EDGE_FORCE_KEY] = edge_forces + + return data + + @model_modifier(persistent=True, private=False) + @classmethod + def enable_ForceStressOutput(cls, model): + """Enable force and stress computation.""" + + def factory(old): + new = cls(func=old.func, do_derivatives=True) + return new + + return replace_submodules(model, cls, factory) + + @model_modifier(persistent=True, private=False) + @classmethod + def disable_ForceStressOutput(cls, model): + """Disable force and stress computation.""" + + def factory(old): + new = cls(func=old.func, do_derivatives=False) + return new + + return replace_submodules(model, cls, factory) diff --git a/model/nn/graph_model.py b/model/nn/graph_model.py new file mode 100644 index 0000000000000000000000000000000000000000..d6e3bff5ae963ee99cc53bba7666661dc9172779 --- /dev/null +++ b/model/nn/graph_model.py @@ -0,0 +1,155 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch + +from onescience.datapipes.materials.nequip import AtomicDataDict +from onescience.utils.nequip.internal.aoti_metadata import NEQUIP_CUSTOM_OPS_LIBS_KEY +from ._graph_mixin import GraphModuleMixin + +from typing import List, Dict, Any, Optional, Final + + +R_MAX_KEY: Final[str] = "r_max" +PER_EDGE_TYPE_CUTOFF_KEY: Final[str] = "per_edge_type_cutoff" +TYPE_NAMES_KEY: Final[str] = "type_names" +NUM_TYPES_KEY: Final[str] = "num_types" +MODEL_DTYPE_KEY: Final[str] = "model_dtype" + + +def _model_metadata_from_config(model_config: Dict[str, str]) -> Dict[str, str]: + model_metadata_dict = {} + # manually process everything + model_metadata_dict[MODEL_DTYPE_KEY] = model_config[MODEL_DTYPE_KEY] + model_metadata_dict[TYPE_NAMES_KEY] = " ".join(model_config[TYPE_NAMES_KEY]) + model_metadata_dict[NUM_TYPES_KEY] = str(len(model_config[TYPE_NAMES_KEY])) + model_metadata_dict[R_MAX_KEY] = str(model_config[R_MAX_KEY]) + + if model_config.get(PER_EDGE_TYPE_CUTOFF_KEY, None) is not None: + from .embedding.utils import cutoff_partialdict_to_str + + model_metadata_dict[PER_EDGE_TYPE_CUTOFF_KEY] = cutoff_partialdict_to_str( + model_config[PER_EDGE_TYPE_CUTOFF_KEY], + model_config[TYPE_NAMES_KEY], + model_config[R_MAX_KEY], + ) + return model_metadata_dict + + +class GraphModel(GraphModuleMixin, torch.nn.Module): + """Top-level module for any complete `nequip` model. + + Manages top-level rescaling, dtypes, and more. + + Args: + model (GraphModuleMixin): model to wrap + model_input_fields (Dict[str, Any]): input fields and their irreps + """ + + model_input_fields: List[str] + is_graph_model: Final[bool] = True + is_compile_graph_model: Final[bool] = False + # ^ to identify `GraphModel` types from `nequip-package`d models (see https://pytorch.org/docs/stable/package.html#torch-package-sharp-edges) + + _metadata: Dict[str, str] + + def __init__( + self, + model: GraphModuleMixin, + model_config: Optional[Dict[str, str]] = None, + model_input_fields: Dict[str, Any] = {}, + ) -> None: + super().__init__() + irreps_in = { + # Things that always make sense as inputs: + AtomicDataDict.POSITIONS_KEY: "1o", + AtomicDataDict.EDGE_INDEX_KEY: None, + AtomicDataDict.EDGE_TRANSPOSE_PERM_KEY: None, + AtomicDataDict.EDGE_CELL_SHIFT_KEY: None, + AtomicDataDict.EDGE_VECTORS_KEY: "1o", + AtomicDataDict.CELL_KEY: "1o", # 3 of them, but still + AtomicDataDict.BATCH_KEY: None, + AtomicDataDict.NUM_NODES_KEY: None, + AtomicDataDict.ATOM_TYPE_KEY: None, + # for LAMMPS ML-IAP + AtomicDataDict.LMP_MLIAP_DATA_KEY: None, + AtomicDataDict.NUM_LOCAL_GHOST_NODES_KEY: None, + } + model_input_fields = AtomicDataDict._fix_irreps_dict(model_input_fields) + irreps_in.update(model_input_fields) + self._init_irreps(irreps_in=irreps_in, irreps_out=model.irreps_out) + for k, irreps in model.irreps_in.items(): + if self.irreps_in.get(k, None) != irreps: + raise RuntimeError( + f"Model has `{k}` in its irreps_in with irreps `{irreps}`, but `{k}` is missing from/has inconsistent irreps in model_input_fields of `{self.irreps_in.get(k, 'missing')}`" + ) + self.model = model + self.model_input_fields = list(self.irreps_in.keys()) + + # the following logic is for backward compatibility and to simplify unittests + self.model_dtype = torch.get_default_dtype() + self._metadata = {} + self.type_names = [] + if model_config is not None: + self._metadata = _model_metadata_from_config(model_config) + self.type_names = self._metadata[TYPE_NAMES_KEY].split(" ") + model_dtype = {"float32": torch.float32, "float64": torch.float64}[ + self._metadata[MODEL_DTYPE_KEY] + ] + assert self.model_dtype == model_dtype + + @property + @torch.jit.unused + def metadata(self) -> Dict[str, str]: + """Get model metadata, including dynamic contributions from modules. + + Collects metadata from all modules that override ``_get_metadata_contributions()``. + Dynamic contributions can override static config values. + Expected to be queried for inference workflows (but not for ``nequip-package``). + """ + out = self._metadata.copy() + + # collect dynamic metadata from module tree + contributed_keys = {} # track which module provided each key + + for name, module in self.model.named_modules(): + if ( + hasattr(module, "_is_graph_module_mixin") + and module._is_graph_module_mixin + ): + contributions = module._get_metadata_contributions() + if not contributions: + continue + + # detect conflicts between multiple modules + for key in contributions: + if key in contributed_keys: + raise ValueError( + f"Metadata conflict: modules '{contributed_keys[key]}' " + f"and '{name}' both contribute key '{key}'" + ) + contributed_keys[key] = name + + # update metadata (overrides static values if keys overlap) + out.update(contributions) + + # update r_max if dynamic per-edge-type cutoffs were contributed + if PER_EDGE_TYPE_CUTOFF_KEY in contributed_keys: + cutoff_values = [float(x) for x in out[PER_EDGE_TYPE_CUTOFF_KEY].split()] + out[R_MAX_KEY] = str(max(cutoff_values)) + + # collect custom ops libs that need to be imported at AOTI load time + custom_ops_libs: set = set() + for m in self.model.modules(): + custom_ops_libs.update(getattr(m, "_nequip_custom_ops_libs", ())) + if custom_ops_libs: + out[NEQUIP_CUSTOM_OPS_LIBS_KEY] = " ".join(sorted(custom_ops_libs)) + + return out + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + # restrict the input data to allowed keys to prevent the model from directly using the dict from the outside, + # preventing weird pass-by-reference bugs + new_data: AtomicDataDict.Type = {} + for k in self.model_input_fields: + if k in data: + new_data[k] = data[k] + return self.model(new_data) diff --git a/model/nn/interaction_block.py b/model/nn/interaction_block.py new file mode 100644 index 0000000000000000000000000000000000000000..58bb0008efef4abb354894b05f4bb4d0895740df --- /dev/null +++ b/model/nn/interaction_block.py @@ -0,0 +1,207 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +"""Interaction Block""" + +import torch + +from e3nn.o3._irreps import Irreps +from e3nn.o3._linear import Linear +from e3nn.o3._tensor_product._sub import FullyConnectedTensorProduct + +from onescience.datapipes.materials.nequip import AtomicDataDict + +from ._graph_mixin import GraphModuleMixin +from .mlp import ScalarMLPFunction +from ._ghost_exchange_base import NoOpGhostExchangeModule +from ._tp_scatter_base import TensorProductScatter +from .norm import AvgNumNeighborsNorm + +from typing import Optional, Sequence, Union, Dict + + +class InteractionBlock(GraphModuleMixin, torch.nn.Module): + use_sc: bool + + def __init__( + self, + irreps_in, + irreps_out, + radial_mlp_depth: int = 1, + radial_mlp_width: int = 8, + use_sc: bool = True, + is_first_layer: bool = False, + type_names: Optional[Sequence[str]] = None, + avg_num_neighbors: Optional[Union[float, Dict[str, float]]] = None, + ) -> None: + """InteractionBlock. + + Args: + irreps_in: input irreps + irreps_out: output irreps + radial_mlp_depth (int): number of radial layers + radial_mlp_width (int): number of hidden neurons in radial function + use_sc (bool): use self-connection or not + is_first_layer (bool): whether to use first layer (default ``False``) + avg_num_neighbors (float/Dict[str, float]): global (float) or per-type (dict) average number of neighbors + type_names (List[str]): list of type names + """ + super().__init__() + + self._init_irreps( + irreps_in=irreps_in, + required_irreps_in=[ + AtomicDataDict.EDGE_EMBEDDING_KEY, + AtomicDataDict.EDGE_ATTRS_KEY, + AtomicDataDict.NODE_FEATURES_KEY, + AtomicDataDict.NODE_ATTRS_KEY, + ], + my_irreps_in={ + AtomicDataDict.EDGE_EMBEDDING_KEY: Irreps( + [ + ( + irreps_in[AtomicDataDict.EDGE_EMBEDDING_KEY].num_irreps, + (0, 1), + ) + ] # (0, 1) is even (invariant) scalars. We are forcing the EDGE_EMBEDDING to be invariant scalars so we can use a dense network + ) + }, + irreps_out={AtomicDataDict.NODE_FEATURES_KEY: irreps_out}, + ) + + # === normalization module === + self.avg_num_neighbors_norm = AvgNumNeighborsNorm( + avg_num_neighbors=avg_num_neighbors, type_names=type_names + ) + + self.use_sc = use_sc + + feature_irreps_in = self.irreps_in[AtomicDataDict.NODE_FEATURES_KEY] + feature_irreps_out = self.irreps_out[AtomicDataDict.NODE_FEATURES_KEY] + irreps_edge_attr = self.irreps_in[AtomicDataDict.EDGE_ATTRS_KEY] + + # - Build modules - + self.linear_1 = Linear( + irreps_in=feature_irreps_in, + irreps_out=feature_irreps_in, + internal_weights=True, + shared_weights=True, + ) + + irreps_mid = [] + instructions = [] + + for i, (mul, ir_in) in enumerate(feature_irreps_in): + for j, (_, ir_edge) in enumerate(irreps_edge_attr): + for ir_out in ir_in * ir_edge: + if ir_out in feature_irreps_out: + k = len(irreps_mid) + irreps_mid.append((mul, ir_out)) + instructions.append((i, j, k, "uvu", True)) + + # We sort the output irreps of the tensor product so that we can simplify them + # when they are provided to the second o3.Linear + irreps_mid = Irreps(irreps_mid) + irreps_mid, p, _ = irreps_mid.sort() + + # Permute the output indexes of the instructions to match the sorted irreps: + instructions = [ + (i_in1, i_in2, p[i_out], mode, train) + for i_in1, i_in2, i_out, mode, train in instructions + ] + + self.tp_scatter = TensorProductScatter( + feature_irreps_in, + irreps_edge_attr, + irreps_mid, + instructions, + ) + + # init_irreps already confirmed that the edge embeddding is all invariant scalars + self.edge_mlp = ScalarMLPFunction( + input_dim=self.irreps_in[AtomicDataDict.EDGE_EMBEDDING_KEY].num_irreps, + output_dim=self.tp_scatter.tp.weight_numel, + hidden_layers_depth=radial_mlp_depth, + hidden_layers_width=radial_mlp_width, + nonlinearity="silu", # hardcode SiLU + bias=False, + forward_weight_init=True, + ) + + self.linear_2 = Linear( + # irreps_mid has uncoallesed irreps because of the uvu instructions, + # but there's no reason to treat them seperately for the Linear + # Note that normalization of o3.Linear changes if irreps are coallesed + # (likely for the better) + irreps_in=irreps_mid.simplify(), + irreps_out=feature_irreps_out, + internal_weights=True, + shared_weights=True, + ) + + self.sc = None + if self.use_sc: + self.sc = FullyConnectedTensorProduct( + feature_irreps_in, + self.irreps_in[AtomicDataDict.NODE_ATTRS_KEY], + feature_irreps_out, + ) + + self.ghost_exchange = NoOpGhostExchangeModule( + field=AtomicDataDict.NODE_FEATURES_KEY, irreps_in=self.irreps_in + ) + + self.is_first_layer = is_first_layer + + @torch.jit.unused + def _get_mliap_num_local(self, data: AtomicDataDict.Type) -> int: + return data[AtomicDataDict.LMP_MLIAP_DATA_KEY].nlocal + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + if AtomicDataDict.LMP_MLIAP_DATA_KEY in data: + num_local_nodes = self._get_mliap_num_local(data) + else: + num_local_nodes = AtomicDataDict.num_nodes(data) + + x = data[AtomicDataDict.NODE_FEATURES_KEY] + + # truncate if not first layer + if not self.is_first_layer: + x = x[:num_local_nodes] + + if self.sc is not None: + node_attrs = data[AtomicDataDict.NODE_ATTRS_KEY] + # truncate if not first layer + if not self.is_first_layer: + node_attrs = node_attrs[:num_local_nodes] + sc = self.sc(x, node_attrs) + + x = self.linear_1(x) + + # normalize before TP-scatter + data[AtomicDataDict.NODE_FEATURES_KEY] = x + data = self.avg_num_neighbors_norm(data) + x = data[AtomicDataDict.NODE_FEATURES_KEY] + + # === comms for ghost-exchange === + # only done if not first layer + # because initial embedding include ghosts since atom types come with ghosts + if not self.is_first_layer: + data[AtomicDataDict.NODE_FEATURES_KEY] = x + data = self.ghost_exchange(data, ghost_included=False) + x = data[AtomicDataDict.NODE_FEATURES_KEY] + + # === TP and scatter === + x = self.tp_scatter( + x=x, + edge_attr=data[AtomicDataDict.EDGE_ATTRS_KEY], + edge_weight=self.edge_mlp(data[AtomicDataDict.EDGE_EMBEDDING_KEY]), + edge_dst=data[AtomicDataDict.EDGE_INDEX_KEY][0], + edge_src=data[AtomicDataDict.EDGE_INDEX_KEY][1], + )[:num_local_nodes] + + x = self.linear_2(x) + + if self.sc is not None: + x = x + sc + + data[AtomicDataDict.NODE_FEATURES_KEY] = x + return data diff --git a/model/nn/misc.py b/model/nn/misc.py new file mode 100644 index 0000000000000000000000000000000000000000..c8c748faea7cc60bcb78635a30274aab5a4ffb57 --- /dev/null +++ b/model/nn/misc.py @@ -0,0 +1,73 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +from typing import List, Optional + +import torch + +from e3nn.o3._irreps import Irreps + +from onescience.datapipes.materials.nequip import AtomicDataDict +from ._graph_mixin import GraphModuleMixin + + +class Concat(GraphModuleMixin, torch.nn.Module): + """Concatenate multiple fields into one.""" + + def __init__(self, in_fields: List[str], out_field: str, irreps_in={}): + super().__init__() + self.in_fields = list(in_fields) + self.out_field = out_field + self._init_irreps(irreps_in=irreps_in, required_irreps_in=self.in_fields) + self.irreps_out[self.out_field] = sum( + (self.irreps_in[k] for k in self.in_fields), Irreps() + ) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + data[self.out_field] = torch.cat([data[k] for k in self.in_fields], dim=-1) + return data + + +class ApplyFactor(GraphModuleMixin, torch.nn.Module): + """Applies factor to field.""" + + def __init__( + self, + in_field: str, + factor: float, + out_field: Optional[str] = None, + irreps_in={}, + ): + super().__init__() + self.in_field = in_field + self.out_field = in_field if out_field is None else out_field + self.factor = factor + self._init_irreps(irreps_in=irreps_in) + self.irreps_out[self.out_field] = self.irreps_in[self.in_field] + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + data[self.out_field] = self.factor * data[self.in_field] + return data + + +class SaveForOutput(torch.nn.Module, GraphModuleMixin): + """Copy a field and disconnect it from the autograd graph. + + Copy a field and disconnect it from the autograd graph, storing it under another key for inspection as part of the models output. + + Args: + field: the field to save + out_field: the key to put the saved copy in + """ + + field: str + out_field: str + + def __init__(self, field: str, out_field: str, irreps_in=None): + super().__init__() + self._init_irreps(irreps_in=irreps_in) + self.irreps_out[out_field] = self.irreps_in[field] + self.field = field + self.out_field = out_field + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + data[self.out_field] = data[self.field].detach().clone() + return data diff --git a/model/nn/mlp.py b/model/nn/mlp.py new file mode 100644 index 0000000000000000000000000000000000000000..d4c16a7d0fd9edee79fa9befab709952260c4d40 --- /dev/null +++ b/model/nn/mlp.py @@ -0,0 +1,271 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +from math import sqrt, prod +import torch + +from e3nn.o3._irreps import Irreps +from e3nn.util.jit import compile_mode + +from onescience.datapipes.materials.nequip import AtomicDataDict +from ._graph_mixin import GraphModuleMixin +from .nonlinearities import ShiftedSoftplus + +from typing import Optional, Final, Dict + + +_NONLINEARITY_MAP: Final[Dict[str, torch.nn.Module]] = { + # NOTE: we include str options for `None` so that the parser always works + None: torch.nn.Identity, + "None": torch.nn.Identity, + "null": torch.nn.Identity, + "silu": torch.nn.SiLU, + "mish": torch.nn.Mish, + "gelu": torch.nn.GELU, + "ssp": ShiftedSoftplus, + "tanh": torch.nn.Tanh, + # not 0 -> 0 + "sigmoid": torch.nn.Sigmoid, + "softplus": torch.nn.Softplus, +} + + +@compile_mode("script") +class ScalarMLP(GraphModuleMixin, torch.nn.Module): + """Apply an MLP to some scalar field.""" + + field: str + out_field: str + + def __init__( + self, + output_dim: int, + hidden_layers_depth: int = 0, + hidden_layers_width: Optional[int] = None, + nonlinearity: Optional[str] = "silu", + bias: bool = False, + forward_weight_init: bool = True, + init_mode: str = "uniform", + parametrization: Optional[str] = None, + field: str = AtomicDataDict.NODE_FEATURES_KEY, + out_field: Optional[str] = None, + irreps_in=None, + ): + super().__init__() + self.field = field + self.out_field = out_field if out_field is not None else field + self._init_irreps( + irreps_in=irreps_in, + required_irreps_in=[self.field], + ) + + assert len(self.irreps_in[self.field]) == 1 + assert self.irreps_in[self.field][0].ir == (0, 1) # scalars + self.mlp_module = ScalarMLPFunction( + input_dim=self.irreps_in[self.field][0].mul, + output_dim=output_dim, + hidden_layers_depth=hidden_layers_depth, + hidden_layers_width=hidden_layers_width, + nonlinearity=nonlinearity, + bias=bias, + forward_weight_init=forward_weight_init, + init_mode=init_mode, + parametrization=parametrization, + ) + self.irreps_out[self.out_field] = Irreps([(self.mlp_module.dims[-1], (0, 1))]) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + data[self.out_field] = self.mlp_module(data[self.field]) + return data + + +@compile_mode("script") +class ScalarMLPFunction(torch.nn.Module): + """Module implementing an MLP according to provided options. + + ``input_dim`` and ``output_dim`` are mandatory arguments. + If only ``input_dim`` and ``output_dim`` are specified, this module defaults to a linear layer (corresponding to the default of ``hidden_layers_depth=0``). + If ``hidden_layers_depth!=0``, ``hidden_layers_width`` must be configured (an error will be raised if the default of ``hidden_layers_width=None`` is used). + + Args: + nonlinearity (str): ``silu`` (default), ``mish``, ``gelu``, ``ssp``, ``tanh``, ``None``, ``null``, or ``"None"`` + bias (bool): whether a bias is included (default ``False``) + forward_weight_init (bool): whether to initialize weights to preserve forward activation variance (default ``True``) or initialize weights to preserve backward gradient variance + """ + + num_layers: int + bias: bool + is_nonlinear: bool + + def __init__( + self, + input_dim: int, + output_dim: int, + hidden_layers_depth: int = 0, + hidden_layers_width: Optional[int] = None, + nonlinearity: Optional[str] = "silu", + bias: bool = False, + forward_weight_init: bool = True, + init_mode: str = "uniform", + parametrization: Optional[str] = None, + ): + super().__init__() + self.bias = bias + + # === process MLP dims === + if hidden_layers_depth != 0: + assert hidden_layers_depth > 0 and hidden_layers_width > 0 + hidden_layers_dims = hidden_layers_depth * [hidden_layers_width] + self.dims = [input_dim] + hidden_layers_dims + [output_dim] + self.num_layers = len(self.dims) - 1 + assert self.num_layers >= 1 + # NOTE: `input_dim` and `output_dim` are always mandatory, which default to at least a linear + # a one-layer MLP is a linear layer + + # === handle nonlinearity === + # TODO: maybe adapt gain to be nonlinearity dependent + if nonlinearity not in _NONLINEARITY_MAP: + available_options = list(_NONLINEARITY_MAP.keys()) + raise ValueError( + f"Unknown nonlinearity '{nonlinearity}'. Available options: {available_options}" + ) + nonlinearity_module = _NONLINEARITY_MAP[nonlinearity] + self.is_nonlinear = False # updated below in loop + + # === build the MLP + weight init === + mlp = torch.nn.Sequential() + for layer, (h_in, h_out) in enumerate(zip(self.dims, self.dims[1:])): + # === weight initialization === + # normalize to preserve variance of forward activations or backward derivatives + # we use "relu" gain (sqrt(2)) as a stand-in for the smooth nonlinearities we use, and only apply them if there is a nonlinearity + # for forward (backward) norm, we don't include the gain for the first (last) layer + # see https://pytorch.org/docs/stable/nn.init.html#torch.nn.init.kaiming_uniform_ + if forward_weight_init: + norm_dim = h_in + gain = 1.0 if nonlinearity is None or (layer == 0) else sqrt(2) + else: + norm_dim = h_out + gain = ( + 1.0 + if nonlinearity is None or (layer == self.num_layers - 1) + else sqrt(2) + ) + # === instantiate `Linear` === + linear_layer = ScalarLinearLayer( + in_features=h_in, + out_features=h_out, + alpha=gain / sqrt(norm_dim), + bias=bias, + init_mode=init_mode, + ) + + # apply parametrization if specified + if parametrization == "spectral_norm": + torch.nn.utils.parametrizations.spectral_norm( + linear_layer, "weight", dim=1 + ) + elif parametrization == "weight_norm": + torch.nn.utils.parametrizations.weight_norm( + linear_layer, "weight", dim=1 + ) + elif parametrization == "orthogonal": + torch.nn.utils.parametrizations.orthogonal(linear_layer, "weight") + elif parametrization not in [None, "None", "null"]: + raise ValueError( + f"Unknown parametrization '{parametrization}'. " + "Available options: None, 'weight_norm', 'orthogonal', 'spectral_norm'" + ) + + mlp.append(linear_layer) + del gain, norm_dim + + # === add nonlinearity (if any) except for last layer === + if (layer != self.num_layers - 1) and (nonlinearity is not None): + # only update `self.is_nonlinear` when a nonlinearity is applied + mlp.append(nonlinearity_module()) + self.is_nonlinear = True + + # use `multidot` based implementation for deep linear net (no nonlinearity, no bias, more than one layer) + # otherwise use the `mlp` built in init + if (not self.is_nonlinear) and (not self.bias) and (self.num_layers > 1): + self.mlp = DeepLinearMLP(mlp) + del mlp + else: + self.mlp = mlp + + def forward(self, x): + return self.mlp(x) + + +class DeepLinearMLP(torch.nn.Module): + def __init__(self, mlp) -> None: + super().__init__() + self.weights = torch.nn.ParameterList() + alphas = [] + for this_idx, mlp_idx in enumerate(range(len(mlp))): + new_weight = torch.clone(mlp[mlp_idx].weight) + self.weights.append(new_weight) + del new_weight + alphas.append(mlp[mlp_idx].alpha) + alpha = prod(alphas) + # the constant has to be a buffer for constant-folding to happen with `torch.compile(...dynamic=True)` + # `persistent=False` for backwards compatibility of checkpoint files + # (and technically preserves the old behavior when using a float in that it's also not persistent) + # `alpha` is already a torch.Tensor here + self.register_buffer("alpha", alpha, persistent=False) + del alphas + + def forward(self, input: torch.Tensor) -> torch.Tensor: + weight = torch.mul( + torch.linalg.multi_dot([weight for weight in self.weights]), self.alpha + ) + return torch.mm(input, weight) + + +class ScalarLinearLayer(torch.nn.Module): + """Module implementing a linear layer with a scaling factor `alpha` applied to the weights.""" + + in_features: int + out_features: int + + def __init__( + self, + in_features: int, + out_features: int, + alpha: float = 1.0, + bias: bool = False, + init_mode: str = "uniform", + ) -> None: + super().__init__() + self.in_features = in_features + self.out_features = out_features + # the constant has to be a buffer for constant-folding to happen with `torch.compile(...dynamic=True)` + # `persistent=False` for backwards compatibility of checkpoint files + # (and technically preserves the old behavior when using a float in that it's also not persistent) + self.register_buffer("alpha", torch.tensor(alpha), persistent=False) + self.weight = torch.nn.Parameter(torch.empty((in_features, out_features))) + # initialize weights based on init_mode + if init_mode == "uniform": + # initialize weights to uniform distribution with mean 0 variance 1 + torch.nn.init.uniform_(self.weight, -sqrt(3), sqrt(3)) + elif init_mode == "normal": + # initialize weights to normal distribution with mean 0 std 1 + torch.nn.init.normal_(self.weight, mean=0.0, std=1.0) + else: + raise ValueError( + f"Unknown init_mode: {init_mode}. Must be 'uniform' or 'normal'." + ) + # initialize bias (if any) to zeros + if bias: + self.bias = torch.nn.Parameter(torch.zeros(out_features)) + else: + self.register_parameter("bias", None) + + def forward(self, input: torch.Tensor) -> torch.Tensor: + # compute scaled weights separately to be constant folded + weight = self.weight * self.alpha + if self.bias is None: + return torch.mm(input, weight) + else: + return torch.addmm(self.bias, input, weight) + + def extra_repr(self) -> str: + return f"in_features={self.in_features}, out_features={self.out_features}, bias={self.bias is not None}, alpha={self.alpha:.6f}" diff --git a/model/nn/model_modifier_utils.py b/model/nn/model_modifier_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..13ebc7516ec57d73ce774ed08840b0b591a5f42e --- /dev/null +++ b/model/nn/model_modifier_utils.py @@ -0,0 +1,107 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch +from typing import Final, Callable, Optional, List + + +# NOTE: persistent modifiers are modifiers that fundamentally change the behavior of the model (same input will lead to different outputs) +# non-persistent modifiers generally refer to accelerations that should preserve similar model behavior, with the only difference being speed +_MODEL_MODIFIER_PERSISTENT_ATTR_NAME: Final[str] = ( + "_nequip_model_modifier_is_persistent" +) +_MODEL_MODIFIER_PRIVATE_ATTR_NAME: Final[str] = "_nequip_model_modifier_is_private" + +# these latter two attributes (unsupported devices and supported compile modes) are meant for acceleration modifiers +_MODEL_MODIFIER_UNSUPPORTED_DEVICES_ATTR_NAME: Final[str] = ( + "_nequip_model_modifier_unsupported_devices" +) +_MODEL_MODIFIER_SUPPORTED_COMPILE_MODES_ATTR_NAME: Final[str] = ( + "_nequip_model_modifier_supported_compile_modes" +) + + +def model_modifier( + persistent: bool, + private: Optional[bool] = None, + unsupported_devices: List[str] = [], + supported_compile_modes: Optional[List[str]] = None, +): + """ + Mark a ``@classmethod`` of an ``nn.Module`` as a "model modifier" that can be applied by the user to modify a packaged or other loaded model on-the-fly. Model modifiers must be a ``@classmethod`` of one of the ``nn.Module`` objects in the model. + + Args: + persistent (bool): Whether the modifier should be applied when building the model for packaging. + private (bool, optional): Whether the modifier is private and should not be exposed in public interfaces. Defaults to None. + unsupported_devices (List[str], optional): List of device types that this modifier does not support. Defaults to []. + supported_compile_modes (List[str], optional): List of compile modes that this modifier supports. Defaults to None. + """ + + def decorator(func): + assert isinstance(func, classmethod), ( + "@model_modifier must be applied after @classmethod" + ) + assert not hasattr(func.__func__, _MODEL_MODIFIER_PERSISTENT_ATTR_NAME) + + setattr(func.__func__, _MODEL_MODIFIER_PERSISTENT_ATTR_NAME, persistent) + + if private is not None: + setattr(func.__func__, _MODEL_MODIFIER_PRIVATE_ATTR_NAME, private) + + setattr( + func.__func__, + _MODEL_MODIFIER_UNSUPPORTED_DEVICES_ATTR_NAME, + unsupported_devices, + ) + + setattr( + func.__func__, + _MODEL_MODIFIER_SUPPORTED_COMPILE_MODES_ATTR_NAME, + supported_compile_modes, + ) + + return func + + return decorator + + +def is_model_modifier(func: callable) -> bool: + # for backwards compatibility, we use the "persistent" flag as a marker for whether the method is a model modifier + return hasattr(func, _MODEL_MODIFIER_PERSISTENT_ATTR_NAME) + + +def is_persistent_model_modifier(func: callable) -> bool: + return getattr(func, _MODEL_MODIFIER_PERSISTENT_ATTR_NAME) + + +def is_private_model_modifier(func: callable) -> Optional[bool]: + # for backwards compatibility of packaged models whose modifier would not have this metadata entry, + # we just default to making it public for convenience of clients + # should be ok since this mechanism is not safety critical and more just a convenience for documenting modifiers + return getattr(func, _MODEL_MODIFIER_PRIVATE_ATTR_NAME, False) + + +def get_model_modifier_unsupported_devices(func: callable) -> List[str]: + """Get the list of unsupported devices for a model modifier. Returns empty list for backwards compatibility.""" + return getattr(func, _MODEL_MODIFIER_UNSUPPORTED_DEVICES_ATTR_NAME, []) + + +def get_model_modifier_supported_compile_modes(func: callable) -> Optional[List[str]]: + """Get the list of supported compile modes for a model modifier. Returns None if not set.""" + return getattr(func, _MODEL_MODIFIER_SUPPORTED_COMPILE_MODES_ATTR_NAME, None) + + +def replace_submodules( + model: torch.nn.Module, + target_cls: type, + factory: Callable[[torch.nn.Module], torch.nn.Module], +) -> torch.nn.Module: + """ + Recursively walk the children of ``model``, and whenever we see an instance of ``target_cls``, replace it (in-place) with ``factory(old_module)`` by mutating ``model._modules[name]``. + """ + for name, child in list(model.named_children()): + if isinstance(child, target_cls): + # build a brand-new one based on `factory` + model._modules[name] = factory(child) + else: + # recurse down + replace_submodules(child, target_cls, factory) + return model diff --git a/model/nn/nonlinearities.py b/model/nn/nonlinearities.py new file mode 100644 index 0000000000000000000000000000000000000000..f22bea2e0f997817cf1538efad461dfd856233e1 --- /dev/null +++ b/model/nn/nonlinearities.py @@ -0,0 +1,20 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch + +import math + +# Technically, use of this module should probably be guarded by conditional_torchscript_jit +# But its use as a drop-in replacement for functions like torch.nn.functional.silu makes that +# difficult, so, given the rarety of its use, we have just removed @torch.jit.script + + +def shifted_softplus(x): + return torch.nn.functional.softplus(x) - math.log(2.0) + + +class ShiftedSoftplus(torch.nn.Module): + def __init__(self): + super().__init__() + + def forward(self, x): + return shifted_softplus(x) diff --git a/model/nn/norm.py b/model/nn/norm.py new file mode 100644 index 0000000000000000000000000000000000000000..072cf2d33960e2478861cc701dac2932c88c62aa --- /dev/null +++ b/model/nn/norm.py @@ -0,0 +1,68 @@ +import torch +from typing import Union, Sequence, Dict +from math import sqrt +from onescience.datapipes.materials.nequip import AtomicDataDict + + +class AvgNumNeighborsNorm(torch.nn.Module): + def __init__( + self, + type_names: Sequence[str], + avg_num_neighbors: Union[float, Dict[str, float]], + ) -> None: + """ + Module to normalize features during training using per type edge sum normalization. + + Args: + type_names (Sequence[str]): list of atom type names + avg_num_neighbors (float/Dict[str, float]): used to normalize edge sums for better numerics + """ + super().__init__() + assert avg_num_neighbors is not None, "avg_num_neighbors must be specified" + + self.in_field = self.out_field = AtomicDataDict.NODE_FEATURES_KEY + self.norm_key = AtomicDataDict.FEATURE_NORM_FACTOR_KEY + + # Put avg_num_neighbors in a list (global or per type) + if isinstance(avg_num_neighbors, (float, int)): + avg_num_neighbors = [avg_num_neighbors] + elif isinstance(avg_num_neighbors, dict): + assert set(type_names) == set(avg_num_neighbors.keys()) + avg_num_neighbors = [avg_num_neighbors[k] for k in type_names] + else: + raise RuntimeError( + "Unrecognized format for `avg_num_neighbors`, only floats or dicts allowed." + ) + assert isinstance(avg_num_neighbors, list) + + # Tensorize avg_num_neighbors and register as buffer + norm_const = torch.tensor([(1.0 / sqrt(N)) for N in avg_num_neighbors]) + norm_const = norm_const.reshape(-1, 1) + # Persistent=False to ensure backwards compatibility of FMs. + # TODO remove this once we're sure FMs are not using this anymore + self.register_buffer("norm_const", norm_const, persistent=False) + + # If global avg_num_neighbors or only one type, no need to do embedding lookup in forward + self.norm_shortcut = self.norm_const.numel() == 1 + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + features = data[self.in_field] + norm_size = features.size(0) + + if self.norm_key in data and data[self.norm_key].size(0) == norm_size: + norm_factor = data[self.norm_key] + else: + # Compute norm factor for the first time + if self.norm_shortcut: + # No need to do embedding lookup in forward + norm_factor = self.norm_const.expand(norm_size, -1) + else: + # Embed each avg_num_neighbors value per type + norm_factor = torch.nn.functional.embedding( + data[AtomicDataDict.ATOM_TYPE_KEY][:norm_size], + self.norm_const, + ) + data[self.norm_key] = norm_factor # shape: (num_local_nodes, 1) + + data[self.out_field] = norm_factor * features + return data diff --git a/model/nn/pair_potential.py b/model/nn/pair_potential.py new file mode 100644 index 0000000000000000000000000000000000000000..a9daaa8f4eb86ac88c9dae330bcd49ba3b15ba59 --- /dev/null +++ b/model/nn/pair_potential.py @@ -0,0 +1,390 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +from typing import Union, Optional, List + +import torch + +from e3nn.o3._irreps import Irreps +from e3nn.util.jit import compile_mode + +from onescience.datapipes.materials.nequip import AtomicDataDict +from onescience.datapipes.materials.nequip.misc import chemical_symbols_to_atomic_numbers_dict +from ._graph_mixin import GraphModuleMixin +from .utils import scatter, with_edge_vectors_ +from onescience.utils.nequip.internal.compile import conditional_torchscript_jit +from .embedding.cutoffs import PolynomialCutoff + + +class _LJParam(torch.nn.Module): + def __init__(self): + super().__init__() + + def forward(self, param, index1, index2): + if param.ndim == 2: + # make it symmetric + param = param.triu() + param.triu(1).transpose(-1, -2) + # get for each atom pair + param = torch.index_select( + param.view(-1), 0, index1 * param.shape[0] + index2 + ) + # make it positive + param = param.relu() # TODO: better way? + return param + + +@compile_mode("script") +class LennardJones(GraphModuleMixin, torch.nn.Module): + """Lennard-Jones and related pair potentials.""" + + lj_style: str + exponent: float + + def __init__( + self, + type_names: List[str], + lj_sigma: Union[torch.Tensor, float], + lj_delta: Union[torch.Tensor, float] = 0, + lj_epsilon: Optional[Union[torch.Tensor, float]] = None, + lj_sigma_trainable: bool = False, + lj_delta_trainable: bool = False, + lj_epsilon_trainable: bool = False, + lj_exponent: Optional[float] = None, + lj_per_type: bool = True, + lj_style: str = "lj", + polynomial_cutoff_p: float = 6.0, + per_atom_energy_field: str = AtomicDataDict.PER_ATOM_ENERGY_KEY, + irreps_in=None, + ) -> None: + super().__init__() + num_types = len(type_names) + self.per_atom_energy_field = per_atom_energy_field + + # === irreps registration === + self._init_irreps( + irreps_in=irreps_in, + required_irreps_in=[AtomicDataDict.NORM_LENGTH_KEY], + irreps_out={self.per_atom_energy_field: "0e"}, + ) + if self.per_atom_energy_field in self.irreps_in: + energy_irreps = Irreps(self.irreps_in[self.per_atom_energy_field]) + assert all(ir.l == 0 for _, ir in energy_irreps), ( + f"{self.per_atom_energy_field} must be scalar irreps, found {energy_irreps}" + ) + self.irreps_out[self.per_atom_energy_field] = energy_irreps + + assert lj_style in ("lj", "lj_repulsive_only", "repulsive") + self.lj_style = lj_style + + for param, (value, trainable) in { + "epsilon": (lj_epsilon, lj_epsilon_trainable), + "sigma": (lj_sigma, lj_sigma_trainable), + "delta": (lj_delta, lj_delta_trainable), + }.items(): + if value is None: + self.register_buffer(param, torch.Tensor()) # torchscript + continue + value = torch.as_tensor(value, dtype=torch.get_default_dtype()) + if value.ndim == 0 and lj_per_type: + # one scalar for all pair types + value = ( + torch.ones( + num_types, num_types, device=value.device, dtype=value.dtype + ) + * value + ) + elif value.ndim == 2: + assert lj_per_type + # one per pair type, check symmetric + assert value.shape == (num_types, num_types) + # per-species square, make sure symmetric + assert torch.equal(value, value.T) + value = torch.triu(value) + else: + raise ValueError + setattr(self, param, torch.nn.Parameter(value, requires_grad=trainable)) + + if lj_exponent is None: + lj_exponent = 6.0 + self.exponent = lj_exponent + + self.cutoff = conditional_torchscript_jit(PolynomialCutoff(polynomial_cutoff_p)) + self.model_dtype = torch.get_default_dtype() + self._param = conditional_torchscript_jit(_LJParam()) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + data = with_edge_vectors_(data, with_lengths=True) + edge_center = data[AtomicDataDict.EDGE_INDEX_KEY][0] + atom_types = data[AtomicDataDict.ATOM_TYPE_KEY] + edge_len = data[AtomicDataDict.EDGE_LENGTH_KEY].unsqueeze(-1) + edge_types = torch.index_select( + atom_types, 0, data[AtomicDataDict.EDGE_INDEX_KEY].reshape(-1) + ).view(2, -1) + index1 = edge_types[0] + index2 = edge_types[1] + + sigma = self._param(self.sigma, index1, index2) + delta = self._param(self.delta, index1, index2) + epsilon = self._param(self.epsilon, index1, index2) + + if self.lj_style == "repulsive": + # 0.5 to assign half and half the energy to each side of the interaction + lj_eng = 0.5 * epsilon * ((sigma * (edge_len - delta)) ** -self.exponent) + else: + lj_eng = (sigma / (edge_len - delta)) ** self.exponent + lj_eng = torch.neg(lj_eng) + lj_eng = lj_eng + lj_eng.square() + # 2.0 because we do the slightly symmetric thing and let + # ij and ji each contribute half of the LJ energy of the pair + # this avoids indexing out certain edges in the general case where + # the edges are not ordered. + lj_eng = (2.0 * epsilon) * lj_eng + + if self.lj_style == "lj_repulsive_only": + # if taking only the repulsive part, shift up so the minima is at eng=0 + lj_eng = lj_eng + epsilon + # this is continuous at the minima, and we mask out everything greater + # TODO: this is probably broken with NaNs at delta + lj_eng = lj_eng * (edge_len < (2 ** (1.0 / self.exponent) + delta)) + + # apply polynomial cutoff from this module's own normalized edge lengths + lj_edge_cutoff = self.cutoff(data[AtomicDataDict.NORM_LENGTH_KEY]).to( + self.model_dtype + ) + lj_eng = lj_eng.to(self.model_dtype) * lj_edge_cutoff + + # sum edge LJ energies onto atoms + atomic_eng = scatter( + lj_eng, + edge_center, + dim=0, + dim_size=AtomicDataDict.num_nodes(data), + ) + if self.per_atom_energy_field in data: + atomic_eng = atomic_eng + data[self.per_atom_energy_field] + data[self.per_atom_energy_field] = atomic_eng + return data + + def __repr__(self) -> str: + def _f(e): + e = e.data + if e.ndim == 0: + return f"{e:.6f}" + elif e.ndim == 2: + return f"{e}" + + return f"PairPotential(lj_style={self.lj_style} | σ={_f(self.sigma)} δ={_f(self.delta)} ε={_f(self.epsilon)} exp={self.exponent:.1f})" + + +@compile_mode("script") +class SimpleLennardJones(GraphModuleMixin, torch.nn.Module): + """Simple Lennard-Jones.""" + + lj_sigma: float + lj_epsilon: float + + def __init__( + self, + lj_sigma: float, + lj_epsilon: float, + polynomial_cutoff_p: float = 6.0, + irreps_in=None, + ) -> None: + super().__init__() + self._init_irreps( + irreps_in=irreps_in, + required_irreps_in=[AtomicDataDict.NORM_LENGTH_KEY], + irreps_out={AtomicDataDict.PER_ATOM_ENERGY_KEY: "0e"}, + ) + self.lj_sigma = lj_sigma + self.lj_epsilon = lj_epsilon + self.cutoff = conditional_torchscript_jit(PolynomialCutoff(polynomial_cutoff_p)) + self.model_dtype = torch.get_default_dtype() + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + data = with_edge_vectors_(data, with_lengths=True) + edge_center = data[AtomicDataDict.EDGE_INDEX_KEY][0] + edge_len = data[AtomicDataDict.EDGE_LENGTH_KEY].unsqueeze(-1) + + lj_eng = (self.lj_sigma / edge_len) ** 6.0 + lj_eng = lj_eng.square() - lj_eng + lj_eng = 2 * self.lj_epsilon * lj_eng + + # apply polynomial cutoff from this module's own normalized edge lengths + lj_edge_cutoff = self.cutoff(data[AtomicDataDict.NORM_LENGTH_KEY]).to( + self.model_dtype + ) + lj_eng = lj_eng.to(self.model_dtype) * lj_edge_cutoff + + # sum edge LJ energies onto atoms + atomic_eng = scatter( + lj_eng, + edge_center, + dim=0, + dim_size=AtomicDataDict.num_nodes(data), + ) + if AtomicDataDict.PER_ATOM_ENERGY_KEY in data: + atomic_eng = atomic_eng + data[AtomicDataDict.PER_ATOM_ENERGY_KEY] + data[AtomicDataDict.PER_ATOM_ENERGY_KEY] = atomic_eng + return data + + +class _ZBL(torch.nn.Module): + def __init__(self): + super().__init__() + + def forward( + self, + Z: torch.Tensor, + r: torch.Tensor, + atom_types: torch.Tensor, + edge_index: torch.Tensor, + qqr2exesquare: float, + ) -> torch.Tensor: + # from LAMMPS pair_zbl_const.h + pzbl: float = 0.23 + a0: float = 0.46850 + c1: float = 0.02817 + c2: float = 0.28022 + c3: float = 0.50986 + c4: float = 0.18175 + d1: float = -0.20162 + d2: float = -0.40290 + d3: float = -0.94229 + d4: float = -3.19980 + # (num_atoms,) -> (num_atoms, 1) + node_Zs = torch.nn.functional.embedding(atom_types.view(-1), Z.view(-1, 1)) + # (num_atoms,) -> (2 * num_edges,) + edge_Zs = torch.nn.functional.embedding(edge_index.view(-1), node_Zs).view( + 2, -1 + ) + Zi = torch.select(edge_Zs, 0, 0) + Zj = torch.select(edge_Zs, 0, 1) + del node_Zs, edge_Zs + x = ((torch.pow(Zi, pzbl) + torch.pow(Zj, pzbl)) * r) / a0 + psi = ( + c1 * (d1 * x).exp() + + c2 * (d2 * x).exp() + + c3 * (d3 * x).exp() + + c4 * (d4 * x).exp() + ) + eng = qqr2exesquare * ((Zi * Zj) / r) * psi + return eng + + +@compile_mode("script") +class ZBL(GraphModuleMixin, torch.nn.Module): + """`ZBL `_ pair potential energy term. + + Useful as a prior for core repulsion to mitigate molecular dynamics failure modes associated with atoms getting too close. + + Args: + type_names (List[str]): list of type names known by the model, ``[atom1, atom2, atom3]`` + chemical_species (List[str]): list of chemical symbols, e.g. ``[C, H, O]`` + units (str): `LAMMPS units `_ that the data is in; ``metal`` and ``real`` are presently supported -- raise a GitHub issue if more is desired + polynomial_cutoff_p (float): exponent used for the polynomial cutoff (default ``6``) + """ + + def __init__( + self, + type_names: List[str], + chemical_species: List[str], + units: str, + polynomial_cutoff_p: float = 6.0, + per_atom_energy_field: str = AtomicDataDict.PER_ATOM_ENERGY_KEY, + irreps_in=None, + ): + super().__init__() + num_types = len(type_names) + self.per_atom_energy_field = per_atom_energy_field + + # === irreps registration === + self._init_irreps( + irreps_in=irreps_in, + required_irreps_in=[AtomicDataDict.NORM_LENGTH_KEY], + irreps_out={self.per_atom_energy_field: "0e"}, + ) + if self.per_atom_energy_field in self.irreps_in: + energy_irreps = Irreps(self.irreps_in[self.per_atom_energy_field]) + assert all(ir.l == 0 for _, ir in energy_irreps), ( + f"{self.per_atom_energy_field} must be scalar irreps, found {energy_irreps}" + ) + self.irreps_out[self.per_atom_energy_field] = energy_irreps + + assert len(chemical_species) == num_types + atomic_numbers: List[int] = [ + chemical_symbols_to_atomic_numbers_dict[chemical_species[type_i]] + for type_i in range(num_types) + ] + if min(atomic_numbers) < 1: + raise ValueError( + f"Your chemical symbols don't seem valid (minimum atomic number is {min(atomic_numbers)} < 1); did you try to use fake chemical symbols for arbitrary atom types?" + ) + + # LAMMPS note on units: + # > The numerical values of the exponential decay constants in the + # > screening function depend on the unit of distance. In the above + # > equation they are given for units of Angstroms. LAMMPS will + # > automatically convert these values to the distance unit of the + # > specified LAMMPS units setting. The values of Z should always be + # > given as multiples of a proton’s charge, e.g. 29.0 for copper. + # So, we store the atomic numbers directly. + self.register_buffer( + "atomic_numbers", + torch.as_tensor(atomic_numbers, dtype=torch.get_default_dtype()), + ) + # And we have to convert our value of prefector into the model's physical units + # Here, prefactor is (electron charge)^2 / (4 * pi * electrical permisivity of vacuum) + # we have a value for that in eV and Angstrom + # See https://github.com/lammps/lammps/blob/c415385ab4b0983fa1c72f9e92a09a8ed7eebe4a/src/update.cpp#L187 for values from LAMMPS + # LAMMPS uses `force->qqr2e * force->qelectron * force->qelectron` + # Make it a buffer so rescalings are persistent, it still acts as a scalar Tensor + self.register_buffer( + "_qqr2exesquare", + torch.as_tensor( + {"metal": 14.399645 * (1.0) ** 2, "real": 332.06371 * (1.0) ** 2}[ + units + ], + dtype=torch.float64, + ) + * 0.5, # Put half the energy on each of ij, ji + ) + self.cutoff = conditional_torchscript_jit(PolynomialCutoff(polynomial_cutoff_p)) + self.model_dtype = torch.get_default_dtype() + self._zbl = conditional_torchscript_jit(_ZBL()) + + def forward(self, data: AtomicDataDict.Type) -> AtomicDataDict.Type: + """""" + data = with_edge_vectors_(data, with_lengths=True) + edge_center = data[AtomicDataDict.EDGE_INDEX_KEY][0] + + # account for possibility of reduced num nodes in atomic energy in a local-ghost atom context + if self.per_atom_energy_field in data: + num_nodes = data[self.per_atom_energy_field].size(0) + else: + num_nodes = AtomicDataDict.num_nodes(data) + + zbl_edge_eng = self._zbl( + Z=self.atomic_numbers, + r=data[AtomicDataDict.EDGE_LENGTH_KEY].view(-1), + atom_types=data[AtomicDataDict.ATOM_TYPE_KEY], + edge_index=data[AtomicDataDict.EDGE_INDEX_KEY], + qqr2exesquare=self._qqr2exesquare, + ).unsqueeze(-1) + + # apply cutoff + zbl_edge_cutoff = self.cutoff(data[AtomicDataDict.NORM_LENGTH_KEY]).to( + self.model_dtype + ) + zbl_edge_eng = zbl_edge_eng * zbl_edge_cutoff + atomic_eng = scatter( + zbl_edge_eng, + edge_center, + dim=0, + dim_size=num_nodes, + ) + if self.per_atom_energy_field in data: + atomic_eng = atomic_eng + data[self.per_atom_energy_field] + data[self.per_atom_energy_field] = atomic_eng + return data + + +__all__ = [LennardJones, ZBL] diff --git a/model/nn/utils.py b/model/nn/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..c1b0effdcb310a989d06e66c7d6c11726ba46b29 --- /dev/null +++ b/model/nn/utils.py @@ -0,0 +1,177 @@ +# This file is a part of the `nequip` package. Please see LICENSE and README at the root for information on using it. +import torch +from e3nn.o3._irreps import Irrep, Irreps +from onescience.datapipes.materials.nequip import AtomicDataDict +from typing import Optional + +""" +Migrated from https://github.com/mir-group/pytorch_runstats +""" + + +def _broadcast(src: torch.Tensor, other: torch.Tensor, dim: int): + if dim < 0: + dim = other.dim() + dim + if src.dim() == 1: + for _ in range(0, dim): + src = src.unsqueeze(0) + for _ in range(src.dim(), other.dim()): + src = src.unsqueeze(-1) + src = src.expand_as(other) + return src + + +def scatter( + src: torch.Tensor, + index: torch.Tensor, + dim: int = -1, + out: Optional[torch.Tensor] = None, + dim_size: Optional[int] = None, + reduce: str = "sum", +) -> torch.Tensor: + assert reduce == "sum" # for now, TODO + index = _broadcast(index, src, dim) + if out is None: + size = list(src.size()) + if dim_size is not None: + size[dim] = dim_size + elif index.numel() == 0: + size[dim] = 0 + else: + size[dim] = int(index.max()) + 1 + out = torch.zeros( + size, + dtype=( + torch.float32 + if src.dtype not in (torch.float32, torch.float64) + else src.dtype + ), + device=src.device, + ) + return out.scatter_add_(dim, index, src.to(out.dtype)) + else: + return out.scatter_add_(dim, index, src) + + +def tp_path_exists(irreps_in1, irreps_in2, ir_out): + irreps_in1 = Irreps(irreps_in1).simplify() + irreps_in2 = Irreps(irreps_in2).simplify() + ir_out = Irrep(ir_out) + + for _, ir1 in irreps_in1: + for _, ir2 in irreps_in2: + if ir_out in ir1 * ir2: + return True + return False + + +def with_edge_vectors_( + data: AtomicDataDict.Type, + with_lengths: bool = True, + edge_index_field: str = AtomicDataDict.EDGE_INDEX_KEY, + edge_cell_shift_field: str = AtomicDataDict.EDGE_CELL_SHIFT_KEY, + edge_vec_field: str = AtomicDataDict.EDGE_VECTORS_KEY, + edge_len_field: str = AtomicDataDict.EDGE_LENGTH_KEY, +) -> AtomicDataDict.Type: + """Compute the edge displacement vectors for a graph.""" + if edge_vec_field in data: + if with_lengths and edge_len_field not in data: + data[edge_len_field] = ( + data[edge_vec_field].square().sum(1, keepdim=True).sqrt() + ) + return data + else: + # Build it dynamically + # Note that this is backwardable, because everything (pos, cell, shifts) is Tensors. + pos = data[AtomicDataDict.POSITIONS_KEY] + edge_index = data[edge_index_field] + edge_vec = torch.index_select(pos, 0, edge_index[1]) - torch.index_select( + pos, 0, edge_index[0] + ) + if AtomicDataDict.CELL_KEY in data: + # ^ note that to save time we don't check that the edge_cell_shifts are trivial if no cell is provided; we just assume they are either not present or all zero. + # NOTE: ASE cell vectors as rows convention + cell = data[AtomicDataDict.CELL_KEY] + edge_cell_shift = data[edge_cell_shift_field] + if AtomicDataDict.BATCH_KEY in data: + # treat batched cell case + edge_indexed_batches = torch.index_select( + data[AtomicDataDict.BATCH_KEY], 0, edge_index[0] + ) + # nj <- n1j <- n1j + n1i @ nij + edge_vec = torch.baddbmm( + edge_vec.view(-1, 1, 3), + edge_cell_shift.view(-1, 1, 3), + torch.index_select(cell, 0, edge_indexed_batches), + ).view(-1, 3) + # TODO: is there a more efficient way to do the above without creating an [n_edge] and [n_edge, 3, 3] tensor? + else: + # if batch key absent, we assume that cell has batch dims 1, + # so we can avoid creating the large intermediate cell tensor + # nj <- nj + ni @ ij + edge_vec = edge_vec + torch.sum( + edge_cell_shift.view(-1, 3, 1) * cell.view(3, 3), 1 + ) + data[edge_vec_field] = edge_vec + if with_lengths: + data[edge_len_field] = edge_vec.square().sum(1, keepdim=True).sqrt() + return data + + +def with_edge_type_( + data: AtomicDataDict.Type, + edge_type_field: str = AtomicDataDict.EDGE_TYPE_KEY, +) -> AtomicDataDict.Type: + """Add edge types to data if not already present.""" + if edge_type_field not in data: + edge_type = torch.index_select( + data[AtomicDataDict.ATOM_TYPE_KEY].view(-1), + 0, + data[AtomicDataDict.EDGE_INDEX_KEY].view(-1), + ).view(2, -1) + data[edge_type_field] = edge_type + return data + + +def mul_ir_to_ir_mul(x: torch.Tensor, irreps) -> torch.Tensor: + irreps = Irreps(irreps) + assert x.size(-1) == irreps.dim + + if all((mul == 1 or ir.dim == 1) for mul, ir in irreps): + return x + + base_shape = x.size()[:-1] + out_chunks = [] + for sl, (mul, ir) in zip(irreps.slices(), irreps): + chunk = x[..., sl] + if mul > 1 and ir.dim > 1: + chunk = ( + chunk.view(*base_shape, mul, ir.dim) + .transpose(-1, -2) + .contiguous() + .view(*base_shape, mul * ir.dim) + ) + out_chunks.append(chunk) + return torch.cat(out_chunks, dim=-1).contiguous() + + +def ir_mul_to_mul_ir(x: torch.Tensor, irreps) -> torch.Tensor: + irreps = Irreps(irreps) + assert x.size(-1) == irreps.dim + + if all((mul == 1 or ir.dim == 1) for mul, ir in irreps): + return x + + base_shape = x.size()[:-1] + out_chunks = [] + for sl, (mul, ir) in zip(irreps.slices(), irreps): + chunk = x[..., sl] + if mul > 1 and ir.dim > 1: + chunk = ( + chunk.view(*base_shape, ir.dim, mul) + .transpose(-1, -2) + .contiguous() + .view(*base_shape, mul * ir.dim) + ) + out_chunks.append(chunk) + return torch.cat(out_chunks, dim=-1).contiguous() diff --git a/model/package_config.py b/model/package_config.py new file mode 100644 index 0000000000000000000000000000000000000000..b069b2f4a6ad33e21dd5fcd629b54ab06eaf7d5f --- /dev/null +++ b/model/package_config.py @@ -0,0 +1,14 @@ +def get_package_data(): + package = "onescience.models.nequip" + data = { + package: [ + "LICENSE", + "NOTICE", + "SOURCE.md", + "AUTHORS.md", + "**/*.yaml", + "**/*.json", + "**/*.pt", + ], + } + return data diff --git a/single_point.py b/single_point.py new file mode 100644 index 0000000000000000000000000000000000000000..d430a00dad2c6a52732b040caf59777e06aae7c3 --- /dev/null +++ b/single_point.py @@ -0,0 +1,189 @@ +"""Run one NequIP energy, force, and stress prediction through ASE. + +This script follows the official NequIP ASE integration style: +https://nequip.readthedocs.io/en/latest/integrations/ase.html +""" + +from __future__ import annotations + +import argparse +import json +import os +import warnings +from pathlib import Path +from typing import Any, Dict + +warnings.filterwarnings("ignore", category=FutureWarning, module="e3nn") + +from ase.build import bulk +from ase.io import read + +from onescience.models.nequip.model import ModelTypeNamesFromPackage +from onescience.models.nequip.model.nequip_models import NequIPGNNModel +from onescience.utils.nequip.internal.global_state import set_global_state +from onescience.utils.nequip import build_nequip_calculator + + +def default_paths() -> Dict[str, str | None]: + """Return default compiled model / checkpoint paths if env var is set.""" + models_dir = os.environ.get("ONESCIENCE_MODELS_DIR") + if not models_dir: + return {"compiled_model": None, "checkpoint": None} + nequip_dir = Path(models_dir) / "NequIP" + return { + "compiled_model": str(nequip_dir / "NequIP-OAM-L-0.1.nequip.pth"), + "checkpoint": None, + } + + +def resolve_model_paths( + compiled_model: str | None, checkpoint: str | None +) -> Dict[str, str | None]: + """Prefer an explicitly selected model source over environment defaults.""" + if compiled_model or checkpoint: + return {"compiled_model": compiled_model, "checkpoint": checkpoint} + return default_paths() + + +def load_structure(path: str | None, index: int): + """Load an ASE structure or use the built-in Cu bulk example.""" + if path: + return read(path, index=index) + return bulk("Cu") + + +def write_workflow_result(result: Dict[str, Any], output_path: str) -> str: + """Write a workflow result dictionary to a JSON file.""" + output = Path(output_path) + output.parent.mkdir(parents=True, exist_ok=True) + with open(output, "w", encoding="utf-8") as f: + json.dump(result, f, indent=2, ensure_ascii=False) + return str(output) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + group = parser.add_mutually_exclusive_group() + group.add_argument( + "--compiled-model", + help="Path to a compiled NequIP model (.nequip.pth or .nequip.pt2).", + ) + group.add_argument( + "--checkpoint", + help="Path to a NequIP checkpoint (.ckpt) or packaged model (.nequip.zip).", + ) + group.add_argument( + "--demo", + action="store_true", + help="Use a small built-in demo model instead of a real checkpoint.", + ) + parser.add_argument( + "--package", + help=( + "Original .nequip.zip package for a fine-tuned checkpoint; its atom " + "types are read automatically." + ), + ) + parser.add_argument( + "--input", + help=( + "CIF, POSCAR, XYZ, trajectory, or another ASE-readable structure; " + "defaults to the built-in periodic Cu example" + ), + ) + parser.add_argument( + "--index", + type=int, + default=0, + help="Zero-based frame index for trajectory inputs (default: 0).", + ) + parser.add_argument("--device", default="cuda") + parser.add_argument("--output", default="outputs/single_point.json") + parser.add_argument( + "--model-type-names", + nargs="+", + default=["C", "H", "O", "Cu"], + help="Chemical species the model knows about (used for demo/checkpoint).", + ) + parser.add_argument( + "--r-max", + type=float, + default=4.0, + help="Neighbor-list cutoff in Angstrom (used for demo/checkpoint models).", + ) + args = parser.parse_args() + + for label, path in ( + ("compiled model", args.compiled_model), + ("checkpoint", args.checkpoint), + ("package", args.package), + ): + if path and not Path(path).expanduser().is_file(): + parser.error(f"{label} not found: {path}") + + model_paths = resolve_model_paths(args.compiled_model, args.checkpoint) + compiled_model = model_paths["compiled_model"] + checkpoint = model_paths["checkpoint"] + if args.package and not checkpoint: + parser.error("--package requires --checkpoint") + + model_type_names = list(args.model_type_names) + package_for_types = args.package + if package_for_types is None and checkpoint and checkpoint.endswith(".nequip.zip"): + package_for_types = checkpoint + if package_for_types: + model_type_names = list(ModelTypeNamesFromPackage(package_for_types)) + + atoms = load_structure(args.input, args.index) + calc_kwargs: Dict[str, Any] = {"device": args.device} + + if args.demo: + set_global_state() + calc_kwargs["model"] = NequIPGNNModel( + seed=123, + model_dtype="float32", + type_names=model_type_names, + num_layers=2, + l_max=1, + num_features=32, + r_max=args.r_max, + parity=False, + avg_num_neighbors=10.0, + ) + elif compiled_model and Path(compiled_model).exists(): + calc_kwargs["compiled_model"] = compiled_model + elif checkpoint and Path(checkpoint).exists(): + calc_kwargs["checkpoint"] = checkpoint + calc_kwargs["model_type_names"] = model_type_names + else: + parser.error( + "no model found; pass --compiled-model, --checkpoint, or --demo" + ) + + atoms.calc = build_nequip_calculator(**calc_kwargs) + + result = { + "formula": atoms.get_chemical_formula(), + "natoms": len(atoms), + "input": str(Path(args.input).expanduser()) if args.input else None, + "input_index": args.index if args.input else None, + "input_source": args.input or "ASE bulk Cu default", + "compiled_model": str(Path(compiled_model).expanduser()) if compiled_model else None, + "checkpoint": str(Path(checkpoint).expanduser()) if checkpoint else None, + "package": str(Path(package_for_types).expanduser()) if package_for_types else None, + "pbc": atoms.pbc.tolist(), + "cell_angstrom": atoms.cell.array.tolist(), + "energy_ev": float(atoms.get_potential_energy()), + "forces_ev_per_angstrom": atoms.get_forces().tolist(), + "stress_ev_per_angstrom_cubed_voigt": atoms.get_stress().tolist(), + } + output = write_workflow_result(result, args.output) + + print("formula:", result["formula"]) + print("atoms:", result["natoms"]) + print("energy (eV):", result["energy_ev"]) + print("result:", output) + + +if __name__ == "__main__": + main() diff --git a/structure_relaxation.py b/structure_relaxation.py new file mode 100644 index 0000000000000000000000000000000000000000..059fa85f997c739a46dd61ecf1757b0413f027f0 --- /dev/null +++ b/structure_relaxation.py @@ -0,0 +1,275 @@ +"""Relax an ASE structure with a compiled NequIP model. + +This follows the official NequIP ASE relaxation example: it supports atomic +and cell relaxation, tracks forces at every ionic step, and aborts exploding +relaxations before they can hang indefinitely. +""" + +from __future__ import annotations + +import argparse +import json +import os +from pathlib import Path +from typing import Any + +import numpy as np +import torch +from ase import Atoms +from ase.build import bulk +from ase.filters import ExpCellFilter, FrechetCellFilter +from ase.io import read, write +import ase.optimize as opt + +from onescience.utils.nequip.integrations.ase import NequIPCalculator + + +OPTIMIZERS = { + "BFGS": opt.BFGS, + "BFGSLineSearch": opt.BFGSLineSearch, + "FIRE": opt.FIRE, + "FIRE2": opt.FIRE2, + "GOQN": opt.GoodOldQuasiNewton, + "GPMin": opt.GPMin, + "LBFGS": opt.LBFGS, + "LBFGSLineSearch": opt.LBFGSLineSearch, +} +CELL_FILTERS = { + "exp": ExpCellFilter, + "frechet": FrechetCellFilter, +} + + +def default_compiled_model() -> str | None: + models_dir = os.environ.get("ONESCIENCE_MODELS_DIR") + if not models_dir: + return None + return str(Path(models_dir) / "NequIP" / "NequIP-OAM-L-0.1.nequip.pth") + + +def load_structure( + input_path: str | None, + index: int, + element: str, + crystal_structure: str, + lattice_constant: float, + displacement: float, +) -> tuple[Atoms, str]: + if input_path: + path = Path(input_path).expanduser().resolve() + if not path.is_file(): + raise FileNotFoundError(f"input structure not found: {path}") + return read(path, index=index), f"{path}[{index}]" + + atoms = bulk( + element, + crystalstructure=crystal_structure, + a=lattice_constant, + cubic=True, + ) + if displacement: + atoms.positions[0, 0] += displacement + return atoms, ( + f"ASE bulk {element} {crystal_structure}, a={lattice_constant} Angstrom, " + f"atom-0 displacement={displacement} Angstrom" + ) + + +def _max_vector_norm(values: np.ndarray) -> float: + array = np.asarray(values) + if array.size == 0: + return 0.0 + return float(np.linalg.norm(array.reshape(-1, 3), axis=1).max()) + + +def relaxation_snapshot(atoms: Atoms, target: Any, step: int) -> dict[str, Any]: + forces = atoms.get_forces() + optimizer_forces = target.get_forces() + stress = atoms.get_stress() + return { + "step": step, + "energy_ev": float(atoms.get_potential_energy()), + "energy_ev_per_atom": float(atoms.get_potential_energy() / len(atoms)), + "volume_angstrom3": float(atoms.get_volume()), + "max_atomic_force_ev_per_angstrom": _max_vector_norm(forces), + "max_optimizer_force": _max_vector_norm(optimizer_forces), + "stress_ev_per_angstrom3_voigt": np.asarray(stress).tolist(), + "max_abs_stress_ev_per_angstrom3": float(np.abs(stress).max()), + } + + +def relax_structure( + atoms: Atoms, + *, + optimizer_name: str, + cell_filter_name: str, + fixed_cell: bool, + fmax: float, + steps: int, + force_limit: float, + logfile: Path, + trajectory: Path, +) -> tuple[bool, int, list[dict[str, Any]]]: + if not fixed_cell and not atoms.pbc.all(): + raise ValueError("cell relaxation requires periodic boundaries; use --fixed-cell") + + target = atoms if fixed_cell else CELL_FILTERS[cell_filter_name](atoms) + optimizer_cls = OPTIMIZERS[optimizer_name] + history: list[dict[str, Any]] = [] + converged = False + + with optimizer_cls( + target, + logfile=str(logfile), + trajectory=str(trajectory), + ) as optimizer: + for converged in optimizer.irun(fmax=fmax, steps=steps): + snapshot = relaxation_snapshot(atoms, target, optimizer.nsteps) + history.append(snapshot) + if max( + snapshot["max_atomic_force_ev_per_angstrom"], + snapshot["max_optimizer_force"], + ) > force_limit: + raise RuntimeError( + f"relaxation force exceeded safety limit {force_limit:g}" + ) + + return bool(converged), int(optimizer.nsteps), history + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--compiled-model", default=default_compiled_model()) + parser.add_argument( + "--input", + help="CIF, POSCAR, XYZ, trajectory, or another ASE-readable structure", + ) + parser.add_argument("--index", type=int, default=0) + parser.add_argument("--device", default="cuda") + parser.add_argument("--optimizer", choices=sorted(OPTIMIZERS), default="GOQN") + parser.add_argument( + "--cell-filter", choices=sorted(CELL_FILTERS), default="frechet" + ) + parser.add_argument( + "--fixed-cell", + action="store_true", + help="relax atomic positions only; the default also relaxes the cell", + ) + parser.add_argument("--fmax", type=float, default=0.05) + parser.add_argument("--steps", type=int, default=500) + parser.add_argument("--force-limit", type=float, default=1.0e6) + parser.add_argument("--element", default="Si") + parser.add_argument("--crystal-structure", default="diamond") + parser.add_argument("--lattice-constant", type=float, default=5.65) + parser.add_argument("--displacement", type=float, default=0.08) + parser.add_argument("--output-dir", default="outputs/structure_relaxation") + parser.add_argument("--output-structure", default="relaxed.cif") + parser.add_argument("--result", default="result.json") + parser.add_argument("--trajectory", default="relax.traj") + parser.add_argument("--log", default="relax.log") + args = parser.parse_args() + + if not args.compiled_model: + parser.error("--compiled-model is required when ONESCIENCE_MODELS_DIR is unset") + compiled_model = Path(args.compiled_model).expanduser().resolve() + if not compiled_model.is_file(): + parser.error(f"compiled model not found: {compiled_model}") + if args.fmax <= 0: + parser.error("--fmax must be positive") + if args.steps < 1: + parser.error("--steps must be positive") + if args.force_limit <= 0: + parser.error("--force-limit must be positive") + + output_dir = Path(args.output_dir).expanduser().resolve() + output_dir.mkdir(parents=True, exist_ok=True) + output_structure = output_dir / args.output_structure + result_path = output_dir / args.result + trajectory_path = output_dir / args.trajectory + log_path = output_dir / args.log + + try: + atoms, input_source = load_structure( + args.input, + args.index, + args.element, + args.crystal_structure, + args.lattice_constant, + args.displacement, + ) + except (FileNotFoundError, IndexError, ValueError) as error: + parser.error(str(error)) + if len(atoms) == 0: + parser.error("input structure has no atoms") + + species = sorted(set(atoms.get_chemical_symbols())) + atoms.calc = NequIPCalculator.from_compiled_model( + compile_path=str(compiled_model), + chemical_species_to_atom_type_map={symbol: symbol for symbol in species}, + device=args.device, + ) + + try: + converged, nsteps, history = relax_structure( + atoms, + optimizer_name=args.optimizer, + cell_filter_name=args.cell_filter, + fixed_cell=args.fixed_cell, + fmax=args.fmax, + steps=args.steps, + force_limit=args.force_limit, + logfile=log_path, + trajectory=trajectory_path, + ) + except ValueError as error: + parser.error(str(error)) + + write(output_structure, atoms) + result = { + "compiled_model": str(compiled_model), + "device": args.device, + "device_name": torch.cuda.get_device_name(0) + if args.device.startswith("cuda") and torch.cuda.is_available() + else "cpu", + "input_source": input_source, + "formula": atoms.get_chemical_formula(), + "num_atoms": len(atoms), + "chemical_species_to_atom_type_map": { + symbol: symbol for symbol in species + }, + "optimizer": args.optimizer, + "cell_filter": None if args.fixed_cell else args.cell_filter, + "fixed_cell": args.fixed_cell, + "fmax_ev_per_angstrom": args.fmax, + "max_steps": args.steps, + "force_safety_limit": args.force_limit, + "converged": converged, + "steps": nsteps, + "initial": history[0], + "final": history[-1], + "energy_change_ev": history[-1]["energy_ev"] - history[0]["energy_ev"], + "volume_change_angstrom3": ( + history[-1]["volume_angstrom3"] - history[0]["volume_angstrom3"] + ), + "history": history, + "relaxed_structure": str(output_structure), + "trajectory": str(trajectory_path), + "log": str(log_path), + } + result_path.write_text(json.dumps(result, indent=2) + "\n", encoding="utf-8") + + print("formula:", result["formula"]) + print("atoms:", result["num_atoms"]) + print("converged:", converged) + print("steps:", nsteps) + print("initial energy (eV):", result["initial"]["energy_ev"]) + print("final energy (eV):", result["final"]["energy_ev"]) + print( + "final max force (eV/Angstrom):", + result["final"]["max_atomic_force_ev_per_angstrom"], + ) + print("result:", result_path) + + +if __name__ == "__main__": + main() diff --git a/train.py b/train.py new file mode 100644 index 0000000000000000000000000000000000000000..0f6c966350a37c963b9f1786d543adef16960090 --- /dev/null +++ b/train.py @@ -0,0 +1,17 @@ +"""NequIP training entry point for OneScience. + +This is a thin wrapper around ``onescience.utils.nequip.cli.train``. OneScience +must already be installed in the active MatChem environment. +""" + +from __future__ import annotations + +import warnings + +warnings.filterwarnings("ignore", category=FutureWarning, module="e3nn") + +from onescience.utils.nequip.cli.train import main + + +if __name__ == "__main__": + main() diff --git a/weight/NequIP-OAM-L-0.1.nequip.pth b/weight/NequIP-OAM-L-0.1.nequip.pth new file mode 100644 index 0000000000000000000000000000000000000000..d94f7718b8be05a730436bc00d9ee687dccae01f --- /dev/null +++ b/weight/NequIP-OAM-L-0.1.nequip.pth @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e83a1d656f8b19b55d2f05708c83e054612f713e9a1b06266aa010db58e56517 +size 39270723 diff --git a/weight/NequIP-OAM-L-0.1.nequip.zip b/weight/NequIP-OAM-L-0.1.nequip.zip new file mode 100644 index 0000000000000000000000000000000000000000..032856a6d1f52253490344eaf2efb81b50c89c4f --- /dev/null +++ b/weight/NequIP-OAM-L-0.1.nequip.zip @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5d01a4fab228abb3cdb6ace0033f93993729956bca6a42234a2a8816825b9a0f +size 78464590