andreasburger commited on
Commit
2646173
·
verified ·
1 Parent(s): 76e4040

Upload folder using huggingface_hub

Browse files
Files changed (2) hide show
  1. equiformer_calculator.py +185 -0
  2. example_inference.py +4 -2
equiformer_calculator.py ADDED
@@ -0,0 +1,185 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Dict, List, Optional, Tuple
2
+ from omegaconf import ListConfig
3
+ import yaml
4
+ import os
5
+ import torch
6
+ from torch_geometric.data import Batch
7
+ from torch_geometric.data import Data as TGData
8
+ from torch_geometric.loader import DataLoader as TGDataLoader
9
+
10
+ from nets.equiformer_v2.equiformer_v2_oc20 import EquiformerV2_OC20
11
+ from nets.prediction_utils import compute_extra_props
12
+ from ocpmodels.common.relaxation.ase_utils import (
13
+ batch_to_atoms,
14
+ ase_atoms_to_torch_geometric,
15
+ )
16
+ from ocpmodels.datasets import data_list_collater
17
+ from ocpmodels.preprocessing import AtomsToGraphs
18
+
19
+ from ase.calculators.calculator import Calculator
20
+ from ase import Atoms
21
+
22
+
23
+ def get_model(config_path):
24
+ with open(config_path, "r") as file:
25
+ config = yaml.safe_load(file)
26
+ model_config = config["model"]
27
+ print("model_config", model_config)
28
+ return EquiformerV2_OC20(**model_config), model_config
29
+
30
+
31
+ class EquiformerCalculator:
32
+ def __init__(
33
+ self,
34
+ checkpoint_path: Optional[str] = None,
35
+ device: Optional[torch.device] = None,
36
+ **kwargs,
37
+ ):
38
+ if device is None:
39
+ device = "cuda" if torch.cuda.is_available() else "cpu"
40
+
41
+ project_root = os.path.dirname(os.path.dirname(__file__))
42
+
43
+ config_path = os.path.join(project_root, "configs/equiformer_v2.yaml")
44
+ self.model, self.model_config = get_model(config_path)
45
+
46
+ if checkpoint_path is None:
47
+ checkpoint_path = os.path.join(project_root, "ckpt/eqv2.ckpt")
48
+ state_dict = torch.load(checkpoint_path, weights_only=True)["state_dict"]
49
+ state_dict = {k.replace("potential.", ""): v for k, v in state_dict.items()}
50
+ self.model.load_state_dict(state_dict, strict=False)
51
+
52
+ self.model.eval()
53
+ self.model.to(device)
54
+
55
+ # ocpmodels/common/relaxation/ase_utils.py
56
+ self.a2g = AtomsToGraphs(
57
+ max_neigh=self.model.max_neighbors,
58
+ radius=self.model.cutoff,
59
+ r_energy=False,
60
+ r_forces=False,
61
+ r_distances=False,
62
+ r_edges=False,
63
+ r_pbc=True,
64
+ )
65
+
66
+ def predict(self, batch):
67
+ """Predict one or multiple batches"""
68
+ batch = batch.to(self.model.device)
69
+ batch = compute_extra_props(batch, pos_require_grad=False)
70
+ energy, forces, eigenpred = self.model.forward(batch, eigen=True)
71
+ return energy, forces, eigenpred
72
+
73
+ def get_forces(self, batch):
74
+ """Get forces from the model"""
75
+ batch = batch.to(self.model.device)
76
+ batch = compute_extra_props(batch, pos_require_grad=False)
77
+ _, forces, _ = self.model.forward(batch, eigen=False)
78
+ return forces
79
+
80
+ def get_energy(self, batch):
81
+ """Get energy from the model"""
82
+ batch = batch.to(self.model.device)
83
+ batch = compute_extra_props(batch, pos_require_grad=False)
84
+ energy, _, _ = self.model.forward(batch, eigen=False)
85
+ return energy
86
+
87
+ def predict_with_hessian(self, batch):
88
+ """Predict one batch with autodiff Hessian"""
89
+ B = batch.batch.max() + 1
90
+ assert B == 1, "Only one batch is supported for Hessian prediction"
91
+
92
+ batch = batch.to(self.model.device)
93
+
94
+ # Prepare batch with extra properties
95
+ batch = compute_extra_props(batch, pos_require_grad=True)
96
+
97
+ # Run prediction
98
+ with torch.enable_grad():
99
+ energy, forces, eigenpred = self.model.forward(batch, eigen=True)
100
+
101
+ # 3D coordinates -> 3N^2 Hessian elements
102
+ N = batch.pos.shape[0]
103
+ forces = forces.reshape(-1)
104
+ num_elements = forces.shape[0]
105
+
106
+ def get_vjp(v):
107
+ return torch.autograd.grad(
108
+ outputs=-1 * forces,
109
+ inputs=batch.pos,
110
+ grad_outputs=v,
111
+ retain_graph=True,
112
+ create_graph=False,
113
+ allow_unused=False,
114
+ )
115
+
116
+ I_N = torch.eye(num_elements, device=forces.device)
117
+ hessian = torch.vmap(get_vjp, in_dims=0, out_dims=0, chunk_size=None)(I_N)[0]
118
+ hessian = hessian.view(N * 3, N * 3)
119
+
120
+ eigenvalues, eigenvectors = torch.linalg.eigh(hessian)
121
+ smallest_eigenvals = eigenvalues[:2]
122
+ smallest_eigenvecs = eigenvectors[:, :2]
123
+ eigenvalues = smallest_eigenvals
124
+ eigenvectors = smallest_eigenvecs.T.view(2, N, 3)
125
+ return energy, forces, hessian, eigenvalues, eigenvectors, eigenpred
126
+
127
+ def predict_gad(self, batch):
128
+ """
129
+ Gentlest Ascent Dynamics (GAD)
130
+ dx/dt = -∇V(x) + 2(∇V, v(x))v(x)
131
+ = F + 2(-F, v(x))v(x)
132
+ since F=-∇V(x)
133
+ where v(x) is the eigenvector of the Hessian with the smallest eigenvalue.
134
+ """
135
+ B = batch.batch.max() + 1
136
+ energy, forces, eigenpred = self.predict(batch)
137
+ v = eigenpred["eigvec_1"].reshape(B, -1)
138
+ # normalize eigenvector
139
+ v = v / torch.norm(v, dim=1, keepdim=True)
140
+ forces = forces.reshape(B, -1)
141
+ # −∇V(x) + 2(∇V, v(x))v(x)
142
+ gad = forces + 2 * torch.einsum("bi,bi->b", -forces, v) * v
143
+ out = {
144
+ "energy": energy,
145
+ "forces": forces,
146
+ }
147
+ out.update(eigenpred)
148
+ return gad, out
149
+
150
+ def predict_gad_with_hessian(self, batch):
151
+ energy, forces, hessian, eigenvalues, eigenvectors, eigenpred = (
152
+ self.predict_with_hessian(batch)
153
+ )
154
+ v = eigenvectors[0].reshape(-1) # N*3
155
+ v = v / torch.norm(v, dim=0, keepdim=True)
156
+ forces = forces.reshape(-1) # N*3
157
+ # Diagnostic prints
158
+ dot_product = torch.dot(-forces, v)
159
+ gad = forces + (2 * dot_product * v)
160
+
161
+ out = {
162
+ "energy": energy,
163
+ "forces": forces,
164
+ "hessian": hessian,
165
+ "eigenvalues": eigenvalues,
166
+ "eigenvectors": eigenvectors,
167
+ }
168
+ out.update(eigenpred)
169
+ return gad, out
170
+
171
+ def find_transitionstate_with_gad_from_hessian(self, batch):
172
+ """Integrate the equations of motion of the GAD vector field."""
173
+ raise NotImplementedError("Not implemented")
174
+
175
+ def ase_to_batch(self, atoms: Atoms):
176
+ # Call base class to set atoms attribute
177
+ Calculator.calculate(self, atoms)
178
+
179
+ # ocpmodels/common/relaxation/ase_utils.py
180
+ # data_object = self.a2g.convert(atoms)
181
+ # batch = data_list_collater([data_object], otf_graph=True)
182
+
183
+ # Convert ASE atoms to torch_geometric format
184
+ batch = ase_atoms_to_torch_geometric(atoms)
185
+ return batch
example_inference.py CHANGED
@@ -68,9 +68,10 @@ def predict_gad(batch, potential):
68
  B = batch.batch.max() + 1
