mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-10-05 15:33:12 +08:00
Compare commits
3 Commits
9810854ac7
...
c0d2c008b0
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c0d2c008b0 | ||
|
|
a7c7d54c61 | ||
|
|
a812694500 |
@ -69,6 +69,16 @@ logs/
|
|||||||
models/finetune_cv/
|
models/finetune_cv/
|
||||||
models/benchmark/
|
models/benchmark/
|
||||||
models/mpnn/
|
models/mpnn/
|
||||||
|
models/qwen2.5-7b-instruct/
|
||||||
|
models/cv5_sample/
|
||||||
|
models/molt5-base/
|
||||||
|
models/nested_cv/
|
||||||
|
models/pretrain/
|
||||||
|
MolE_ckpt/
|
||||||
|
MolE_repo/
|
||||||
|
models/abl/
|
||||||
|
models/biot5-plus-base/
|
||||||
|
models/abl_full/
|
||||||
models/*.pt
|
models/*.pt
|
||||||
models/*.json
|
models/*.json
|
||||||
!models/final/
|
!models/final/
|
||||||
|
|||||||
@ -36,8 +36,8 @@ COPY app/ ./app/
|
|||||||
# 安装项目包
|
# 安装项目包
|
||||||
RUN pip install -e .
|
RUN pip install -e .
|
||||||
|
|
||||||
# 复制模型文件
|
# 模型、Qwen 权重和 RAG 数据在运行时以只读 volume 挂载。
|
||||||
COPY models/final/ ./models/final/
|
# 不把数 GB 权重烘进镜像,也避免误把 Git LFS pointer 当成模型。
|
||||||
|
|
||||||
# ============ API 服务 ============
|
# ============ API 服务 ============
|
||||||
FROM base AS api
|
FROM base AS api
|
||||||
@ -60,4 +60,3 @@ ENV STREAMLIT_SERVER_PORT=8501 \
|
|||||||
STREAMLIT_BROWSER_GATHER_USAGE_STATS=false
|
STREAMLIT_BROWSER_GATHER_USAGE_STATS=false
|
||||||
|
|
||||||
CMD ["streamlit", "run", "app/app.py"]
|
CMD ["streamlit", "run", "app/app.py"]
|
||||||
|
|
||||||
|
|||||||
59
SERVER_DEPLOY.md
Normal file
59
SERVER_DEPLOY.md
Normal file
@ -0,0 +1,59 @@
|
|||||||
|
# GPU 服务器测试与部署
|
||||||
|
|
||||||
|
这个代码包不包含大模型权重。解压后,请先在项目根目录放好以下运行时文件:
|
||||||
|
|
||||||
|
```text
|
||||||
|
models/final/model.pt
|
||||||
|
models/mpnn/all_amine_split_for_LiON/cv_0/fold_0/model_0/model.pt
|
||||||
|
models/mpnn/all_amine_split_for_LiON/cv_1/fold_0/model_0/model.pt
|
||||||
|
models/mpnn/all_amine_split_for_LiON/cv_2/fold_0/model_0/model.pt
|
||||||
|
models/mpnn/all_amine_split_for_LiON/cv_3/fold_0/model_0/model.pt
|
||||||
|
models/mpnn/all_amine_split_for_LiON/cv_4/fold_0/model_0/model.pt
|
||||||
|
models/qwen2.5-7b-instruct/config.json
|
||||||
|
models/qwen2.5-7b-instruct/ # 其余 tokenizer/权重文件
|
||||||
|
data/interim/internal.csv
|
||||||
|
```
|
||||||
|
|
||||||
|
`.pt` 必须是真实二进制权重,不能是 100 多字节的 Git LFS pointer。
|
||||||
|
|
||||||
|
## 1. 服务器预检
|
||||||
|
|
||||||
|
```bash
|
||||||
|
nvidia-smi
|
||||||
|
docker compose -f docker-compose-gpu.yml config
|
||||||
|
bash server_gpu_test.sh preflight
|
||||||
|
```
|
||||||
|
|
||||||
|
## 2. 构建并启动测试服务
|
||||||
|
|
||||||
|
```bash
|
||||||
|
LLM_INFERENCE_BATCH_SIZE=4 bash server_gpu_test.sh start
|
||||||
|
```
|
||||||
|
|
||||||
|
对 24GB 显卡可将微批设为 `2`;48GB/80GB 显卡通常可使用 `4`。这个参数只控制 Qwen 内部微批,不会限制 API 一次提交的条数。
|
||||||
|
|
||||||
|
## 3. 单条与 32 条 LLM 冒烟测试
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bash server_gpu_test.sh smoke
|
||||||
|
```
|
||||||
|
|
||||||
|
另开一个终端观察峰值:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
watch -n 0.5 nvidia-smi
|
||||||
|
```
|
||||||
|
|
||||||
|
预期:两个请求都返回 HTTP 200,日志中出现 `LLM inference micro-batching`,32 条请求不再把 32 条同时送入 Qwen。
|
||||||
|
|
||||||
|
## 4. 确认后上线
|
||||||
|
|
||||||
|
如果当前机器就是正式机,`start` 启动的已是 Compose 中的正式服务。检查:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
docker compose -f docker-compose-gpu.yml ps
|
||||||
|
docker compose -f docker-compose-gpu.yml logs --tail=200 api
|
||||||
|
curl -fsS http://127.0.0.1:18000/
|
||||||
|
```
|
||||||
|
|
||||||
|
回滚时保留旧代码目录,然后在旧目录里重新执行 `docker compose up -d --build`。
|
||||||
161
app/api.py
161
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,37 @@ 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="外层模型前向批大小;7B LLM 会按显存安全值内部微批",
|
||||||
|
)
|
||||||
|
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 +295,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 +439,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 +617,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
|
||||||
|
|
||||||
# 转换结果
|
# 转换结果
|
||||||
@ -537,4 +673,3 @@ if __name__ == "__main__":
|
|||||||
port=8000,
|
port=8000,
|
||||||
reload=True,
|
reload=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
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
|
||||||
@ -577,7 +571,9 @@ def predict_all(
|
|||||||
all_ee_preds = []
|
all_ee_preds = []
|
||||||
all_toxic_preds = []
|
all_toxic_preds = []
|
||||||
|
|
||||||
with torch.no_grad():
|
# inference_mode 比 no_grad 还会关闭 version counter/view tracking,
|
||||||
|
# 对大批量 LLM 推理可再减少一部分临时内存。
|
||||||
|
with torch.inference_mode():
|
||||||
for batch in dataloader:
|
for batch in dataloader:
|
||||||
smiles = batch["smiles"]
|
smiles = batch["smiles"]
|
||||||
tabular = {k: v.to(device) for k, v in batch["tabular"].items()}
|
tabular = {k: v.to(device) for k, v in batch["tabular"].items()}
|
||||||
@ -1151,4 +1147,3 @@ def main(
|
|||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
app()
|
app()
|
||||||
|
|
||||||
|
|||||||
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', '?')}"
|
||||||
|
)
|
||||||
@ -6,12 +6,20 @@ services:
|
|||||||
dockerfile: Dockerfile
|
dockerfile: Dockerfile
|
||||||
target: api
|
target: api
|
||||||
container_name: lnp-api
|
container_name: lnp-api
|
||||||
|
ports:
|
||||||
|
# 仅绑定本机,供服务器上的冒烟测试使用,不直接暴露到公网。
|
||||||
|
- "127.0.0.1:${API_PORT:-18000}:8000"
|
||||||
environment:
|
environment:
|
||||||
- MODEL_PATH=/app/models/final/model.pt
|
- MODEL_PATH=/app/models/final/model.pt
|
||||||
|
- RAG_POOL_CSV=/app/data/interim/internal.csv
|
||||||
|
# 7B LLM 内部微批;外层批量筛选仍可一次提交更多条。
|
||||||
|
- LLM_INFERENCE_BATCH_SIZE=${LLM_INFERENCE_BATCH_SIZE:-4}
|
||||||
volumes:
|
volumes:
|
||||||
# 挂载模型目录以便更新模型
|
# 权重和 RAG 数据不进镜像,服务器上需预先放到这些路径。
|
||||||
- ./models/final:/app/models/final:ro
|
- ./models/final:/app/models/final:ro
|
||||||
- ./models/mpnn:/app/models/mpnn:ro
|
- ./models/mpnn:/app/models/mpnn:ro
|
||||||
|
- ./models/qwen2.5-7b-instruct:/app/models/qwen2.5-7b-instruct:ro
|
||||||
|
- ./data/interim/internal.csv:/app/data/interim/internal.csv:ro
|
||||||
restart: unless-stopped
|
restart: unless-stopped
|
||||||
healthcheck:
|
healthcheck:
|
||||||
test: ["CMD", "curl", "-f", "http://localhost:8000/"]
|
test: ["CMD", "curl", "-f", "http://localhost:8000/"]
|
||||||
|
|||||||
@ -9,6 +9,16 @@ import torch.nn as nn
|
|||||||
|
|
||||||
DEFAULT_MOLT5_PATH = os.environ.get("MOLT5_PATH", "models/molt5-base")
|
DEFAULT_MOLT5_PATH = os.environ.get("MOLT5_PATH", "models/molt5-base")
|
||||||
|
|
||||||
|
|
||||||
|
def _positive_int_env(name: str, default: int) -> int:
|
||||||
|
"""Read a positive integer environment variable with a safe fallback."""
|
||||||
|
try:
|
||||||
|
value = int(os.environ.get(name, str(default)))
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return default
|
||||||
|
return value if value > 0 else default
|
||||||
|
|
||||||
|
|
||||||
# 检索池支持的额外多任务(除 delivery 外)
|
# 检索池支持的额外多任务(除 delivery 外)
|
||||||
_EXTRA_TASKS = ["size", "pdi", "ee", "toxic", "biodist"]
|
_EXTRA_TASKS = ["size", "pdi", "ee", "toxic", "biodist"]
|
||||||
|
|
||||||
@ -58,6 +68,9 @@ class LLMPromptEncoder(nn.Module):
|
|||||||
self.use_rag = use_rag
|
self.use_rag = use_rag
|
||||||
self.rag_top_k = rag_top_k
|
self.rag_top_k = rag_top_k
|
||||||
self.use_soft_prompt = use_soft_prompt
|
self.use_soft_prompt = use_soft_prompt
|
||||||
|
# 7B 模型即使在 no_grad 下,大批量长 prompt 的临时激活也很大。
|
||||||
|
# 外层 DataLoader 仍可用大 batch;这里只把 LLM 前向切成小块。
|
||||||
|
self.inference_batch_size = _positive_int_env("LLM_INFERENCE_BATCH_SIZE", 4)
|
||||||
|
|
||||||
_name_l = model_name_or_path.lower()
|
_name_l = model_name_or_path.lower()
|
||||||
_is_t5 = "t5" in _name_l
|
_is_t5 = "t5" in _name_l
|
||||||
@ -145,16 +158,18 @@ class LLMPromptEncoder(nn.Module):
|
|||||||
def _apply_lora(self, r, alpha, dropout, prepare_kbit=False):
|
def _apply_lora(self, r, alpha, dropout, prepare_kbit=False):
|
||||||
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
|
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
|
||||||
|
|
||||||
if prepare_kbit:
|
# 4-bit/8-bit 基座只 prepare 一次。PEFT 会处理输入梯度并开启
|
||||||
|
# gradient checkpointing;后面再显式开一次,兼容不同 PEFT 版本。
|
||||||
|
is_kbit = (
|
||||||
|
prepare_kbit
|
||||||
|
or getattr(self.encoder, "is_loaded_in_4bit", False)
|
||||||
|
or getattr(self.encoder, "is_loaded_in_8bit", False)
|
||||||
|
)
|
||||||
|
if is_kbit:
|
||||||
self.encoder = prepare_model_for_kbit_training(self.encoder)
|
self.encoder = prepare_model_for_kbit_training(self.encoder)
|
||||||
else:
|
else:
|
||||||
for p in self.encoder.parameters():
|
for p in self.encoder.parameters():
|
||||||
p.requires_grad = False
|
p.requires_grad = False
|
||||||
|
|
||||||
# 4-bit 量化模型(QLoRA)需先 prepare,才能正确接收梯度
|
|
||||||
if getattr(self.encoder, "is_loaded_in_4bit", False) or getattr(self.encoder, "is_loaded_in_8bit", False):
|
|
||||||
from peft import prepare_model_for_kbit_training
|
|
||||||
self.encoder = prepare_model_for_kbit_training(self.encoder)
|
|
||||||
for p in self.encoder.parameters():
|
for p in self.encoder.parameters():
|
||||||
p.requires_grad = False
|
p.requires_grad = False
|
||||||
# 不同架构注意力层命名不同:
|
# 不同架构注意力层命名不同:
|
||||||
@ -170,6 +185,24 @@ class LLMPromptEncoder(nn.Module):
|
|||||||
target_modules=_targets, bias="none")
|
target_modules=_targets, bias="none")
|
||||||
self.encoder = get_peft_model(self.encoder, cfg)
|
self.encoder = get_peft_model(self.encoder, cfg)
|
||||||
|
|
||||||
|
# Decoder-only LLM 做表征编码不需要 KV cache。训练时保留 cache
|
||||||
|
# 只会额外占显存;梯度检查点则用计算换显存。
|
||||||
|
if hasattr(self.encoder, "config"):
|
||||||
|
self.encoder.config.use_cache = False
|
||||||
|
if hasattr(self.encoder, "gradient_checkpointing_enable"):
|
||||||
|
try:
|
||||||
|
self.encoder.gradient_checkpointing_enable(
|
||||||
|
gradient_checkpointing_kwargs={"use_reentrant": False}
|
||||||
|
)
|
||||||
|
except TypeError:
|
||||||
|
self.encoder.gradient_checkpointing_enable()
|
||||||
|
|
||||||
|
def _encoder_forward(self, **kwargs):
|
||||||
|
"""Forward the backbone without allocating an unused decoder KV cache."""
|
||||||
|
if self._is_qwen:
|
||||||
|
kwargs["use_cache"] = False
|
||||||
|
return self.encoder(**kwargs)
|
||||||
|
|
||||||
# ---------- 检索池 ----------
|
# ---------- 检索池 ----------
|
||||||
def set_retrieval_pool(self, smiles_list, labels, pool_id="train", extra_labels=None):
|
def set_retrieval_pool(self, smiles_list, labels, pool_id="train", extra_labels=None):
|
||||||
"""设置 RAG 检索池(防泄漏:只传训练集分子与标签)。
|
"""设置 RAG 检索池(防泄漏:只传训练集分子与标签)。
|
||||||
@ -326,14 +359,15 @@ class LLMPromptEncoder(nn.Module):
|
|||||||
def _rag_encode_batch(self, prompts, device):
|
def _rag_encode_batch(self, prompts, device):
|
||||||
"""编码一批 RAG prompt,取最后有效 token。grad 由调用方上下文决定。"""
|
"""编码一批 RAG prompt,取最后有效 token。grad 由调用方上下文决定。"""
|
||||||
outs = []
|
outs = []
|
||||||
for i in range(0, len(prompts), 4):
|
chunk_size = self.inference_batch_size
|
||||||
bt = prompts[i:i+4]
|
for i in range(0, len(prompts), chunk_size):
|
||||||
|
bt = prompts[i:i + chunk_size]
|
||||||
enc = self.tokenizer(
|
enc = self.tokenizer(
|
||||||
bt, padding=True, truncation=True,
|
bt, padding=True, truncation=True,
|
||||||
max_length=self.max_length, return_tensors="pt",
|
max_length=self.max_length, return_tensors="pt",
|
||||||
).to(device)
|
).to(device)
|
||||||
self._warn_if_truncated(enc)
|
self._warn_if_truncated(enc)
|
||||||
out = self.encoder(**enc).last_hidden_state # [B,L,H]
|
out = self._encoder_forward(**enc).last_hidden_state # [B,L,H]
|
||||||
lengths = enc["attention_mask"].sum(1) - 1
|
lengths = enc["attention_mask"].sum(1) - 1
|
||||||
b = torch.arange(out.size(0), device=device)
|
b = torch.arange(out.size(0), device=device)
|
||||||
outs.append(out[b, lengths.long(), :].float()) # [B,H]
|
outs.append(out[b, lengths.long(), :].float()) # [B,H]
|
||||||
@ -359,7 +393,8 @@ class LLMPromptEncoder(nn.Module):
|
|||||||
return self._rag_encode_batch(prompts, device)
|
return self._rag_encode_batch(prompts, device)
|
||||||
|
|
||||||
# ---------- soft-prompt 编码(带梯度,不缓存特征)----------
|
# ---------- soft-prompt 编码(带梯度,不缓存特征)----------
|
||||||
def _encode_softrag(self, smiles, chem, tab, device) -> torch.Tensor:
|
def _encode_softrag_chunk(self, smiles, chem, tab, device) -> torch.Tensor:
|
||||||
|
"""单个 LLM micro-batch 的 soft-RAG 前向。"""
|
||||||
prompts = [self._get_prompt(s) for s in smiles]
|
prompts = [self._get_prompt(s) for s in smiles]
|
||||||
enc = self.tokenizer(
|
enc = self.tokenizer(
|
||||||
prompts, padding=True, truncation=True,
|
prompts, padding=True, truncation=True,
|
||||||
@ -393,7 +428,7 @@ class LLMPromptEncoder(nn.Module):
|
|||||||
position_ids = attn_mask.long().cumsum(-1) - 1
|
position_ids = attn_mask.long().cumsum(-1) - 1
|
||||||
position_ids = position_ids.masked_fill(attn_mask == 0, 1)
|
position_ids = position_ids.masked_fill(attn_mask == 0, 1)
|
||||||
|
|
||||||
out = self.encoder(
|
out = self._encoder_forward(
|
||||||
inputs_embeds=inputs_embeds,
|
inputs_embeds=inputs_embeds,
|
||||||
attention_mask=attn_mask,
|
attention_mask=attn_mask,
|
||||||
position_ids=position_ids,
|
position_ids=position_ids,
|
||||||
@ -407,6 +442,43 @@ class LLMPromptEncoder(nn.Module):
|
|||||||
feat = out[b, lengths.long(), :]
|
feat = out[b, lengths.long(), :]
|
||||||
return self.proj_down(feat.float())
|
return self.proj_down(feat.float())
|
||||||
|
|
||||||
|
def _encode_softrag(self, smiles, chem, tab, device) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
soft-RAG 编码。推理时自动微批,避免 batch API 把 7B LLM 的
|
||||||
|
临时激活按请求 batch_size 线性放大。训练仍保持原 batch,
|
||||||
|
由 gradient checkpointing 降低激活显存。
|
||||||
|
"""
|
||||||
|
batch_size = len(smiles)
|
||||||
|
micro_batch = self.inference_batch_size
|
||||||
|
should_chunk = (
|
||||||
|
not self.training
|
||||||
|
and not torch.is_grad_enabled()
|
||||||
|
and batch_size > micro_batch
|
||||||
|
)
|
||||||
|
if not should_chunk:
|
||||||
|
return self._encode_softrag_chunk(smiles, chem, tab, device)
|
||||||
|
|
||||||
|
if not getattr(self, "_microbatch_logged", False):
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
f"LLM inference micro-batching: batch={batch_size}, "
|
||||||
|
f"micro_batch={micro_batch} (override with LLM_INFERENCE_BATCH_SIZE)"
|
||||||
|
)
|
||||||
|
self._microbatch_logged = True
|
||||||
|
|
||||||
|
outputs = []
|
||||||
|
for start in range(0, batch_size, micro_batch):
|
||||||
|
end = min(start + micro_batch, batch_size)
|
||||||
|
chem_chunk = chem[start:end] if chem is not None else None
|
||||||
|
tab_chunk = tab[start:end] if tab is not None else None
|
||||||
|
outputs.append(
|
||||||
|
self._encode_softrag_chunk(
|
||||||
|
smiles[start:end], chem_chunk, tab_chunk, device
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return torch.cat(outputs, dim=0)
|
||||||
|
|
||||||
# ---------- 旧路径(保留,向后兼容)----------
|
# ---------- 旧路径(保留,向后兼容)----------
|
||||||
def _mean_pool(self, last_hidden, mask):
|
def _mean_pool(self, last_hidden, mask):
|
||||||
m = mask.unsqueeze(-1).float()
|
m = mask.unsqueeze(-1).float()
|
||||||
@ -417,12 +489,13 @@ class LLMPromptEncoder(nn.Module):
|
|||||||
missing = [s for s in smiles if s not in self._cache]
|
missing = [s for s in smiles if s not in self._cache]
|
||||||
if missing:
|
if missing:
|
||||||
uniq = list(dict.fromkeys(missing))
|
uniq = list(dict.fromkeys(missing))
|
||||||
for i in range(0, len(uniq), 256):
|
chunk_size = self.inference_batch_size if self._is_qwen else 256
|
||||||
chunk = uniq[i:i + 256]
|
for i in range(0, len(uniq), chunk_size):
|
||||||
|
chunk = uniq[i:i + chunk_size]
|
||||||
mols = [self._fmt_mol(s) for s in chunk]
|
mols = [self._fmt_mol(s) for s in chunk]
|
||||||
enc = self.tokenizer(mols, padding=True, truncation=True,
|
enc = self.tokenizer(mols, padding=True, truncation=True,
|
||||||
max_length=self.max_length, return_tensors="pt").to(device)
|
max_length=self.max_length, return_tensors="pt").to(device)
|
||||||
out = self.encoder(**enc).last_hidden_state
|
out = self._encoder_forward(**enc).last_hidden_state
|
||||||
pooled = self._mean_pool(out, enc["attention_mask"])
|
pooled = self._mean_pool(out, enc["attention_mask"])
|
||||||
for s, v in zip(chunk, pooled):
|
for s, v in zip(chunk, pooled):
|
||||||
self._cache[s] = v.cpu()
|
self._cache[s] = v.cpu()
|
||||||
@ -432,7 +505,7 @@ class LLMPromptEncoder(nn.Module):
|
|||||||
mols = [self._fmt_mol(s) for s in smiles]
|
mols = [self._fmt_mol(s) for s in smiles]
|
||||||
enc = self.tokenizer(mols, padding=True, truncation=True,
|
enc = self.tokenizer(mols, padding=True, truncation=True,
|
||||||
max_length=self.max_length, return_tensors="pt").to(device)
|
max_length=self.max_length, return_tensors="pt").to(device)
|
||||||
out = self.encoder(**enc).last_hidden_state
|
out = self._encoder_forward(**enc).last_hidden_state
|
||||||
return self._mean_pool(out, enc["attention_mask"])
|
return self._mean_pool(out, enc["attention_mask"])
|
||||||
|
|
||||||
def forward(self, smiles: List[str], chem: Optional[torch.Tensor] = None,
|
def forward(self, smiles: List[str], chem: Optional[torch.Tensor] = None,
|
||||||
@ -448,4 +521,4 @@ class LLMPromptEncoder(nn.Module):
|
|||||||
|
|
||||||
def clear_cache(self) -> None:
|
def clear_cache(self) -> None:
|
||||||
self._cache.clear()
|
self._cache.clear()
|
||||||
self._prompt_cache.clear()
|
self._prompt_cache.clear()
|
||||||
|
|||||||
@ -118,7 +118,7 @@ def load_model(
|
|||||||
return model
|
return model
|
||||||
|
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.inference_mode()
|
||||||
def predict_batch(
|
def predict_batch(
|
||||||
model: Union[LNPModel, LNPModelWithoutMPNN],
|
model: Union[LNPModel, LNPModelWithoutMPNN],
|
||||||
loader: DataLoader,
|
loader: DataLoader,
|
||||||
|
|||||||
@ -31,5 +31,6 @@ uvicorn = ">=0.33.0, <0.34"
|
|||||||
optuna = ">=4.5.0, <5"
|
optuna = ">=4.5.0, <5"
|
||||||
captum = ">=0.7.0, <0.8"
|
captum = ">=0.7.0, <0.8"
|
||||||
transformers = ">=4.30, <4.46"
|
transformers = ">=4.30, <4.46"
|
||||||
|
peft = "==0.13.2"
|
||||||
sentencepiece = "*"
|
sentencepiece = "*"
|
||||||
protobuf = "*"
|
protobuf = "*"
|
||||||
|
|||||||
@ -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"]
|
||||||
|
|||||||
@ -21,5 +21,8 @@ uvicorn>=0.33.0,<0.34
|
|||||||
optuna>=4.5.0,<5
|
optuna>=4.5.0,<5
|
||||||
captum>=0.7.0
|
captum>=0.7.0
|
||||||
transformers>=4.30,<4.46
|
transformers>=4.30,<4.46
|
||||||
|
peft==0.13.2
|
||||||
sentencepiece
|
sentencepiece
|
||||||
protobuf
|
protobuf
|
||||||
|
bitsandbytes
|
||||||
|
accelerate
|
||||||
|
|||||||
108
server_gpu_test.sh
Executable file
108
server_gpu_test.sh
Executable file
@ -0,0 +1,108 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
set -euo pipefail
|
||||||
|
|
||||||
|
ROOT_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)
|
||||||
|
cd "${ROOT_DIR}"
|
||||||
|
|
||||||
|
COMPOSE_FILE=${COMPOSE_FILE:-docker-compose-gpu.yml}
|
||||||
|
API_PORT=${API_PORT:-18000}
|
||||||
|
API_URL=${API_URL:-http://127.0.0.1:${API_PORT}}
|
||||||
|
LLM_INFERENCE_BATCH_SIZE=${LLM_INFERENCE_BATCH_SIZE:-4}
|
||||||
|
export API_PORT
|
||||||
|
export LLM_INFERENCE_BATCH_SIZE
|
||||||
|
|
||||||
|
fail() {
|
||||||
|
echo "ERROR: $*" >&2
|
||||||
|
exit 1
|
||||||
|
}
|
||||||
|
|
||||||
|
assert_real_file() {
|
||||||
|
local path=$1
|
||||||
|
[ -f "${path}" ] || fail "missing file: ${path}"
|
||||||
|
if head -n 1 "${path}" 2>/dev/null | grep -q "git-lfs.github.com/spec"; then
|
||||||
|
fail "${path} is a Git LFS pointer, not the real artifact"
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
preflight() {
|
||||||
|
command -v docker >/dev/null || fail "docker is not installed"
|
||||||
|
command -v nvidia-smi >/dev/null || fail "nvidia-smi is not installed"
|
||||||
|
nvidia-smi >/dev/null || fail "NVIDIA driver/GPU is unavailable"
|
||||||
|
docker compose version >/dev/null || fail "docker compose plugin is unavailable"
|
||||||
|
|
||||||
|
assert_real_file models/final/model.pt
|
||||||
|
assert_real_file data/interim/internal.csv
|
||||||
|
assert_real_file models/qwen2.5-7b-instruct/config.json
|
||||||
|
|
||||||
|
local fold
|
||||||
|
for fold in 0 1 2 3 4; do
|
||||||
|
assert_real_file "models/mpnn/all_amine_split_for_LiON/cv_${fold}/fold_0/model_0/model.pt"
|
||||||
|
done
|
||||||
|
|
||||||
|
find models/qwen2.5-7b-instruct -maxdepth 1 -type f \
|
||||||
|
\( -name '*.safetensors' -o -name '*.bin' \) -print -quit | grep -q . \
|
||||||
|
|| fail "Qwen weight shards (*.safetensors or *.bin) are missing"
|
||||||
|
|
||||||
|
docker compose -f "${COMPOSE_FILE}" config >/dev/null
|
||||||
|
echo "preflight: OK"
|
||||||
|
}
|
||||||
|
|
||||||
|
wait_for_api() {
|
||||||
|
local attempt
|
||||||
|
for attempt in $(seq 1 120); do
|
||||||
|
if curl -fsS "${API_URL}/" >/dev/null 2>&1; then
|
||||||
|
echo "API ready after ${attempt} checks"
|
||||||
|
return 0
|
||||||
|
fi
|
||||||
|
sleep 2
|
||||||
|
done
|
||||||
|
docker compose -f "${COMPOSE_FILE}" logs --tail=200 api || true
|
||||||
|
fail "API did not become healthy within 240 seconds"
|
||||||
|
}
|
||||||
|
|
||||||
|
start() {
|
||||||
|
preflight
|
||||||
|
docker compose -f "${COMPOSE_FILE}" build
|
||||||
|
docker compose -f "${COMPOSE_FILE}" up -d
|
||||||
|
wait_for_api
|
||||||
|
docker compose -f "${COMPOSE_FILE}" ps
|
||||||
|
curl -fsS "${API_URL}/"
|
||||||
|
echo
|
||||||
|
}
|
||||||
|
|
||||||
|
request_item() {
|
||||||
|
printf '%s' '{"smiles":"CC(C)NCCNC(C)C","cationic_lipid_to_mrna_ratio":10.0,"cationic_lipid_mol_ratio":35.0,"phospholipid_mol_ratio":16.0,"cholesterol_mol_ratio":46.5,"peg_lipid_mol_ratio":2.5,"helper_lipid":"DOPE","route":"intravenous"}'
|
||||||
|
}
|
||||||
|
|
||||||
|
smoke() {
|
||||||
|
wait_for_api
|
||||||
|
|
||||||
|
local item payload i
|
||||||
|
item=$(request_item)
|
||||||
|
curl -fsS -H 'Content-Type: application/json' \
|
||||||
|
-d "${item}" "${API_URL}/predict" >/tmp/lnp-single-response.json
|
||||||
|
echo "single prediction: OK"
|
||||||
|
|
||||||
|
payload='{"items":['
|
||||||
|
for i in $(seq 1 32); do
|
||||||
|
[ "${i}" -eq 1 ] || payload+=','
|
||||||
|
payload+="${item}"
|
||||||
|
done
|
||||||
|
payload+='],"batch_size":32,"use_llm":true}'
|
||||||
|
|
||||||
|
curl -fsS -H 'Content-Type: application/json' \
|
||||||
|
-d "${payload}" "${API_URL}/predict/batch" >/tmp/lnp-batch-response.json
|
||||||
|
grep -q '"n_succeeded":32' /tmp/lnp-batch-response.json \
|
||||||
|
|| fail "batch response did not report 32 successes"
|
||||||
|
echo "32-item LLM prediction: OK"
|
||||||
|
|
||||||
|
docker compose -f "${COMPOSE_FILE}" logs --tail=200 api \
|
||||||
|
| grep -E 'LLM inference micro-batching|ERROR|CUDA out of memory' || true
|
||||||
|
}
|
||||||
|
|
||||||
|
case "${1:-}" in
|
||||||
|
preflight) preflight ;;
|
||||||
|
start) start ;;
|
||||||
|
smoke) smoke ;;
|
||||||
|
*) echo "Usage: $0 {preflight|start|smoke}" >&2; exit 2 ;;
|
||||||
|
esac
|
||||||
88
server_test_cases/README.md
Normal file
88
server_test_cases/README.md
Normal file
@ -0,0 +1,88 @@
|
|||||||
|
# LNP API 与 GPU 显存测试
|
||||||
|
|
||||||
|
这里有两套测试:
|
||||||
|
|
||||||
|
- `run_lnp_batch_stress.sh`:专门复现 LNP 批量调用的显存问题。
|
||||||
|
- `run_cases.sh`:检查健康状态、器官列表、单条预测、批量预测、异常输入和小规模配方优化。
|
||||||
|
|
||||||
|
## 1. 复现 LNP 批量显存场景
|
||||||
|
|
||||||
|
推荐先运行与问题最接近的场景:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bash server_test_cases/run_lnp_batch_stress.sh repro
|
||||||
|
```
|
||||||
|
|
||||||
|
`repro` 会依次提交:
|
||||||
|
|
||||||
|
1. 32 条不同 LNP 配方,`use_llm=true`、外层 `batch_size=32`。这是原先最容易从约 10GB 涨到 63GB 的场景。
|
||||||
|
2. 500 条不同 LNP 配方,仍使用外层 `batch_size=32`,检查多轮处理后是否持续涨显存。
|
||||||
|
|
||||||
|
逐级观察批量大小与显存的关系:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bash server_test_cases/run_lnp_batch_stress.sh ladder
|
||||||
|
```
|
||||||
|
|
||||||
|
完整测试(关闭 LLM 对照 + 阶梯 + 500 条持续压力):
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bash server_test_cases/run_lnp_batch_stress.sh full
|
||||||
|
```
|
||||||
|
|
||||||
|
每个场景都会打印并写入 CSV:运行前显存、采样峰值、显存增量、耗时、HTTP 状态。结果文件名形如 `vram-results-20260922-160000.csv`。
|
||||||
|
|
||||||
|
可以设置显存上限,让超过阈值时测试直接失败。例如期望不超过 20GB:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
MAX_GPU_MIB=20480 bash server_test_cases/run_lnp_batch_stress.sh repro
|
||||||
|
```
|
||||||
|
|
||||||
|
这里采样的是整张卡的显存;测试时请尽量不要让其他任务共用该卡。
|
||||||
|
|
||||||
|
## 2. Mock 数据
|
||||||
|
|
||||||
|
已附带三组可直接提交的数据:
|
||||||
|
|
||||||
|
- `mock_batch_32_llm.json`:核心显存回归场景。
|
||||||
|
- `mock_batch_500_llm.json`:最大请求量与持续处理场景。
|
||||||
|
- `mock_batch_128_no_llm.json`:关闭 LLM 的对照组。
|
||||||
|
|
||||||
|
这些数据从项目 `data/interim/internal.csv` 抽取不同 SMILES,再随机生成满足总和 100% 的组分比例。也可以自行生成:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python3 server_test_cases/generate_mock_batch.py \
|
||||||
|
--count 200 \
|
||||||
|
--batch-size 32 \
|
||||||
|
--use-llm \
|
||||||
|
--source-csv data/interim/internal.csv \
|
||||||
|
--output /tmp/lnp-batch-200.json
|
||||||
|
```
|
||||||
|
|
||||||
|
生成过程使用固定随机种子,便于前后版本用完全相同的数据对比。
|
||||||
|
|
||||||
|
## 3. 其他功能测试
|
||||||
|
|
||||||
|
服务启动后执行:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
bash server_test_cases/run_cases.sh
|
||||||
|
```
|
||||||
|
|
||||||
|
如果 API 端口或 GPU 编号不同:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
API_PORT=18000 GPU_INDEX=0 bash server_test_cases/run_cases.sh
|
||||||
|
```
|
||||||
|
|
||||||
|
用例包含:
|
||||||
|
|
||||||
|
1. 健康检查与模型状态。
|
||||||
|
2. 可用器官列表。
|
||||||
|
3. 启用 LLM 的单条预测,检查输出完整性、有限数和 biodistribution 归一化。
|
||||||
|
4. 64 条关闭 LLM 的吞吐对照。
|
||||||
|
5. 32 条启用 LLM 的显存回归测试。
|
||||||
|
6. 2 条合法 + 3 条非法输入,检查批量接口的局部容错。
|
||||||
|
7. 小范围 `/optimize` 配方搜索,检查优化功能和返回结果。
|
||||||
|
|
||||||
|
GPU 显存采样来自整张卡;如果还有其他进程使用同一张卡,请以 `nvidia-smi` 中 API 进程的变化为准。
|
||||||
102
server_test_cases/generate_mock_batch.py
Executable file
102
server_test_cases/generate_mock_batch.py
Executable file
@ -0,0 +1,102 @@
|
|||||||
|
#!/usr/bin/env python3
|
||||||
|
"""Generate reproducible /predict/batch payloads for GPU memory tests."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import csv
|
||||||
|
import json
|
||||||
|
import random
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
# Used only when data/interim/internal.csv is not available. These are valid,
|
||||||
|
# deliberately varied SMILES; the project dataset is preferred for realistic RAG prompts.
|
||||||
|
FALLBACK_SMILES = [
|
||||||
|
"CC(C)NCCNC(C)C",
|
||||||
|
"CCCCN(CC)CC",
|
||||||
|
"CCCCCCCCN(C)C",
|
||||||
|
"CCN(CC)CCOC(=O)CCCC",
|
||||||
|
"CCCCCCCCCCN(CC)CC",
|
||||||
|
"CC(C)(C)OC(=O)NCCN(C)C",
|
||||||
|
"CCCCCCCCCCCCN(C)C",
|
||||||
|
"CCOC(=O)CN(C)C",
|
||||||
|
"CCCCCCN(CC)CC",
|
||||||
|
"CN(C)CCOC(=O)CCCCCCCC",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def load_smiles(source: Path | None) -> list[str]:
|
||||||
|
if source and source.is_file():
|
||||||
|
with source.open(encoding="utf-8-sig", newline="") as handle:
|
||||||
|
rows = csv.DictReader(handle)
|
||||||
|
values = []
|
||||||
|
seen = set()
|
||||||
|
for row in rows:
|
||||||
|
value = (row.get("smiles") or "").strip()
|
||||||
|
if value and value.lower() != "nan" and value not in seen:
|
||||||
|
seen.add(value)
|
||||||
|
values.append(value)
|
||||||
|
if values:
|
||||||
|
return values
|
||||||
|
return FALLBACK_SMILES
|
||||||
|
|
||||||
|
|
||||||
|
def make_composition(rng: random.Random) -> tuple[float, float, float, float]:
|
||||||
|
# Match the ranges used by the optimizer and guarantee an exact 100% sum.
|
||||||
|
for _ in range(1000):
|
||||||
|
peg = round(rng.uniform(1.0, 5.0), 2)
|
||||||
|
cationic = round(rng.uniform(25.0, 53.0), 2)
|
||||||
|
phospholipid = round(rng.uniform(8.0, 30.0), 2)
|
||||||
|
cholesterol = round(100.0 - peg - cationic - phospholipid, 2)
|
||||||
|
if 15.0 <= cholesterol <= 46.0:
|
||||||
|
return cationic, phospholipid, cholesterol, peg
|
||||||
|
raise RuntimeError("unable to generate a valid composition")
|
||||||
|
|
||||||
|
|
||||||
|
def main() -> None:
|
||||||
|
parser = argparse.ArgumentParser(description=__doc__)
|
||||||
|
parser.add_argument("--count", type=int, required=True, choices=range(1, 501))
|
||||||
|
parser.add_argument("--batch-size", type=int, default=32, choices=range(1, 257))
|
||||||
|
parser.add_argument("--use-llm", action=argparse.BooleanOptionalAction, default=True)
|
||||||
|
parser.add_argument("--seed", type=int, default=20260922)
|
||||||
|
parser.add_argument("--source-csv", type=Path)
|
||||||
|
parser.add_argument("--output", type=Path, required=True)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
rng = random.Random(args.seed)
|
||||||
|
smiles_pool = load_smiles(args.source_csv)
|
||||||
|
rng.shuffle(smiles_pool)
|
||||||
|
items = []
|
||||||
|
for index in range(args.count):
|
||||||
|
cationic, phospholipid, cholesterol, peg = make_composition(rng)
|
||||||
|
items.append(
|
||||||
|
{
|
||||||
|
"smiles": smiles_pool[index % len(smiles_pool)],
|
||||||
|
"cationic_lipid_to_mrna_ratio": round(rng.uniform(7.0, 20.0), 2),
|
||||||
|
"cationic_lipid_mol_ratio": cationic,
|
||||||
|
"phospholipid_mol_ratio": phospholipid,
|
||||||
|
"cholesterol_mol_ratio": cholesterol,
|
||||||
|
"peg_lipid_mol_ratio": peg,
|
||||||
|
"helper_lipid": ("DOPE", "DSPC")[index % 2],
|
||||||
|
"route": ("intravenous", "intramuscular")[index % 2],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = {
|
||||||
|
"items": items,
|
||||||
|
"batch_size": args.batch_size,
|
||||||
|
"use_llm": args.use_llm,
|
||||||
|
}
|
||||||
|
args.output.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
args.output.write_text(
|
||||||
|
json.dumps(payload, ensure_ascii=False, indent=2) + "\n", encoding="utf-8"
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"generated {args.output}: items={args.count}, unique_smiles={len(set(x['smiles'] for x in items))}, "
|
||||||
|
f"batch_size={args.batch_size}, use_llm={args.use_llm}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
56
server_test_cases/mixed_valid_invalid.json
Normal file
56
server_test_cases/mixed_valid_invalid.json
Normal file
@ -0,0 +1,56 @@
|
|||||||
|
{
|
||||||
|
"items": [
|
||||||
|
{
|
||||||
|
"smiles": "CC(C)NCCNC(C)C",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 10.0,
|
||||||
|
"cationic_lipid_mol_ratio": 35.0,
|
||||||
|
"phospholipid_mol_ratio": 16.0,
|
||||||
|
"cholesterol_mol_ratio": 46.5,
|
||||||
|
"peg_lipid_mol_ratio": 2.5,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "CCCCCCCCCCCCCCCCN",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 12.0,
|
||||||
|
"cationic_lipid_mol_ratio": 40.0,
|
||||||
|
"phospholipid_mol_ratio": 15.0,
|
||||||
|
"cholesterol_mol_ratio": 42.5,
|
||||||
|
"peg_lipid_mol_ratio": 2.5,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "this-is-not-smiles",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 10.0,
|
||||||
|
"cationic_lipid_mol_ratio": 35.0,
|
||||||
|
"phospholipid_mol_ratio": 16.0,
|
||||||
|
"cholesterol_mol_ratio": 46.5,
|
||||||
|
"peg_lipid_mol_ratio": 2.5,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "CCN(CC)CC",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 10.0,
|
||||||
|
"cationic_lipid_mol_ratio": 35.0,
|
||||||
|
"phospholipid_mol_ratio": 16.0,
|
||||||
|
"cholesterol_mol_ratio": 30.0,
|
||||||
|
"peg_lipid_mol_ratio": 2.5,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "C1CCN(CC1)CCO",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 10.0,
|
||||||
|
"cationic_lipid_mol_ratio": 35.0,
|
||||||
|
"phospholipid_mol_ratio": 16.0,
|
||||||
|
"cholesterol_mol_ratio": 46.5,
|
||||||
|
"peg_lipid_mol_ratio": 2.5,
|
||||||
|
"helper_lipid": "INVALID_HELPER",
|
||||||
|
"route": "intravenous"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"batch_size": 5,
|
||||||
|
"use_llm": false
|
||||||
|
}
|
||||||
1286
server_test_cases/mock_batch_128_no_llm.json
Normal file
1286
server_test_cases/mock_batch_128_no_llm.json
Normal file
File diff suppressed because it is too large
Load Diff
326
server_test_cases/mock_batch_32_llm.json
Normal file
326
server_test_cases/mock_batch_32_llm.json
Normal file
@ -0,0 +1,326 @@
|
|||||||
|
{
|
||||||
|
"items": [
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O)CCNC(C[C@@H](CN)CC(C)C)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 7.46,
|
||||||
|
"cationic_lipid_mol_ratio": 37.58,
|
||||||
|
"phospholipid_mol_ratio": 23.5,
|
||||||
|
"cholesterol_mol_ratio": 36.53,
|
||||||
|
"peg_lipid_mol_ratio": 2.39,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCCCCCCCC)CCN(CCC(OCCCCCCCCCCCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCCCCCCCC)=O)CCC(OCCCCCCCCCCCCCC)=O)CCNC(CCC(NCC1=CC=C(B(O)O)C=C1)=O)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 16.82,
|
||||||
|
"cationic_lipid_mol_ratio": 50.09,
|
||||||
|
"phospholipid_mol_ratio": 13.01,
|
||||||
|
"cholesterol_mol_ratio": 33.14,
|
||||||
|
"peg_lipid_mol_ratio": 3.76,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCN([C@H](C(NCCCN(CCCCCCOC(CCCCCCC)=O)CCCCCCOC(CCCCCCC)=O)=O)CO)CCCCCCOC(CCCCCCC)=O)CCCCCCC",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 9.92,
|
||||||
|
"cationic_lipid_mol_ratio": 46.82,
|
||||||
|
"phospholipid_mol_ratio": 16.37,
|
||||||
|
"cholesterol_mol_ratio": 33.48,
|
||||||
|
"peg_lipid_mol_ratio": 3.33,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCN([C@H](C(NCCCCCCN(CCCCCCOC(CCCCC)=O)CCCCCCOC(CCCCC)=O)=O)CO)CCCCCCOC(CCCCC)=O)CCCCC",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 10.73,
|
||||||
|
"cationic_lipid_mol_ratio": 36.31,
|
||||||
|
"phospholipid_mol_ratio": 23.42,
|
||||||
|
"cholesterol_mol_ratio": 35.49,
|
||||||
|
"peg_lipid_mol_ratio": 4.78,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCN(CC1=NC=CS1)CC2=NC=CS2)CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 19.35,
|
||||||
|
"cationic_lipid_mol_ratio": 51.39,
|
||||||
|
"phospholipid_mol_ratio": 10.87,
|
||||||
|
"cholesterol_mol_ratio": 35.21,
|
||||||
|
"peg_lipid_mol_ratio": 2.53,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCCCCCCCC)CCN(CCC(OCCCCCCCCCCCCCC)=O)CCCCCCN(CCCCCCN(CCC(OCCCCCCCCCCCCCC)=O)CCC(OCCCCCCCCCCCCCC)=O)CCCN1CCOCC1",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 17.21,
|
||||||
|
"cationic_lipid_mol_ratio": 46.17,
|
||||||
|
"phospholipid_mol_ratio": 19.64,
|
||||||
|
"cholesterol_mol_ratio": 31.96,
|
||||||
|
"peg_lipid_mol_ratio": 2.23,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "CCCCCCCCCCCC(OCCCCCCN(CCCCCCOC(CCCCCCCCCCC)=O)CCCOC([C@@H](N(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)CO)=O)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 10.34,
|
||||||
|
"cationic_lipid_mol_ratio": 44.28,
|
||||||
|
"phospholipid_mol_ratio": 17.74,
|
||||||
|
"cholesterol_mol_ratio": 34.79,
|
||||||
|
"peg_lipid_mol_ratio": 3.19,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCN([C@H](C(NCCN(C)CCN(CCCCCCOC(CCCCCCCCC=C)=O)CCCCCCOC(CCCCCCCCC=C)=O)=O)CO)CCCCCCOC(CCCCCCCCC=C)=O)CCCCCCCCC=C",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 9.15,
|
||||||
|
"cationic_lipid_mol_ratio": 42.88,
|
||||||
|
"phospholipid_mol_ratio": 10.3,
|
||||||
|
"cholesterol_mol_ratio": 43.76,
|
||||||
|
"peg_lipid_mol_ratio": 3.06,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "CC(CNC([C@@H](N(CCCCCCOC(CCCCCCCCC=C)=O)CCCCCCOC(CCCCCCCCC=C)=O)CO)=O)(CN(CCCCCCOC(CCCCCCCCC=C)=O)CCCCCCOC(CCCCCCCCC=C)=O)C",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 10.12,
|
||||||
|
"cationic_lipid_mol_ratio": 49.74,
|
||||||
|
"phospholipid_mol_ratio": 12.42,
|
||||||
|
"cholesterol_mol_ratio": 35.63,
|
||||||
|
"peg_lipid_mol_ratio": 2.21,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(NCCN(CCCCN(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)CCCCN(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)CCCN",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 19.72,
|
||||||
|
"cationic_lipid_mol_ratio": 32.8,
|
||||||
|
"phospholipid_mol_ratio": 25.26,
|
||||||
|
"cholesterol_mol_ratio": 38.21,
|
||||||
|
"peg_lipid_mol_ratio": 3.73,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCN1CC(C(OC)=O)CC1=O)CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 13.7,
|
||||||
|
"cationic_lipid_mol_ratio": 43.91,
|
||||||
|
"phospholipid_mol_ratio": 21.35,
|
||||||
|
"cholesterol_mol_ratio": 31.92,
|
||||||
|
"peg_lipid_mol_ratio": 2.82,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "CCCCCCCCCCCCOC(CCN([C@H](C(NCCCCCCN(CCC(OCCCCCCCCCCCC)=O)CCC(OCCCCCCCCCCCC)=O)=O)CO)CCC(OCCCCCCCCCCCC)=O)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 10.47,
|
||||||
|
"cationic_lipid_mol_ratio": 32.49,
|
||||||
|
"phospholipid_mol_ratio": 24.48,
|
||||||
|
"cholesterol_mol_ratio": 40.96,
|
||||||
|
"peg_lipid_mol_ratio": 2.07,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C([C@H](CO)N(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)N(CC1)CCN1C2=NC=CC=N2",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 12.86,
|
||||||
|
"cationic_lipid_mol_ratio": 51.8,
|
||||||
|
"phospholipid_mol_ratio": 19.97,
|
||||||
|
"cholesterol_mol_ratio": 25.07,
|
||||||
|
"peg_lipid_mol_ratio": 3.16,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(CN(CC(O)CCCCCCCC)CC(O)CCCCCCCC)NCCCN(CC(O)CCCCCCCC)CC(O)CCCCCCCC",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 18.73,
|
||||||
|
"cationic_lipid_mol_ratio": 45.45,
|
||||||
|
"phospholipid_mol_ratio": 13.35,
|
||||||
|
"cholesterol_mol_ratio": 38.13,
|
||||||
|
"peg_lipid_mol_ratio": 3.07,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "CCCCCCCCCCCC(OCCCCCCN(CCN(C)CCNC([C@@H](N(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)CO)=O)CCCCCCOC(CCCCCCCCCCC)=O)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 10.09,
|
||||||
|
"cationic_lipid_mol_ratio": 37.43,
|
||||||
|
"phospholipid_mol_ratio": 11.94,
|
||||||
|
"cholesterol_mol_ratio": 45.78,
|
||||||
|
"peg_lipid_mol_ratio": 4.85,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "OC[C@@H](C(NCCCCCCN(CC(O)CCCCCCCCCC)CC(O)CCCCCCCCCC)=O)N(CC(CCCCCCCCCC)O)CC(CCCCCCCCCC)O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 13.05,
|
||||||
|
"cationic_lipid_mol_ratio": 52.84,
|
||||||
|
"phospholipid_mol_ratio": 8.6,
|
||||||
|
"cholesterol_mol_ratio": 37.06,
|
||||||
|
"peg_lipid_mol_ratio": 1.5,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCCCCCCCC)CCN(CCC(OCCCCCCCCCCCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCCCCCCCC)=O)CCC(OCCCCCCCCCCCCCC)=O)CCNC(C[C@@H](CN)CC(C)C)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 11.09,
|
||||||
|
"cationic_lipid_mol_ratio": 48.25,
|
||||||
|
"phospholipid_mol_ratio": 14.52,
|
||||||
|
"cholesterol_mol_ratio": 32.66,
|
||||||
|
"peg_lipid_mol_ratio": 4.57,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCC/C=C\\C/C=C\\CCCCC)CCN(CCC(OCCCCCCCC/C=C\\C/C=C\\CCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCC/C=C\\C/C=C\\CCCCC)=O)CCC(OCCCCCCCC/C=C\\C/C=C\\CCCCC)=O)CCNC(CCC(NCCC1=CC=C(O)C(O)=C1)=O)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 19.63,
|
||||||
|
"cationic_lipid_mol_ratio": 31.98,
|
||||||
|
"phospholipid_mol_ratio": 25.5,
|
||||||
|
"cholesterol_mol_ratio": 40.57,
|
||||||
|
"peg_lipid_mol_ratio": 1.95,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(NCCCCCCN(CCCCCCCCCCCC)CCCCCCCCCCCC)[C@H](C)N(CCCCCCCCCCCC)CCCCCCCCCCCC",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 16.96,
|
||||||
|
"cationic_lipid_mol_ratio": 28.21,
|
||||||
|
"phospholipid_mol_ratio": 28.57,
|
||||||
|
"cholesterol_mol_ratio": 39.51,
|
||||||
|
"peg_lipid_mol_ratio": 3.71,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCN(CCC(F)(F)F)CCC(F)(F)F)CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 19.73,
|
||||||
|
"cationic_lipid_mol_ratio": 37.62,
|
||||||
|
"phospholipid_mol_ratio": 22.08,
|
||||||
|
"cholesterol_mol_ratio": 36.42,
|
||||||
|
"peg_lipid_mol_ratio": 3.88,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "CCCCCCCCC(CN(CCCNC([C@@H](N(CC(CCCCCCCC)O)CC(CCCCCCCC)O)CO)=O)CC(CCCCCCCC)O)O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 15.62,
|
||||||
|
"cationic_lipid_mol_ratio": 51.29,
|
||||||
|
"phospholipid_mol_ratio": 27.02,
|
||||||
|
"cholesterol_mol_ratio": 18.83,
|
||||||
|
"peg_lipid_mol_ratio": 2.86,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C([C@H](CO)N(CCCCCCOC(CCCCCCCCC=C)=O)CCCCCCOC(CCCCCCCCC=C)=O)N(CC1)CCN1C2=NC=CC=N2",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 14.57,
|
||||||
|
"cationic_lipid_mol_ratio": 38.85,
|
||||||
|
"phospholipid_mol_ratio": 18.16,
|
||||||
|
"cholesterol_mol_ratio": 40.17,
|
||||||
|
"peg_lipid_mol_ratio": 2.82,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "CCCCCCCCCCCC(OCCCCCCN([C@H](C(NCCCCCCN(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)=O)CO)CCCCCCOC(CCCCCCCCCCC)=O)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 17.0,
|
||||||
|
"cationic_lipid_mol_ratio": 41.61,
|
||||||
|
"phospholipid_mol_ratio": 21.12,
|
||||||
|
"cholesterol_mol_ratio": 35.9,
|
||||||
|
"peg_lipid_mol_ratio": 1.37,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCCCCCC)CCN(CCC(OCCCCCCCCCCCC)=O)CCCCCCN(CCCCCCN(CCC(OCCCCCCCCCCCC)=O)CCC(OCCCCCCCCCCCC)=O)CCN1CCOCC1",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 12.75,
|
||||||
|
"cationic_lipid_mol_ratio": 31.73,
|
||||||
|
"phospholipid_mol_ratio": 29.19,
|
||||||
|
"cholesterol_mol_ratio": 37.55,
|
||||||
|
"peg_lipid_mol_ratio": 1.53,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(CN(CCCCCCCCCCCCCC)CCCCCCCCCCCCCC)NCCCN(CCCCCCCCCCCCCC)CCCCCCCCCCCCCC",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 19.09,
|
||||||
|
"cationic_lipid_mol_ratio": 41.63,
|
||||||
|
"phospholipid_mol_ratio": 16.23,
|
||||||
|
"cholesterol_mol_ratio": 39.43,
|
||||||
|
"peg_lipid_mol_ratio": 2.71,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCCCN(CCCCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O)CCNC(CCC(NCCC1=CC=C(O)C(O)=C1)=O)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 16.96,
|
||||||
|
"cationic_lipid_mol_ratio": 34.46,
|
||||||
|
"phospholipid_mol_ratio": 28.65,
|
||||||
|
"cholesterol_mol_ratio": 32.06,
|
||||||
|
"peg_lipid_mol_ratio": 4.83,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O)CCNC([C@@H](N)CCCCN)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 14.76,
|
||||||
|
"cationic_lipid_mol_ratio": 34.58,
|
||||||
|
"phospholipid_mol_ratio": 25.73,
|
||||||
|
"cholesterol_mol_ratio": 35.03,
|
||||||
|
"peg_lipid_mol_ratio": 4.66,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(CN(CCCCCCCCCCCC)CCCCCCCCCCCC)NCCCN(CCCCCCCCCCCC)CCCCCCCCCCCC",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 7.01,
|
||||||
|
"cationic_lipid_mol_ratio": 42.86,
|
||||||
|
"phospholipid_mol_ratio": 23.71,
|
||||||
|
"cholesterol_mol_ratio": 30.73,
|
||||||
|
"peg_lipid_mol_ratio": 2.7,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCN(CC1=COC=N1)CC2=COC=N2)CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 7.11,
|
||||||
|
"cationic_lipid_mol_ratio": 43.3,
|
||||||
|
"phospholipid_mol_ratio": 22.32,
|
||||||
|
"cholesterol_mol_ratio": 31.55,
|
||||||
|
"peg_lipid_mol_ratio": 2.83,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCNC(CCCN)=O)CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 18.85,
|
||||||
|
"cationic_lipid_mol_ratio": 45.22,
|
||||||
|
"phospholipid_mol_ratio": 17.28,
|
||||||
|
"cholesterol_mol_ratio": 33.01,
|
||||||
|
"peg_lipid_mol_ratio": 4.49,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(CN(CCCCCCCCCC)CCCCCCCCCC)NCCCN(CCCCCCCCCC)CCCCCCCCCC",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 10.74,
|
||||||
|
"cationic_lipid_mol_ratio": 38.47,
|
||||||
|
"phospholipid_mol_ratio": 25.72,
|
||||||
|
"cholesterol_mol_ratio": 32.81,
|
||||||
|
"peg_lipid_mol_ratio": 3.0,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"smiles": "O=C(OCCCCCCCC/C=C\\C/C=C\\CCCCC)CCN(CCC(OCCCCCCCC/C=C\\C/C=C\\CCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCC/C=C\\C/C=C\\CCCCC)=O)CCC(OCCCCCCCC/C=C\\C/C=C\\CCCCC)=O)CCNC(CCCN)=O",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 17.44,
|
||||||
|
"cationic_lipid_mol_ratio": 42.07,
|
||||||
|
"phospholipid_mol_ratio": 17.36,
|
||||||
|
"cholesterol_mol_ratio": 36.89,
|
||||||
|
"peg_lipid_mol_ratio": 3.68,
|
||||||
|
"helper_lipid": "DSPC",
|
||||||
|
"route": "intramuscular"
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"batch_size": 32,
|
||||||
|
"use_llm": true
|
||||||
|
}
|
||||||
5006
server_test_cases/mock_batch_500_llm.json
Normal file
5006
server_test_cases/mock_batch_500_llm.json
Normal file
File diff suppressed because it is too large
Load Diff
31
server_test_cases/optimize_smoke.json
Normal file
31
server_test_cases/optimize_smoke.json
Normal file
@ -0,0 +1,31 @@
|
|||||||
|
{
|
||||||
|
"smiles": "CC(C)NCCNC(C)C",
|
||||||
|
"organ": "liver",
|
||||||
|
"top_k": 3,
|
||||||
|
"num_seeds": 5,
|
||||||
|
"top_per_seed": 1,
|
||||||
|
"rerank_top_n": 8,
|
||||||
|
"step_sizes": [5],
|
||||||
|
"wr_step_sizes": [5],
|
||||||
|
"routes": ["intravenous"],
|
||||||
|
"comp_ranges": {
|
||||||
|
"weight_ratio_min": 10,
|
||||||
|
"weight_ratio_max": 15,
|
||||||
|
"cationic_mol_min": 40,
|
||||||
|
"cationic_mol_max": 50,
|
||||||
|
"phospholipid_mol_min": 10,
|
||||||
|
"phospholipid_mol_max": 20,
|
||||||
|
"cholesterol_mol_min": 28,
|
||||||
|
"cholesterol_mol_max": 45,
|
||||||
|
"peg_mol_min": 1,
|
||||||
|
"peg_mol_max": 5
|
||||||
|
},
|
||||||
|
"scoring_weights": {
|
||||||
|
"biodist_weight": 1,
|
||||||
|
"delivery_weight": 0,
|
||||||
|
"size_weight": 0,
|
||||||
|
"ee_class_weights": [0, 0, 0],
|
||||||
|
"pdi_class_weights": [0, 0],
|
||||||
|
"toxic_class_weights": [0, 0]
|
||||||
|
}
|
||||||
|
}
|
||||||
187
server_test_cases/run_cases.sh
Executable file
187
server_test_cases/run_cases.sh
Executable file
@ -0,0 +1,187 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
set -euo pipefail
|
||||||
|
|
||||||
|
CASE_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)
|
||||||
|
API_PORT=${API_PORT:-18000}
|
||||||
|
API_URL=${API_URL:-http://127.0.0.1:${API_PORT}}
|
||||||
|
GPU_INDEX=${GPU_INDEX:-0}
|
||||||
|
CURL_TIMEOUT=${CURL_TIMEOUT:-1800}
|
||||||
|
|
||||||
|
PASS=0
|
||||||
|
FAIL=0
|
||||||
|
|
||||||
|
run_case() {
|
||||||
|
local name=$1
|
||||||
|
shift
|
||||||
|
echo
|
||||||
|
echo "===== ${name} ====="
|
||||||
|
if "$@"; then
|
||||||
|
PASS=$((PASS + 1))
|
||||||
|
echo "PASS: ${name}"
|
||||||
|
else
|
||||||
|
FAIL=$((FAIL + 1))
|
||||||
|
echo "FAIL: ${name}" >&2
|
||||||
|
fi
|
||||||
|
}
|
||||||
|
|
||||||
|
validate_health() {
|
||||||
|
local response=/tmp/lnp-test-health.json
|
||||||
|
curl -fsS --max-time 30 "${API_URL}/" >"${response}"
|
||||||
|
python3 - "${response}" <<'PY'
|
||||||
|
import json, sys
|
||||||
|
data = json.load(open(sys.argv[1], encoding="utf-8"))
|
||||||
|
assert data["status"] == "healthy", data
|
||||||
|
assert data["model_loaded"] is True, data
|
||||||
|
assert data["use_llm"] is True, data
|
||||||
|
print(json.dumps(data, ensure_ascii=False, indent=2))
|
||||||
|
PY
|
||||||
|
}
|
||||||
|
|
||||||
|
validate_organs() {
|
||||||
|
local response=/tmp/lnp-test-organs.json
|
||||||
|
curl -fsS --max-time 30 "${API_URL}/organs" >"${response}"
|
||||||
|
python3 - "${response}" <<'PY'
|
||||||
|
import json, sys
|
||||||
|
data = json.load(open(sys.argv[1], encoding="utf-8"))
|
||||||
|
expected = {"lymph_nodes", "heart", "liver", "spleen", "lung", "kidney", "muscle"}
|
||||||
|
assert set(data) == expected, data
|
||||||
|
print(json.dumps(data, ensure_ascii=False))
|
||||||
|
PY
|
||||||
|
}
|
||||||
|
|
||||||
|
validate_single() {
|
||||||
|
local response=/tmp/lnp-test-single.json
|
||||||
|
curl -fsS --max-time "${CURL_TIMEOUT}" \
|
||||||
|
-H 'Content-Type: application/json' \
|
||||||
|
--data-binary "@${CASE_DIR}/single_valid.json" \
|
||||||
|
"${API_URL}/predict" >"${response}"
|
||||||
|
python3 - "${response}" <<'PY'
|
||||||
|
import json, math, sys
|
||||||
|
data = json.load(open(sys.argv[1], encoding="utf-8"))
|
||||||
|
required = ["biodist", "quantified_delivery", "pdi_class", "ee_class", "toxic_class"]
|
||||||
|
assert all(k in data for k in required), data
|
||||||
|
values = list(data["biodist"].values())
|
||||||
|
assert len(values) == 7, values
|
||||||
|
assert all(math.isfinite(float(x)) for x in values), values
|
||||||
|
assert abs(sum(values) - 1.0) < 1e-4, sum(values)
|
||||||
|
assert math.isfinite(float(data["quantified_delivery"]))
|
||||||
|
print(json.dumps(data, ensure_ascii=False, indent=2))
|
||||||
|
PY
|
||||||
|
}
|
||||||
|
|
||||||
|
make_batch_payload() {
|
||||||
|
local count=$1
|
||||||
|
local use_llm=$2
|
||||||
|
local batch_size=$3
|
||||||
|
python3 - "${CASE_DIR}/single_valid.json" "${count}" "${use_llm}" "${batch_size}" <<'PY'
|
||||||
|
import json, sys
|
||||||
|
item = json.load(open(sys.argv[1], encoding="utf-8"))
|
||||||
|
print(json.dumps({
|
||||||
|
"items": [item for _ in range(int(sys.argv[2]))],
|
||||||
|
"batch_size": int(sys.argv[4]),
|
||||||
|
"use_llm": sys.argv[3].lower() == "true",
|
||||||
|
}))
|
||||||
|
PY
|
||||||
|
}
|
||||||
|
|
||||||
|
validate_batch_response() {
|
||||||
|
local response=$1
|
||||||
|
local expected=$2
|
||||||
|
python3 - "${response}" "${expected}" <<'PY'
|
||||||
|
import json, math, sys
|
||||||
|
data = json.load(open(sys.argv[1], encoding="utf-8"))
|
||||||
|
expected = int(sys.argv[2])
|
||||||
|
assert data["n_requested"] == expected, data
|
||||||
|
assert data["n_succeeded"] == expected, data
|
||||||
|
assert data["errors"] == [], data["errors"]
|
||||||
|
assert len(data["predictions"]) == expected
|
||||||
|
for pred in data["predictions"]:
|
||||||
|
values = list(pred["biodist"].values())
|
||||||
|
assert len(values) == 7
|
||||||
|
assert all(math.isfinite(float(x)) for x in values)
|
||||||
|
assert abs(sum(values) - 1.0) < 1e-4
|
||||||
|
print(f"n_succeeded={data['n_succeeded']}, errors={len(data['errors'])}")
|
||||||
|
PY
|
||||||
|
}
|
||||||
|
|
||||||
|
run_batch() {
|
||||||
|
local count=$1
|
||||||
|
local use_llm=$2
|
||||||
|
local batch_size=$3
|
||||||
|
local response="/tmp/lnp-test-batch-${count}-${use_llm}.json"
|
||||||
|
local payload before peak current curl_pid start_ms end_ms
|
||||||
|
|
||||||
|
payload=$(make_batch_payload "${count}" "${use_llm}" "${batch_size}")
|
||||||
|
before=$(nvidia-smi --id="${GPU_INDEX}" --query-gpu=memory.used \
|
||||||
|
--format=csv,noheader,nounits | head -n 1 | tr -d ' ')
|
||||||
|
peak=${before}
|
||||||
|
start_ms=$(date +%s%3N)
|
||||||
|
|
||||||
|
curl -fsS --max-time "${CURL_TIMEOUT}" \
|
||||||
|
-H 'Content-Type: application/json' \
|
||||||
|
-d "${payload}" "${API_URL}/predict/batch" >"${response}" &
|
||||||
|
curl_pid=$!
|
||||||
|
while kill -0 "${curl_pid}" 2>/dev/null; do
|
||||||
|
current=$(nvidia-smi --id="${GPU_INDEX}" --query-gpu=memory.used \
|
||||||
|
--format=csv,noheader,nounits | head -n 1 | tr -d ' ')
|
||||||
|
if [ "${current}" -gt "${peak}" ]; then peak=${current}; fi
|
||||||
|
sleep 0.2
|
||||||
|
done
|
||||||
|
wait "${curl_pid}"
|
||||||
|
end_ms=$(date +%s%3N)
|
||||||
|
|
||||||
|
validate_batch_response "${response}" "${count}"
|
||||||
|
echo "use_llm=${use_llm}, count=${count}, outer_batch=${batch_size}"
|
||||||
|
echo "GPU used before=${before} MiB, observed peak=${peak} MiB, delta=$((peak - before)) MiB"
|
||||||
|
echo "elapsed=$((end_ms - start_ms)) ms"
|
||||||
|
}
|
||||||
|
|
||||||
|
validate_mixed_errors() {
|
||||||
|
local response=/tmp/lnp-test-mixed.json
|
||||||
|
curl -fsS --max-time "${CURL_TIMEOUT}" \
|
||||||
|
-H 'Content-Type: application/json' \
|
||||||
|
--data-binary "@${CASE_DIR}/mixed_valid_invalid.json" \
|
||||||
|
"${API_URL}/predict/batch" >"${response}"
|
||||||
|
python3 - "${response}" <<'PY'
|
||||||
|
import json, sys
|
||||||
|
data = json.load(open(sys.argv[1], encoding="utf-8"))
|
||||||
|
assert data["n_requested"] == 5, data
|
||||||
|
assert data["n_succeeded"] == 2, data
|
||||||
|
assert len(data["errors"]) == 3, data
|
||||||
|
assert {x["index"] for x in data["errors"]} == {2, 3, 4}, data["errors"]
|
||||||
|
print(json.dumps(data["errors"], ensure_ascii=False, indent=2))
|
||||||
|
PY
|
||||||
|
}
|
||||||
|
|
||||||
|
validate_optimize() {
|
||||||
|
local response=/tmp/lnp-test-optimize.json
|
||||||
|
curl -fsS --max-time "${CURL_TIMEOUT}" \
|
||||||
|
-H 'Content-Type: application/json' \
|
||||||
|
--data-binary "@${CASE_DIR}/optimize_smoke.json" \
|
||||||
|
"${API_URL}/optimize" >"${response}"
|
||||||
|
python3 - "${response}" <<'PY'
|
||||||
|
import json, math, sys
|
||||||
|
data = json.load(open(sys.argv[1], encoding="utf-8"))
|
||||||
|
assert data["target_organ"] == "liver", data
|
||||||
|
assert 1 <= len(data["formulations"]) <= 3, data
|
||||||
|
for row in data["formulations"]:
|
||||||
|
assert math.isfinite(float(row["target_biodist"])), row
|
||||||
|
assert abs(sum(float(x) for x in row["all_biodist"].values()) - 1.0) < 1e-4, row
|
||||||
|
print(f"optimize returned {len(data['formulations'])} formulations")
|
||||||
|
PY
|
||||||
|
}
|
||||||
|
|
||||||
|
run_case "health" validate_health
|
||||||
|
run_case "available organs" validate_organs
|
||||||
|
run_case "single prediction with LLM" validate_single
|
||||||
|
run_case "64 items without LLM (throughput baseline)" run_batch 64 false 64
|
||||||
|
run_case "32 items with LLM (VRAM regression)" run_batch 32 true 32
|
||||||
|
run_case "mixed valid/invalid items" validate_mixed_errors
|
||||||
|
run_case "small optimize search" validate_optimize
|
||||||
|
|
||||||
|
echo
|
||||||
|
echo "===== SUMMARY ====="
|
||||||
|
echo "passed=${PASS}, failed=${FAIL}"
|
||||||
|
if [ "${FAIL}" -ne 0 ]; then
|
||||||
|
exit 1
|
||||||
|
fi
|
||||||
146
server_test_cases/run_lnp_batch_stress.sh
Executable file
146
server_test_cases/run_lnp_batch_stress.sh
Executable file
@ -0,0 +1,146 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
set -uo pipefail
|
||||||
|
|
||||||
|
CASE_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)
|
||||||
|
API_PORT=${API_PORT:-18000}
|
||||||
|
API_URL=${API_URL:-http://127.0.0.1:${API_PORT}}
|
||||||
|
GPU_INDEX=${GPU_INDEX:-0}
|
||||||
|
CURL_TIMEOUT=${CURL_TIMEOUT:-3600}
|
||||||
|
MAX_GPU_MIB=${MAX_GPU_MIB:-0}
|
||||||
|
PROFILE=${1:-repro}
|
||||||
|
RESULT_CSV=${RESULT_CSV:-"${CASE_DIR}/vram-results-$(date +%Y%m%d-%H%M%S).csv"}
|
||||||
|
|
||||||
|
if ! command -v curl >/dev/null || ! command -v python3 >/dev/null; then
|
||||||
|
echo "curl and python3 are required" >&2
|
||||||
|
exit 2
|
||||||
|
fi
|
||||||
|
if ! command -v nvidia-smi >/dev/null; then
|
||||||
|
echo "nvidia-smi is required for the GPU memory test" >&2
|
||||||
|
exit 2
|
||||||
|
fi
|
||||||
|
if ! curl -fsS --max-time 30 "${API_URL}/" >/tmp/lnp-stress-health.json; then
|
||||||
|
echo "API is not reachable: ${API_URL}" >&2
|
||||||
|
exit 2
|
||||||
|
fi
|
||||||
|
|
||||||
|
gpu_used_mib() {
|
||||||
|
nvidia-smi --id="${GPU_INDEX}" --query-gpu=memory.used --format=csv,noheader,nounits \
|
||||||
|
| head -n 1 | tr -d ' '
|
||||||
|
}
|
||||||
|
|
||||||
|
now_ms() {
|
||||||
|
python3 -c 'import time; print(int(time.time() * 1000))'
|
||||||
|
}
|
||||||
|
|
||||||
|
validate_response() {
|
||||||
|
local response=$1 expected=$2
|
||||||
|
python3 - "${response}" "${expected}" <<'PY'
|
||||||
|
import json, math, sys
|
||||||
|
data = json.load(open(sys.argv[1], encoding="utf-8"))
|
||||||
|
expected = int(sys.argv[2])
|
||||||
|
assert data["n_requested"] == expected, data
|
||||||
|
assert data["n_succeeded"] == expected, data
|
||||||
|
assert data["errors"] == [], data["errors"]
|
||||||
|
assert len(data["predictions"]) == expected, len(data["predictions"])
|
||||||
|
for pred in data["predictions"]:
|
||||||
|
values = list(pred["biodist"].values())
|
||||||
|
assert len(values) == 7
|
||||||
|
assert all(math.isfinite(float(value)) for value in values)
|
||||||
|
assert abs(sum(values) - 1.0) < 1e-4, sum(values)
|
||||||
|
print(f"validated {expected} predictions")
|
||||||
|
PY
|
||||||
|
}
|
||||||
|
|
||||||
|
run_payload() {
|
||||||
|
local label=$1 payload=$2 expected=$3
|
||||||
|
local response="/tmp/lnp-${label}-response.json"
|
||||||
|
local http_file="/tmp/lnp-${label}-http.txt"
|
||||||
|
local before peak current curl_pid curl_status http_code start_ms end_ms elapsed delta status
|
||||||
|
|
||||||
|
echo
|
||||||
|
echo "===== ${label} ====="
|
||||||
|
before=$(gpu_used_mib)
|
||||||
|
peak=${before}
|
||||||
|
start_ms=$(now_ms)
|
||||||
|
|
||||||
|
curl -sS --max-time "${CURL_TIMEOUT}" -o "${response}" -w '%{http_code}' \
|
||||||
|
-H 'Content-Type: application/json' --data-binary "@${payload}" \
|
||||||
|
"${API_URL}/predict/batch" >"${http_file}" &
|
||||||
|
curl_pid=$!
|
||||||
|
while kill -0 "${curl_pid}" 2>/dev/null; do
|
||||||
|
current=$(gpu_used_mib 2>/dev/null || echo 0)
|
||||||
|
if [[ "${current}" =~ ^[0-9]+$ ]] && [ "${current}" -gt "${peak}" ]; then
|
||||||
|
peak=${current}
|
||||||
|
fi
|
||||||
|
sleep 0.1
|
||||||
|
done
|
||||||
|
wait "${curl_pid}"
|
||||||
|
curl_status=$?
|
||||||
|
end_ms=$(now_ms)
|
||||||
|
http_code=$(cat "${http_file}" 2>/dev/null || echo 000)
|
||||||
|
elapsed=$((end_ms - start_ms))
|
||||||
|
delta=$((peak - before))
|
||||||
|
status=PASS
|
||||||
|
|
||||||
|
if [ "${curl_status}" -ne 0 ] || [ "${http_code}" != "200" ]; then
|
||||||
|
status=FAIL
|
||||||
|
echo "request failed: curl_status=${curl_status}, HTTP=${http_code}" >&2
|
||||||
|
sed -n '1,30p' "${response}" 2>/dev/null || true
|
||||||
|
elif ! validate_response "${response}" "${expected}"; then
|
||||||
|
status=FAIL
|
||||||
|
elif [ "${MAX_GPU_MIB}" -gt 0 ] && [ "${peak}" -gt "${MAX_GPU_MIB}" ]; then
|
||||||
|
status=FAIL
|
||||||
|
echo "peak ${peak} MiB exceeds MAX_GPU_MIB=${MAX_GPU_MIB}" >&2
|
||||||
|
fi
|
||||||
|
|
||||||
|
echo "${label}: status=${status}, before=${before} MiB, peak=${peak} MiB, delta=${delta} MiB, elapsed=${elapsed} ms"
|
||||||
|
printf '%s,%s,%s,%s,%s,%s,%s\n' \
|
||||||
|
"${label}" "${status}" "${before}" "${peak}" "${delta}" "${elapsed}" "${http_code}" >>"${RESULT_CSV}"
|
||||||
|
[ "${status}" = PASS ]
|
||||||
|
}
|
||||||
|
|
||||||
|
make_payload() {
|
||||||
|
local count=$1 outer_batch=$2 use_llm=$3 output=$4
|
||||||
|
local llm_flag=--use-llm
|
||||||
|
[ "${use_llm}" = false ] && llm_flag=--no-use-llm
|
||||||
|
python3 "${CASE_DIR}/generate_mock_batch.py" \
|
||||||
|
--count "${count}" --batch-size "${outer_batch}" "${llm_flag}" \
|
||||||
|
--source-csv "${CASE_DIR}/../data/interim/internal.csv" --output "${output}"
|
||||||
|
}
|
||||||
|
|
||||||
|
printf 'case,status,before_mib,peak_mib,delta_mib,elapsed_ms,http_code\n' >"${RESULT_CSV}"
|
||||||
|
FAILED=0
|
||||||
|
|
||||||
|
case "${PROFILE}" in
|
||||||
|
repro)
|
||||||
|
run_payload "llm-32-outer32" "${CASE_DIR}/mock_batch_32_llm.json" 32 || FAILED=1
|
||||||
|
run_payload "llm-500-outer32" "${CASE_DIR}/mock_batch_500_llm.json" 500 || FAILED=1
|
||||||
|
;;
|
||||||
|
ladder)
|
||||||
|
for count in 1 8 16 32 64 128; do
|
||||||
|
payload="/tmp/lnp-ladder-${count}.json"
|
||||||
|
make_payload "${count}" "${count}" true "${payload}"
|
||||||
|
run_payload "llm-${count}-outer${count}" "${payload}" "${count}" || FAILED=1
|
||||||
|
curl -fsS --max-time 30 "${API_URL}/" >/dev/null || break
|
||||||
|
done
|
||||||
|
;;
|
||||||
|
full)
|
||||||
|
run_payload "baseline-no-llm-128" "${CASE_DIR}/mock_batch_128_no_llm.json" 128 || FAILED=1
|
||||||
|
for count in 1 8 16 32 64 128; do
|
||||||
|
payload="/tmp/lnp-ladder-${count}.json"
|
||||||
|
make_payload "${count}" "${count}" true "${payload}"
|
||||||
|
run_payload "llm-${count}-outer${count}" "${payload}" "${count}" || FAILED=1
|
||||||
|
curl -fsS --max-time 30 "${API_URL}/" >/dev/null || break
|
||||||
|
done
|
||||||
|
curl -fsS --max-time 30 "${API_URL}/" >/dev/null \
|
||||||
|
&& run_payload "llm-500-outer32" "${CASE_DIR}/mock_batch_500_llm.json" 500 || FAILED=1
|
||||||
|
;;
|
||||||
|
*)
|
||||||
|
echo "Usage: $0 [repro|ladder|full]" >&2
|
||||||
|
exit 2
|
||||||
|
;;
|
||||||
|
esac
|
||||||
|
|
||||||
|
echo
|
||||||
|
echo "Results: ${RESULT_CSV}"
|
||||||
|
exit "${FAILED}"
|
||||||
10
server_test_cases/single_valid.json
Normal file
10
server_test_cases/single_valid.json
Normal file
@ -0,0 +1,10 @@
|
|||||||
|
{
|
||||||
|
"smiles": "CC(C)NCCNC(C)C",
|
||||||
|
"cationic_lipid_to_mrna_ratio": 10.0,
|
||||||
|
"cationic_lipid_mol_ratio": 35.0,
|
||||||
|
"phospholipid_mol_ratio": 16.0,
|
||||||
|
"cholesterol_mol_ratio": 46.5,
|
||||||
|
"peg_lipid_mol_ratio": 2.5,
|
||||||
|
"helper_lipid": "DOPE",
|
||||||
|
"route": "intravenous"
|
||||||
|
}
|
||||||
33
server_test_cases/ui_batch_screening_32.csv
Normal file
33
server_test_cases/ui_batch_screening_32.csv
Normal file
@ -0,0 +1,33 @@
|
|||||||
|
name,SMILES
|
||||||
|
LNP-001,O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O)CCNC(C[C@@H](CN)CC(C)C)=O
|
||||||
|
LNP-002,O=C(OCCCCCCCCCCCCCC)CCN(CCC(OCCCCCCCCCCCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCCCCCCCC)=O)CCC(OCCCCCCCCCCCCCC)=O)CCNC(CCC(NCC1=CC=C(B(O)O)C=C1)=O)=O
|
||||||
|
LNP-003,O=C(OCCCCCCN([C@H](C(NCCCN(CCCCCCOC(CCCCCCC)=O)CCCCCCOC(CCCCCCC)=O)=O)CO)CCCCCCOC(CCCCCCC)=O)CCCCCCC
|
||||||
|
LNP-004,O=C(OCCCCCCN([C@H](C(NCCCCCCN(CCCCCCOC(CCCCC)=O)CCCCCCOC(CCCCC)=O)=O)CO)CCCCCCOC(CCCCC)=O)CCCCC
|
||||||
|
LNP-005,O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCN(CC1=NC=CS1)CC2=NC=CS2)CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O
|
||||||
|
LNP-006,O=C(OCCCCCCCCCCCCCC)CCN(CCC(OCCCCCCCCCCCCCC)=O)CCCCCCN(CCCCCCN(CCC(OCCCCCCCCCCCCCC)=O)CCC(OCCCCCCCCCCCCCC)=O)CCCN1CCOCC1
|
||||||
|
LNP-007,CCCCCCCCCCCC(OCCCCCCN(CCCCCCOC(CCCCCCCCCCC)=O)CCCOC([C@@H](N(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)CO)=O)=O
|
||||||
|
LNP-008,O=C(OCCCCCCN([C@H](C(NCCN(C)CCN(CCCCCCOC(CCCCCCCCC=C)=O)CCCCCCOC(CCCCCCCCC=C)=O)=O)CO)CCCCCCOC(CCCCCCCCC=C)=O)CCCCCCCCC=C
|
||||||
|
LNP-009,CC(CNC([C@@H](N(CCCCCCOC(CCCCCCCCC=C)=O)CCCCCCOC(CCCCCCCCC=C)=O)CO)=O)(CN(CCCCCCOC(CCCCCCCCC=C)=O)CCCCCCOC(CCCCCCCCC=C)=O)C
|
||||||
|
LNP-010,O=C(NCCN(CCCCN(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)CCCCN(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)CCCN
|
||||||
|
LNP-011,O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCN1CC(C(OC)=O)CC1=O)CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O
|
||||||
|
LNP-012,CCCCCCCCCCCCOC(CCN([C@H](C(NCCCCCCN(CCC(OCCCCCCCCCCCC)=O)CCC(OCCCCCCCCCCCC)=O)=O)CO)CCC(OCCCCCCCCCCCC)=O)=O
|
||||||
|
LNP-013,O=C([C@H](CO)N(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)N(CC1)CCN1C2=NC=CC=N2
|
||||||
|
LNP-014,O=C(CN(CC(O)CCCCCCCC)CC(O)CCCCCCCC)NCCCN(CC(O)CCCCCCCC)CC(O)CCCCCCCC
|
||||||
|
LNP-015,CCCCCCCCCCCC(OCCCCCCN(CCN(C)CCNC([C@@H](N(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)CO)=O)CCCCCCOC(CCCCCCCCCCC)=O)=O
|
||||||
|
LNP-016,OC[C@@H](C(NCCCCCCN(CC(O)CCCCCCCCCC)CC(O)CCCCCCCCCC)=O)N(CC(CCCCCCCCCC)O)CC(CCCCCCCCCC)O
|
||||||
|
LNP-017,O=C(OCCCCCCCCCCCCCC)CCN(CCC(OCCCCCCCCCCCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCCCCCCCC)=O)CCC(OCCCCCCCCCCCCCC)=O)CCNC(C[C@@H](CN)CC(C)C)=O
|
||||||
|
LNP-018,O=C(OCCCCCCCC/C=C\C/C=C\CCCCC)CCN(CCC(OCCCCCCCC/C=C\C/C=C\CCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCC/C=C\C/C=C\CCCCC)=O)CCC(OCCCCCCCC/C=C\C/C=C\CCCCC)=O)CCNC(CCC(NCCC1=CC=C(O)C(O)=C1)=O)=O
|
||||||
|
LNP-019,O=C(NCCCCCCN(CCCCCCCCCCCC)CCCCCCCCCCCC)[C@H](C)N(CCCCCCCCCCCC)CCCCCCCCCCCC
|
||||||
|
LNP-020,O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCN(CCC(F)(F)F)CCC(F)(F)F)CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O
|
||||||
|
LNP-021,CCCCCCCCC(CN(CCCNC([C@@H](N(CC(CCCCCCCC)O)CC(CCCCCCCC)O)CO)=O)CC(CCCCCCCC)O)O
|
||||||
|
LNP-022,O=C([C@H](CO)N(CCCCCCOC(CCCCCCCCC=C)=O)CCCCCCOC(CCCCCCCCC=C)=O)N(CC1)CCN1C2=NC=CC=N2
|
||||||
|
LNP-023,CCCCCCCCCCCC(OCCCCCCN([C@H](C(NCCCCCCN(CCCCCCOC(CCCCCCCCCCC)=O)CCCCCCOC(CCCCCCCCCCC)=O)=O)CO)CCCCCCOC(CCCCCCCCCCC)=O)=O
|
||||||
|
LNP-024,O=C(OCCCCCCCCCCCC)CCN(CCC(OCCCCCCCCCCCC)=O)CCCCCCN(CCCCCCN(CCC(OCCCCCCCCCCCC)=O)CCC(OCCCCCCCCCCCC)=O)CCN1CCOCC1
|
||||||
|
LNP-025,O=C(CN(CCCCCCCCCCCCCC)CCCCCCCCCCCCCC)NCCCN(CCCCCCCCCCCCCC)CCCCCCCCCCCCCC
|
||||||
|
LNP-026,O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCCCN(CCCCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O)CCNC(CCC(NCCC1=CC=C(O)C(O)=C1)=O)=O
|
||||||
|
LNP-027,O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O)CCNC([C@@H](N)CCCCN)=O
|
||||||
|
LNP-028,O=C(CN(CCCCCCCCCCCC)CCCCCCCCCCCC)NCCCN(CCCCCCCCCCCC)CCCCCCCCCCCC
|
||||||
|
LNP-029,O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCN(CC1=COC=N1)CC2=COC=N2)CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O
|
||||||
|
LNP-030,O=C(OCCCCCCCC)CCN(CCC(OCCCCCCCC)=O)CCCCN(CCNC(CCCN)=O)CCCCN(CCC(OCCCCCCCC)=O)CCC(OCCCCCCCC)=O
|
||||||
|
LNP-031,O=C(CN(CCCCCCCCCC)CCCCCCCCCC)NCCCN(CCCCCCCCCC)CCCCCCCCCC
|
||||||
|
LNP-032,O=C(OCCCCCCCC/C=C\C/C=C\CCCCC)CCN(CCC(OCCCCCCCC/C=C\C/C=C\CCCCC)=O)CCCCN(CCCCN(CCC(OCCCCCCCC/C=C\C/C=C\CCCCC)=O)CCC(OCCCCCCCC/C=C\C/C=C\CCCCC)=O)CCNC(CCCN)=O
|
||||||
|
3
server_test_cases/ui_batch_screening_matrix.csv
Normal file
3
server_test_cases/ui_batch_screening_matrix.csv
Normal file
@ -0,0 +1,3 @@
|
|||||||
|
head,Tail-A,Tail-B
|
||||||
|
H-01,CC(C)NCCNC(C)C,CCCCCCCCCCN(C)C
|
||||||
|
H-02,CCN(CC)CCOC(=O)CCCC,CCCCCCCCCCCCN(C)C
|
||||||
|
7
server_test_cases/ui_batch_screening_mixed.csv
Normal file
7
server_test_cases/ui_batch_screening_mixed.csv
Normal file
@ -0,0 +1,7 @@
|
|||||||
|
name,SMILES
|
||||||
|
valid-1,CC(C)NCCNC(C)C
|
||||||
|
valid-2,CCCCCCCCCCN(C)C
|
||||||
|
invalid-backend-1,C1INVALIDCCCC
|
||||||
|
filtered-short,abc
|
||||||
|
duplicate-valid-1,CC(C)NCCNC(C)C
|
||||||
|
invalid-backend-2,NOT_A_VALID_CCCC
|
||||||
|
58
tests/test_llm_memory.py
Normal file
58
tests/test_llm_memory.py
Normal file
@ -0,0 +1,58 @@
|
|||||||
|
"""LLM 显存保护逻辑的轻量测试(不需要下载模型)。"""
|
||||||
|
|
||||||
|
from types import MethodType, SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from lnp_ml.modeling.layers.llm_prompt import LLMPromptEncoder, _positive_int_env
|
||||||
|
|
||||||
|
|
||||||
|
def _empty_encoder(micro_batch: int = 2) -> LLMPromptEncoder:
|
||||||
|
encoder = LLMPromptEncoder.__new__(LLMPromptEncoder)
|
||||||
|
torch.nn.Module.__init__(encoder)
|
||||||
|
encoder.inference_batch_size = micro_batch
|
||||||
|
return encoder
|
||||||
|
|
||||||
|
|
||||||
|
def test_invalid_inference_batch_env_falls_back() -> None:
|
||||||
|
with patch.dict("os.environ", {"LLM_INFERENCE_BATCH_SIZE": "invalid"}):
|
||||||
|
assert _positive_int_env("LLM_INFERENCE_BATCH_SIZE", 4) == 4
|
||||||
|
with patch.dict("os.environ", {"LLM_INFERENCE_BATCH_SIZE": "0"}):
|
||||||
|
assert _positive_int_env("LLM_INFERENCE_BATCH_SIZE", 4) == 4
|
||||||
|
|
||||||
|
|
||||||
|
def test_softrag_inference_is_microbatched() -> None:
|
||||||
|
encoder = _empty_encoder(micro_batch=2).eval()
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def fake_chunk(self, smiles, chem, tab, device):
|
||||||
|
calls.append((list(smiles), chem.clone(), tab.clone()))
|
||||||
|
return chem[:, 0, :] + tab[:, 0, :]
|
||||||
|
|
||||||
|
encoder._encode_softrag_chunk = MethodType(fake_chunk, encoder)
|
||||||
|
chem = torch.arange(5 * 2 * 3, dtype=torch.float32).reshape(5, 2, 3)
|
||||||
|
tab = torch.ones(5, 1, 3)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
output = encoder._encode_softrag(
|
||||||
|
[f"SMILES-{i}" for i in range(5)], chem, tab, torch.device("cpu")
|
||||||
|
)
|
||||||
|
|
||||||
|
assert [len(smiles) for smiles, _, _ in calls] == [2, 2, 1]
|
||||||
|
assert torch.equal(output, chem[:, 0, :] + tab[:, 0, :])
|
||||||
|
|
||||||
|
|
||||||
|
def test_qwen_forward_disables_kv_cache() -> None:
|
||||||
|
encoder = _empty_encoder()
|
||||||
|
encoder._is_qwen = True
|
||||||
|
received = {}
|
||||||
|
|
||||||
|
class FakeBackbone(torch.nn.Module):
|
||||||
|
def forward(self, **kwargs):
|
||||||
|
received.update(kwargs)
|
||||||
|
return SimpleNamespace(last_hidden_state=torch.zeros(1, 1, 1))
|
||||||
|
|
||||||
|
encoder.encoder = FakeBackbone()
|
||||||
|
encoder._encoder_forward(input_ids=torch.ones(1, 1, dtype=torch.long))
|
||||||
|
assert received["use_cache"] is False
|
||||||
Loading…
x
Reference in New Issue
Block a user