Spaces:
Sleeping
Sleeping
File size: 11,306 Bytes
4d0da28 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 | import atexit
import logging
import os
import shutil
import dirsync
import joblib
import torch
import pytorch_lightning as pl
from pytorch_lightning import seed_everything
from torch.utils.data import DataLoader, Subset
from pricePrediction import config
from pricePrediction.ArgParser_base import ArgParseable
from pricePrediction.dataManager.dataset import Dataset_graphPrice
from pricePrediction.preprocessData.prepareDataMol2Price import DataBuilder
from pricePrediction.utils import EncodedDirNamesAndTemplates
class GraphPriceDatamodule(pl.LightningDataModule, ArgParseable):
DESIRED_PARAMS_TO_ASK= ['encodedDir', 'batch_size', 'num_workers', 'use_lds', 'use_weights', 'force_compute_data',
'do_undersample', 'augment_labels', 'copy_data_to_wdir', 'debug']
def __init__(self, encodedDir=config.ENCODED_DIR, deg_fname=None, batch_size: int = config.BATCH_SIZE,
num_workers: int = config.NUM_WORKERS_PER_GPU, random_seed: int = config.RANDOM_SEED,
use_log_for_target: bool = True, use_lds: bool = False,
use_weights: bool = False, force_compute_data: bool = False, do_undersample: bool = False,
augment_labels: bool =True, shuffle_training: bool =True, copy_data_to_wdir: str = None, debug =False):
'''The DataManager that will create several DataLoader(s)
:param str encodedDir: The directory where the dataset has been prepared as lmdb files.
:param str deg_fname: The file containing the statistics about the nodes degree. By default is automatically searched within
encodedDir.
:param int batch_size: Batch size.
:param int num_workers: Number of workers as in "pytorch.DataLoader".
:param int random_seed: Random seed.
:param bool use_log_for_target: Apply log to the target (log(price)).
:param bool use_lds: Use label density smoothing to reduce data imbalance.
:param bool use_weights: Apply weights to reduce data imbalance.
:param bool force_compute_data: Recompute all data before using the DataManager. Takes some time.
:param bool do_undersample: Use undersampling to reduce data imbalance.
:param bool augment_labels: Add random perturbation to training labels to reduce overfitting.
:param bool shuffle_training: Shuffle training data.
:param str copy_data_to_wdir: An optional working dir where data can be copied. Useful to avoid network disk usage.
:param bool debug: Use only a small subset of the dataset
'''
seed_everything(random_seed)
super().__init__()
if encodedDir is None:
encodedDir = config.ENCODED_DIR
if deg_fname is None:
names = EncodedDirNamesAndTemplates(encodedDir)
deg_fname = names.DEGREES_FNAME
self.deg_fname = deg_fname
self.do_undersample = do_undersample
self.use_weights = use_weights or use_lds
self.use_lds = use_lds
self.use_log_for_target = use_log_for_target
self.random_seed = random_seed
self.copy_data_to_wdir = copy_data_to_wdir
self.encodedDir = encodedDir
self.batch_size = batch_size
self.num_workers = num_workers
self.force_compute_data = force_compute_data
self.augment_labels = augment_labels
self.shuffle_training = shuffle_training
self.debug = debug
###### Delayed variables
self._dims = None
self._train_n_instances = None
self.setup_done = False
self.dataset_train = None
self.dataset_val = None
self.dataset_test = None
def _getBucketsFname(self, training_phase):
names = EncodedDirNamesAndTemplates(self.encodedDir)
return names.BUCKETS_FNAME_TEMPLATE % training_phase
@property
def dims(self):
if self._dims is None:
metadata_fname = EncodedDirNamesAndTemplates(self.encodedDir).DATASET_METADATA_FNAME_TEMPLATE % "train"
size_metadata_info = joblib.load(metadata_fname)
self._dims = (size_metadata_info["nodes_n_features"], size_metadata_info["edges_n_features"])
self._train_n_instances = size_metadata_info["total_size"]
return self._dims
def get_nodes_degree(self):
if os.path.isfile( self.deg_fname ):
return joblib.load(self.deg_fname)
else:
self.prepare_data()
return self.get_nodes_degree()
def prepare_data(self):
# Compute if not already done
if self.force_compute_data:
dataBuilder = DataBuilder()
dataBuilder.prepareDataset(encodedDir=self.encodedDir, datasetSplit="train")
dataBuilder.prepareDataset(encodedDir=self.encodedDir, datasetSplit="test")
dataBuilder.prepareDataset(encodedDir=self.encodedDir, datasetSplit="val")
else:
dataBuilder = DataBuilder()
if not os.path.exists(os.path.join(self.encodedDir, "train")):
dataBuilder.prepareDataset(encodedDir=self.encodedDir, datasetSplit="train")
if not os.path.exists(os.path.join(self.encodedDir, "test")):
dataBuilder.prepareDataset(encodedDir=self.encodedDir, datasetSplit="test")
if not os.path.exists(os.path.join(self.encodedDir, "val")):
dataBuilder.prepareDataset(encodedDir=self.encodedDir, datasetSplit="val")
# if self.copy_data_to_wdir:
# tmpArgs = {}
# print(self.copy_data_to_wdir, os.path.isdir(self.copy_data_to_wdir))
# if not os.path.isdir(self.copy_data_to_wdir):
# tmpArgs["prefix"] = self.copy_data_to_wdir
# self._tmpdir = TemporaryDirectory(**tmpArgs)
# def cleanTmpdir():
# del self._tmpdir
# atexit.register( cleanTmpdir)
#
# self.copy_data_to_wdir = self._tmpdir.name
if self.copy_data_to_wdir:
if not os.path.isdir(self.copy_data_to_wdir):
os.makedirs(self.copy_data_to_wdir)
_tmpdir = self.copy_data_to_wdir
def cleanTmpdir():
shutil.rmtree( _tmpdir)
atexit.register(cleanTmpdir)
print("copying training data from %s to %s" % (self.encodedDir, self.copy_data_to_wdir))
dirsync.sync(self.encodedDir, self.copy_data_to_wdir, 'sync', verbose=True, logger=logging.getLogger('dummy'))
def setup(self, stage=None):
if self.dataset_train is not None and self.dataset_val is not None and self.dataset_test is not None:
return
if self.copy_data_to_wdir is not None:
encodedDir = self.copy_data_to_wdir
else:
encodedDir = self.encodedDir
self.dataset_train = Dataset_graphPrice(encodedDir, datasetSplit="train",
use_log_for_target=self.use_log_for_target,
use_weights=self.use_weights, augment_labels = self.augment_labels,
buckets_fname=self._getBucketsFname("train"),
use_lds=self.use_lds, do_undersample=self.do_undersample)
self.dataset_val = Dataset_graphPrice(encodedDir, datasetSplit="val",
use_log_for_target=self.use_log_for_target,
use_weights=self.use_weights, augment_labels = False,
buckets_fname=self._getBucketsFname("val"),
use_lds=self.use_lds, do_undersample=False)
self.dataset_test = Dataset_graphPrice(encodedDir, datasetSplit="test",
use_log_for_target=self.use_log_for_target,
use_weights=self.use_weights, augment_labels = False,
buckets_fname=self._getBucketsFname("test"),
use_lds=self.use_lds, do_undersample=False)
atexit.register( self.dataset_train.close )
atexit.register( self.dataset_val.close )
atexit.register( self.dataset_test.close )
def _generic_dataloader(self, dataset, num_workers=None, **kwargs):
if num_workers is None:
num_workers = self.num_workers
if num_workers>0 and "timeout" not in kwargs:
kwargs["timeout"] = config.MULTIPROC_TIMEOUT
if num_workers > 0 and "prefetch_factor" not in kwargs:
kwargs["prefetch_factor"] = config.MULTIPROC_PREFETCH_FACTOR
if num_workers > 0 and "persistent_workers" not in kwargs:
kwargs["persistent_workers"] = config.PERSISTENT_WORKERS
print("Using num_workers:", num_workers)
#print(self.trainer.reload_dataloaders_every_epoch)
from torch_geometric.data import Batch
def from_list_of_examples_to_batch(examples):
graphs, ys = zip(* examples)
g_batch = Batch.from_data_list(graphs)
ys = torch.stack(ys)
return [g_batch, ys]
if self.debug:
dataset = Subset(dataset, range(21*self.batch_size))
return DataLoader(dataset, num_workers=num_workers, batch_size=self.batch_size,
pin_memory=num_workers>0, collate_fn = from_list_of_examples_to_batch,
**kwargs)
def train_dataloader(self, **kwargs):
if self.trainer:
print("Generating a train_dataloader at epoch: %d"% ( self.trainer.current_epoch ))
if "persistent_workers" not in kwargs:
kwargs["persistent_workers"] = False
return self._generic_dataloader(self.dataset_train, shuffle=True, **kwargs)
def val_dataloader(self, **kwargs):
return self._generic_dataloader(self.dataset_val, shuffle=False, **kwargs)
def test_dataloader(self, **kwargs):
return self._generic_dataloader(self.dataset_test, shuffle=False, **kwargs)
def predict_dataloader(self, dataset=None):
if dataset == "train":
raise NotImplementedError()
elif dataset == "test":
raise NotImplementedError()
elif dataset == "val":
raise NotImplementedError()
if dataset is not None: # TODO: check type
raise NotImplementedError()
else:
raise NotImplementedError()
if __name__ == "__main__":
from pricePrediction.ArgParser_base import ArgParseable, MyArgParser
parser = MyArgParser(prog=None, usage=None, description=None, )
group = parser.add_argument_group(title="data")
GraphPriceDatamodule.addParamsToArgParse(group)
cmd_args = parser.parse_args()
dm = GraphPriceDatamodule(**cmd_args.get("data", {}))
dm.setup()
allPrices = []
for batchG, batchY in dm.test_dataloader():
allPrices.extend( batchY.tolist() )
import matplotlib.pyplot as plt
plt.hist(allPrices, bins=200)
plt.show()
'''
python -m pricePrediction.dataManager.dataManager
'''
|