Download scripts/run_external_validation.py from NeerajCodz/aiBatteryLifeCycle: direct link, hf CLI and curl.
- Browser
- Download file 7.59 kB
-
https://huggingface.co/spaces/NeerajCodz/aiBatteryLifeCycle/resolve/main/scripts/run_external_validation.py
- Command line
-
hf download hf://spaces/NeerajCodz/aiBatteryLifeCycle/scripts/run_external_validation.py
-
curl -L -o run_external_validation.py https://huggingface.co/spaces/NeerajCodz/aiBatteryLifeCycle/resolve/main/scripts/run_external_validation.py
7.59 kB
| """Within-dataset and NASA-to-external validation for CALCE or 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 src.experiments.classical import run_grouped_tabular_benchmark, run_zero_shot_tabular | |
| from src.experiments.deep import run_grouped_sequence_benchmark, run_zero_shot_sequence | |
| from scripts.run_sequence_benchmark import ( | |
| MODEL_FAMILIES, | |
| _atomic_write_csv, | |
| _model_is_complete, | |
| _read_csv_or_empty, | |
| ) | |
| from scripts.run_zero_shot_benchmark import _write_target_results, _zero_shot_model_is_complete | |
| 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 _ensure_grouped_results( | |
| root: Path, | |
| dataset: str, | |
| seeds: tuple[int, ...], | |
| max_epochs: int, | |
| patience: int, | |
| batch_size: int, | |
| ) -> None: | |
| results = root / "artifacts" / "v3" / "results" | |
| features = root / "artifacts" / "v3" / "features" / dataset | |
| classical_metrics = results / f"{dataset}_classical_fold_metrics.csv" | |
| classical_predictions = results / f"{dataset}_classical_predictions.csv" | |
| if not classical_metrics.exists() or not classical_predictions.exists(): | |
| frame = pd.read_csv(features / "features.csv") | |
| metrics, predictions = run_grouped_tabular_benchmark( | |
| frame, | |
| dataset_name=dataset.title(), | |
| n_splits=min(5, frame["battery_id"].nunique()), | |
| seeds=seeds, | |
| ) | |
| metrics.to_csv(classical_metrics, index=False) | |
| predictions.to_csv(classical_predictions, index=False) | |
| X, index = _load_sequence(root, dataset) | |
| for family, model_ids in MODEL_FAMILIES.items(): | |
| metric_path = results / f"{dataset}_{family}_fold_metrics.csv" | |
| prediction_path = results / f"{dataset}_{family}_predictions.csv" | |
| metrics = _read_csv_or_empty(metric_path) | |
| predictions = _read_csv_or_empty(prediction_path) | |
| n_splits = min(5, index["battery_id"].nunique()) | |
| for model_id in model_ids: | |
| if _model_is_complete(metrics, predictions, model_id, seeds, n_splits): | |
| continue | |
| model_metrics, model_predictions = run_grouped_sequence_benchmark( | |
| X, | |
| index, | |
| dataset_name=dataset.title(), | |
| n_splits=n_splits, | |
| seeds=seeds, | |
| max_epochs=max_epochs, | |
| patience=patience, | |
| batch_size=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, metric_path) | |
| _atomic_write_csv(predictions, prediction_path) | |
| def _ensure_zero_shot_results( | |
| root: Path, | |
| dataset: str, | |
| seeds: tuple[int, ...], | |
| max_epochs: int, | |
| patience: int, | |
| batch_size: int, | |
| ) -> None: | |
| results = root / "artifacts" / "v3" / "results" | |
| feature_root = root / "artifacts" / "v3" / "features" | |
| source = pd.read_csv(feature_root / "nasa" / "features.csv") | |
| target = pd.read_csv(feature_root / dataset / "features.csv") | |
| tabular_metrics = results / f"nasa_to_{dataset}_classical_metrics.csv" | |
| tabular_predictions = results / f"nasa_to_{dataset}_classical_predictions.csv" | |
| if not tabular_metrics.exists() or not tabular_predictions.exists(): | |
| metrics, predictions = run_zero_shot_tabular( | |
| source, | |
| target, | |
| target_name=dataset.title(), | |
| random_state=42, | |
| ) | |
| metrics.to_csv(tabular_metrics, index=False) | |
| predictions.to_csv(tabular_predictions, index=False) | |
| source_X, source_index = _load_sequence(root, "nasa") | |
| targets = { | |
| target_name.title(): _load_sequence(root, target_name) | |
| for target_name in ("calce", "oxford") | |
| } | |
| for family, model_ids in MODEL_FAMILIES.items(): | |
| for model_id in model_ids: | |
| if _zero_shot_model_is_complete(results, family, model_id, seeds): | |
| continue | |
| metrics, predictions = run_zero_shot_sequence( | |
| source_X, | |
| source_index, | |
| targets, | |
| seeds=seeds, | |
| max_epochs=max_epochs, | |
| patience=patience, | |
| batch_size=batch_size, | |
| model_ids=(model_id,), | |
| ) | |
| _write_target_results(results, family, metrics, predictions, merge=True) | |
| def run_external_validation( | |
| project_root: str | Path, | |
| *, | |
| dataset: str, | |
| seeds: tuple[int, ...] = (17, 42, 2026), | |
| max_epochs: int = 200, | |
| patience: int = 20, | |
| batch_size: int = 64, | |
| ) -> pd.DataFrame: | |
| dataset = dataset.lower() | |
| if dataset not in {"calce", "oxford"}: | |
| raise ValueError("dataset must be 'calce' or 'oxford'") | |
| root = Path(project_root) | |
| result_dir = root / "artifacts" / "v3" / "results" | |
| result_dir.mkdir(parents=True, exist_ok=True) | |
| _ensure_grouped_results(root, dataset, seeds, max_epochs, patience, batch_size) | |
| _ensure_zero_shot_results(root, dataset, seeds, max_epochs, patience, batch_size) | |
| frames = [] | |
| for path in sorted(result_dir.glob(f"{dataset}_*_fold_metrics.csv")): | |
| frame = pd.read_csv(path) | |
| frame["validation"] = "within_dataset_grouped" | |
| frames.append(frame) | |
| for path in sorted(result_dir.glob(f"nasa_to_{dataset}_*_metrics.csv")): | |
| frame = pd.read_csv(path) | |
| frame["validation"] = "nasa_zero_shot" | |
| frames.append(frame) | |
| combined = pd.concat(frames, ignore_index=True) | |
| summary = ( | |
| combined.groupby(["validation", "model"], as_index=False) | |
| .agg( | |
| mae=("mae", "mean"), | |
| rmse=("rmse", "mean"), | |
| r2=("r2", "mean"), | |
| adjusted_r2=("adjusted_r2", "mean"), | |
| mape=("mape", "mean"), | |
| within_5pp=("within_5pp", "mean"), | |
| ) | |
| .sort_values(["validation", "mae"]) | |
| ) | |
| return summary | |
| def main() -> None: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("dataset", choices=("calce", "oxford")) | |
| 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() | |
| summary = run_external_validation( | |
| args.project_root, | |
| dataset=args.dataset, | |
| seeds=tuple(args.seeds), | |
| max_epochs=args.max_epochs, | |
| patience=args.patience, | |
| batch_size=args.batch_size, | |
| ) | |
| output = args.project_root / "artifacts" / "v3" / "results" / f"{args.dataset}_validation_summary.csv" | |
| summary.to_csv(output, index=False) | |
| print(summary.to_string(index=False)) | |
| if __name__ == "__main__": | |
| main() | |