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()