diff --git a/rdagent/scenarios/data_mining/experiment/model_template/train.py b/rdagent/scenarios/data_mining/experiment/model_template/train.py index 05a670c9..f42fa347 100644 --- a/rdagent/scenarios/data_mining/experiment/model_template/train.py +++ b/rdagent/scenarios/data_mining/experiment/model_template/train.py @@ -95,6 +95,9 @@ for data in test_dataloader: acc = roc_auc_score(y_test, np.concatenate(y_pred)) print(acc) + +res = pd.Series(data=[acc], index=['AUROC']) +res.to_csv("./submission.csv") # Save the predictions to submission.csv -with open("./submission.txt", "w") as f: - f.write(str(acc)) +# with open("./submission.txt", "w") as f: +# f.write(str(acc)) diff --git a/rdagent/scenarios/data_mining/experiment/workspace.py b/rdagent/scenarios/data_mining/experiment/workspace.py index e2e9d3a1..0f9baa6b 100644 --- a/rdagent/scenarios/data_mining/experiment/workspace.py +++ b/rdagent/scenarios/data_mining/experiment/workspace.py @@ -23,10 +23,9 @@ class DMFBWorkspace(FBWorkspace): env=run_env, ) - csv_path = self.workspace_path / "submission.txt" + csv_path = self.workspace_path / "submission.csv" if not csv_path.exists(): logger.error(f"File {csv_path} does not exist.") return None - with open(self.workspace_path / "submission.txt", "r") as f: - return f.read() + return pd.read_csv(csv_path, index_col=0).iloc[:, 0]