flashcart 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- flashcart/__init__.py +3 -0
- flashcart/calculators/__init__.py +12 -0
- flashcart/calculators/ase.py +230 -0
- flashcart/calculators/lammps_mliap.py +803 -0
- flashcart/cli/__init__.py +1 -0
- flashcart/cli/lammps_mliap.py +93 -0
- flashcart/cli/profile.py +561 -0
- flashcart/cli/test.py +154 -0
- flashcart/cli/train.py +242 -0
- flashcart/configs/default.yaml +131 -0
- flashcart/configs/profile.yaml +20 -0
- flashcart/data/__init__.py +1 -0
- flashcart/data/data.py +323 -0
- flashcart/data/dataset.py +372 -0
- flashcart/data/graph.py +92 -0
- flashcart/data/neighbors.py +143 -0
- flashcart/data/padding.py +325 -0
- flashcart/data/samplers.py +261 -0
- flashcart/data/statistics.py +158 -0
- flashcart/data/utils.py +218 -0
- flashcart/model/__init__.py +9 -0
- flashcart/model/atomistic.py +527 -0
- flashcart/model/flashcart.py +483 -0
- flashcart/nn/__init__.py +24 -0
- flashcart/nn/layers.py +1138 -0
- flashcart/nn/radial.py +245 -0
- flashcart/o3/__init__.py +8 -0
- flashcart/o3/_codegen_cache.py +110 -0
- flashcart/o3/_codegen_common.py +123 -0
- flashcart/o3/_codegen_irreps.py +601 -0
- flashcart/o3/_codegen_linear.py +451 -0
- flashcart/o3/_codegen_tensor_product.py +3636 -0
- flashcart/o3/_irreps.py +254 -0
- flashcart/o3/_linear.py +971 -0
- flashcart/o3/_tensor_product.py +1369 -0
- flashcart/o3/_triton_launch.py +29 -0
- flashcart/o3/irreps.py +80 -0
- flashcart/o3/kernel_config.py +24 -0
- flashcart/o3/linear.py +137 -0
- flashcart/o3/tensor_product.py +182 -0
- flashcart/o3/utils.py +496 -0
- flashcart/py.typed +0 -0
- flashcart/training/__init__.py +15 -0
- flashcart/training/loss.py +469 -0
- flashcart/training/muon.py +313 -0
- flashcart/training/muon_norm.py +229 -0
- flashcart/training/optimizers.py +148 -0
- flashcart/training/schedulers.py +119 -0
- flashcart/training/tasks.py +740 -0
- flashcart/utils/__init__.py +4 -0
- flashcart/utils/compile.py +37 -0
- flashcart/utils/config.py +160 -0
- flashcart/utils/distributed.py +55 -0
- flashcart/utils/env.py +34 -0
- flashcart/utils/geometry.py +77 -0
- flashcart/utils/lammps.py +194 -0
- flashcart/utils/logging.py +36 -0
- flashcart/utils/parameter_groups.py +30 -0
- flashcart/utils/scatter.py +68 -0
- flashcart/utils/torch_geometric/__init__.py +6 -0
- flashcart/utils/torch_geometric/batch.py +140 -0
- flashcart/utils/torch_geometric/data.py +291 -0
- flashcart/utils/torch_geometric/dataloader.py +136 -0
- flashcart/utils/torch_geometric/dataset.py +71 -0
- flashcart-0.1.0.dist-info/METADATA +194 -0
- flashcart-0.1.0.dist-info/RECORD +71 -0
- flashcart-0.1.0.dist-info/WHEEL +5 -0
- flashcart-0.1.0.dist-info/entry_points.txt +5 -0
- flashcart-0.1.0.dist-info/licenses/LICENSE +202 -0
- flashcart-0.1.0.dist-info/licenses/NOTICE +155 -0
- flashcart-0.1.0.dist-info/top_level.txt +1 -0
flashcart/__init__.py
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
"""Interfaces for atomistic simulations with trained FlashCart models.
|
|
2
|
+
|
|
3
|
+
The package exposes ``FlashCartCalculator`` for use with the Atomic Simulation
|
|
4
|
+
Environment (ASE). The LAMMPS ML-IAP interface is available separately through
|
|
5
|
+
``flashcart.calculators.lammps_mliap``. Importing that module sets a default PyTorch
|
|
6
|
+
CUDA allocator configuration and attempts to load the optional LAMMPS dependency.
|
|
7
|
+
Keeping this import separate avoids these effects when using the ASE calculator.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from flashcart.calculators.ase import FlashCartCalculator
|
|
11
|
+
|
|
12
|
+
__all__ = ["FlashCartCalculator"]
|
|
@@ -0,0 +1,230 @@
|
|
|
1
|
+
from pathlib import Path
|
|
2
|
+
from typing import List, Optional, Union
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
import torch
|
|
6
|
+
from ase.calculators.calculator import Calculator, all_changes
|
|
7
|
+
from ase.stress import full_3x3_to_voigt_6_stress
|
|
8
|
+
|
|
9
|
+
from flashcart.data.graph import graph_from_ase, update_graph_positions
|
|
10
|
+
from flashcart.data.padding import PadAtomicData, slice_padded_outputs
|
|
11
|
+
from flashcart.model.flashcart import FlashCartPotential
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class FlashCartCalculator(Calculator):
|
|
15
|
+
"""Evaluate energies, forces, and stress with a FlashCart model in ASE.
|
|
16
|
+
|
|
17
|
+
A positive ``skin`` extends the neighbor-list radius beyond the model cutoff,
|
|
18
|
+
allowing graph connectivity to be reused between evaluations. The cutoff envelope
|
|
19
|
+
removes contributions from edges outside the model cutoff.
|
|
20
|
+
|
|
21
|
+
Setting ``compile_mode`` enables compiled prediction and graph padding. Padding
|
|
22
|
+
keeps tensor shapes unchanged while the graph fits within the allocated capacities,
|
|
23
|
+
reducing recompilation as the number of edges varies. The capacities grow when
|
|
24
|
+
required. Positive ``pad_n_atoms`` or ``pad_n_edges`` also enables padding without
|
|
25
|
+
compilation.
|
|
26
|
+
|
|
27
|
+
Args:
|
|
28
|
+
checkpoint (Union[str, Path], optional): Inference checkpoint directory.
|
|
29
|
+
Required when ``model`` is not provided. Otherwise ignored.
|
|
30
|
+
model (FlashCartPotential, optional): Model to use for prediction. If set, it
|
|
31
|
+
overrides ``checkpoint``.
|
|
32
|
+
device (Union[str, torch.device], optional): Device used for prediction. If
|
|
33
|
+
None, retain the device of a provided model. Checkpoints are loaded on CPU
|
|
34
|
+
unless a device is specified.
|
|
35
|
+
skin (float, optional): Additional neighbor-list radius. The list is constructed
|
|
36
|
+
at ``r_max + skin``. For a fixed cell, it can be reused while every atom has
|
|
37
|
+
moved less than ``skin / 2`` from its position at construction, provided no
|
|
38
|
+
other atomic properties have changed. Cell changes are checked separately.
|
|
39
|
+
Default: 0.0, which rebuilds the list for each calculation.
|
|
40
|
+
compile_mode (str, optional): ``torch.compile`` mode used for prediction.
|
|
41
|
+
Setting a mode also enables graph padding. Default: None, which evaluates
|
|
42
|
+
the model without compilation.
|
|
43
|
+
compile_fullgraph (bool, optional): Whether compilation requires a single graph
|
|
44
|
+
without graph breaks. Used when ``compile_mode`` is set. Default: True.
|
|
45
|
+
pad_n_atoms (int, optional): Initial capacity for atoms before adding reserved
|
|
46
|
+
padding atoms. The capacity grows when required. Default: 0.
|
|
47
|
+
pad_n_edges (int, optional): Initial capacity for edges before rounding to
|
|
48
|
+
``pad_edge_multiple``. The capacity grows when required. Default: 0.
|
|
49
|
+
pad_extra_atoms (int, optional): Number of additional padding atoms reserved as
|
|
50
|
+
endpoints of padding edges, with a minimum of two. Default: None, which
|
|
51
|
+
determines the number from the atom capacity and ``pad_edge_headroom``.
|
|
52
|
+
pad_atom_multiple (int, optional): Round the total number of atoms, including
|
|
53
|
+
reserved padding atoms, up to this multiple. Default: 1.
|
|
54
|
+
pad_edge_headroom (float, optional): Multiplicative factor applied to the
|
|
55
|
+
required edge count when increasing the edge capacity. Default: 1.1.
|
|
56
|
+
pad_edge_multiple (int, optional): Round the total number of edges, including
|
|
57
|
+
padding edges, up to this multiple. Default: 128.
|
|
58
|
+
add_atomic_offsets (bool, optional): Restore the per-element energy shifts
|
|
59
|
+
subtracted from the reference energies during training. These shifts affect
|
|
60
|
+
the returned energy but not forces or stress. Default: False.
|
|
61
|
+
**kwargs: Additional keyword arguments passed to ``ase.calculators.Calculator``.
|
|
62
|
+
"""
|
|
63
|
+
|
|
64
|
+
implemented_properties = ["energy", "forces", "stress"]
|
|
65
|
+
|
|
66
|
+
def __init__(
|
|
67
|
+
self,
|
|
68
|
+
checkpoint: Optional[Union[str, Path]] = None,
|
|
69
|
+
model: Optional[FlashCartPotential] = None,
|
|
70
|
+
device: Optional[Union[str, torch.device]] = None,
|
|
71
|
+
skin: float = 0.0,
|
|
72
|
+
compile_mode: Optional[str] = None,
|
|
73
|
+
compile_fullgraph: bool = True,
|
|
74
|
+
pad_n_atoms: int = 0,
|
|
75
|
+
pad_n_edges: int = 0,
|
|
76
|
+
pad_extra_atoms: Optional[int] = None,
|
|
77
|
+
pad_atom_multiple: int = 1,
|
|
78
|
+
pad_edge_headroom: float = 1.1,
|
|
79
|
+
pad_edge_multiple: int = 128,
|
|
80
|
+
add_atomic_offsets: bool = False,
|
|
81
|
+
**kwargs,
|
|
82
|
+
):
|
|
83
|
+
super().__init__(**kwargs)
|
|
84
|
+
if model is None:
|
|
85
|
+
if checkpoint is None:
|
|
86
|
+
raise ValueError("Provide either checkpoint or model.")
|
|
87
|
+
model = FlashCartPotential.from_checkpoint(checkpoint, device=device)
|
|
88
|
+
elif device is not None:
|
|
89
|
+
model = model.to(device)
|
|
90
|
+
self.model = model
|
|
91
|
+
self._device = next(model.parameters()).device
|
|
92
|
+
self.add_atomic_offsets = bool(add_atomic_offsets)
|
|
93
|
+
self._atomic_shifts = model.scale_shift.shifts.detach().to(self._device)
|
|
94
|
+
self.skin = float(skin)
|
|
95
|
+
self.compile_mode = compile_mode
|
|
96
|
+
self.compile_fullgraph = bool(compile_fullgraph)
|
|
97
|
+
self._use_compile = compile_mode is not None
|
|
98
|
+
self._use_padding = self._use_compile or pad_n_atoms > 0 or pad_n_edges > 0
|
|
99
|
+
self._padder = PadAtomicData(
|
|
100
|
+
r_max=model.r_max,
|
|
101
|
+
atom_budget=pad_n_atoms,
|
|
102
|
+
edge_budget=pad_n_edges,
|
|
103
|
+
extra_atoms=pad_extra_atoms,
|
|
104
|
+
atom_multiple=pad_atom_multiple,
|
|
105
|
+
edge_headroom=pad_edge_headroom,
|
|
106
|
+
edge_multiple=pad_edge_multiple,
|
|
107
|
+
)
|
|
108
|
+
self._cached_graph = None
|
|
109
|
+
self._cached_positions: Optional[np.ndarray] = None
|
|
110
|
+
self._cached_cell: Optional[np.ndarray] = None
|
|
111
|
+
|
|
112
|
+
@classmethod
|
|
113
|
+
def from_checkpoint(
|
|
114
|
+
cls,
|
|
115
|
+
checkpoint: Union[str, Path],
|
|
116
|
+
device: Optional[Union[str, torch.device]] = None,
|
|
117
|
+
skin: float = 0.0,
|
|
118
|
+
**kwargs,
|
|
119
|
+
) -> "FlashCartCalculator":
|
|
120
|
+
"""Build a calculator from a saved inference checkpoint.
|
|
121
|
+
|
|
122
|
+
Args:
|
|
123
|
+
checkpoint (Union[str, Path]): Inference checkpoint directory.
|
|
124
|
+
device (Union[str, torch.device], optional): Device used for prediction.
|
|
125
|
+
Default: None, which loads the model on CPU.
|
|
126
|
+
skin (float, optional): Additional neighbor-list radius, as described in
|
|
127
|
+
``FlashCartCalculator``. Default: 0.0.
|
|
128
|
+
**kwargs: Additional arguments passed to ``FlashCartCalculator``.
|
|
129
|
+
|
|
130
|
+
Returns:
|
|
131
|
+
FlashCartCalculator: Calculator initialized with the saved model.
|
|
132
|
+
"""
|
|
133
|
+
return cls(checkpoint=checkpoint, device=device, skin=skin, **kwargs)
|
|
134
|
+
|
|
135
|
+
def calculate(self, atoms=None, properties: Optional[List[str]] = None, system_changes=all_changes):
|
|
136
|
+
"""Evaluate the requested properties and store them in ``self.results``.
|
|
137
|
+
|
|
138
|
+
Graph connectivity is reused when the neighbor-list reuse conditions are
|
|
139
|
+
satisfied. Otherwise, the graph is rebuilt. The results contain only the
|
|
140
|
+
requested properties. Energy is stored as a scalar, forces as an array of shape
|
|
141
|
+
``(n_atoms, 3)``, and stress as a six-component array in ASE Voigt order:
|
|
142
|
+
``xx, yy, zz, yz, xz, xy``.
|
|
143
|
+
|
|
144
|
+
Args:
|
|
145
|
+
atoms (ase.Atoms, optional): Atomic configuration to evaluate. If None, use
|
|
146
|
+
the configuration already stored by the calculator.
|
|
147
|
+
properties (list[str], optional): Properties to evaluate, chosen from
|
|
148
|
+
``"energy"``, ``"forces"``, and ``"stress"``. Default:
|
|
149
|
+
``["energy", "forces"]``.
|
|
150
|
+
system_changes (list[str], optional): Changes reported by ASE since the
|
|
151
|
+
preceding evaluation. Changes other than positions or cell require the
|
|
152
|
+
neighbor list to be rebuilt. Default: ``all_changes``.
|
|
153
|
+
"""
|
|
154
|
+
if properties is None:
|
|
155
|
+
properties = ["energy", "forces"]
|
|
156
|
+
super().calculate(atoms, properties, system_changes)
|
|
157
|
+
assert self.atoms is not None
|
|
158
|
+
|
|
159
|
+
graph = self._get_graph(system_changes)
|
|
160
|
+
compute_forces = "forces" in properties
|
|
161
|
+
compute_stress = "stress" in properties
|
|
162
|
+
n_real_atoms = int(graph.positions.shape[0])
|
|
163
|
+
n_real_graphs = 1
|
|
164
|
+
predict_graph = graph
|
|
165
|
+
if self._use_padding:
|
|
166
|
+
predict_graph, n_real_atoms, n_real_graphs = self._padder(graph)
|
|
167
|
+
|
|
168
|
+
out = self.model.predict(
|
|
169
|
+
predict_graph,
|
|
170
|
+
compute_forces=compute_forces,
|
|
171
|
+
compute_stress=compute_stress,
|
|
172
|
+
use_compile=self._use_compile,
|
|
173
|
+
compile_mode=self.compile_mode or "reduce-overhead",
|
|
174
|
+
fullgraph=self.compile_fullgraph,
|
|
175
|
+
dynamic=False,
|
|
176
|
+
)
|
|
177
|
+
if self._use_padding:
|
|
178
|
+
out = slice_padded_outputs(out, n_real_atoms, n_real_graphs)
|
|
179
|
+
|
|
180
|
+
self.results = {}
|
|
181
|
+
if "energy" in properties:
|
|
182
|
+
energy = float(out["energy"].detach().cpu().reshape(-1)[0])
|
|
183
|
+
if self.add_atomic_offsets:
|
|
184
|
+
energy += float(self._atomic_shifts[graph.atom_types].double().sum().cpu())
|
|
185
|
+
self.results["energy"] = energy
|
|
186
|
+
if "forces" in properties:
|
|
187
|
+
self.results["forces"] = out["forces"].detach().cpu().numpy()
|
|
188
|
+
if "stress" in properties:
|
|
189
|
+
stress = out["stress"].detach().cpu().numpy()
|
|
190
|
+
self.results["stress"] = stress.reshape(-1) if stress.size == 6 else full_3x3_to_voigt_6_stress(stress)
|
|
191
|
+
|
|
192
|
+
def _get_graph(self, system_changes: List[str]):
|
|
193
|
+
"""Build an atomic graph or update the geometry of a cached graph.
|
|
194
|
+
|
|
195
|
+
Reuse requires a positive skin and changes limited to positions and cell. Each
|
|
196
|
+
atomic displacement and the Frobenius norm of the cell change, measured from the
|
|
197
|
+
last graph construction, must be less than ``skin / 2``. Reuse updates positions
|
|
198
|
+
and cell while retaining the cached connectivity and integer periodic shifts.
|
|
199
|
+
|
|
200
|
+
Args:
|
|
201
|
+
system_changes (list[str]): Changes reported by ASE since the preceding
|
|
202
|
+
evaluation.
|
|
203
|
+
|
|
204
|
+
Returns:
|
|
205
|
+
AtomicData: Graph for the current atomic configuration.
|
|
206
|
+
"""
|
|
207
|
+
reuse = (
|
|
208
|
+
self.skin > 0.0
|
|
209
|
+
and self._cached_graph is not None
|
|
210
|
+
and self._cached_positions is not None
|
|
211
|
+
and self._cached_cell is not None
|
|
212
|
+
and set(system_changes) <= {"positions", "cell"}
|
|
213
|
+
)
|
|
214
|
+
if reuse:
|
|
215
|
+
disp = np.linalg.norm(self.atoms.get_positions() - self._cached_positions, axis=-1)
|
|
216
|
+
cell_disp = np.linalg.norm(np.asarray(self.atoms.get_cell()) - self._cached_cell)
|
|
217
|
+
if np.all(disp < self.skin / 2) and cell_disp < self.skin / 2:
|
|
218
|
+
return update_graph_positions(self._cached_graph, self.atoms)
|
|
219
|
+
|
|
220
|
+
graph = graph_from_ase(
|
|
221
|
+
self.atoms,
|
|
222
|
+
self.model.elements,
|
|
223
|
+
self.model.r_max,
|
|
224
|
+
skin=self.skin,
|
|
225
|
+
device=self._device,
|
|
226
|
+
)
|
|
227
|
+
self._cached_graph = graph
|
|
228
|
+
self._cached_positions = self.atoms.get_positions().copy()
|
|
229
|
+
self._cached_cell = np.asarray(self.atoms.get_cell(), dtype=np.float64).copy()
|
|
230
|
+
return graph
|