260 lines
9.0 KiB
Python
260 lines
9.0 KiB
Python
|
|
"""Generate versioned trial metrics and paper source-data tables."""
|
||
|
|
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import argparse
|
||
|
|
import csv
|
||
|
|
import json
|
||
|
|
import os
|
||
|
|
import tempfile
|
||
|
|
from pathlib import Path
|
||
|
|
from typing import Any, Mapping, Sequence
|
||
|
|
|
||
|
|
import numpy as np
|
||
|
|
|
||
|
|
from experiments.hashing import file_sha256, stable_hash, to_jsonable
|
||
|
|
from experiments.io import atomic_write_json, atomic_write_jsonl, utc_now
|
||
|
|
from experiments.plan import load_document
|
||
|
|
from experiments.schema import SCHEMA_VERSION, STORAGE_FORMATS
|
||
|
|
from experiments.validate import require_valid_batch
|
||
|
|
|
||
|
|
from .metrics import derive_trial_metrics
|
||
|
|
|
||
|
|
|
||
|
|
def _method_selected(
|
||
|
|
family: str,
|
||
|
|
family_configuration: Mapping[str, Any],
|
||
|
|
method_id: str,
|
||
|
|
) -> bool:
|
||
|
|
"""Apply explicit per-family method filters without silent typos."""
|
||
|
|
include = family_configuration.get("include_methods")
|
||
|
|
exclude = family_configuration.get("exclude_methods")
|
||
|
|
for name, value in (
|
||
|
|
("include_methods", include),
|
||
|
|
("exclude_methods", exclude),
|
||
|
|
):
|
||
|
|
if value is not None and (
|
||
|
|
not isinstance(value, list)
|
||
|
|
or any(not isinstance(item, str) or not item for item in value)
|
||
|
|
):
|
||
|
|
raise ValueError(f"{family}.{name} must be a list of method IDs")
|
||
|
|
if include is not None and method_id not in include:
|
||
|
|
return False
|
||
|
|
if exclude is not None and method_id in exclude:
|
||
|
|
return False
|
||
|
|
return True
|
||
|
|
|
||
|
|
|
||
|
|
def _trial_metric_configuration(
|
||
|
|
configuration: Mapping[str, Any],
|
||
|
|
method_id: str,
|
||
|
|
) -> dict[str, Any] | None:
|
||
|
|
enabled = configuration.get("enabled")
|
||
|
|
if not isinstance(enabled, list) or not enabled:
|
||
|
|
raise ValueError("metric configuration needs a non-empty enabled list")
|
||
|
|
selected = []
|
||
|
|
trial_configuration: dict[str, Any] = {}
|
||
|
|
for family in enabled:
|
||
|
|
if not isinstance(family, str) or not family:
|
||
|
|
raise ValueError("enabled metric family names must be non-empty strings")
|
||
|
|
family_configuration = configuration.get(family, {})
|
||
|
|
if not isinstance(family_configuration, Mapping):
|
||
|
|
raise ValueError(f"{family} metric configuration must be a mapping")
|
||
|
|
if _method_selected(family, family_configuration, method_id):
|
||
|
|
selected.append(family)
|
||
|
|
trial_configuration[family] = dict(family_configuration)
|
||
|
|
if not selected:
|
||
|
|
return None
|
||
|
|
trial_configuration["enabled"] = selected
|
||
|
|
return trial_configuration
|
||
|
|
|
||
|
|
|
||
|
|
def _scalar_cell(value: Any) -> Any:
|
||
|
|
if value is None or isinstance(value, (str, int, float, bool)):
|
||
|
|
return value
|
||
|
|
return json.dumps(
|
||
|
|
to_jsonable(value),
|
||
|
|
sort_keys=True,
|
||
|
|
separators=(",", ":"),
|
||
|
|
ensure_ascii=False,
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def _identity_row(trial: Mapping[str, Any]) -> dict[str, Any]:
|
||
|
|
trajectory = trial["trajectory"]
|
||
|
|
method = trial["method"]
|
||
|
|
row = {
|
||
|
|
"trial_id": trial["trial_id"],
|
||
|
|
"pair_id": trial["pair_id"],
|
||
|
|
"study_id": trial["study_id"],
|
||
|
|
"split": trial["split"],
|
||
|
|
"method_id": method["method_id"],
|
||
|
|
"trajectory_id": trajectory["trajectory_id"],
|
||
|
|
"trajectory_family": trajectory.get("family", ""),
|
||
|
|
"replicate": trial["replicate"],
|
||
|
|
}
|
||
|
|
for name, value in sorted(trial.get("factors", {}).items()):
|
||
|
|
row[f"factor_{name}"] = _scalar_cell(value)
|
||
|
|
return row
|
||
|
|
|
||
|
|
|
||
|
|
def _atomic_write_csv(path: Path, rows: Sequence[Mapping[str, Any]]) -> None:
|
||
|
|
path.parent.mkdir(parents=True, exist_ok=True)
|
||
|
|
identity_order = [
|
||
|
|
"trial_id",
|
||
|
|
"pair_id",
|
||
|
|
"study_id",
|
||
|
|
"split",
|
||
|
|
"method_id",
|
||
|
|
"trajectory_id",
|
||
|
|
"trajectory_family",
|
||
|
|
"replicate",
|
||
|
|
]
|
||
|
|
all_fields = {key for row in rows for key in row}
|
||
|
|
fieldnames = [name for name in identity_order if name in all_fields]
|
||
|
|
fieldnames.extend(sorted(all_fields - set(fieldnames)))
|
||
|
|
fd, temporary_name = tempfile.mkstemp(
|
||
|
|
prefix=f".{path.name}.", suffix=".tmp", dir=path.parent
|
||
|
|
)
|
||
|
|
temporary = Path(temporary_name)
|
||
|
|
try:
|
||
|
|
with os.fdopen(fd, "w", encoding="utf-8", newline="") as stream:
|
||
|
|
writer = csv.DictWriter(stream, fieldnames=fieldnames)
|
||
|
|
writer.writeheader()
|
||
|
|
for row in rows:
|
||
|
|
writer.writerow(
|
||
|
|
{name: _scalar_cell(row.get(name)) for name in fieldnames}
|
||
|
|
)
|
||
|
|
stream.flush()
|
||
|
|
os.fsync(stream.fileno())
|
||
|
|
os.replace(temporary, path)
|
||
|
|
except Exception:
|
||
|
|
temporary.unlink(missing_ok=True)
|
||
|
|
raise
|
||
|
|
|
||
|
|
|
||
|
|
def generate_paper_source_data(
|
||
|
|
batch_dir: Path,
|
||
|
|
metric_configuration: Mapping[str, Any],
|
||
|
|
*,
|
||
|
|
output_root: Path | None = None,
|
||
|
|
) -> dict[str, Any]:
|
||
|
|
"""Recompute all configured metrics and emit traceable source-data tables."""
|
||
|
|
batch_dir = Path(batch_dir)
|
||
|
|
require_valid_batch(batch_dir, require_complete=True)
|
||
|
|
plan = load_document(batch_dir / "plan.json")
|
||
|
|
output_root = batch_dir if output_root is None else Path(output_root)
|
||
|
|
derived_dir = output_root / "derived"
|
||
|
|
source_dir = output_root / "paper" / "source_data"
|
||
|
|
derived_dir.mkdir(parents=True, exist_ok=True)
|
||
|
|
source_dir.mkdir(parents=True, exist_ok=True)
|
||
|
|
|
||
|
|
rows: list[dict[str, Any]] = []
|
||
|
|
input_files: dict[str, str] = {}
|
||
|
|
for trial in plan["trials"]:
|
||
|
|
trial_dir = batch_dir / "raw" / trial["trial_id"]
|
||
|
|
sample_path = trial_dir / "samples.npz"
|
||
|
|
input_files[str(sample_path.relative_to(batch_dir))] = file_sha256(sample_path)
|
||
|
|
with np.load(sample_path, allow_pickle=False) as archive:
|
||
|
|
samples = {name: archive[name] for name in archive.files}
|
||
|
|
method_id = trial["method"]["method_id"]
|
||
|
|
trial_configuration = _trial_metric_configuration(
|
||
|
|
metric_configuration,
|
||
|
|
method_id,
|
||
|
|
)
|
||
|
|
metrics = (
|
||
|
|
{}
|
||
|
|
if trial_configuration is None
|
||
|
|
else derive_trial_metrics(samples, trial_configuration)
|
||
|
|
)
|
||
|
|
rows.append({**_identity_row(trial), **metrics})
|
||
|
|
|
||
|
|
metric_path = derived_dir / "trial_metrics.jsonl"
|
||
|
|
atomic_write_jsonl(metric_path, rows)
|
||
|
|
table_files: dict[str, str] = {}
|
||
|
|
family_row_counts: dict[str, int] = {}
|
||
|
|
for family in metric_configuration["enabled"]:
|
||
|
|
family_rows: list[dict[str, Any]] = []
|
||
|
|
prefix = f"{family}_"
|
||
|
|
for row in rows:
|
||
|
|
if not any(key.startswith(prefix) for key in row):
|
||
|
|
continue
|
||
|
|
selected = {
|
||
|
|
key: value
|
||
|
|
for key, value in row.items()
|
||
|
|
if key in {
|
||
|
|
"trial_id",
|
||
|
|
"pair_id",
|
||
|
|
"study_id",
|
||
|
|
"split",
|
||
|
|
"method_id",
|
||
|
|
"trajectory_id",
|
||
|
|
"trajectory_family",
|
||
|
|
"replicate",
|
||
|
|
}
|
||
|
|
or key.startswith("factor_")
|
||
|
|
or key.startswith(prefix)
|
||
|
|
}
|
||
|
|
family_rows.append(selected)
|
||
|
|
if not family_rows:
|
||
|
|
raise ValueError(
|
||
|
|
f"metric family {family!r} selected no methods in this plan"
|
||
|
|
)
|
||
|
|
table_path = source_dir / f"{family}.csv"
|
||
|
|
_atomic_write_csv(table_path, family_rows)
|
||
|
|
table_files[str(table_path.relative_to(output_root))] = file_sha256(table_path)
|
||
|
|
family_row_counts[family] = len(family_rows)
|
||
|
|
|
||
|
|
manifest = {
|
||
|
|
"kind": "paper_artifact_manifest",
|
||
|
|
"schema_version": SCHEMA_VERSION,
|
||
|
|
"created_utc": utc_now(),
|
||
|
|
"source_batch": str(batch_dir),
|
||
|
|
"source_plan_hash": plan["plan_hash"],
|
||
|
|
"metric_configuration": to_jsonable(metric_configuration),
|
||
|
|
"metric_configuration_hash": stable_hash(
|
||
|
|
metric_configuration, prefix="metric-configuration"
|
||
|
|
),
|
||
|
|
"row_count": len(rows),
|
||
|
|
"family_row_counts": family_row_counts,
|
||
|
|
"row_hash": stable_hash(rows, prefix="trial-metric-rows"),
|
||
|
|
"storage_formats": {
|
||
|
|
"trial_metrics": STORAGE_FORMATS["trial_metrics"],
|
||
|
|
"paper_source_data": STORAGE_FORMATS["paper_source_data"],
|
||
|
|
},
|
||
|
|
"inputs": input_files,
|
||
|
|
"outputs": {
|
||
|
|
str(metric_path.relative_to(output_root)): file_sha256(metric_path),
|
||
|
|
**table_files,
|
||
|
|
},
|
||
|
|
}
|
||
|
|
manifest_path = output_root / "paper" / "artifact_manifest.json"
|
||
|
|
atomic_write_json(manifest_path, manifest)
|
||
|
|
return manifest
|
||
|
|
|
||
|
|
|
||
|
|
def build_parser() -> argparse.ArgumentParser:
|
||
|
|
parser = argparse.ArgumentParser(
|
||
|
|
description="Recompute trial metrics and paper source-data tables"
|
||
|
|
)
|
||
|
|
parser.add_argument("--batch-dir", type=Path, required=True)
|
||
|
|
parser.add_argument("--metric-config", type=Path, required=True)
|
||
|
|
parser.add_argument("--output-root", type=Path, default=None)
|
||
|
|
return parser
|
||
|
|
|
||
|
|
|
||
|
|
def main(argv: list[str] | None = None) -> int:
|
||
|
|
args = build_parser().parse_args(argv)
|
||
|
|
configuration = load_document(args.metric_config)
|
||
|
|
manifest = generate_paper_source_data(
|
||
|
|
args.batch_dir,
|
||
|
|
configuration,
|
||
|
|
output_root=args.output_root,
|
||
|
|
)
|
||
|
|
print(json.dumps(manifest, indent=2, ensure_ascii=False, sort_keys=True))
|
||
|
|
return 0
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
raise SystemExit(main())
|