From e76f9fe696c4eca51849c7d88802325e4c2cd6a2 Mon Sep 17 00:00:00 2001 From: liuht Date: Mon, 15 Jan 2024 08:41:09 +0800 Subject: [PATCH] [MNT] fix random seed repeatedly --- examples/dataset_table_workflow/base.py | 1 + examples/dataset_table_workflow/hetero.py | 6 ++++-- examples/dataset_table_workflow/workflow.py | 6 ++---- 3 files changed, 7 insertions(+), 6 deletions(-) diff --git a/examples/dataset_table_workflow/base.py b/examples/dataset_table_workflow/base.py index 571f4fa..a10c69a 100644 --- a/examples/dataset_table_workflow/base.py +++ b/examples/dataset_table_workflow/base.py @@ -76,6 +76,7 @@ class TableWorkflow: self.user_semantic["Name"]["Values"] = "" if len(self.market) == 0 or rebuild == True: + if retrain: set_seed(0) for learnware_id in self.benchmark.learnware_ids: with tempfile.TemporaryDirectory(prefix="table_benchmark_") as tempdir: zip_path = os.path.join(tempdir, f"{learnware_id}.zip") diff --git a/examples/dataset_table_workflow/hetero.py b/examples/dataset_table_workflow/hetero.py index 0099e63..68842fc 100644 --- a/examples/dataset_table_workflow/hetero.py +++ b/examples/dataset_table_workflow/hetero.py @@ -11,13 +11,14 @@ from learnware.reuse import AveragingReuser, FeatureAlignLearnware from methods import * from base import TableWorkflow from config import align_model_params, user_semantic, hetero_n_labeled_list, hetero_n_repeat_list -from utils import Recorder, plot_performance_curves +from utils import Recorder, plot_performance_curves, set_seed logger = get_module_logger("hetero_test", level="INFO") class HeterogeneousDatasetWorkflow(TableWorkflow): def unlabeled_hetero_table_example(self): + set_seed(0) logger.info("Total Item: %d" % len(self.market)) learnware_rmse_list = [] single_score_list = [] @@ -39,7 +40,7 @@ class HeterogeneousDatasetWorkflow(TableWorkflow): ) logger.info(f"Searching Market for user: {user}_{idx}") - search_result = self.market.search_learnware(user_info, max_search_num=10) + search_result = self.market.search_learnware(user_info, search_method="auto") single_result = search_result.get_single_results() multiple_result = search_result.get_multiple_results() @@ -107,6 +108,7 @@ class HeterogeneousDatasetWorkflow(TableWorkflow): def labeled_hetero_table_example(self, skip_test): + set_seed(0) logger.info("Total Items: %d" % len(self.market)) methods = ["user_model", "hetero_single_aug", "hetero_multiple_avg", "hetero_ensemble_pruning"] recorders = {method: Recorder() for method in methods} diff --git a/examples/dataset_table_workflow/workflow.py b/examples/dataset_table_workflow/workflow.py index 93a6fcc..dee3fc9 100644 --- a/examples/dataset_table_workflow/workflow.py +++ b/examples/dataset_table_workflow/workflow.py @@ -27,7 +27,6 @@ class TableDatasetWorkflow: workflow.labeled_homo_table_example(skip_test=skip_test) def cross_feat_eng_hetero_table_example(self): - set_seed(0) workflow = HeterogeneousDatasetWorkflow( benchmark_config=hetero_cross_feat_eng_benchmark_config, name="hetero", @@ -37,12 +36,11 @@ class TableDatasetWorkflow: workflow.unlabeled_hetero_table_example() def cross_task_hetero_table_example(self, skip_test=False): - set_seed(0) workflow = HeterogeneousDatasetWorkflow( benchmark_config=hetero_cross_task_benchmark_config, name="hetero", - rebuild=False, - retrain=False + rebuild=True, + retrain=True ) workflow.labeled_hetero_table_example(skip_test=skip_test)