From cdded157c00f112476ea74aea4246935518310aa Mon Sep 17 00:00:00 2001 From: Frozenmad Date: Sun, 11 Jul 2021 10:17:57 +0000 Subject: [PATCH] add default parameters for budget of SAINT, set num_workers to 0. --- .../node_classification_sampled_trainer.py | 17 +++++++++++------ 1 file changed, 11 insertions(+), 6 deletions(-) diff --git a/autogl/module/train/node_classification_trainer/node_classification_sampled_trainer.py b/autogl/module/train/node_classification_trainer/node_classification_sampled_trainer.py index 9ccdfce..415e9bd 100644 --- a/autogl/module/train/node_classification_trainer/node_classification_sampled_trainer.py +++ b/autogl/module/train/node_classification_trainer/node_classification_sampled_trainer.py @@ -166,7 +166,7 @@ class NodeClassificationGraphSAINTTrainer(BaseNodeClassificationTrainer): self.__num_graphs_per_epoch: int = num_graphs_per_epoch " Set sampled_budget " - sampled_budget: int = kwargs.get("sampled_budget") + sampled_budget: int = kwargs.get("sampled_budget", 1e4) # todo: This is a version caused by current unreasonable initialization process # todo: Refactor the framework for trainer to fix in future version # if type(sampled_budget) != int: @@ -197,11 +197,16 @@ class NodeClassificationGraphSAINTTrainer(BaseNodeClassificationTrainer): __cpu_count: _typing.Optional[int] = os.cpu_count() return __cpu_count if __cpu_count else 0 - self.__training_sampler_num_workers: int = kwargs.get( - "training_sampler_num_workers", _cpu_count() - ) - if not 0 <= self.__training_sampler_num_workers <= _cpu_count(): - self.__training_sampler_num_workers: int = _cpu_count() + # self.__training_sampler_num_workers: int = kwargs.get( + # "training_sampler_num_workers", _cpu_count() + # ) + + # if not 0 <= self.__training_sampler_num_workers <= _cpu_count(): + # self.__training_sampler_num_workers: int = _cpu_count() + + # force to be 0 to be compactible with current pyg solution. + self.__training_sampler_num_workers: int = 0 + super(NodeClassificationGraphSAINTTrainer, self).__init__( model, num_features, num_classes, device, init, feval, loss )