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