69
  energy, forces, eigenpred = predict(batch, potential)
70
  v = eigenpred["eigvec_1"].reshape(B, -1)
 
71
  forces = forces.reshape(B, -1)
72
  # −∇V(x) + 2(∇V, v(x))v(x)
73
- gad = -forces + 2 * torch.einsum("bi,bi->b", forces, v) * v
74
  return gad
75
 
76
 
@@ -79,9 +80,10 @@ def predict_gad_with_hessian(batch, potential):
79
  predict_with_hessian(batch, potential)
80
  )
81
  v = eigenvectors[0].reshape(-1) # N*3
 
82
  forces = forces.reshape(-1) # N*3
83
  # −∇V(x) + 2(∇V, v(x))v(x)
84
- gad = -forces + 2 * torch.einsum("i,i->", forces, v) * v
85
  return gad
86
 
87
 
 
68
  B = batch.batch.max() + 1
69
  energy, forces, eigenpred = predict(batch, potential)
70
  v = eigenpred["eigvec_1"].reshape(B, -1)
71
+ v = v / torch.norm(v, dim=1, keepdim=True)
72
  forces = forces.reshape(B, -1)
73
  # −∇V(x) + 2(∇V, v(x))v(x)
74
+ gad = forces + 2 * torch.einsum("bi,bi->b", -forces, v) * v
75
  return gad
76
 
77
 
 
80
  predict_with_hessian(batch, potential)
81
  )
82
  v = eigenvectors[0].reshape(-1) # N*3
83
+ v = v / torch.norm(v, dim=0, keepdim=True)
84
  forces = forces.reshape(-1) # N*3
85
  # −∇V(x) + 2(∇V, v(x))v(x)
86
+ gad = forces + 2 * torch.einsum("i,i->", -forces, v) * v
87
  return gad
88
 
89