MS2-SMILES-AlignNet / infer_predict111.py
monaaaaaa's picture
Upload 26 files
eeabcff verified
Raw
History Blame Contribute Delete
15.4 kB
# -*- 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))