feat(app): 新增批量筛选,结果与模式可跨刷新恢复

This commit is contained in:
Michelle0574 2026-09-05 17:15:37 +00:00
parent 9810854ac7
commit a812694500
6 changed files with 804 additions and 109 deletions

View File

@ -5,32 +5,30 @@ FastAPI 配方优化 API
uvicorn app.api:app --host 0.0.0.0 --port 8000 --reload
"""
from contextlib import asynccontextmanager
import os
from pathlib import Path
from typing import List, Dict, Optional
from contextlib import asynccontextmanager
from typing import Dict, List, Optional
import torch
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, Field
from loguru import logger
from pydantic import BaseModel, Field
import torch
from lnp_ml.config import MODELS_DIR
from lnp_ml.modeling.predict import load_model
from app.optimize import (
optimize,
format_results,
AVAILABLE_ORGANS,
HELPER_LIPID_OPTIONS,
ROUTE_OPTIONS,
TARGET_BIODIST,
CompRanges,
ScoringWeights,
HELPER_LIPID_OPTIONS,
ROUTE_OPTIONS,
create_dataframe_from_formulations,
optimize,
predict_all,
)
from lnp_ml.config import MODELS_DIR
from lnp_ml.modeling.predict import load_model
# ============ Pydantic Models ============
@ -220,6 +218,32 @@ class PredictResponse(BaseModel):
}
class BatchPredictItemError(BaseModel):
"""批量请求中单个非法条目的错误信息"""
index: int
detail: str
class BatchPredictRequest(BaseModel):
"""批量配方预测请求"""
items: List[PredictRequest] = Field(
..., min_length=1, max_length=500, description="配方列表,单次最多 500 条"
)
batch_size: int = Field(default=32, ge=1, le=256, description="模型前向批大小")
use_llm: bool = Field(
default=True,
description="是否启用 LLM 分支;关闭可大幅提速,但会损失 delivery 精度",
)
class BatchPredictResponse(BaseModel):
"""批量配方预测响应。单条非法不中断整批,汇总在 errors 中返回"""
n_requested: int
n_succeeded: int
predictions: List[PredictResponse]
errors: List[BatchPredictItemError] = []
# ============ Global State ============
class ModelState:
@ -266,6 +290,7 @@ async def lifespan(app: FastAPI):
if _llm is not None and getattr(_llm, "use_rag", False):
import numpy as np
import pandas as pd
from lnp_ml.dataset import LNPDataset, process_dataframe
from lnp_ml.modeling.nested_cv_optuna import _build_rag_pool
@ -409,6 +434,112 @@ async def predict_formulation(request: PredictRequest):
)
@app.post("/predict/batch", response_model=BatchPredictResponse)
async def predict_formulations_batch(request: BatchPredictRequest):
"""批量配方预测:一次请求多条,模型侧批量前向,用于虚拟筛选。"""
import pandas as pd
from rdkit import Chem, RDLogger
if state.model is None:
raise HTTPException(status_code=503, detail="Model not loaded")
RDLogger.DisableLog("rdApp.*")
valid_idx: List[int] = []
frames = []
errors: List[BatchPredictItemError] = []
for i, item in enumerate(request.items):
if item.helper_lipid not in HELPER_LIPID_OPTIONS:
errors.append(BatchPredictItemError(
index=i, detail=f"Invalid helper_lipid: {item.helper_lipid}"))
continue
if item.route not in ROUTE_OPTIONS:
errors.append(BatchPredictItemError(
index=i, detail=f"Invalid route: {item.route}"))
continue
mol_sum = (
item.cationic_lipid_mol_ratio + item.phospholipid_mol_ratio
+ item.cholesterol_mol_ratio + item.peg_lipid_mol_ratio
)
if abs(mol_sum - 100.0) > 0.5:
errors.append(BatchPredictItemError(
index=i, detail=f"四项 mol 比例之和须为 100(当前 {mol_sum:.2f})"))
continue
if Chem.MolFromSmiles(item.smiles) is None:
errors.append(BatchPredictItemError(
index=i, detail=f"无法解析的 SMILES: {item.smiles[:80]}"))
continue
frames.append(create_dataframe_from_formulations(
item.smiles,
[(
item.cationic_lipid_to_mrna_ratio,
item.cationic_lipid_mol_ratio,
item.phospholipid_mol_ratio,
item.cholesterol_mol_ratio,
item.peg_lipid_mol_ratio,
)],
[item.helper_lipid],
[item.route],
))
valid_idx.append(i)
if not frames:
raise HTTPException(
status_code=400,
detail=f"无合法条目;首个错误: {errors[0].detail if errors else 'unknown'}",
)
logger.info(f"Batch predict: {len(frames)}/{len(request.items)} valid, "
f"use_llm={request.use_llm}, batch_size={request.batch_size}")
df = pd.concat(frames, ignore_index=True)
_has_llm = getattr(state.model, "llm_prompt", None) is not None
try:
if _has_llm:
state.model.set_llm_enabled(request.use_llm)
df = predict_all(state.model, df, state.device, batch_size=request.batch_size)
except Exception as e:
logger.error(f"Batch prediction failed: {e}")
if torch.cuda.is_available():
torch.cuda.empty_cache()
raise HTTPException(status_code=500, detail=str(e))
finally:
# /optimize 与本接口共享同一模型实例,用完必须复位,避免污染后续请求
if _has_llm:
state.model.set_llm_enabled(True)
predictions: List[PredictResponse] = []
for j, i in enumerate(valid_idx):
row = df.iloc[j]
item = request.items[i]
_size, _unnorm = row.get("pred_size"), row.get("pred_unnorm_delivery")
predictions.append(PredictResponse(
smiles=item.smiles,
helper_lipid=item.helper_lipid,
route=item.route,
biodist={
c.replace("Biodistribution_", ""): float(row[f"pred_{c}"])
for c in TARGET_BIODIST
},
size=float(_size) if pd.notna(_size) else None,
quantified_delivery=float(row["pred_delivery"]),
unnormalized_delivery=float(_unnorm) if pd.notna(_unnorm) else None,
pdi_class=int(row["pred_pdi_class"]),
ee_class=int(row["pred_ee_class"]),
toxic_class=int(row["pred_toxic_class"]),
))
return BatchPredictResponse(
n_requested=len(request.items),
n_succeeded=len(predictions),
predictions=predictions,
errors=errors,
)
@app.post("/optimize", response_model=OptimizeResponse)
async def optimize_formulation(request: OptimizeRequest):
"""
@ -481,7 +612,7 @@ async def optimize_formulation(request: OptimizeRequest):
)
# 用于计算综合评分的权重
from app.optimize import compute_formulation_score, DEFAULT_SCORING_WEIGHTS
from app.optimize import DEFAULT_SCORING_WEIGHTS, compute_formulation_score
actual_scoring_weights = scoring_weights if scoring_weights is not None else DEFAULT_SCORING_WEIGHTS
# 转换结果

