Download evaluate.py from OneScience-Group/eSEN: direct link, hf CLI and curl.
- Browser
- Download file 3.87 kB
-
https://huggingface.co/OneScience-Group/eSEN/resolve/main/evaluate.py
- Command line
-
hf download hf://OneScience-Group/eSEN/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/OneScience-Group/eSEN/resolve/main/evaluate.py
3.87 kB
| """Evaluate an eSEN checkpoint in physical units on an independent ASE DB.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import os | |
| from pathlib import Path | |
| os.environ.setdefault( | |
| "ONESCIENCE_ESEN_JD_PATH", | |
| os.path.join(os.path.dirname(__file__), "weight", "Jd.pt"), | |
| ) | |
| import torch | |
| from onescience.utils.esen.checkpoint import ESENCheckpointTransforms | |
| from onescience.utils.uma.common.utils import load_model_and_weights_from_checkpoint | |
| from finetune import _loader | |
| def _stats(error: torch.Tensor) -> dict[str, float]: | |
| error = error.detach().reshape(-1).double() | |
| return { | |
| "mae": float(error.abs().mean()), | |
| "rmse": float(error.square().mean().sqrt()), | |
| } | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--checkpoint", required=True) | |
| parser.add_argument("--data", required=True, help="Independent ASE DB/ASE-LMDB") | |
| parser.add_argument("--device", default="cuda") | |
| parser.add_argument("--batch-size", type=int, default=2) | |
| parser.add_argument("--workers", type=int, default=0) | |
| parser.add_argument("--max-samples", type=int) | |
| parser.add_argument("--seed", type=int, default=0) | |
| parser.add_argument("--output") | |
| args = parser.parse_args() | |
| device = torch.device(args.device) | |
| if device.type == "cuda" and not torch.cuda.is_available(): | |
| raise RuntimeError("CUDA/DCU was requested but is unavailable") | |
| import onescience.models.esen # noqa: F401 | |
| model = load_model_and_weights_from_checkpoint(args.checkpoint).to(device) | |
| model.eval() | |
| transforms = ESENCheckpointTransforms.from_checkpoint(args.checkpoint).to(device) | |
| loader = _loader( | |
| args.data, | |
| args.batch_size, | |
| args.workers, | |
| max_samples=args.max_samples, | |
| train=False, | |
| seed=args.seed, | |
| ) | |
| energy_errors = [] | |
| energy_per_atom_errors = [] | |
| force_errors = [] | |
| stress_errors = [] | |
| with torch.enable_grad(): | |
| for batch in loader: | |
| batch = batch.to(device) | |
| prediction = model(batch) | |
| pred_energy = transforms.denormalize_prediction("energy", prediction["energy"], batch) | |
| target_energy = batch.energy.reshape_as(pred_energy) | |
| energy_errors.append((pred_energy - target_energy).detach().cpu()) | |
| natoms = batch.natoms.to(pred_energy).reshape((-1,) + (1,) * (pred_energy.ndim - 1)) | |
| energy_per_atom_errors.append(((pred_energy - target_energy) / natoms).detach().cpu()) | |
| pred_forces = transforms.denormalize_prediction("forces", prediction["forces"], batch) | |
| target_forces = batch.forces.reshape_as(pred_forces) | |
| force_errors.append((pred_forces - target_forces).detach().cpu()) | |
| if "stress" in prediction and hasattr(batch, "stress"): | |
| pred_stress = transforms.denormalize_prediction("stress", prediction["stress"], batch) | |
| target_stress = batch.stress.reshape_as(pred_stress) | |
| stress_errors.append((pred_stress - target_stress).detach().cpu()) | |
| result = { | |
| "checkpoint": str(Path(args.checkpoint).expanduser()), | |
| "data": str(Path(args.data).expanduser()), | |
| "samples": len(loader.dataset), | |
| "energy_total_eV": _stats(torch.cat(energy_errors)), | |
| "energy_per_atom_eV": _stats(torch.cat(energy_per_atom_errors)), | |
| "forces_eV_per_A": _stats(torch.cat(force_errors)), | |
| } | |
| if stress_errors: | |
| result["stress_eV_per_A3"] = _stats(torch.cat(stress_errors)) | |
| print(json.dumps(result, indent=2, sort_keys=True)) | |
| if args.output: | |
| output = Path(args.output).expanduser() | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| output.write_text(json.dumps(result, indent=2, sort_keys=True) + "\n") | |
| if __name__ == "__main__": | |
| main() | |