mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-10-04 23:19:44 +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
|
uvicorn app.api:app --host 0.0.0.0 --port 8000 --reload
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from contextlib import asynccontextmanager
|
||||||
import os
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import List, Dict, Optional
|
from typing import Dict, List, Optional
|
||||||
from contextlib import asynccontextmanager
|
|
||||||
|
|
||||||
import torch
|
|
||||||
from fastapi import FastAPI, HTTPException
|
from fastapi import FastAPI, HTTPException
|
||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
from pydantic import BaseModel, Field
|
|
||||||
from loguru import logger
|
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 (
|
from app.optimize import (
|
||||||
optimize,
|
|
||||||
format_results,
|
|
||||||
AVAILABLE_ORGANS,
|
AVAILABLE_ORGANS,
|
||||||
|
HELPER_LIPID_OPTIONS,
|
||||||
|
ROUTE_OPTIONS,
|
||||||
TARGET_BIODIST,
|
TARGET_BIODIST,
|
||||||
CompRanges,
|
CompRanges,
|
||||||
ScoringWeights,
|
ScoringWeights,
|
||||||
HELPER_LIPID_OPTIONS,
|
|
||||||
ROUTE_OPTIONS,
|
|
||||||
create_dataframe_from_formulations,
|
create_dataframe_from_formulations,
|
||||||
|
optimize,
|
||||||
predict_all,
|
predict_all,
|
||||||
)
|
)
|
||||||
|
from lnp_ml.config import MODELS_DIR
|
||||||
|
from lnp_ml.modeling.predict import load_model
|
||||||
|
|
||||||
# ============ Pydantic Models ============
|
# ============ 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 ============
|
# ============ Global State ============
|
||||||
|
|
||||||
class ModelState:
|
class ModelState:
|
||||||
@ -266,6 +290,7 @@ async def lifespan(app: FastAPI):
|
|||||||
if _llm is not None and getattr(_llm, "use_rag", False):
|
if _llm is not None and getattr(_llm, "use_rag", False):
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
|
|
||||||
from lnp_ml.dataset import LNPDataset, process_dataframe
|
from lnp_ml.dataset import LNPDataset, process_dataframe
|
||||||
from lnp_ml.modeling.nested_cv_optuna import _build_rag_pool
|
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)
|
@app.post("/optimize", response_model=OptimizeResponse)
|
||||||
async def optimize_formulation(request: OptimizeRequest):
|
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
|
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 环境变量:
|
Docker 环境变量:
|
||||||
API_URL: API 服务地址 (默认: http://localhost:8000)
|
API_URL: API 服务地址 (默认: http://localhost:8000)
|
||||||
|
APP_PASSWORD: 访问口令,设置后进入页面需先输入;不设则不启用
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import io
|
|
||||||
import os
|
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import streamlit as st
|
import streamlit as st
|
||||||
|
|
||||||
# ============ 配置 ============
|
import batch_screening
|
||||||
|
from ui_common import (
|
||||||
# 从环境变量读取 API 地址,支持 Docker 环境
|
API_URL,
|
||||||
API_URL = os.environ.get("API_URL", "http://localhost:8000")
|
AVAILABLE_ORGANS,
|
||||||
|
AVAILABLE_ROUTES,
|
||||||
AVAILABLE_ORGANS = [
|
EE_CLASS_LABELS,
|
||||||
"liver",
|
MODE_LABELS,
|
||||||
"spleen",
|
MODE_OPTIMIZE,
|
||||||
"lung",
|
MODE_SCREEN,
|
||||||
"heart",
|
ORGAN_LABELS,
|
||||||
"kidney",
|
PDI_CLASS_LABELS,
|
||||||
"muscle",
|
ROUTE_LABELS,
|
||||||
"lymph_nodes",
|
TOXIC_CLASS_LABELS,
|
||||||
]
|
check_api_status,
|
||||||
|
render_api_status,
|
||||||
ORGAN_LABELS = {
|
require_auth,
|
||||||
"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)",
|
|
||||||
}
|
|
||||||
|
|
||||||
# ============ 页面配置 ============
|
# ============ 页面配置 ============
|
||||||
|
|
||||||
@ -60,6 +42,10 @@ st.set_page_config(
|
|||||||
initial_sidebar_state="expanded",
|
initial_sidebar_state="expanded",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# ============ 访问口令 ============
|
||||||
|
|
||||||
|
require_auth()
|
||||||
|
|
||||||
# ============ 自定义样式 ============
|
# ============ 自定义样式 ============
|
||||||
|
|
||||||
st.markdown("""
|
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(
|
def call_optimize_api(
|
||||||
smiles: str,
|
smiles: str,
|
||||||
@ -174,26 +151,6 @@ def call_optimize_api(
|
|||||||
return response.json()
|
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:
|
def format_results_dataframe(results: dict, smiles_label: str = None) -> pd.DataFrame:
|
||||||
"""将 API 结果转换为 DataFrame"""
|
"""将 API 结果转换为 DataFrame"""
|
||||||
formulations = results["formulations"]
|
formulations = results["formulations"]
|
||||||
@ -265,26 +222,32 @@ def main():
|
|||||||
# 检查 API 状态
|
# 检查 API 状态
|
||||||
api_online = check_api_status()
|
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:
|
with st.sidebar:
|
||||||
# st.header("⚙️ 参数设置")
|
# 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()
|
# st.divider()
|
||||||
|
|
||||||
# SMILES 输入
|
# SMILES 输入
|
||||||
@ -579,7 +542,7 @@ def main():
|
|||||||
|
|
||||||
# 优化按钮
|
# 优化按钮
|
||||||
optimize_button = st.button(
|
optimize_button = st.button(
|
||||||
"🚀 开始配方优选",
|
"开始配方优选",
|
||||||
type="primary",
|
type="primary",
|
||||||
use_container_width=True,
|
use_container_width=True,
|
||||||
disabled=not api_online or not smiles_input.strip() or not selected_routes,
|
disabled=not api_online or not smiles_input.strip() or not selected_routes,
|
||||||
@ -636,7 +599,7 @@ def main():
|
|||||||
except httpx.HTTPStatusError as e:
|
except httpx.HTTPStatusError as e:
|
||||||
try:
|
try:
|
||||||
error_detail = e.response.json().get("detail", str(e))
|
error_detail = e.response.json().get("detail", str(e))
|
||||||
except:
|
except ValueError:
|
||||||
error_detail = str(e)
|
error_detail = str(e)
|
||||||
errors.append(f"SMILES {idx + 1}: {error_detail}")
|
errors.append(f"SMILES {idx + 1}: {error_detail}")
|
||||||
except httpx.RequestError as e:
|
except httpx.RequestError as e:
|
||||||
@ -740,7 +703,7 @@ def main():
|
|||||||
target_organ = results["target_organ"]
|
target_organ = results["target_organ"]
|
||||||
|
|
||||||
st.download_button(
|
st.download_button(
|
||||||
label="📥 导出 CSV",
|
label="导出 CSV",
|
||||||
data=csv_content,
|
data=csv_content,
|
||||||
file_name=f"lnp_optimization_{target_organ}_{datetime.now().strftime('%Y%m%d_%H%M%S')}.csv",
|
file_name=f"lnp_optimization_{target_organ}_{datetime.now().strftime('%Y%m%d_%H%M%S')}.csv",
|
||||||
mime="text/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
|
python -m app.optimize --smiles "CC(C)..." --organ liver
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import itertools
|
from dataclasses import dataclass, field
|
||||||
import json
|
import json
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Dict, List, Optional, Tuple
|
from typing import Dict, List, Optional, Tuple
|
||||||
from dataclasses import dataclass, field
|
|
||||||
|
|
||||||
|
from loguru import logger
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import torch
|
import torch
|
||||||
from torch.utils.data import DataLoader
|
from torch.utils.data import DataLoader
|
||||||
from loguru import logger
|
|
||||||
from tqdm import tqdm
|
|
||||||
import typer
|
import typer
|
||||||
|
|
||||||
from lnp_ml.config import MODELS_DIR
|
from lnp_ml.config import MODELS_DIR
|
||||||
from lnp_ml.dataset import (
|
from lnp_ml.dataset import (
|
||||||
LNPDataset,
|
|
||||||
LNPDatasetConfig,
|
|
||||||
collate_fn,
|
|
||||||
SMILES_COL,
|
SMILES_COL,
|
||||||
COMP_COLS,
|
|
||||||
HELP_COLS,
|
|
||||||
TARGET_BIODIST,
|
TARGET_BIODIST,
|
||||||
get_phys_cols,
|
LNPDataset,
|
||||||
|
collate_fn,
|
||||||
get_exp_cols,
|
get_exp_cols,
|
||||||
)
|
)
|
||||||
from lnp_ml.modeling.predict import load_model
|
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]
|
[tool.ruff]
|
||||||
line-length = 99
|
line-length = 99
|
||||||
src = ["lnp_ml"]
|
src = ["lnp_ml", "app"]
|
||||||
include = ["pyproject.toml", "lnp_ml/**/*.py"]
|
include = ["pyproject.toml", "lnp_ml/**/*.py", "app/**/*.py"]
|
||||||
|
|
||||||
[tool.ruff.lint]
|
[tool.ruff.lint]
|
||||||
extend-select = ["I"] # Add import sorting
|
select = ["E4", "E7", "E9", "F", "I"]
|
||||||
|
|
||||||
[tool.ruff.lint.isort]
|
[tool.ruff.lint.isort]
|
||||||
known-first-party = ["lnp_ml"]
|
known-first-party = ["lnp_ml"]
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user