aiBatteryLifeCycle / scripts /run_sequence_benchmark.py
NeerajCodz's picture
Complete reviewer 2026-09 revision
8b37c3f
Raw History Blame Contribute Delete
4.29 kB
"""Run a restartable sequence-model benchmark family from the command line."""
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 src.experiments.deep import run_grouped_sequence_benchmark
MODEL_FAMILIES = {
"recurrent": ("vanilla_lstm", "bidirectional_lstm", "gru", "attention_lstm"),
"transformer": ("battery_gpt", "temporal_fusion_transformer", "itransformer", "physics_itransformer"),
"graph_variational": ("dynamic_graph_itransformer", "vae_lstm"),
}
def _read_csv_or_empty(path: Path) -> pd.DataFrame:
return pd.read_csv(path) if path.exists() else pd.DataFrame()
def _atomic_write_csv(frame: pd.DataFrame, path: Path) -> None:
temporary = path.with_suffix(path.suffix + ".tmp")
frame.to_csv(temporary, index=False)
temporary.replace(path)
def _model_is_complete(
metrics: pd.DataFrame,
predictions: pd.DataFrame,
model_id: str,
seeds: tuple[int, ...],
n_splits: int,
) -> bool:
if metrics.empty or predictions.empty or "model" not in metrics or "model" not in predictions:
return False
subset = metrics[metrics["model"] == model_id]
pred_subset = predictions[predictions["model"] == model_id]
expected = {(seed, fold) for seed in seeds for fold in range(1, n_splits + 1)}
observed = set(zip(subset.get("seed", []), subset.get("fold", [])))
return observed == expected and not pred_subset.empty
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("dataset", choices=("nasa", "calce", "oxford"))
parser.add_argument("family", choices=tuple(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()
feature_dir = args.project_root / "artifacts" / "v3" / "features" / args.dataset
result_dir = args.project_root / "artifacts" / "v3" / "results"
result_dir.mkdir(parents=True, exist_ok=True)
X = np.load(feature_dir / "sequences.npz")["X"]
index = pd.read_csv(feature_dir / "sequence_index.csv")
prefix = result_dir / f"{args.dataset}_{args.family}"
metrics_path = Path(f"{prefix}_fold_metrics.csv")
predictions_path = Path(f"{prefix}_predictions.csv")
metrics = _read_csv_or_empty(metrics_path)
predictions = _read_csv_or_empty(predictions_path)
seeds = tuple(args.seeds)
n_splits = min(5, index["battery_id"].nunique())
dataset_name = args.dataset.upper() if args.dataset == "nasa" else args.dataset.title()
for model_id in MODEL_FAMILIES[args.family]:
if _model_is_complete(metrics, predictions, model_id, seeds, n_splits):
print(f"[{dataset_name}] model={model_id} checkpoint complete; skipping", flush=True)
continue
model_metrics, model_predictions = run_grouped_sequence_benchmark(
X,
index,
dataset_name=dataset_name,
n_splits=n_splits,
seeds=seeds,
max_epochs=args.max_epochs,
patience=args.patience,
batch_size=args.batch_size,
model_ids=(model_id,),
)
if not metrics.empty and "model" in metrics:
metrics = metrics[metrics["model"] != model_id]
if not predictions.empty and "model" in predictions:
predictions = predictions[predictions["model"] != model_id]
metrics = pd.concat([metrics, model_metrics], ignore_index=True)
predictions = pd.concat([predictions, model_predictions], ignore_index=True)
_atomic_write_csv(metrics, metrics_path)
_atomic_write_csv(predictions, predictions_path)
print(f"[{dataset_name}] model={model_id} checkpoint written", flush=True)
print(metrics.groupby("model")[["mae", "rmse", "r2", "within_5pp", "epochs"]].mean().sort_values("mae"))
if __name__ == "__main__":
main()