File size: 2,378 Bytes
8b37c3f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 | """Normalize all raw datasets and build leakage-safe benchmark bundles."""
from __future__ import annotations
import argparse
from pathlib import Path
from src.data.adapters import (
iter_calce_xlsx_cycles,
iter_nasa_mat_cycles,
iter_oxford_mat_cycles,
)
from src.data.benchmark import build_benchmark_bundle, save_benchmark_bundle
from src.utils.config import EXCLUDED_BATTERIES
def load_all_cycles(raw_root: str | Path) -> dict[str, list]:
raw_root = Path(raw_root)
nasa_paths = sorted((raw_root / "nasa" / "5_Battery_Data_Set").rglob("B*.mat"))
if not nasa_paths:
raise FileNotFoundError("NASA MATLAB files are missing; run the benchmark downloader")
oxford_path = raw_root / "oxford" / "Oxford_Battery_Degradation_Dataset_1.mat"
if not oxford_path.exists():
raise FileNotFoundError("Oxford MATLAB file is missing; run the benchmark downloader")
return {
"NASA": list(iter_nasa_mat_cycles(nasa_paths)),
"CALCE": list(iter_calce_xlsx_cycles(raw_root / "calce")),
"Oxford": list(iter_oxford_mat_cycles(oxford_path)),
}
def build_all_benchmarks(raw_root: str | Path, output_root: str | Path) -> list[dict[str, object]]:
output_root = Path(output_root)
summaries = []
for name, cycles in load_all_cycles(raw_root).items():
bundle = build_benchmark_bundle(
cycles,
excluded_batteries=EXCLUDED_BATTERIES if name == "NASA" else set(),
)
destination = output_root / name.lower()
save_benchmark_bundle(bundle, destination)
summaries.append({
"dataset": name,
"cycles_loaded": len(cycles),
"cycles_retained": len(bundle.features),
"cycles_excluded": len(bundle.exclusions),
"batteries": bundle.features["battery_id"].nunique(),
"features": bundle.features.shape[1],
"sequence_shape": tuple(bundle.sequences.shape),
})
return summaries
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--raw-root", type=Path, default=Path("datasets/raw"))
parser.add_argument("--output-root", type=Path, default=Path("artifacts/v3/features"))
args = parser.parse_args()
for row in build_all_benchmarks(args.raw_root, args.output_root):
print(row)
if __name__ == "__main__":
main()
|