# -*- conding: utf-8 -*- # @Time : 2025/12/14 10:58 # @Author : psi from utils import * from modules import * import os, sys import numpy as np from tqdm import tqdm import random import torch from torch import nn from config import CFG from dataset import * import torch.utils.data import copy, json, pickle import itertools as it import glob import torch.nn.functional as F def my_collate(batch): batch = list(filter(lambda x: (x is not None), batch)) msbinl, molfpl, molfml, vl, al, msl = [], [], [], [], [], [] bat = {} msbinl1, msbinl2 = [], [] for b in batch: if 'ms_bins' in b: msbinl.append(b['ms_bins']) if 'ms_bins1' in b: msbinl1.append(b['ms_bins1']) if 'ms_bins2' in b: msbinl2.append(b['ms_bins2']) if 'mol_fps' in b: molfpl.append(b['mol_fps']) if 'mol_fmvec' in b: molfml.append(b['mol_fmvec']) if 'V' in b: vl.append(b['V']) if 'A' in b: al.append(b['A']) if 'mol_size' in b: msl.append(b['mol_size']) if msbinl: bat['ms_bins'] = torch.stack(msbinl) if msbinl1: bat['ms_bins1'] = torch.stack(msbinl1) if msbinl2: bat['ms_bins2'] = torch.stack(msbinl2) if molfpl: bat['mol_fps'] = torch.stack(molfpl) if molfml: bat['mol_fmvec'] = torch.stack(molfml) if vl and al and msl: max_n = max(map(lambda x:x.shape[0], vl)) vl1, al1 = [], [] for v in vl: vl1.append(pad_V(v, max_n)) for a in al: al1.append(pad_A(a, max_n)) bat['V'] = torch.stack(vl1) bat['A'] = torch.stack(al1) bat['mol_size'] = torch.cat(msl, dim=0) # return torch.utils.data.dataloader.default_collate(batch) return bat def build_loaders(inp, mode, cfg, num_workers): if type(inp[0]) is dict: dataset = Dataset(inp, cfg) else: dataset = PathDataset(inp, cfg) dataloader = torch.utils.data.DataLoader( dataset, batch_size=len(dataset), num_workers=num_workers, shuffle=True if mode == "train" else False, collate_fn=my_collate ) return dataloader class Predictor(): def __init__(self, file, model_file): CFG.load(file) cfg = CFG self.cfg = cfg self.device = "cuda" if torch.cuda.is_available() else "cpu" model = FragSimiModelNew(cfg).to(cfg.device) encmodel = torch.load(model_file) # model.mol_gnn_encoder.load_state_dict(encmodel.mol_gnn_encoder.state_dict()) model.load_state_dict(encmodel['state_dict']) self.model = model self.model.eval() def process_file(self, ms): # d = json.load(open(file, 'r', encoding='utf-8')) # ms = d['ms'] smi = 'Br.C=CC1CN2CCC1CC2C(O)c1ccnc2ccc(OC)cc12' # out = {'ms': ms, 'smiles': smi} # ms = self.data[idx]['ms'] # smi = self.data[idx]['smiles'] nls = [] item = calc_feats(smi, ms, nls, self.cfg) return item def process(self, data): res = [] res.append(self.process_file(data)) batch = my_collate(res) return batch def get_eval_info(self, ms_embeddings, mol_embeddings, top_ks=(1, 3, 5, 10)): N = ms_embeddings.shape[0] # 1. L2 归一化(非常关键) # ms_norm = F.normalize(ms_embeddings, dim=1) # mol_norm = F.normalize(mol_embeddings, dim=1) ms_norm = ms_embeddings mol_norm = mol_embeddings recalls = {k: 0 for k in top_ks} # 2. 对每个样本做检索 for i in range(N): query = ms_norm[i] # (256,) sims = torch.matmul(mol_norm, query) # (N,) ranked_indices = torch.argsort(sims, descending=True) for k in top_ks: if i in ranked_indices[:k]: recalls[k] += 1 # 3. 取平均 for k in recalls: recalls[k] /= N return recalls def predict(self, ms): # data_files = [] # for root, _, files in os.walk(file_path): # for f in files: # if f.endswith(('.json', '.pkl', '.mgf')): # data_files.append(os.path.join(root, f)) # data = sorted( # data_files, # key=lambda x: int(os.path.splitext(os.path.basename(x))[0]) # ) batch = self.process(ms) for k, v in batch.items(): batch[k] = v.to(self.cfg.device) with torch.no_grad(): loss, loss_infonce, loss_mse, ms_embeddings, mol_embeddings = self.model(batch, is_predict=True) # recalls_info = self.get_eval_info(ms_embeddings, mol_embeddings) # print(recalls_info) # # return loss, loss_infonce, loss_mse, recalls_info return ms_embeddings def process_file_1(self, ms, smi): # d = json.load(open(file, 'r', encoding='utf-8')) # ms = d['ms'] # smi = d['smiles'] # out = {'ms': ms, 'smiles': smi} # ms = self.data[idx]['ms'] # smi = self.data[idx]['smiles'] nls = [] item = calc_feats(smi, ms, nls, self.cfg) return item def process_1(self, data): res = [] for d in data: try: res.append([d['smiles'], self.process_file_1(d['ms'], d['smiles'])]) except Exception as e: print(e) if len(res) == 0: return None res1 = [x[1] for x in res] res2 = [x[0] for x in res] batch = my_collate(res1) return batch, res2 def get_eval_info_1(self, ms_embeddings, mol_embeddings, top_ks=(1, 3, 5, 10)): N = ms_embeddings.shape[0] # 1. L2 归一化(非常关键) # ms_norm = F.normalize(ms_embeddings, dim=1) # mol_norm = F.normalize(mol_embeddings, dim=1) ms_norm = ms_embeddings mol_norm = mol_embeddings recalls = {k: 0 for k in top_ks} # 2. 对每个样本做检索 for i in range(N): query = ms_norm[i] # (256,) sims = torch.matmul(mol_norm, query) # (N,) ranked_indices = torch.argsort(sims, descending=True) for k in top_ks: if i in ranked_indices[:k]: recalls[k] += 1 # 3. 取平均 for k in recalls: recalls[k] /= N return recalls def topk_similarity_1(self, ms_embedding, res_embeddings, batch_size=128, top_k=10): ms_embedding = ms_embedding.to(self.device) res_embeddings = res_embeddings.to(self.device) # 保证是 float tensor ms_embedding = ms_embedding.float() res_embeddings = res_embeddings.float() # 归一化(用于余弦相似度) # ms_embedding = F.normalize(ms_embedding, dim=1) # res_embeddings = F.normalize(res_embeddings, dim=1) similarities = [] # 分 batch 计算 for i in range(0, res_embeddings.size(0), batch_size): batch = res_embeddings[i:i + batch_size] # (B, 256) # (1, 256) @ (256, B) -> (1, B) sim = torch.matmul(ms_embedding, batch.T) # 余弦相似度 similarities.append(sim.squeeze(0)) # (B,) # 拼接成 (15511,) similarities = torch.cat(similarities, dim=0) top_k = min(top_k, res_embeddings.size(0)) # 取 top_k topk_sim, topk_idx = torch.topk(similarities, k=top_k) return topk_sim, topk_idx def predict_1(self, data): # data_files = [] # for root, _, files in os.walk(file_path): # for f in files: # if f.endswith(('.json', '.pkl', '.mgf')): # data_files.append(os.path.join(root, f)) # data = sorted( # data_files, # key=lambda x: int(os.path.splitext(os.path.basename(x))[0]) # ) batch, res_smiles = self.process_1(data) if batch is None: return None for k, v in batch.items(): batch[k] = v.to(self.cfg.device) loss, loss_infonce, loss_mse, ms_embeddings, mol_embeddings = self.model(batch, is_predict=True) # recalls_info = self.get_eval_info(ms_embeddings, mol_embeddings) # print(recalls_info) # # return loss, loss_infonce, loss_mse, recalls_info # ms_embedding = ms_embeddings.to(self.device) # res_embeddings = mol_embeddings.to(self.device) ms_embedding = ms_embeddings[0:1, :] topk_sim, topk_idx = self.topk_similarity_1(ms_embedding, mol_embeddings) topk_idx = topk_idx.to("cpu").numpy().tolist() res_pred_name = [] for x, i in enumerate(topk_idx): res_pred_name.append([res_smiles[i], topk_sim[x].item()]) return res_pred_name class InferOnline(): def __init__(self, model_file): pred = Predictor('config.json', model_file) self.pred_model = pred model_name = model_file.split('/')[-1][:-4] emb_file = f'/dev/shm/data/tongji_data/all_pos_pred_emb_{model_name}.pt' file_path = '/dev/shm/data/tongji_data/all_pos.json' self.device = "cuda" if torch.cuda.is_available() else "cpu" self.emb_data= torch.load(emb_file) print("self.emb_data shape ,,,", self.emb_data.shape) self.all_data = json.load(open(file_path, 'r', encoding='utf-8')) print("self.all_data shape ...", len(self.all_data)) pred_file = "/dev/shm/data/tongji_data/all_pos_pred.pt" self.pred_name_data = [x[0] for x in torch.load(pred_file)] print("self.pred_name_data shape ,,,", len(self.pred_name_data)) self.pred_smiles_name_2_id = {x: i for i, x in enumerate(self.pred_name_data)} print("self.pred_smiles_name_2_id shape ,,,", len(self.pred_smiles_name_2_id)) all_data_res3 = {} for k, v in self.all_data['res3'].items(): v1 = [x for x in v if x in self.pred_name_data] if len(v1) > 0: all_data_res3[k] = v1 self.all_data['res3'] = all_data_res3 def select_smiles(self, parent_mz, bn=50): res_smiles = [] for k, v in self.all_data['res3'].items(): k = float(k) if k > parent_mz - bn and k < parent_mz + bn: res_smiles += v res_embeddings = [] for x in res_smiles: i = self.pred_smiles_name_2_id[x] res_embeddings.append(self.emb_data[i, :].unsqueeze(0)) res_embeddings = torch.cat(res_embeddings, dim=0) return res_embeddings, res_smiles def topk_similarity(self, ms_embedding, res_embeddings, batch_size=128, top_k=10): ms_embedding = ms_embedding.to(self.device) res_embeddings = res_embeddings.to(self.device) # 保证是 float tensor ms_embedding = ms_embedding.float() res_embeddings = res_embeddings.float() # 归一化(用于余弦相似度) # ms_embedding = F.normalize(ms_embedding, dim=1) # res_embeddings = F.normalize(res_embeddings, dim=1) similarities = [] # 分 batch 计算 for i in range(0, res_embeddings.size(0), batch_size): batch = res_embeddings[i:i + batch_size] # (B, 256) # (1, 256) @ (256, B) -> (1, B) sim = torch.matmul(ms_embedding, batch.T) # 余弦相似度 similarities.append(sim.squeeze(0)) # (B,) # 拼接成 (15511,) similarities = torch.cat(similarities, dim=0) top_k = min(top_k, res_embeddings.size(0)) # 取 top_k topk_sim, topk_idx = torch.topk(similarities, k=top_k) return topk_sim, topk_idx def infer(self, ms, parent_mz, bn=50): ms_embeddings = self.pred_model.predict(ms) # (1, 256) res_embeddings, res_smiles = self.select_smiles(parent_mz, bn=bn) if len(res_smiles) == 0: return [] topk_sim, topk_idx = self.topk_similarity(ms_embeddings, res_embeddings) topk_idx = topk_idx.to("cpu").numpy().tolist() res_pred_name = [] for x, i in enumerate(topk_idx): res_pred_name.append([res_smiles[i], topk_sim[x].item()]) return res_pred_name model_file = ["model-tloss3.437-vloss2.907-epoch0.pth", "model-tloss2.495-vloss2.253-epoch1.pth", "model-tloss1.987-vloss1.866-epoch2.pth", "model-tloss1.597-vloss1.573-epoch3.pth", "model-tloss1.332-vloss1.384-epoch4.pth", "model-tloss1.088-vloss1.255-epoch5.pth", "model-tloss0.899-vloss1.068-epoch6.pth", "/root/代码/out_data/train-018/model-tloss0.76-vloss0.986-epoch0.pth", "/root/代码/out_data/train-019/model-tloss0.608-vloss0.943-epoch0.pth", "/root/代码/out_data/train-019/model-tloss0.577-vloss0.852-epoch1.pth", "/root/代码/out_data/train-019/model-tloss0.503-vloss0.811-epoch2.pth", '/root/代码/out_data/train-019/model-tloss0.448-vloss0.763-epoch3.pth', '/root/代码/out_data/train-019/model-tloss0.405-vloss0.74-epoch4.pth', '/root/代码/out_data/train-019/model-tloss0.367-vloss0.722-epoch5.pth', '/root/代码/out_data/train-019/model-tloss0.337-vloss0.705-epoch6.pth', '/root/代码/out_data/train-019/model-tloss0.317-vloss0.671-epoch7.pth'][-1] infer_online = InferOnline(model_file) if __name__ == '__main__': # pred = Predictor('config.json', model_file) # # file_path = ["../../data/CASMI2016/data/pos", "../../data/CASMI2017/data/pos"][1] # loss, loss_infonce, loss_mse, recalls_info = pred.predict(file_path) # # print("loss ...", loss) # print("loss_infonce ...", loss_infonce) # print("loss_mse ...", loss_mse) # print("recalls_info ...", recalls_info) ms = [[41.038587, 880600.0], [42.033833, 1973400.0], [43.041651, 2117400.0], [44.049388, 925150.0], [44.979347, 4397200.0], [51.022884, 593400.0], [53.038537, 9694400.0], [54.033783, 415000.0], [55.054152, 1911200.0], [56.049325, 4400500.0], [56.979301, 487200.0], [65.038474, 449400.0], [67.041567, 1667200.0], [67.054107, 786000.0], [68.049268, 6593250.0], [68.979253, 836000.0], [69.056975, 17628050.0], [69.069744, 290600.0], [70.06482, 1716900.0], [70.994926, 276400.0], [73.010627, 215600.0], [77.038437, 417800.0], [79.054112, 1319600.0], [80.049378, 3584000.0], [81.057194, 2957200.0], [82.064812, 25688800.0], [82.070396, 669000.0], [82.073249, 528650.0], [82.994909, 3564600.0], [83.072653, 7343000.0], [84.080502, 821400.0], [94.065026, 1006000.0], [95.049053, 230800.0], [96.080647, 13938800.0], [97.010531, 20776600.0], [97.013298, 339600.0], [98.989853, 367600.0], [110.096244, 1418800.0], [110.989833, 48727600.0], [110.991981, 1024000.0], [110.994248, 515600.0], [111.001067, 985400.0], [111.103933, 13806250.0], [112.111972, 17873400.0], [112.114998, 263400.0], [115.054168, 518400.0], [117.069717, 320200.0], [134.018401, 474400.0], [194.099834, 1533000.0], [194.993274, 21076400.0], [306.09803, 51809350.0], [306.181335, 516550.0]] parent_mz = 306.181335 res_pred_name = infer_online.infer(ms, parent_mz) print(res_pred_name) print(len(res_pred_name))