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)