Skip to content

Commit 504a2c6

Browse files
committed
feat: apply copilot review
1 parent d5d9edf commit 504a2c6

4 files changed

Lines changed: 347 additions & 15 deletions

File tree

epftoolbox2/evaluators/rmae.py

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
import math
12
from typing import Dict
23

34
import pandas as pd
@@ -12,12 +13,15 @@ def __init__(self, base_model: str):
1213

1314
def compute(self, df: pd.DataFrame, **kwargs) -> float:
1415
model_dfs: Dict[str, pd.DataFrame] = kwargs.get("model_dfs", {})
15-
if self.base_model not in model_dfs:
16+
if not model_dfs:
1617
raise ValueError(
17-
f"rMAE base model '{self.base_model}' not found in pipeline models. "
18-
f"Available: {list(model_dfs)}"
18+
f"rMAE base model '{self.base_model}' not found in pipeline models."
1919
)
20+
if self.base_model not in model_dfs:
21+
return math.nan
2022
base_df = model_dfs[self.base_model]
23+
if base_df.empty:
24+
return math.nan
2125
base_mae = (base_df["prediction"] - base_df["actual"]).abs().mean()
2226
if base_mae == 0:
2327
return float("inf")

epftoolbox2/exporters/csv.py

Lines changed: 17 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,8 @@
66
from .base import Exporter
77
from ..results.report import EvaluationReport
88

9+
_MERGE_KEYS = ["run_date", "target_date", "hour", "horizon"]
10+
911

1012
class CsvExporter(Exporter):
1113
def __init__(self, path: str, extra_columns: Optional[List[str]] = None):
@@ -25,21 +27,25 @@ def export(self, report: EvaluationReport) -> None:
2527

2628
self.path.parent.mkdir(parents=True, exist_ok=True)
2729

28-
sort_keys = ["target_date", "hour", "horizon"]
2930
base_cols = ["run_date", "target_date", "hour", "horizon", "day_in_test", "actual"]
30-
3131
base_df: Optional[pd.DataFrame] = None
32-
model_names: List[str] = []
3332

3433
for model_name, model_df in report.iter_details():
35-
model_df = model_df.sort_values(by=sort_keys).reset_index(drop=True)
36-
if base_df is None:
37-
base_df = model_df[base_cols].copy()
38-
base_df[f"{model_name}_prediction"] = model_df["prediction"].values
39-
base_df[f"{model_name}_error"] = (
40-
model_df["prediction"].values - base_df["actual"].values
34+
model_df = model_df.rename(columns={
35+
"prediction": f"{model_name}_prediction",
36+
})
37+
model_df[f"{model_name}_error"] = (
38+
model_df[f"{model_name}_prediction"] - model_df["actual"]
4139
)
42-
model_names.append(model_name)
40+
keep = _MERGE_KEYS + [f"{model_name}_prediction", f"{model_name}_error"]
41+
if base_df is None:
42+
base_df = model_df[base_cols + [f"{model_name}_prediction", f"{model_name}_error"]].copy()
43+
else:
44+
base_df = base_df.merge(
45+
model_df[keep],
46+
on=_MERGE_KEYS,
47+
how="outer",
48+
)
4349
del model_df
4450

4551
if base_df is None:
@@ -48,6 +54,7 @@ def export(self, report: EvaluationReport) -> None:
4854
if self.extra_columns and report.source_data is not None:
4955
base_df = self._join_extra_columns(base_df, report.source_data)
5056

57+
sort_keys = ["target_date", "hour", "horizon"]
5158
base_df = base_df.sort_values(by=sort_keys).reset_index(drop=True)
5259
base_df.to_csv(self.path, index=False)
5360

epftoolbox2/results/report.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
from __future__ import annotations
22

3-
from typing import Dict, Iterator, List, Tuple, Union
3+
from typing import Dict, Iterator, List, Optional, Tuple, Union
44

55
import pandas as pd
66

@@ -13,7 +13,7 @@ def __init__(
1313
self,
1414
results_or_refs: Dict[str, Union[ModelResultRef, List[Dict]]],
1515
evaluators: List[Evaluator],
16-
source_data: pd.DataFrame = None,
16+
source_data: Optional[pd.DataFrame] = None,
1717
):
1818
self.evaluators = evaluators
1919
self.source_data = source_data

0 commit comments

Comments
 (0)