MS2-SMILES-AlignNet / embedding_predict222.py
monaaaaaa's picture
Upload 26 files
eeabcff verified
Raw
History Blame Contribute Delete
7.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
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, smi):
# d = json.load(open(file, 'r', encoding='utf-8'))
# ms = d['ms']
# smi = d['smiles']
ms = [[41.998214, 491000.0], [107.049799, 633000.0], [131.049708, 1969500.0], [134.040609, 827000.0], [145.052688, 300000.0], [161.051274, 270500.0], [309.102783, 2374000.0]]
# 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, file):
if isinstance(file, str):
res = json.load(open(file, 'r', encoding='utf-8'))
data = res['smiles']
res = []
for d in tqdm(data, desc='process ...'):
try:
res.append([d, self.process_file(d)])
except:
res.append([d, None])
o_file = '/dev/shm/data/tongji_data/all_neg_pred.pt'
torch.save(res, o_file)
os._exit(0)
else:
res = file
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, file_path):
# 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(file_path)
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 mol_embeddings
if __name__ == '__main__':
model_file = ["/root/代码/out_data/train-020/model-tloss3.239-vloss2.752-epoch0.pth",
"/root/代码/out_data/train-020/model-tloss2.487-vloss2.279-epoch1.pth",
"/root/代码/out_data/train-020/model-tloss2.086-vloss1.924-epoch2.pth",
"/root/代码/out_data/train-020/model-tloss1.716-vloss1.609-epoch3.pth",
"/root/代码/out_data/train-020/model-tloss1.414-vloss1.361-epoch4.pth",
'/root/代码/out_data/train-020/model-tloss1.177-vloss1.173-epoch5.pth',
"/root/代码/out_data/train-020/model-tloss0.99-vloss1.036-epoch6.pth",
'/root/代码/out_data/train-020/model-tloss0.849-vloss0.929-epoch7.pth',
'/root/代码/out_data/train-020/model-tloss0.736-vloss0.851-epoch8.pth',
'/root/代码/out_data/train-020/model-tloss0.647-vloss0.779-epoch9.pth',
'/root/代码/out_data/train-020/model-tloss0.575-vloss0.728-epoch10.pth'][-1]
model_name = model_file.split('/')[-1][:-4]
pred = Predictor('config.json', model_file)
# file_path = '/dev/shm/data/tongji_data/all_neg.json'
# loss, loss_infonce, loss_mse, recalls_info = pred.predict(file_path)
o_file = '/dev/shm/data/tongji_data/all_neg_pred.pt'
data = torch.load(o_file)
batch_size = 128
res = []
for i in tqdm(range(0, len(data), batch_size), desc='predict ...'):
batch = data[i:i + batch_size]
p = [x[1] for x in batch]
mol_embeddings = pred.predict(p)
mol_embeddings = mol_embeddings.to("cpu")
res.append(mol_embeddings)
result = torch.cat(res, dim=0)
print(f"len is : {len(data)} ...")
print(f"result shape is : {result.shape} ...")
out_file = f'/dev/shm/data/tongji_data/all_neg_pred_emb_{model_name}.pt'
torch.save(result, out_file)