Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions docs/source/user_guide/benchmarks/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -17,4 +17,5 @@ Benchmarks
tm_complexes
conformers
molecular_dynamics
molecular_reactions
defect
45 changes: 45 additions & 0 deletions docs/source/user_guide/benchmarks/molecular_reactions.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
===================
Molecular reactions
===================

Tautomers
=========

Summary
-------

Performance in predicting the relative energy of tautomer pairs. Each system is
a pair of tautomers (constitutional isomers differing in the position of a
proton and an associated double bond), and the benchmark measures how well a
model reproduces the energy difference between the two forms. The structures are
taken from the Tautobase dataset and are pre-optimised; only single-point
energies are evaluated.

Metrics
-------

1. Reaction energy MAE

For each pair the reaction energy is the energy difference between the two
tautomers. The mean absolute error (MAE) between the predicted and reference
reaction energies is reported in kcal/mol, averaged across all pairs. Pairs on
which inference fails are excluded from the average.

Computational cost
------------------

Low: only single-point energies are evaluated, so tests run quickly even for the

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

cpu/gpu estimate? can be reasonbly vaguge e.g. minutes on gpu

full dataset.

Data availability
-----------------

Input structures:

* Tautobase: an open tautomer database.
Wahl, O.; Sander, T. *J. Chem. Inf. Model.* 2020, 60 (3), 1085-1089.
DOI: 10.1021/acs.jcim.0c00035

Reference data:

* Same as input data
241 changes: 241 additions & 0 deletions ml_peg/analysis/molecular_reactions/tautomers/analyse_tautomers.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,241 @@
"""Analyse Tautobase tautomer benchmark."""

from __future__ import annotations

import json
from pathlib import Path

from ase.calculators.calculator import Calculator
from mlipaudit.benchmarks.tautomers.tautomers import TautomersModelOutput
import pytest

from ml_peg.analysis.utils.decorators import build_table, plot_parity
from ml_peg.analysis.utils.utils import build_dispersion_name_map, load_metrics_config
from ml_peg.app import APP_ROOT
from ml_peg.calcs import CALCS_ROOT
from ml_peg.calcs.utils.mlipaudit import MlPegTautomersBenchmark
from ml_peg.calcs.utils.utils import download_s3_data
from ml_peg.models import current_models
from ml_peg.models.get_models import load_models

MODELS = load_models(current_models)
DISPERSION_NAME_MAP = build_dispersion_name_map(MODELS)

CALC_PATH = CALCS_ROOT / "molecular_reactions" / "tautomers" / "outputs"
OUT_PATH = APP_ROOT / "data" / "molecular_reactions" / "tautomers"

METRICS_CONFIG_PATH = Path(__file__).with_name("metrics.yml")
DEFAULT_THRESHOLDS, DEFAULT_TOOLTIPS, DEFAULT_WEIGHTS = load_metrics_config(
METRICS_CONFIG_PATH
)


def labels() -> list:
"""
Get the ordered list of tautomer pair IDs.

Returns
-------
list
List of all tautomer pair structure IDs.
"""
for model_name in MODELS:
path = CALC_PATH / model_name / "model_output.json"
if path.exists():
output = TautomersModelOutput.model_validate_json(path.read_text())
return output.structure_ids
return []


@pytest.fixture
def analyze_results() -> dict:
"""
Run the mlipaudit analysis for each model.

Returns
-------
dict
Mapping of model name to its ``TautomersResult``.
"""
data_input_dir = download_s3_data(
key="inputs/molecular_reactions/tautomers/tautomers.zip",
filename="tautomers.zip",
)

results = {}
for model_name in MODELS:
path = CALC_PATH / model_name / "model_output.json"
if not path.exists():
continue
benchmark = MlPegTautomersBenchmark(
force_field=Calculator(),
data_input_dir=data_input_dir,
run_mode="standard",
)
benchmark.model_output = TautomersModelOutput.model_validate_json(
path.read_text()
)
results[model_name] = benchmark.analyze()
return results


@pytest.fixture
def struct_info() -> dict:
"""
Write the combined element set to ``info.json`` for filtering.

Returns
-------
dict
Mapping with the sorted list of elements present in the dataset.
"""
data_input_dir = download_s3_data(
key="inputs/molecular_reactions/tautomers/tautomers.zip",
filename="tautomers.zip",
)
benchmark = MlPegTautomersBenchmark(
force_field=Calculator(),
data_input_dir=data_input_dir,
run_mode="standard",
)
elements = sorted(
{
symbol
for pair in benchmark._tautomers_data.values()
for symbols in pair.atom_symbols
for symbol in symbols
}
)
info = {"elements": elements}
OUT_PATH.mkdir(parents=True, exist_ok=True)
(OUT_PATH / "info.json").write_text(json.dumps(info, indent=1))
return info


@pytest.fixture
@plot_parity(
filename=OUT_PATH / "figure_tautomers.json",
title="Tautomer reaction energies",
x_label="Predicted reaction energy / kcal/mol",
y_label="Reference reaction energy / kcal/mol",
hoverdata={
"Labels": labels(),
},
)
def tautomer_energies(analyze_results) -> dict[str, list]:
"""
Get predicted and reference tautomer reaction energies.

Parameters
----------
analyze_results
Mapping of model name to its ``TautomersResult``.

Returns
-------
dict[str, list]
Dictionary of reference and predicted reaction energies, aligned to the
ordered tautomer pair IDs.
"""
ids = labels()
ref_map: dict[str, float] = {}
pred_maps: dict[str, dict[str, float]] = {}

