File size: 3,866 Bytes
fecdc11
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
"""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()