File size: 1,530 Bytes
bf314e8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from torch.utils.data import DataLoader

# 自动定位 UMA 旋转基文件 Jd.pt
import os
_REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
_JD_PATH = os.path.join(_REPO_ROOT, "weight", "Jd.pt")
if os.path.isfile(_JD_PATH):
    os.environ.setdefault("ONESCIENCE_UMA_JD_PATH", _JD_PATH)

from onescience.datapipes.materials.custom_stack.core.atomic_data import (
    atomicdata_list_to_batch,
)
from onescience.datapipes.materials.custom_stack.storage.ase_datasets import AseDBDataset
from onescience.utils.uma.units.mlip_unit import load_predict_unit


def main() -> None:
    # Update these two paths before running.
    db_path = "../dataset/omat24/val/rattled-300-subsampled/data.aselmdb"
    checkpoint_path = "../weight/uma-s-1p1_converted.pt"

    dataset = AseDBDataset(
        config={
            "src": db_path,
            "a2g_args": {"task_name": "omat"},
        }
    )

    loader = DataLoader(
        dataset,
        batch_size=16,
        collate_fn=atomicdata_list_to_batch,
    )

    predictor = load_predict_unit(checkpoint_path, device="cuda")

    for i, batch in enumerate(loader):
        preds = predictor.predict(batch)

        for j in range(len(preds["energy"])):
            energy = preds["energy"][j].item()
            forces = preds["forces"][batch.batch == j].cpu().numpy()

            print(f"\\n[Batch {i} | Structure {j}]")
            print("Predicted energy:", energy)
            print("Predicted forces:\\n", forces)


if __name__ == "__main__":
    main()