mirror of
https://github.com/RYDE-WORK/lnp_ml.git
synced 2026-09-18 14:23:20 +08:00
fix: 收窄优化搜索域,新增 /predict 端点
This commit is contained in:
parent
aeddf8e1ee
commit
48b22cfe86
170
app/api.py
170
app/api.py
@ -25,23 +25,28 @@ from app.optimize import (
|
||||
TARGET_BIODIST,
|
||||
CompRanges,
|
||||
ScoringWeights,
|
||||
HELPER_LIPID_OPTIONS,
|
||||
ROUTE_OPTIONS,
|
||||
create_dataframe_from_formulations,
|
||||
predict_all,
|
||||
)
|
||||
|
||||
|
||||
# ============ Pydantic Models ============
|
||||
|
||||
class CompRangesRequest(BaseModel):
|
||||
"""组分范围配置(mol 比例为百分数 0-100)"""
|
||||
weight_ratio_min: float = Field(default=5.0, ge=1.0, le=50.0, description="阳离子脂质/mRNA 重量比最小值")
|
||||
weight_ratio_max: float = Field(default=30.0, ge=1.0, le=50.0, description="阳离子脂质/mRNA 重量比最大值")
|
||||
cationic_mol_min: float = Field(default=5.0, ge=0.0, le=100.0, description="阳离子脂质 mol 比例最小值 (%)")
|
||||
cationic_mol_max: float = Field(default=80.0, ge=0.0, le=100.0, description="阳离子脂质 mol 比例最大值 (%)")
|
||||
phospholipid_mol_min: float = Field(default=0.0, ge=0.0, le=100.0, description="磷脂 mol 比例最小值 (%)")
|
||||
phospholipid_mol_max: float = Field(default=80.0, ge=0.0, le=100.0, description="磷脂 mol 比例最大值 (%)")
|
||||
cholesterol_mol_min: float = Field(default=0.0, ge=0.0, le=100.0, description="胆固醇 mol 比例最小值 (%)")
|
||||
cholesterol_mol_max: float = Field(default=80.0, ge=0.0, le=100.0, description="胆固醇 mol 比例最大值 (%)")
|
||||
peg_mol_min: float = Field(default=0.0, ge=0.0, le=20.0, description="PEG 脂质 mol 比例最小值 (%)")
|
||||
peg_mol_max: float = Field(default=5.0, ge=0.0, le=20.0, description="PEG 脂质 mol 比例最大值 (%)")
|
||||
"""组分范围配置(mol 比例为百分数 0-100)
|
||||
"""
|
||||
weight_ratio_min: float = Field(default=7.0, ge=1.0, le=50.0, description="阳离子脂质/mRNA 重量比最小值")
|
||||
weight_ratio_max: float = Field(default=20.0, ge=1.0, le=50.0, description="阳离子脂质/mRNA 重量比最大值")
|
||||
cationic_mol_min: float = Field(default=22.0, ge=0.0, le=100.0, description="阳离子脂质 mol 比例最小值 (%)")
|
||||
cationic_mol_max: float = Field(default=55.0, ge=0.0, le=100.0, description="阳离子脂质 mol 比例最大值 (%)")
|
||||
phospholipid_mol_min: float = Field(default=7.0, ge=0.0, le=100.0, description="磷脂 mol 比例最小值 (%)")
|
||||
phospholipid_mol_max: float = Field(default=42.0, ge=0.0, le=100.0, description="磷脂 mol 比例最大值 (%)")
|
||||
cholesterol_mol_min: float = Field(default=15.0, ge=0.0, le=100.0, description="胆固醇 mol 比例最小值 (%)")
|
||||
cholesterol_mol_max: float = Field(default=46.0, ge=0.0, le=100.0, description="胆固醇 mol 比例最大值 (%)")
|
||||
peg_mol_min: float = Field(default=1.0, ge=0.0, le=20.0, description="PEG 脂质 mol 比例最小值 (%)")
|
||||
peg_mol_max: float = Field(default=6.0, ge=0.0, le=20.0, description="PEG 脂质 mol 比例最大值 (%)")
|
||||
|
||||
def to_comp_ranges(self) -> CompRanges:
|
||||
"""转换为 CompRanges 对象"""
|
||||
@ -151,6 +156,70 @@ class HealthResponse(BaseModel):
|
||||
use_rag: bool = False
|
||||
|
||||
|
||||
class PredictRequest(BaseModel):
|
||||
"""单配方预测请求"""
|
||||
smiles: str = Field(..., description="Cationic lipid SMILES")
|
||||
cationic_lipid_to_mrna_ratio: float = Field(..., gt=0, le=50, description="阳离子脂质/mRNA 重量比")
|
||||
cationic_lipid_mol_ratio: float = Field(..., ge=0, le=100, description="阳离子脂质 mol 比例 (%)")
|
||||
phospholipid_mol_ratio: float = Field(..., ge=0, le=100, description="磷脂 mol 比例 (%)")
|
||||
cholesterol_mol_ratio: float = Field(..., ge=0, le=100, description="胆固醇 mol 比例 (%)")
|
||||
peg_lipid_mol_ratio: float = Field(..., ge=0, le=20, description="PEG 脂质 mol 比例 (%)")
|
||||
helper_lipid: str = Field(default="DOPE", description=f"辅助脂质,可选 {HELPER_LIPID_OPTIONS}")
|
||||
route: str = Field(default="intravenous", description=f"给药途径,可选 {ROUTE_OPTIONS}")
|
||||
|
||||
class Config:
|
||||
json_schema_extra = {
|
||||
"example": {
|
||||
"smiles": "CC(C)NCCNC(C)C",
|
||||
"cationic_lipid_to_mrna_ratio": 10.0,
|
||||
"cationic_lipid_mol_ratio": 50.0,
|
||||
"phospholipid_mol_ratio": 10.0,
|
||||
"cholesterol_mol_ratio": 38.5,
|
||||
"peg_lipid_mol_ratio": 1.5,
|
||||
"helper_lipid": "DOPE",
|
||||
"route": "intravenous",
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
class PredictResponse(BaseModel):
|
||||
"""单配方预测响应"""
|
||||
smiles: str
|
||||
helper_lipid: str
|
||||
route: str
|
||||
biodist: Dict[str, float]
|
||||
size: Optional[float] = None
|
||||
quantified_delivery: Optional[float] = None
|
||||
unnormalized_delivery: Optional[float] = None
|
||||
pdi_class: Optional[int] = None
|
||||
ee_class: Optional[int] = None
|
||||
toxic_class: Optional[int] = None
|
||||
|
||||
class Config:
|
||||
json_schema_extra = {
|
||||
"example": {
|
||||
"smiles": "CC(C)NCCNC(C)C",
|
||||
"helper_lipid": "DOPE",
|
||||
"route": "intravenous",
|
||||
"biodist": {
|
||||
"lymph_nodes": 0.0048,
|
||||
"heart": 0.0044,
|
||||
"liver": 0.6591,
|
||||
"spleen": 0.2817,
|
||||
"lung": 0.0245,
|
||||
"kidney": 0.0052,
|
||||
"muscle": 0.0203,
|
||||
},
|
||||
"size": 100.7,
|
||||
"quantified_delivery": 0.1169,
|
||||
"unnormalized_delivery": 0.3432,
|
||||
"pdi_class": 0,
|
||||
"ee_class": 2,
|
||||
"toxic_class": 0,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
# ============ Global State ============
|
||||
|
||||
class ModelState:
|
||||
@ -261,6 +330,85 @@ async def get_available_organs():
|
||||
return AVAILABLE_ORGANS
|
||||
|
||||
|
||||
@app.post("/predict", response_model=PredictResponse)
|
||||
async def predict_formulation(request: PredictRequest):
|
||||
"""对单个指定配方做属性预测(不搜索,约 0.2 秒)。"""
|
||||
import pandas as pd
|
||||
from rdkit import Chem, RDLogger
|
||||
|
||||
if state.model is None:
|
||||
raise HTTPException(status_code=503, detail="Model not loaded")
|
||||
|
||||
if request.helper_lipid not in HELPER_LIPID_OPTIONS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Invalid helper_lipid: {request.helper_lipid}. Available: {HELPER_LIPID_OPTIONS}",
|
||||
)
|
||||
if request.route not in ROUTE_OPTIONS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Invalid route: {request.route}. Available: {ROUTE_OPTIONS}",
|
||||
)
|
||||
|
||||
mol_sum = (
|
||||
request.cationic_lipid_mol_ratio
|
||||
+ request.phospholipid_mol_ratio
|
||||
+ request.cholesterol_mol_ratio
|
||||
+ request.peg_lipid_mol_ratio
|
||||
)
|
||||
if abs(mol_sum - 100.0) > 0.5:
|
||||
raise HTTPException(
|
||||
status_code=400, detail=f"四项 mol 比例之和须为 100(当前 {mol_sum:.2f})"
|
||||
)
|
||||
|
||||
RDLogger.DisableLog("rdApp.*")
|
||||
if Chem.MolFromSmiles(request.smiles) is None:
|
||||
raise HTTPException(status_code=400, detail=f"无法解析的 SMILES: {request.smiles[:80]}")
|
||||
|
||||
logger.info(f"Predict request: helper={request.helper_lipid}, route={request.route}, "
|
||||
f"smiles={request.smiles[:50]}...")
|
||||
|
||||
df = create_dataframe_from_formulations(
|
||||
request.smiles,
|
||||
[(
|
||||
request.cationic_lipid_to_mrna_ratio,
|
||||
request.cationic_lipid_mol_ratio,
|
||||
request.phospholipid_mol_ratio,
|
||||
request.cholesterol_mol_ratio,
|
||||
request.peg_lipid_mol_ratio,
|
||||
)],
|
||||
[request.helper_lipid],
|
||||
[request.route],
|
||||
)
|
||||
|
||||
try:
|
||||
if getattr(state.model, "llm_prompt", None) is not None:
|
||||
state.model.set_llm_enabled(True)
|
||||
df = predict_all(state.model, df, state.device, batch_size=1)
|
||||
except Exception as e:
|
||||
logger.error(f"Prediction failed: {e}")
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
row = df.iloc[0]
|
||||
_size, _unnorm = row.get("pred_size"), row.get("pred_unnorm_delivery")
|
||||
return PredictResponse(
|
||||
smiles=request.smiles,
|
||||
helper_lipid=request.helper_lipid,
|
||||
route=request.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"]),
|
||||
)
|
||||
|
||||
|
||||
@app.post("/optimize", response_model=OptimizeResponse)
|
||||
async def optimize_formulation(request: OptimizeRequest):
|
||||
"""
|
||||
|
||||
20
app/app.py
20
app/app.py
@ -460,37 +460,37 @@ def main():
|
||||
st.caption("阳离子脂质/mRNA 重量比")
|
||||
col1, col2 = st.columns(2)
|
||||
with col1:
|
||||
weight_ratio_min = st.number_input("最小", min_value=1.0, max_value=50.0, value=5.0, step=1.0, format="%.1f", key="wr_min")
|
||||
weight_ratio_min = st.number_input("最小", min_value=1.0, max_value=50.0, value=7.0, step=1.0, format="%.1f", key="wr_min")
|
||||
with col2:
|
||||
weight_ratio_max = st.number_input("最大", min_value=1.0, max_value=50.0, value=30.0, step=1.0, format="%.1f", key="wr_max")
|
||||
weight_ratio_max = st.number_input("最大", min_value=1.0, max_value=50.0, value=20.0, step=1.0, format="%.1f", key="wr_max")
|
||||
|
||||
st.caption("阳离子脂质 mol 比例 (%)")
|
||||
col1, col2 = st.columns(2)
|
||||
with col1:
|
||||
cationic_mol_min = st.number_input("最小", min_value=0.0, max_value=100.0, value=5.0, step=5.0, format="%.1f", key="cat_min")
|
||||
cationic_mol_min = st.number_input("最小", min_value=0.0, max_value=100.0, value=22.0, step=5.0, format="%.1f", key="cat_min")
|
||||
with col2:
|
||||
cationic_mol_max = st.number_input("最大", min_value=0.0, max_value=100.0, value=80.0, step=5.0, format="%.1f", key="cat_max")
|
||||
cationic_mol_max = st.number_input("最大", min_value=0.0, max_value=100.0, value=55.0, step=5.0, format="%.1f", key="cat_max")
|
||||
|
||||
st.caption("磷脂 mol 比例 (%)")
|
||||
col1, col2 = st.columns(2)
|
||||
with col1:
|
||||
phospholipid_mol_min = st.number_input("最小", min_value=0.0, max_value=100.0, value=0.0, step=5.0, format="%.1f", key="phos_min")
|
||||
phospholipid_mol_min = st.number_input("最小", min_value=0.0, max_value=100.0, value=7.0, step=5.0, format="%.1f", key="phos_min")
|
||||
with col2:
|
||||
phospholipid_mol_max = st.number_input("最大", min_value=0.0, max_value=100.0, value=80.0, step=5.0, format="%.1f", key="phos_max")
|
||||
phospholipid_mol_max = st.number_input("最大", min_value=0.0, max_value=100.0, value=42.0, step=5.0, format="%.1f", key="phos_max")
|
||||
|
||||
st.caption("胆固醇 mol 比例 (%)")
|
||||
col1, col2 = st.columns(2)
|
||||
with col1:
|
||||
cholesterol_mol_min = st.number_input("最小", min_value=0.0, max_value=100.0, value=0.0, step=5.0, format="%.1f", key="chol_min")
|
||||
cholesterol_mol_min = st.number_input("最小", min_value=0.0, max_value=100.0, value=15.0, step=5.0, format="%.1f", key="chol_min")
|
||||
with col2:
|
||||
cholesterol_mol_max = st.number_input("最大", min_value=0.0, max_value=100.0, value=80.0, step=5.0, format="%.1f", key="chol_max")
|
||||
cholesterol_mol_max = st.number_input("最大", min_value=0.0, max_value=100.0, value=46.0, step=5.0, format="%.1f", key="chol_max")
|
||||
|
||||
st.caption("PEG 脂质 mol 比例 (%)")
|
||||
col1, col2 = st.columns(2)
|
||||
with col1:
|
||||
peg_mol_min = st.number_input("最小", min_value=0.0, max_value=20.0, value=0.0, step=1.0, format="%.1f", key="peg_min")
|
||||
peg_mol_min = st.number_input("最小", min_value=0.0, max_value=20.0, value=1.0, step=1.0, format="%.1f", key="peg_min")
|
||||
with col2:
|
||||
peg_mol_max = st.number_input("最大", min_value=0.0, max_value=20.0, value=5.0, step=1.0, format="%.1f", key="peg_max")
|
||||
peg_mol_max = st.number_input("最大", min_value=0.0, max_value=20.0, value=6.0, step=1.0, format="%.1f", key="peg_max")
|
||||
|
||||
comp_ranges = {
|
||||
"weight_ratio_min": weight_ratio_min,
|
||||
|
||||
@ -45,22 +45,23 @@ AVAILABLE_ORGANS = ["lymph_nodes", "heart", "liver", "spleen", "lung", "kidney",
|
||||
|
||||
@dataclass
|
||||
class CompRanges:
|
||||
"""组分参数范围配置(mol 比例为百分数 0-100)"""
|
||||
# 阳离子脂质/mRNA 重量比
|
||||
weight_ratio_min: float = 5.0
|
||||
weight_ratio_max: float = 30.0
|
||||
# 阳离子脂质 mol 比例 (%)
|
||||
cationic_mol_min: float = 5.0
|
||||
cationic_mol_max: float = 80.0
|
||||
# 磷脂 mol 比例 (%)
|
||||
phospholipid_mol_min: float = 0.0
|
||||
phospholipid_mol_max: float = 80.0
|
||||
# 胆固醇 mol 比例 (%)
|
||||
cholesterol_mol_min: float = 0.0
|
||||
cholesterol_mol_max: float = 80.0
|
||||
# PEG 脂质 mol 比例 (%)
|
||||
peg_mol_min: float = 0.0
|
||||
peg_mol_max: float = 5.0
|
||||
"""组分参数范围配置(mol 比例为百分数 0-100)
|
||||
"""
|
||||
# 阳离子脂质/mRNA 重量比(数据 1%–99%: 7.52–19.35)
|
||||
weight_ratio_min: float = 7.0
|
||||
weight_ratio_max: float = 20.0
|
||||
# 阳离子脂质 mol 比例 (%)(数据 1%–99%: 22.38–52.83)
|
||||
cationic_mol_min: float = 22.0
|
||||
cationic_mol_max: float = 55.0
|
||||
# 磷脂 mol 比例 (%)(数据 1%–99%: 7.50–40.50)
|
||||
phospholipid_mol_min: float = 7.0
|
||||
phospholipid_mol_max: float = 42.0
|
||||
# 胆固醇 mol 比例 (%)(数据 1%–99%: 16.00–45.00)
|
||||
cholesterol_mol_min: float = 15.0
|
||||
cholesterol_mol_max: float = 46.0
|
||||
# PEG 脂质 mol 比例 (%)(数据 1%–99%: 1.0–6.0)
|
||||
peg_mol_min: float = 1.0
|
||||
peg_mol_max: float = 6.0
|
||||
|
||||
def to_dict(self) -> Dict:
|
||||
"""转换为字典"""
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user