aiBatteryLifeCycle / scripts /run_zero_shot_benchmark.py
NeerajCodz's picture
Complete reviewer 2026-09 revision
8b37c3f
Raw History Blame Contribute Delete
5.72 kB
"""Train once on NASA and score frozen models on CALCE and Oxford."""
from __future__ import annotations
import argparse
from pathlib import Path
import sys
import numpy as np
import pandas as pd
PROJECT_ROOT = Path(__file__).resolve().parents[1]
if str(PROJECT_ROOT) not in sys.path:
sys.path.insert(0, str(PROJECT_ROOT))
from scripts.run_sequence_benchmark import MODEL_FAMILIES, _atomic_write_csv, _read_csv_or_empty
from src.experiments.classical import run_zero_shot_tabular
from src.experiments.deep import run_zero_shot_sequence
def _load_sequence(root: Path, dataset: str) -> tuple[np.ndarray, pd.DataFrame]:
folder = root / "artifacts" / "v3" / "features" / dataset
return np.load(folder / "sequences.npz")["X"], pd.read_csv(folder / "sequence_index.csv")
def _write_target_results(
result_dir: Path,
family: str,
metrics: pd.DataFrame,
predictions: pd.DataFrame,
*,
merge: bool = False,
) -> None:
for target in ("calce", "oxford"):
label = target.title()
target_metrics = metrics[metrics["target_dataset"] == label].copy()
target_predictions = predictions[predictions["target_dataset"] == label].copy()
metrics_path = result_dir / f"nasa_to_{target}_{family}_metrics.csv"
predictions_path = result_dir / f"nasa_to_{target}_{family}_predictions.csv"
if merge:
model_ids = set(target_metrics["model"])
old_metrics = _read_csv_or_empty(metrics_path)
old_predictions = _read_csv_or_empty(predictions_path)
if not old_metrics.empty and "model" in old_metrics:
old_metrics = old_metrics[~old_metrics["model"].isin(model_ids)]
if not old_predictions.empty and "model" in old_predictions:
old_predictions = old_predictions[~old_predictions["model"].isin(model_ids)]
target_metrics = pd.concat([old_metrics, target_metrics], ignore_index=True)
target_predictions = pd.concat([old_predictions, target_predictions], ignore_index=True)
_atomic_write_csv(target_metrics, metrics_path)
_atomic_write_csv(target_predictions, predictions_path)
def _zero_shot_model_is_complete(
result_dir: Path,
family: str,
model_id: str,
seeds: tuple[int, ...],
) -> bool:
for target in ("calce", "oxford"):
metrics = _read_csv_or_empty(result_dir / f"nasa_to_{target}_{family}_metrics.csv")
predictions = _read_csv_or_empty(result_dir / f"nasa_to_{target}_{family}_predictions.csv")
if metrics.empty or predictions.empty or "model" not in metrics or "model" not in predictions:
return False
model_metrics = metrics[metrics["model"] == model_id]
model_predictions = predictions[predictions["model"] == model_id]
if set(model_metrics.get("seed", [])) != set(seeds) or len(model_metrics) != len(seeds):
return False
if model_predictions.empty:
return False
return True
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("family", choices=("classical", *MODEL_FAMILIES))
parser.add_argument("--project-root", type=Path, default=PROJECT_ROOT)
parser.add_argument("--max-epochs", type=int, default=200)
parser.add_argument("--patience", type=int, default=20)
parser.add_argument("--batch-size", type=int, default=64)
parser.add_argument("--seeds", type=int, nargs="+", default=(17, 42, 2026))
args = parser.parse_args()
root = args.project_root
result_dir = root / "artifacts" / "v3" / "results"
result_dir.mkdir(parents=True, exist_ok=True)
if args.family == "classical":
source = pd.read_csv(root / "artifacts" / "v3" / "features" / "nasa" / "features.csv")
for target in ("calce", "oxford"):
frame = pd.read_csv(root / "artifacts" / "v3" / "features" / target / "features.csv")
metrics, predictions = run_zero_shot_tabular(
source, frame, target_name=target.title(), random_state=42,
)
metrics.to_csv(result_dir / f"nasa_to_{target}_classical_metrics.csv", index=False)
predictions.to_csv(result_dir / f"nasa_to_{target}_classical_predictions.csv", index=False)
print(target, metrics[["model", "mae", "rmse", "r2", "within_5pp"]].sort_values("mae").to_string(index=False))
return
source_X, source_index = _load_sequence(root, "nasa")
targets = {
dataset.title(): _load_sequence(root, dataset)
for dataset in ("calce", "oxford")
}
seeds = tuple(args.seeds)
for model_id in MODEL_FAMILIES[args.family]:
if _zero_shot_model_is_complete(result_dir, args.family, model_id, seeds):
print(f"[NASA->external] model={model_id} checkpoint complete; skipping", flush=True)
continue
metrics, predictions = run_zero_shot_sequence(
source_X,
source_index,
targets,
seeds=seeds,
max_epochs=args.max_epochs,
patience=args.patience,
batch_size=args.batch_size,
model_ids=(model_id,),
)
_write_target_results(result_dir, args.family, metrics, predictions, merge=True)
print(f"[NASA->external] model={model_id} checkpoint written", flush=True)
all_metrics = []
for target in ("calce", "oxford"):
all_metrics.append(_read_csv_or_empty(result_dir / f"nasa_to_{target}_{args.family}_metrics.csv"))
metrics = pd.concat(all_metrics, ignore_index=True)
print(metrics.groupby(["target_dataset", "model"])[["mae", "rmse", "r2", "within_5pp", "epochs"]].mean().to_string())
if __name__ == "__main__":
main()