View File

@ -6,50 +6,32 @@ Streamlit 配方优化交互界面
Docker 环境变量:
API_URL: API 服务地址 (默认: http://localhost:8000)
APP_PASSWORD: 访问口令,设置后进入页面需先输入;不设则不启用
"""
import io
import os
from datetime import datetime
import httpx
import pandas as pd
import streamlit as st
# ============ 配置 ============
# 从环境变量读取 API 地址,支持 Docker 环境
API_URL = os.environ.get("API_URL", "http://localhost:8000")
AVAILABLE_ORGANS = [
"liver",
"spleen",
"lung",
"heart",
"kidney",
"muscle",
"lymph_nodes",
]
ORGAN_LABELS = {
"liver": "肝脏 (Liver)",
"spleen": "脾脏 (Spleen)",
"lung": "肺 (Lung)",
"heart": "心脏 (Heart)",
"kidney": "肾脏 (Kidney)",
"muscle": "肌肉 (Muscle)",
"lymph_nodes": "淋巴结 (Lymph Nodes)",
}
AVAILABLE_ROUTES = [
"intravenous",
"intramuscular",
]
ROUTE_LABELS = {
"intravenous": "静脉注射 (Intravenous)",
"intramuscular": "肌肉注射 (Intramuscular)",
}
import batch_screening
from ui_common import (
API_URL,
AVAILABLE_ORGANS,
AVAILABLE_ROUTES,
EE_CLASS_LABELS,
MODE_LABELS,
MODE_OPTIMIZE,
MODE_SCREEN,
ORGAN_LABELS,
PDI_CLASS_LABELS,
ROUTE_LABELS,
TOXIC_CLASS_LABELS,
check_api_status,
render_api_status,
require_auth,
)
# ============ 页面配置 ============
@ -60,6 +42,10 @@ st.set_page_config(
initial_sidebar_state="expanded",
)
# ============ 访问口令 ============
require_auth()
# ============ 自定义样式 ============
st.markdown("""
@ -127,15 +113,6 @@ st.markdown("""
# ============ 辅助函数 ============
def check_api_status() -> bool:
"""检查 API 状态"""
try:
with httpx.Client(timeout=5) as client:
response = client.get(f"{API_URL}/")
return response.status_code == 200
except:
return False
def call_optimize_api(
smiles: str,
@ -174,26 +151,6 @@ def call_optimize_api(
return response.json()
# PDI 分类标签
PDI_CLASS_LABELS = {
0: "<0.2 (优)",
1: "≥0.2 (欠佳)",
}
# EE 分类标签
EE_CLASS_LABELS = {
0: "<50% (低)",
1: "50-80% (中)",
2: ">80% (高)",
}
# 毒性分类标签
TOXIC_CLASS_LABELS = {
0: "无毒 ✓",
1: "有毒 ⚠",
}
def format_results_dataframe(results: dict, smiles_label: str = None) -> pd.DataFrame:
"""将 API 结果转换为 DataFrame"""
formulations = results["formulations"]
@ -265,26 +222,32 @@ def main():
# 检查 API 状态
api_online = check_api_status()
# ========== 功能选择 ==========
_modes = [MODE_OPTIMIZE, MODE_SCREEN]
if "mode" not in st.session_state:
_from_url = st.query_params.get("mode")
st.session_state["mode"] = _from_url if _from_url in _modes else MODE_OPTIMIZE
mode = st.sidebar.radio(
"功能",
options=_modes,
format_func=lambda x: MODE_LABELS[x],
label_visibility="collapsed",
key="mode",
)
if st.query_params.get("mode") != mode:
st.query_params["mode"] = mode
with st.sidebar:
render_api_status(api_online)
if mode == MODE_SCREEN:
batch_screening.render(api_online)
return
# ========== 侧边栏 ==========
with st.sidebar:
# st.header("⚙️ 参数设置")
# API 状态
if api_online:
st.success("🟢 API 服务在线")
try:
with httpx.Client(timeout=5) as _c:
_info = _c.get(f"{API_URL}/").json()
_caps = [n for n, k in (("MoE", "use_moe"), ("LLM", "use_llm"), ("RAG", "use_rag"))
if _info.get(k)]
st.caption(f"模型: {' + '.join(_caps) if _caps else '仅 backbone'}|{_info.get('device', '?')}")
except Exception:
pass
else:
st.error("🔴 API 服务离线")
st.info(f"请先启动 API 服务:\n```\nuvicorn app.api:app --port 8010\n```\n"
f"当前 API_URL: {API_URL}")
# st.divider()
# SMILES 输入
@ -579,7 +542,7 @@ def main():
# 优化按钮
optimize_button = st.button(
"🚀 开始配方优选",
"开始配方优选",
type="primary",
use_container_width=True,
disabled=not api_online or not smiles_input.strip() or not selected_routes,
@ -636,7 +599,7 @@ def main():
except httpx.HTTPStatusError as e:
try:
error_detail = e.response.json().get("detail", str(e))
except:
except ValueError:
error_detail = str(e)
errors.append(f"SMILES {idx + 1}: {error_detail}")
except httpx.RequestError as e:
@ -740,7 +703,7 @@ def main():
target_organ = results["target_organ"]
st.download_button(
label="📥 导出 CSV",
label="导出 CSV",
data=csv_content,
file_name=f"lnp_optimization_{target_organ}_{datetime.now().strftime('%Y%m%d_%H%M%S')}.csv",
mime="text/csv",

466
app/batch_screening.py Normal file
View File

@ -0,0 +1,466 @@
"""批量筛选页面。
在一组固定的标准配方下,对上传的 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"])

View File

@ -7,30 +7,24 @@
python -m app.optimize --smiles "CC(C)..." --organ liver
"""
import itertools
from dataclasses import dataclass, field
import json
from pathlib import Path
from typing import Dict, List, Optional, Tuple
from dataclasses import dataclass, field
from loguru import logger
import numpy as np
import pandas as pd
import torch
from torch.utils.data import DataLoader
from loguru import logger
from tqdm import tqdm
import typer
from lnp_ml.config import MODELS_DIR
from lnp_ml.dataset import (
LNPDataset,
LNPDatasetConfig,
collate_fn,
SMILES_COL,
COMP_COLS,
HELP_COLS,
TARGET_BIODIST,
get_phys_cols,
LNPDataset,
collate_fn,
get_exp_cols,
)
from lnp_ml.modeling.predict import load_model

141
app/ui_common.py Normal file
View File

@ -0,0 +1,141 @@
"""Streamlit 前端共享配置与工具。
"""
import os
import httpx
import streamlit as st
# ============ API ============
API_URL = os.environ.get("API_URL", "http://localhost:8000")
# ============ 选项与标签 ============
AVAILABLE_ORGANS = [
"liver",
"spleen",
"lung",
"heart",
"kidney",
"muscle",
"lymph_nodes",
]
ORGAN_LABELS = {
"liver": "肝脏 (Liver)",
"spleen": "脾脏 (Spleen)",
"lung": "肺 (Lung)",
"heart": "心脏 (Heart)",
"kidney": "肾脏 (Kidney)",
"muscle": "肌肉 (Muscle)",
"lymph_nodes": "淋巴结 (Lymph Nodes)",
}
AVAILABLE_ROUTES = [
"intravenous",
"intramuscular",
]
ROUTE_LABELS = {
"intravenous": "静脉注射 (Intravenous)",
"intramuscular": "肌肉注射 (Intramuscular)",
}
HELPER_LIPID_OPTIONS = ["DOPE", "DSPC"]
# PDI 分类标签
PDI_CLASS_LABELS = {
0: "<0.2 (优)",
1: "≥0.2 (欠佳)",
}
EE_CLASS_LABELS = {
0: "<50% (低)",
1: "50-80% (中)",
2: ">80% (高)",
}
TOXIC_CLASS_LABELS = {
0: "无毒 ✓",
1: "有毒 ⚠",
}
# ============ 功能模式 ============
MODE_OPTIMIZE = "optimize"
MODE_SCREEN = "screen"
MODE_LABELS = {
MODE_OPTIMIZE: "配方优选",
MODE_SCREEN: "批量筛选",
}
# ============ 访问口令 ============
def check_password() -> bool:
"""校验访问口令,未设置 APP_PASSWORD 时直接放行,内网访问不受影响。"""
expected = os.environ.get("APP_PASSWORD")
if not expected:
return True
if st.session_state.get("_authed"):
return True
gate = st.empty()
with gate.container():
st.markdown("### LNP 配方优化")
pwd = st.text_input("🔒 访问口令", type="password")
if pwd and pwd != expected:
st.error("口令不正确")
if pwd == expected:
st.session_state["_authed"] = True
gate.empty()
return True
return False
def require_auth() -> None:
"""口令未通过时终止页面渲染。应用只有 app.py 一个入口,在其顶部调用一次即可。"""
if not check_password():
st.stop()
# ============ API 客户端 ============
def check_api_status(timeout: float = 5.0) -> bool:
"""检查 API 是否在线。"""
try:
with httpx.Client(timeout=timeout) as client:
return client.get(f"{API_URL}/").status_code == 200
except httpx.HTTPError:
return False
def render_api_status(api_online: bool) -> None:
"""在侧边栏渲染 API 状态与模型能力,两种功能模式共用。"""
if not api_online:
st.error("🔴 API 服务离线")
st.info(
f"请先启动 API 服务:\n```\nuvicorn app.api:app --port 8010\n```\n"
f"当前 API_URL: {API_URL}"
)
return
st.success("🟢 API 服务在线")
try:
with httpx.Client(timeout=5) as client:
info = client.get(f"{API_URL}/").json()
except httpx.HTTPError:
return
caps = [
name
for name, key in (("MoE", "use_moe"), ("LLM", "use_llm"), ("RAG", "use_rag"))
if info.get(key)
]
st.caption(
f"模型: {' + '.join(caps) if caps else '仅 backbone'}|{info.get('device', '?')}"
)

View File

@ -20,11 +20,11 @@ requires-python = ">=3.8"
[tool.ruff]
line-length = 99
src = ["lnp_ml"]
include = ["pyproject.toml", "lnp_ml/**/*.py"]
src = ["lnp_ml", "app"]
include = ["pyproject.toml", "lnp_ml/**/*.py", "app/**/*.py"]
[tool.ruff.lint]
extend-select = ["I"] # Add import sorting
select = ["E4", "E7", "E9", "F", "I"]
[tool.ruff.lint.isort]
known-first-party = ["lnp_ml"]