Spaces:
Sleeping
Sleeping
File size: 7,922 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 | import gzip
import os
import re
import sys
import time
from functools import reduce
from itertools import chain
from multiprocessing import cpu_count
import lmdb
import psutil
import joblib
import torch
from joblib import Parallel, delayed
import numpy as np
from pricePrediction import config
from pricePrediction.config import USE_MMOL_INSTEAD_GRAM
from pricePrediction.preprocessData.serializeDatapoints import getExampleId, serializeExample
from pricePrediction.utils import tryMakedir, getBucketRanges, search_buckedId, EncodedDirNamesAndTemplates
from .smilesToGraph import compute_nodes_degree, fromPerGramToPerMMolPrice
if config.USE_FEATURES_NET:
from .smilesToDescriptors import smiles_to_graph
else:
from .smilesToGraph import smiles_to_graph
PER_WORKER_MEMORY_GB = 2
class DataBuilder():
def __init__(self, n_cpus= config.N_CPUS):
if n_cpus is None:
mem_gib = psutil.virtual_memory().available / (1024. ** 3)
n_cpus = int(max(1, min(cpu_count(), mem_gib // PER_WORKER_MEMORY_GB)))
self.n_cpus = n_cpus
def processOneFileOfSmiles(self, encodedDir, fileNum, datasetSplit, fname, nrows=None):
print("processing %s"%fname)
if fname.endswith(".csv"):
open_fun = open
decode_line = lambda line: line
elif fname.endswith(".csv.gz"):
open_fun = gzip.open
decode_line = lambda line : line.decode('utf-8')
else:
raise ValueError("Bad file format")
cols = ['SMILES','price']
names = EncodedDirNamesAndTemplates(encodedDir)
outFname_base = datasetSplit + "_" + str(fileNum) + "_lmdb"
outFname = os.path.join(encodedDir, datasetSplit,outFname_base)
env = lmdb.open(outFname, map_size=10737418240)
one_example_graph = smiles_to_graph("CCOCCC")
degs = compute_nodes_degree(None)
min_in_x = torch.ones(one_example_graph.x.shape[-1])*float('inf')
max_in_x = -torch.ones(one_example_graph.x.shape[-1])*float('inf')
num_examples = 0
bucket_ranges = getBucketRanges()
n_per_bucket = np.zeros(len(bucket_ranges), dtype= np.int64)
with open_fun(fname) as f_in:
header = decode_line(f_in.readline()).strip().split(",")
try:
smi_index, price_index = [ header.index(col) for col in cols]
except ValueError:
smi_index, price_index = 0,1
with env.begin(write=True) as sink, \
gzip.open(names.SELECTED_DATAPOINTS_TEMPLATE % (datasetSplit, fileNum), "wt") as f_out:
f_out.write("SMILES,price\n")
cur_time = time.time()
for i, line in enumerate(f_in):
lineArray = decode_line(line).strip().split(",")
smi, price = lineArray[smi_index], lineArray[price_index]
price = float(price)
graph = smiles_to_graph(smi)
if graph is None:
continue
#Save the original smiles-price
f_out.write("%s,%s\n"%(smi, price))
# Use the per mmol price
if USE_MMOL_INSTEAD_GRAM:
price = fromPerGramToPerMMolPrice(price, smi)
bucketId = search_buckedId( np.log(price), bucket_ranges)
n_per_bucket[bucketId] += 1
degs += compute_nodes_degree([graph])
min_in_x = torch.stack([min_in_x, torch.min(graph.x, 0)[0]]).min(0)[0]
max_in_x = torch.stack([max_in_x, torch.max(graph.x, 0)[0]]).max(0)[0]
fileId= getExampleId(outFname_base, num_examples)
sink.put(fileId, serializeExample(price, graph))
num_examples += 1
if nrows is not None and num_examples > nrows:
break
if i % 10000 == 0 and fileNum % self.n_cpus == 0:
new_time = time.time()
print("Current iteration: %d # task: %d (%.2f s) " % (i, fileNum, new_time - cur_time), end="\r")
cur_time = new_time
if fileNum % self.n_cpus == 0:
print()
return ((outFname, degs, min_in_x, max_in_x, num_examples, n_per_bucket),)
def getNFeatures(self):
one_graph = smiles_to_graph("CCCCCCO")
# print(one_graph)
return dict(nodes_n_features=one_graph["x"].shape[-1], edges_n_features=one_graph["edge_attr"].shape[-1])
def prepareDataset(self, inputDir=config.DATASET_DIRNAME, encodedDir=config.ENCODED_DIR, datasetSplit="train",
nrows=None, **kwargs):
assert datasetSplit in ["train", "val", "test"]
print("Computing %s dataset" % datasetSplit)
print("Using %d workers for data preparation"%self.n_cpus)
# os.environ["OMP_NUM_THREADS"] = "1"
# os.environ["MKL_NUM_THREADS"] = "1"
names = EncodedDirNamesAndTemplates(encodedDir)
tryMakedir(encodedDir, remove=False)
tryMakedir(os.path.join(encodedDir, datasetSplit))
tryMakedir(names.DIR_RAW_DATA_SELECTED)
fnames = [os.path.join(inputDir, fname) for fname in os.listdir(inputDir) if
re.match(config.RAW_DATA_FILE_SUFFIX, fname) and datasetSplit in fname]
assert len(fnames) > 0
results = Parallel(n_jobs=self.n_cpus, batch_size=1,
verbose=10)(delayed(self.processOneFileOfSmiles)(encodedDir, i, datasetSplit, fname, nrows=nrows)
for i, fname in enumerate(fnames))
results = chain.from_iterable(results)
results = list(results)
# print( results )
fnames_list, degrees, min_in_x, max_in_x, sizes_lis, n_per_bucket = zip(*results)
degrees = reduce(lambda prev, x: prev + x, degrees).numpy().tolist()
min_in_x = reduce(lambda prev, x: torch.stack([prev, x]).min(0)[0], min_in_x).numpy().tolist()
max_in_x = reduce(lambda prev, x: torch.stack([prev, x]).max(0)[0], max_in_x).numpy().tolist()
n_per_bucket = reduce(lambda prev, x: prev + x, n_per_bucket)
metadata_dict = {"name":datasetSplit, "fnames_list":fnames_list, "sizes_list": sizes_lis,
"total_size": sum(sizes_lis), "feature_stats":{"min":min_in_x, "max":max_in_x}}
metadata_dict.update(self.getNFeatures())
# print(metadata_dict)
joblib.dump(metadata_dict,
names.DATASET_METADATA_FNAME_TEMPLATE % datasetSplit)
joblib.dump(degrees, names.DEGREES_FNAME)
# print(n_per_bucket)
joblib.dump({"bucket_ranges": getBucketRanges(), "n_per_bucket":n_per_bucket}, names.BUCKETS_FNAME_TEMPLATE % datasetSplit)
print("Dataset %s computed" % datasetSplit)
return fnames_list
if __name__ == "__main__":
print( " ".join(sys.argv))
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("-i", "--inputDir", type=str, default=config.DATASET_DIRNAME, help="Directory where smiles-price pairs are located")
parser.add_argument("-o", "--encodedDir", type=str, default=config.ENCODED_DIR)
parser.add_argument("-n", "--ncpus", type=int, default=config.N_CPUS)
parser.add_argument("--nrows", type=int, default=None, help="The number of rows to process in each file")
args = vars( parser.parse_args())
config.N_CPUS = args.get("ncpus", config.N_CPUS)
dataBuilder = DataBuilder(n_cpus=config.N_CPUS)
dataBuilder.prepareDataset(datasetSplit="train", **args)
dataBuilder.prepareDataset(datasetSplit="val", **args)
dataBuilder.prepareDataset(datasetSplit="test", **args)
'''
python -m pricePrediction.preprocessData.prepareDataMol2Price
''' |