"""roomtsc 0.1: estimates the Hopfield sum of a crystal structure and, for a hydride, the scattering strength per proton, from the cell alone. Usage: .venv/bin/python model/predict.py structure.cif [more files] (any format pymatgen reads: CIF, POSCAR, ...; prints one JSON object per file) The estimate is the mean of two estimators trained on the Alexandria electron-phonon release at one atmosphere: gradient-boosted trees on descriptors of the cell, and an ensemble of graph networks. What it is worth is measured in model/results2/: on structure types held out of training the error of each estimator in ln h is 0.40 to 0.54 of that of a fitted constant, and neither places more than 2 of 23 compounds stronger than its training range above that range. Every result carries the flags that say when a structure lies outside what the estimators have seen. """ import json import math import pathlib import sys import joblib import numpy as np import torch from pymatgen.core import Structure HERE = pathlib.Path(__file__).resolve().parent sys.path.insert(0, str(HERE)) import dataset as D import structure_types as ST import train as T W = HERE / 'weights' _cache = {} def models(): if not _cache: meta = json.load(open(W / 'meta.json')) nets = [] for seed in meta['graph_seeds']: net = T.Net() net.load_state_dict(torch.load(W / f'graph_{seed}.pt', map_location='cpu')) net.eval() nets.append(net) index = json.load(open(W / 'training_index.json')) if (W / 'training_index.json').exists() else {} _cache.update(meta=meta, nets=nets, trees=joblib.load(W / 'trees.joblib'), index=index, order=json.load(open(W / 'feature_order.json'))) return _cache def as_record(st): if not st.is_ordered: raise ValueError('the structure has partial occupancies; give an ordered cell') species = [str(s.specie.symbol) for s in st] n_h = species.count('H') return dict(lattice=st.lattice.matrix.tolist(), frac=[list(map(float, s.frac_coords)) for s in st], species=species, volume=float(st.volume), n_H=n_h, rho_H=n_h / float(st.volume)) @torch.no_grad() def predict(st): m = models() meta = m['meta'] r = as_record(st) if len(r['species']) > 200: raise ValueError('cells of more than 200 atoms are not accepted; the largest training cell has %d' % meta['largest_training_cell_atoms']) unknown = sorted(set(r['species']) - set(meta['elements_seen'])) if unknown: raise ValueError('no training data for ' + ', '.join(unknown)) i, j, d = D.graph(r['lattice'], r['frac'], r['species']) sg, _ = D.prototype(r['lattice'], r['frac'], r['species']) stype = ST.structure_type(r['lattice'], r['frac'], r['species']) feats = D.describe(r) feats.update(D.geometry(r['species'], i, j, d)) feats['spacegroup'] = sg ev = [T.element_vector(s) for s in r['species']] g = dict(x=np.stack([e[0] for e in ev]), z=np.array([e[1] for e in ev]), inv_m=np.array([T.M_H / e[2] for e in ev], dtype=np.float32), is_h=np.array([s == 'H' for s in r['species']]), i=i, j=j, d=d, vol=r['volume'], ln_S=0.0, ln_h=float('nan')) batch = T.collate([dict(g=g)], 'cpu') ln_s, ln_h, per_atom = [], [], [] for net in m['nets']: s, h, a = net(batch) ln_s.append(float(s[0])); ln_h.append(float(h[0])); per_atom.append(a.numpy()) out = dict(model='roomtsc ' + meta['version'], formula=st.composition.reduced_formula, n_sites=len(r['species']), spacegroup=sg, structure_type=stype, volume_A3=round(r['volume'], 3), hydrogen_density_per_A3=round(r['rho_H'], 4)) s_whole = math.exp(np.mean(ln_s)) out['hopfield_sum_eV_per_A2'] = dict(value=round(s_whole, 3), members=[round(math.exp(v), 3) for v in ln_s], note='whole compound, graph ensemble, hydrogen mass as the unit of mass') flags = [] if r['n_H']: x = np.array([[feats[k] for k in m['order']]], dtype=float) lt = float(m['trees'].predict(x)[0]) lg = float(np.mean(ln_h)) lh = 0.5 * (lt + lg) h = math.exp(lh) eta = r['rho_H'] * h cal = meta.get('calibration') or {} out['scattering_strength_per_proton_eV_A'] = dict(value=round(h, 2), trees=round(math.exp(lt), 2), graph=round(math.exp(lg), 2), graph_members=[round(math.exp(v), 2) for v in ln_h]) if cal.get('q90_ln'): f = math.exp(cal['q90_ln']) out['scattering_strength_per_proton_eV_A']['interval_90'] = [round(h / f, 2), round(h * f, 2)] out['scattering_strength_per_proton_eV_A']['interval_note'] = cal.get('note', '') out['hydrogen_hopfield_parameter_eV_per_A2'] = round(eta, 3) k = meta['asymptote_K_per_sqrt_eV_A2'] lo, hi = meta['efficiency_range'] asym = k * math.sqrt(eta) out['upper_scale_of_Tc_K'] = dict(asymptote=round(asym), at_efficiency_0_35=round(asym * lo), at_efficiency_0_58=round(asym * hi), note='upper scale from the hydrogen part of the Hopfield sum, and the range at the efficiencies of published hydride spectra; not a calculated transition temperature') base = meta['baselines_h'] out['baselines_eV_A'] = dict(training_mean=round(base['training_geometric_mean'], 2), electron_gas=base['electron_gas'], note='the estimators failed the test of predicting above their training range, so the estimate cannot show that a structure is stronger than any in the training set; the baselines fixed in advance are given beside it') req = meta['requirement_eta_H'] out['against_300K'] = dict(needed_single_mode=req['single_mode'], needed_at_ambient_efficiencies=req['at_ambient_efficiencies'], floor=req['floor'], fraction_of_lowest_single_mode_requirement=round(eta / req['single_mode'][0], 2)) out['per_hydrogen_strength_eV_A'] = [round(float(np.mean([a[q] for a in per_atom])), 2) for q in np.flatnonzero(g['is_h'])][:64] tr = meta['training_range'] if stype not in meta['structure_types_seen_in_labelled_hydrides']: flags.append('structure type not among the labelled hydrides of the training set; the tests on held-out structure types apply') missing = sorted(set(r['species']) - set(meta['elements_seen_in_labelled_hydrides'])) if missing: flags.append('no labelled hydride in the training set contains ' + ', '.join(missing)) if not tr['rho_H'][0] <= r['rho_H'] <= tr['rho_H'][1]: flags.append('hydrogen density outside the training range of %.3f to %.3f per A^3' % tuple(tr['rho_H'])) if h > tr['h_percentiles']['99']: flags.append('estimate above the 99th percentile of the training labels') flags.append('the estimators failed the test of predicting above their training range (largest training value %.0f eV A); a compound stronger than that would be underestimated' % tr['h'][1]) if abs(lt - lg) > math.log(1.5): flags.append('the two estimators differ by more than a factor of 1.5') else: flags.append('no hydrogen: only the whole Hopfield sum is estimated') if len(r['species']) > meta['largest_training_cell_atoms']: flags.append('cell larger than any in the training set (%d atoms)' % meta['largest_training_cell_atoms']) seen = m['index'].get(st.composition.reduced_formula + '|' + stype) if seen: out['in_training_set'] = dict(label_h_eV_A=seen[0], label_hopfield_sum_eV_per_A2=seen[1]) flags.insert(0, 'a compound of this composition and structure type is in the training set, so this is not an out-of-sample estimate; its label is given') flags.append('labels are harmonic calculations at one atmosphere and a smearing of 0.030 Ry; persistence of the structure is not estimated') out['flags'] = flags return out def main(): for path in sys.argv[1:]: try: print(json.dumps(predict(Structure.from_file(path)), indent=1)) except Exception as err: print(json.dumps(dict(file=path, error=str(err)))) if __name__ == '__main__': main()