#!/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())