From 47991000614d91b21132b9a7b40f781de0306a59 Mon Sep 17 00:00:00 2001 From: yeze5421 <81737316+yeze5421@users.noreply.github.com> Date: Sun, 12 Apr 2026 14:23:07 +0800 Subject: [PATCH] Build production-style MLIPX project with PaiNN training stack --- .gitignore | 4 + README.md | 129 ++++++++++++++ configs/minimal.yaml | 34 ++++ configs/standard.yaml | 34 ++++ mlip.py | 288 -------------------------------- pyproject.toml | 28 ++++ scripts/evaluate.sh | 5 + scripts/infer.sh | 6 + scripts/make_demo_dataset.py | 39 +++++ scripts/package.sh | 8 + scripts/train.sh | 3 + src/mlipx/__init__.py | 3 + src/mlipx/ase_ext/__init__.py | 0 src/mlipx/ase_ext/calculator.py | 75 +++++++++ src/mlipx/config.py | 76 +++++++++ src/mlipx/data/__init__.py | 0 src/mlipx/data/datamodule.py | 122 ++++++++++++++ src/mlipx/data/dataset.py | 63 +++++++ src/mlipx/data/graph.py | 52 ++++++ src/mlipx/data/types.py | 18 ++ src/mlipx/evaluate.py | 86 ++++++++++ src/mlipx/infer.py | 39 +++++ src/mlipx/models/__init__.py | 0 src/mlipx/models/layers.py | 28 ++++ src/mlipx/models/model.py | 49 ++++++ src/mlipx/models/painn.py | 78 +++++++++ src/mlipx/train.py | 50 ++++++ src/mlipx/training/__init__.py | 0 src/mlipx/training/engine.py | 148 ++++++++++++++++ src/mlipx/utils/__init__.py | 0 src/mlipx/utils/metrics.py | 19 +++ src/mlipx/utils/seed.py | 13 ++ tests/conftest.py | 8 + tests/test_data_pipeline.py | 50 ++++++ tests/test_model_and_train.py | 87 ++++++++++ 35 files changed, 1354 insertions(+), 288 deletions(-) create mode 100644 .gitignore create mode 100644 README.md create mode 100644 configs/minimal.yaml create mode 100644 configs/standard.yaml delete mode 100644 mlip.py create mode 100644 pyproject.toml create mode 100755 scripts/evaluate.sh create mode 100755 scripts/infer.sh create mode 100755 scripts/make_demo_dataset.py create mode 100755 scripts/package.sh create mode 100755 scripts/train.sh create mode 100644 src/mlipx/__init__.py create mode 100644 src/mlipx/ase_ext/__init__.py create mode 100644 src/mlipx/ase_ext/calculator.py create mode 100644 src/mlipx/config.py create mode 100644 src/mlipx/data/__init__.py create mode 100644 src/mlipx/data/datamodule.py create mode 100644 src/mlipx/data/dataset.py create mode 100644 src/mlipx/data/graph.py create mode 100644 src/mlipx/data/types.py create mode 100644 src/mlipx/evaluate.py create mode 100644 src/mlipx/infer.py create mode 100644 src/mlipx/models/__init__.py create mode 100644 src/mlipx/models/layers.py create mode 100644 src/mlipx/models/model.py create mode 100644 src/mlipx/models/painn.py create mode 100644 src/mlipx/train.py create mode 100644 src/mlipx/training/__init__.py create mode 100644 src/mlipx/training/engine.py create mode 100644 src/mlipx/utils/__init__.py create mode 100644 src/mlipx/utils/metrics.py create mode 100644 src/mlipx/utils/seed.py create mode 100644 tests/conftest.py create mode 100644 tests/test_data_pipeline.py create mode 100644 tests/test_model_and_train.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..77bf3ad --- /dev/null +++ b/.gitignore @@ -0,0 +1,4 @@ +__pycache__/ +*.pyc +.pytest_cache/ +*.zip diff --git a/README.md b/README.md new file mode 100644 index 0000000..2d0dfbe --- /dev/null +++ b/README.md @@ -0,0 +1,129 @@ +# MLIPX: PaiNN-style 机器学习原子势工程项目 + +MLIPX 是一个可训练、可评估、可推理、可扩展的原子势项目,面向 ASE/extxyz 数据,支持能量+力联合训练,默认采用 **PaiNN-style message passing**(标量+向量通道,支持从能量自动微分得到力)。 + +## 1. 项目结构 + +```text +. +├── configs/ +│ ├── minimal.yaml +│ └── standard.yaml +├── data/ +├── scripts/ +│ ├── make_demo_dataset.py +│ ├── train.sh +│ ├── evaluate.sh +│ └── infer.sh +├── src/mlipx/ +│ ├── config.py +│ ├── train.py +│ ├── evaluate.py +│ ├── infer.py +│ ├── data/ +│ │ ├── dataset.py +│ │ ├── datamodule.py +│ │ ├── graph.py +│ │ └── types.py +│ ├── models/ +│ │ ├── layers.py +│ │ ├── painn.py +│ │ └── model.py +│ ├── training/engine.py +│ ├── ase_ext/calculator.py +│ └── utils/ +└── tests/ +``` + +## 2. 安装 + +```bash +python -m venv .venv +source .venv/bin/activate +pip install -e .[dev] +``` + +## 3. 数据格式(ASE/extxyz) + +每个结构需要至少包含: +- 原子序数(ASE 自动包含) +- positions +- `info['energy']`(总能) +- `arrays['forces']` +- cell / pbc(若有) + +可先生成可跑通数据: + +```bash +python scripts/make_demo_dataset.py --out data/demo.extxyz --n-samples 80 +``` + +## 4. 训练 + +最小实验: + +```bash +python -m mlipx.train --config configs/minimal.yaml +# 或 bash scripts/train.sh configs/minimal.yaml +``` + +正式实验建议: + +```bash +python -m mlipx.train --config configs/standard.yaml +``` + +可命令行覆盖超参数: + +```bash +python -m mlipx.train --config configs/minimal.yaml --epochs 10 --batch-size 8 --lr 3e-4 +``` + +训练输出: +- `artifacts/.../best.pt` +- `artifacts/.../last.pt` +- `artifacts/.../history.json` + +## 5. 评估 + +```bash +python -m mlipx.evaluate \ + --config configs/minimal.yaml \ + --checkpoint artifacts/minimal/best.pt \ + --split test \ + --out artifacts/test_predictions.csv +``` + +输出指标:Energy/Force 的 MAE 与 RMSE。 + +## 6. 推理 + +```bash +python -m mlipx.infer \ + --config configs/minimal.yaml \ + --checkpoint artifacts/minimal/best.pt \ + --input data/demo.extxyz \ + --output artifacts/infer.txt +``` + +## 7. ASE Calculator 接入 + +```python +from mlipx.config import load_config +from mlipx.ase_ext.calculator import build_calculator +from ase.io import read + +cfg = load_config("configs/minimal.yaml") +calc = build_calculator(cfg, "artifacts/minimal/best.pt") +atoms = read("data/demo.extxyz", index=0) +atoms.calc = calc +print(atoms.get_potential_energy()) +print(atoms.get_forces()) +``` + +## 8. 当前实现说明 + +- 使用 cutoff 邻域图,处理不同原子数 batch。 +- 使用 energy normalization(可选 per-atom)。 +- 支持 AdamW、CosineLR、gradient clipping、AMP(CUDA时)、early stopping、checkpoint resume。 +- stress/virial 目前未启用(可在后续版本加入基于应变导数的实现)。 diff --git a/configs/minimal.yaml b/configs/minimal.yaml new file mode 100644 index 0000000..4efb80e --- /dev/null +++ b/configs/minimal.yaml @@ -0,0 +1,34 @@ +name: mlipx_minimal +data: + path: data/demo.extxyz + format: extxyz + energy_key: energy + force_key: forces + cutoff: 5.0 + max_neighbors: 32 + train_ratio: 0.8 + val_ratio: 0.1 + test_ratio: 0.1 + batch_size: 4 + num_workers: 0 + normalize_energy_per_atom: true +model: + hidden_dim: 64 + n_interactions: 3 + n_rbf: 24 + cutoff: 5.0 +train: + seed: 42 + device: auto + epochs: 3 + lr: 0.0005 + weight_decay: 1.0e-6 + energy_weight: 1.0 + force_weight: 20.0 + grad_clip: 5.0 + amp: false + scheduler: cosine + scheduler_tmax: 3 + early_stopping_patience: 10 + output_dir: artifacts/minimal + resume_checkpoint: null diff --git a/configs/standard.yaml b/configs/standard.yaml new file mode 100644 index 0000000..658a80c --- /dev/null +++ b/configs/standard.yaml @@ -0,0 +1,34 @@ +name: mlipx_standard +data: + path: data/demo.extxyz + format: extxyz + energy_key: energy + force_key: forces + cutoff: 5.0 + max_neighbors: 64 + train_ratio: 0.8 + val_ratio: 0.1 + test_ratio: 0.1 + batch_size: 16 + num_workers: 2 + normalize_energy_per_atom: true +model: + hidden_dim: 128 + n_interactions: 4 + n_rbf: 32 + cutoff: 5.0 +train: + seed: 42 + device: auto + epochs: 50 + lr: 0.0002 + weight_decay: 1.0e-6 + energy_weight: 1.0 + force_weight: 50.0 + grad_clip: 5.0 + amp: true + scheduler: cosine + scheduler_tmax: 50 + early_stopping_patience: 12 + output_dir: artifacts/standard + resume_checkpoint: null diff --git a/mlip.py b/mlip.py deleted file mode 100644 index 6cac4e0..0000000 --- a/mlip.py +++ /dev/null @@ -1,288 +0,0 @@ -import math -import random -from dataclasses import dataclass -from typing import List, Tuple - -import torch -import torch.nn as nn -import torch.optim as optim - -# ============================================================ -# Tiny MLIP from scratch (toy version) -# ------------------------------------------------------------ -# What this file does: -# 1. Generate a synthetic atomic dataset using Lennard-Jones energy -# 2. Build simple radial descriptors for each atom -# 3. Predict total energy with a small neural network -# 4. Obtain forces by automatic differentiation -# -# This is a teaching demo, not a production MLIP. -# It is intentionally small and readable. -# ============================================================ - -# ------------------------ -# Reproducibility -# ------------------------ -random.seed(42) -torch.manual_seed(42) - -DEVICE = "cuda" if torch.cuda.is_available() else "cpu" -DTYPE = torch.float32 - - -@dataclass -class Config: - n_atoms: int = 5 - box_size: float = 4.0 - n_samples: int = 800 - train_ratio: float = 0.8 - cutoff: float = 3.0 - sigma_values: Tuple[float, ...] = (0.4, 0.7, 1.0, 1.3, 1.6) - hidden_dim: int = 64 - batch_size: int = 32 - lr: float = 1e-3 - epochs: int = 80 - lj_epsilon: float = 0.5 - lj_sigma: float = 1.0 - min_dist: float = 0.85 # avoid atom overlap - - -CFG = Config() - - -# ------------------------ -# Physics: Lennard-Jones toy label -# ------------------------ -def pairwise_vectors(positions: torch.Tensor) -> torch.Tensor: - # positions: [N, 3] - return positions[:, None, :] - positions[None, :, :] - - -def pairwise_distances(positions: torch.Tensor) -> torch.Tensor: - rij = pairwise_vectors(positions) - dij = torch.sqrt(torch.sum(rij**2, dim=-1) + 1e-12) - return dij - - -def smooth_cutoff(r: torch.Tensor, rc: float) -> torch.Tensor: - # cosine cutoff - x = 0.5 * (torch.cos(math.pi * r / rc) + 1.0) - return torch.where(r < rc, x, torch.zeros_like(r)) - - -def lennard_jones_energy(positions: torch.Tensor, epsilon: float, sigma: float, cutoff: float) -> torch.Tensor: - # total energy for one structure - d = pairwise_distances(positions) - mask = torch.triu(torch.ones_like(d), diagonal=1) > 0 - rij = d[mask] - valid = rij < cutoff - rij = rij[valid] - sr6 = (sigma / rij) ** 6 - sr12 = sr6**2 - e_pair = 4.0 * epsilon * (sr12 - sr6) - return torch.sum(e_pair) - - -def sample_structure(n_atoms: int, box_size: float, min_dist: float) -> torch.Tensor: - pts: List[torch.Tensor] = [] - max_trials = 5000 - trials = 0 - while len(pts) < n_atoms and trials < max_trials: - cand = torch.rand(3) * box_size - if len(pts) == 0: - pts.append(cand) - else: - ok = True - for p in pts: - if torch.norm(cand - p) < min_dist: - ok = False - break - if ok: - pts.append(cand) - trials += 1 - if len(pts) < n_atoms: - raise RuntimeError("Could not sample a valid structure. Try reducing min_dist.") - return torch.stack(pts, dim=0) - - -# ------------------------ -# Descriptor: simple radial Gaussian sums -# ------------------------ -def atomic_descriptors(positions: torch.Tensor, cutoff: float, sigma_values: Tuple[float, ...]) -> torch.Tensor: - # output: [N, n_features] - d = pairwise_distances(positions) # [N, N] - n_atoms = d.shape[0] - eye = torch.eye(n_atoms, device=d.device, dtype=torch.bool) - desc_list = [] - - for center in sigma_values: - # Gaussian in distance, summed over neighbors - g = torch.exp(-((d - center) ** 2) / (2.0 * 0.25**2)) * smooth_cutoff(d, cutoff) - g = torch.where(eye, torch.zeros_like(g), g) - desc_list.append(torch.sum(g, dim=1, keepdim=True)) - - # also add inverse-distance-like feature - inv_d = torch.where(eye, torch.zeros_like(d), 1.0 / (d + 1e-6)) - inv_d = inv_d * smooth_cutoff(d, cutoff) - desc_list.append(torch.sum(inv_d, dim=1, keepdim=True)) - - return torch.cat(desc_list, dim=1) - - -# ------------------------ -# Dataset -# ------------------------ -class ToyAtomicDataset(torch.utils.data.Dataset): - def __init__(self, cfg: Config): - self.items = [] - for _ in range(cfg.n_samples): - pos = sample_structure(cfg.n_atoms, cfg.box_size, cfg.min_dist) - energy = lennard_jones_energy(pos, cfg.lj_epsilon, cfg.lj_sigma, cfg.cutoff) - self.items.append((pos, energy.unsqueeze(0))) - - def __len__(self): - return len(self.items) - - def __getitem__(self, idx): - return self.items[idx] - - -def collate_fn(batch): - positions = torch.stack([item[0] for item in batch], dim=0) # [B, N, 3] - energies = torch.stack([item[1] for item in batch], dim=0) # [B, 1] - return positions, energies - - -# ------------------------ -# Tiny MLIP model -# ------------------------ -class AtomicNetwork(nn.Module): - def __init__(self, in_dim: int, hidden_dim: int): - super().__init__() - self.net = nn.Sequential( - nn.Linear(in_dim, hidden_dim), - nn.SiLU(), - nn.Linear(hidden_dim, hidden_dim), - nn.SiLU(), - nn.Linear(hidden_dim, 1), - ) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - # x: [B, N, F] - atomic_e = self.net(x) # [B, N, 1] - total_e = torch.sum(atomic_e, dim=1) # [B, 1] - return total_e - - -class TinyMLIP(nn.Module): - def __init__(self, cfg: Config): - super().__init__() - self.cfg = cfg - in_dim = len(cfg.sigma_values) + 1 - self.atomic_net = AtomicNetwork(in_dim, cfg.hidden_dim) - - def descriptors(self, positions: torch.Tensor) -> torch.Tensor: - # positions: [B, N, 3] - feats = [] - for b in range(positions.shape[0]): - feat_b = atomic_descriptors(positions[b], self.cfg.cutoff, self.cfg.sigma_values) - feats.append(feat_b) - return torch.stack(feats, dim=0) # [B, N, F] - - def forward(self, positions: torch.Tensor) -> torch.Tensor: - feats = self.descriptors(positions) - return self.atomic_net(feats) - - def predict_energy_forces(self, positions: torch.Tensor): - # positions: [B, N, 3] - positions = positions.clone().detach().requires_grad_(True) - energy = self.forward(positions) - forces = -torch.autograd.grad(energy.sum(), positions, create_graph=False)[0] - return energy, forces - - -# ------------------------ -# Training and evaluation -# ------------------------ -def split_dataset(dataset, train_ratio=0.8): - n_total = len(dataset) - n_train = int(n_total * train_ratio) - n_val = n_total - n_train - return torch.utils.data.random_split(dataset, [n_train, n_val]) - - -def mae(x: torch.Tensor, y: torch.Tensor) -> float: - return torch.mean(torch.abs(x - y)).item() - - -def train_model(cfg: Config): - dataset = ToyAtomicDataset(cfg) - train_set, val_set = split_dataset(dataset, cfg.train_ratio) - - train_loader = torch.utils.data.DataLoader( - train_set, batch_size=cfg.batch_size, shuffle=True, collate_fn=collate_fn - ) - val_loader = torch.utils.data.DataLoader( - val_set, batch_size=cfg.batch_size, shuffle=False, collate_fn=collate_fn - ) - - model = TinyMLIP(cfg).to(DEVICE) - optimizer = optim.Adam(model.parameters(), lr=cfg.lr) - criterion = nn.MSELoss() - - for epoch in range(1, cfg.epochs + 1): - model.train() - train_loss = 0.0 - for positions, energies in train_loader: - positions = positions.to(DEVICE, dtype=DTYPE) - energies = energies.to(DEVICE, dtype=DTYPE) - - pred = model(positions) - loss = criterion(pred, energies) - - optimizer.zero_grad() - loss.backward() - optimizer.step() - train_loss += loss.item() * positions.size(0) - - train_loss /= len(train_loader.dataset) - - model.eval() - val_loss = 0.0 - val_mae = 0.0 - with torch.no_grad(): - for positions, energies in val_loader: - positions = positions.to(DEVICE, dtype=DTYPE) - energies = energies.to(DEVICE, dtype=DTYPE) - pred = model(positions) - loss = criterion(pred, energies) - val_loss += loss.item() * positions.size(0) - val_mae += torch.sum(torch.abs(pred - energies)).item() - - val_loss /= len(val_loader.dataset) - val_mae /= len(val_loader.dataset) - - if epoch % 10 == 0 or epoch == 1: - print(f"Epoch {epoch:03d} | train_loss={train_loss:.6f} | val_loss={val_loss:.6f} | val_MAE={val_mae:.6f}") - - return model, val_set - - -def demo_prediction(model: TinyMLIP, sample_item): - positions, true_energy = sample_item - positions = positions.unsqueeze(0).to(DEVICE, dtype=DTYPE) - true_energy = true_energy.to(DEVICE, dtype=DTYPE) - - pred_energy, forces = model.predict_energy_forces(positions) - - print("\n===== Demo prediction =====") - print(f"True energy : {true_energy.item():.6f}") - print(f"Pred energy : {pred_energy.item():.6f}") - print(f"Forces shape: {tuple(forces.shape)}") - print("First atom force:", forces[0, 0].detach().cpu().numpy()) - - -if __name__ == "__main__": - print(f"Using device: {DEVICE}") - model, val_set = train_model(CFG) - demo_prediction(model, val_set[0]) diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..de85103 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,28 @@ +[build-system] +requires = ["setuptools", "wheel"] +build-backend = "setuptools.build_meta" + +[project] +name = "mlipx" +version = "0.2.0" +description = "Research-engineering MLIP project with PaiNN-style architecture" +readme = "README.md" +requires-python = ">=3.10" +dependencies = [ + "torch>=2.1", + "ase>=3.23", + "numpy>=1.24", + "pyyaml>=6.0", +] + +[project.optional-dependencies] +dev = [ + "pytest>=8.0", +] + +[tool.setuptools] +package-dir = {"" = "src"} + +[tool.setuptools.packages.find] +where = ["src"] +include = ["mlipx*"] diff --git a/scripts/evaluate.sh b/scripts/evaluate.sh new file mode 100755 index 0000000..34cbf4b --- /dev/null +++ b/scripts/evaluate.sh @@ -0,0 +1,5 @@ +#!/usr/bin/env bash +set -euo pipefail +CFG=${1:-configs/minimal.yaml} +CKPT=${2:-artifacts/minimal/best.pt} +python -m mlipx.evaluate --config "$CFG" --checkpoint "$CKPT" --split test --out artifacts/test_predictions.csv diff --git a/scripts/infer.sh b/scripts/infer.sh new file mode 100755 index 0000000..907e123 --- /dev/null +++ b/scripts/infer.sh @@ -0,0 +1,6 @@ +#!/usr/bin/env bash +set -euo pipefail +CFG=${1:-configs/minimal.yaml} +CKPT=${2:-artifacts/minimal/best.pt} +INP=${3:-data/demo.extxyz} +python -m mlipx.infer --config "$CFG" --checkpoint "$CKPT" --input "$INP" --output artifacts/infer.txt diff --git a/scripts/make_demo_dataset.py b/scripts/make_demo_dataset.py new file mode 100755 index 0000000..488cc4f --- /dev/null +++ b/scripts/make_demo_dataset.py @@ -0,0 +1,39 @@ +#!/usr/bin/env python +from __future__ import annotations + +import argparse +import random + +import numpy as np +from ase import Atoms +from ase.calculators.emt import EMT +from ase.io import write + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--out", default="data/demo.extxyz") + p.add_argument("--n-samples", type=int, default=80) + p.add_argument("--seed", type=int, default=42) + args = p.parse_args() + + random.seed(args.seed) + np.random.seed(args.seed) + + frames = [] + for _ in range(args.n_samples): + n = random.randint(2, 6) + symbols = random.choices(["H", "C", "N", "O"], k=n) + pos = np.random.uniform(0.0, 3.5, size=(n, 3)) + atoms = Atoms(symbols=symbols, positions=pos, cell=np.eye(3) * 8.0, pbc=False) + atoms.calc = EMT() + atoms.info["energy"] = float(atoms.get_potential_energy()) + atoms.arrays["forces"] = atoms.get_forces() + frames.append(atoms) + + write(args.out, frames, format="extxyz") + print(f"wrote {len(frames)} samples to {args.out}") + + +if __name__ == "__main__": + main() diff --git a/scripts/package.sh b/scripts/package.sh new file mode 100755 index 0000000..38d2121 --- /dev/null +++ b/scripts/package.sh @@ -0,0 +1,8 @@ +#!/usr/bin/env bash +set -euo pipefail + +OUT=${1:-mlip_project.zip} +zip -r "$OUT" . \ + -x '.git/*' '__pycache__/*' '*.pyc' '.pytest_cache/*' '*.zip' + +echo "Created archive: $OUT" diff --git a/scripts/train.sh b/scripts/train.sh new file mode 100755 index 0000000..23bbd33 --- /dev/null +++ b/scripts/train.sh @@ -0,0 +1,3 @@ +#!/usr/bin/env bash +set -euo pipefail +python -m mlipx.train --config ${1:-configs/minimal.yaml} diff --git a/src/mlipx/__init__.py b/src/mlipx/__init__.py new file mode 100644 index 0000000..2801c52 --- /dev/null +++ b/src/mlipx/__init__.py @@ -0,0 +1,3 @@ +from mlipx.config import ExperimentConfig, load_config + +__all__ = ["ExperimentConfig", "load_config"] diff --git a/src/mlipx/ase_ext/__init__.py b/src/mlipx/ase_ext/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/mlipx/ase_ext/calculator.py b/src/mlipx/ase_ext/calculator.py new file mode 100644 index 0000000..d9fa5fc --- /dev/null +++ b/src/mlipx/ase_ext/calculator.py @@ -0,0 +1,75 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import numpy as np +import torch +from ase.calculators.calculator import Calculator, all_changes + +from mlipx.config import ExperimentConfig +from mlipx.data.datamodule import MLIPDataModule +from mlipx.data.dataset import ASEDataset, StructureSample +from mlipx.models.model import EnergyScaler, MLIPEnergyModel + + +@dataclass +class CalculatorContext: + cfg: ExperimentConfig + dm: MLIPDataModule + model: MLIPEnergyModel + device: torch.device + + +class MLIPXCalculator(Calculator): + implemented_properties = ["energy", "forces"] + + def __init__(self, context: CalculatorContext, **kwargs): + super().__init__(**kwargs) + self.ctx = context + + def calculate(self, atoms=None, properties=("energy", "forces"), system_changes=all_changes): + super().calculate(atoms, properties, system_changes) + assert atoms is not None + sample = StructureSample( + z=torch.tensor(atoms.numbers, dtype=torch.long), + pos=torch.tensor(atoms.positions, dtype=torch.float32), + energy=torch.tensor([0.0], dtype=torch.float32), + forces=torch.zeros((len(atoms), 3), dtype=torch.float32), + cell=torch.tensor(np.asarray(atoms.cell.array), dtype=torch.float32), + pbc=torch.tensor(np.asarray(atoms.pbc, dtype=np.int64), dtype=torch.bool), + ) + batch = self.ctx.dm.collate([sample]) + batch = type(batch)(**{k: getattr(batch, k).to(self.ctx.device) for k in batch.__dataclass_fields__.keys()}) + self.ctx.model.eval() + en_n, forces = self.ctx.model.predict_energy_forces(batch) + en = en_n * self.ctx.dm.stats.energy_std + self.ctx.dm.stats.energy_mean + self.results["energy"] = float(en.item()) + self.results["forces"] = forces.detach().cpu().numpy() + + +def build_calculator(cfg: ExperimentConfig, checkpoint: str) -> MLIPXCalculator: + ds = ASEDataset(path=cfg.data.path, energy_key=cfg.data.energy_key, force_key=cfg.data.force_key) + dm = MLIPDataModule( + dataset=ds, + cutoff=cfg.data.cutoff, + max_neighbors=cfg.data.max_neighbors, + batch_size=cfg.data.batch_size, + num_workers=cfg.data.num_workers, + train_ratio=cfg.data.train_ratio, + val_ratio=cfg.data.val_ratio, + test_ratio=cfg.data.test_ratio, + normalize_energy_per_atom=cfg.data.normalize_energy_per_atom, + seed=cfg.train.seed, + ) + device = torch.device("cuda" if (cfg.train.device == "auto" and torch.cuda.is_available()) else cfg.train.device) + model = MLIPEnergyModel( + hidden_dim=cfg.model.hidden_dim, + n_interactions=cfg.model.n_interactions, + n_rbf=cfg.model.n_rbf, + cutoff=cfg.model.cutoff, + scaler=EnergyScaler(mean=dm.stats.energy_mean, std=dm.stats.energy_std), + ).to(device) + ckpt = torch.load(checkpoint, map_location=device) + model.load_state_dict(ckpt["model"]) + ctx = CalculatorContext(cfg=cfg, dm=dm, model=model, device=device) + return MLIPXCalculator(context=ctx) diff --git a/src/mlipx/config.py b/src/mlipx/config.py new file mode 100644 index 0000000..b53f3c5 --- /dev/null +++ b/src/mlipx/config.py @@ -0,0 +1,76 @@ +"""Configuration dataclasses and YAML loading utilities.""" +from __future__ import annotations + +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import yaml + + +@dataclass +class DataConfig: + path: str + format: str = "extxyz" + energy_key: str = "energy" + force_key: str = "forces" + cutoff: float = 5.0 + max_neighbors: int = 64 + train_ratio: float = 0.8 + val_ratio: float = 0.1 + test_ratio: float = 0.1 + batch_size: int = 8 + num_workers: int = 0 + normalize_energy_per_atom: bool = True + + +@dataclass +class ModelConfig: + hidden_dim: int = 128 + n_interactions: int = 4 + n_rbf: int = 32 + cutoff: float = 5.0 + + +@dataclass +class TrainConfig: + seed: int = 42 + device: str = "auto" + epochs: int = 20 + lr: float = 2e-4 + weight_decay: float = 1e-6 + energy_weight: float = 1.0 + force_weight: float = 50.0 + grad_clip: float = 5.0 + amp: bool = False + scheduler: str = "cosine" + scheduler_tmax: int = 20 + early_stopping_patience: int = 10 + output_dir: str = "artifacts/run_default" + resume_checkpoint: str | None = None + + +@dataclass +class ExperimentConfig: + name: str + data: DataConfig + model: ModelConfig + train: TrainConfig + + +def _build_cfg(raw: dict[str, Any]) -> ExperimentConfig: + return ExperimentConfig( + name=raw.get("name", "mlipx"), + data=DataConfig(**raw["data"]), + model=ModelConfig(**raw["model"]), + train=TrainConfig(**raw["train"]), + ) + + +def load_config(path: str | Path) -> ExperimentConfig: + cfg_path = Path(path) + if not cfg_path.exists(): + raise FileNotFoundError(f"Config file not found: {cfg_path}") + with cfg_path.open("r", encoding="utf-8") as f: + raw = yaml.safe_load(f) + return _build_cfg(raw) diff --git a/src/mlipx/data/__init__.py b/src/mlipx/data/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/mlipx/data/datamodule.py b/src/mlipx/data/datamodule.py new file mode 100644 index 0000000..c364960 --- /dev/null +++ b/src/mlipx/data/datamodule.py @@ -0,0 +1,122 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import torch + +from .dataset import ASEDataset, StructureSample +from .graph import build_radius_graph +from .types import GraphBatch + + +@dataclass +class NormalizationStats: + energy_mean: float + energy_std: float + + +class MLIPDataModule: + def __init__( + self, + dataset: ASEDataset, + cutoff: float, + max_neighbors: int, + batch_size: int, + num_workers: int, + train_ratio: float, + val_ratio: float, + test_ratio: float, + normalize_energy_per_atom: bool, + seed: int, + ): + if abs((train_ratio + val_ratio + test_ratio) - 1.0) > 1e-6: + raise ValueError("train/val/test ratios must sum to 1.0") + self.dataset = dataset + self.cutoff = cutoff + self.max_neighbors = max_neighbors + self.batch_size = batch_size + self.num_workers = num_workers + self.normalize_energy_per_atom = normalize_energy_per_atom + self.seed = seed + + n_total = len(dataset) + n_train = int(n_total * train_ratio) + n_val = int(n_total * val_ratio) + n_test = n_total - n_train - n_val + generator = torch.Generator().manual_seed(seed) + self.train_set, self.val_set, self.test_set = torch.utils.data.random_split( + dataset, [n_train, n_val, n_test], generator=generator + ) + self.stats = self.compute_train_stats() + + def _energy_target(self, sample: StructureSample) -> float: + e = sample.energy.item() + if self.normalize_energy_per_atom: + e /= float(sample.z.numel()) + return e + + def compute_train_stats(self) -> NormalizationStats: + targets = [] + for idx in self.train_set.indices: + sample = self.dataset[idx] + targets.append(self._energy_target(sample)) + mean = float(torch.tensor(targets).mean()) + std = float(torch.tensor(targets).std(unbiased=False).clamp(min=1e-8)) + return NormalizationStats(energy_mean=mean, energy_std=std) + + def collate(self, batch: list[StructureSample]) -> GraphBatch: + z_list, pos_list, energy_list, forces_list, natoms = [], [], [], [], [] + batch_index = [] + offset = 0 + for i, item in enumerate(batch): + n = item.z.numel() + natoms.append(n) + z_list.append(item.z) + pos_list.append(item.pos) + forces_list.append(item.forces) + et = self._energy_target(item) + et = (et - self.stats.energy_mean) / self.stats.energy_std + energy_list.append(torch.tensor([et], dtype=torch.float32)) + batch_index.append(torch.full((n,), i, dtype=torch.long)) + offset += n + + z = torch.cat(z_list, dim=0) + pos = torch.cat(pos_list, dim=0) + forces = torch.cat(forces_list, dim=0) + batch_vec = torch.cat(batch_index, dim=0) + energy = torch.stack(energy_list, dim=0) + edge_index, edge_vec, edge_dist = build_radius_graph( + pos=pos, + batch=batch_vec, + cutoff=self.cutoff, + max_neighbors=self.max_neighbors, + ) + return GraphBatch( + z=z, + pos=pos, + batch=batch_vec, + edge_index=edge_index, + edge_vec=edge_vec, + edge_dist=edge_dist, + energy=energy, + forces=forces, + natoms=torch.tensor(natoms, dtype=torch.long), + ) + + def _loader(self, split, shuffle: bool) -> torch.utils.data.DataLoader: + return torch.utils.data.DataLoader( + split, + batch_size=self.batch_size, + shuffle=shuffle, + num_workers=self.num_workers, + collate_fn=self.collate, + ) + + def train_loader(self): + return self._loader(self.train_set, shuffle=True) + + def val_loader(self): + return self._loader(self.val_set, shuffle=False) + + def test_loader(self): + return self._loader(self.test_set, shuffle=False) diff --git a/src/mlipx/data/dataset.py b/src/mlipx/data/dataset.py new file mode 100644 index 0000000..aa10873 --- /dev/null +++ b/src/mlipx/data/dataset.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Sequence + +import numpy as np +import torch +from ase import Atoms +from ase.io import read + + +@dataclass +class StructureSample: + z: torch.Tensor + pos: torch.Tensor + energy: torch.Tensor + forces: torch.Tensor + cell: torch.Tensor + pbc: torch.Tensor + + +class ASEDataset(torch.utils.data.Dataset): + """ASE extxyz dataset for atomistic ML.""" + + def __init__(self, path: str, energy_key: str = "energy", force_key: str = "forces"): + self.frames: Sequence[Atoms] = read(path, index=":") + if len(self.frames) == 0: + raise ValueError(f"No structures found in dataset: {path}") + self.energy_key = energy_key + self.force_key = force_key + + def __len__(self) -> int: + return len(self.frames) + + def __getitem__(self, idx: int) -> StructureSample: + atoms = self.frames[idx] + z = torch.tensor(atoms.numbers, dtype=torch.long) + pos = torch.tensor(atoms.positions, dtype=torch.float32) + + energy = atoms.info.get(self.energy_key) + if energy is None: + try: + energy = atoms.get_potential_energy() + except Exception as exc: # noqa: BLE001 + raise KeyError(f"Missing energy for sample {idx}") from exc + + forces = atoms.arrays.get(self.force_key) + if forces is None: + try: + forces = atoms.get_forces() + except Exception as exc: # noqa: BLE001 + raise KeyError(f"Missing forces for sample {idx}") from exc + + cell = torch.tensor(np.asarray(atoms.cell.array), dtype=torch.float32) + pbc = torch.tensor(np.asarray(atoms.pbc, dtype=np.int64), dtype=torch.bool) + return StructureSample( + z=z, + pos=pos, + energy=torch.tensor([float(energy)], dtype=torch.float32), + forces=torch.tensor(forces, dtype=torch.float32), + cell=cell, + pbc=pbc, + ) diff --git a/src/mlipx/data/graph.py b/src/mlipx/data/graph.py new file mode 100644 index 0000000..7e34655 --- /dev/null +++ b/src/mlipx/data/graph.py @@ -0,0 +1,52 @@ +"""Neighbor graph construction utilities.""" +from __future__ import annotations + +import torch + + +def build_radius_graph( + pos: torch.Tensor, + batch: torch.Tensor, + cutoff: float, + max_neighbors: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Build directed radius graph. + + Returns edge_index (2,E), edge_vec (E,3), edge_dist (E,). + """ + device = pos.device + edge_src = [] + edge_dst = [] + edge_vec = [] + edge_dist = [] + + for mol_id in torch.unique(batch): + idx = torch.where(batch == mol_id)[0] + p = pos[idx] + dmat = torch.cdist(p, p) + n = p.shape[0] + for i in range(n): + neighbors = torch.where((dmat[i] < cutoff) & (dmat[i] > 0))[0] + if neighbors.numel() > max_neighbors: + neighbor_dist = dmat[i, neighbors] + _, order = torch.topk(neighbor_dist, k=max_neighbors, largest=False) + neighbors = neighbors[order] + for j in neighbors.tolist(): + src = idx[i].item() + dst = idx[j].item() + vec = pos[dst] - pos[src] + edge_src.append(src) + edge_dst.append(dst) + edge_vec.append(vec) + edge_dist.append(torch.norm(vec)) + + if not edge_src: + edge_index = torch.empty((2, 0), dtype=torch.long, device=device) + evec = torch.empty((0, 3), dtype=pos.dtype, device=device) + edist = torch.empty((0,), dtype=pos.dtype, device=device) + return edge_index, evec, edist + + edge_index = torch.tensor([edge_src, edge_dst], dtype=torch.long, device=device) + evec = torch.stack(edge_vec).to(device) + edist = torch.stack(edge_dist).to(device) + return edge_index, evec, edist diff --git a/src/mlipx/data/types.py b/src/mlipx/data/types.py new file mode 100644 index 0000000..3934e34 --- /dev/null +++ b/src/mlipx/data/types.py @@ -0,0 +1,18 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import torch + + +@dataclass +class GraphBatch: + z: torch.Tensor + pos: torch.Tensor + batch: torch.Tensor + edge_index: torch.Tensor + edge_vec: torch.Tensor + edge_dist: torch.Tensor + energy: torch.Tensor + forces: torch.Tensor + natoms: torch.Tensor diff --git a/src/mlipx/evaluate.py b/src/mlipx/evaluate.py new file mode 100644 index 0000000..ef500b0 --- /dev/null +++ b/src/mlipx/evaluate.py @@ -0,0 +1,86 @@ +from __future__ import annotations + +import argparse +import csv +from pathlib import Path + +import torch + +from mlipx.config import load_config +from mlipx.data.datamodule import MLIPDataModule +from mlipx.data.dataset import ASEDataset +from mlipx.models.model import EnergyScaler, MLIPEnergyModel +from mlipx.utils.metrics import mae_rmse + + +def parse_args(): + p = argparse.ArgumentParser(description="Evaluate trained MLIP model") + p.add_argument("--config", required=True) + p.add_argument("--checkpoint", required=True) + p.add_argument("--split", choices=["val", "test"], default="test") + p.add_argument("--out", default="artifacts/predictions.csv") + return p.parse_args() + + +def main(): + args = parse_args() + cfg = load_config(args.config) + ds = ASEDataset(path=cfg.data.path, energy_key=cfg.data.energy_key, force_key=cfg.data.force_key) + dm = MLIPDataModule( + dataset=ds, + cutoff=cfg.data.cutoff, + max_neighbors=cfg.data.max_neighbors, + batch_size=cfg.data.batch_size, + num_workers=cfg.data.num_workers, + train_ratio=cfg.data.train_ratio, + val_ratio=cfg.data.val_ratio, + test_ratio=cfg.data.test_ratio, + normalize_energy_per_atom=cfg.data.normalize_energy_per_atom, + seed=cfg.train.seed, + ) + + device = torch.device("cuda" if (cfg.train.device == "auto" and torch.cuda.is_available()) else cfg.train.device) + ckpt = torch.load(args.checkpoint, map_location=device) + scaler = EnergyScaler(mean=dm.stats.energy_mean, std=dm.stats.energy_std) + model = MLIPEnergyModel( + hidden_dim=cfg.model.hidden_dim, + n_interactions=cfg.model.n_interactions, + n_rbf=cfg.model.n_rbf, + cutoff=cfg.model.cutoff, + scaler=scaler, + ).to(device) + model.load_state_dict(ckpt["model"]) + model.eval() + + loader = dm.test_loader() if args.split == "test" else dm.val_loader() + pred_e, true_e, pred_f, true_f = [], [], [], [] + for batch in loader: + batch = type(batch)(**{k: getattr(batch, k).to(device) for k in batch.__dataclass_fields__.keys()}) + e_pred_n, f_pred = model.predict_energy_forces(batch) + e_pred = e_pred_n * dm.stats.energy_std + dm.stats.energy_mean + e_true = batch.energy * dm.stats.energy_std + dm.stats.energy_mean + pred_e.append(e_pred.detach().cpu()) + true_e.append(e_true.detach().cpu()) + pred_f.append(f_pred.detach().cpu()) + true_f.append(batch.forces.detach().cpu()) + + pred_e_t = torch.cat(pred_e, dim=0) + true_e_t = torch.cat(true_e, dim=0) + pred_f_t = torch.cat(pred_f, dim=0) + true_f_t = torch.cat(true_f, dim=0) + + e_metrics = mae_rmse(pred_e_t, true_e_t) + f_metrics = mae_rmse(pred_f_t, true_f_t) + print(f"{args.split} | E_MAE={e_metrics.mae:.6f} E_RMSE={e_metrics.rmse:.6f} | F_MAE={f_metrics.mae:.6f} F_RMSE={f_metrics.rmse:.6f}") + + out_path = Path(args.out) + out_path.parent.mkdir(parents=True, exist_ok=True) + with out_path.open("w", newline="", encoding="utf-8") as f: + w = csv.writer(f) + w.writerow(["idx", "true_energy", "pred_energy"]) + for i, (t, p) in enumerate(zip(true_e_t.squeeze(-1).tolist(), pred_e_t.squeeze(-1).tolist())): + w.writerow([i, t, p]) + + +if __name__ == "__main__": + main() diff --git a/src/mlipx/infer.py b/src/mlipx/infer.py new file mode 100644 index 0000000..154d122 --- /dev/null +++ b/src/mlipx/infer.py @@ -0,0 +1,39 @@ +from __future__ import annotations + +import argparse +from pathlib import Path + +from ase.io import read + +from mlipx.ase_ext.calculator import build_calculator +from mlipx.config import load_config + + +def parse_args(): + p = argparse.ArgumentParser(description="Inference with trained MLIP") + p.add_argument("--config", required=True) + p.add_argument("--checkpoint", required=True) + p.add_argument("--input", required=True, help="single structure file or extxyz trajectory") + p.add_argument("--output", default="artifacts/infer_results.txt") + return p.parse_args() + + +def main(): + args = parse_args() + cfg = load_config(args.config) + calc = build_calculator(cfg, args.checkpoint) + frames = read(args.input, index=":") + + out = Path(args.output) + out.parent.mkdir(parents=True, exist_ok=True) + with out.open("w", encoding="utf-8") as f: + for i, atoms in enumerate(frames): + atoms.calc = calc + e = atoms.get_potential_energy() + forces = atoms.get_forces() + f.write(f"frame={i} energy={e:.8f} f_norm={float((forces**2).sum()**0.5):.8f}\n") + print(f"Saved inference results to {out}") + + +if __name__ == "__main__": + main() diff --git a/src/mlipx/models/__init__.py b/src/mlipx/models/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/mlipx/models/layers.py b/src/mlipx/models/layers.py new file mode 100644 index 0000000..bee5e22 --- /dev/null +++ b/src/mlipx/models/layers.py @@ -0,0 +1,28 @@ +from __future__ import annotations + +import math + +import torch +import torch.nn as nn + + +def scatter_add(src: torch.Tensor, index: torch.Tensor, dim_size: int) -> torch.Tensor: + out = torch.zeros((dim_size, src.shape[-1]), device=src.device, dtype=src.dtype) + out.index_add_(0, index, src) + return out + + +class GaussianRBF(nn.Module): + def __init__(self, n_rbf: int, cutoff: float): + super().__init__() + centers = torch.linspace(0.0, cutoff, n_rbf) + self.register_buffer("centers", centers) + self.gamma = nn.Parameter(torch.tensor(10.0)) + self.cutoff = cutoff + + def forward(self, distances: torch.Tensor) -> torch.Tensor: + d = distances.unsqueeze(-1) + rbf = torch.exp(-torch.abs(self.gamma) * (d - self.centers) ** 2) + cutoff = 0.5 * (torch.cos(math.pi * distances / self.cutoff) + 1.0) + cutoff = torch.where(distances < self.cutoff, cutoff, torch.zeros_like(cutoff)) + return rbf * cutoff.unsqueeze(-1) diff --git a/src/mlipx/models/model.py b/src/mlipx/models/model.py new file mode 100644 index 0000000..01a20f0 --- /dev/null +++ b/src/mlipx/models/model.py @@ -0,0 +1,49 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import torch +import torch.nn as nn + +from mlipx.data.types import GraphBatch +from mlipx.models.painn import PaiNN + + +@dataclass +class EnergyScaler: + mean: float + std: float + + +class MLIPEnergyModel(nn.Module): + def __init__(self, hidden_dim: int, n_interactions: int, n_rbf: int, cutoff: float, scaler: EnergyScaler): + super().__init__() + self.core = PaiNN( + n_atom_embeddings=100, + hidden_dim=hidden_dim, + n_interactions=n_interactions, + n_rbf=n_rbf, + cutoff=cutoff, + ) + self.scaler = scaler + + def forward(self, batch: GraphBatch) -> torch.Tensor: + return self.core(batch) + + def predict_energy_forces(self, batch: GraphBatch) -> tuple[torch.Tensor, torch.Tensor]: + pos = batch.pos.clone().detach().requires_grad_(True) + with torch.enable_grad(): + batch = GraphBatch( + z=batch.z, + pos=pos, + batch=batch.batch, + edge_index=batch.edge_index, + edge_vec=(pos[batch.edge_index[1]] - pos[batch.edge_index[0]]), + edge_dist=torch.norm(pos[batch.edge_index[1]] - pos[batch.edge_index[0]], dim=-1), + energy=batch.energy, + forces=batch.forces, + natoms=batch.natoms, + ) + en = self.forward(batch) + forces = -torch.autograd.grad(en.sum(), pos, create_graph=False)[0] + return en, forces diff --git a/src/mlipx/models/painn.py b/src/mlipx/models/painn.py new file mode 100644 index 0000000..1fcda8f --- /dev/null +++ b/src/mlipx/models/painn.py @@ -0,0 +1,78 @@ +"""PaiNN-style message passing network for atomistic potentials.""" +from __future__ import annotations + +import torch +import torch.nn as nn + +from mlipx.data.types import GraphBatch + +from .layers import GaussianRBF, scatter_add + + +class PaiNNInteraction(nn.Module): + def __init__(self, hidden_dim: int, n_rbf: int): + super().__init__() + self.filter_net = nn.Sequential( + nn.Linear(n_rbf, hidden_dim), + nn.SiLU(), + nn.Linear(hidden_dim, 3 * hidden_dim), + ) + self.scalar_update = nn.Sequential( + nn.Linear(hidden_dim, hidden_dim), + nn.SiLU(), + nn.Linear(hidden_dim, hidden_dim), + ) + + def forward( + self, + q: torch.Tensor, + mu: torch.Tensor, + edge_index: torch.Tensor, + edge_rbf: torch.Tensor, + edge_vec: torch.Tensor, + edge_dist: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + src, dst = edge_index + filters = self.filter_net(edge_rbf) + f_q, f_mu_r, f_mu_v = torch.chunk(filters, chunks=3, dim=-1) + + q_msg = q[dst] * f_q + dq = scatter_add(q_msg, src, q.shape[0]) + + unit = edge_vec / edge_dist.unsqueeze(-1).clamp(min=1e-8) + mu_dst = mu[dst] + radial = f_mu_r.unsqueeze(-1) * unit.unsqueeze(1) + mu_msg = f_mu_v.unsqueeze(-1) * mu_dst + radial + dmu = torch.zeros_like(mu) + dmu.index_add_(0, src, mu_msg) + + q = q + self.scalar_update(dq) + mu = mu + dmu + return q, mu + + +class PaiNN(nn.Module): + def __init__(self, n_atom_embeddings: int, hidden_dim: int, n_interactions: int, n_rbf: int, cutoff: float): + super().__init__() + self.embedding = nn.Embedding(n_atom_embeddings, hidden_dim) + self.rbf = GaussianRBF(n_rbf=n_rbf, cutoff=cutoff) + self.interactions = nn.ModuleList([PaiNNInteraction(hidden_dim=hidden_dim, n_rbf=n_rbf) for _ in range(n_interactions)]) + self.readout = nn.Sequential( + nn.Linear(hidden_dim, hidden_dim), + nn.SiLU(), + nn.Linear(hidden_dim, 1), + ) + + def forward(self, batch: GraphBatch) -> torch.Tensor: + q = self.embedding(batch.z) + mu = torch.zeros((q.shape[0], q.shape[1], 3), device=q.device, dtype=q.dtype) + + edge_rbf = self.rbf(batch.edge_dist) + for block in self.interactions: + q, mu = block(q, mu, batch.edge_index, edge_rbf, batch.edge_vec, batch.edge_dist) + + atom_e = self.readout(q).squeeze(-1) + n_struct = int(batch.batch.max().item()) + 1 + total_e = torch.zeros((n_struct,), device=q.device, dtype=q.dtype) + total_e.index_add_(0, batch.batch, atom_e) + return total_e.unsqueeze(-1) diff --git a/src/mlipx/train.py b/src/mlipx/train.py new file mode 100644 index 0000000..94a9fb8 --- /dev/null +++ b/src/mlipx/train.py @@ -0,0 +1,50 @@ +from __future__ import annotations + +import argparse + +from mlipx.config import load_config +from mlipx.data.datamodule import MLIPDataModule +from mlipx.data.dataset import ASEDataset +from mlipx.training.engine import Trainer +from mlipx.utils.seed import set_seed + + +def parse_args(): + parser = argparse.ArgumentParser(description="Train MLIP model") + parser.add_argument("--config", type=str, required=True) + parser.add_argument("--epochs", type=int, default=None) + parser.add_argument("--batch-size", type=int, default=None) + parser.add_argument("--lr", type=float, default=None) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + cfg = load_config(args.config) + if args.epochs is not None: + cfg.train.epochs = args.epochs + if args.batch_size is not None: + cfg.data.batch_size = args.batch_size + if args.lr is not None: + cfg.train.lr = args.lr + + set_seed(cfg.train.seed) + ds = ASEDataset(path=cfg.data.path, energy_key=cfg.data.energy_key, force_key=cfg.data.force_key) + dm = MLIPDataModule( + dataset=ds, + cutoff=cfg.data.cutoff, + max_neighbors=cfg.data.max_neighbors, + batch_size=cfg.data.batch_size, + num_workers=cfg.data.num_workers, + train_ratio=cfg.data.train_ratio, + val_ratio=cfg.data.val_ratio, + test_ratio=cfg.data.test_ratio, + normalize_energy_per_atom=cfg.data.normalize_energy_per_atom, + seed=cfg.train.seed, + ) + trainer = Trainer(cfg=cfg, datamodule=dm) + trainer.fit() + + +if __name__ == "__main__": + main() diff --git a/src/mlipx/training/__init__.py b/src/mlipx/training/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/mlipx/training/engine.py b/src/mlipx/training/engine.py new file mode 100644 index 0000000..b6b9396 --- /dev/null +++ b/src/mlipx/training/engine.py @@ -0,0 +1,148 @@ +from __future__ import annotations + +import json +from dataclasses import asdict +from pathlib import Path + +import torch +import torch.nn.functional as F +from torch.optim import AdamW + +from mlipx.config import ExperimentConfig +from mlipx.data.datamodule import MLIPDataModule +from mlipx.data.types import GraphBatch +from mlipx.models.model import EnergyScaler, MLIPEnergyModel +from mlipx.utils.metrics import mae_rmse + + +class Trainer: + def __init__(self, cfg: ExperimentConfig, datamodule: MLIPDataModule): + self.cfg = cfg + self.dm = datamodule + self.device = torch.device("cuda" if (cfg.train.device == "auto" and torch.cuda.is_available()) else cfg.train.device) + scaler = EnergyScaler(mean=self.dm.stats.energy_mean, std=self.dm.stats.energy_std) + self.model = MLIPEnergyModel( + hidden_dim=cfg.model.hidden_dim, + n_interactions=cfg.model.n_interactions, + n_rbf=cfg.model.n_rbf, + cutoff=cfg.model.cutoff, + scaler=scaler, + ).to(self.device) + self.opt = AdamW(self.model.parameters(), lr=cfg.train.lr, weight_decay=cfg.train.weight_decay) + self.scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(self.opt, T_max=cfg.train.scheduler_tmax) + self.use_amp = cfg.train.amp and self.device.type == "cuda" + self.grad_scaler = torch.amp.GradScaler("cuda", enabled=self.use_amp) + self.out_dir = Path(cfg.train.output_dir) + self.out_dir.mkdir(parents=True, exist_ok=True) + self.best_val = float("inf") + self.bad_epochs = 0 + self.start_epoch = 1 + + if cfg.train.resume_checkpoint: + self._load_checkpoint(cfg.train.resume_checkpoint) + + def _to_device(self, b: GraphBatch) -> GraphBatch: + return GraphBatch( + z=b.z.to(self.device), + pos=b.pos.to(self.device), + batch=b.batch.to(self.device), + edge_index=b.edge_index.to(self.device), + edge_vec=b.edge_vec.to(self.device), + edge_dist=b.edge_dist.to(self.device), + energy=b.energy.to(self.device), + forces=b.forces.to(self.device), + natoms=b.natoms.to(self.device), + ) + + def _compute_loss(self, batch: GraphBatch) -> tuple[torch.Tensor, dict[str, float]]: + pred_energy_norm, pred_forces = self.model.predict_energy_forces(batch) + loss_e = F.mse_loss(pred_energy_norm, batch.energy) + pred_forces_norm = pred_forces / self.dm.stats.energy_std + target_forces_norm = batch.forces / self.dm.stats.energy_std + loss_f = F.mse_loss(pred_forces_norm, target_forces_norm) + total = self.cfg.train.energy_weight * loss_e + self.cfg.train.force_weight * loss_f + + pred_e_phys = pred_energy_norm * self.dm.stats.energy_std + self.dm.stats.energy_mean + target_e_phys = batch.energy * self.dm.stats.energy_std + self.dm.stats.energy_mean + e_metrics = mae_rmse(pred_e_phys, target_e_phys) + f_metrics = mae_rmse(pred_forces, batch.forces) + logs = { + "loss": float(total.item()), + "e_mae": e_metrics.mae, + "e_rmse": e_metrics.rmse, + "f_mae": f_metrics.mae, + "f_rmse": f_metrics.rmse, + } + return total, logs + + def _run_epoch(self, loader, train: bool) -> dict[str, float]: + self.model.train(train) + totals = {"loss": 0.0, "e_mae": 0.0, "e_rmse": 0.0, "f_mae": 0.0, "f_rmse": 0.0} + n_batches = 0 + for b in loader: + batch = self._to_device(b) + if train: + self.opt.zero_grad(set_to_none=True) + with torch.autocast(device_type=self.device.type, enabled=self.use_amp): + loss, logs = self._compute_loss(batch) + self.grad_scaler.scale(loss).backward() + self.grad_scaler.unscale_(self.opt) + torch.nn.utils.clip_grad_norm_(self.model.parameters(), self.cfg.train.grad_clip) + self.grad_scaler.step(self.opt) + self.grad_scaler.update() + else: + loss, logs = self._compute_loss(batch) + for k in totals: + totals[k] += logs[k] + n_batches += 1 + for k in totals: + totals[k] /= max(1, n_batches) + return totals + + def _save_checkpoint(self, epoch: int, is_best: bool): + ckpt = { + "epoch": epoch, + "model": self.model.state_dict(), + "optimizer": self.opt.state_dict(), + "scheduler": self.scheduler.state_dict(), + "config": asdict(self.cfg), + "stats": asdict(self.dm.stats), + } + torch.save(ckpt, self.out_dir / "last.pt") + if is_best: + torch.save(ckpt, self.out_dir / "best.pt") + + def _load_checkpoint(self, path: str): + ckpt = torch.load(path, map_location=self.device) + self.model.load_state_dict(ckpt["model"]) + self.opt.load_state_dict(ckpt["optimizer"]) + self.scheduler.load_state_dict(ckpt["scheduler"]) + self.start_epoch = int(ckpt["epoch"]) + 1 + + def fit(self) -> None: + history = [] + for epoch in range(self.start_epoch, self.cfg.train.epochs + 1): + train_logs = self._run_epoch(self.dm.train_loader(), train=True) + val_logs = self._run_epoch(self.dm.val_loader(), train=False) + self.scheduler.step() + is_best = val_logs["loss"] < self.best_val + if is_best: + self.best_val = val_logs["loss"] + self.bad_epochs = 0 + else: + self.bad_epochs += 1 + self._save_checkpoint(epoch, is_best=is_best) + row = {"epoch": epoch, "train": train_logs, "val": val_logs} + history.append(row) + print( + f"Epoch {epoch:03d} | " + f"train_loss={train_logs['loss']:.4f} val_loss={val_logs['loss']:.4f} | " + f"E_MAE={val_logs['e_mae']:.4f} E_RMSE={val_logs['e_rmse']:.4f} | " + f"F_MAE={val_logs['f_mae']:.4f} F_RMSE={val_logs['f_rmse']:.4f}" + ) + if self.bad_epochs >= self.cfg.train.early_stopping_patience: + print("Early stopping triggered") + break + + with (self.out_dir / "history.json").open("w", encoding="utf-8") as f: + json.dump(history, f, indent=2) diff --git a/src/mlipx/utils/__init__.py b/src/mlipx/utils/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/mlipx/utils/metrics.py b/src/mlipx/utils/metrics.py new file mode 100644 index 0000000..2c5cd00 --- /dev/null +++ b/src/mlipx/utils/metrics.py @@ -0,0 +1,19 @@ +from __future__ import annotations + +import math +from dataclasses import dataclass + +import torch + + +@dataclass +class RegressionMetrics: + mae: float + rmse: float + + +def mae_rmse(pred: torch.Tensor, target: torch.Tensor) -> RegressionMetrics: + err = pred - target + mae = torch.mean(torch.abs(err)).item() + rmse = math.sqrt(torch.mean(err * err).item()) + return RegressionMetrics(mae=mae, rmse=rmse) diff --git a/src/mlipx/utils/seed.py b/src/mlipx/utils/seed.py new file mode 100644 index 0000000..c1977b7 --- /dev/null +++ b/src/mlipx/utils/seed.py @@ -0,0 +1,13 @@ +from __future__ import annotations + +import random + +import numpy as np +import torch + + +def set_seed(seed: int) -> None: + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..7a33e64 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,8 @@ +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +SRC = ROOT / "src" +for p in (ROOT, SRC): + if str(p) not in sys.path: + sys.path.insert(0, str(p)) diff --git a/tests/test_data_pipeline.py b/tests/test_data_pipeline.py new file mode 100644 index 0000000..a453b27 --- /dev/null +++ b/tests/test_data_pipeline.py @@ -0,0 +1,50 @@ +import pytest +pytest.importorskip("torch") +pytest.importorskip("ase") +pytest.importorskip("numpy") + +from pathlib import Path + +import numpy as np +import pytest +from ase import Atoms +from ase.calculators.emt import EMT +from ase.io import write + +from mlipx.data.datamodule import MLIPDataModule +from mlipx.data.dataset import ASEDataset + + +def build_tmp_dataset(path: Path, n_samples: int = 6): + frames = [] + for _ in range(n_samples): + atoms = Atoms("H2O", positions=np.random.rand(3, 3), cell=np.eye(3) * 6.0, pbc=False) + atoms.calc = EMT() + atoms.info["energy"] = atoms.get_potential_energy() + atoms.arrays["forces"] = atoms.get_forces() + frames.append(atoms) + write(path, frames, format="extxyz") + + +@pytest.mark.parametrize("batch_size", [2]) +def test_loader_and_graph(batch_size, tmp_path): + data_path = tmp_path / "tiny.extxyz" + build_tmp_dataset(data_path) + ds = ASEDataset(str(data_path)) + dm = MLIPDataModule( + dataset=ds, + cutoff=5.0, + max_neighbors=16, + batch_size=batch_size, + num_workers=0, + train_ratio=0.7, + val_ratio=0.2, + test_ratio=0.1, + normalize_energy_per_atom=True, + seed=123, + ) + batch = next(iter(dm.train_loader())) + assert batch.z.ndim == 1 + assert batch.pos.shape[1] == 3 + assert batch.edge_index.shape[0] == 2 + assert batch.energy.ndim == 2 diff --git a/tests/test_model_and_train.py b/tests/test_model_and_train.py new file mode 100644 index 0000000..86e4856 --- /dev/null +++ b/tests/test_model_and_train.py @@ -0,0 +1,87 @@ +import pytest +pytest.importorskip("torch") +pytest.importorskip("ase") +pytest.importorskip("numpy") + +import numpy as np +from ase import Atoms +from ase.calculators.emt import EMT +from ase.io import write + +from mlipx.config import load_config +from mlipx.data.datamodule import MLIPDataModule +from mlipx.data.dataset import ASEDataset +from mlipx.models.model import EnergyScaler, MLIPEnergyModel +from mlipx.training.engine import Trainer + + +def _make_dataset(path, n=8): + frames = [] + for _ in range(n): + atoms = Atoms("H2", positions=np.random.rand(2, 3), cell=np.eye(3) * 5.0, pbc=False) + atoms.calc = EMT() + atoms.info["energy"] = atoms.get_potential_energy() + atoms.arrays["forces"] = atoms.get_forces() + frames.append(atoms) + write(path, frames, format="extxyz") + + +def test_forward_and_force(tmp_path): + path = tmp_path / "d.extxyz" + _make_dataset(path, n=6) + ds = ASEDataset(str(path)) + dm = MLIPDataModule(ds, 5.0, 16, 2, 0, 0.8, 0.1, 0.1, True, 42) + batch = next(iter(dm.train_loader())) + model = MLIPEnergyModel(hidden_dim=32, n_interactions=2, n_rbf=16, cutoff=5.0, scaler=EnergyScaler(dm.stats.energy_mean, dm.stats.energy_std)) + e, f = model.predict_energy_forces(batch) + assert e.shape[0] == batch.energy.shape[0] + assert f.shape == batch.forces.shape + + +def test_smoke_train(tmp_path): + path = tmp_path / "d2.extxyz" + _make_dataset(path, n=10) + cfg_path = tmp_path / "cfg.yaml" + cfg_path.write_text( + f""" +name: smoke +data: + path: {path} + format: extxyz + energy_key: energy + force_key: forces + cutoff: 5.0 + max_neighbors: 16 + train_ratio: 0.8 + val_ratio: 0.1 + test_ratio: 0.1 + batch_size: 2 + num_workers: 0 + normalize_energy_per_atom: true +model: + hidden_dim: 32 + n_interactions: 2 + n_rbf: 16 + cutoff: 5.0 +train: + seed: 42 + device: cpu + epochs: 1 + lr: 0.001 + weight_decay: 1.0e-6 + energy_weight: 1.0 + force_weight: 5.0 + grad_clip: 5.0 + amp: false + scheduler: cosine + scheduler_tmax: 1 + early_stopping_patience: 5 + output_dir: {tmp_path}/artifacts + resume_checkpoint: null +""", + encoding="utf-8", + ) + cfg = load_config(cfg_path) + dm = MLIPDataModule(ASEDataset(cfg.data.path), cfg.data.cutoff, cfg.data.max_neighbors, cfg.data.batch_size, cfg.data.num_workers, cfg.data.train_ratio, cfg.data.val_ratio, cfg.data.test_ratio, cfg.data.normalize_energy_per_atom, cfg.train.seed) + trainer = Trainer(cfg, dm) + trainer.fit()