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()
|