From 1e46fa001c3f60b3523fa5dcd13412dc8237b3dc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ole=20Sch=C3=BCtt?= Date: Wed, 8 Jan 2025 22:13:53 +0100 Subject: [PATCH] PAO-ML: Stabilize test and add TZV2P basis set --- tools/docker/scripts/test_misc.sh | 2 +- tools/pao-ml/pao/model.py | 11 +++++++---- 2 files changed, 8 insertions(+), 5 deletions(-) diff --git a/tools/docker/scripts/test_misc.sh b/tools/docker/scripts/test_misc.sh index 8c2bc46ca5..bee1132d15 100755 --- a/tools/docker/scripts/test_misc.sh +++ b/tools/docker/scripts/test_misc.sh @@ -35,7 +35,7 @@ run_test ./tools/docker/generate_dockerfiles.py --check # Test pao-ml training. run_test ./tools/pao-ml/pao-train.py --kind=H --epochs=200 ./tools/pao-ml/example.pao run_test ./tools/pao-ml/pao-retrain.py --model="DZVP-MOLOPT-GTH-PAO4-H.pt" --epochs=200 ./tools/pao-ml/example.pao -run_test ./tools/pao-ml/pao-validate.py --threshold=1e-2 --model="DZVP-MOLOPT-GTH-PAO4-H.pt" ./tools/pao-ml/example.pao +run_test ./tools/pao-ml/pao-validate.py --threshold=1e-1 --model="DZVP-MOLOPT-GTH-PAO4-H.pt" ./tools/pao-ml/example.pao run_test ./tools/pao-ml/pao-validate.py --threshold=1e-6 --model="tests/QS/regtest-pao-5/DZVP-MOLOPT-GTH-PAO4-H.pt" ./tools/pao-ml/example.pao run_test ./tools/pao-ml/pao-validate.py --threshold=1e-5 --model="tests/QS/regtest-pao-5/DZVP-MOLOPT-GTH-PAO4-O.pt" ./tools/pao-ml/example.pao diff --git a/tools/pao-ml/pao/model.py b/tools/pao-ml/pao/model.py index d3e524484f..234ec4dc3d 100644 --- a/tools/pao-ml/pao/model.py +++ b/tools/pao-ml/pao/model.py @@ -51,12 +51,15 @@ class PaoModel(torch.nn.Module): self.cutoff = cutoff # Irreps of primary basis - assert prim_basis_name == "DZVP-MOLOPT-GTH" # TODO support more basis sets + # TODO: Export the specs directly from cp2k as part of the .pao files. prim_basis_specs = { - "O": "2x0e + 2x1o + 1x2e", # two s-shells, two p-shells, one d-shell - "H": "2x0e + 1x1o", # two s-shells, one p-shell + "DZVP-MOLOPT-GTH/H": "2x0e + 1x1o", # two s-shells, one p-shell + "DZVP-MOLOPT-GTH/O": "2x0e + 2x1o + 1x2e", # two s, two p, one d-shell + "TZV2P-MOLOPT-GGA-GTH-q1/H": "3x0e + 2x1o + 1x2e", + "TZV2P-MOLOPT-GGA-GTH-q6/O": "3x0e + 3x1o + 2x2e + 1x3o", } - prim_basis_irreps = e3nn.o3.Irreps(prim_basis_specs[kind_name]) + basis_specs_key = f"{prim_basis_name}/{kind_name}" + prim_basis_irreps = e3nn.o3.Irreps(prim_basis_specs[basis_specs_key]) assert self.prim_basis_size == prim_basis_irreps.dim # auxiliary Hamiltonian