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.
Files changed (71) hide show
  1. flashcart/__init__.py +3 -0
  2. flashcart/calculators/__init__.py +12 -0
  3. flashcart/calculators/ase.py +230 -0
  4. flashcart/calculators/lammps_mliap.py +803 -0
  5. flashcart/cli/__init__.py +1 -0
  6. flashcart/cli/lammps_mliap.py +93 -0
  7. flashcart/cli/profile.py +561 -0
  8. flashcart/cli/test.py +154 -0
  9. flashcart/cli/train.py +242 -0
  10. flashcart/configs/default.yaml +131 -0
  11. flashcart/configs/profile.yaml +20 -0
  12. flashcart/data/__init__.py +1 -0
  13. flashcart/data/data.py +323 -0
  14. flashcart/data/dataset.py +372 -0
  15. flashcart/data/graph.py +92 -0
  16. flashcart/data/neighbors.py +143 -0
  17. flashcart/data/padding.py +325 -0
  18. flashcart/data/samplers.py +261 -0
  19. flashcart/data/statistics.py +158 -0
  20. flashcart/data/utils.py +218 -0
  21. flashcart/model/__init__.py +9 -0
  22. flashcart/model/atomistic.py +527 -0
  23. flashcart/model/flashcart.py +483 -0
  24. flashcart/nn/__init__.py +24 -0
  25. flashcart/nn/layers.py +1138 -0
  26. flashcart/nn/radial.py +245 -0
  27. flashcart/o3/__init__.py +8 -0
  28. flashcart/o3/_codegen_cache.py +110 -0
  29. flashcart/o3/_codegen_common.py +123 -0
  30. flashcart/o3/_codegen_irreps.py +601 -0
  31. flashcart/o3/_codegen_linear.py +451 -0
  32. flashcart/o3/_codegen_tensor_product.py +3636 -0
  33. flashcart/o3/_irreps.py +254 -0
  34. flashcart/o3/_linear.py +971 -0
  35. flashcart/o3/_tensor_product.py +1369 -0
  36. flashcart/o3/_triton_launch.py +29 -0
  37. flashcart/o3/irreps.py +80 -0
  38. flashcart/o3/kernel_config.py +24 -0
  39. flashcart/o3/linear.py +137 -0
  40. flashcart/o3/tensor_product.py +182 -0
  41. flashcart/o3/utils.py +496 -0
  42. flashcart/py.typed +0 -0
  43. flashcart/training/__init__.py +15 -0
  44. flashcart/training/loss.py +469 -0
  45. flashcart/training/muon.py +313 -0
  46. flashcart/training/muon_norm.py +229 -0
  47. flashcart/training/optimizers.py +148 -0
  48. flashcart/training/schedulers.py +119 -0
  49. flashcart/training/tasks.py +740 -0
  50. flashcart/utils/__init__.py +4 -0
  51. flashcart/utils/compile.py +37 -0
  52. flashcart/utils/config.py +160 -0
  53. flashcart/utils/distributed.py +55 -0
  54. flashcart/utils/env.py +34 -0
  55. flashcart/utils/geometry.py +77 -0
  56. flashcart/utils/lammps.py +194 -0
  57. flashcart/utils/logging.py +36 -0
  58. flashcart/utils/parameter_groups.py +30 -0
  59. flashcart/utils/scatter.py +68 -0
  60. flashcart/utils/torch_geometric/__init__.py +6 -0
  61. flashcart/utils/torch_geometric/batch.py +140 -0
  62. flashcart/utils/torch_geometric/data.py +291 -0
  63. flashcart/utils/torch_geometric/dataloader.py +136 -0
  64. flashcart/utils/torch_geometric/dataset.py +71 -0
  65. flashcart-0.1.0.dist-info/METADATA +194 -0
  66. flashcart-0.1.0.dist-info/RECORD +71 -0
  67. flashcart-0.1.0.dist-info/WHEEL +5 -0
  68. flashcart-0.1.0.dist-info/entry_points.txt +5 -0
  69. flashcart-0.1.0.dist-info/licenses/LICENSE +202 -0
  70. flashcart-0.1.0.dist-info/licenses/NOTICE +155 -0
  71. flashcart-0.1.0.dist-info/top_level.txt +1 -0
flashcart/__init__.py ADDED
@@ -0,0 +1,3 @@
1
+ """FlashCart: Fast Cartesian tensor products for equivariant interatomic potentials."""
2
+
3
+ __version__ = "0.1.0"
@@ -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