pino-source-code / scripts /evaluate_substantivity.py
mattbitzesty's picture
Upload scripts/evaluate_substantivity.py with huggingface_hub
65da34a verified
Raw
History Blame Contribute Delete
7.7 kB
"""Evaluate a trained PIMT checkpoint on MEASURED Poucher substantivity.
This is the E1 metric. Substantivity (olfactive duration) is physics-causal —
it IS evaporation — and the Poucher coefficients are measured, i.e.
engine-independent (not UNIFAC-derived). A non-redundant physics channel
(gamma blend-interaction) should improve prediction of this target on
held-out molecules; a redundant one should not.
Method:
* Rebuild the exact molecule-disjoint train/val split (same seed + ratio as
training) and run the substantivity head over the VALIDATION records.
* Keep only records with a measured-Poucher target (mask == 1).
* Report masked-MSE, MAE, Pearson and Spearman correlation between predicted
log10 substantivity and the measured Poucher target, plus n.
* Compare physics_off vs physics_gamma arms on the same split/seed.
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import numpy as np
import torch
from pino.heads import PIMTHeads
from pino.pimt_model import (
FragranceTrajectoryDataset,
PhysicsInformedMixtureTransformer,
)
from pino.upload_data import create_molecule_disjoint_split_from_records
def load_model(ckpt_path: Path, objective_dim: int):
sd = torch.load(ckpt_path, map_location="cpu", weights_only=False)
state = sd.get("model_state_dict", sd.get("model", sd))
heads_state = sd.get("heads_state_dict", sd.get("heads", {}))
state_dim = int(state["gating.physics_scaler"].shape[0])
embedding_dim = int(state["input_proj.weight"].shape[1])
hidden_dim = int(state["input_proj.weight"].shape[0])
num_layers = sum(1 for k in state if k.endswith("self_attn.in_proj_weight"))
# num_heads is not directly stored; default to the train config value (4 for
# full runs, 2 for smoke). It does not affect load_state_dict correctness
# because attention weights are stored shape-complete; only num_layers/dims
# must match. Use 4 but rely on strict=False for safety.
model = PhysicsInformedMixtureTransformer(
embedding_dim=embedding_dim, state_dim=state_dim,
hidden_dim=hidden_dim, num_heads=4, num_layers=max(num_layers, 1),
)
model.load_state_dict(state, strict=False)
heads = PIMTHeads(hidden_dim=hidden_dim, objective_dim=objective_dim)
if heads_state:
heads.load_state_dict(heads_state, strict=False)
model.eval(); heads.eval()
return model, heads, {"embedding_dim": embedding_dim, "state_dim": state_dim,
"hidden_dim": hidden_dim, "num_layers": num_layers}
def pearson(x, y):
if len(x) < 3:
return float("nan")
return float(np.corrcoef(x, y)[0, 1])
def spearman(x, y):
if len(x) < 3:
return float("nan")
rx = np.argsort(np.argsort(x)); ry = np.argsort(np.argsort(y))
return float(np.corrcoef(rx, ry)[0, 1])
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--checkpoint", required=True)
ap.add_argument("--data", required=True)
ap.add_argument("--structural-source", default="morgan")
ap.add_argument("--train-ratio", type=float, default=0.85)
ap.add_argument("--seed", type=int, default=42)
ap.add_argument("--arm-label", default=None)
ap.add_argument("--use-gamma", action="store_true")
ap.add_argument("--split", choices=["molecule", "formula"], default="formula",
help="molecule = strict molecule-disjoint holdout (clean but n~10 measured); "
"formula = random formula-level holdout (~460 measured, but molecules leak). "
"E1 uses formula-level for power and discloses the leakage.")
ap.add_argument("--output", required=True)
args = ap.parse_args()
objective_dim = 575 if args.structural_source == "pom_alltags" else 138
engine_source = "openpom_256" if args.structural_source == "pom_alltags" else args.structural_source
model, heads, dims = load_model(Path(args.checkpoint), objective_dim)
model_state_dim = dims["state_dim"]
with open(args.data) as f:
all_records = [json.loads(line) for line in f if line.strip()]
if args.split == "molecule":
split = create_molecule_disjoint_split_from_records(
all_records, train_ratio=args.train_ratio, seed=args.seed)
val_records = split["validation"]
split_note = "molecule-disjoint holdout (strict, low measured-n)"
else:
# formula-level random holdout: higher measured-Poucher n for power, at the
# cost of molecule leakage between train/val. Disclosed in the output.
# MUST match train.py's formula split exactly (same rng.permutation + seed).
rng = np.random.default_rng(args.seed)
n = len(all_records)
n_val = max(1, int(round(n * (1.0 - args.train_ratio))))
perm = rng.permutation(n)
val_idx = set(perm[:n_val].tolist())
val_records = [all_records[i] for i in sorted(val_idx)]
split_note = ("formula-level random holdout; molecules may leak between splits "
"(disclosed). Chosen for measured-Poucher power.")
ds = FragranceTrajectoryDataset(
records=val_records, use_embedding_fallback=True,
structural_source=engine_source, objective_dim=objective_dim,
use_gamma=args.use_gamma,
)
preds, targs = [], []
device = "cuda" if torch.cuda.is_available() else "cpu"
model.to(device); heads.to(device)
with torch.no_grad():
for i in range(len(ds)):
item = ds[i]
mask = float(item.get("target_substantivity_mask", 0.0))
if mask < 0.5:
continue
tokens = item["tokens"].unsqueeze(0).to(device) # (1,S,E)
physics = item["physics"].unsqueeze(0).to(device) # (1,T,S,C)
# match model state_dim
if physics.size(-1) != model_state_dim:
if physics.size(-1) < model_state_dim:
pad = torch.zeros(1, physics.size(1), physics.size(2), model_state_dim - physics.size(-1), device=device)
physics = torch.cat([physics, pad], dim=-1)
else:
physics = physics[..., :model_state_dim]
latent = model(tokens, physics)
out = heads(latent, physics)
sub = out["subjective"]["substantivity"]
pred = float(sub.reshape(-1)[0].cpu())
preds.append(pred)
targs.append(float(item["target_substantivity"]))
preds = np.array(preds); targs = np.array(targs)
result = {
"arm": args.arm_label or Path(args.checkpoint).stem,
"use_gamma": bool(args.use_gamma),
"seed": args.seed,
"model_dims": dims,
"task": "measured Poucher substantivity (engine-independent, physics-causal)",
"n_measured_val": int(len(preds)),
"mse": float(((preds - targs) ** 2).mean()) if len(preds) else None,
"mae": float(np.abs(preds - targs).mean()) if len(preds) else None,
"pearson": pearson(preds, targs),
"spearman": spearman(preds, targs),
"target_mean": float(targs.mean()) if len(targs) else None,
"target_std": float(targs.std()) if len(targs) else None,
"baseline_mse_predict_mean": float(((targs - targs.mean()) ** 2).mean()) if len(targs) else None,
}
out = Path(args.output); out.parent.mkdir(parents=True, exist_ok=True)
out.write_text(json.dumps(result, indent=2))
print(json.dumps({k: result[k] for k in ["arm", "use_gamma", "n_measured_val", "mse", "mae", "pearson", "spearman"]}, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())