File size: 4,292 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
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
"""Run a restartable sequence-model benchmark family from the command line."""

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.deep import run_grouped_sequence_benchmark

MODEL_FAMILIES = {
    "recurrent": ("vanilla_lstm", "bidirectional_lstm", "gru", "attention_lstm"),
    "transformer": ("battery_gpt", "temporal_fusion_transformer", "itransformer", "physics_itransformer"),
    "graph_variational": ("dynamic_graph_itransformer", "vae_lstm"),
}


def _read_csv_or_empty(path: Path) -> pd.DataFrame:
    return pd.read_csv(path) if path.exists() else pd.DataFrame()


def _atomic_write_csv(frame: pd.DataFrame, path: Path) -> None:
    temporary = path.with_suffix(path.suffix + ".tmp")
    frame.to_csv(temporary, index=False)
    temporary.replace(path)


def _model_is_complete(
    metrics: pd.DataFrame,
    predictions: pd.DataFrame,
    model_id: str,
    seeds: tuple[int, ...],
    n_splits: int,
) -> bool:
    if metrics.empty or predictions.empty or "model" not in metrics or "model" not in predictions:
        return False
    subset = metrics[metrics["model"] == model_id]
    pred_subset = predictions[predictions["model"] == model_id]
    expected = {(seed, fold) for seed in seeds for fold in range(1, n_splits + 1)}
    observed = set(zip(subset.get("seed", []), subset.get("fold", [])))
    return observed == expected and not pred_subset.empty


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("dataset", choices=("nasa", "calce", "oxford"))
    parser.add_argument("family", choices=tuple(MODEL_FAMILIES))
    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()

    feature_dir = args.project_root / "artifacts" / "v3" / "features" / args.dataset
    result_dir = args.project_root / "artifacts" / "v3" / "results"
    result_dir.mkdir(parents=True, exist_ok=True)
    X = np.load(feature_dir / "sequences.npz")["X"]
    index = pd.read_csv(feature_dir / "sequence_index.csv")
    prefix = result_dir / f"{args.dataset}_{args.family}"
    metrics_path = Path(f"{prefix}_fold_metrics.csv")
    predictions_path = Path(f"{prefix}_predictions.csv")
    metrics = _read_csv_or_empty(metrics_path)
    predictions = _read_csv_or_empty(predictions_path)
    seeds = tuple(args.seeds)
    n_splits = min(5, index["battery_id"].nunique())
    dataset_name = args.dataset.upper() if args.dataset == "nasa" else args.dataset.title()

    for model_id in MODEL_FAMILIES[args.family]:
        if _model_is_complete(metrics, predictions, model_id, seeds, n_splits):
            print(f"[{dataset_name}] model={model_id} checkpoint complete; skipping", flush=True)
            continue
        model_metrics, model_predictions = run_grouped_sequence_benchmark(
            X,
            index,
            dataset_name=dataset_name,
            n_splits=n_splits,
            seeds=seeds,
            max_epochs=args.max_epochs,
            patience=args.patience,
            batch_size=args.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, metrics_path)
        _atomic_write_csv(predictions, predictions_path)
        print(f"[{dataset_name}] model={model_id} checkpoint written", flush=True)

    print(metrics.groupby("model")[["mae", "rmse", "r2", "within_5pp", "epochs"]].mean().sort_values("mae"))


if __name__ == "__main__":
    main()