# -*- conding: utf-8 -*- # @Time : 2025/12/22 18:48 # @Author : psi # -*- coding: utf-8 -*- # @Time : 2025/12/22 # @Author : psi import json import time import gradio as gr import pandas as pd from rdkit import Chem from rdkit.Chem import Draw import os from PIL import Image import base64 from smiles_to_pubchem_2d_image import smiles_to_pubchem_2d_image flag_t = True if flag_t: from infer_predict111 import infer_online as pos_infer_online # from infer_predict222 import infer_online as neg_infer_online # from a_predict111 import pred as pos_pred # from a_predict222 import pred as neg_pred base_image_path = "/Users/xiaojie/Documents/lunwen/代码/data/" logo_image_local_path = "logo_1.png" # ========================= # 系统 SMILES 库(示例) # ========================= SYSTEM_SMILES_DB = [ 'Br.C=CC1CN2CCC1CC2C(O)c1ccnc2ccc(OC)cc12', "CCOC(=O)c1ccc(CCN)cc1", "CN1CCN(CC1)C2=CC=CC=C2", "COC1=CC=CC=C1C(=O)O", "CCN(CC)C(=O)C1=CC=CC=C1", "CCC1=CC=CC=C1O", 'Br.C=CC1CN2CCC1CC2C(O)c1ccnc2ccc(OC)cc12', ] def infer_ms(msp_file, ion_mode, parent_mass, parent_mass_bn=50): if ion_mode == "pos": return pos_infer_online.infer(msp_file, parent_mass, parent_mass_bn) else: return neg_infer_online.infer(msp_file, parent_mass, parent_mass_bn) # ========================= # Mock Cross-modal Retrieval # ========================= def cross_modal_retrieval(parent_mass, ion_mode, msp_file, user_smiles, parent_mass_bn=50): """ 真实版本中替换为: - MS2 embedding - SMILES embedding - cosine similarity top-k """ if not flag_t: candidates = SYSTEM_SMILES_DB + user_smiles return candidates[:8] # 故意 <10,用于测试留空逻辑 res_pred_name1 = infer_ms(msp_file, ion_mode, parent_mass, parent_mass_bn) if len(user_smiles) == 0: return res_pred_name1 data = [] for x in user_smiles: data.append({'ms': msp_file, "smiles": x}) if ion_mode == "pos": # results = pos_pred.predict(data) results = pos_infer_online.pred_model.predict_1(data) else: # results = neg_pred.predict(data) results = neg_infer_online.pred_model.predict_1(data) if results is None: return res_pred_name1 res_pred_name1 = res_pred_name1 + results res_pred_name1 = sorted(res_pred_name1, key=lambda x: x[1], reverse=True) return res_pred_name1[:10] # ========================= # SMILES → Image # ========================= def smiles_to_image1(smiles, size=(300, 300)): mol = Chem.MolFromSmiles(smiles) if mol is None: return None return Draw.MolToImage(mol, size=size) def smiles_to_image(smiles, size=(300, 300)): """ 通过读取本地文件并 Resize 后返回 PIL 对象,避免 Gradio 路径权限报错 """ # 你的图片基础路径 # base_image_path = "/Users/xiaojie/Documents/lunwen/代码/data/" # 构造图片完整路径(这里假设文件名就是 smiles + .png) # 如果你的文件名逻辑不同,请根据实际情况修改 img_path = os.path.join(base_image_path, f"{smiles}.png") # img_path = '/Users/xiaojie/Documents/lunwen/代码/data/Br.C=CC1CN2CCC1CC2C(O)c1ccnc2ccc(OC)cc12.png' try: if os.path.exists(img_path): # 打开图片文件 img = Image.open(img_path) # 强制调整为 300x300 大小 img = img.resize(size, Image.Resampling.LANCZOS) return img else: print(f"警告: 文件未找到 {img_path}") # time.sleep(1) img = smiles_to_image1(smiles) return img except Exception as e: print(f"处理图片时出错: {e}") return None # ========================= # 加载用户 SMILES # ========================= def load_user_smiles(file): if file is None: return [] if file.name.endswith(".txt"): with open(file.name, "r") as f: smiles = [l.strip() for l in f if l.strip()] elif file.name.endswith(".csv"): df = pd.read_csv(file.name) if "smiles" not in df.columns: raise ValueError("CSV 中必须包含 smiles 列") smiles = df["smiles"].dropna().tolist() else: smiles = [] return [s for s in smiles if Chem.MolFromSmiles(s)] def parse_msp(file): if file.name.endswith(".msp"): with open(file.name, "r", encoding='utf-8') as f: lines = f.readlines() start_index = 0 for i, line in enumerate(lines): if line.startswith("Num Peaks:"): start_index = i + 1 break # 3. 提取数据行并转换为二维数组 peaks_array = [] for line in lines[start_index:]: # split() 会自动处理空格和制表符 (\t) parts = line.replace("\n", "").split() if len(parts) == 2: peaks_array.append([float(parts[0]), float(parts[1])]) return peaks_array elif file.name.endswith(".mgf"): with open(file.name, "r", encoding='utf-8') as f: lines = f.readlines() peaks_array = [] # 2. 遍历每一行进行判断 for line in lines: # 排除掉不包含数值的标签行 if not line or any(tag in line for tag in ["BEGIN", "END", "=", "NAME"]): continue # 尝试将行内容切分为两部分 parts = line.replace("\n", "").split() # 如果分割后有两个元素,且第一个元素是数字,则加入数组 if len(parts) == 2: try: mz = float(parts[0]) intensity = float(parts[1]) peaks_array.append([mz, intensity]) except ValueError: # 如果转换数字失败(比如标题行),则跳过 continue return peaks_array elif file.name.endswith(".json"): with open(file.name, "r", encoding='utf-8') as f: line = json.load(f) peaks_array = None if 'ms' in line: peaks_array = line['ms'] elif 'Ms' in line: peaks_array = line['Ms'] elif 'MS' in line: peaks_array = line['MS'] elif 'mS' in line: peaks_array = line['mS'] return peaks_array elif file.name.endswith(".mzXML"): with open(file.name, "r", encoding='utf-8') as f: lines = f.readlines() return None elif file.name.endswith(".mzML"): with open(file.name, "r", encoding='utf-8') as f: lines = f.readlines() return None else: return None # ========================= # 第一类:业务逻辑(只算结果) # ========================= def run_retrieval(msp_file, ion_mode, parent_mass, user_smiles_file, parent_mass_bn=50): if msp_file is None: return [], None print("msp_file ...", msp_file) print("ion_mode ...", ion_mode) print("parent_mass ...", parent_mass) print("user_smiles_file ...", user_smiles_file) user_smiles = load_user_smiles(user_smiles_file) msp_file = parse_msp(msp_file) print("msp_file ...", msp_file) if msp_file is None: return [], None print("user_smiles ...", user_smiles) smiles_list = cross_modal_retrieval(parent_mass, ion_mode, msp_file, user_smiles, parent_mass_bn) print("smiles_list ...", smiles_list) # scores_list = [] if isinstance(smiles_list[0], list): scores_list = [x[1] for x in smiles_list] smiles_list = [x[0] for x in smiles_list] else: scores_list = [1.0] * len(smiles_list) results = [] rows = [] for i, smi in enumerate(smiles_list): if i >= 10: break img = smiles_to_image(smi) results.append({ "rank": i + 1, "img": img, "smiles": smi, "score": scores_list[i] }) rows.append({ "rank": i + 1, "parent_ion_mass": parent_mass, "ion_mode": ion_mode, "smiles": smi, "score": scores_list[i] }) print("len results ...") print(len(results)) csv_path = "cross_modal_results.csv" pd.DataFrame(rows).to_csv(csv_path, index=False) return results, csv_path # ========================= # 第一类:业务逻辑(只算结果) # ========================= def run_retrieval_2(msp_file, ion_mode, parent_mass, user_smiles_file, compound_num_min, compound_name_max, pr=10, parent_mass_bn=50): if msp_file is None: return [], None print("msp_file ...", msp_file) print("ion_mode ...", ion_mode) print("parent_mass ...", parent_mass) print("user_smiles_file ...", user_smiles_file) user_smiles = load_user_smiles(user_smiles_file) msp_file_s1, msp_file_s2 = saixuan_example(msp_file, compound_num_min, compound_name_max, pr) if msp_file_s1 is None or len(msp_file_s1) == 0: return [], None print("msp_file_s ...", len(msp_file_s1)) print("user_smiles ...", user_smiles) results = [] rows = [] for i11, msp_mass_file in enumerate(msp_file_s1): msp_file = msp_mass_file['ms'] parent_mass = msp_mass_file['mass'][0] smiles_list = cross_modal_retrieval(parent_mass, ion_mode, msp_file, user_smiles, parent_mass_bn) print("smiles_list ...", smiles_list) # scores_list = [] if isinstance(smiles_list[0], list): scores_list = [x[1] for x in smiles_list] smiles_list = [x[0] for x in smiles_list] else: scores_list = [1.0] * len(smiles_list) for i, smi in enumerate(smiles_list): if i >= 10: break if i11 == 0: img = smiles_to_image(smi) results.append({ "rank": i + 1, "img": img, "smiles": smi, "score": scores_list[i] }) rows.append({ "id": i11, "rank": i + 1, "parent_ion_mass": parent_mass, "ion_mode": ion_mode, "smiles": smi, "score": scores_list[i] }) print("len results ...") print(len(results)) import datetime timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") csv_path = f"cross_modal_results_{timestamp}.csv" pd.DataFrame(rows).to_csv(csv_path, index=False) return results[:10], csv_path # ========================= # 第二类:UI 适配(固定铺 10 个) # ========================= def fill_top10(results, csv_path): images = [] smiles_texts = [] score_texts = [] print("results ...") print(results) for i in range(10): if i < len(results): images.append(results[i]["img"]) smiles_texts.append(results[i]["smiles"]) score_texts.append(str(results[i]['score'])) else: images.append(None) smiles_texts.append("") score_texts.append("") return images + smiles_texts + score_texts + [csv_path] def unified_run_retrieval(msp_file, ion_mode, parent_mass, user_smiles_file, compound_mode, compound_num_min, compound_name_max, pr=10, parent_mass_bn=50): # 1. 执行原有的 retrieval 逻辑 if compound_mode == "单化合物": results, csv_path = run_retrieval(msp_file, ion_mode, parent_mass, user_smiles_file, parent_mass_bn) else: results, csv_path = run_retrieval_2(msp_file, ion_mode, parent_mass, user_smiles_file, compound_num_min, compound_name_max, pr, parent_mass_bn) # 2. 调用原有的 fill_top10 逻辑获取 UI 输出列表 # 假设 fill_top10 返回的是 image_outputs + smiles_outputs + [csv_file] 对应的值 ui_outputs = fill_top10(results, csv_path) return ui_outputs # ========================= # 示例填充 # ========================= def load_example(): time.sleep(1) return "/Users/xiaojie/Documents/lunwen/代码/data/29579-06-0.mgf", "pos", 336.1735 def parse_ms_data_simple(lines): results = [] current_entry = None # 按行分割文本 # lines = raw_text.strip().split('\n') for line in lines: line = line.strip() if not line: continue # 识别样本开始 if line == "BEGIN IONS": current_entry = {"ms": [], "mass": []} continue # 识别样本结束 if line == "END IONS": if current_entry is not None: results.append(current_entry) current_entry = None continue # 处于样本块内部时 if current_entry is not None: if line.startswith("PEPMASS="): # 提取 PEPMASS= 之后的部分并按空格切分 mass_values = line.split('=')[1].split() current_entry["mass"] = [float(v) for v in mass_values] elif "=" in line: # 跳过其他元数据行,如 TITLE, CHARGE, RTINSECONDS 等 continue else: # 处理碎片数据行 (m/z intensity) parts = line.split() if len(parts) >= 2: try: mz = float(parts[0]) intensity = float(parts[1]) current_entry["ms"].append([mz, intensity]) except ValueError: # 容错处理:如果行首不是数字则跳过 continue return results def parse_ms_data_msp(lines): results = [] # lines = text.splitlines() current_entry = None capture_ms = False peaks_to_collect = 0 for line in lines: line = line.strip() if not line: continue # 遇到新条目开始 if line.startswith("NAME:"): if current_entry and current_entry["ms"]: results.append(current_entry) current_entry = {"mass": None, "ms": []} capture_ms = False continue # 解析 PRECURSORMZ if line.startswith("PRECURSORMZ:"): mass_values = line.split(":", 1)[1].strip().split() current_entry["mass"] = [float(v) for v in mass_values] # 检测 Num Peaks,开始捕获 MS 数据 elif line.startswith("Num Peaks:"): val = line.split(":", 1)[1].strip() peaks_to_collect = int(val) capture_ms = True # 捕获峰数据 elif capture_ms and peaks_to_collect > 0: parts = line.split() if len(parts) >= 2: mz = float(parts[0]) intensity = float(parts[1]) current_entry["ms"].append([mz, intensity]) peaks_to_collect -= 1 if peaks_to_collect == 0: capture_ms = False # 添加最后一个条目 if current_entry and current_entry["ms"]: results.append(current_entry) return results def saixuan_example(file, min_val, max_val, pr=10): if file.name.endswith(".msp"): try: with open(file.name, "r", encoding='utf-8') as f: lines = f.readlines() res2 = parse_ms_data_msp(lines) res = [] for x in res2: ms = [] for c in x['ms']: if float(c[1]) <= pr: continue ms.append(c) if len(ms) > 0: res.append({ "ms": ms, "mass": x["mass"] }) res1 = [] for x in res: pep_mass = x["mass"] if len(pep_mass) == 2: if float(pep_mass[1]) >= float(min_val) and float(pep_mass[1]) <= float(max_val): res1.append(x) else: res1.append(x) return res1, res except Exception as e: print(e) return None, None elif file.name.endswith(".mgf"): try: with open(file.name, "r", encoding='utf-8') as f: lines = f.readlines() res2 = parse_ms_data_simple(lines) res = [] for x in res2: ms = [] for c in x['ms']: if float(c[1]) <= pr: continue ms.append(c) if len(ms) > 0: res.append({ "ms": ms, "mass": x["mass"] }) res1 = [] for x in res: pep_mass = x["mass"] if float(pep_mass[1]) >= float(min_val) and float(pep_mass[1]) <= float(max_val): res1.append(x) return res1, res except Exception as e: print(e) return None, None elif file.name.endswith(".json"): try: with open(file.name, "r", encoding='utf-8') as f: lines = json.load(f) res2 = [] for line in lines: d = {"ms": None, "mass": None} if 'ms' in line: d['ms'] = line['ms'] elif 'Ms' in line: d['ms'] = line['Ms'] elif 'MS' in line: d['ms'] = line['MS'] elif 'mS' in line: d['ms'] = line['mS'] if 'parent_mz' in line: d['mass'] = [line['parent_mz']] else: d['mass'] = [0] res2.append(d) res = [] for x in res2: ms = [] for c in x['ms']: if float(c[1]) <= pr: continue ms.append(c) if len(ms) > 0: res.append({ "ms": ms, "mass": x["mass"] }) res1 = [] for x in res: pep_mass = x["mass"] if len(pep_mass) == 2: if float(pep_mass[1]) >= float(min_val) and float(pep_mass[1]) <= float(max_val): res1.append(x) else: res1.append(x) return res1, res except Exception as e: print(e) return None, None else: return None, None # ---- 模拟文件解析逻辑 ---- def process_compounds(file, mode, min_val, max_val, pr=10): if not file: return "未上传文件", "未上传文件" if mode == "单化合物": return "1", "1" # 这里放置你的解析逻辑,例如读取 msp 文件内容 # 假设我们通过某种逻辑得到了总数和筛选后的数量 # 以下为示例数值: try: # demo 逻辑:模拟从文件中读到了 50 个,经过响应值筛选剩下 10 个 res1, res2 = saixuan_example(file, min_val, max_val, pr=pr) if res1 is not None and res2 is not None: parsed_selected = len(res1) parsed_total = len(res2) return str(parsed_selected), str(parsed_total) else: return '0', '0' except Exception as e: return "解析错误", str(e) # 1. 图片转 Base64 函数(解决路径显示不出的问题) def get_base64_image(image_path): try: with open(image_path, "rb") as img_file: return base64.b64encode(img_file.read()).decode('utf-8') except Exception as e: print(f"图片读取失败: {e}") return "" # CSS 样式:定义背景颜色、文字颜色和悬停效果 custom_css = """ #yellow_btn { background-color: #FEBA02 !important; /* 黄色背景 */ color: black !important; /* 黑色文字,确保对比度 */ border: none; } #yellow_btn:hover { background-color: #e6b800 !important; /* 鼠标悬停时稍微深一点的黄色 */ } /* 新增:图片中的四种颜色按钮样式 */ /* 1. 深绿色 (#3D9F3C) */ #btn_green_dark { background-color: #3D9F3C !important; color: white !important; border: none; } #btn_green_dark:hover { background-color: #328532 !important; /* 悬停颜色稍深 */ } /* 2. 浅绿色 (#9ED17B) */ #btn_green_light { background-color: #9ED17B !important; color: black !important; border: none; } #btn_green_light:hover { background-color: #89c065 !important; } /* 3. 中蓝色 (#367DB0) */ #btn_blue_medium { background-color: #367DB0 !important; color: white !important; border: none; } #btn_blue_medium:hover { background-color: #2b658f !important; } /* 4. 天蓝色 (#9DC7DD) */ #btn_blue_light { background-color: #9DC7DD !important; color: black !important; border: none; } #btn_blue_light:hover { background-color: #85b5cc !important; } """ # ========================= # Gradio UI # ========================= with gr.Blocks(title="MS² → SMILES Cross-Modal Retrieval", css=custom_css) as demo: # gr.Markdown("## 🔬 MS² → SMILES Cross-Modal Retrieval") # gr.Markdown( # """ #
# Mass Spectrometry to Chemical Structure #
#