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 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"]