bgc-setnet / source /src /bgc_retrieval /multiseed_analysis.py
whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
11.4 kB
"""Aggregate completed training-seed campaigns without treating seeds as replicates."""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any, Iterable
import pandas as pd
from .artifacts import sha256_file, write_json_immutable
def _normalise_method(method: str) -> str:
if method.startswith("ensemble_validation_alpha_"):
return "ensemble_validation_selected"
return method
def _load_json(path: Path) -> dict[str, Any]:
with path.open("r", encoding="utf-8") as handle:
return json.load(handle)
def _append(
rows: list[dict[str, Any]],
*,
campaign_tag: str,
training_seed: int,
metric_id: str,
domain: str,
run: str,
method: str,
subset: str,
metric: str,
statistic: str,
analysis_status: str,
value: float,
) -> None:
rows.append(
{
"campaign_tag": campaign_tag,
"training_seed": training_seed,
"metric_id": metric_id,
"domain": domain,
"run": run,
"method": method,
"subset": subset,
"metric": metric,
"statistic": statistic,
"analysis_status": analysis_status,
"value": float(value),
}
)
def aggregate_seed_campaigns(
artifact_root: str | Path,
campaign_tags: Iterable[str],
output_dir: str | Path,
) -> Path:
"""Create tidy per-seed and descriptive across-seed summaries.
Within-campaign confidence intervals and tests are intentionally not pooled.
The across-seed standard deviation describes optimization-seed sensitivity;
five seeds are not treated as independent biological replicates.
"""
root = Path(artifact_root)
output = Path(output_dir)
tags = list(campaign_tags)
if len(tags) < 2:
raise ValueError("At least two campaign tags are required")
if len(tags) != len(set(tags)):
raise ValueError("Campaign tags must be unique")
output.mkdir(parents=True, exist_ok=False)
rows: list[dict[str, Any]] = []
inputs: list[Path] = []
tag_seeds: list[dict[str, Any]] = []
for tag in tags:
config_path = root / f"{tag}-main/config.json"
config = _load_json(config_path)
training_seed = int(config["project"]["seed"])
tag_seeds.append({"campaign_tag": tag, "training_seed": training_seed})
inputs.append(config_path)
for run in ("main", "no-phase1"):
normalized_run = run.replace("-", "_")
external_path = root / f"{tag}-{run}-external/external_similarity_summary.csv"
inputs.append(external_path)
for record in pd.read_csv(external_path).to_dict("records"):
method = _normalise_method(str(record["method"]))
subset = str(record["subset"])
_append(
rows,
campaign_tag=tag,
training_seed=training_seed,
metric_id=(
f"external.{normalized_run}.{method}.{subset}.spearman_r"
),
domain="external",
run=normalized_run,
method=method,
subset=subset,
metric="structural_similarity",
statistic="spearman_r",
analysis_status="campaign_output",
value=record["spearman_r"],
)
internal_path = root / f"{tag}-{run}-evaluation/summary.csv"
metadata_path = root / f"{tag}-{run}-evaluation/metadata.json"
inputs.extend([internal_path, metadata_path])
for record in pd.read_csv(internal_path).to_dict("records"):
if str(record["metric"]) == "tie_fraction":
continue
method = _normalise_method(str(record["method"]))
metric = str(record["metric"])
_append(
rows,
campaign_tag=tag,
training_seed=training_seed,
metric_id=f"internal.{normalized_run}.{method}.{metric}.mean",
domain="internal_silver",
run=normalized_run,
method=method,
subset="test",
metric=metric,
statistic="mean",
analysis_status="development_only",
value=record["mean"],
)
selected_alpha = float(_load_json(metadata_path)["selected_alpha"])
_append(
rows,
campaign_tag=tag,
training_seed=training_seed,
metric_id=(
f"internal.{normalized_run}.ensemble_validation_selected.alpha"
),
domain="internal_silver",
run=normalized_run,
method="ensemble_validation_selected",
subset="validation",
metric="alpha",
statistic="selected_value",
analysis_status="development_only",
value=selected_alpha,
)
analysis_dir = root / f"{tag}-analysis"
external_paired_path = analysis_dir / "external_paired_comparisons.csv"
internal_paired_path = analysis_dir / "internal_paired_comparisons.csv"
exact_path = analysis_dir / "exact_product_summary.csv"
training_path = analysis_dir / "training_summary.json"
inputs.extend(
[
external_paired_path,
internal_paired_path,
exact_path,
training_path,
]
)
for record in pd.read_csv(external_paired_path).to_dict("records"):
comparison = str(record["comparison"])
subset = str(record["subset"])
_append(
rows,
campaign_tag=tag,
training_seed=training_seed,
metric_id=(
f"external_paired.{comparison}.{subset}.delta_spearman"
),
domain="external_paired",
run="comparison",
method=comparison,
subset=subset,
metric="structural_similarity",
statistic="delta_spearman",
analysis_status=str(record["analysis_status"]),
value=record["delta_spearman"],
)
for record in pd.read_csv(internal_paired_path).to_dict("records"):
family = str(record["family"])
metric = str(record["metric"])
_append(
rows,
campaign_tag=tag,
training_seed=training_seed,
metric_id=f"internal_paired.{family}.{metric}.mean_delta",
domain="internal_silver_paired",
run="comparison",
method=family,
subset="test",
metric=metric,
statistic="mean_delta",
analysis_status=str(record["analysis_status"]),
value=record["mean_delta"],
)
for record in pd.read_csv(exact_path).to_dict("records"):
run = "main" if "-main-external" in str(record["run"]) else "no_phase1"
for metric in ("recall@50", "mrr", "map", "ndcg@50", "precision@50"):
_append(
rows,
campaign_tag=tag,
training_seed=training_seed,
metric_id=f"exact_product.{run}.{record['method']}.{metric}.mean",
domain="exact_product",
run=run,
method=str(record["method"]),
subset="eligible_references",
metric=metric,
statistic="mean",
analysis_status="post_hoc",
value=record[metric],
)
training = _load_json(training_path)
objectives = {
"phase1": "validation_loss",
"phase2_main": "validation_recall@50",
"phase2_no_phase1": "validation_recall@50",
}
for phase, objective in objectives.items():
summary = training[phase]
for statistic, value in (
("epochs", summary["epochs"]),
("best_epoch", summary["best"]["epoch"]),
(f"best_{objective}", summary["best"][objective]),
):
_append(
rows,
campaign_tag=tag,
training_seed=training_seed,
metric_id=f"training.{phase}.{statistic}",
domain="training",
run=phase,
method="setnet",
subset="validation",
metric=objective,
statistic=statistic,
analysis_status="campaign_output",
value=value,
)
seeds = [entry["training_seed"] for entry in tag_seeds]
if len(seeds) != len(set(seeds)):
raise ValueError(f"Training seeds must be unique; found {seeds}")
seed_level = pd.DataFrame(rows).sort_values(
["metric_id", "training_seed"], kind="stable"
)
counts = seed_level.groupby(["metric_id", "training_seed"]).size()
if int(counts.max()) != 1:
duplicates = counts[counts > 1].index.tolist()
raise ValueError(f"Duplicate per-seed metrics found: {duplicates[:5]}")
seed_level.to_csv(output / "seed_level_metrics.csv", index=False)
descriptors = [
"metric_id",
"domain",
"run",
"method",
"subset",
"metric",
"statistic",
"analysis_status",
]
aggregate = (
seed_level.groupby(descriptors, as_index=False, dropna=False)["value"]
.agg(
n="count",
mean="mean",
sample_std="std",
median="median",
minimum="min",
maximum="max",
)
.sort_values("metric_id", kind="stable")
)
expected = len(tags)
if not (aggregate["n"] == expected).all():
incomplete = aggregate.loc[aggregate["n"] != expected, ["metric_id", "n"]]
raise ValueError(
"Every normalized metric must be present for every seed: "
+ incomplete.to_dict("records").__repr__()
)
aggregate.to_csv(output / "aggregate_summary.csv", index=False)
metadata = {
"schema_version": 1,
"analysis_type": "descriptive_training_seed_sensitivity",
"campaigns": tag_seeds,
"campaign_count": len(tags),
"interpretation_warning": (
"Across-seed standard deviations describe optimization sensitivity. "
"They are not biological-replicate uncertainty or inferential confidence intervals."
),
"within_campaign_inference": (
"Within-campaign bootstrap intervals and paired tests remain in each "
"campaign analysis directory and are not pooled here."
),
"input_sha256": {
str(path.relative_to(root)): sha256_file(path) for path in sorted(set(inputs))
},
}
write_json_immutable(output / "multiseed_metadata.json", metadata)
return output