andy88836's picture
Deploy MOFScreen-Agent FastAPI backend
4d0da28 verified
Raw
History Blame Contribute Delete
6.61 kB
import operator
import os
import joblib
import lmdb
import numpy as np
import torch
from torch.utils.data import Dataset
from pricePrediction import config
from pricePrediction.preprocessData.serializeDatapoints import getExampleId, deserializeExample
from pricePrediction.utils import search_buckedId, EncodedDirNamesAndTemplates
class Dataset_graphPrice(Dataset):
def __init__(self, encodedDir, datasetSplit, use_log_for_target=True, use_weights=False, buckets_fname=None,
use_lds=False, do_undersample=False, augment_labels=False, store_in_memory=config.DATASET_IN_MEMORY,
normalize_node_feats=config.NORMALIZE_NODE_FEATS):
self.use_log_for_target = use_log_for_target
self.use_weights = use_weights
self.buckets_fname = buckets_fname
self.use_lds = use_lds
self.do_undersample = do_undersample
self.augment_labels = augment_labels
self.store_in_memory = store_in_memory
self.normalize_node_feats = normalize_node_feats
assert (self.buckets_fname is not None and self.use_weights) or not self.use_weights
if self.do_undersample: raise NotImplementedError()
dataNames = EncodedDirNamesAndTemplates(encodedDir)
size_metadata_info = joblib.load( dataNames.DATASET_METADATA_FNAME_TEMPLATE % datasetSplit)
if self.normalize_node_feats:
self.min_x = torch.FloatTensor(size_metadata_info["feature_stats"]["min"])
self.max_x = torch.FloatTensor(size_metadata_info["feature_stats"]["max"])
self.range_x = self.max_x - self.min_x
assert not torch.isclose(self.range_x, torch.zeros(1)).any(), f"Error, range close to 0, {torch.where(torch.isclose(self.range_x, torch.zeros(1)))}"
assert not torch.isinf(self.range_x).any()
self.total_size = size_metadata_info["total_size"]
self.size_per_file = size_metadata_info["sizes_list"]
self.db_fnames = [ os.path.join( encodedDir, datasetSplit, os.path.basename(fname))
for fname in size_metadata_info["fnames_list"]]
self.cum_size_at_file = np.cumsum(np.array(size_metadata_info["sizes_list"], dtype=int))
self.last_cum_elemIdx_at_file = self.cum_size_at_file - 1
self.fileIdx_to_fileBasename = [os.path.basename(fname) for fname in self.db_fnames]
#Delayed load
self.db_managers= [None]*len(self.size_per_file)
self._bucket_ranges = None
self._n_per_bucket = None
self._cache = {}
def _idx_to_file_and_num(self, idx):
fileIdx = search_buckedId(idx, self.last_cum_elemIdx_at_file)
idx = idx - (self.cum_size_at_file[fileIdx - 1] if fileIdx - 1 >= 0 else 0)
return fileIdx, idx
def _init_db(self, fileIdx):
env = lmdb.open(self.db_fnames[fileIdx],
readonly=True, lock=False,
readahead=False, meminit=False)
txn = env.begin()
self.db_managers[fileIdx] = (env, txn)
return env, txn
@property
def bucket_ranges(self):
if self._bucket_ranges is None:
self._getBuckets()
return self._bucket_ranges
@property
def n_per_bucket(self):
if self._n_per_bucket is None:
self._getBuckets()
return self._n_per_bucket
def n_batches(self, batch_size):
return len(self)//batch_size + int(bool(len(self)%batch_size))
def _getBuckets(self):
buckets = joblib.load(self.buckets_fname)
bucket_ranges, n_per_bucket = buckets["bucket_ranges"], buckets["n_per_bucket"]
if self.use_lds:
print("using lds")
from scipy.ndimage import gaussian_filter1d
n_per_bucket = gaussian_filter1d(n_per_bucket, sigma=2)
self._bucket_ranges = bucket_ranges
self._n_per_bucket = n_per_bucket
def compute_weight(self, label):
if not self.use_log_for_target:
label = np.exp(label)
bucket_num = search_buckedId(label, self.bucket_ranges)
w = self.total_size / (1 + self.n_per_bucket[bucket_num])
return w
def prepare_undersample(self, maximum_percentile=80):
maximum_per_bucket = np.percentile(self.n_per_bucket, maximum_percentile)
prob_per_bucket = maximum_per_bucket / (1e-10 + self.n_per_bucket)
prob_per_bucket = np.where(prob_per_bucket < 1, prob_per_bucket, 1)
print("maximum_per_bucket while sampling: %d" % maximum_per_bucket)
new_len = sum([min(maximum_per_bucket, n_elems) for n_elems in self.n_per_bucket])
def predicate(graph_y):
y = graph_y[1]
bucket_num = search_buckedId(y, self.bucket_ranges)
return np.random.rand() < prob_per_bucket[bucket_num]
# dataset = dataset.select(predicate)
raise NotImplementedError() #This code is obsolete and requires refactoring
def normalize_x(self, x):
return (x-self.min_x)/self.range_x
def __getitem__(self, index):
out = self._cache.get(index, None)
if out is not None:
return out
else:
if index >= self.total_size:
raise IndexError()
fileIdx, pointIdx = self._idx_to_file_and_num(index)
db_manager = self.db_managers[fileIdx]
if db_manager is None:
txn= self._init_db(fileIdx)[-1]
else:
txn = db_manager[-1]
dataBytes = txn.get(getExampleId(self.fileIdx_to_fileBasename[fileIdx], pointIdx))
if dataBytes is None:
msg = "Error:\n%s %s %s %s"%(index, fileIdx, pointIdx, getExampleId(self.fileIdx_to_fileBasename[fileIdx], pointIdx))
raise Exception(msg)
graph, label = deserializeExample(dataBytes)
if self.normalize_node_feats:
graph.x = self.normalize_x(graph.x)
if self.use_log_for_target:
label = np.log(label)
if self.augment_labels:
label = label * (1+0.05*np.random.rand()-0.025)
if self.use_weights:
graph.w = self.compute_weight(label)
else:
graph.w = 1
graph.y = label
if self.store_in_memory:
self._cache[index] = (graph, label)
return graph, label
def __len__(self):
return self.total_size
def close(self):
for db_handler in self.db_managers:
if db_handler:
env, txn = db_handler
env.close()