-
Notifications
You must be signed in to change notification settings - Fork 9
Expand file tree
/
Copy pathoptimize.py
More file actions
140 lines (119 loc) 路 5.03 KB
/
Copy pathoptimize.py
File metadata and controls
140 lines (119 loc) 路 5.03 KB
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
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
"""Define structure optimization tasks."""
from __future__ import annotations
from ase import Atoms
from ase.calculators.calculator import BaseCalculator
from ase.constraints import FixSymmetry
from ase.filters import ExpCellFilter, Filter, FrechetCellFilter, StrainFilter, UnitCellFilter
from ase.optimize import (
BFGS,
FIRE,
FIRE2,
LBFGS,
BFGSLineSearch,
CellAwareBFGS,
GPMin,
LBFGSLineSearch,
MDMin,
ODE12r,
QuasiNewton,
)
from ase.optimize.optimize import Optimizer
from prefect import task
from prefect.runtime import task_run
from mlip_arena.models import MLIPEnum
from mlip_arena.tasks.utils import ARENA_TASK_CACHE_POLICY, logger, pformat, get_calculator, resolve_calculator_name
_valid_filters: dict[str, Filter] = {
"Filter": Filter,
"UnitCell": UnitCellFilter,
"ExpCell": ExpCellFilter,
"Strain": StrainFilter,
"FrechetCell": FrechetCellFilter,
}
_valid_optimizers: dict[str, Optimizer] = {
"MDMin": MDMin,
"FIRE": FIRE,
"FIRE2": FIRE2,
"LBFGS": LBFGS,
"LBFGSLineSearch": LBFGSLineSearch,
"BFGS": BFGS,
"BFGSLineSearch": BFGSLineSearch,
"QuasiNewton": QuasiNewton,
"GPMin": GPMin,
"CellAwareBFGS": CellAwareBFGS,
"ODE12r": ODE12r,
}
def _generate_task_run_name():
task_name = task_run.task_name
parameters = task_run.parameters
atoms = parameters["atoms"]
calculator_name = resolve_calculator_name(parameters.get("calculator"))
return f"{task_name}: {atoms.get_chemical_formula()} - {calculator_name}"
@task(name="OPT", task_run_name=_generate_task_run_name, cache_policy=ARENA_TASK_CACHE_POLICY)
def run(
atoms: Atoms,
calculator: str | MLIPEnum | BaseCalculator | None = None,
calculator_kwargs: dict | None = None,
dispersion: bool = False,
dispersion_kwargs: dict | None = None,
optimizer: Optimizer | str = BFGSLineSearch,
optimizer_kwargs: dict | None = None,
filter: Filter | str | None = None,
filter_kwargs: dict | None = None,
criterion: dict | None = None,
symmetry: bool = False,
):
"""Run structure optimization.
Args:
atoms (Atoms): ASE Atoms object to optimize.
calculator (str | MLIPEnum | BaseCalculator, optional): ASE calculator or model name/enum.
calculator_kwargs (dict, optional): Keyword arguments to pass to the calculator. Defaults to None.
dispersion (bool, optional): Whether to use dispersion correction. Defaults to False.
dispersion_kwargs (dict, optional): Keyword arguments for dispersion correction.
optimizer (Optimizer | str, optional): ASE optimizer class or name. Defaults to BFGSLineSearch.
optimizer_kwargs (dict, optional): Keyword arguments to pass to the optimizer class. Defaults to None.
filter (Filter | str, optional): ASE filter class or name to apply to the atoms. Defaults to None.
filter_kwargs (dict, optional): Keyword arguments to pass to the filter class. Defaults to None.
criterion (dict, optional): Termination criteria for the optimizer (e.g., {'fmax': 0.01, 'steps': 100}). Defaults to {'steps': 1000}.
symmetry (bool, optional): Whether to use FixSymmetry constraint. Defaults to False.
Returns:
dict: A dictionary containing:
- "atoms": Optimized ASE Atoms object.
- "steps": Number of steps taken by the optimizer.
- "converged": Whether the optimization converged.
"""
atoms = atoms.copy()
calculator_obj = get_calculator(calculator, calculator_kwargs, dispersion, dispersion_kwargs)
atoms.calc = calculator_obj
if isinstance(filter, str):
if filter not in _valid_filters:
raise ValueError(f"Invalid filter: {filter}")
filter = _valid_filters[filter]
if isinstance(optimizer, str):
if optimizer not in _valid_optimizers:
raise ValueError(f"Invalid optimizer: {optimizer}")
optimizer = _valid_optimizers[optimizer]
filter_kwargs = filter_kwargs or {}
optimizer_kwargs = optimizer_kwargs or {}
criterion = criterion or dict(steps=1000)
if symmetry:
atoms.set_constraint(FixSymmetry(atoms))
if isinstance(filter, type) and issubclass(filter, Filter):
filter_instance = filter(atoms, **filter_kwargs)
logger.info(f"Using filter: {filter_instance}")
logger.info(pformat(filter_kwargs))
optimizer_instance = optimizer(filter_instance, **optimizer_kwargs)
logger.info(f"Using optimizer: {optimizer_instance}")
logger.info(pformat(optimizer_kwargs))
logger.info(f"Criterion: {pformat(criterion)}")
converged = optimizer_instance.run(**criterion)
elif filter is None:
optimizer_instance = optimizer(atoms, **optimizer_kwargs)
logger.info(f"Using optimizer: {optimizer_instance}")
logger.info(pformat(optimizer_kwargs))
logger.info(f"Criterion: {pformat(criterion)}")
converged = optimizer_instance.run(**criterion)
return {
"atoms": atoms,
"steps": optimizer_instance.nsteps,
"converged": converged,
}