#!/usr/bin/env python3 """ sevennet_md.py SevenNet + ASE molecular dynamics example. Default: - 216-atom Si diamond supercell (3x3x3 conventional cell) - SevenNet-Omni, PBE(+U) "mpa" task - NVT Langevin MD - CUDA is used automatically when available Examples -------- # 216-atom Si benchmark, 300 K, 10 ps python sevennet_md.py # Short benchmark python sevennet_md.py --steps 1000 --log-interval 20 # Use a VASP POSCAR python sevennet_md.py -i POSCAR --temperature 600 --steps 20000 # Use the lightweight SevenNet-Nano model python sevennet_md.py --model 7net-nano-5.5 --modal none # Force CPU (for comparison) python sevennet_md.py --device cpu --steps 1000 Requirements ------------ pip install sevenn ase PyTorch must be installed with a CUDA-enabled build to use an NVIDIA GPU. """ from __future__ import annotations import argparse import csv import time from pathlib import Path import numpy as np import torch from ase import Atoms, units from ase.build import bulk from ase.constraints import FixCom from ase.io import read, write from ase.io.trajectory import Trajectory from ase.md.langevin import Langevin from ase.md.velocitydistribution import MaxwellBoltzmannDistribution, Stationary from sevenn.calculator import SevenNetCalculator def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser( description="Run NVT molecular dynamics with SevenNet + ASE." ) p.add_argument( "-i", "--input", default=None, help="Input structure readable by ASE (POSCAR, CONTCAR, CIF, xyz, ...). " "If omitted, a 216-atom Si diamond cell is generated." ) p.add_argument( "--repeat", nargs=3, type=int, metavar=("NX", "NY", "NZ"), default=None, help="Repeat the input structure, e.g. --repeat 2 2 2." ) p.add_argument( "--model", default="7net-omni", help="SevenNet model keyword or checkpoint path (default: 7net-omni)." ) p.add_argument( "--modal", default="mpa", help="Task/modal for multi-task models (default: mpa). " "Use 'none' for single-task models such as 7net-nano-5.5." ) p.add_argument( "--device", choices=("auto", "cuda", "cpu"), default="auto", help="Device used by SevenNet (default: auto)." ) p.add_argument( "--accelerator", choices=("none", "cueq", "flash", "oeq"), default="none", help="Optional SevenNet tensor-product accelerator." ) p.add_argument("--temperature", type=float, default=300.0, help="Target temperature in K (default: 300).") p.add_argument("--timestep", type=float, default=1.0, help="MD timestep in fs (default: 1.0).") p.add_argument("--steps", type=int, default=10000, help="Number of MD steps (default: 10000 = 10 ps at 1 fs).") p.add_argument( "--friction", type=float, default=0.01, help="Langevin friction coefficient in fs^-1 (default: 0.01)." ) p.add_argument("--seed", type=int, default=12345, help="Random seed for initial velocities (default: 12345).") p.add_argument("--log-interval", type=int, default=100, help="Write thermodynamic data every N steps (default: 100).") p.add_argument("--traj-interval", type=int, default=100, help="Write trajectory every N steps (default: 100).") p.add_argument("--prefix", default="sevennet_md", help="Output filename prefix (default: sevennet_md).") return p.parse_args() def make_atoms(args: argparse.Namespace) -> Atoms: if args.input is None: # ASE conventional cubic diamond Si cell has 8 atoms. # 3x3x3 => 216 atoms. atoms = bulk("Si", "diamond", a=5.431, cubic=True).repeat((3, 3, 3)) source = "generated Si diamond 3x3x3" else: atoms = read(args.input) source = args.input if args.repeat is not None: atoms = atoms.repeat(tuple(args.repeat)) print(f"Structure : {source}") print(f"Atoms : {len(atoms)}") print(f"Formula : {atoms.get_chemical_formula()}") print(f"PBC : {atoms.pbc.tolist()}") if np.any(atoms.pbc): print(f"Cell volume: {atoms.get_volume():.3f} A^3") return atoms def make_calculator(args: argparse.Namespace) -> SevenNetCalculator: kwargs = { "model": args.model, "device": args.device, } if args.modal.lower() not in ("none", "null", "-"): kwargs["modal"] = args.modal if args.accelerator == "cueq": kwargs["enable_cueq"] = True elif args.accelerator == "flash": kwargs["enable_flash"] = True elif args.accelerator == "oeq": kwargs["enable_oeq"] = True return SevenNetCalculator(**kwargs) def print_device_info(args: argparse.Namespace) -> None: print("\n=== Runtime ===") print(f"PyTorch : {torch.__version__}") print(f"CUDA avail: {torch.cuda.is_available()}") if torch.cuda.is_available(): print(f"CUDA : {torch.version.cuda}") print(f"GPU : {torch.cuda.get_device_name(0)}") prop = torch.cuda.get_device_properties(0) print(f"VRAM : {prop.total_memory / 1024**3:.2f} GiB") print(f"Requested : {args.device}") print(f"Model : {args.model}") print(f"Modal : {args.modal}") print(f"Accelerator: {args.accelerator}") def main() -> None: args = parse_args() if args.steps <= 0: raise ValueError("--steps must be > 0") if args.timestep <= 0: raise ValueError("--timestep must be > 0") if args.temperature < 0: raise ValueError("--temperature must be >= 0") if args.log_interval <= 0 or args.traj_interval <= 0: raise ValueError("Output intervals must be > 0") print_device_info(args) atoms = make_atoms(args) print("\nLoading SevenNet model ...") t_model = time.perf_counter() atoms.calc = make_calculator(args) # Trigger one calculation here so model loading / first CUDA kernel setup # is separated from the MD timing. epot0 = atoms.get_potential_energy() forces0 = atoms.get_forces() load_elapsed = time.perf_counter() - t_model print(f"Initial Epot : {epot0:.6f} eV") print(f"Max |F| : {np.linalg.norm(forces0, axis=1).max():.6f} eV/A") print(f"Init time : {load_elapsed:.3f} s") # Initial velocities. np.random.seed(args.seed) MaxwellBoltzmannDistribution(atoms, temperature_K=args.temperature) Stationary(atoms) # Recommended replacement for Langevin's deprecated fixcm=True behavior. old_constraints = list(atoms.constraints) atoms.set_constraint(old_constraints + [FixCom()]) prefix = Path(args.prefix) traj_path = prefix.with_suffix(".traj") csv_path = prefix.with_suffix(".csv") final_path = prefix.parent / f"{prefix.name}_final.xyz" traj = Trajectory(str(traj_path), "w", atoms) dyn = Langevin( atoms, timestep=args.timestep * units.fs, temperature_K=args.temperature, friction=args.friction / units.fs, fixcm=False, ) csv_file = open(csv_path, "w", newline="", encoding="utf-8") writer = csv.writer(csv_file) writer.writerow([ "step", "time_fs", "temperature_K", "Epot_eV", "Ekin_eV", "Etot_eV", "Epot_eV_per_atom", "Ekin_eV_per_atom" ]) def log_status() -> None: step = dyn.nsteps epot = atoms.get_potential_energy() ekin = atoms.get_kinetic_energy() temp = atoms.get_temperature() nat = len(atoms) writer.writerow([ step, step * args.timestep, temp, epot, ekin, epot + ekin, epot / nat, ekin / nat, ]) csv_file.flush() print( f"step={step:7d} " f"t={step * args.timestep:10.2f} fs " f"T={temp:8.2f} K " f"Epot/N={epot / nat:12.6f} eV " f"Etot/N={(epot + ekin) / nat:12.6f} eV" ) dyn.attach(log_status, interval=args.log_interval) dyn.attach(traj.write, interval=args.traj_interval) print("\n=== MD ===") print(f"Ensemble : NVT (Langevin)") print(f"Temperature: {args.temperature:.1f} K") print(f"Timestep : {args.timestep:.3f} fs") print(f"Friction : {args.friction:.5f} fs^-1") print(f"Steps : {args.steps}") print(f"Total time : {args.steps * args.timestep / 1000.0:.3f} ps") print() # Log/write initial state explicitly. log_status() traj.write() if torch.cuda.is_available(): torch.cuda.synchronize() t0 = time.perf_counter() try: dyn.run(args.steps) finally: if torch.cuda.is_available(): torch.cuda.synchronize() elapsed = time.perf_counter() - t0 traj.close() csv_file.close() write(final_path, atoms) steps_per_s = args.steps / elapsed atom_steps_per_s = len(atoms) * steps_per_s simulated_ps = args.steps * args.timestep / 1000.0 ns_per_day = args.timestep * steps_per_s * 86400.0 / 1.0e6 print("\n=== Performance ===") print(f"Elapsed : {elapsed:.3f} s") print(f"MD steps/s : {steps_per_s:.3f}") print(f"atom-steps/s : {atom_steps_per_s:.1f}") print(f"Simulated time: {simulated_ps:.4f} ps") print(f"Throughput : {ns_per_day:.4f} ns/day") print("\n=== Output ===") print(f"Trajectory : {traj_path}") print(f"Thermo CSV : {csv_path}") print(f"Final XYZ : {final_path}") if __name__ == "__main__": main()