"""批量筛选页面。 在一组固定的标准配方下,对上传的 SMILES 列表批量调用 /predict/batch 并排序, 用于从大量候选脂质中初筛。本页不做配方搜索——搜索请用主页的配方优选。 """ import io import json import os from pathlib import Path import time from typing import Optional import uuid import httpx import pandas as pd import streamlit as st from ui_common import ( API_URL, AVAILABLE_ORGANS, AVAILABLE_ROUTES, EE_CLASS_LABELS, HELPER_LIPID_OPTIONS, ORGAN_LABELS, PDI_CLASS_LABELS, ROUTE_LABELS, TOXIC_CLASS_LABELS, ) CHUNK_SIZE = 100 RUNS_DIR = Path( os.environ.get("SCREENING_RUNS_DIR", Path.home() / ".cache" / "lnp_screening_runs") ) RUN_TTL_HOURS = 48 # 标准配方默认值,取自 data/interim/internal.csv 的中位配方 DEFAULT_FORMULATION = { "weight_ratio": 10.0, "cationic_mol": 35.0, "phospholipid_mol": 16.0, "cholesterol_mol": 46.5, "peg_mol": 2.5, } MOL_SUM_TOLERANCE = 0.5 _SMILES_MIN_LEN = 10 _SMILES_COL_RATIO = 0.5 LAYOUT_LONG = "long" LAYOUT_MATRIX = "matrix" LAYOUT_LABELS = { LAYOUT_LONG: "长表:某一列是 SMILES,一行一个分子", LAYOUT_MATRIX: "矩阵表:行列各代表一个结构维度,单元格是 SMILES", } def looks_like_smiles(value) -> bool: """判断一个单元格是否为 SMILES。""" text = str(value).strip() if len(text) < _SMILES_MIN_LEN or " " in text: return False return any(atom in text for atom in "CNOcno") def smiles_like_columns(df: pd.DataFrame) -> list: """返回过半单元格看起来像 SMILES 的列。""" return [ col for col in df.columns if df[col].map(looks_like_smiles).mean() > _SMILES_COL_RATIO ] def guess_layout(df: pd.DataFrame) -> str: """两列以上都装着 SMILES,基本只能是矩阵表。""" return LAYOUT_MATRIX if len(smiles_like_columns(df)) >= 2 else LAYOUT_LONG def melt_matrix( df: pd.DataFrame, index_col, value_cols: list, row_name: str, col_name: str, ) -> pd.DataFrame: """把矩阵表展开为 [row_name, col_name, SMILES] 三列的长表。""" long_df = df.melt( id_vars=[index_col], value_vars=value_cols, var_name=col_name, value_name="SMILES", ).rename(columns={index_col: row_name}) long_df["SMILES"] = long_df["SMILES"].astype(str).str.strip() long_df[row_name] = long_df[row_name].astype(str).str.strip() return long_df[long_df["SMILES"].map(looks_like_smiles)].reset_index(drop=True) def save_run(payload: dict) -> str: """把一次筛选结果写盘,返回可放进 URL 的 token。""" RUNS_DIR.mkdir(parents=True, exist_ok=True) os.chmod(RUNS_DIR, 0o700) token = uuid.uuid4().hex[:12] # 标签列可能来自 Excel 的数值列,取出来是 numpy 标量,json 无法直接序列化, # 故用 default=str 兜底。 (RUNS_DIR / f"{token}.json").write_text( json.dumps(payload, ensure_ascii=False, default=str), encoding="utf-8" ) return token def load_run(token: str) -> Optional[dict]: """按 token 读回结果。token 来自 URL,属于外部输入,必须校验以防目录穿越。""" if not token or not token.isalnum(): return None path = RUNS_DIR / f"{token}.json" if not path.is_file(): return None try: return json.loads(path.read_text(encoding="utf-8")) except (OSError, ValueError): return None def purge_old_runs() -> None: """清掉过期结果,避免长期堆积。""" cutoff = time.time() - RUN_TTL_HOURS * 3600 for path in RUNS_DIR.glob("*.json"): try: if path.stat().st_mtime < cutoff: path.unlink() except OSError: continue def render_results(run: dict) -> None: """渲染一次筛选的结果表与导出按钮。run 可能来自 session_state,也可能来自磁盘。""" df = pd.DataFrame(run["rows"]) # 排序是纯事后操作,放在这里而不是运行前,改排序就不必重跑整批 sort_options = ["delivery"] + AVAILABLE_ORGANS stored = run.get("sort_key", "delivery") sort_key = st.selectbox( "结果排序依据", options=sort_options, index=sort_options.index(stored) if stored in sort_options else 0, format_func=lambda x: ( "递送效率 (delivery)" if x == "delivery" else f"{ORGAN_LABELS[x]} 分布占比" ), key="_result_sort", ) sort_col = "delivery" if sort_key == "delivery" else f"分布_{ORGAN_LABELS[sort_key]}" df = df.sort_values(sort_col, ascending=False).reset_index(drop=True) df.insert(0, "排名", range(1, len(df) + 1)) st.subheader(f"筛选结果(成功 {len(df)} / {run['n_requested']} 条)") st.caption(f"运行于 {run['finished_at']}|{run['meta_line']}") if run["failures"]: with st.expander(f"⚠️ {len(run['failures'])} 条失败,点击查看"): st.code("\n".join(run["failures"][:200])) st.dataframe( df, use_container_width=True, hide_index=True, column_config={"SMILES": st.column_config.TextColumn("SMILES", width="medium")}, ) buffer = io.StringIO() buffer.write(f"# {run['meta_line']}, sorted_by={sort_col}\n") df.to_csv(buffer, index=False) st.download_button( "下载完整结果 CSV", data=buffer.getvalue().encode("utf-8-sig"), file_name=f"screening_{len(df)}_lipids.csv", mime="text/csv", use_container_width=True, ) if st.button("清除结果,开始新一批", use_container_width=True): st.session_state.pop("screen_run", None) st.query_params.pop("run", None) st.rerun() def render(api_online: bool) -> None: """渲染批量筛选界面。由 app.py 在「批量筛选」模式下调用。""" st.subheader("批量筛选") st.caption( "在同一组标准配方下批量预测多个脂质的性能并排序,用于从大量候选中初筛。" "初筛出的前几十条可切回「配方优选」做配方搜索。" ) if not api_online: st.error(f"API 服务离线,无法筛选。当前 API_URL: {API_URL}") st.stop() if "screen_run" not in st.session_state: purge_old_runs() restored = load_run(st.query_params.get("run", "")) if restored: st.session_state["screen_run"] = restored # ============ 1. 上传与解析 ============ uploaded = st.file_uploader( "脂质库文件", type=["xlsx", "xls", "csv"], help="支持两种布局:一列 SMILES 的长表,或行列各代表一个结构维度的矩阵表。", ) if uploaded is None: # 刷新后文件对象丢失,但结果已从磁盘还原,此处直接展示,不必重新上传 if st.session_state.get("screen_run"): st.info("已恢复上一次的筛选结果。要跑新一批,重新上传文件即可。") render_results(st.session_state["screen_run"]) st.stop() try: if uploaded.name.lower().endswith(".csv"): df_in = pd.read_csv(uploaded) else: df_in = pd.read_excel(uploaded) except (ValueError, pd.errors.ParserError, OSError) as exc: st.error(f"读取文件失败: {exc}") st.stop() if df_in.empty: st.error("文件内容为空") st.stop() st.success(f"已读取 {len(df_in)} 行、{len(df_in.columns)} 列") with st.expander("预览前 5 行"): st.dataframe(df_in.head(), use_container_width=True) _columns = list(df_in.columns) _smiles_cols = smiles_like_columns(df_in) _layouts = [LAYOUT_LONG, LAYOUT_MATRIX] layout = st.radio( "表格布局", options=_layouts, index=_layouts.index(guess_layout(df_in)), format_func=lambda x: LAYOUT_LABELS[x], help="已按文件内容自动判断,判断错误时可手动切换。", ) if layout == LAYOUT_MATRIX: col_left, col_right = st.columns(2) with col_left: index_col = st.selectbox( "行标签所在列", options=_columns, index=0, help="通常是最左侧那一列,例如头部基团编号。", ) row_name = st.text_input("行维度名称", value="头部基团") with col_right: value_cols = st.multiselect( "SMILES 单元格所在列", options=[c for c in _columns if c != index_col], default=[c for c in _smiles_cols if c != index_col], ) col_name = st.text_input("列维度名称", value="尾链") if not value_cols: st.error("请至少选择一列 SMILES 单元格") st.stop() parsed = melt_matrix(df_in, index_col, value_cols, row_name, col_name) meta_cols = [row_name, col_name] else: _guess = next( iter(_smiles_cols), next((c for c in _columns if "smiles" in str(c).lower()), _columns[0]), ) smiles_col = st.selectbox( "SMILES 所在列", options=_columns, index=_columns.index(_guess) ) meta_cols = st.multiselect( "一并带入结果的标签列", options=[c for c in _columns if c != smiles_col], help="例如分子名称、编号、来源。这些列不参与预测,只原样写入结果表。", ) parsed = df_in[[smiles_col] + meta_cols].rename(columns={smiles_col: "SMILES"}) parsed["SMILES"] = parsed["SMILES"].astype(str).str.strip() parsed = parsed[parsed["SMILES"].map(looks_like_smiles)].reset_index(drop=True) if parsed.empty: st.error("没有解析出任何有效的 SMILES,请检查布局设置与列选择") st.stop() n_parsed = len(parsed) if st.checkbox("去除重复 SMILES", value=True): parsed = parsed.drop_duplicates(subset="SMILES").reset_index(drop=True) records = parsed.to_dict("records") _dedup_note = f"(原始 {n_parsed} 条,已去重)" if len(records) != n_parsed else "" st.write(f"待筛选:**{len(records)}** 条{_dedup_note}") with st.expander("确认解析结果(前 10 条)"): st.dataframe(parsed.head(10), use_container_width=True, hide_index=True) # ============ 2. 标准配方 ============ st.subheader("统一使用的标准配方") st.caption("所有 SMILES 都在同一组配方下预测,便于横向比较。默认值取自内部数据的中位配方。") col_a, col_b, col_c = st.columns(3) with col_a: weight_ratio = st.number_input( "脂质/mRNA 重量比", 0.1, 50.0, DEFAULT_FORMULATION["weight_ratio"], 0.5 ) cationic_mol = st.number_input( "阳离子脂质 mol%", 0.0, 100.0, DEFAULT_FORMULATION["cationic_mol"], 0.5 ) with col_b: phospholipid_mol = st.number_input( "磷脂 mol%", 0.0, 100.0, DEFAULT_FORMULATION["phospholipid_mol"], 0.5 ) cholesterol_mol = st.number_input( "胆固醇 mol%", 0.0, 100.0, DEFAULT_FORMULATION["cholesterol_mol"], 0.5 ) with col_c: peg_mol = st.number_input( "PEG 脂质 mol%", 0.0, 20.0, DEFAULT_FORMULATION["peg_mol"], 0.1 ) helper_lipid = st.selectbox("辅助脂质", HELPER_LIPID_OPTIONS) route = st.selectbox( "给药途径", AVAILABLE_ROUTES, format_func=lambda x: ROUTE_LABELS[x] ) mol_sum = cationic_mol + phospholipid_mol + cholesterol_mol + peg_mol if abs(mol_sum - 100.0) > MOL_SUM_TOLERANCE: st.error(f"四项 mol 比例之和须为 100%,当前为 {mol_sum:.1f}%") st.stop() st.caption(f"mol 比例合计 {mol_sum:.1f}% ✓") # ============ 3. 运行选项 ============ use_llm = st.checkbox( "启用 LLM 分支(更准,但明显更慢)", value=True, help="关闭后 delivery 精度会下降,但吞吐大幅提升。条数很多时可先关闭做粗筛。", ) if st.button("开始批量筛选", type="primary", use_container_width=True): # ============ 4. 分块调用 /predict/batch ============ def build_item(smiles: str) -> dict: """构造 /predict/batch 的单个 items 元素。""" return { "smiles": smiles, "cationic_lipid_to_mrna_ratio": weight_ratio, "cationic_lipid_mol_ratio": cationic_mol, "phospholipid_mol_ratio": phospholipid_mol, "cholesterol_mol_ratio": cholesterol_mol, "peg_lipid_mol_ratio": peg_mol, "helper_lipid": helper_lipid, "route": route, } def flatten_prediction(pred: dict) -> dict: """把 PredictResponse 展平为一行表格数据。""" row = { "SMILES": pred["smiles"], "delivery": pred["quantified_delivery"], "delivery_原始值": pred.get("unnormalized_delivery"), "粒径_nm": pred.get("size"), "PDI": PDI_CLASS_LABELS.get(pred.get("pdi_class")), "EE": EE_CLASS_LABELS.get(pred.get("ee_class")), "毒性": TOXIC_CLASS_LABELS.get(pred.get("toxic_class")), } for organ, value in pred["biodist"].items(): row[f"分布_{ORGAN_LABELS.get(organ, organ)}"] = value return row chunks = [records[i:i + CHUNK_SIZE] for i in range(0, len(records), CHUNK_SIZE)] progress = st.progress(0.0) status = st.empty() rows: list[dict] = [] failures: list[str] = [] started = time.time() with httpx.Client(timeout=1800) as client: for idx, chunk in enumerate(chunks): done = idx * CHUNK_SIZE eta = "" if done: per_item = (time.time() - started) / done eta = f",预计剩余 {per_item * (len(records) - done) / 60:.1f} 分钟" status.text( f"批次 {idx + 1}/{len(chunks)}" f"(已完成 {done}/{len(records)} 条){eta}" ) try: response = client.post( f"{API_URL}/predict/batch", json={ "items": [build_item(r["SMILES"]) for r in chunk], "batch_size": 32, "use_llm": use_llm, }, ) response.raise_for_status() payload = response.json() except httpx.HTTPError as exc: failures.append(f"批次 {idx + 1} 整体失败: {exc}") progress.progress((idx + 1) / len(chunks)) continue error_idx = {err["index"] for err in payload.get("errors", [])} for err in payload.get("errors", []): failures.append( f"{chunk[err['index']]['SMILES'][:60]}: {err['detail']}" ) valid = [r for j, r in enumerate(chunk) if j not in error_idx] for record, pred in zip(valid, payload["predictions"]): row = {col: record.get(col) for col in meta_cols} row.update(flatten_prediction(pred)) rows.append(row) progress.progress((idx + 1) / len(chunks)) progress.progress(1.0) status.text(f"完成,用时 {(time.time() - started) / 60:.1f} 分钟") # ============ 5. 结果入库 ============ if not rows: st.error("没有任何成功的预测结果") if failures: st.code("\n".join(failures[:50])) st.stop() run = { "rows": rows, "failures": failures, "n_requested": len(records), "sort_key": "delivery", "finished_at": time.strftime("%Y-%m-%d %H:%M"), "meta_line": ( f"标准配方: weight_ratio={weight_ratio}, cationic={cationic_mol}, " f"phospholipid={phospholipid_mol}, cholesterol={cholesterol_mol}, " f"peg={peg_mol}, helper_lipid={helper_lipid}, route={route}, " f"use_llm={use_llm}" ), } st.session_state["screen_run"] = run st.query_params["run"] = save_run(run) # ============ 6. 结果展示 ============ if st.session_state.get("screen_run"): render_results(st.session_state["screen_run"])