"""Builds the training table from the extracted records. Input: data/alexandria_ph/extract/*.jsonl.gz (written by model/extract.py) Output: model/data/labels.csv one row per compound with a spectral function data/alexandria_ph/dataset.pkl the same rows with cells, descriptors and graphs Labels are taken at one smearing width, 0.030 Ry, the width the release uses for its own transition temperatures. S Hopfield sum of the compound, M_H * lambda * , eV/A^2 S_top the part of S in the 3 n_H highest branches h S_top / rho_H, eV A, defined only where every one of those branches has at least 90% of its eigenvector weight on hydrogen at every stored wavevector (the branch criterion of the paper) Each compound gets a structural prototype: space group, reduced stoichiometry and the Wyckoff letters occupied by each element, with element names removed. Splits assign whole prototypes. Run: .venv/bin/python model/dataset.py """ import csv import gzip import json import pathlib import pickle import sys from collections import Counter from concurrent.futures import ProcessPoolExecutor from functools import reduce from math import gcd import numpy as np import spglib from pymatgen.core import Element, Lattice, Structure ROOT = pathlib.Path(__file__).resolve().parents[1] EXTRACT = ROOT / 'data' / 'alexandria_ph' / 'extract' WIDTH = '0.030' BRANCH_MIN = 0.90 CUTOFF = 5.0 # A, neighbour list for the graph model M_H = 1.00794 def valence(el): """Electrons outside the noble-gas core, by group; 3 for lanthanides and actinides.""" if el.is_lanthanoid or el.is_actinoid: return 3 g = el.group return g if g <= 12 else g - 10 def element_row(sym): el = Element(sym) return dict(Z=el.Z, X=el.X if el.X == el.X and el.X is not None else 1.0, r=float(el.atomic_radius) if el.atomic_radius else 1.5, m=float(el.atomic_mass), group=el.group, row=el.row, val=valence(el)) ELEMENTS = {} def el(sym): if sym not in ELEMENTS: ELEMENTS[sym] = element_row(sym) return ELEMENTS[sym] def prototype(lattice, frac, species): """Space group number, and a label that is the same for isostructural compounds.""" numbers = [Element(s).Z for s in species] ds = spglib.get_symmetry_dataset((lattice, frac, numbers), symprec=0.05) if ds is None: return 1, 'unresolved-' + '-'.join(sorted(Counter(species).values().__str__())) counts = Counter(species) g = reduce(gcd, counts.values()) per = {} for s, w in zip(species, ds.wyckoffs): per.setdefault(s, Counter())[w] += 1 parts = sorted('%d%s' % (counts[s] // g, ''.join(sorted(per[s]))) for s in counts) return int(ds.number), '%d_%s' % (ds.number, '_'.join(parts)) def stats(vals, weights=None): v = np.asarray(vals, dtype=float) return [float(v.mean()), float(v.min()), float(v.max()), float(v.std())] def describe(r): """Descriptors of the cell alone; nothing here needs an electronic-structure calculation.""" species = r['species'] n = len(species) rows = [el(s) for s in species] heavy = [x for s, x in zip(species, rows) if s != 'H'] or rows vol = r['volume'] n_val = sum(x['val'] for x in rows) f = dict(n_sites=n, vol_per_atom=vol / n, rho_H=r['rho_H'], frac_H=r['n_H'] / n, n_species=len(set(species)), val_density=n_val / vol, r_s=(3 * vol / (4 * np.pi * n_val)) ** (1 / 3) / 0.529177, heavy_density=(n - r['n_H']) / vol, mass_density=sum(x['m'] for x in rows) / vol) for key in ('Z', 'X', 'r', 'm', 'group', 'row', 'val'): for name, v in zip(('mean', 'min', 'max', 'std'), stats([x[key] for x in rows])): f[f'{key}_{name}'] = v for name, v in zip(('mean', 'min', 'max', 'std'), stats([x[key] for x in heavy])): f[f'heavy_{key}_{name}'] = v return f def graph(lattice, frac, species): st = Structure(Lattice(lattice), species, frac) i, j, _, d = st.get_neighbor_list(CUTOFF) return i.astype(np.int32), j.astype(np.int32), d.astype(np.float32) def geometry(species, i, j, d): """Shortest distances and hydrogen coordination, from the neighbour list.""" sp = np.array(species) hi, hj = sp[i] == 'H', sp[j] == 'H' out = dict(d_min=float(d.min()) if len(d) else CUTOFF) for name, mask in (('d_HH', hi & hj), ('d_HM', hi & ~hj), ('d_MM', ~hi & ~hj)): out[name] = float(d[mask].min()) if mask.any() else CUTOFF n_h = int((sp == 'H').sum()) if n_h: out['H_coord_M'] = float((hi & ~hj & (d < 2.4)).sum() / n_h) out['H_coord_H'] = float((hi & hj & (d < 2.4)).sum() / n_h) else: out['H_coord_M'] = out['H_coord_H'] = 0.0 return out def build(r): a = r['a2f'].get(WIDTH) if not a: return None lattice, frac, species = r['lattice'], r['frac'], r['species'] sg, proto = prototype(lattice, frac, species) i, j, d = graph(lattice, frac, species) branch_ok = bool(r['n_H']) and r.get('h_top_min', 0.0) >= BRANCH_MIN row = dict(id=r['id'], formula=r['formula'], spacegroup=sg, prototype=proto, n_sites=len(species), n_H=r['n_H'], volume=r['volume'], rho_H=r['rho_H'], S=a['S'], S_top=a.get('S_top'), lam=a['lam'], wlog_meV=a['wlog'] * 1000, w2_meV=a['w2'] * 1000, dos_Ef=a.get('dos'), branch_ok=branch_ok, h_top_min=r.get('h_top_min'), h=(a['S_top'] / r['rho_H']) if branch_ok else None, tc_ad=(r['tc']['allen_dynes'] or {}).get('0.10'), tc_el=(r['tc']['eliashberg'] or {}).get('0.10'), w_max_cm=r.get('w_max')) for tag, w in (('005', '0.005'), ('050', '0.050')): # the narrowest and widest smearing b = r['a2f'].get(w) or {} row['S_' + tag], row['S_top_' + tag], row['lam_' + tag] = b.get('S'), b.get('S_top'), b.get('lam') feats = describe(r) feats.update(geometry(species, i, j, d)) feats['spacegroup'] = sg return dict(row=row, feats=feats, lattice=lattice, frac=frac, species=species, edges=(i, j, d), S_branch=r.get('S_branch'), lam_branch=r.get('lam_branch'), h_weight_branch=r.get('h_weight_branch'), freqs_q=r.get('freqs_q'), a2f_widths={k: (v['S'], v['lam']) for k, v in r['a2f'].items()}) def load(path): out, n, n_imag, n_h = [], 0, 0, 0 with gzip.open(path, 'rt') as f: for line in f: r = json.loads(line) n += 1 n_imag += r['imag'] n_h += bool(r['n_H']) if 'a2f' in r: b = build(r) if b: out.append(b) return out, (n, n_imag, n_h) if __name__ == '__main__': files = sorted(EXTRACT.glob('*.jsonl.gz')) rows, total = [], np.zeros(3, dtype=int) with ProcessPoolExecutor(int(sys.argv[1]) if len(sys.argv) > 1 else 8) as pool: for part, counts in pool.map(load, files): rows.extend(part) total += counts seen, unique = set(), [] for r in rows: # one row per compound if r['row']['id'] not in seen: seen.add(r['row']['id']) unique.append(r) rows = unique out = ROOT / 'model' / 'data' out.mkdir(exist_ok=True) cols = list(rows[0]['row'].keys()) with open(out / 'labels.csv', 'w', newline='') as f: w = csv.writer(f) w.writerow(cols) for r in rows: w.writerow([('%.6g' % v if isinstance(v, float) else ('' if v is None else v)) for v in (r['row'][c] for c in cols)]) with open(ROOT / 'data' / 'alexandria_ph' / 'dataset.pkl', 'wb') as f: pickle.dump(rows, f, protocol=pickle.HIGHEST_PROTOCOL) hyd = [r for r in rows if r['row']['n_H']] ok = [r for r in hyd if r['row']['branch_ok']] summary = dict(files=len(files), records=int(total[0]), records_with_imaginary_modes=int(total[1]), hydride_records=int(total[2]), with_spectrum=len(rows), hydrides_with_spectrum=len(hyd), hydrides_passing_branch_criterion=len(ok), prototypes=len({r['row']['prototype'] for r in rows}), hydride_prototypes=len({r['row']['prototype'] for r in hyd})) json.dump(summary, open(out / 'summary.json', 'w'), indent=1) print(json.dumps(summary, indent=1))