MS2-SMILES-AlignNet / dataset.py
monaaaaaa's picture
Upload 26 files
eeabcff verified
Raw
History Blame Contribute Delete
5.78 kB
import os, json
import torch
import utils
import numpy as np
def calc_feats(smi, ms, nls, cfg, ms_bins3=None, precursor_mz=490.28+18.01, diagnostic_ions=[102.05, 135.08],
neutral_losses=[18.01]):
item = {}
item['ms_bins'] = utils.ms_binner(ms, nls,
min_mz=cfg.min_mz,
max_mz=cfg.max_mz,
bin_size=cfg.bin_size,
add_nl=cfg.add_nl,
binary_intn=cfg.binary_intn)
# precursor = 490.28 + 18.01 # 模拟丢失水
precursor = precursor_mz
meta = np.zeros(25)
meta[0] = 1 # 假设第一位是 Orbitrap
item['ms_bins1'], item['ms_bins2'] = utils.ms_feature_processor(ms,
precursor,
meta,
diagnostic_ions=diagnostic_ions,
neutral_losses=neutral_losses)
item['ms_bins3'] = torch.tensor(ms_bins3) if ms_bins3 is not None else None
fmcalced = False
if 'fp' in cfg.mol_encoder:
if not 'fm' in cfg.mol_encoder:
item['mol_fps'] = utils.mol_fp_encoder(smi,
tp=cfg.fptype,
nbits=cfg.mol_embedding_dim)
else:
item['mol_fps'], item['mol_fmvec'] = utils.mol_fp_fm_encoder(smi,
tp=cfg.fptype,
nbits=cfg.mol_embedding_dim)
fmcalced = True
if 'gnn' in cfg.mol_encoder:
f = utils.mol_graph_featurizer(smi)
if not f:
return None
item.update(f)
if 'fm' in cfg.mol_encoder and not fmcalced:
item['mol_fmvec'] = utils.smi2fmvec(smi)
return item
class Dataset(torch.utils.data.Dataset):
def __init__(self, inp, cfg):
if type(inp) is str:
self.data = json.load(open(inp))
else:
self.data = inp
self.cfg = cfg
def __getitem__(self, idx):
d = self.data[idx]
if "ms_bins" in d:
return d
else:
item = {}
try:
if 'nls' in self.data[idx]:
nls = self.data[idx]['nls']
else:
nls = []
ms = self.data[idx]['ms']
smi = self.data[idx]['smiles']
item = calc_feats(smi, ms, nls, self.cfg)
except Exception as e:
print('='*50, idx, str(e))
return None
return item
def __len__(self):
return len(self.data)
class DatasetGNNFP(torch.utils.data.Dataset):
def __init__(self, inp, cfg):
if type(inp) is str:
self.data = json.load(open(inp))
else:
self.data = inp
self.cfg = cfg
def __getitem__(self, idx):
try:
smi = self.data[idx]['smiles']
item = {}
item['mol_fps'] = utils.mol_fp_encoder(smi,
tp=self.cfg.fptype,
nbits=self.cfg.mol_embedding_dim)
item.update(utils.mol_graph_featurizer(smi))
except Exception as e:
print('='*50, idx, str(e))
return None
return item
def __len__(self):
return len(self.data)
class PathDataset(torch.utils.data.Dataset):
def __init__(self, pathlist, cfg):
self.fns = pathlist
self.cfg = cfg
self.data = {}
def __getitem__(self, idx):
fn = self.fns[idx]
if fn.endswith(".pt"):
item = torch.load(fn)
return item
else:
try:
item = {}
nls = []
if not idx in self.data:
out = self.proc_data(self.fns[idx], self.cfg.energy)
if out is None:
return None
self.data[idx] = out
ms = self.data[idx]['ms']
smi = self.data[idx]['smiles']
item = calc_feats(smi, ms, nls, self.cfg)
except Exception as e:
print('='*50, idx, str(e))
return None
return item
def proc_data(self, fn, energy='Energy1'):
if fn.endswith('.json'):
d = json.load(open(fn, 'r', encoding='utf-8'))
l = d['ms']
smi = d['smiles']
out = {'ms': l, 'smiles': smi}
return out
else:
tl = open(fn).readlines()
l = []
try:
flag = False
for i in tl:
if energy in i:
smi = i.split(';')[-2]
flag = True
continue
if 'END IONS' in i:
if flag:
break
if flag:
mz, intn = i.split(' ')
l.append((float(mz), float(intn)))
except:
return None
out = {'ms': l, 'smiles': smi}
return out
def __len__(self):
return len(self.fns)