pino-source-code / src /pino /dataset_builder.py
mattbitzesty's picture
feat(data): generate OAV targets in dataset builder and validate UNIFAC subgroups
884ed70 unverified
Raw
History Blame Contribute Delete
10.8 kB
from __future__ import annotations
import argparse
import json
import logging
import signal
import threading
import uuid
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
from pathlib import Path
from typing import Any
import numpy as np
from pino.draft_engine import FormulaGenerator
from pino.verifier import FragrancePipelineVerifier
logger = logging.getLogger("pino.dataset_builder")
GENRES = ["citrus_cologne", "fougere", "floral_woody", "amber_oriental", "wildcard"]
# Worker-thread local storage for generator and verifier reuse.
_worker_local = threading.local()
def _get_worker_generator(
genre: str,
rules_path: Path,
registry_path: Path,
literature_path: Path,
base_seed: int,
) -> FormulaGenerator:
"""Return a thread-local FormulaGenerator for this genre."""
current_genre = getattr(_worker_local, "genre", None)
if current_genre != genre or not hasattr(_worker_local, "generator"):
thread_id = threading.current_thread().ident or 0
_worker_local.generator = FormulaGenerator(
genre=genre,
rules_path=rules_path,
registry_path=registry_path,
literature_path=literature_path,
seed=base_seed + hash(genre) + thread_id,
min_k=5,
max_k=20,
)
_worker_local.genre = genre
return _worker_local.generator
def _get_worker_verifier() -> FragrancePipelineVerifier:
"""Return a thread-local FragrancePipelineVerifier."""
if not hasattr(_worker_local, "verifier"):
_worker_local.verifier = FragrancePipelineVerifier()
return _worker_local.verifier
class DatasetBuilder:
"""
Streaming, multi-threaded dataset builder for PINO synthetic trajectories.
Runs one stratified sweep per genre, writes verified records append-only to
a JSON-Lines file, and emits a final analytics card.
"""
def __init__(
self,
output_path: str | Path,
rules_path: str | Path,
registry_path: str | Path,
literature_path: str | Path,
per_genre: int = 1000,
max_workers: int = 16,
seed: int = 2026,
) -> None:
self.output_path = Path(output_path)
self.output_path.parent.mkdir(parents=True, exist_ok=True)
self.rules_path = Path(rules_path)
self.registry_path = Path(registry_path)
self.literature_path = Path(literature_path)
self.per_genre = per_genre
self.max_workers = max_workers
self.seed = seed
self._shutdown = False
def _signal_handler(self, signum: int, frame: Any) -> None:
logger.warning("Shutdown signal received; finishing in-flight work ...")
self._shutdown = True
def _task(self, idx: int, genre: str) -> dict[str, Any] | None:
"""Generate and verify a single formula in a worker thread."""
generator = _get_worker_generator(
genre, self.rules_path, self.registry_path, self.literature_path, self.seed
)
formula, formula_id = generator.generate(idx=idx)
if not generator.light_ifra_check(formula)["passed"]:
return None
try:
verifier = _get_worker_verifier()
result = verifier.run_sim(formula, duration_seconds=28800, interval_seconds=600)
except Exception as exc:
logger.debug("Verifier rejected %s: %s", formula_id, exc)
return None
if result.get("status") not in ("passed", "depleted"):
return None
# Build OAV-weighted descriptor targets using the new semantics module.
from pino import semantics
composition = [
{"cas": c.get("cas"), "smiles": c.get("smiles")}
for c in formula
]
C_gas = np.array([step["C_gas_mg_m3"] for step in result["trajectory"]])
# Align concentration matrix by canonical CAS ordering.
cas_order = [c["cas"] for c in composition]
concentrations = np.zeros((C_gas.shape[0], len(cas_order)), dtype=np.float32)
for t_idx, step in enumerate(result["trajectory"]):
for c_idx, cas in enumerate(cas_order):
concentrations[t_idx, c_idx] = step["C_gas_mg_m3"].get(cas, 0.0)
objective_targets = semantics.compute_oav_targets(composition, concentrations)
psychometric_targets = semantics.get_psychometric_targets(genre)
return {
"status": result.get("status"),
"formula_id": formula_id,
"genre": genre,
"metadata": {
"generation_strategy": genre,
"active_components_count": len(formula) - 1, # exclude solvent
"formula_id": formula_id,
},
"formula": formula,
"depletion_rates": result.get("depletion_rates", {}),
"trajectory": result.get("trajectory", []),
"objective_targets": objective_targets.tolist(),
"psychometric_targets": psychometric_targets.tolist(),
}
def _append_record(self, record: dict[str, Any]) -> None:
with self.output_path.open("a") as f:
f.write(json.dumps(record) + "\n")
def _count_existing(self, genre: str) -> int:
if not self.output_path.exists():
return 0
count = 0
with self.output_path.open("r") as f:
for line in f:
try:
rec = json.loads(line)
if rec.get("metadata", {}).get("generation_strategy") == genre:
count += 1
except Exception:
continue
return count
def _run_genre(self, genre: str, global_start_idx: int) -> dict[str, Any]:
"""Generate and verify `per_genre` records for a single genre."""
existing = self._count_existing(genre)
target = self.per_genre - existing
if target <= 0:
logger.info("Genre %s already has %d records; skipping", genre, existing)
return {"genre": genre, "verified": existing, "drafted": 0, "rejected": 0}
logger.info("Starting genre %s | need %d more records", genre, target)
verified = existing
drafted = 0
rejected = 0
next_idx = global_start_idx
with ThreadPoolExecutor(max_workers=self.max_workers, thread_name_prefix=f"pino-{genre}") as executor:
futures: set[Any] = set()
while verified < self.per_genre and not self._shutdown:
# Keep the worker pool saturated.
while len(futures) < self.max_workers * 4 and not self._shutdown:
futures.add(executor.submit(self._task, next_idx, genre))
next_idx += 1
drafted += 1
if drafted >= target * 5:
# Safety valve: stop drafting if rejection rate is too high.
break
if not futures:
break
done, futures = wait(futures, return_when=FIRST_COMPLETED)
for future in done:
result = future.result()
if result is None:
rejected += 1
continue
self._append_record(result)
verified += 1
if verified % 100 == 0 and verified > existing:
logger.info(
"Genre %s: verified %d/%d | drafted %d | rejected %d",
genre,
verified,
self.per_genre,
drafted,
rejected,
)
if verified >= self.per_genre:
break
# Safety valve: if rejection rate is high, move on anyway.
if drafted > max(100, verified * 3) and verified < self.per_genre:
logger.warning(
"Genre %s has high rejection rate (verified %d, drafted %d); continuing",
genre, verified, drafted
)
logger.info(
"Genre %s complete: verified %d/%d | drafted %d | rejected %d",
genre,
verified,
self.per_genre,
drafted,
rejected,
)
return {"genre": genre, "verified": verified, "drafted": drafted, "rejected": rejected}
def build(self) -> dict[str, Any]:
"""Run the full stratified sweep and return the analytics card."""
signal.signal(signal.SIGINT, self._signal_handler)
signal.signal(signal.SIGTERM, self._signal_handler)
# Wipe only if empty; otherwise resume.
if not self.output_path.exists() or self.output_path.stat().st_size == 0:
self.output_path.write_text("")
card = {"genres": [], "total_verified": 0, "total_drafted": 0, "total_rejected": 0}
global_idx = 0
for genre in GENRES:
summary = self._run_genre(genre, global_idx)
card["genres"].append(summary)
card["total_verified"] += summary["verified"]
card["total_drafted"] += summary["drafted"]
card["total_rejected"] += summary["rejected"]
global_idx += self.per_genre
logger.info(
"Dataset complete: %d verified | %d drafted | %d rejected",
card["total_verified"],
card["total_drafted"],
card["total_rejected"],
)
return card
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Build PINO synthetic dataset v2")
parser.add_argument("--output", default="data/synthetic_dataset_v2.jsonl")
parser.add_argument("--rules", default="data/genre_rules.json")
parser.add_argument("--registry", default="src/pino/registry.db")
parser.add_argument("--literature", default="data/literature_formulas.json")
parser.add_argument("--per-genre", type=int, default=1000)
parser.add_argument("--workers", type=int, default=4)
parser.add_argument("--seed", type=int, default=2026)
parser.add_argument("--log-level", default="INFO")
args = parser.parse_args()
logging.basicConfig(
level=getattr(logging, args.log_level.upper()),
format="%(asctime)s %(levelname)s %(name)s: %(message)s",
)
builder = DatasetBuilder(
output_path=args.output,
rules_path=args.rules,
registry_path=args.registry,
literature_path=args.literature,
per_genre=args.per_genre,
max_workers=args.workers,
seed=args.seed,
)
card = builder.build()
print(json.dumps(card, indent=2))