From 6726cd5364a319b28f54a6ecfabfb0f86e1bf9bc Mon Sep 17 00:00:00 2001 From: SwiftieH Date: Thu, 25 Feb 2021 08:44:37 +0000 Subject: [PATCH] fixed dayaset with no masks and gbm config problem --- autogl/datasets/utils.py | 21 ++++++++++++++++++--- autogl/module/ensemble/stacking.py | 2 +- 2 files changed, 19 insertions(+), 4 deletions(-) diff --git a/autogl/datasets/utils.py b/autogl/datasets/utils.py index 8b677ee..46b68ae 100644 --- a/autogl/datasets/utils.py +++ b/autogl/datasets/utils.py @@ -61,8 +61,15 @@ def random_splits_mask(dataset, train_ratio=0.2, val_ratio=0.4, seed=None): torch.set_rng_state(r_s) if torch.cuda.is_available(): torch.cuda.set_rng_state(r_s_cuda) - - dataset.data, dataset.slices = dataset.collate([d for d in dataset]) + datalist = [] + for d in dataset: + setattr(d, "train_mask", data.train_mask) + setattr(d, "val_mask", data.val_mask) + setattr(d, "test_mask", data.test_mask) + datalist.append(d) + dataset.data, dataset.slices = dataset.collate(datalist) + if hasattr(dataset, '__data_list__'): + delattr(dataset, '__data_list__') # while type(dataset.data.num_nodes) == list: # dataset.data.num_nodes = dataset.data.num_nodes[0] # dataset.data.num_nodes = dataset.data.num_nodes[0] @@ -160,7 +167,15 @@ def random_splits_mask_class( if torch.cuda.is_available(): torch.cuda.set_rng_state(r_s_cuda) - dataset.data, dataset.slices = dataset.collate([d for d in dataset]) + datalist = [] + for d in dataset: + setattr(d, "train_mask", data.train_mask) + setattr(d, "val_mask", data.val_mask) + setattr(d, "test_mask", data.test_mask) + datalist.append(d) + dataset.data, dataset.slices = dataset.collate(datalist) + if hasattr(dataset, '__data_list__'): + delattr(dataset, '__data_list__') # while type(dataset.data.num_nodes) == list: # dataset.data.num_nodes = dataset.data.num_nodes[0] # dataset.data.num_nodes = dataset.data.num_nodes[0] diff --git a/autogl/module/ensemble/stacking.py b/autogl/module/ensemble/stacking.py index c29f849..08337b2 100644 --- a/autogl/module/ensemble/stacking.py +++ b/autogl/module/ensemble/stacking.py @@ -100,7 +100,7 @@ class Stacking(BaseEnsembler): torch.tensor(predictions).transpose(0, 1).flatten(start_dim=1).numpy() ) meta_Y = np.array(label) - + config = {} model = GradientBoostingClassifier(**config) model.fit(meta_X, meta_Y)