pino-source-code / scripts /freeze_genre_evaluation.py
Matthew Ford
feat: v11.6 repair/enrichment scripts + independent-candidate prep + integrity tests
1bb570d
Raw
History Blame Contribute Delete
10.9 kB
#!/usr/bin/env python3
"""Freeze molecule-disjoint genre splits and export PIMT representations.
The prepare phase is label-aware only for the preregistered coverage check. The
export phase never runs the benchmark and writes embeddings in manifest order.
"""
from __future__ import annotations
import argparse
from collections import Counter
import hashlib
import json
import random
import subprocess
from pathlib import Path
from typing import Any
import numpy as np
from pino.genre_benchmark import SOLVENTS, active_compounds, formula_fingerprint
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def genre(row: dict[str, Any]) -> str:
return str(row.get("genre") or row.get("metadata", {}).get("generation_strategy") or "wildcard")
def lean_row(row: dict[str, Any]) -> dict[str, Any]:
"""Retain exactly the fields consumed by the benchmark and row alignment."""
return {
"source_index": row["source_index"],
"formula_id": row.get("formula_id") or row.get("metadata", {}).get("formula_id"),
"genre": genre(row),
"formula": row.get("formula", []),
}
def load_index(dataset: Path) -> tuple[list[dict[str, Any]], list[str]]:
rows: list[dict[str, Any]] = []
molecules: set[str] = set()
with dataset.open(encoding="utf-8") as handle:
for index, line in enumerate(handle):
row = json.loads(line)
row["source_index"] = index
compounds = active_compounds(row)
rows.append({
"source_index": index,
"genre": genre(row),
"compounds": compounds,
"fingerprint": formula_fingerprint(row),
})
molecules.update(compounds)
return rows, sorted(molecules)
def split_indices(rows: list[dict[str, Any]], molecules: list[str], seed: int, train_ratio: float):
shuffled = molecules.copy()
random.Random(seed).shuffle(shuffled)
cut = min(max(int(len(shuffled) * train_ratio), 1), len(shuffled) - 1)
train_molecules = set(shuffled[:cut])
test_molecules = set(shuffled[cut:])
train, test, excluded = [], [], []
for row in rows:
compounds = row["compounds"]
if not compounds:
excluded.append(row["source_index"])
elif compounds <= train_molecules:
train.append(row["source_index"])
elif compounds <= test_molecules:
test.append(row["source_index"])
else:
excluded.append(row["source_index"])
return train, test, excluded
def counts(rows_by_index: dict[int, dict[str, Any]], indices: list[int]) -> dict[str, int]:
return dict(sorted(Counter(rows_by_index[i]["genre"] for i in indices).items()))
def prepare(args: argparse.Namespace) -> None:
dataset, checkpoint, output = Path(args.dataset), Path(args.checkpoint), Path(args.output_dir)
output.mkdir(parents=True, exist_ok=True)
rows, molecules = load_index(dataset)
by_index = {row["source_index"]: row for row in rows}
split_specs = []
wanted: dict[int, tuple[str, str]] = {}
for seed in args.seeds:
train, test, excluded = split_indices(rows, molecules, seed, args.train_ratio)
train_counts, test_counts = counts(by_index, train), counts(by_index, test)
genres = sorted(set(train_counts) | set(test_counts))
if any(train_counts.get(g, 0) < args.minimum_per_genre or test_counts.get(g, 0) < args.minimum_per_genre for g in genres):
raise SystemExit(f"seed {seed} fails minimum-per-genre={args.minimum_per_genre}: train={train_counts}, test={test_counts}")
name = f"seed_{seed}"
train_path, test_path = output / f"{name}_train.jsonl", output / f"{name}_test.jsonl"
for index in train:
wanted[index] = wanted.get(index, ("", ""))
split_specs.append((name, seed, train, test, excluded, train_path, test_path, train_counts, test_counts))
# Stream the large source once and write compact benchmark records.
handles = {}
memberships: dict[int, list[tuple[Path, str]]] = {}
for _name, _seed, train, test, _excluded, train_path, test_path, _tc, _vc in split_specs:
handles[train_path] = train_path.open("w", encoding="utf-8")
handles[test_path] = test_path.open("w", encoding="utf-8")
for i in train:
memberships.setdefault(i, []).append((train_path, "train"))
for i in test:
memberships.setdefault(i, []).append((test_path, "test"))
with dataset.open(encoding="utf-8") as source:
for index, line in enumerate(source):
if index not in memberships:
continue
row = json.loads(line)
row["source_index"] = index
encoded = json.dumps(lean_row(row), sort_keys=True, separators=(",", ":")) + "\n"
for path, _partition in memberships[index]:
handles[path].write(encoded)
for handle in handles.values():
handle.close()
manifest, audits = [], []
for name, seed, train, test, excluded, train_path, test_path, train_counts, test_counts in split_specs:
manifest.append({
"name": name,
"seed": seed,
"train_records": str(train_path),
"test_records": str(test_path),
"learned_train": str(output / f"{name}_train_embeddings.npy"),
"learned_test": str(output / f"{name}_test_embeddings.npy"),
})
audits.append({
"name": name, "seed": seed, "train_records": len(train), "test_records": len(test),
"excluded_records": len(excluded), "train_genres": train_counts, "test_genres": test_counts,
"minimum_per_genre": args.minimum_per_genre, "coverage_passed": True,
})
manifest_path = output / "genre_split_manifest.json"
manifest_path.write_text(json.dumps(manifest, indent=2) + "\n")
commit = subprocess.check_output(["git", "rev-parse", "HEAD"], text=True).strip()
protocol = {
"protocol_version": 1,
"status": "frozen_before_embedding_export_and_benchmark",
"dataset": {"path": str(dataset), "sha256": sha256(dataset), "records": len(rows)},
"checkpoint": {"path": str(checkpoint), "sha256": sha256(checkpoint)},
"git_commit": commit,
"split": {"method": "active-molecule assignment; mixed-boundary formulas excluded", "train_ratio": args.train_ratio,
"seeds": args.seeds, "minimum_records_per_genre_per_partition": args.minimum_per_genre},
"benchmark": {"probe": "training-standardized nearest centroid", "bootstrap_samples": 10000,
"bootstrap_seed": 8675309, "minimum_valid_splits": 3, "required_margin": 0.0,
"decision": "support iff every valid split has paired CI95 lower bound > required_margin"},
"manifest": str(manifest_path),
"coverage_audit": audits,
}
(output / "frozen_protocol.json").write_text(json.dumps(protocol, indent=2) + "\n")
print(json.dumps(protocol, indent=2))
def pooled_embedding(model, item, torch):
tokens = item["tokens"].unsqueeze(0)
states = item["physics"].unsqueeze(0)
with torch.inference_mode():
latent = model(tokens, states)
return latent.mean(dim=(1, 2)).squeeze(0).cpu().numpy().astype(np.float32)
def export(args: argparse.Namespace) -> None:
import torch
from pino.pimt_model import FragranceTrajectoryDataset, PhysicsInformedMixtureTransformer
protocol_path = Path(args.protocol)
protocol = json.loads(protocol_path.read_text())
dataset, checkpoint = Path(protocol["dataset"]["path"]), Path(protocol["checkpoint"]["path"])
if sha256(dataset) != protocol["dataset"]["sha256"] or sha256(checkpoint) != protocol["checkpoint"]["sha256"]:
raise SystemExit("frozen dataset or checkpoint hash mismatch")
manifest = json.loads(Path(protocol["manifest"]).read_text())
needed = set()
for split in manifest:
for key in ("train_records", "test_records"):
with Path(split[key]).open() as handle:
needed.update(json.loads(line)["source_index"] for line in handle if line.strip())
state = torch.load(checkpoint, map_location="cpu", weights_only=False)["model_state_dict"]
hidden = state["input_proj.weight"].shape[0]
layers = len({key.split(".")[2] for key in state if key.startswith("encoder.layers.")})
model = PhysicsInformedMixtureTransformer(embedding_dim=151, state_dim=2, hidden_dim=hidden, num_heads=4, num_layers=layers)
model.load_state_dict(state)
model.eval()
embeddings: dict[int, np.ndarray] = {}
with dataset.open(encoding="utf-8") as source:
for index, line in enumerate(source):
if index not in needed:
continue
row = json.loads(line)
ds = FragranceTrajectoryDataset(data_path=None, records=[row], state_dim=2, use_embedding_fallback=True, label_noise=False)
embeddings[index] = pooled_embedding(model, ds[0], torch)
if len(embeddings) % 100 == 0:
print(f"exported {len(embeddings)}/{len(needed)} unique rows", flush=True)
if embeddings.keys() != needed:
raise SystemExit(f"missing {len(needed - embeddings.keys())} source rows")
for split in manifest:
for partition in ("train", "test"):
record_path = Path(split[f"{partition}_records"])
indices = [json.loads(line)["source_index"] for line in record_path.open() if line.strip()]
np.save(split[f"learned_{partition}"], np.stack([embeddings[i] for i in indices]))
protocol["status"] = "embeddings_exported_benchmark_not_run"
protocol["embedding_export"] = {"pooling": "unmasked mean over time and active ingredient tokens", "dimension": hidden,
"unique_records": len(embeddings)}
protocol_path.write_text(json.dumps(protocol, indent=2) + "\n")
def main() -> None:
parser = argparse.ArgumentParser()
sub = parser.add_subparsers(dest="command", required=True)
prep = sub.add_parser("prepare")
prep.add_argument("--dataset", required=True); prep.add_argument("--checkpoint", required=True); prep.add_argument("--output-dir", required=True)
prep.add_argument("--seeds", type=int, nargs="+", default=[1729, 2718, 3141]); prep.add_argument("--train-ratio", type=float, default=.85)
prep.add_argument("--minimum-per-genre", type=int, default=5); prep.set_defaults(func=prepare)
exp = sub.add_parser("export"); exp.add_argument("--protocol", required=True); exp.set_defaults(func=export)
args = parser.parse_args(); args.func(args)
if __name__ == "__main__":
main()