mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-10-01 21:33:22 +08:00
466 lines
16 KiB
Python
466 lines
16 KiB
Python
"""批量筛选页面。
|
||
|
||
在一组固定的标准配方下,对上传的 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"]) |