pino-source-code / scripts /evaluate_frozen_benchmarks.py
mattbitzesty's picture
round-10: physics-factorial + direct-2048 chain (scripts/evaluate_frozen_benchmarks.py)
63f6c7c verified
Raw
History Blame Contribute Delete
11.9 kB
#!/usr/bin/env python3
"""Evaluate a trained PIMT checkpoint ONCE on the frozen benchmarks.
Benchmarks (all label_access=evaluation_only, never used in training/selection):
1. Discordant substitution triplets: the model should rank the perceptually
preferred substitute CLOSER to the target than the structurally-closer
distractor. Scored in the model's learned odor space.
2. Prospective formulas: predicted family profile vs sequestered intended
profile (cosine similarity per formula).
Reports every outcome including negative/inconclusive. This is the single,
preregistered readout for the two-arm representation A/B.
"""
from __future__ import annotations
import argparse
import csv
import hashlib
import json
from pathlib import Path
from typing import Any
import numpy as np
import torch
from scipy.spatial.distance import cosine
from pino.embeddings import OlfactoryEmbeddingEngine
from pino.heads import PIMTHeads
from pino.pimt_model import PhysicsInformedMixtureTransformer
ROOT = Path(__file__).resolve().parents[1]
def _canon_smiles(smiles: str) -> str | None:
try:
from rdkit import Chem
m = Chem.MolFromSmiles(smiles or "")
return Chem.MolToSmiles(m, isomericSmiles=False) if m else None
except Exception: # noqa: BLE001
return None
def _load_openpom_proxy() -> set[str]:
"""Public GoodScents+Leffingwell corpus OpenPOM was trained on (proxy bound)."""
f = ROOT / "data/openpom_curated_4983.csv"
keys: set[str] = set()
if not f.exists():
return keys
with open(f, newline="") as fh:
for row in csv.DictReader(fh):
k = _canon_smiles(row.get("nonStereoSMILES", ""))
if k:
keys.add(k)
return keys
def sha256_file(path: Path) -> str:
return hashlib.sha256(path.read_bytes()).hexdigest()
def infer_model_dims(state_dict: dict[str, torch.Tensor]) -> dict[str, int]:
hidden_dim, embedding_dim = state_dict["input_proj.weight"].shape
state_dim = state_dict["gating.physics_scaler"].shape[0]
num_layers = max(
int(k.split(".")[2]) for k in state_dict
if k.startswith("encoder.layers.") and k.endswith(".self_attn.in_proj_weight")
) + 1
num_heads = 8 if hidden_dim % 8 == 0 and hidden_dim >= 512 else 4
return {"embedding_dim": int(embedding_dim), "state_dim": int(state_dim),
"hidden_dim": int(hidden_dim), "num_heads": num_heads, "num_layers": int(num_layers)}
def load_model(ckpt_path: Path, objective_dim: int = 138):
ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
dims = infer_model_dims(ckpt["model_state_dict"])
model = PhysicsInformedMixtureTransformer(**dims)
model.load_state_dict(ckpt["model_state_dict"])
model.eval()
heads = PIMTHeads(hidden_dim=dims["hidden_dim"], objective_dim=objective_dim)
heads.load_state_dict(ckpt["heads_state_dict"], strict=False)
heads.eval()
return model, heads, dims
def odor_pyramid(model, heads, engine, smiles: str, cas: str) -> np.ndarray:
"""Encode one molecule to the model's 138-dim odor pyramid (mean of tiers).
This is the learned perceptual space used for triplet ranking."""
z = torch.from_numpy(engine.get_embedding(smiles, cas=cas)).float()
tokens = z.unsqueeze(0).unsqueeze(0) # (1, T=1, S=1, E)
physics = torch.zeros(1, 1, 1, 2) # (1, T=1, S=1, 2)
with torch.no_grad():
latent = model(tokens, physics) # (1, T, S, H)
out = heads(latent, physics)
pyramid = out["objective"].squeeze(0).numpy() # (3, 138)
return pyramid.mean(axis=0) # (138,)
def eval_triplets(model, heads, engine) -> dict[str, Any]:
f = ROOT / "data/benchmarks/substitution_triplets/triplets.jsonl"
rows = [json.loads(l) for l in f.read_text().splitlines() if l.strip()]
correct, details = 0, []
for r in rows:
# Triplets store CAS/SMILES: keys; the engine resolves structure via CAS.
t_emb = odor_pyramid(model, heads, engine, "", r["target_cas"])
p_emb = odor_pyramid(model, heads, engine, "", r["preferred_cas"])
d_emb = odor_pyramid(model, heads, engine, "", r["structural_distractor_cas"])
d_pref = cosine(t_emb, p_emb)
d_dist = cosine(t_emb, d_emb)
ok = bool(d_pref < d_dist) # preferred should be perceptually closer
correct += int(ok)
details.append({"triplet_id": r["triplet_id"], "target": r["target"],
"cos_dist_preferred": round(float(d_pref), 4),
"cos_dist_distractor": round(float(d_dist), 4), "correct": ok})
n = len(rows)
return {"n_triplets": n, "n_correct": correct,
"accuracy": round(correct / n, 4) if n else None,
"chance": 0.5, "details": details}
# Family -> Pyrfume 138 tags that belong to that olfactory family. Used to project
# the model's predicted pyramid onto a family profile.
FAMILY_TAGS = {
"citrus_fresh": ["citrus", "lemon", "orange", "grapefruit", "fresh", "ozone", "terpenic", "bergamot"],
"aromatic_herbal": ["herbal", "lavender", "chamomile", "camphoreous", "mentholic", "mint", "minty", "aromatic"],
"floral_sweet": ["floral", "rose", "jasmin", "jasmine", "muguet", "violet", "hyacinth", "lily", "geranium", "sweet"],
"woody_amber": ["woody", "cedar", "pine", "amber", "sandalwood", "vetiver", "patchouli", "mossy"],
"gourmand": ["vanilla", "chocolate", "caramellic", "honey", "cocoa", "coffee", "nutty", "coumarinic", "lactonic", "creamy", "balsamic"],
"musk_clean": ["musk", "clean", "powdery", "soapy", "aldehydic", "animal"],
"green": ["green", "grassy", "leafy", "cucumber", "vegetable", "hay", "weedy"],
}
FAMS = list(FAMILY_TAGS)
def _family_projection(vocab: list[str], pyramid: np.ndarray) -> np.ndarray:
"""Project a predicted 138-dim pyramid onto family mass via tag membership."""
idx = {t: i for i, t in enumerate(vocab)}
prof = np.zeros(len(FAMS))
for fi, fam in enumerate(FAMS):
for tag in FAMILY_TAGS[fam]:
if tag in idx:
prof[fi] += pyramid[idx[tag]]
return prof
def eval_prospective(model, heads, engine, openpom_proxy: set[str] | None = None) -> dict[str, Any]:
vocab = json.loads((ROOT / "data/pyrfume_vocabulary.json").read_text())["vocabulary"]
ff = ROOT / "data/benchmarks/prospective_formulas/formulas.jsonl"
lf = ROOT / "data/benchmarks/prospective_formulas/labels.sequestered.json"
formulas = [json.loads(l) for l in ff.read_text().splitlines() if l.strip()]
labels = {l["formula_id"]: l["target_family_profile"]
for l in json.loads(lf.read_text())["labels"]}
proxy = openpom_proxy if openpom_proxy is not None else set()
sims, weights, details = [], [], []
for form in formulas:
# Model prediction: weight-average each ingredient's predicted pyramid by
# the model's concentration-normalised contribution (weight_fraction here
# as the physical dose), then project onto family tag-space.
blend = np.zeros(138)
ing_keys = []
for ing in form["ingredients"]:
emb = odor_pyramid(model, heads, engine, ing.get("smiles", ""), ing.get("cas") or "")
blend += ing["weight_fraction"] * emb
ing_keys.append(_canon_smiles(ing.get("smiles", "")))
pred = _family_projection(vocab, blend)
pred = pred / pred.sum() if pred.sum() else pred
tgt = np.array([labels[form["formula_id"]].get(f, 0.0) for f in FAMS])
tgt = tgt / tgt.sum() if tgt.sum() else tgt
sim = 1.0 - cosine(pred, tgt) if (pred.any() and tgt.any()) else 0.0
# Novelty weight: fraction of ingredients absent from OpenPOM's public training
# corpus. On this benchmark overlap is near-total, so this bounds the honest,
# overlap-free agreement. 1.0 when no proxy is available (unweighted).
valid = [k for k in ing_keys if k]
novelty = (sum(1 for k in valid if k not in proxy) / len(valid)) if valid and proxy else 1.0
sims.append(sim)
weights.append(novelty)
details.append({**{"formula_id": form["formula_id"], "genre": form["genre"],
"family_profile_cosine": round(float(sim), 4),
"novelty_weight": round(float(novelty), 4)}})
w = np.array(weights)
s = np.array(sims)
wmean = float((w * s).sum() / w.sum()) if w.sum() > 0 else None
result = {"n_formulas": len(formulas),
"mean_family_profile_cosine": round(float(np.mean(sims)), 4) if sims else None,
"note": "model-predicted pyramid projected onto family tag-space vs sequestered intended profile",
"details": details}
if proxy:
result["openpom_overlap"] = {
"novelty_weighted_mean_cosine": round(wmean, 4) if wmean is not None else None,
"n_formulas_with_any_novelty": int((w > 0).sum()),
"mean_novelty_weight": round(float(w.mean()), 4),
"note": ("novelty_weight = fraction of ingredients NOT in OpenPOM's public training corpus; "
"the novelty-weighted cosine is the honest overlap-bounded readout (unweighted mean is "
"an upper bound inflated by pretraining memorization)."),
}
return result
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--checkpoint", required=True)
ap.add_argument("--structural-source", choices=["morgan", "morgan_2048_rp", "morgan_2048_direct", "openpom_256", "pom_alltags", "disjoint_256"], required=True)
ap.add_argument("--arm-label", default=None)
ap.add_argument("--output", required=True)
args = ap.parse_args()
ckpt = Path(args.checkpoint)
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(ckpt, objective_dim=objective_dim)
engine = OlfactoryEmbeddingEngine(structural_source=engine_source)
openpom_proxy = _load_openpom_proxy()
result = {
"arm": args.arm_label or args.structural_source,
"structural_source": args.structural_source,
"objective_dim": objective_dim,
"checkpoint": {"path": str(ckpt), "sha256": sha256_file(ckpt)},
"model_dims": dims,
"input_embedding_dim": engine.embedding_dim,
"openpom_overlap_proxy_corpus_size": len(openpom_proxy),
"substitution_triplets": eval_triplets(model, heads, engine),
# Prospective family-profile projection is defined over the 138-dim
# Pyrfume vocabulary; not comparable for the 575-dim all-tags arm.
"prospective_formulas": (eval_prospective(model, heads, engine, openpom_proxy)
if objective_dim == 138
else {"not_applicable": "575-dim tag target; family projection is 138-vocab specific"}),
"label_access": "evaluation_only; results reported for all outcomes incl. negative",
}
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", "input_embedding_dim"]}, indent=2))
print("triplets:", result["substitution_triplets"]["accuracy"],
f"({result['substitution_triplets']['n_correct']}/{result['substitution_triplets']['n_triplets']})")
if objective_dim == 138:
print("prospective mean cosine:", result["prospective_formulas"]["mean_family_profile_cosine"])
else:
print("prospective: not applicable (575-dim tag target)")
return 0
if __name__ == "__main__":
raise SystemExit(main())