"""Trains the estimators that are released, on every compound of the release with a label. The tests in model/results/ score estimators trained with part of the data held out. The released ones use the same code and settings with nothing held out. Output, in model/weights/: trees.joblib gradient-boosted trees for ln h of a hydride graph_.pt one member of the graph ensemble (state dictionary) meta.json training ranges, elements and structure types seen, settings, the calibrated range Run: .venv/bin/python model/train_final.py trees .venv/bin/python model/train_final.py graph [epochs] .venv/bin/python model/train_final.py meta """ import json import math import pathlib import sys import time import joblib import numpy as np import torch sys.path.insert(0, str(pathlib.Path(__file__).resolve().parent)) import train as T OUT = T.ROOT / 'model' / 'weights' def main(): what = sys.argv[1] OUT.mkdir(exist_ok=True) rows = T.load() hyd = [r for r in rows if r['has_h']] if what == 'trees': X, keys = T.feature_matrix(hyd) _, model = T.fit_trees(hyd, hyd[:1], 'h') joblib.dump(model, OUT / 'trees.joblib') json.dump(keys, open(OUT / 'feature_order.json', 'w')) print('trees trained on', len(hyd), 'hydrides,', len(keys), 'features') elif what == 'graph': seed = int(sys.argv[2]) epochs = int(sys.argv[3]) if len(sys.argv) > 3 else 60 t = time.time() net = T.fit_graph(rows, epochs=epochs, seed=seed, log=print) torch.save({k: v.cpu() for k, v in net.state_dict().items()}, OUT / f'graph_{seed}.pt') print(f'graph seed {seed} trained on {len(rows)} compounds in {time.time() - t:.0f} s') elif what == 'meta': h = np.array([r['row']['h'] for r in hyd]) rho = np.array([r['row']['rho_H'] for r in hyd]) elements = sorted({s for r in rows for s in r['species']}) hyd_elements = sorted({s for r in hyd for s in r['species']}) seeds = sorted(int(p.stem.split('_')[1]) for p in OUT.glob('graph_*.pt')) meta = dict( name='roomtsc', version='0.1', trained_on=dict(source='Alexandria phonon and electron-phonon release of 2025-08-11 (PBEsol, 3D, 1 atm), CC BY 4.0', compounds_with_spectrum=len(rows), hydrides_with_h=len(hyd), smearing_Ry=0.030), label='h = eta_H / rho_H in eV A, where eta_H is the part of the Hopfield sum in the 3 n_H highest phonon branches; S is the whole Hopfield sum in eV/A^2 with the hydrogen mass as the unit', graph_seeds=seeds, graph=dict(dim=96, layers=4, n_rbf=32, cutoff_A=T.Net().centers[-1].item(), epochs=60), training_range=dict(h=[float(h.min()), float(h.max())], h_percentiles={str(q): float(np.percentile(h, q)) for q in (1, 5, 50, 95, 99)}, rho_H=[float(rho.min()), float(rho.max())]), elements_seen=elements, elements_seen_in_labelled_hydrides=hyd_elements, prototypes_seen_in_labelled_hydrides=sorted({r['row']['prototype'] for r in hyd}), electron_gas_h=T.GAS, asymptote_K_per_sqrt_eV_A2=136.54, efficiency_range=[0.35, 0.58], requirement_eta_H=dict(single_mode=[8.7, 12.0], at_ambient_efficiencies=[18, 39], floor=4.83)) # 90% error factor of the mean of the two estimators on the hardest held-out test # (the corrected run if it exists, else the first run) top = next((p for p in (T.ROOT / 'model' / 'results2' / 'top.json', T.ROOT / 'model' / 'results' / 'top.json') if p.exists()), None) if top: r = json.load(open(top)) c = r['compounds_held_out'] est = 0.5 * (np.log([x['trees'] for x in c]) + np.log([x['graph'] for x in c])) y = np.log([x['h'] for x in c]) res = np.abs(est - y) q90 = float(np.percentile(res, 90)) above = y > math.log(r['training_maximum_h']) outside = int((y[above] > est[above] + q90).sum()) meta['calibration'] = dict(q90_ln=q90, q50_ln=float(np.percentile(res, 50)), n=len(c), source=str(top.relative_to(T.ROOT)), note=f'90% of the {len(c)} hydrides held out in the hardest test (structure types with the largest values) lay within this range of the estimate; of the {int(above.sum())} stronger than anything in training, {outside} lay above it') # compounds of the training set by composition and structure type, with their labels import csv st = {x['id']: x for x in csv.DictReader(open(T.ROOT / 'model' / 'data' / 'structure_types.csv'))} index = {} for r in rows: index.setdefault(st[r['row']['id']]['compound'], [round(r['row']['h'], 2) if r['has_h'] else None, round(r['row']['S'], 3)]) json.dump(index, open(OUT / 'training_index.json', 'w'), separators=(',', ':')) meta['structure_types_seen_in_labelled_hydrides'] = sorted({st[r['row']['id']]['structure_type'] for r in hyd}) meta['largest_training_cell_atoms'] = max(len(r['species']) for r in rows) meta['baselines_h'] = dict(training_geometric_mean=float(np.exp(np.log(h).mean())), electron_gas=T.GAS) del meta['prototypes_seen_in_labelled_hydrides'] json.dump(meta, open(OUT / 'meta.json', 'w'), indent=1) print('meta written:', len(elements), 'elements,', len(meta['structure_types_seen_in_labelled_hydrides']), 'hydride structure types, graph seeds', seeds) else: raise SystemExit('trees | graph [epochs] | meta') if __name__ == '__main__': main()