Skip to content

Commit b5eb5ff

Browse files
committed
add normlized RMSE, no-jit flag
1 parent f5ffb19 commit b5eb5ff

3 files changed

Lines changed: 25 additions & 3 deletions

File tree

deepmd/entrypoints/test.py

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -160,7 +160,7 @@ def test(
160160
dp_random.seed(rand_seed % (2**32))
161161

162162
# init model
163-
dp = DeepEval(model, head=head)
163+
dp = DeepEval(model, head=head, no_jit=kwargs.get("no_jit", False))
164164

165165
for cc, system in enumerate(all_sys):
166166
log.info("# ---------------output of dp test--------------- ")
@@ -469,6 +469,16 @@ def test_ener(
469469
mae_f = mae(diff_f)
470470
rmse_f = rmse(diff_f)
471471
size_f = diff_f.size
472+
nrmse_e = (
473+
(rmse_e / np.sqrt(np.mean(test_data["energy"][:numb_test].reshape([-1, 1]) ** 2))) * 100
474+
if find_energy == 1
475+
else None
476+
)
477+
nrmse_f = (
478+
(rmse_f / np.sqrt(np.mean(test_data["force"][:numb_test] ** 2))) * 100
479+
if not out_put_spin and find_force == 1
480+
else None
481+
)
472482
if find_atom_pref == 1:
473483
atom_weight = test_data["atom_pref"][:numb_test]
474484
weight_sum = np.sum(atom_weight)
@@ -505,15 +515,19 @@ def test_ener(
505515
log.info(f"Energy RMSE : {rmse_e:e} eV")
506516
log.info(f"Energy MAE/Natoms : {mae_ea:e} eV")
507517
log.info(f"Energy RMSE/Natoms : {rmse_ea:e} eV")
518+
log.info(f"Energy NRMSE : {nrmse_e:.4f} %")
508519
dict_to_return["mae_e"] = (mae_e, energy.size)
509520
dict_to_return["mae_ea"] = (mae_ea, energy.size)
510521
dict_to_return["rmse_e"] = (rmse_e, energy.size)
511522
dict_to_return["rmse_ea"] = (rmse_ea, energy.size)
523+
dict_to_return["nrmse_e"] = (nrmse_e, energy.size)
512524
if not out_put_spin and find_force == 1:
513525
log.info(f"Force MAE : {mae_f:e} eV/Å")
514526
log.info(f"Force RMSE : {rmse_f:e} eV/Å")
527+
log.info(f"Force NRMSE : {nrmse_f:.4f} %")
515528
dict_to_return["mae_f"] = (mae_f, size_f)
516529
dict_to_return["rmse_f"] = (rmse_f, size_f)
530+
dict_to_return["nrmse_f"] = (nrmse_f, size_f)
517531
if find_atom_pref == 1:
518532
log.info(f"Force weighted MAE : {mae_fw:e} eV/Å")
519533
log.info(f"Force weighted RMSE: {rmse_fw:e} eV/Å")
@@ -661,9 +675,11 @@ def print_ener_sys_avg(avg: dict[str, float]) -> None:
661675
log.info(f"Energy RMSE : {avg['rmse_e']:e} eV")
662676
log.info(f"Energy MAE/Natoms : {avg['mae_ea']:e} eV")
663677
log.info(f"Energy RMSE/Natoms : {avg['rmse_ea']:e} eV")
678+
log.info(f"Energy NRMSE : {avg['nrmse_e']:.4f} %")
664679
if "rmse_f" in avg:
665680
log.info(f"Force MAE : {avg['mae_f']:e} eV/Å")
666681
log.info(f"Force RMSE : {avg['rmse_f']:e} eV/Å")
682+
log.info(f"Force NRMSE : {avg['nrmse_f']:.4f} %")
667683
if "rmse_fw" in avg:
668684
log.info(f"Force weighted MAE : {avg['mae_fw']:e} eV/Å")
669685
log.info(f"Force weighted RMSE: {avg['rmse_fw']:e} eV/Å")

deepmd/main.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -448,6 +448,12 @@ def main_parser() -> argparse.ArgumentParser:
448448
type=str,
449449
help="(Supported backend: PyTorch) Task head (alias: model branch) to test if in multi-task mode.",
450450
)
451+
parser_tst.add_argument(
452+
"--no-jit",
453+
action="store_true",
454+
default=False,
455+
help="(Supported backend: PyTorch) Disable JIT compilation when loading the model.",
456+
)
451457

452458
# * eval_desc script ***************************************************************
453459
parser_eval_desc = subparsers.add_parser(

deepmd/utils/weight_avg.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -26,15 +26,15 @@ def weighted_average(errors: list[dict[str, tuple[float, float]]]) -> dict:
2626
sum_siz = defaultdict(int)
2727
for err in errors:
2828
for kk, (ee, ss) in err.items():
29-
if kk.startswith("mae"):
29+
if kk.startswith("mae") or kk.startswith("nrmse"):
3030
sum_err[kk] += ee * ss
3131
elif kk.startswith("rmse"):
3232
sum_err[kk] += ee * ee * ss
3333
else:
3434
raise RuntimeError("unknown error type")
3535
sum_siz[kk] += ss
3636
for kk in sum_err.keys():
37-
if kk.startswith("mae"):
37+
if kk.startswith("mae") or kk.startswith("nrmse"):
3838
sum_err[kk] = sum_err[kk] / sum_siz[kk]
3939
elif kk.startswith("rmse"):
4040
sum_err[kk] = np.sqrt(sum_err[kk] / sum_siz[kk])

0 commit comments

Comments
 (0)