"""Estimators of the Hopfield sum and of the scattering strength per proton, and the tests the paper fixed for them before any training. Two estimators are scored against the same baselines on the same splits. trees gradient-boosted trees on descriptors of the cell (composition, densities, shortest distances, space group) graph a message-passing network on the crystal graph. It returns a strength a_j for each atom, and the compound's values follow from the definitions in the paper: S = (1/V) * sum_j a_j * M_H / M_j h = mean of a_j over hydrogen atoms Baselines for h (eV A): gas 36.6 for every compound, the value for one proton in an electron gas dos a constant fitted on the training side times the density of states at the Fermi level per unit volume. The paper specifies the hydrogen-projected density of states; the release stores only the total, which is used here. const the geometric mean of h on the training side. Added to the paper's two in the deposit (model/deposit/manifest.json says why); the mark is against the best of the three. Run: .venv/bin/python model/train.py [options] (see the bottom of the file) """ import json import math import os import pathlib import pickle import sys import time import numpy as np import torch import torch.nn as nn from pymatgen.core import Element from sklearn.ensemble import HistGradientBoostingRegressor ROOT = pathlib.Path(__file__).resolve().parents[1] DATA = ROOT / 'data' / 'alexandria_ph' / 'dataset.pkl' OUT = ROOT / 'model' / 'results' M_H = 1.00794 GAS = 36.6 A2MH6 = '225_1a_2c_6e' # ---------------------------------------------------------------- data def load(): rows = pickle.load(open(DATA, 'rb')) rows = [r for r in rows if r['row']['S'] > 0] for r in rows: x = r['row'] r['has_h'] = bool(x['branch_ok'] and x['h'] and x['h'] > 0) return rows def feature_matrix(rows): keys = sorted(rows[0]['feats']) return np.array([[r['feats'][k] for k in keys] for r in rows], dtype=float), keys # ---------------------------------------------------------------- baselines and scoring def baselines(train, test): """Predicted ln h for the test hydrides from each baseline.""" y_tr = np.array([math.log(r['row']['h']) for r in train]) out = {'gas': np.full(len(test), math.log(GAS)), 'const': np.full(len(test), y_tr.mean())} dv = lambda r: (r['row']['dos_Ef'] or float('nan')) / r['row']['volume'] d_tr = np.array([dv(r) for r in train]) ok = np.isfinite(d_tr) & (d_tr > 0) c = (y_tr[ok] - np.log(d_tr[ok])).mean() d_te = np.array([dv(r) for r in test]) out['dos'] = np.where(np.isfinite(d_te) & (d_te > 0), c + np.log(np.where(d_te > 0, d_te, 1.0)), y_tr.mean()) return out def score(y, preds, groups, n_boot=2000, seed=0): """Mean absolute error in ln h for each predictor. For each estimator, its ratio to the best of the three baselines (the mark of the deposit) and to the better of the two the paper fixed, each with an interval from resampling held-out groups.""" err = {k: np.abs(np.asarray(p) - y) for k, p in preds.items()} res = {k: float(e.mean()) for k, e in err.items()} best = min(('gas', 'dos', 'const'), key=lambda k: res[k]) paper = min(('gas', 'dos'), key=lambda k: res[k]) rng = np.random.default_rng(seed) ug = np.unique(groups) idx = {g: np.flatnonzero(groups == g) for g in ug} draws = [np.concatenate([idx[g] for g in rng.choice(ug, len(ug))]) for _ in range(n_boot)] out = dict(n=int(len(y)), groups=int(len(ug)), mae=res, best_baseline=best, better_of_the_papers_two=paper) for name in [k for k in preds if k not in ('gas', 'dos', 'const')]: def ratio_to(base): r = [err[name][d].mean() / err[base][d].mean() for d in draws] lo, hi = np.percentile(r, [2.5, 97.5]) ratio = res[name] / res[base] return dict(baseline=base, ratio=float(ratio), interval=[float(lo), float(hi)], passes=bool(ratio <= 0.5 and hi < 1.0)) out[name] = dict(mae=res[name], against_best_baseline=ratio_to(best), against_papers_baselines=ratio_to(paper), within_30pct=float((err[name] < math.log(1.3)).mean()), spearman=float(spearman(preds[name], y))) return out def spearman(a, b): ra, rb = np.argsort(np.argsort(a)), np.argsort(np.argsort(b)) return np.corrcoef(ra, rb)[0, 1] if len(a) > 2 else float('nan') # ---------------------------------------------------------------- trees def fit_trees(train, test, target, seed=0): X, _ = feature_matrix(train + test) Xtr, Xte = X[:len(train)], X[len(train):] y = np.array([math.log(r['row'][target]) for r in train]) m = HistGradientBoostingRegressor(max_iter=600, learning_rate=0.04, max_leaf_nodes=24, min_samples_leaf=8, l2_regularization=1.0, early_stopping=True, validation_fraction=0.15, n_iter_no_change=40, random_state=seed) m.fit(Xtr, y) return m.predict(Xte), m # ---------------------------------------------------------------- graph network GROUPS, ROWS = 18, 9 _EL = {} def element_vector(sym): if sym not in _EL: e = Element(sym) v = np.zeros(GROUPS + ROWS + 5, dtype=np.float32) v[e.group - 1] = 1 v[GROUPS + e.row - 1] = 1 x = e.X if (e.X is not None and e.X == e.X) else 1.0 rad = float(e.atomic_radius) if e.atomic_radius else 1.5 val = 3 if (e.is_lanthanoid or e.is_actinoid) else (e.group if e.group <= 12 else e.group - 10) v[-5:] = [x / 4, rad / 2.5, math.log(float(e.atomic_mass)) / 5.5, val / 12, 1.0 if sym == 'H' else 0.0] _EL[sym] = (v, e.Z, float(e.atomic_mass)) return _EL[sym] def tensors(r): """One compound as arrays: element vectors, atomic numbers, masses, edges and labels.""" if 'g' in r: return r['g'] ev = [element_vector(s) for s in r['species']] i, j, d = r['edges'] x = r['row'] r['g'] = dict(x=np.stack([e[0] for e in ev]), z=np.array([e[1] for e in ev]), inv_m=np.array([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=x['volume'], ln_S=math.log(x['S']), ln_h=math.log(x['h']) if r['has_h'] else float('nan')) return r['g'] def collate(batch, device): xs, zs, ms, hs, ii, jj, dd, gid, vol, ln_s, ln_h = [], [], [], [], [], [], [], [], [], [], [] off = 0 for k, r in enumerate(batch): g = tensors(r) n = len(g['z']) xs.append(g['x']); zs.append(g['z']); ms.append(g['inv_m']); hs.append(g['is_h']) ii.append(g['i'] + off); jj.append(g['j'] + off); dd.append(g['d']) gid.append(np.full(n, k)); vol.append(g['vol']); ln_s.append(g['ln_S']); ln_h.append(g['ln_h']) off += n t = lambda a, dt=torch.float32: torch.as_tensor(np.concatenate(a) if isinstance(a[0], np.ndarray) else np.array(a), dtype=dt, device=device) return dict(x=t(xs), z=t(zs, torch.long), inv_m=t(ms), is_h=t(hs, torch.bool), i=t(ii, torch.long), j=t(jj, torch.long), d=t(dd), gid=t(gid, torch.long), vol=t(vol), ln_S=t(ln_s), ln_h=t(ln_h), n=len(batch)) class Conv(nn.Module): def __init__(self, dim, n_rbf): super().__init__() self.lin = nn.Linear(2 * dim + n_rbf, 2 * dim) self.norm = nn.LayerNorm(dim) def forward(self, h, i, j, e): z = self.lin(torch.cat([h[i], h[j], e], dim=1)) gate, core = z.chunk(2, dim=1) m = torch.zeros_like(h).index_add_(0, i, torch.sigmoid(gate) * nn.functional.softplus(core)) deg = torch.zeros(h.shape[0], 1, device=h.device).index_add_(0, i, torch.ones(len(i), 1, device=h.device)).clamp(min=1) return h + self.norm(m / deg.sqrt()) class Net(nn.Module): def __init__(self, dim=96, layers=4, n_rbf=32, cutoff=5.0): super().__init__() self.embed = nn.Embedding(100, dim) self.desc = nn.Linear(GROUPS + ROWS + 5, dim) self.register_buffer('centers', torch.linspace(0.5, cutoff, n_rbf)) self.gamma = (n_rbf / cutoff) ** 2 self.convs = nn.ModuleList(Conv(dim, n_rbf) for _ in range(layers)) self.head = nn.Sequential(nn.Linear(dim, dim), nn.SiLU(), nn.Linear(dim, 1)) def forward(self, b): h = self.embed(b['z']) + self.desc(b['x']) e = torch.exp(-self.gamma * (b['d'][:, None] - self.centers[None, :]) ** 2) for conv in self.convs: h = conv(h, b['i'], b['j'], e) a = nn.functional.softplus(self.head(h).squeeze(1)) + 1e-4 # strength per atom, eV A s = torch.zeros(b['n'], device=a.device).index_add_(0, b['gid'], a * b['inv_m']) / b['vol'] hsum = torch.zeros(b['n'], device=a.device).index_add_(0, b['gid'], a * b['is_h']) hcnt = torch.zeros(b['n'], device=a.device).index_add_(0, b['gid'], b['is_h'].float()) return torch.log(s), torch.log(hsum / hcnt.clamp(min=1) + 1e-8), a def fit_graph(train, epochs=120, seed=0, device=None, log=None, batch_size=96, lr=2e-3): device = device or os.environ.get('ROOMTSC_DEVICE') or ('mps' if torch.backends.mps.is_available() else 'cpu') torch.manual_seed(seed) rng = np.random.default_rng(seed) net = Net().to(device) opt = torch.optim.AdamW(net.parameters(), lr=lr, weight_decay=1e-4) steps = epochs * math.ceil(len(train) / batch_size) sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=lr, total_steps=steps, pct_start=0.1) huber = nn.HuberLoss(delta=1.0) for ep in range(epochs): net.train() order = rng.permutation(len(train)) tot = 0.0 for k in range(0, len(order), batch_size): b = collate([train[q] for q in order[k:k + batch_size]], device) ln_s, ln_h, _ = net(b) loss = huber(ln_s, b['ln_S']) m = torch.isfinite(b['ln_h']) if m.any(): loss = loss + huber(ln_h[m], b['ln_h'][m]) opt.zero_grad() loss.backward() nn.utils.clip_grad_norm_(net.parameters(), 5.0) opt.step() sched.step() tot += loss.item() * b['n'] if log and (ep % 10 == 9 or ep == epochs - 1): log(f' epoch {ep + 1}/{epochs} loss {tot / len(train):.4f}') return net @torch.no_grad() def predict_graph(net, rows, batch_size=256): device = next(net.parameters()).device net.eval() ls, lh = [], [] for k in range(0, len(rows), batch_size): b = collate(rows[k:k + batch_size], device) s, h, _ = net(b) ls.append(s.cpu().numpy()); lh.append(h.cpu().numpy()) return np.concatenate(ls), np.concatenate(lh) # ---------------------------------------------------------------- splits def split_top(rows): """Third test of the paper: prototypes ranked by their largest h; the leading ones are removed with all their members until a tenth of the hydrides with a label are held out.""" hyd = [r for r in rows if r['has_h']] best = {} for r in hyd: p = r['row']['prototype'] best[p] = max(best.get(p, 0.0), r['row']['h']) count = {} for r in hyd: count[r['row']['prototype']] = count.get(r['row']['prototype'], 0) + 1 held, n = set(), 0 for p in sorted(best, key=best.get, reverse=True): if n >= 0.1 * len(hyd): break held.add(p) n += count[p] return held def deposit(): """The deposited splits, checked against the label file they were written for.""" import csv, hashlib d = ROOT / 'model' / 'deposit' if not (d / 'manifest.json').exists(): raise SystemExit('no deposit: run model/deposit.py, commit it, then train') man = json.load(open(d / 'manifest.json')) labels = ROOT / 'model' / 'data' / 'labels.csv' if hashlib.sha256(open(labels, 'rb').read()).hexdigest() != man['checksums']['labels.csv']: raise SystemExit('labels.csv differs from the deposited one') return {r['id']: r for r in csv.DictReader(open(d / 'splits.csv'))}, man def run_split(rows, held_protos, name, log, epochs, seeds=(0,), with_graph=True): train_all = [r for r in rows if r['row']['prototype'] not in held_protos] test = [r for r in rows if r['row']['prototype'] in held_protos and r['has_h']] train_h = [r for r in train_all if r['has_h']] y = np.array([math.log(r['row']['h']) for r in test]) groups = np.array([r['row']['prototype'] for r in test]) log(f'{name}: train {len(train_all)} compounds ({len(train_h)} hydrides with h), test {len(test)} hydrides in {len(set(groups))} prototypes') preds = baselines(train_h, test) preds['trees'], _ = fit_trees(train_h, test, 'h') if with_graph: ph = [] for s in seeds: t = time.time() net = fit_graph(train_all, epochs=epochs, seed=s, log=log) ph.append(predict_graph(net, test)[1]) log(f' graph seed {s} trained in {time.time() - t:.0f} s') preds['graph'] = np.mean(ph, axis=0) return test, y, groups, preds, train_h # ---------------------------------------------------------------- experiments def main(): exp = sys.argv[1] epochs = int(sys.argv[2]) if len(sys.argv) > 2 else 120 OUT.mkdir(exist_ok=True) logf = open(OUT / f'{exp}.log', 'a') def log(s): print(s, flush=True) logf.write(s + '\n'); logf.flush() rows = load() dep, man = deposit() hyd = [r for r in rows if r['has_h']] log(f'== {exp}: {len(rows)} compounds with a spectrum, {len(hyd)} hydrides with h; deposit of {man["deposited_utc"]}') res = dict(experiment=exp, compounds=len(rows), hydrides_with_h=len(hyd), epochs=epochs, deposit=man['deposited_utc']) if exp == 'top': held = {r['row']['prototype'] for r in rows if dep[r['row']['id']]['top'] == 'test'} test, y, groups, preds, train_h = run_split(rows, held, 'hold out the top', log, epochs) res['score'] = score(y, preds, groups) ceiling = max(math.log(r['row']['h']) for r in train_h) above = y > ceiling res['training_maximum_h'] = math.exp(ceiling) res['held_out_above_training_maximum'] = int(above.sum()) for k in ('trees', 'graph'): if k in preds: res[k + '_extrapolation'] = dict( predicted_above_among_those_above=float((preds[k][above] > ceiling).mean()) if above.any() else None, predicted_above_among_the_rest=float((preds[k][~above] > ceiling).mean()) if (~above).any() else None) res['compounds_held_out'] = [dict(formula=r['row']['formula'], prototype=r['row']['prototype'], h=r['row']['h'], **{k: float(math.exp(p[i])) for k, p in preds.items()}) for i, r in enumerate(test)] elif exp == 'a2mh6': held = {A2MH6} test, y, groups, preds, _ = run_split(rows, held, 'hold out cubic A2MH6', log, epochs) groups = np.array([r['row']['id'] for r in test]) # resample compounds: one prototype res['score'] = score(y, preds, groups) res['compounds_held_out'] = [dict(formula=r['row']['formula'], h=r['row']['h'], **{k: float(math.exp(p[i])) for k, p in preds.items()}) for i, r in enumerate(test)] elif exp in ('grouped', 'random'): key = exp + '_fold' ys, gs, ps = [], [], {} for n in range(5): te_all = [r for r in rows if int(dep[r['row']['id']][key]) == n] tr_all = [r for r in rows if int(dep[r['row']['id']][key]) != n] te = [r for r in te_all if r['has_h']] tr_h = [r for r in tr_all if r['has_h']] log(f'fold {n + 1}: train {len(tr_all)} ({len(tr_h)} hydrides with h), test {len(te)}') p = baselines(tr_h, te) p['trees'], _ = fit_trees(tr_h, te, 'h') t = time.time() net = fit_graph(tr_all, epochs=epochs, seed=n, log=log) p['graph'] = predict_graph(net, te)[1] log(f' graph trained in {time.time() - t:.0f} s') ys.append([math.log(r['row']['h']) for r in te]) gs.append([r['row']['prototype'] for r in te]) for k, v in p.items(): ps.setdefault(k, []).append(v) y = np.concatenate(ys) allp = {k: np.concatenate(v) for k, v in ps.items()} res['score'] = score(y, allp, np.concatenate(gs)) res['predictions'] = dict(ln_h=[round(float(v), 4) for v in y], **{k: [round(float(v), 4) for v in p] for k, p in allp.items()}) else: raise SystemExit('experiments: top, a2mh6, grouped, random') log(json.dumps(res['score'], indent=1)) json.dump(res, open(OUT / f'{exp}.json', 'w'), indent=1) if __name__ == '__main__': main()