mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-10-01 13:25:46 +08:00
feat(app): 新增批量筛选,结果与模式可跨刷新恢复
This commit is contained in:
parent
9810854ac7
commit
a812694500
155
app/api.py
155
app/api.py
@ -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
|
||||
|
||||
# 转换结果
|
||||
|
||||
131
app/app.py
131
app/app.py
@ -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
466
app/batch_screening.py
Normal 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"])
|
||||
@ -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
141
app/ui_common.py
Normal 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', '?')}"
|
||||
)
|
||||
@ -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"]
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user