From a8126945008ca82eeb9a668c979fe814cc9e4dba Mon Sep 17 00:00:00 2001 From: Michelle0574 <2170308303@qq.com> Date: Sat, 5 Sep 2026 17:15:37 +0000 Subject: [PATCH] =?UTF-8?q?feat(app):=20=E6=96=B0=E5=A2=9E=E6=89=B9?= =?UTF-8?q?=E9=87=8F=E7=AD=9B=E9=80=89=EF=BC=8C=E7=BB=93=E6=9E=9C=E4=B8=8E?= =?UTF-8?q?=E6=A8=A1=E5=BC=8F=E5=8F=AF=E8=B7=A8=E5=88=B7=E6=96=B0=E6=81=A2?= =?UTF-8?q?=E5=A4=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/api.py | 155 ++++++++++++-- app/app.py | 131 +++++------- app/batch_screening.py | 466 +++++++++++++++++++++++++++++++++++++++++ app/optimize.py | 14 +- app/ui_common.py | 141 +++++++++++++ pyproject.toml | 6 +- 6 files changed, 804 insertions(+), 109 deletions(-) create mode 100644 app/batch_screening.py create mode 100644 app/ui_common.py diff --git a/app/api.py b/app/api.py index 426f366..9dee506 100644 --- a/app/api.py +++ b/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 # 转换结果 diff --git a/app/app.py b/app/app.py index 6939c67..d11cac0 100644 --- a/app/app.py +++ b/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", diff --git a/app/batch_screening.py b/app/batch_screening.py new file mode 100644 index 0000000..1a92d3c --- /dev/null +++ b/app/batch_screening.py @@ -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"]) \ No newline at end of file diff --git a/app/optimize.py b/app/optimize.py index ff0ebf5..a37f6be 100644 --- a/app/optimize.py +++ b/app/optimize.py @@ -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 diff --git a/app/ui_common.py b/app/ui_common.py new file mode 100644 index 0000000..84ed255 --- /dev/null +++ b/app/ui_common.py @@ -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', '?')}" + ) \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index 86b7890..698dcff 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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"]