Download scripts/run_protocol_comparison.py from NeerajCodz/aiBatteryLifeCycle: direct link, hf CLI and curl.
- Browser
- Download file 4.79 kB
-
https://huggingface.co/spaces/NeerajCodz/aiBatteryLifeCycle/resolve/main/scripts/run_protocol_comparison.py
- Command line
-
hf download hf://spaces/NeerajCodz/aiBatteryLifeCycle/scripts/run_protocol_comparison.py
-
curl -L -o run_protocol_comparison.py https://huggingface.co/spaces/NeerajCodz/aiBatteryLifeCycle/resolve/main/scripts/run_protocol_comparison.py
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() | |