aiBatteryLifeCycle / scripts /run_protocol_comparison.py
NeerajCodz's picture
Complete reviewer 2026-09 revision
8b37c3f
Raw History Blame Contribute Delete
4.79 kB
"""Quantify split and target-proxy effects for the v1/v2/v3 discussion."""
from __future__ import annotations
from pathlib import Path
import argparse
import sys
PROJECT_ROOT = Path(__file__).resolve().parents[1]
if str(PROJECT_ROOT) not in sys.path:
sys.path.insert(0, str(PROJECT_ROOT))
import numpy as np
import pandas as pd
from sklearn.ensemble import ExtraTreesRegressor
from sklearn.impute import SimpleImputer
from sklearn.model_selection import train_test_split
from sklearn.pipeline import Pipeline
from src.evaluation.metrics import regression_metrics
from src.evaluation.protocol import grouped_train_val_test_folds
from src.utils.config import FEATURE_COLS_V3
LEAKAGE_PROXY_COLS = [
"current_capacity_retention",
"current_cumulative_capacity",
"current_delta_capacity",
"current_soh_rolling_mean",
]
def _add_legacy_target_proxies(frame: pd.DataFrame) -> pd.DataFrame:
out = frame.sort_values(["battery_id", "cycle_number"]).copy()
out["current_capacity_retention"] = out["capacity_ah"] / out["reference_capacity_ah"]
out["current_cumulative_capacity"] = out.groupby("battery_id")["capacity_ah"].cumsum()
out["current_delta_capacity"] = out.groupby("battery_id")["capacity_ah"].diff().fillna(0)
out["current_soh_rolling_mean"] = out.groupby("battery_id")["SoH"].transform(
lambda values: values.rolling(5, min_periods=1).mean()
)
return out.reset_index(drop=True)
def _chronological_indices(frame: pd.DataFrame) -> tuple[np.ndarray, np.ndarray]:
train, test = [], []
for _, group in frame.groupby("battery_id"):
ordered = group.sort_values("cycle_number").index.to_numpy()
split = max(1, int(0.8 * len(ordered)))
train.extend(ordered[:split])
test.extend(ordered[split:])
return np.asarray(train), np.asarray(test)
def _score(frame: pd.DataFrame, features: list[str], train_idx: np.ndarray, test_idx: np.ndarray, seed: int) -> dict[str, float]:
model = Pipeline([
("imputer", SimpleImputer(strategy="median")),
("model", ExtraTreesRegressor(
n_estimators=300, min_samples_leaf=2, random_state=seed, n_jobs=-1
)),
])
model.fit(frame.loc[train_idx, features], frame.loc[train_idx, "SoH"])
pred = model.predict(frame.loc[test_idx, features])
return regression_metrics(
frame.loc[test_idx, "SoH"].to_numpy(), pred, n_predictors=len(features)
)
def run_protocol_comparison(
project_root: str | Path,
*,
seeds: tuple[int, ...] = (17, 42, 2026),
) -> pd.DataFrame:
project_root = Path(project_root)
frame = pd.read_csv(project_root / "artifacts" / "v3" / "features" / "nasa" / "features.csv")
frame = _add_legacy_target_proxies(frame)
groups = frame["battery_id"].astype(str).to_numpy()
rows = []
for seed in seeds:
all_idx = np.arange(len(frame))
v1_train, v1_test = train_test_split(all_idx, test_size=0.2, random_state=seed)
v2_train, v2_test = _chronological_indices(frame)
protocols = [("v1_random_cycle", 1, v1_train, v1_test), ("v2_chronological", 1, v2_train, v2_test)]
for fold, (train, _, test) in enumerate(
grouped_train_val_test_folds(groups, n_splits=5, random_state=seed), start=1
):
protocols.append(("v3_grouped_battery", fold, train, test))
for protocol, fold, train_idx, test_idx in protocols:
for feature_contract, features in [
("safe_partial_cycle", FEATURE_COLS_V3),
("legacy_target_proxies", FEATURE_COLS_V3 + LEAKAGE_PROXY_COLS),
]:
metrics = _score(frame, features, train_idx, test_idx, seed)
rows.append({
"seed": seed,
"protocol": protocol,
"fold": fold,
"feature_contract": feature_contract,
"train_batteries": frame.loc[train_idx, "battery_id"].nunique(),
"test_batteries": frame.loc[test_idx, "battery_id"].nunique(),
**metrics,
})
return pd.DataFrame(rows)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--project-root", type=Path, default=PROJECT_ROOT)
parser.add_argument("--output", type=Path, default=None)
args = parser.parse_args()
output = args.output or args.project_root / "artifacts" / "v3" / "results" / "protocol_comparison.csv"
output.parent.mkdir(parents=True, exist_ok=True)
frame = run_protocol_comparison(args.project_root)
frame.to_csv(output, index=False)
print(frame.groupby(["protocol", "feature_contract"])[["mae", "rmse", "r2", "within_5pp"]].mean())
print(f"Saved {output}")
if __name__ == "__main__":
main()