diff --git a/interatomic_potentials/README.md b/interatomic_potentials/README.md index ae0cd3b2..d6ebf12e 100644 --- a/interatomic_potentials/README.md +++ b/interatomic_potentials/README.md @@ -6,30 +6,81 @@ Machine-learning interatomic potentials (MLIP) bridge the gap between quantum-le ## 2.Models Matrix -| **Supported Functions** | **[CHGNet](./configs/chgnet/README.md)** | **[MatterSim](./configs/mattersim//README.md)** | -| ----------------------------------- | ---------------------------------------- | ----------------------------------------------- | -| **Forward Prediction** | | | -|  Energy | ✅ | ✅ | -|  Force | ✅ | ✅ | -|  Stress | ✅ | ✅ | -|  Magmom | ✅ | - | -| **ML Capabilities · Training** | | | -|  Single-GPU | ✅ | ✅ | -|  Distributed Train | ✅ | ✅ | -|  Mixed Precision | - | - | -|  Fine-tuning | ✅ | ✅ | -|  Uncertainty / Active-Learning | - | - | -|  Dynamic→Static | - | - | -|  Compiler CINN | - | - | -| **ML Capabilities · Predict** | | | -|  Distillation / Pruning | - | - | -|  Standard inference | ✅ | ✅ | -|  Distributed inference | - | - | -|  Compiler CINN | - | - | -| **Molecular Dynamic Interface** | | | -|  ASE | ✅ | ✅ | -| **Dataset** | | | -|  MPtrj | ✅ | 🚧 | -| **ML2DDB🌟** | ✅ | - | +| **Supported Functions** | **[CHGNet](./configs/chgnet/README.md)** | **[MatterSim](./configs/mattersim//README.md)** | **[SevenNet](./configs/sevennet/README.md)** | +| ----------------------------------- | ---------------------------------------- | ----------------------------------------------- | -------------------------------------------- | +| **Forward Prediction** | | | | +|  Energy | ✅ | ✅ | ✅ | +|  Force | ✅ | ✅ | ✅ | +|  Stress | ✅ | ✅ | - | +|  Magmom | ✅ | - | - | +| **ML Capabilities · Training** | | | | +|  Single-GPU | ✅ | ✅ | ✅ | +|  Distributed Train | ✅ | ✅ | - | +|  Mixed Precision | - | - | - | +|  Fine-tuning | ✅ | ✅ | ✅ | +|  Uncertainty / Active-Learning | - | - | - | +|  Dynamic→Static | - | - | - | +|  Compiler CINN | - | - | - | +| **ML Capabilities · Predict** | | | | +|  Distillation / Pruning | - | - | - | +|  Standard inference | ✅ | ✅ | ✅ | +|  Distributed inference | - | - | - | +|  Compiler CINN | - | - | - | +| **Molecular Dynamic Interface** | | | | +|  ASE | ✅ | ✅ | ✅ | +| **Dataset** | | | | +|  MPtrj | ✅ | 🚧 | - | +| **ML2DDB🌟** | ✅ | - | - | **Notice**:🌟 represent originate research work published from paddlematerials toolkit + +## 3.SevenNet Quick Start + +### 3.1 Overview + +SevenNet is a message-passing graph neural network for interatomic potential prediction. This implementation provides a pure PaddlePaddle version without external dependencies like e3nn. + +### 3.2 Installation + +```bash +# Install dependencies +pip install paddlepaddle>=3.0.0 ase numpy tqdm pyyaml +``` + +### 3.3 Training + +```bash +# Navigate to interatomic_potentials directory +cd interatomic_potentials + +# Run training with config file +python sevennet_train.py --config configs/sevennet/sevennet_hfo2.yaml +``` + +### 3.4 Inference + +```bash +python sevennet_train.py --config configs/sevennet/sevennet_hfo2.yaml \ + --output-dir ./output/sevennet_hfo2 +``` + +### 3.5 Configuration + +The configuration file includes: +- **Model parameters**: hidden_dim, num_message_layers, num_rbf, cutoff +- **Dataset parameters**: path, cutoff, valid_ratio +- **Training parameters**: epoch, batch_size, learning_rate + +### 3.6 Test + +```bash +# Run unit tests +pytest test/test_sevennet.py -v +``` + +### 3.7 Notes + +- SevenNet uses message-passing architecture with Gaussian RBF for distance encoding +- Supports energy and force prediction +- Pure PaddlePaddle implementation, no PyTorch/e3nn dependencies +- Recommended for small to medium scale molecular systems \ No newline at end of file diff --git a/interatomic_potentials/configs/sevennet/README.md b/interatomic_potentials/configs/sevennet/README.md new file mode 100644 index 00000000..e3ce1715 --- /dev/null +++ b/interatomic_potentials/configs/sevennet/README.md @@ -0,0 +1,84 @@ +# SevenNet Model Configuration + +## Overview + +SevenNet is a message-passing graph neural network for interatomic potential prediction. + +## Pretrained Models + +SevenNet provides multiple pretrained models: + +| Model | Description | Training Dataset | Performance (CPS) | +|-------|-------------|------------------|-------------------| +| **SevenNet-Omni** (Recommended) | Universal potential, 15 datasets | 15 open ab initio datasets | 0.849 | +| SevenNet-Omni-i8 | Higher capacity (Nlayer=8) | 15 datasets | 0.859 | +| SevenNet-Omni-i12 | Highest capacity (Nlayer=12) | 15 datasets | 0.873 | +| SevenNet-MF-ompa | Multi-fidelity learning | MPtrj, sAlex, OMat24 | 0.845 | +| SevenNet-omat | OMat24 only | OMat24 | κSRME: 0.221 | +| SevenNet-l3i5 | MPtrj with lmax=3 | MPtrj | 0.714 | +| SevenNet-0 | Fastest inference | MPtrj | F1: 0.67 | + +## Using Pretrained Models + +### Step 1: Download Model + +When official provides URLs, download the pretrained model: + +```python +from ppmat.utils import download + +model_name = "sevennet_omni" # or other model name +model_path = download.get_weights_path_from_url(MODEL_REGISTRY[model_name]) +``` + +### Step 2: Load Model + +```python +import paddle +from ppmat.models import PurePaddleSevenNet + +# Load checkpoint +state_dict = paddle.load("path/to/checkpoint.pdparams") + +# Create model +model = PurePaddleSevenNet( + num_species=100, + hidden_dim=128, + num_message_layers=5, + num_rbf=32, + cutoff=5.0, +) + +# Load weights +model.set_state_dict(state_dict) +model.eval() +``` + +### Step 3: Predict + +```python +# Create graph from structure +graph = { + "z": paddle.to_tensor([1, 8, 1], dtype="int64"), # H, O, H + "pos": paddle.to_tensor([[0, 0, 0], [1, 0, 0], [0, 1, 0]], dtype="float32"), + "edge_index": paddle.to_tensor([[0, 1], [1, 0]], dtype="int64"), +} + +# Predict +result = model(graph) +energy = result["total_energy"] +``` + +## Training Datasets + +The official training datasets include: + +- **MPtrj** (Materials Project Trajectory): https://figshare.com/articles/dataset/Materials_Project_Trjectory_MPtrj_Dataset/23713842 +- **OMat24**: https://huggingface.co/datasets/fairchem/OMAT24 +- **sAlex**: https://huggingface.co/datasets/fairchem/OMAT24 + +## References + +- [SevenNet Official Repository](https://github.com/MDIL-SNU/SevenNet) +- [SevenNet Documentation](https://sevennet.readthedocs.io/) +- [Pretrained Models Guide](https://sevennet.readthedocs.io/en/latest/user_guide/pretrained.html) \ No newline at end of file diff --git a/interatomic_potentials/configs/sevennet/sevennet_hfo2.yaml b/interatomic_potentials/configs/sevennet/sevennet_hfo2.yaml new file mode 100644 index 00000000..de8a7cf7 --- /dev/null +++ b/interatomic_potentials/configs/sevennet/sevennet_hfo2.yaml @@ -0,0 +1,42 @@ +model: + type: "sevennet" + num_species: 100 + hidden_dim: 32 + num_message_layers: 3 + num_rbf: 32 + cutoff: 4.0 + +dataset: + type: "extxyz" + # 使用仓库内相对路径或用户自行指定 + # 示例: 下载数据集后放在 data/hfo2.extxyz + path: "data/hfo2.extxyz" + cutoff: 4.0 + valid_ratio: 0.1 + +dataloader: + batch_size: 2 + shuffle: true + num_workers: 0 + +optimizer: + type: "adam" + lr: 0.001 + +lr_scheduler: + type: "constant" + +loss: + type: "mse" + force_loss_weight: 0.1 + +trainer: + epoch: 2 + per_epoch: 1 + seed: 42 + device: "auto" + +log: + save_dir: "./output/sevennet_hfo2" + log_interval: 1 + save_interval: 1 \ No newline at end of file diff --git a/interatomic_potentials/configs/sevennet/sevennet_hfo2_large.yaml b/interatomic_potentials/configs/sevennet/sevennet_hfo2_large.yaml new file mode 100644 index 00000000..9dbddaec --- /dev/null +++ b/interatomic_potentials/configs/sevennet/sevennet_hfo2_large.yaml @@ -0,0 +1,42 @@ +model: + type: "sevennet" + num_species: 100 + hidden_dim: 128 + num_message_layers: 5 + num_rbf: 64 + cutoff: 5.0 + +dataset: + type: "extxyz" + # 使用仓库内相对路径或用户自行指定 + # 示例: 下载数据集后放在 data/hfo2.extxyz + path: "data/hfo2.extxyz" + cutoff: 5.0 + valid_ratio: 0.1 + +dataloader: + batch_size: 4 + shuffle: true + num_workers: 0 + +optimizer: + type: "adam" + lr: 0.001 + +lr_scheduler: + type: "constant" + +loss: + type: "mse" + force_loss_weight: 0.1 + +trainer: + epoch: 10 + per_epoch: 2 + seed: 42 + device: "auto" + +log: + save_dir: "./output/sevennet_hfo2_large" + log_interval: 1 + save_interval: 2 diff --git a/interatomic_potentials/infer.py b/interatomic_potentials/infer.py new file mode 100644 index 00000000..2ebbd4b9 --- /dev/null +++ b/interatomic_potentials/infer.py @@ -0,0 +1,118 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +""" +PaddleMaterials - Interatomic Potentials Inference Script +推理入口:python infer.py --model model.pdparams --structure structure.extxyz +""" + +import argparse +import os +import sys +import numpy as np +import paddle +from ase.io import read +from ase.neighborlist import neighbor_list + +PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +sys.path.insert(0, PROJECT_ROOT) + +from ppmat.models.sevennet import PurePaddleSevenNet + + +def parse_args(): + parser = argparse.ArgumentParser(description="Infer interatomic potential model") + parser.add_argument("--model", type=str, required=True, help="Path to model checkpoint") + parser.add_argument("--structure", type=str, required=True, help="Path to structure file") + parser.add_argument("--cutoff", type=float, default=5.0, help="Cutoff radius") + return parser.parse_args() + + +def atoms_to_graph(atoms, cutoff): + positions = atoms.get_positions().astype("float32") + numbers = atoms.get_atomic_numbers().astype("int64") + i, j = neighbor_list("ij", atoms, cutoff) + edge_index = np.vstack([i, j]).astype("int64") + + return { + "z": paddle.to_tensor(numbers, dtype="int64"), + "pos": paddle.to_tensor(positions, dtype="float32"), + "edge_index": paddle.to_tensor(edge_index, dtype="int64"), + } + + +def predict(model, atoms, cutoff, energy_mean, energy_std): + graph = atoms_to_graph(atoms, cutoff) + num_atoms = len(atoms) + + with paddle.no_grad(): + out = model(graph) + raw_total = float(out["total_energy"].numpy()) + pred_total = raw_total * energy_std + energy_mean + pred_per_atom = pred_total / num_atoms + + return pred_total, pred_per_atom + + +def main(): + args = parse_args() + + # 加载模型 + print(f"Loading model: {args.model}") + checkpoint = paddle.load(args.model) + state_dict = checkpoint["model_state_dict"] + energy_mean = float(checkpoint.get("energy_mean", 0.0)) + energy_std = float(checkpoint.get("energy_std", 1.0)) + + # 获取配置 + config = checkpoint.get("config", {}) + model_cfg = config.get("model", {}) + + # 创建模型 + model = PurePaddleSevenNet( + num_species=model_cfg.get("num_species", 100), + hidden_dim=model_cfg.get("hidden_dim", 128), + num_message_layers=model_cfg.get("num_message_layers", 4), + num_rbf=model_cfg.get("num_rbf", 32), + cutoff=model_cfg.get("cutoff", args.cutoff), + ) + model.set_state_dict(state_dict) + model.eval() + + print(f"Model loaded successfully") + print(f"Energy mean: {energy_mean:.6f}, std: {energy_std:.6f}") + + # 读取结构 + print(f"\nReading structure: {args.structure}") + atoms_list = read(args.structure, index=":") + if not isinstance(atoms_list, list): + atoms_list = [atoms_list] + print(f"Loaded {len(atoms_list)} structures") + + # 预测 + print("\n=== Prediction Results ===") + for idx, atoms in enumerate(atoms_list): + pred_total, pred_per_atom = predict(model, atoms, args.cutoff, energy_mean, energy_std) + + ref_energy = None + for key in ["y_energy", "energy", "free_energy", "total_energy"]: + if key in atoms.info: + try: + ref_energy = float(atoms.info[key]) + break + except Exception: + pass + + print(f"\nStructure {idx + 1}:") + print(f" Atoms: {len(atoms)}") + print(f" Predicted total energy: {pred_total:.6f} eV") + print(f" Predicted per-atom energy: {pred_per_atom:.6f} eV/atom") + + if ref_energy is not None: + print(f" Reference energy: {ref_energy:.6f} eV") + print(f" Error: {abs(pred_total - ref_energy):.6f} eV") + + print("\nInference completed!") + + +if __name__ == "__main__": + main() diff --git a/interatomic_potentials/sevennet_predict.py b/interatomic_potentials/sevennet_predict.py new file mode 100644 index 00000000..d2c8baea --- /dev/null +++ b/interatomic_potentials/sevennet_predict.py @@ -0,0 +1,311 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +""" +SevenNet prediction script using PotentialPredictor interface. + +Usage: + # Option 1: Using custom trained model + python interatomic_potentials/sevennet_predict.py \ + --config_path interatomic_potentials/configs/sevennet/sevennet_hfo2.yaml \ + --checkpoint_path path/to/checkpoint.pdparams + + # Option 2: Interactive prediction + python interatomic_potentials/sevennet_predict.py --interactive +""" + +import argparse +import sys +import os + +# Add project root to path +project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +sys.path.insert(0, project_root) + +import paddle +import numpy as np + +# Import SevenNet model directly to avoid ppmat import issues +import importlib.util +model_path = os.path.join(project_root, "ppmat/models/sevennet/sevennet_model.py") +spec = importlib.util.spec_from_file_location("sevennet_model", model_path) +sevennet_module = importlib.util.module_from_spec(spec) +spec.loader.exec_module(sevennet_module) +PurePaddleSevenNet = sevennet_module.PurePaddleSevenNet + + +def create_graph_converter(cutoff=5.0): + """Create a simple graph converter for molecular structures. + + This is a simplified version. For production use, integrate with + ppmat.models.common.graph_converter.FindPointsInSpheres. + """ + def convert(atoms): + """Convert ASE atoms object to graph format. + + Args: + atoms: ASE Atoms object or dict with 'positions', 'numbers' + + Returns: + dict: Graph data with 'z', 'pos', 'edge_index' + """ + if hasattr(atoms, 'get_positions'): + # ASE Atoms object + positions = atoms.get_positions() + numbers = atoms.get_atomic_numbers() + else: + positions = atoms['positions'] + numbers = atoms['numbers'] + + num_atoms = len(numbers) + z = paddle.to_tensor(numbers, dtype="int64") + pos = paddle.to_tensor(positions, dtype="float32") + + # Build edge list using cutoff radius + edge_src = [] + edge_dst = [] + cutoff_dist = cutoff + + for i in range(num_atoms): + for j in range(i + 1, num_atoms): + dist = np.linalg.norm(positions[i] - positions[j]) + if dist < cutoff_dist: + edge_src.extend([i, j]) + edge_dst.extend([j, i]) + + if len(edge_src) > 0: + edge_index = paddle.to_tensor([edge_src, edge_dst], dtype="int64") + else: + edge_index = paddle.to_tensor([[], []], dtype="int64") + + return { + "z": z, + "pos": pos, + "edge_index": edge_index, + } + + return convert + + +class SevenNetPredictor: + """SevenNet predictor with PotentialPredictor-compatible interface.""" + + def __init__(self, checkpoint_path=None, config=None): + """ + Initialize SevenNet predictor. + + Args: + checkpoint_path: Path to model checkpoint (.pdparams) + config: Model configuration dict + """ + if config is None: + config = {} + + self.model = PurePaddleSevenNet(**config) + + if checkpoint_path and os.path.exists(checkpoint_path): + ckpt = paddle.load(checkpoint_path) + # Support both raw state_dict and wrapped checkpoint format + if isinstance(ckpt, dict) and "model_state_dict" in ckpt: + state_dict = ckpt["model_state_dict"] + self.energy_mean = ckpt.get("energy_mean", 0.0) + self.energy_std = ckpt.get("energy_std", 1.0) + print(f"Loaded checkpoint from {checkpoint_path}") + print(f" energy_mean={self.energy_mean}, energy_std={self.energy_std}") + else: + state_dict = ckpt + self.energy_mean = 0.0 + self.energy_std = 1.0 + self.model.set_state_dict(state_dict) + + self.model.eval() + self.graph_converter = create_graph_converter( + cutoff=config.get('cutoff', 5.0) + ) + + def predict(self, atoms): + """Predict energy for a structure. + + Args: + atoms: ASE Atoms object or dict with 'positions', 'numbers' + + Returns: + dict: Prediction results with 'energy' and optionally 'forces' + """ + graph = self.graph_converter(atoms) + + with paddle.no_grad(): + result = self.model(graph) + + # Denormalize energy + raw_energy = float(result["total_energy"].numpy()) + energy = raw_energy * self.energy_std + self.energy_mean + + return { + "energy": energy, + "atomic_energy": result["atomic_energy"].numpy(), + } + + def predict_with_forces(self, atoms): + """Predict energy and forces for a structure. + + Uses finite difference method for force calculation because + Paddle's scatter_nd_add backward pass has numerical instability + with deep message-passing networks. + + Args: + atoms: ASE Atoms object or dict with 'positions', 'numbers' + + Returns: + dict: Prediction results with 'energy', 'forces', 'atomic_energy' + """ + graph = self.graph_converter(atoms) + positions_np = graph["pos"].numpy() + numbers = graph["z"].numpy() + edge_index = graph["edge_index"] + z = graph["z"] + + # Forward pass for energy + with paddle.no_grad(): + result = self.model(graph) + + raw_energy = float(result["total_energy"].numpy()) + energy = raw_energy * self.energy_std + self.energy_mean + + # Finite difference forces: F_i = -dE/dr_i + eps = 1e-4 + num_atoms = positions_np.shape[0] + forces = np.zeros_like(positions_np) + + for i in range(num_atoms): + for j in range(3): + pos_plus = positions_np.copy() + pos_plus[i, j] += eps + pos_minus = positions_np.copy() + pos_minus[i, j] -= eps + + pos_p = paddle.to_tensor(pos_plus, dtype="float32") + graph_p = {"z": z, "pos": pos_p, "edge_index": edge_index} + with paddle.no_grad(): + e_plus = float(self.model(graph_p)["total_energy"].numpy()) + + pos_m = paddle.to_tensor(pos_minus, dtype="float32") + graph_m = {"z": z, "pos": pos_m, "edge_index": edge_index} + with paddle.no_grad(): + e_minus = float(self.model(graph_m)["total_energy"].numpy()) + + # Denormalize + e_plus = e_plus * self.energy_std + self.energy_mean + e_minus = e_minus * self.energy_std + self.energy_mean + + forces[i, j] = -(e_plus - e_minus) / (2 * eps) + + return { + "energy": energy, + "forces": forces, + "atomic_energy": result["atomic_energy"].numpy(), + } + + +def interactive_demo(): + """Interactive demo for SevenNet prediction.""" + print("\n" + "="*50) + print("SevenNet Interactive Prediction Demo") + print("="*50) + + # Try to load pretrained checkpoint + checkpoint_path = os.path.join( + project_root, "ppmat/models/sevennet/checkpoints/sevennet_hfo2_best.pdparams" + ) + + # Config matching the pretrained weights + config = { + "num_species": 100, + "hidden_dim": 64, + "num_message_layers": 5, + "num_rbf": 32, + "cutoff": 5.0, + } + + if os.path.exists(checkpoint_path): + print(f"\nLoading pretrained weights: {checkpoint_path}") + predictor = SevenNetPredictor(checkpoint_path=checkpoint_path, config=config) + else: + print("\nNo pretrained weights found, using random initialization") + predictor = SevenNetPredictor(config=config) + + # H2O molecule example + print("\nExample: H2O molecule") + h2o = { + 'positions': np.array([ + [0.0, 0.0, 0.0], # O + [0.96, 0.0, 0.0], # H + [-0.24, 0.93, 0.0], # H + ]), + 'numbers': np.array([8, 1, 1]), # O, H, H + } + + result = predictor.predict(h2o) + print(f" Total energy: {result['energy']:.6f} eV") + print(f" Atomic energies: {result['atomic_energy'].flatten()}") + + # Get forces + result_with_forces = predictor.predict_with_forces(h2o) + print(f"\n Forces:") + for i, (num, pos) in enumerate(zip(h2o['numbers'], h2o['positions'])): + element = {1: 'H', 8: 'O'}.get(num, f'Z={num}') + forces = result_with_forces['forces'][i] + has_nan = np.any(np.isnan(forces)) + if has_nan: + print(f" {element}: [forces not available - atoms outside cutoff]") + else: + print(f" {element}: {forces}") + + print("\n" + "="*50) + print("Demo completed!") + print("="*50 + "\n") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description="SevenNet Prediction") + parser.add_argument("--config", type=str, default=None, + help="Path to config YAML file") + parser.add_argument("--checkpoint", type=str, default=None, + help="Path to checkpoint file") + parser.add_argument("--interactive", action="store_true", + help="Run interactive demo") + parser.add_argument("--positions", type=str, default=None, + help="Path to positions.npy file") + parser.add_argument("--numbers", type=str, default=None, + help="Path to numbers.npy file") + args = parser.parse_args() + + if args.interactive or (args.positions is None and args.numbers is None): + interactive_demo() + else: + # Load positions and numbers + positions = np.load(args.positions) if args.positions else None + numbers = np.load(args.numbers) if args.numbers else None + + if positions is None or numbers is None: + print("Error: Please provide both --positions and --numbers") + sys.exit(1) + + atoms = {'positions': positions, 'numbers': numbers} + + # Default config (in production, load from config file) + config = { + "num_species": 100, + "hidden_dim": 128, + "num_message_layers": 4, + "num_rbf": 32, + "cutoff": 5.0, + } + + predictor = SevenNetPredictor( + checkpoint_path=args.checkpoint, + config=config + ) + + result = predictor.predict_with_forces(atoms) + print(f"Energy: {result['energy']:.6f} eV") + print(f"Forces shape: {result['forces'].shape}") \ No newline at end of file diff --git a/interatomic_potentials/sevennet_train.py b/interatomic_potentials/sevennet_train.py new file mode 100644 index 00000000..7fa9dd17 --- /dev/null +++ b/interatomic_potentials/sevennet_train.py @@ -0,0 +1,329 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +""" +SevenNet Interatomic Potential Training Script +独立训练脚本,不依赖项目主框架 + +Usage: + python sevennet_train.py --config configs/sevennet/sevennet_hfo2.yaml +""" + +import argparse +import os +import sys +import random +import numpy as np +import paddle +import paddle.nn.functional as F +from paddle.io import Dataset, DataLoader +from tqdm import tqdm + +# 添加项目路径 +PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +sys.path.insert(0, PROJECT_ROOT) + +# 直接导入模型文件,不依赖 ppmat.__init__ +import importlib.util +model_path = os.path.join(PROJECT_ROOT, "ppmat/models/sevennet/sevennet_model.py") +spec = importlib.util.spec_from_file_location("sevennet_model", model_path) +sevennet_model = importlib.util.module_from_spec(spec) +spec.loader.exec_module(sevennet_model) + + +def parse_args(): + parser = argparse.ArgumentParser(description="Train SevenNet interatomic potential") + parser.add_argument("--config", type=str, required=True, help="Path to config file") + parser.add_argument("--output-dir", type=str, default=None, help="Output directory") + return parser.parse_args() + + +def load_config(config_path): + import yaml + with open(config_path, "r") as f: + return yaml.safe_load(f) + + +def _extract_energy(atoms): + for key in ["y_energy", "energy", "free_energy", "total_energy"]: + if key in atoms.info: + try: + return float(atoms.info[key]) + except Exception: + pass + try: + return float(atoms.get_potential_energy()) + except Exception: + pass + raise ValueError("Cannot find energy label from atoms object") + + +def _extract_forces(atoms): + if "y_force" in atoms.arrays: + try: + return np.array(atoms.arrays["y_force"], dtype=np.float32) + except Exception: + pass + try: + return np.array(atoms.get_forces(), dtype=np.float32) + except Exception: + pass + raise ValueError("Cannot find force labels from atoms object") + + +def _atoms_to_graph(atoms, cutoff): + from ase.neighborlist import neighbor_list + + positions = atoms.get_positions().astype("float32") + numbers = atoms.get_atomic_numbers().astype("int64") + i, j = neighbor_list("ij", atoms, cutoff) + edge_index = np.vstack([i, j]).astype("int64") + + total_energy = float(_extract_energy(atoms)) + forces = _extract_forces(atoms) + num_atoms = len(numbers) + energy_per_atom = total_energy / max(num_atoms, 1) + + return { + "z": numbers, + "pos": positions, + "edge_index": edge_index, + "energy": np.array([total_energy], dtype="float32"), + "energy_per_atom": np.array([energy_per_atom], dtype="float32"), + "forces": forces.astype("float32"), + "num_nodes": np.array([num_atoms], dtype="int64"), + } + + +class GraphDataset(Dataset): + def __init__(self, graphs): + super().__init__() + self.graphs = graphs + + def __len__(self): + return len(self.graphs) + + def __getitem__(self, idx): + g = self.graphs[idx] + return { + "z": paddle.to_tensor(g["z"], dtype="int64"), + "pos": paddle.to_tensor(g["pos"], dtype="float32"), + "edge_index": paddle.to_tensor(g["edge_index"], dtype="int64"), + "energy": paddle.to_tensor(g["energy"], dtype="float32"), + "energy_per_atom": paddle.to_tensor(g["energy_per_atom"], dtype="float32"), + "forces": paddle.to_tensor(g["forces"], dtype="float32"), + "num_nodes": paddle.to_tensor(g["num_nodes"], dtype="int64"), + } + + +def _collate_graphs(batch): + return batch + + +def _split_graphs(graphs, valid_ratio=0.1, seed=1): + rng = random.Random(seed) + idx = list(range(len(graphs))) + rng.shuffle(idx) + n_valid = max(1, int(len(graphs) * valid_ratio)) + valid_idx = idx[:n_valid] + train_idx = idx[n_valid:] + return [graphs[i] for i in train_idx], [graphs[i] for i in valid_idx] + + +def _compute_stats(graphs): + energies = np.array([float(g["energy_per_atom"][0]) for g in graphs], dtype=np.float32) + return float(np.mean(energies)), float(max(np.std(energies), 1e-8)) + + +def _run_epoch(model, loader, energy_mean, energy_std, force_weight, optimizer=None): + is_train = optimizer is not None + model.train() if is_train else model.eval() + + total_loss, total_e_loss, total_f_loss, count = 0.0, 0.0, 0.0, 0 + + for batch in tqdm(loader, leave=False): + if is_train: + optimizer.clear_grad() + + losses, e_losses, f_losses = [], [], [] + + for graph in batch: + pos = graph["pos"].detach().clone() + pos.stop_gradient = False + + out = model({ + "z": graph["z"], + "pos": pos, + "edge_index": graph["edge_index"], + }) + pred_total = out["total_energy"].reshape([1]) + + num_atoms = paddle.cast(graph["num_nodes"].reshape([1]), "float32") + pred_per_atom = pred_total / num_atoms + true_per_atom = graph["energy_per_atom"].reshape([1]) + + true_norm = (true_per_atom - energy_mean) / energy_std + pred_norm = (pred_per_atom - energy_mean) / energy_std + e_loss = F.mse_loss(pred_norm, true_norm) + + pred_forces = -paddle.grad( + outputs=[pred_total], + inputs=[pos], + create_graph=False, + retain_graph=False, + )[0] + f_loss = F.mse_loss(pred_forces, graph["forces"]) + + loss = e_loss + force_weight * f_loss + losses.append(loss) + e_losses.append(e_loss.detach()) + f_losses.append(f_loss.detach()) + + loss = paddle.stack(losses).mean() + if is_train: + loss.backward() + optimizer.step() + + total_loss += float(loss.item()) + total_e_loss += float(paddle.stack(e_losses).mean().item()) + total_f_loss += float(paddle.stack(f_losses).mean().item()) + count += 1 + + return { + "loss": total_loss / max(count, 1), + "e_loss": total_e_loss / max(count, 1), + "f_loss": total_f_loss / max(count, 1), + } + + +def build_model(config): + model_cfg = config["model"] + return sevennet_model.PurePaddleSevenNet( + num_species=model_cfg.get("num_species", 100), + hidden_dim=model_cfg.get("hidden_dim", 128), + num_message_layers=model_cfg.get("num_message_layers", 4), + num_rbf=model_cfg.get("num_rbf", 32), + cutoff=model_cfg.get("cutoff", 5.0), + ) + + +def build_optimizer(config, model): + optim_cfg = config["optimizer"] + lr = optim_cfg.get("lr", 0.001) + return paddle.optimizer.Adam(learning_rate=lr, parameters=model.parameters()) + + +def main(): + args = parse_args() + config = load_config(args.config) + + # 设置种子 + seed = config["trainer"].get("seed", 42) + random.seed(seed) + np.random.seed(seed) + paddle.seed(seed) + + # 设置设备 + device = config["trainer"].get("device", "auto") + if device == "auto": + device = "gpu" if paddle.is_compiled_with_cuda() else "cpu" + paddle.set_device(device) + print(f"Using device: {device}") + + # 创建输出目录 + save_dir = args.output_dir or config["log"].get("save_dir", "./output/sevennet") + os.makedirs(save_dir, exist_ok=True) + print(f"Output directory: {save_dir}") + + # 加载数据 + print("Loading dataset...") + from ase.io import read + data_path = config["dataset"]["path"] + + # 处理相对路径 - 支持仓库内路径 + if not os.path.isabs(data_path): + data_path = os.path.join(PROJECT_ROOT, data_path) + + # 如果路径不存在,使用内置测试数据 + if not os.path.exists(data_path): + print(f"Warning: Dataset not found at {data_path}, using built-in test data") + data_path = os.path.join(PROJECT_ROOT, "interatomic_potentials/example_data", "hfo2.extxyz") + + atoms_list = read(data_path, index=":") + if not isinstance(atoms_list, list): + atoms_list = [atoms_list] + print(f"Loaded {len(atoms_list)} structures") + + # 构建图数据集 + cutoff = config["dataset"]["cutoff"] + graphs = [_atoms_to_graph(atoms, cutoff) for atoms in tqdm(atoms_list)] + train_graphs, valid_graphs = _split_graphs( + graphs, + valid_ratio=config["dataset"].get("valid_ratio", 0.1), + seed=seed + ) + print(f"Train: {len(train_graphs)}, Valid: {len(valid_graphs)}") + + # 计算统计信息 + energy_mean, energy_std = _compute_stats(train_graphs) + print(f"Energy mean: {energy_mean:.6f}, std: {energy_std:.6f}") + + # 构建数据加载器 + train_dataset = GraphDataset(train_graphs) + valid_dataset = GraphDataset(valid_graphs) + + train_loader = DataLoader( + train_dataset, + batch_size=config["dataloader"]["batch_size"], + shuffle=config["dataloader"].get("shuffle", True), + collate_fn=_collate_graphs, + num_workers=config["dataloader"].get("num_workers", 0), + ) + valid_loader = DataLoader( + valid_dataset, + batch_size=config["dataloader"]["batch_size"], + shuffle=False, + collate_fn=_collate_graphs, + num_workers=config["dataloader"].get("num_workers", 0), + ) + + # 构建模型和优化器 + model = build_model(config) + optimizer = build_optimizer(config, model) + print(f"Model parameters: {sum(p.size for p in model.parameters())}") + + # 训练循环 + epoch = config["trainer"]["epoch"] + force_weight = config["loss"].get("force_loss_weight", 0.1) + best_valid = float("inf") + + for ep in range(1, epoch + 1): + print(f"\nEpoch {ep}/{epoch}") + + train_stats = _run_epoch(model, train_loader, energy_mean, energy_std, force_weight, optimizer) + valid_stats = _run_epoch(model, valid_loader, energy_mean, energy_std, force_weight) + + print(f"Train: loss={train_stats['loss']:.6f}, e_loss={train_stats['e_loss']:.6f}, f_loss={train_stats['f_loss']:.6f}") + print(f"Valid: loss={valid_stats['loss']:.6f}, e_loss={valid_stats['e_loss']:.6f}, f_loss={valid_stats['f_loss']:.6f}") + + # 保存模型 + save_obj = { + "model_state_dict": model.state_dict(), + "energy_mean": energy_mean, + "energy_std": energy_std, + "config": config, + } + + if valid_stats["loss"] < best_valid: + best_valid = valid_stats["loss"] + paddle.save(save_obj, os.path.join(save_dir, "best_model.pdparams")) + print(f"Best model saved") + + if ep % config["log"].get("save_interval", 1) == 0: + paddle.save(save_obj, os.path.join(save_dir, f"model_epoch_{ep}.pdparams")) + + paddle.save(save_obj, os.path.join(save_dir, "model_final.pdparams")) + print(f"\nTraining finished. Final model saved to {save_dir}") + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/interatomic_potentials/train.py b/interatomic_potentials/train.py index 09fd8880..ba61ec4e 100644 --- a/interatomic_potentials/train.py +++ b/interatomic_potentials/train.py @@ -148,4 +148,4 @@ time_info, loss_info, metric_info = trainer.eval(val_loader) if config["Global"].get("do_test", False): logger.info("Evaluating on test set") - time_info, loss_info, metric_info = trainer.eval(test_loader) + time_info, loss_info, metric_info = trainer.eval(test_loader) \ No newline at end of file diff --git a/ppmat/models/__init__.py b/ppmat/models/__init__.py index 95d73232..81f93f21 100644 --- a/ppmat/models/__init__.py +++ b/ppmat/models/__init__.py @@ -38,6 +38,7 @@ from ppmat.models.mattergen.mattergen import MatterGen from ppmat.models.mattergen.mattergen import MatterGenWithCondition from ppmat.models.mattersim.m3gnet import M3GNet +from ppmat.models.sevennet.sevennet_model import PurePaddleSevenNet from ppmat.models.mattersim.m3gnet_graph_converter import M3GNetGraphConvertor from ppmat.models.megnet.megnet import MEGNetPlus from ppmat.models.infgcn.infgcn import InfGCN @@ -67,6 +68,7 @@ "DiffNMR", "InfGCN", "MatENO", + "PurePaddleSevenNet", ] # Warning: The key of the dictionary must be consistent with the file name of the value @@ -111,6 +113,12 @@ "mattergen_ml2ddb": "https://paddle-org.bj.bcebos.com/paddlematerial/workflow/ml2ddb/mattergen_ml2ddb.zip", "mattergen_ml2ddb_chemical_system": "https://paddle-org.bj.bcebos.com/paddlematerial/workflow/ml2ddb/mattergen_ml2ddb_chemical_system.zip", "mattergen_ml2ddb_space_group": "https://paddle-org.bj.bcebos.com/paddlematerial/workflow/ml2ddb/mattergen_ml2ddb_space_group.zip", + # SevenNet pretrained models - URLs to be provided by official + # "sevennet_omni": "https://paddle-org.bj.bcebos.com/paddlematerial/checkpoints/interatomic_potentials/sevennet/sevennet_omni.zip", + # "sevennet_mf_ompa": "https://paddle-org.bj.bcebos.com/paddlematerial/checkpoints/interatomic_potentials/sevennet/sevennet_mf_ompa.zip", + # "sevennet_omat": "https://paddle-org.bj.bcebos.com/paddlematerial/checkpoints/interatomic_potentials/sevennet/sevennet_omat.zip", + # "sevennet_l3i5": "https://paddle-org.bj.bcebos.com/paddlematerial/checkpoints/interatomic_potentials/sevennet/sevennet_l3i5.zip", + # "sevennet_0": "https://paddle-org.bj.bcebos.com/paddlematerial/checkpoints/interatomic_potentials/sevennet/sevennet_0.zip", } diff --git a/ppmat/models/sevennet/__init__.py b/ppmat/models/sevennet/__init__.py new file mode 100644 index 00000000..8580803c --- /dev/null +++ b/ppmat/models/sevennet/__init__.py @@ -0,0 +1,3 @@ +from .sevennet_model import PurePaddleSevenNet + +__all__ = ["PurePaddleSevenNet"] \ No newline at end of file diff --git a/ppmat/models/sevennet/checkpoints/sevennet_hfo2_best.pdparams b/ppmat/models/sevennet/checkpoints/sevennet_hfo2_best.pdparams new file mode 100644 index 00000000..6bcb9824 Binary files /dev/null and b/ppmat/models/sevennet/checkpoints/sevennet_hfo2_best.pdparams differ diff --git a/ppmat/models/sevennet/checkpoints/sevennet_hfo2_final.pdparams b/ppmat/models/sevennet/checkpoints/sevennet_hfo2_final.pdparams new file mode 100644 index 00000000..6bcb9824 Binary files /dev/null and b/ppmat/models/sevennet/checkpoints/sevennet_hfo2_final.pdparams differ diff --git a/ppmat/models/sevennet/pretrained_model_registry.py b/ppmat/models/sevennet/pretrained_model_registry.py new file mode 100644 index 00000000..a9b3a2e0 --- /dev/null +++ b/ppmat/models/sevennet/pretrained_model_registry.py @@ -0,0 +1,33 @@ +# SevenNet Pretrained Model Registry +# +# When official provides URLs, add them here: +# Format: +# "sevennet_${model_name}": "https://official-url/model.zip" +# +# Example: +# "sevennet_mptrj": "https://paddle-org.bj.bcebos.com/paddlematerial/checkpoints/interatomic_potentials/sevennet/sevennet_mptrj.zip" + +# ==================== To be filled by official ==================== + +# SevenNet-Omni (Recommended) - Universal potential trained on 15 datasets +# sevennet_omni: "" + +# SevenNet-MF-ompa - Multi-fidelity model +# sevennet_mf_ompa: "" + +# SevenNet-omat - OMat24 dataset +# sevennet_omat: "" + +# SevenNet-l3i5 - MPtrj dataset +# sevennet_l3i5: "" + +# SevenNet-0 - Early version, fastest inference +# sevennet_0: "" + +# ==================== Dataset Information ==================== +# Training Datasets: +# - MPtrj (Materials Project Trajectory): https://figshare.com/articles/dataset/Materials_Project_Trjectory_MPtrj_Dataset/23713842 +# - OMat24: https://huggingface.co/datasets/fairchem/OMAT24 +# - sAlex: https://huggingface.co/datasets/fairchem/OMAT24 +# +# For more details, see: https://sevennet.readthedocs.io/en/latest/user_guide/pretrained.html \ No newline at end of file diff --git a/ppmat/models/sevennet/sevennet_model.py b/ppmat/models/sevennet/sevennet_model.py new file mode 100644 index 00000000..59f9465c --- /dev/null +++ b/ppmat/models/sevennet/sevennet_model.py @@ -0,0 +1,201 @@ +import paddle +import paddle.nn as nn +import paddle.nn.functional as F + + +class GaussianRBF(nn.Layer): + def __init__(self, num_basis=32, cutoff=5.0, gamma=None): + super().__init__() + centers = paddle.linspace(0.0, cutoff, num_basis) + self.register_buffer("centers", centers) + if gamma is None: + gamma = 10.0 / max(cutoff, 1e-6) + self.gamma = gamma + + def forward(self, distances): + diff = distances - self.centers.reshape([1, -1]) + return paddle.exp(-self.gamma * diff * diff) + + +class PolynomialCutoff(nn.Layer): + def __init__(self, cutoff=5.0, p=6): + super().__init__() + self.cutoff = cutoff + self.p = p + + def forward(self, distances): + x = distances / self.cutoff + x = paddle.clip(x, 0.0, 1.0) + weight = 1.0 - 6.0 * x**5 + 15.0 * x**4 - 10.0 * x**3 + mask = (distances < self.cutoff).astype("float32") + return weight * mask + + +class MessageBlock(nn.Layer): + def __init__(self, hidden_dim, rbf_dim): + super().__init__() + self.msg_mlp = nn.Sequential( + nn.Linear(hidden_dim * 2 + rbf_dim, hidden_dim), + nn.Silu(), + nn.Linear(hidden_dim, hidden_dim), + nn.Silu(), + ) + self.upd_mlp = nn.Sequential( + nn.Linear(hidden_dim * 2, hidden_dim), + nn.Silu(), + nn.Linear(hidden_dim, hidden_dim), + ) + self.norm = nn.LayerNorm(hidden_dim) + + def forward(self, h, edge_src, edge_dst, rbf, edge_weight): + src_h = paddle.gather(h, edge_src, axis=0) + dst_h = paddle.gather(h, edge_dst, axis=0) + + msg_in = paddle.concat([src_h, dst_h, rbf], axis=-1) + msg = self.msg_mlp(msg_in) + msg = msg * edge_weight + + # 使用 scatter_nd_add 实现消息聚合 + num_nodes = h.shape[0] + agg = paddle.zeros([num_nodes, h.shape[1]], dtype=h.dtype) + edge_dst_2d = edge_dst.unsqueeze(1) + agg = paddle.scatter_nd_add(agg, edge_dst_2d, msg) + + upd_in = paddle.concat([h, agg], axis=-1) + dh = self.upd_mlp(upd_in) + return self.norm(h + dh) + + +class PurePaddleSevenNet(nn.Layer): + def __init__( + self, + num_species=100, + hidden_dim=128, + num_message_layers=4, + num_rbf=32, + cutoff=5.0, + ): + super().__init__() + self.cutoff = cutoff + self.embedding = nn.Embedding(num_species, hidden_dim) + self.rbf = GaussianRBF(num_basis=num_rbf, cutoff=cutoff) + self.cutoff_fn = PolynomialCutoff(cutoff=cutoff, p=6) + self.blocks = nn.LayerList( + [MessageBlock(hidden_dim, num_rbf) for _ in range(num_message_layers)] + ) + self.energy_head = nn.Sequential( + nn.Linear(hidden_dim, hidden_dim), + nn.Silu(), + nn.Linear(hidden_dim, 1), + ) + + def forward(self, graph): + z = graph["z"] + pos = graph["pos"] + edge_index = graph["edge_index"] + + edge_src = edge_index[0] + edge_dst = edge_index[1] + + h = self.embedding(z) + + src_pos = paddle.gather(pos, edge_src, axis=0) + dst_pos = paddle.gather(pos, edge_dst, axis=0) + edge_vec = src_pos - dst_pos + dist = paddle.sqrt(paddle.sum(edge_vec * edge_vec, axis=-1, keepdim=True) + 1e-12) + + rbf = self.rbf(dist) + edge_weight = self.cutoff_fn(dist) + + for block in self.blocks: + h = block(h, edge_src, edge_dst, rbf, edge_weight) + + atomic_energy = self.energy_head(h) + total_energy = paddle.sum(atomic_energy) + + return { + "total_energy": total_energy, + "atomic_energy": atomic_energy, + } + + def predict(self, data): + """Predict energy and forces for a batch of structures. + + This method provides a compatible interface for PotentialPredictor. + + Args: + data: Dictionary containing graph data with keys: + - z: atomic numbers (Tensor) + - pos: positions (Tensor) + - edge_index: edge connectivity (Tensor) + + Returns: + Dictionary with predictions: + - energy: total energy per structure (float) + - forces: atomic forces (numpy array) + """ + # Support both direct graph dict and data object with graph attribute + if hasattr(data, 'graph'): + graph = data.graph + else: + graph = data + + # Handle batch dimension + if isinstance(graph, list): + results = [] + for g in graph: + result = self.forward(g) + # Convert to numpy and format output + prediction = { + "energy": float(result["total_energy"].numpy()), + "atomic_energy": result["atomic_energy"].numpy(), + } + results.append(prediction) + return results + else: + result = self.forward(graph) + return { + "energy": float(result["total_energy"].numpy()), + "atomic_energy": result["atomic_energy"].numpy(), + } + + def compute_forces(self, graph, eps=1e-4): + """Compute forces by finite difference of energy w.r.t. positions. + + Uses finite difference instead of autograd because Paddle's + scatter_nd_add backward pass has numerical instability with + deep message-passing networks. + + Args: + graph: Dictionary containing graph data + eps: Finite difference step size + + Returns: + numpy array of forces with shape [num_atoms, 3] + """ + import numpy as np + + positions_np = graph["pos"].numpy() + z = graph["z"] + edge_index = graph["edge_index"] + num_atoms = positions_np.shape[0] + forces = np.zeros_like(positions_np) + + for i in range(num_atoms): + for j in range(3): + pos_plus = positions_np.copy() + pos_plus[i, j] += eps + pos_minus = positions_np.copy() + pos_minus[i, j] -= eps + + pos_p = paddle.to_tensor(pos_plus, dtype="float32") + with paddle.no_grad(): + e_plus = float(self.forward({"z": z, "pos": pos_p, "edge_index": edge_index})["total_energy"].numpy()) + + pos_m = paddle.to_tensor(pos_minus, dtype="float32") + with paddle.no_grad(): + e_minus = float(self.forward({"z": z, "pos": pos_m, "edge_index": edge_index})["total_energy"].numpy()) + + forces[i, j] = -(e_plus - e_minus) / (2 * eps) + + return forces \ No newline at end of file diff --git a/test/test_sevennet.py b/test/test_sevennet.py new file mode 100644 index 00000000..26a14de9 --- /dev/null +++ b/test/test_sevennet.py @@ -0,0 +1,203 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +""" +Test file for SevenNet model + +Usage: + python test/test_sevennet.py + pytest test/test_sevennet.py -v +""" + +import sys +import os + +# 添加项目路径 +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +import paddle +import numpy as np + + +def import_model(): + """Import SevenNet model dynamically to avoid ppmat dependencies""" + import importlib.util + + project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + model_path = os.path.join(project_root, "ppmat/models/sevennet/sevennet_model.py") + + spec = importlib.util.spec_from_file_location("sevennet_model", model_path) + sevennet_model = importlib.util.module_from_spec(spec) + spec.loader.exec_module(sevennet_model) + + return sevennet_model.PurePaddleSevenNet + + +def test_forward_pass(): + """Test basic forward pass""" + PurePaddleSevenNet = import_model() + + model = PurePaddleSevenNet( + num_species=100, + hidden_dim=16, + num_message_layers=2, + num_rbf=16, + cutoff=4.0, + ) + + z = paddle.to_tensor([1, 8, 1], dtype="int64") + pos = paddle.to_tensor([[0, 0, 0], [1, 0, 0], [0, 1, 0]], dtype="float32") + edge_index = paddle.to_tensor([[0, 1, 1, 2, 0, 2], [1, 0, 2, 1, 2, 0]], dtype="int64") + + graph = {"z": z, "pos": pos, "edge_index": edge_index} + out = model(graph) + + assert "total_energy" in out, "Missing total_energy" + assert "atomic_energy" in out, "Missing atomic_energy" + assert out["total_energy"].shape == (), f"Expected scalar, got {out['total_energy'].shape}" + assert out["atomic_energy"].shape == (3, 1), f"Expected (3, 1), got {out['atomic_energy'].shape}" + assert not paddle.isnan(out["total_energy"]), "Total energy is NaN" + assert not paddle.isinf(out["total_energy"]), "Total energy is Inf" + + print("✓ test_forward_pass passed") + + +def test_backward_pass(): + """Test backward pass""" + PurePaddleSevenNet = import_model() + + model = PurePaddleSevenNet( + num_species=100, + hidden_dim=16, + num_message_layers=2, + num_rbf=16, + cutoff=4.0, + ) + + z = paddle.to_tensor([1, 8, 1], dtype="int64") + pos = paddle.to_tensor([[0, 0, 0], [1, 0, 0], [0, 1, 0]], dtype="float32") + pos.stop_gradient = False + edge_index = paddle.to_tensor([[0, 1, 1, 2, 0, 2], [1, 0, 2, 1, 2, 0]], dtype="int64") + + graph = {"z": z, "pos": pos, "edge_index": edge_index} + out = model(graph) + loss = out["total_energy"] + loss.backward() + + assert pos.grad is not None, "Position gradient is None" + assert pos.grad.shape == pos.shape, "Gradient shape mismatch" + assert not paddle.any(paddle.isnan(pos.grad)), "Gradient contains NaN" + assert not paddle.any(paddle.isinf(pos.grad)), "Gradient contains Inf" + + print("✓ test_backward_pass passed") + + +def test_force_computation(): + """Test force computation via gradient""" + PurePaddleSevenNet = import_model() + + model = PurePaddleSevenNet( + num_species=100, + hidden_dim=16, + num_message_layers=2, + num_rbf=16, + cutoff=4.0, + ) + + z = paddle.to_tensor([1, 8, 1], dtype="int64") + pos = paddle.to_tensor([[0, 0, 0], [1, 0, 0], [0, 1, 0]], dtype="float32") + pos.stop_gradient = False + edge_index = paddle.to_tensor([[0, 1, 1, 2, 0, 2], [1, 0, 2, 1, 2, 0]], dtype="int64") + + graph = {"z": z, "pos": pos, "edge_index": edge_index} + out = model(graph) + forces = -paddle.grad( + outputs=[out["total_energy"]], + inputs=[pos], + create_graph=False, + retain_graph=False, + )[0] + + assert forces.shape == pos.shape, f"Forces shape mismatch: {forces.shape} vs {pos.shape}" + assert not paddle.any(paddle.isnan(forces)), "Forces contain NaN" + assert not paddle.any(paddle.isinf(forces)), "Forces contain Inf" + + print("✓ test_force_computation passed") + + +def test_model_reproducibility(): + """Test model reproducibility with same seed""" + PurePaddleSevenNet = import_model() + + paddle.seed(42) + model1 = PurePaddleSevenNet( + num_species=100, + hidden_dim=16, + num_message_layers=2, + num_rbf=16, + cutoff=4.0, + ) + + paddle.seed(42) + model2 = PurePaddleSevenNet( + num_species=100, + hidden_dim=16, + num_message_layers=2, + num_rbf=16, + cutoff=4.0, + ) + + z = paddle.to_tensor([1, 8, 1], dtype="int64") + pos = paddle.to_tensor([[0, 0, 0], [1, 0, 0], [0, 1, 0]], dtype="float32") + edge_index = paddle.to_tensor([[0, 1, 1, 2, 0, 2], [1, 0, 2, 1, 2, 0]], dtype="int64") + graph = {"z": z, "pos": pos, "edge_index": edge_index} + + out1 = model1(graph) + out2 = model2(graph) + + diff = paddle.abs(out1["total_energy"] - out2["total_energy"]) + assert diff < 1e-6, f"Models not reproducible: diff={diff}" + + print("✓ test_model_reproducibility passed") + + +def test_empty_graph(): + """Test behavior with empty graph""" + PurePaddleSevenNet = import_model() + + model = PurePaddleSevenNet( + num_species=100, + hidden_dim=16, + num_message_layers=2, + num_rbf=16, + cutoff=4.0, + ) + + z = paddle.to_tensor([], dtype="int64") + pos = paddle.to_tensor([], dtype="float32").reshape([0, 3]) + edge_index = paddle.to_tensor([[], []], dtype="int64") + + graph = {"z": z, "pos": pos, "edge_index": edge_index} + out = model(graph) + + assert out["total_energy"].shape == () + + print("✓ test_empty_graph passed") + + +if __name__ == "__main__": + print("Running SevenNet tests...\n") + + try: + test_forward_pass() + test_backward_pass() + test_force_computation() + test_model_reproducibility() + test_empty_graph() + + print("\n✅ All tests passed!") + sys.exit(0) + except Exception as e: + print(f"\n❌ Test failed: {e}") + import traceback + traceback.print_exc() + sys.exit(1) \ No newline at end of file diff --git a/test/test_sevennet_simple.py b/test/test_sevennet_simple.py new file mode 100644 index 00000000..7de8849f --- /dev/null +++ b/test/test_sevennet_simple.py @@ -0,0 +1,74 @@ +#!/usr/bin/env python +# -*- coding: utf-8 -*- +"""Simple test for SevenNet model""" +import sys +import os +import paddle +import numpy as np + +# 添加项目路径 +PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +sys.path.insert(0, PROJECT_ROOT) + +# 直接导入模型文件 +import importlib.util +model_path = os.path.join(PROJECT_ROOT, "ppmat/models/sevennet/sevennet_model.py") +spec = importlib.util.spec_from_file_location("sevennet_model", model_path) +sevennet_model = importlib.util.module_from_spec(spec) +spec.loader.exec_module(sevennet_model) + + +def test_sevennet(): + print("Testing SevenNet model...") + + # 创建模型 + model = sevennet_model.PurePaddleSevenNet( + num_species=100, + hidden_dim=32, + num_message_layers=3, + num_rbf=32, + cutoff=4.0, + ) + + # 创建测试输入 + z = paddle.to_tensor([1, 8, 1], dtype="int64") + pos = paddle.to_tensor([[0, 0, 0], [1, 0, 0], [0, 1, 0]], dtype="float32") + edge_index = paddle.to_tensor([[0, 1, 1, 2, 0, 2], [1, 0, 2, 1, 2, 0]], dtype="int64") + + graph = {"z": z, "pos": pos, "edge_index": edge_index} + + # 前向传播 + out = model(graph) + + print(f"✓ Forward pass OK") + print(f" Total energy: {float(out['total_energy'].item()):.6f} eV") + + # 测试梯度 + pos_grad = paddle.to_tensor([[0, 0, 0], [1, 0, 0], [0, 1, 0]], dtype="float32") + pos_grad.stop_gradient = False + graph_grad = {"z": z, "pos": pos_grad, "edge_index": edge_index} + + out_grad = model(graph_grad) + forces = -paddle.grad( + outputs=[out_grad["total_energy"]], + inputs=[pos_grad], + create_graph=False, + retain_graph=False, + )[0] + + print(f"✓ Gradient computation OK") + print(f" Forces shape: {forces.shape}") + + return True + + +if __name__ == "__main__": + try: + test_sevennet() + print("\n✅ All tests passed!") + sys.exit(0) + except Exception as e: + print(f"\n❌ Test failed: {e}") + import traceback + traceback.print_exc() + sys.exit(1)