meidnet 2.2.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.
- meidnet/__init__.py +20 -0
- meidnet/benchmark.py +545 -0
- meidnet/checkpoint.py +167 -0
- meidnet/chem.py +46 -0
- meidnet/cli.py +394 -0
- meidnet/config.py +302 -0
- meidnet/constraints.py +346 -0
- meidnet/data.py +369 -0
- meidnet/designspace.py +114 -0
- meidnet/families/double_perovskite_a2bbx6.yaml +100 -0
- meidnet/families/perovskite_abx3.yaml +175 -0
- meidnet/family.py +229 -0
- meidnet/generate.py +556 -0
- meidnet/model.py +176 -0
- meidnet/pipeline.py +186 -0
- meidnet/registry.py +49 -0
- meidnet/report.py +427 -0
- meidnet/screen.py +182 -0
- meidnet/studio/__init__.py +1 -0
- meidnet/studio/ask_prism.js +340 -0
- meidnet/studio/chemiscope.py +232 -0
- meidnet/studio/landing.html +244 -0
- meidnet/studio/server.py +1355 -0
- meidnet/studio/studio.html +1585 -0
- meidnet/svg.py +215 -0
- meidnet/terms.py +265 -0
- meidnet/train.py +212 -0
- meidnet-2.2.0.dist-info/METADATA +222 -0
- meidnet-2.2.0.dist-info/RECORD +33 -0
- meidnet-2.2.0.dist-info/WHEEL +5 -0
- meidnet-2.2.0.dist-info/entry_points.txt +2 -0
- meidnet-2.2.0.dist-info/licenses/LICENSE +21 -0
- meidnet-2.2.0.dist-info/top_level.txt +1 -0
meidnet/__init__.py
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
"""
|
|
2
|
+
MEIDNet — Multimodal Equivariant Inverse Design Network
|
|
3
|
+
=======================================================
|
|
4
|
+
|
|
5
|
+
Learn a shared latent space between crystal structures and their properties,
|
|
6
|
+
then search it for new materials that hit property targets while obeying the
|
|
7
|
+
chemical and structural rules of a material family.
|
|
8
|
+
|
|
9
|
+
Typical use (the same steps the ``meidnet`` command runs)::
|
|
10
|
+
|
|
11
|
+
from meidnet.config import load_config
|
|
12
|
+
from meidnet.pipeline import check, train, generate
|
|
13
|
+
|
|
14
|
+
cfg = load_config("meidnet.yaml")
|
|
15
|
+
check(cfg) # is my data usable? (writes check_report.html)
|
|
16
|
+
train(cfg) # learn the latent space (writes model.pt + training_report.html)
|
|
17
|
+
generate(cfg) # design candidates (writes CIFs + generation_report.html)
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
__version__ = "2.2.0"
|
meidnet/benchmark.py
ADDED
|
@@ -0,0 +1,545 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Benchmark tasks and their reference implementation.
|
|
3
|
+
|
|
4
|
+
A dataset's protocol (``benchmarks/datasets/<id>.json``, key ``protocol``) fixes the splits, the property targets,
|
|
5
|
+
the candidate budget and the metrics of each task. This module computes those metrics for a MEIDNet checkpoint and
|
|
6
|
+
for the baselines, and scores a set of candidate structures whatever model produced them, so that any method can
|
|
7
|
+
be compared under the same rules.
|
|
8
|
+
|
|
9
|
+
Tasks of the Perov-5 protocol:
|
|
10
|
+
|
|
11
|
+
inverse_design candidates for property targets in several chemical families, scored for stability
|
|
12
|
+
(MLIP formation energy), uniqueness, novelty and, where a DFT value exists, the target
|
|
13
|
+
property_prediction properties predicted from the crystal structure, on the test split
|
|
14
|
+
representation cross-modal retrieval, latent agreement and a k-nearest-neighbour probe, on the test split
|
|
15
|
+
|
|
16
|
+
Stability needs an MLIP (MACE-MP-0); it runs in a separate process (``scripts/mlip_stability.py``) so that the
|
|
17
|
+
environment holding MACE can differ from the one holding MEIDNet.
|
|
18
|
+
"""
|
|
19
|
+
from __future__ import annotations
|
|
20
|
+
|
|
21
|
+
import csv
|
|
22
|
+
import json
|
|
23
|
+
import math
|
|
24
|
+
import os
|
|
25
|
+
import random
|
|
26
|
+
from dataclasses import dataclass, field
|
|
27
|
+
|
|
28
|
+
import numpy as np
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
# ───────────────────────── small helpers ─────────────────────────
|
|
32
|
+
def formula_key(formula: str) -> str:
|
|
33
|
+
"""Reduced formula used for novelty and uniqueness (BaTiO3, not Ba1Ti1O3)."""
|
|
34
|
+
from pymatgen.core import Composition
|
|
35
|
+
try:
|
|
36
|
+
return Composition(str(formula)).reduced_formula
|
|
37
|
+
except Exception:
|
|
38
|
+
return str(formula).strip()
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _finite(v):
|
|
42
|
+
return v is not None and isinstance(v, (int, float)) and math.isfinite(v)
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def regression(y: np.ndarray, p: np.ndarray) -> dict:
|
|
46
|
+
"""MAE, RMSE and R2 of predictions p for targets y."""
|
|
47
|
+
y = np.asarray(y, dtype=float)
|
|
48
|
+
p = np.asarray(p, dtype=float)
|
|
49
|
+
err = p - y
|
|
50
|
+
ss_tot = float(((y - y.mean()) ** 2).sum())
|
|
51
|
+
return {"mae": float(np.abs(err).mean()), "rmse": float(np.sqrt((err ** 2).mean())),
|
|
52
|
+
"r2": float(1.0 - (err ** 2).sum() / ss_tot) if ss_tot > 0 else float("nan")}
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def property_metrics(columns: list[str], Y: np.ndarray, P: np.ndarray) -> dict:
|
|
56
|
+
"""mae_/rmse_/r2_<column>, plus mae_<column>_nonzero on the materials whose true value is not zero (for
|
|
57
|
+
properties such as a band gap that is zero for most materials, the error that matters for design)."""
|
|
58
|
+
out = {}
|
|
59
|
+
for j, c in enumerate(columns):
|
|
60
|
+
r = regression(Y[:, j], P[:, j])
|
|
61
|
+
out[f"mae_{c}"], out[f"rmse_{c}"], out[f"r2_{c}"] = r["mae"], r["rmse"], r["r2"]
|
|
62
|
+
nz = Y[:, j] != 0
|
|
63
|
+
if 0 < nz.sum() < len(nz):
|
|
64
|
+
out[f"mae_{c}_nonzero"] = float(np.abs(P[nz, j] - Y[nz, j]).mean())
|
|
65
|
+
out["n_evaluated"] = float(len(Y))
|
|
66
|
+
return out
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def retrieval(Zc: np.ndarray, Zp: np.ndarray, Y: np.ndarray | None = None) -> dict:
|
|
70
|
+
"""Structure -> property retrieval (unit-length latents).
|
|
71
|
+
|
|
72
|
+
With Y (the property values), the candidates are the distinct property profiles of the split: materials with
|
|
73
|
+
identical values have identical property latents, so retrieving either is equally right, and counting them
|
|
74
|
+
separately would make the score depend on how ties are broken. Chance level is then 1 / n_profiles."""
|
|
75
|
+
cos = (Zc * Zp).sum(1)
|
|
76
|
+
if Y is not None:
|
|
77
|
+
_, first, inv = np.unique(np.round(np.asarray(Y, dtype=float), 8), axis=0, return_index=True, return_inverse=True)
|
|
78
|
+
inv = np.asarray(inv).reshape(-1)
|
|
79
|
+
S = Zc @ Zp[first].T # one column per distinct profile
|
|
80
|
+
own = S[np.arange(len(Zc)), inv]
|
|
81
|
+
n_profiles = len(first)
|
|
82
|
+
else:
|
|
83
|
+
S = Zc @ Zp.T
|
|
84
|
+
own = np.diag(S)
|
|
85
|
+
n_profiles = len(Zp)
|
|
86
|
+
rank = (S > own[:, None]).sum(1) # how many profiles beat the material's own
|
|
87
|
+
return {"retrieval_top1": float((rank == 0).mean()), "retrieval_top5": float((rank < 5).mean()),
|
|
88
|
+
"cosine_matched": float(cos.mean()),
|
|
89
|
+
"l2_matched": float(np.sqrt(np.maximum(0.0, 2.0 - 2.0 * cos)).mean()), "n_profiles": float(n_profiles)}
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def knn_predict(train_X: np.ndarray, train_Y: np.ndarray, test_X: np.ndarray, k: int = 5,
|
|
93
|
+
metric: str = "cosine", chunk: int = 512) -> np.ndarray:
|
|
94
|
+
"""Mean property of the k nearest training points (cosine for unit latents, euclidean for features)."""
|
|
95
|
+
out = np.zeros((len(test_X), train_Y.shape[1]))
|
|
96
|
+
if metric == "cosine":
|
|
97
|
+
A = train_X / np.maximum(np.linalg.norm(train_X, axis=1, keepdims=True), 1e-12)
|
|
98
|
+
else:
|
|
99
|
+
a2 = (train_X ** 2).sum(1)
|
|
100
|
+
for i in range(0, len(test_X), chunk):
|
|
101
|
+
x = test_X[i:i + chunk]
|
|
102
|
+
if metric == "cosine":
|
|
103
|
+
xn = x / np.maximum(np.linalg.norm(x, axis=1, keepdims=True), 1e-12)
|
|
104
|
+
d = -(xn @ A.T)
|
|
105
|
+
else:
|
|
106
|
+
d = a2[None, :] - 2 * x @ train_X.T + (x ** 2).sum(1)[:, None]
|
|
107
|
+
idx = np.argpartition(d, kth=min(k, d.shape[1] - 1) - 1, axis=1)[:, :k]
|
|
108
|
+
out[i:i + chunk] = train_Y[idx].mean(1)
|
|
109
|
+
return out
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
def composition_features(formulas: list[str], elements: list[str] | None = None) -> tuple[np.ndarray, list[str]]:
|
|
113
|
+
"""Element-fraction vectors (the usual composition-only representation)."""
|
|
114
|
+
from pymatgen.core import Composition
|
|
115
|
+
comps = []
|
|
116
|
+
for f in formulas:
|
|
117
|
+
try:
|
|
118
|
+
comps.append(Composition(str(f)).fractional_composition.get_el_amt_dict())
|
|
119
|
+
except Exception:
|
|
120
|
+
comps.append({})
|
|
121
|
+
if elements is None:
|
|
122
|
+
elements = sorted({e for c in comps for e in c})
|
|
123
|
+
pos = {e: i for i, e in enumerate(elements)}
|
|
124
|
+
X = np.zeros((len(formulas), len(elements)))
|
|
125
|
+
for r, c in enumerate(comps):
|
|
126
|
+
for e, v in c.items():
|
|
127
|
+
if e in pos:
|
|
128
|
+
X[r, pos[e]] = v
|
|
129
|
+
return X, elements
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
# ───────────────────────── data and latents ─────────────────────────
|
|
133
|
+
def load_split(data_dir: str, split: str, columns: list[str], max_sites: int):
|
|
134
|
+
"""Records of one split, read exactly as training reads them (atoms in file order, as in the paper)."""
|
|
135
|
+
from meidnet.config import DataSection, PropertyColumn
|
|
136
|
+
from meidnet.data import load_records, read_table
|
|
137
|
+
table = os.path.join(data_dir, f"{split}.csv")
|
|
138
|
+
if not os.path.exists(table):
|
|
139
|
+
raise SystemExit(f"{table} is missing: run `meidnet download-data` first")
|
|
140
|
+
df = read_table(table)
|
|
141
|
+
dcfg = DataSection(table=table, id_column="material_id", cif_column="cif", align_to_prototype=False,
|
|
142
|
+
max_sites=max_sites, properties=[PropertyColumn(column=c) for c in columns])
|
|
143
|
+
recs, _ = load_records(df, dcfg, lambda p: p, None, source=f"{split} split")
|
|
144
|
+
formulas = dict(zip(df["material_id"].astype(str), df["formula"].astype(str))) if "formula" in df else {}
|
|
145
|
+
return recs, formulas
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
SPACES = ("projection", "encoder")
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
def encode(lm, records, batch: int = 128, space: str = "projection"):
|
|
152
|
+
"""Structure latents, property latents and structure-only property predictions (physical units).
|
|
153
|
+
|
|
154
|
+
`space` is where the two latents are read: "projection" (after the projection heads: what the decoders read) or
|
|
155
|
+
"encoder" (the normalised encoder outputs, before the heads). A model aligns its modalities in one of the two,
|
|
156
|
+
the one its contrastive loss acts on; measured in the other, the same model can look unaligned."""
|
|
157
|
+
import torch
|
|
158
|
+
import torch.nn.functional as F
|
|
159
|
+
from torch.utils.data import DataLoader
|
|
160
|
+
from meidnet.data import MaterialsDataset
|
|
161
|
+
if space not in SPACES:
|
|
162
|
+
raise ValueError(f"space must be one of {SPACES}")
|
|
163
|
+
model = lm.model.eval()
|
|
164
|
+
dev = next(model.parameters()).device
|
|
165
|
+
Zc, Zp, P = [], [], []
|
|
166
|
+
with torch.no_grad():
|
|
167
|
+
for b in DataLoader(MaterialsDataset(records, lm.stats), batch_size=batch, shuffle=False):
|
|
168
|
+
cv, props = b["crystal_vec"].to(dev), b["props"].to(dev)
|
|
169
|
+
zc, zp, *_ = model.encode_modalities(cv, props)
|
|
170
|
+
P.append(lm.stats.denormalize_tensor(model.property_decoder(zc)).cpu().numpy())
|
|
171
|
+
if space == "encoder":
|
|
172
|
+
zc = F.normalize(model.crystal_encoder(cv)[0], p=2, dim=1)
|
|
173
|
+
zp = F.normalize(model.property_encoder(props), p=2, dim=1)
|
|
174
|
+
Zc.append(zc.cpu().numpy())
|
|
175
|
+
Zp.append(zp.cpu().numpy())
|
|
176
|
+
return np.concatenate(Zc), np.concatenate(Zp), np.concatenate(P)
|
|
177
|
+
|
|
178
|
+
|
|
179
|
+
@dataclass
|
|
180
|
+
class ModelEvaluation:
|
|
181
|
+
property_prediction: dict
|
|
182
|
+
representation: dict
|
|
183
|
+
predictions: list[dict] = field(default_factory=list) # per test material: id, true_/pred_<col>, cos
|
|
184
|
+
|
|
185
|
+
|
|
186
|
+
def evaluate_checkpoint(lm, data_dir: str, k: int = 5, space: str = "projection") -> ModelEvaluation:
|
|
187
|
+
"""Property prediction and representation metrics of a checkpoint on the test split; the k-NN probe uses the
|
|
188
|
+
training split's structure latents. The representation is read in `space` (see `encode`); the matched cosine
|
|
189
|
+
in both spaces is reported next to it."""
|
|
190
|
+
cols = list(lm.stats.columns)
|
|
191
|
+
ms = lm.model.max_sites
|
|
192
|
+
train, _ = load_split(data_dir, "train", cols, ms)
|
|
193
|
+
test, _ = load_split(data_dir, "test", cols, ms)
|
|
194
|
+
Zc_tr, _, _ = encode(lm, train, space=space)
|
|
195
|
+
Zc, Zp, P = encode(lm, test, space=space)
|
|
196
|
+
Y = np.array([r.properties for r in test], dtype=float)
|
|
197
|
+
Y_tr = np.array([r.properties for r in train], dtype=float)
|
|
198
|
+
prop = property_metrics(cols, Y, P)
|
|
199
|
+
rep = retrieval(Zc, Zp, Y)
|
|
200
|
+
knn = knn_predict(Zc_tr, Y_tr, Zc, k=k, metric="cosine")
|
|
201
|
+
for j, c in enumerate(cols):
|
|
202
|
+
rep[f"knn_mae_{c}"] = float(np.abs(knn[:, j] - Y[:, j]).mean())
|
|
203
|
+
rep["n_evaluated"] = float(len(test))
|
|
204
|
+
cos = (Zc * Zp).sum(1)
|
|
205
|
+
other = next(s for s in SPACES if s != space)
|
|
206
|
+
Zc_o, Zp_o, _ = encode(lm, test, space=other)
|
|
207
|
+
rep[f"cosine_{space}"] = rep["cosine_matched"]
|
|
208
|
+
rep[f"cosine_{other}"] = float((Zc_o * Zp_o).sum(1).mean())
|
|
209
|
+
preds = [{"id": r.material_id, **{f"true_{c}": float(Y[i, j]) for j, c in enumerate(cols)},
|
|
210
|
+
**{f"pred_{c}": float(P[i, j]) for j, c in enumerate(cols)}, "cos": float(cos[i])}
|
|
211
|
+
for i, r in enumerate(test)]
|
|
212
|
+
return ModelEvaluation(prop, rep, preds)
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
def evaluate_baselines(data_dir: str, columns: list[str], max_sites: int = 20, k: int = 5) -> dict:
|
|
216
|
+
"""Composition k-NN and training-mean baselines (property prediction), the composition k-NN probe and the
|
|
217
|
+
chance level of retrieval (representation)."""
|
|
218
|
+
train, f_tr = load_split(data_dir, "train", columns, max_sites)
|
|
219
|
+
test, f_te = load_split(data_dir, "test", columns, max_sites)
|
|
220
|
+
Y_tr = np.array([r.properties for r in train], dtype=float)
|
|
221
|
+
Y = np.array([r.properties for r in test], dtype=float)
|
|
222
|
+
X_tr, els = composition_features([f_tr.get(r.material_id, "") for r in train])
|
|
223
|
+
X_te, _ = composition_features([f_te.get(r.material_id, "") for r in test], els)
|
|
224
|
+
P_knn = knn_predict(X_tr, Y_tr, X_te, k=k, metric="euclidean")
|
|
225
|
+
P_mean = np.repeat(Y_tr.mean(0, keepdims=True), len(Y), axis=0)
|
|
226
|
+
n = len(Y)
|
|
227
|
+
n_prof = len(np.unique(np.round(Y, 8), axis=0))
|
|
228
|
+
knn_rep = {f"knn_mae_{c}": float(np.abs(P_knn[:, j] - Y[:, j]).mean()) for j, c in enumerate(columns)}
|
|
229
|
+
knn_rep["n_evaluated"] = float(n)
|
|
230
|
+
return {
|
|
231
|
+
"composition_knn": {"property_prediction": property_metrics(columns, Y, P_knn), "representation": knn_rep},
|
|
232
|
+
"train_mean": {"property_prediction": property_metrics(columns, Y, P_mean)},
|
|
233
|
+
"chance": {"representation": {"retrieval_top1": 1.0 / n_prof, "retrieval_top5": min(1.0, 5.0 / n_prof),
|
|
234
|
+
"cosine_matched": 0.0, "n_evaluated": float(n), "n_profiles": float(n_prof)}},
|
|
235
|
+
"predictions": {"composition_knn": [{"id": r.material_id, **{f"true_{c}": float(Y[i, j]) for j, c in enumerate(columns)},
|
|
236
|
+
**{f"pred_{c}": float(P_knn[i, j]) for j, c in enumerate(columns)}}
|
|
237
|
+
for i, r in enumerate(test)]},
|
|
238
|
+
}
|
|
239
|
+
|
|
240
|
+
|
|
241
|
+
# ───────────────────────── inverse design: candidates ─────────────────────────
|
|
242
|
+
@dataclass
|
|
243
|
+
class Candidate:
|
|
244
|
+
id: str
|
|
245
|
+
variant: str
|
|
246
|
+
target: int # 1-based index into the protocol's targets
|
|
247
|
+
targets: dict # the target values
|
|
248
|
+
formula: str
|
|
249
|
+
elements: dict
|
|
250
|
+
cif: str # path relative to the run folder
|
|
251
|
+
passes_rules: bool = True
|
|
252
|
+
|
|
253
|
+
|
|
254
|
+
def _family(family: str, variant: str):
|
|
255
|
+
from meidnet.family import load_family
|
|
256
|
+
return load_family(family, variant=variant)
|
|
257
|
+
|
|
258
|
+
|
|
259
|
+
def _write_cif(cand_struct, path: str):
|
|
260
|
+
from pymatgen.io.cif import CifWriter
|
|
261
|
+
os.makedirs(os.path.dirname(path), exist_ok=True)
|
|
262
|
+
CifWriter(cand_struct).write_file(path)
|
|
263
|
+
|
|
264
|
+
|
|
265
|
+
def passing_compositions(family: str, variant: str) -> list[dict]:
|
|
266
|
+
"""Every composition of the family variant that passes all of its rules, in a fixed order."""
|
|
267
|
+
from meidnet.designspace import enumerate_space
|
|
268
|
+
space = enumerate_space(_family(family, variant), None)
|
|
269
|
+
return [r["e"] for r in space["rows"] if all(r["ok"].values())]
|
|
270
|
+
|
|
271
|
+
|
|
272
|
+
def _materialise(family: str, variant: str, elements: dict, path: str) -> bool:
|
|
273
|
+
"""Build the prototype structure of a composition, check the rules, write the CIF."""
|
|
274
|
+
from meidnet.constraints import build_candidate, evaluate
|
|
275
|
+
fam = _family(family, variant)
|
|
276
|
+
cand = build_candidate(fam, elements)
|
|
277
|
+
evaluate(cand, fam.constraints)
|
|
278
|
+
ok = all(r.passed for r in cand.results)
|
|
279
|
+
_write_cif(getattr(cand, "structure", None) or cand.raw, path)
|
|
280
|
+
return ok
|
|
281
|
+
|
|
282
|
+
|
|
283
|
+
def design_random(settings: dict, run_dir: str, seed: int = 0) -> list[Candidate]:
|
|
284
|
+
"""Baseline: compositions drawn at random (without replacement) from those that pass the family's rules."""
|
|
285
|
+
out = []
|
|
286
|
+
K, T = settings["per_target"], settings["targets"]
|
|
287
|
+
for vi, variant in enumerate(settings["variants"]):
|
|
288
|
+
pool = passing_compositions(settings["family"], variant)
|
|
289
|
+
rng = random.Random(seed * 1000 + vi)
|
|
290
|
+
picked = rng.sample(pool, min(len(pool), K * len(T)))
|
|
291
|
+
for n, el in enumerate(picked):
|
|
292
|
+
t = n // K + 1
|
|
293
|
+
cid = f"{variant}_T{t}_{n % K + 1}"
|
|
294
|
+
path = os.path.join("cifs", cid + ".cif")
|
|
295
|
+
ok = _materialise(settings["family"], variant, el, os.path.join(run_dir, path))
|
|
296
|
+
out.append(Candidate(cid, variant, t, T[t - 1], _formula_of(settings["family"], variant, el), el, path, ok))
|
|
297
|
+
return out
|
|
298
|
+
|
|
299
|
+
|
|
300
|
+
def _formula_of(family: str, variant: str, elements: dict) -> str:
|
|
301
|
+
from meidnet.constraints import build_candidate
|
|
302
|
+
return build_candidate(_family(family, variant), elements).formula()
|
|
303
|
+
|
|
304
|
+
|
|
305
|
+
def design_screening(lm, settings: dict, objectives: list[dict], run_dir: str) -> list[Candidate]:
|
|
306
|
+
"""Baseline: rank the rule-passing compositions by the structure encoder's predictions (the model's forward
|
|
307
|
+
direction, no latent search) and keep the closest to each target, using the generator's own selection score."""
|
|
308
|
+
from meidnet.designspace import enumerate_space
|
|
309
|
+
from meidnet.generate import objective_distance
|
|
310
|
+
out = []
|
|
311
|
+
K, T = settings["per_target"], settings["targets"]
|
|
312
|
+
for variant in settings["variants"]:
|
|
313
|
+
space = enumerate_space(_family(settings["family"], variant), lm)
|
|
314
|
+
rows = [r for r in space["rows"] if all(r["ok"].values())]
|
|
315
|
+
taken = set()
|
|
316
|
+
for t, target in enumerate(T, start=1):
|
|
317
|
+
def score(r):
|
|
318
|
+
return sum(o.get("select_weight", 1.0) * objective_distance(r["p"][o["property"]], target[o["property"]],
|
|
319
|
+
o.get("select_loss") or o.get("loss", "l2"))
|
|
320
|
+
for o in objectives)
|
|
321
|
+
n = 0
|
|
322
|
+
for r in sorted(rows, key=score):
|
|
323
|
+
key = tuple(sorted(r["e"].items()))
|
|
324
|
+
if key in taken:
|
|
325
|
+
continue
|
|
326
|
+
taken.add(key)
|
|
327
|
+
n += 1
|
|
328
|
+
cid = f"{variant}_T{t}_{n}"
|
|
329
|
+
path = os.path.join("cifs", cid + ".cif")
|
|
330
|
+
ok = _materialise(settings["family"], variant, r["e"], os.path.join(run_dir, path))
|
|
331
|
+
out.append(Candidate(cid, variant, t, target, r["f"], r["e"], path, ok))
|
|
332
|
+
if n >= K:
|
|
333
|
+
break
|
|
334
|
+
return out
|
|
335
|
+
|
|
336
|
+
|
|
337
|
+
def design_meidnet(lm, generation, settings: dict, run_dir: str, log=print, device=None) -> list[Candidate]:
|
|
338
|
+
"""MEIDNet's latent search, run per chemical family with the protocol's targets and budget."""
|
|
339
|
+
import torch
|
|
340
|
+
from meidnet.checkpoint import property_ranges
|
|
341
|
+
from meidnet.generate import Designer
|
|
342
|
+
device = device or torch.device("cpu")
|
|
343
|
+
out = []
|
|
344
|
+
for variant in settings["variants"]:
|
|
345
|
+
g = generation.model_copy(deep=True)
|
|
346
|
+
g.family, g.variant = settings["family"], variant
|
|
347
|
+
g.targets = [dict(t) for t in settings["targets"]]
|
|
348
|
+
g.per_target = settings["per_target"]
|
|
349
|
+
g.rounds = settings["max_rounds"]
|
|
350
|
+
g.output_prefix = variant
|
|
351
|
+
fam = _family(settings["family"], variant)
|
|
352
|
+
gen_dir = os.path.join(run_dir, "generation", variant)
|
|
353
|
+
res = Designer(lm, fam, g, device=device, log=log).run(gen_dir, ranges=property_ranges(lm))
|
|
354
|
+
for c in res.saved:
|
|
355
|
+
n = sum(1 for x in out if x.variant == variant and x.target == c.target_index) + 1
|
|
356
|
+
cid = f"{variant}_T{c.target_index}_{n}"
|
|
357
|
+
src = os.path.join(gen_dir, c.file)
|
|
358
|
+
path = os.path.join("cifs", cid + ".cif")
|
|
359
|
+
os.makedirs(os.path.join(run_dir, "cifs"), exist_ok=True)
|
|
360
|
+
with open(src, encoding="utf-8") as fi, open(os.path.join(run_dir, path), "w", encoding="utf-8", newline="\n") as fo:
|
|
361
|
+
fo.write(fi.read())
|
|
362
|
+
ok = all(r.get("passed", True) for r in c.constraint_results)
|
|
363
|
+
out.append(Candidate(cid, variant, c.target_index, settings["targets"][c.target_index - 1], c.formula,
|
|
364
|
+
dict(c.elements), path, ok))
|
|
365
|
+
return out
|
|
366
|
+
|
|
367
|
+
|
|
368
|
+
def write_candidates(cands: list[Candidate], path: str) -> None:
|
|
369
|
+
keys = sorted({k for c in cands for k in c.targets})
|
|
370
|
+
with open(path, "w", newline="", encoding="utf-8") as f:
|
|
371
|
+
w = csv.writer(f)
|
|
372
|
+
w.writerow(["id", "variant", "target", *[f"target_{k}" for k in keys], "formula", "elements", "cif", "passes_rules"])
|
|
373
|
+
for c in cands:
|
|
374
|
+
w.writerow([c.id, c.variant, c.target, *[c.targets.get(k, "") for k in keys], c.formula,
|
|
375
|
+
json.dumps(c.elements, sort_keys=True), c.cif.replace(os.sep, "/"), c.passes_rules])
|
|
376
|
+
|
|
377
|
+
|
|
378
|
+
def read_candidates(path: str) -> list[dict]:
|
|
379
|
+
with open(path, newline="", encoding="utf-8") as f:
|
|
380
|
+
return list(csv.DictReader(f))
|
|
381
|
+
|
|
382
|
+
|
|
383
|
+
# ───────────────────────── inverse design: scoring ─────────────────────────
|
|
384
|
+
ANIONS = {"O", "N", "F", "Cl", "Br", "I", "S", "Se", "Te", "H", "C", "P", "As", "Sb"}
|
|
385
|
+
|
|
386
|
+
|
|
387
|
+
def site_key(elements: dict) -> str | None:
|
|
388
|
+
"""A|B|X key of an ABX3 composition with one anion, from a candidate's site groups."""
|
|
389
|
+
if all(g in elements for g in ("A", "B", "X")):
|
|
390
|
+
return f"{elements['A']}|{elements['B']}|{elements['X']}"
|
|
391
|
+
return None
|
|
392
|
+
|
|
393
|
+
|
|
394
|
+
def _structure_site_key(structure) -> str | None:
|
|
395
|
+
"""A|B|X key of a 5-atom cubic perovskite cell: the anion is the element on three sites, B the cation closest
|
|
396
|
+
to an anion (centre of the octahedron, at a/2) and A the other cation (at a/sqrt 2)."""
|
|
397
|
+
species = [str(sp) for sp in structure.species]
|
|
398
|
+
anion_sites = [i for i, e in enumerate(species) if e in ANIONS]
|
|
399
|
+
cations = [i for i in range(len(species)) if i not in anion_sites]
|
|
400
|
+
if len(species) != 5 or len(anion_sites) != 3 or len({species[i] for i in anion_sites}) != 1 or len(cations) != 2:
|
|
401
|
+
return None
|
|
402
|
+
d = [min(structure.get_distance(i, j) for j in anion_sites) for i in cations]
|
|
403
|
+
b = cations[int(np.argmin(d))]
|
|
404
|
+
a = cations[1 - int(np.argmin(d))]
|
|
405
|
+
return f"{species[a]}|{species[b]}|{species[anion_sites[0]]}"
|
|
406
|
+
|
|
407
|
+
|
|
408
|
+
def known_by_site(data_dir: str, column: str, cache: str | None = None) -> dict[str, list[float]]:
|
|
409
|
+
"""A|B|X -> DFT values of `column` over all splits, for the single-anion entries of the dataset. Perov-5 holds
|
|
410
|
+
both site orderings of many compositions (BaTiO3 with Ba on A and with Ti on A), with very different band gaps,
|
|
411
|
+
so the check matches sites, not formulas. Cached as JSON because it parses every CIF once."""
|
|
412
|
+
if cache and os.path.exists(cache):
|
|
413
|
+
with open(cache, encoding="utf-8") as f:
|
|
414
|
+
return json.load(f)
|
|
415
|
+
import pandas as pd
|
|
416
|
+
from pymatgen.core import Structure
|
|
417
|
+
out: dict[str, list[float]] = {}
|
|
418
|
+
for split in ("train", "val", "test"):
|
|
419
|
+
p = os.path.join(data_dir, f"{split}.csv")
|
|
420
|
+
if not os.path.exists(p):
|
|
421
|
+
continue
|
|
422
|
+
df = pd.read_csv(p, usecols=["cif", column])
|
|
423
|
+
for cif, v in zip(df["cif"], df[column]):
|
|
424
|
+
try:
|
|
425
|
+
key = _structure_site_key(Structure.from_str(cif, fmt="cif"))
|
|
426
|
+
except Exception:
|
|
427
|
+
key = None
|
|
428
|
+
if key:
|
|
429
|
+
out.setdefault(key, []).append(float(v))
|
|
430
|
+
if cache:
|
|
431
|
+
os.makedirs(os.path.dirname(cache), exist_ok=True)
|
|
432
|
+
with open(cache, "w", encoding="utf-8") as f:
|
|
433
|
+
json.dump(out, f, sort_keys=True)
|
|
434
|
+
return out
|
|
435
|
+
|
|
436
|
+
|
|
437
|
+
def known_values(data_dir: str, column: str) -> dict[str, list[float]]:
|
|
438
|
+
"""Reduced formula -> DFT values of `column` over all splits (the materials whose property is known)."""
|
|
439
|
+
import pandas as pd
|
|
440
|
+
out: dict[str, list[float]] = {}
|
|
441
|
+
for split in ("train", "val", "test"):
|
|
442
|
+
p = os.path.join(data_dir, f"{split}.csv")
|
|
443
|
+
if os.path.exists(p):
|
|
444
|
+
df = pd.read_csv(p, usecols=["formula", column])
|
|
445
|
+
for f, v in zip(df["formula"], df[column]):
|
|
446
|
+
out.setdefault(formula_key(f), []).append(float(v))
|
|
447
|
+
return out
|
|
448
|
+
|
|
449
|
+
|
|
450
|
+
def training_formulas(data_dir: str) -> set[str]:
|
|
451
|
+
import pandas as pd
|
|
452
|
+
df = pd.read_csv(os.path.join(data_dir, "train.csv"), usecols=["formula"])
|
|
453
|
+
return {formula_key(f) for f in df["formula"]}
|
|
454
|
+
|
|
455
|
+
|
|
456
|
+
def dataset_formulas(data_dir: str) -> set[str]:
|
|
457
|
+
"""Every composition of the data set (training, validation and test splits): the novelty reference. A candidate
|
|
458
|
+
is novel when the data set does not contain it, whichever part of it a model was trained on."""
|
|
459
|
+
import pandas as pd
|
|
460
|
+
out = set()
|
|
461
|
+
for split in ("train", "val", "test"):
|
|
462
|
+
p = os.path.join(data_dir, f"{split}.csv")
|
|
463
|
+
if os.path.exists(p):
|
|
464
|
+
out |= {formula_key(f) for f in pd.read_csv(p, usecols=["formula"])["formula"]}
|
|
465
|
+
return out
|
|
466
|
+
|
|
467
|
+
|
|
468
|
+
def score_inverse_design(candidates: list[dict], stability: dict[str, dict], settings: dict, train_formulas: set[str],
|
|
469
|
+
known: dict[str, list[float]], gap_column: str = "dir_gap") -> tuple[dict, list[dict]]:
|
|
470
|
+
"""Metrics of a candidate set under the protocol. The SUN rate is over the requested budget
|
|
471
|
+
(variants x targets x per_target): a candidate that was not delivered counts as a failure. Stable, unique and
|
|
472
|
+
novel are fractions of the delivered candidates.
|
|
473
|
+
|
|
474
|
+
stability: candidate id -> {"dHf": float, "artefact": bool} (from scripts/mlip_stability.py).
|
|
475
|
+
known: A|B|X site key (or reduced formula) -> DFT values of the gap column (see known_by_site)."""
|
|
476
|
+
budget = settings.get("budget") or len(settings["variants"]) * len(settings["targets"]) * settings["per_target"]
|
|
477
|
+
thr, tol = settings["stability_threshold"], settings["gap_tolerance"]
|
|
478
|
+
seen, rows = set(), []
|
|
479
|
+
n_valid = n_unique = n_novel = n_stable = n_sun = known_n = hits = 0
|
|
480
|
+
dhf = []
|
|
481
|
+
for c in candidates:
|
|
482
|
+
key = formula_key(c["formula"])
|
|
483
|
+
valid = str(c.get("passes_rules", "True")) in ("True", "true", "1")
|
|
484
|
+
unique = key not in seen
|
|
485
|
+
seen.add(key)
|
|
486
|
+
novel = key not in train_formulas
|
|
487
|
+
st = stability.get(c["id"], {})
|
|
488
|
+
e = st.get("dHf")
|
|
489
|
+
artefact = bool(st.get("artefact", True)) if st else True
|
|
490
|
+
stable = (not artefact) and _finite(e) and e <= thr
|
|
491
|
+
if _finite(e) and not artefact:
|
|
492
|
+
dhf.append(e)
|
|
493
|
+
sun = valid and stable and unique and novel
|
|
494
|
+
target_gap = float(c.get(f"target_{gap_column}", "nan") or "nan")
|
|
495
|
+
el = c.get("elements") or {}
|
|
496
|
+
if isinstance(el, str):
|
|
497
|
+
try:
|
|
498
|
+
el = json.loads(el)
|
|
499
|
+
except ValueError:
|
|
500
|
+
el = {}
|
|
501
|
+
dft = known.get(site_key(el) or key)
|
|
502
|
+
hit = None
|
|
503
|
+
if dft is not None and math.isfinite(target_gap):
|
|
504
|
+
known_n += 1
|
|
505
|
+
hit = min(abs(v - target_gap) for v in dft) <= tol
|
|
506
|
+
hits += int(hit)
|
|
507
|
+
n_valid += valid
|
|
508
|
+
n_unique += unique
|
|
509
|
+
n_novel += novel
|
|
510
|
+
n_stable += stable
|
|
511
|
+
n_sun += sun
|
|
512
|
+
rows.append({**c, "formula_key": key, "valid": valid, "unique": unique, "novel": novel,
|
|
513
|
+
"dHf": e if _finite(e) else None, "stable": stable, "sun": sun,
|
|
514
|
+
"dft_" + gap_column: (min(dft, key=lambda v: abs(v - target_gap)) if dft else None), "dft_hit": hit})
|
|
515
|
+
nd = max(len(candidates), 1)
|
|
516
|
+
metrics = {
|
|
517
|
+
"sun_rate": n_sun / budget, "stable_rate": n_stable / nd, "unique_rate": n_unique / nd,
|
|
518
|
+
"novel_rate": n_novel / nd, "valid_rate": n_valid / budget,
|
|
519
|
+
"dhf_median": float(np.median(dhf)) if dhf else float("nan"),
|
|
520
|
+
"dft_hit_rate": hits / known_n if known_n else float("nan"), "dft_known": float(known_n),
|
|
521
|
+
"n_delivered": float(len(candidates)), "n_sun": float(n_sun), "n_budget": float(budget),
|
|
522
|
+
}
|
|
523
|
+
return metrics, rows
|
|
524
|
+
|
|
525
|
+
|
|
526
|
+
def read_stability(path: str) -> dict[str, dict]:
|
|
527
|
+
out = {}
|
|
528
|
+
if not os.path.exists(path):
|
|
529
|
+
return out
|
|
530
|
+
with open(path, newline="", encoding="utf-8") as f:
|
|
531
|
+
for r in csv.DictReader(f):
|
|
532
|
+
try:
|
|
533
|
+
e = float(r.get("dHf", "nan"))
|
|
534
|
+
except ValueError:
|
|
535
|
+
e = float("nan")
|
|
536
|
+
out[r["id"]] = {"dHf": e, "artefact": str(r.get("artefact", "True")) in ("True", "true", "1"),
|
|
537
|
+
"converged": str(r.get("converged", "False")) in ("True", "true", "1")}
|
|
538
|
+
return out
|
|
539
|
+
|
|
540
|
+
|
|
541
|
+
__all__ = ["formula_key", "regression", "property_metrics", "retrieval", "knn_predict", "composition_features",
|
|
542
|
+
"load_split", "encode", "evaluate_checkpoint", "evaluate_baselines", "Candidate", "passing_compositions",
|
|
543
|
+
"design_random", "design_screening", "design_meidnet", "write_candidates", "read_candidates",
|
|
544
|
+
"known_values", "known_by_site", "site_key", "training_formulas", "dataset_formulas", "score_inverse_design", "read_stability",
|
|
545
|
+
"SPACES"]
|