From b18ed7ce276717bec563b24f06362b051a5eaac1 Mon Sep 17 00:00:00 2001 From: you-n-g Date: Tue, 15 Apr 2025 17:10:31 +0800 Subject: [PATCH] fix: update metric direction to return bool (#791) --- rdagent/app/data_science/loop.py | 1 + rdagent/scenarios/data_science/scen/__init__.py | 6 ++++-- 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/rdagent/app/data_science/loop.py b/rdagent/app/data_science/loop.py index abb65c6b..cf7d91ee 100644 --- a/rdagent/app/data_science/loop.py +++ b/rdagent/app/data_science/loop.py @@ -200,6 +200,7 @@ def main( DS_RD_SETTING.competition = competition if DS_RD_SETTING.competition: + if DS_RD_SETTING.scen.endswith("KaggleScen"): download_data(competition=DS_RD_SETTING.competition, settings=DS_RD_SETTING) else: diff --git a/rdagent/scenarios/data_science/scen/__init__.py b/rdagent/scenarios/data_science/scen/__init__.py index 5f8538e5..a6ca2b93 100644 --- a/rdagent/scenarios/data_science/scen/__init__.py +++ b/rdagent/scenarios/data_science/scen/__init__.py @@ -31,7 +31,9 @@ class DataScienceScen(Scenario): self.raw_description = self._get_description() self.processed_data_folder_description = self._get_data_folder_description() self._analysis_competition_description() - self.metric_direction = self._get_direction() + self.metric_direction: bool = ( + self._get_direction() + ) # True indicates higher is better, False indicates lower is better def _get_description(self): if (fp := Path(f"{DS_RD_SETTING.local_data_path}/{self.competition}.json")).exists(): @@ -148,7 +150,7 @@ class KaggleScen(DataScienceScen): def _get_direction(self): leaderboard = leaderboard_scores(self.competition) - return "maximize" if float(leaderboard[0]) > float(leaderboard[-1]) else "minimize" + return float(leaderboard[0]) > float(leaderboard[-1]) @property def rich_style_description(self) -> str: