Matthew Ford
feat: v11.6 repair/enrichment scripts + independent-candidate prep + integrity tests
1bb570d | #!/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() | |