for model_name, result in analyze_results.items():
pred_maps[model_name] = {}
for molecule in result.molecules:
if molecule.failed:
continue
pred_maps[model_name][molecule.structure_id] = (
molecule.predicted_energy_diff
)
ref_map.setdefault(molecule.structure_id, molecule.ref_energy_diff)

results = {"ref": [ref_map.get(i) for i in ids]}
for model_name in MODELS:
model_preds = pred_maps.get(model_name, {})
results[model_name] = [model_preds.get(i) for i in ids]
return results


@pytest.fixture
def get_mae(analyze_results) -> dict[str, float]:
"""
Get the reaction energy mean absolute error for each model.

Parameters
----------
analyze_results
Mapping of model name to its ``TautomersResult``.

Returns
-------
dict[str, float]
Mean absolute error in kcal/mol for each model.
"""
return {model_name: result.mae for model_name, result in analyze_results.items()}


@pytest.fixture
def get_score(analyze_results) -> dict[str, float]:
"""
Get the mlipaudit benchmark score for each model.

Parameters
----------
analyze_results
Mapping of model name to its ``TautomersResult``.

Returns
-------
dict[str, float]
The mlipaudit per-molecule soft-threshold score (0 to 1) for each model.
"""
return {model_name: result.score for model_name, result in analyze_results.items()}


@pytest.fixture
@build_table(
filename=OUT_PATH / "tautomers_metrics_table.json",
metric_tooltips=DEFAULT_TOOLTIPS,
thresholds=DEFAULT_THRESHOLDS,
weights=DEFAULT_WEIGHTS,
mlip_name_map=DISPERSION_NAME_MAP,
)
def metrics(
tautomer_energies, get_mae: dict[str, float], get_score: dict[str, float]
) -> dict[str, dict]:
"""
Get all metrics.

Parameters
----------
tautomer_energies
Reference and predicted reaction energies (triggers the parity plot).
get_mae
Mean absolute errors for all models.
get_score
The mlipaudit benchmark scores for all models.

Returns
-------
dict[str, dict]
Metric names and values for all models.
"""
return {
"MAE": get_mae,
"Tautomer Score": get_score,
}


def test_tautomers(metrics: dict[str, dict], struct_info: dict) -> None:
"""
Run tautomers analysis.

Parameters
----------
metrics : dict[str, dict]
Tautomers metric results provided by fixtures.
struct_info : dict
Element info written to ``info.json`` for filtering.
"""
13 changes: 13 additions & 0 deletions ml_peg/analysis/molecular_reactions/tautomers/metrics.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,13 @@
metrics:
MAE:
good: 0.0
bad: 2.0
unit: kcal/mol
weight: 0
tooltip: Mean absolute error of tautomer relative (reaction) energies.
Tautomer Score:
good: 1.0
bad: 0.0
unit: null
weight: 1
tooltip: mlipaudit per-molecule soft-threshold score (MAE, failed molecules score 0).
65 changes: 65 additions & 0 deletions ml_peg/app/molecular_reactions/tautomers/app_tautomers.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
"""Run Tautobase tautomer benchmark app."""

from __future__ import annotations

from dash import Dash
from dash.html import Div

from ml_peg.app import APP_ROOT
from ml_peg.app.base_app import BaseApp
from ml_peg.app.utils.build_callbacks import plot_from_table_column
from ml_peg.app.utils.load import read_plot

BENCHMARK_NAME = "Tautomers"
DOCS_URL = "https://ddmms.github.io/ml-peg/user_guide/benchmarks/molecular_reactions.html#tautomers"
DATA_PATH = APP_ROOT / "data" / "molecular_reactions" / "tautomers"


class TautomersApp(BaseApp):
"""Tautomers benchmark app layout and callbacks."""

def register_callbacks(self) -> None:
"""Register callbacks to app."""
scatter = read_plot(
DATA_PATH / "figure_tautomers.json",
id=f"{BENCHMARK_NAME}-figure",
)

plot_from_table_column(
table_id=self.table_id,
plot_id=f"{BENCHMARK_NAME}-figure-placeholder",
column_to_plot={"MAE": scatter},
)


def get_app() -> TautomersApp:
"""
Get tautomers benchmark app layout and callback registration.

Returns
-------
TautomersApp
Benchmark layout and callback registration.
"""
return TautomersApp(
name=BENCHMARK_NAME,
framework_ids="mlip_audit",
description=(
"Performance in predicting relative energies of tautomer pairs. "
"Reference data from the Tautobase dataset."
),
docs_url=DOCS_URL,
table_path=DATA_PATH / "tautomers_metrics_table.json",
info_path=DATA_PATH / "info.json",
extra_components=[
Div(id=f"{BENCHMARK_NAME}-figure-placeholder"),
],
)


if __name__ == "__main__":
full_app = Dash(__name__, assets_folder=DATA_PATH.parent.parent)
benchmark_app = get_app()
full_app.layout = benchmark_app.layout
benchmark_app.register_callbacks()
full_app.run(port=8068, debug=True)
8 changes: 8 additions & 0 deletions ml_peg/app/utils/frameworks.yml
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,14 @@ mlip_arena:
url: "https://huggingface.co/spaces/atomind/mlip-arena"
logo: "https://huggingface.co/front/assets/huggingface_logo-noborder.svg"


mlip_audit:
label: MLIP Audit
color: "#1d4ed8"
text_color: "#ffffff"
url: "https://github.com/instadeepai/mlipaudit"
logo: "https://raw.githubusercontent.com/instadeepai/mlipaudit/mlpeg-migration/InstaDeep_Logo.png"

mace-multihead:
label: Multihead Cross Learning
color: "#7c3aed"
Expand Down
Loading
Loading