Skip to content
Merged
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
9 changes: 9 additions & 0 deletions examples/input_yamls/minimize_energy.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
minimizer:
fmax: 0.05
steps: 500
device: cpu
dtype: torch.float32
optimizer_cls: torch.optim.LBFGS
model_file: SOME_PATH
structure_file: SOME_PATH
output_file: minimized_structures.pt
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,7 @@ mlcg-train = "mlcg.scripts.mlcg_train:main"
mlcg-train_multischeduler = "mlcg.scripts.mlcg_train_multischeduler:main"
mlcg-nvt_langevin = "mlcg.scripts.mlcg_nvt_langevin:main"
mlcg-nvt_pt_langevin = "mlcg.scripts.mlcg_nvt_pt_langevin:main"
mlcg-minimize_energy = "mlcg.scripts.mlcg_minimize_energy:main"

[tool.black]
line-length = 80
Expand Down
24 changes: 24 additions & 0 deletions src/mlcg/scripts/mlcg_minimize_energy.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
#! /usr/bin/env python

from time import ctime
import torch

from mlcg.simulation import parse_minimizer_config, minimize_energy


def main():
print(f"Starting minimization at {ctime()} with {minimize_energy}")
(
model,
initial_data_list,
minimizer_kwargs,
output_file,
) = parse_minimizer_config()

minimized = minimize_energy(model, initial_data_list, **minimizer_kwargs)
torch.save(minimized, output_file)
print(f"Ending minimization at {ctime()}")


if __name__ == "__main__":
main()
3 changes: 2 additions & 1 deletion src/mlcg/simulation/__init__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from .langevin import LangevinSimulation, OverdampedSimulation
from .parallel_tempering import PTSimulation
from .base import _Simulation
from .cli import parse_simulation_config
from .cli import parse_simulation_config, parse_minimizer_config
from .specialize_prior import condense_all_priors_for_simulation
from .minimizer import minimize_energy
108 changes: 108 additions & 0 deletions src/mlcg/simulation/cli.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import os.path as osp
from typing import Any, List, Dict, Tuple, Sequence
import torch
from jsonargparse import (
Expand All @@ -13,6 +14,7 @@
PTSimulation,
OverdampedSimulation,
)
from .minimizer import minimize_energy
from ..data import AtomicData
from ..nn import load_and_adapt_old_checkpoint
from ..utils import dump_yaml
Expand Down Expand Up @@ -124,6 +126,112 @@ def parse_simulation_config(
return model, initial_data_list, betas, simulation, profile


def parse_minimizer_config(
description: str = "Energy minimization command line tool",
parser_kwargs: Dict[str, Any] = None,
) -> Tuple[torch.nn.Module, List[AtomicData], Dict[str, Any], str]:
"""Utility to parse a configuration file to run :py:func:`minimize_energy`
from the command line, mirroring :py:func:`parse_simulation_config`.

Parameters
----------
description : str, optional
cli description, by default "Energy minimization command line tool"
parser_kwargs : Dict[str, Any], optional
more arguments to the parser, by default None

Returns
-------
model, initial_data_list, minimizer_kwargs, output_file
"""
parser_kwargs = {} if parser_kwargs is None else parser_kwargs
parser_kwargs.update({"description": description})
parser = SimulationParser(**parser_kwargs)
# minimize_energy is a plain function rather than a _Simulation subclass,
# so its arguments are added directly instead of via add_simulation_args.
# model/configurations are skipped here: they are loaded from
# model_file/structure_file below instead of being part of the config.
parser.add_function_arguments(
minimize_energy,
"minimizer",
skip={"model", "configurations"},
fail_untyped=False,
)

parser.add_argument(
"-mf",
"--model_file",
metavar="FN",
type=Path_fr,
help="path to the pytorch model file (including the priors) in pytorch format",
)

parser.add_argument(
"-sf",
"--structure_file",
metavar="FN",
type=Path_fr,
help="path to the starting configurations (a list of un-collated "
"AtomicData) in pytorch format",
)

parser.add_argument(
"-o",
"--output_file",
metavar="FN",
type=str,
help="path at which the minimized configurations (a list of "
"AtomicData) will be saved, in pytorch format",
)

config = parser.parse_args()
# save config
exported_config = {}
for k, v in config.items():
# Path_fr must be converted to string otherwise they can't be saved
if isinstance(v, Path_fr):
exported_config[k] = str(v)
# redundant to save the path to the original config
elif k == "config":
continue
elif k == "minimizer":
# dtype and optimizer_cls are resolved to actual torch.dtype/type
# objects by this point, neither of which ruamel.yaml can
# represent, so they are dumped back out as import path strings.
exported_config[k] = {
mk: (
str(mv)
if isinstance(mv, torch.dtype)
else (
f"{mv.__module__}.{mv.__qualname__}"
if isinstance(mv, type)
else mv
)
)
for mk, mv in v.items()
}
else:
exported_config[k] = v
out_name = osp.splitext(config["output_file"])[0]
dump_yaml(f"{out_name}_config.yaml", exported_config)

model_fn = config.pop("model_file")
model = load_and_adapt_old_checkpoint(
(model_fn if isinstance(model_fn, str) else model_fn())
)

structures_fn = config.pop("structure_file")
initial_data_list = torch.load(
(structures_fn if isinstance(structures_fn, str) else structures_fn()),
weights_only=False,
)

minimizer_kwargs = config.pop("minimizer")
output_file = config.pop("output_file")

return model, initial_data_list, minimizer_kwargs, output_file


class ConfigurationException(Exception):
"""
Exception used to inform users
Expand Down
Loading
Loading