mofscreen-agent-api / tools /six_step_screening.py
andy88836's picture
Deploy MOFScreen-Agent FastAPI backend
e5bcce8 verified
Raw
History Blame Contribute Delete
21.3 kB
"""Deterministic six-step MOF screening tools backed by precomputed tables."""
from __future__ import annotations
import math
import re
from functools import lru_cache
from pathlib import Path
from typing import Any
import joblib
import numpy as np
import pandas as pd
ROOT = Path(__file__).resolve().parent.parent
ALL_DES_ROOT = ROOT.parent
DOC_ROOT = ALL_DES_ROOT.parent
REVISE_ROOT = ALL_DES_ROOT / "0629_revise"
ORDER_ROOT = REVISE_ROOT / "Screening" / "筛选顺序" / "筛选顺序"
PRICE_ROOT = DOC_ROOT / "price_filter" / "mof_price_0707"
STEP1_FULL = REVISE_ROOT / "screen_result" / "all_mofs_sorted_by_balanced_score_with_YEYVOO_clean5b_only.csv"
STEP1_TOP20 = REVISE_ROOT / "screen_result" / "top20_percent_mofs_by_balanced_score_with_YEYVOO_clean5b_only.csv"
STEP2_DIR = ORDER_ROOT / "1重金属结果"
STEP3_SA = ORDER_ROOT / "2 配体可合成性" / "All_MOF_SA_Ranking.csv"
STEP3_TOP6000 = ORDER_ROOT / "3 水生物毒性" / "Top6000_MOF_MaxSA_Linker.csv"
STEP4_TOX = ORDER_ROOT / "3 水生物毒性" / "Strict_EasySynth_Linkers_Optimized_Predictions.csv"
STEP5_DIR = ORDER_ROOT / "4 PMT"
STEP5_DESC = STEP5_DIR / "PMT描述符.csv"
STEP6_PRICE = PRICE_ROOT / "Strict_EasySynth_Linkers_Toxicity_Ranking0705_CoPriNet_price_ranked.csv"
SIX_STEP_TOOLS = [
"adsorption_screen",
"heavy_metal",
"ligand_sa",
"aquatic_toxicity",
"pmt",
"price",
]
ALLOWED_METALS = {"Mg", "Al", "Ca", "Ti", "Mn", "Fe", "Cu", "Zn", "Zr", "Ag"}
ALL_METALS = {
"Li", "Be", "Na", "Mg", "Al", "K", "Ca", "Sc", "Ti", "V", "Cr", "Mn", "Fe", "Co", "Ni",
"Cu", "Zn", "Ga", "Rb", "Sr", "Y", "Zr", "Nb", "Mo", "Tc", "Ru", "Rh", "Pd", "Ag", "Cd",
"In", "Sn", "Cs", "Ba", "La", "Ce", "Pr", "Nd", "Pm", "Sm", "Eu", "Gd", "Tb", "Dy", "Ho",
"Er", "Tm", "Yb", "Lu", "Hf", "Ta", "W", "Re", "Os", "Ir", "Pt", "Au", "Hg", "Tl", "Pb",
"Bi", "Po", "Fr", "Ra", "Ac", "Th", "Pa", "U", "Np", "Pu",
}
def normalize_mof_name(value: Any) -> str:
"""Normalize MOF ids for table and filename matching."""
if value is None:
return ""
return re.sub(r"[\s\-_()]+", "", str(value)).lower()
def _smiles_key(value: Any) -> str:
return "" if value is None else re.sub(r"\s+", "", str(value))
@lru_cache(maxsize=16)
def _csv(path: str) -> pd.DataFrame:
return pd.read_csv(path)
def _read(path: Path) -> pd.DataFrame:
return _csv(str(path))
def _find_by_name(df: pd.DataFrame, name: str | None, columns: tuple[str, ...] = ("MOF_Name", "MOF")) -> pd.Series | None:
if not name:
return None
key = normalize_mof_name(name)
for column in columns:
if column not in df.columns:
continue
mask = df[column].map(normalize_mof_name) == key
if mask.any():
return df.loc[mask].iloc[0]
return None
def _find_by_smiles(df: pd.DataFrame, smiles: str | None, column: str = "SMILES") -> pd.Series | None:
if not smiles or column not in df.columns:
return None
key = _smiles_key(smiles)
mask = df[column].map(_smiles_key) == key
if mask.any():
return df.loc[mask].iloc[0]
return None
def _num(value: Any) -> float | None:
try:
if pd.isna(value):
return None
value = float(value)
if math.isnan(value) or math.isinf(value):
return None
return value
except Exception:
return None
def _int(value: Any) -> int | None:
number = _num(value)
return None if number is None else int(number)
def _jsonable(value: Any) -> Any:
if isinstance(value, dict):
return {str(k): _jsonable(v) for k, v in value.items()}
if isinstance(value, list | tuple):
return [_jsonable(v) for v in value]
if isinstance(value, np.ndarray):
return _jsonable(value.tolist())
if isinstance(value, np.generic):
return _jsonable(value.item())
if isinstance(value, float) and (math.isnan(value) or math.isinf(value)):
return None
return value
def _unknown(tool: str, reason: str, step: int) -> dict[str, Any]:
return {"tool": tool, "step": step, "status": "unknown", "reason": reason, "pass": None}
def resolve_candidate(cif_path: str | None = None, user_text: str = "") -> dict[str, Any]:
"""Resolve uploaded filename or user text to a known MOF candidate."""
query_names: list[str] = []
if cif_path:
query_names.append(Path(cif_path).stem)
query_names.extend(re.findall(r"[A-Za-z0-9]+(?:[-_][A-Za-z0-9]+)*(?:-\(id[:_]\d+\))?", user_text or ""))
candidate: dict[str, Any] = {
"cif_path": cif_path,
"query_names": list(dict.fromkeys([q for q in query_names if q])),
"matched_mof": None,
"linker_smiles": None,
"match_source": None,
}
for path, columns, source in [
(STEP1_FULL, ("MOF_Name", "MOF"), "adsorption_screen"),
(STEP3_SA, ("MOF", "MOF_Name"), "ligand_sa"),
(STEP3_TOP6000, ("MOF", "MOF_Name"), "top6000_linker"),
(STEP6_PRICE, ("MOF", "MOF_Name"), "price"),
]:
if not path.exists():
continue
df = _read(path)
for name in candidate["query_names"]:
row = _find_by_name(df, name, columns)
if row is not None:
matched = row.get("MOF") if "MOF" in row.index else row.get("MOF_Name")
candidate["matched_mof"] = str(matched)
candidate["match_source"] = source
if "SMILES" in row.index and pd.notna(row.get("SMILES")):
candidate["linker_smiles"] = str(row.get("SMILES"))
return candidate
return candidate
def adsorption_screen_tool(candidate: dict[str, Any]) -> dict[str, Any]:
"""Step 1: precomputed top-20% adsorption screen."""
if not STEP1_FULL.exists() or not STEP1_TOP20.exists():
return _unknown("adsorption_screen", "adsorption ranking table is missing", 1)
full = _read(STEP1_FULL)
top = _read(STEP1_TOP20)
name = candidate.get("matched_mof") or next(iter(candidate.get("query_names", [])), None)
row = _find_by_name(full, name)
if row is None:
return _unknown("adsorption_screen", "MOF was not found in the precomputed adsorption ranking", 1)
mof = str(row.get("MOF_Name"))
top_row = _find_by_name(top, mof)
rank = _int(row.get("rank"))
cutoff = len(top)
passed = top_row is not None
candidate["matched_mof"] = mof
return {
"tool": "adsorption_screen",
"step": 1,
"status": "pass" if passed else "fail",
"pass": passed,
"mof": mof,
"rank": rank,
"top20_cutoff_rank": cutoff,
"predicted_benzene_adsorption": _num(row.get("predicted_benzene_adsorption")),
"predicted_toluene_adsorption": _num(row.get("predicted_toluene_adsorption")),
"balanced_score": _num(row.get("balanced_score")),
"model": "precomputed benzene/toluene adsorption ranking from 0629_revise/screen_result",
}
def heavy_metal_tool(candidate: dict[str, Any]) -> dict[str, Any]:
"""Step 2: allowed-metal filter from uploaded CIF."""
cif_path = candidate.get("cif_path")
if not cif_path:
return _unknown("heavy_metal", "no CIF file was provided for metal parsing", 2)
try:
from pymatgen.core import Structure
structure = Structure.from_file(cif_path)
elements = sorted({str(site.specie.symbol) for site in structure})
except Exception as exc:
return _unknown("heavy_metal", f"failed to parse CIF metals: {exc}", 2)
metals = [element for element in elements if element in ALL_METALS]
illegal = [element for element in metals if element not in ALLOWED_METALS]
passed = len(illegal) == 0
return {
"tool": "heavy_metal",
"step": 2,
"status": "pass" if passed else "fail",
"pass": passed,
"detected_metals": metals,
"illegal_metals": illegal,
"allowed_metals": sorted(ALLOWED_METALS),
"model": "pymatgen CIF parser + fixed allowed-metal list",
}
def ligand_sa_tool(candidate: dict[str, Any]) -> dict[str, Any]:
"""Step 3: ligand synthesizability lookup."""
if not STEP3_SA.exists() or not STEP3_TOP6000.exists():
return _unknown("ligand_sa", "ligand SA tables are missing", 3)
name = candidate.get("matched_mof") or next(iter(candidate.get("query_names", [])), None)
sa = _read(STEP3_SA)
top = _read(STEP3_TOP6000)
row = _find_by_name(sa, name)
if row is None:
return _unknown("ligand_sa", "MOF was not found in All_MOF_SA_Ranking.csv", 3)
mof = str(row.get("MOF"))
top_row = _find_by_name(top, mof)
passed = top_row is not None
smiles = str(top_row.get("SMILES")) if top_row is not None and pd.notna(top_row.get("SMILES")) else None
candidate["matched_mof"] = mof
if smiles:
candidate["linker_smiles"] = smiles
return {
"tool": "ligand_sa",
"step": 3,
"status": "pass" if passed else "fail",
"pass": passed,
"mof": mof,
"sa_rank": _int(row.get("SA_Rank")),
"sa_score": _num(row.get("Max_SA")),
"mean_sa": _num(row.get("Mean_SA")),
"n_linker": _int(row.get("NLinker")),
"in_top6000": passed,
"linker_id": _int(top_row.get("Linker_ID")) if top_row is not None else None,
"linker_smiles": smiles,
"model": "All_MOF_SA_Ranking.csv + Top6000_MOF_MaxSA_Linker.csv lookup",
}
def aquatic_toxicity_tool(candidate: dict[str, Any]) -> dict[str, Any]:
"""Step 4: aquatic toxicity lookup by linker SMILES."""
if not STEP4_TOX.exists():
return _unknown("aquatic_toxicity", "aquatic toxicity prediction table is missing", 4)
if not candidate.get("linker_smiles"):
ligand_sa_tool(candidate)
smiles = candidate.get("linker_smiles")
if not smiles:
return _unknown("aquatic_toxicity", "no linker SMILES was available for toxicity lookup", 4)
row = _find_by_smiles(_read(STEP4_TOX), smiles)
if row is None:
return _unknown("aquatic_toxicity", "linker SMILES was not found in toxicity predictions", 4)
values = [
_num(row.get("Predicted_Tox_IBC50")),
_num(row.get("Predicted_Tox_IGC50")),
_num(row.get("Predicted_Tox_LC50")),
_num(row.get("Predicted_Tox_LC50DM")),
]
present = [v for v in values if v is not None]
return {
"tool": "aquatic_toxicity",
"step": 4,
"status": "pass" if len(present) == 4 else "unknown",
"pass": True if len(present) == 4 else None,
"linker_smiles": smiles,
"linker_id": _int(row.get("Linker_ID")),
"IBC50": values[0],
"IGC50": values[1],
"LC50": values[2],
"LC50DM": values[3],
"mean_toxicity": round(sum(present) / len(present), 6) if present else None,
"worst_toxicity": min(present) if present else None,
"model": "Strict_EasySynth_Linkers_Optimized_Predictions.csv lookup",
}
def pmt_tool(candidate: dict[str, Any]) -> dict[str, Any]:
"""Step 5: PMT classifier lookup + model inference from precomputed descriptors."""
if not STEP5_DESC.exists():
return _unknown("pmt", "PMT descriptor table is missing", 5)
if not candidate.get("linker_smiles"):
ligand_sa_tool(candidate)
smiles = candidate.get("linker_smiles")
if not smiles:
return _unknown("pmt", "no linker SMILES was available for PMT descriptor lookup", 5)
desc = _read(STEP5_DESC)
row = _find_by_smiles(desc, smiles)
if row is None:
return _unknown("pmt", "linker SMILES was not found in PMT描述符.csv", 5)
try:
imputer = joblib.load(STEP5_DIR / "PMT_imputer.pkl")
scaler = joblib.load(STEP5_DIR / "PMT_scaler.pkl")
selector = joblib.load(STEP5_DIR / "PMT_selector.pkl")
model = joblib.load(STEP5_DIR / "PMT_xgb_model.pkl")
feature_row = row.drop(labels=["SMILES"], errors="ignore").to_frame().T
feature_row = feature_row.apply(pd.to_numeric, errors="coerce")
x = imputer.transform(feature_row)
x = scaler.transform(x)
x = selector.transform(x)
proba = model.predict_proba(x)[0]
classes = list(getattr(model, "classes_", [0, 1]))
positive_index = classes.index(1) if 1 in classes else len(proba) - 1
probability = float(proba[positive_index])
except Exception as exc:
return _unknown("pmt", f"PMT model inference failed: {exc}", 5)
threshold = 0.4
passed = probability >= threshold
return {
"tool": "pmt",
"step": 5,
"status": "pass" if passed else "fail",
"pass": passed,
"linker_smiles": smiles,
"pmt_probability": round(probability, 6),
"pmt_class": "non_PMT" if passed else "PMT_risk",
"threshold": threshold,
"threshold_note": "Class 1 is treated as non-PMT/pass; probability >= 0.4 passes.",
"model": "PMT_xgb_model.pkl with PMT_imputer/scaler/selector",
}
def price_tool(candidate: dict[str, Any]) -> dict[str, Any]:
"""Step 6: CoPriNet price lookup."""
if not STEP6_PRICE.exists():
return _unknown("price", "price ranking table is missing", 6)
price = _read(STEP6_PRICE)
row = _find_by_name(price, candidate.get("matched_mof"))
if row is None and candidate.get("linker_smiles"):
row = _find_by_smiles(price, candidate.get("linker_smiles"))
if row is None:
for name in candidate.get("query_names", []):
row = _find_by_name(price, name)
if row is not None:
break
if row is None:
return _unknown("price", "MOF/linker was not found in the CoPriNet price ranking", 6)
status = str(row.get("Price_Prediction_Status", "unknown"))
passed = status.upper() == "OK"
if pd.notna(row.get("SMILES")):
candidate["linker_smiles"] = str(row.get("SMILES"))
if pd.notna(row.get("MOF")):
candidate["matched_mof"] = str(row.get("MOF"))
return {
"tool": "price",
"step": 6,
"status": "pass" if passed else "fail",
"pass": passed,
"mof": str(row.get("MOF")),
"linker_smiles": str(row.get("SMILES")) if pd.notna(row.get("SMILES")) else None,
"coprinet_price_rank": _int(row.get("CoPriNet_Price_Rank")),
"usd_per_g": _num(row.get("CoPriNet_USD_per_g")),
"usd_per_mmol": _num(row.get("CoPriNet_USD_per_mmol")),
"price_status": status,
"model": "CoPriNet price ranking CSV lookup; independent of PMT",
}
def _run_tool(name: str, candidate: dict[str, Any]) -> dict[str, Any]:
tools = {
"adsorption_screen": adsorption_screen_tool,
"heavy_metal": heavy_metal_tool,
"ligand_sa": ligand_sa_tool,
"aquatic_toxicity": aquatic_toxicity_tool,
"pmt": pmt_tool,
"price": price_tool,
}
if name not in tools:
return _unknown(name, "unknown six-step tool name", 0)
return _jsonable(tools[name](candidate))
def run_selected_six_step_tools(cif_path: str | None, user_text: str, tools: list[str]) -> dict[str, Any]:
candidate = resolve_candidate(cif_path, user_text)
results: dict[str, Any] = {}
trace: list[dict[str, Any]] = []
for name in tools:
if name in {"aquatic_toxicity", "pmt", "price"} and not candidate.get("linker_smiles"):
ligand = results.get("ligand_sa") or _run_tool("ligand_sa", candidate)
results.setdefault("ligand_sa", ligand)
trace.append({"agent": "Tool Execution Agent", "action": "resolved linker SMILES via ligand_sa", "output": ligand})
result = _run_tool(name, candidate)
results[name] = result
trace.append({"agent": "Tool Execution Agent", "action": f"ran {name}", "output": result})
payload = _final_payload(candidate, results, full=False)
payload["agent_trace"] = trace + payload["agent_trace"]
return payload
def run_six_step_screening(cif_path: str | None, user_text: str = "") -> dict[str, Any]:
candidate = resolve_candidate(cif_path, user_text)
results: dict[str, Any] = {}
trace: list[dict[str, Any]] = []
for name in SIX_STEP_TOOLS:
result = _run_tool(name, candidate)
results[name] = result
trace.append({"agent": "Tool Execution Agent", "action": f"ran step {result.get('step')}: {name}", "output": result})
payload = _final_payload(candidate, results, full=True)
payload["agent_trace"] = trace + payload["agent_trace"]
return payload
def _final_payload(candidate: dict[str, Any], results: dict[str, Any], full: bool) -> dict[str, Any]:
failed = [r for r in results.values() if r.get("status") == "fail"]
unknown = [r for r in results.values() if r.get("status") == "unknown"]
if failed:
failed_first = sorted(failed, key=lambda r: r.get("step", 999))[0]
gate_status = "failed"
failed_at_step = failed_first.get("step")
recommendation = f"failed_at_step_{failed_at_step}_{failed_first.get('tool')}"
elif unknown:
gate_status = "incomplete"
failed_at_step = None
recommendation = "incomplete_due_to_unknown_evidence"
else:
gate_status = "pass_full_screening" if full else "pass_selected_tools"
failed_at_step = None
recommendation = gate_status
decision_record = {
"decision_class": gate_status,
"recommendation": recommendation,
"failed_at_step": failed_at_step,
"unknown_steps": [r.get("step") for r in unknown],
"blocking_tools": [r.get("tool") for r in failed],
"full_screening": full,
}
payload = {
"mof_id": candidate.get("matched_mof") or next(iter(candidate.get("query_names", [])), "unknown"),
"candidate": candidate,
"six_step": results,
"results": results,
"gate_status": gate_status,
"failed_at_step": failed_at_step,
"recommendation": recommendation,
"decision_record": decision_record,
"final_score": _score_from_status(gate_status),
"explanation": _summary(results, gate_status, failed_at_step),
"warnings": [],
"errors": [],
"agent_trace": [{
"agent": "Decision Agent",
"action": "computed deterministic six-step gate status",
"output": decision_record,
}],
}
payload["row"] = build_result_row(payload)
payload["evidence_ledger"] = [
{"tool": name, "status": result.get("status"), "outputs": result, "confidence": "precomputed_or_deterministic"}
for name, result in results.items()
]
return _jsonable(payload)
def _score_from_status(status: str) -> float | None:
if status.startswith("pass"):
return 10.0
if status == "incomplete":
return 5.0
if status == "failed":
return 0.0
return None
def _summary(results: dict[str, Any], gate_status: str, failed_at_step: int | None) -> str:
parts = []
for name in SIX_STEP_TOOLS:
result = results.get(name)
if not result:
continue
parts.append(f"step {result.get('step')} {name}: {result.get('status')}")
if failed_at_step:
tail = f"Final gate status is failed at step {failed_at_step}."
elif gate_status == "incomplete":
tail = "Final gate status is incomplete because at least one required tool returned unknown."
else:
tail = f"Final gate status is {gate_status}."
return "; ".join(parts + [tail])
def _fmt_status(result: dict[str, Any] | None) -> str:
if not result:
return "N/A"
return str(result.get("status", "unknown"))
def build_result_row(result: dict[str, Any]) -> dict[str, Any]:
steps = result.get("six_step") or result.get("results") or {}
adsorption = steps.get("adsorption_screen") or {}
metal = steps.get("heavy_metal") or {}
sa = steps.get("ligand_sa") or {}
tox = steps.get("aquatic_toxicity") or {}
pmt = steps.get("pmt") or {}
price = steps.get("price") or {}
return {
"MOF ID": result.get("mof_id", "unknown"),
"Step 1 adsorption": (
f"{_fmt_status(adsorption)}; rank={adsorption.get('rank')}; "
f"B={adsorption.get('predicted_benzene_adsorption')}; T={adsorption.get('predicted_toluene_adsorption')}"
),
"Step 2 metal": (
f"{_fmt_status(metal)}; metals={', '.join(metal.get('detected_metals', []) or [])}; "
f"illegal={', '.join(metal.get('illegal_metals', []) or [])}"
),
"Step 3 SA": f"{_fmt_status(sa)}; rank={sa.get('sa_rank')}; score={sa.get('sa_score')}",
"Step 4 toxicity": (
f"{_fmt_status(tox)}; mean={tox.get('mean_toxicity')}; worst={tox.get('worst_toxicity')}"
),
"Step 5 PMT": (
f"{_fmt_status(pmt)}; class={pmt.get('pmt_class')}; prob={pmt.get('pmt_probability')}"
),
"Step 6 price": (
f"{_fmt_status(price)}; rank={price.get('coprinet_price_rank')}; USD/g={price.get('usd_per_g')}"
),
"Final gate status": result.get("gate_status", "unknown"),
"Recommendation": result.get("recommendation", "N/A"),
"Score": result.get("final_score"),
}