Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion docs/dd/tools/admet.md
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ df = admet.run()
from deeporigin.drug_discovery import LigandSet

ligand_set = LigandSet.from_csv("ligands.csv") # "smiles" column
large = Admet(ligands=ligand_set)
large = Admet(ligands=ligand_set, batch_size=100) # optional; smaller batches run more in parallel (min 50)
large.start()
large.wait()
df = large.get_results()
Expand Down
146 changes: 49 additions & 97 deletions src/drug_discovery/admet.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,12 +22,7 @@

from __future__ import annotations

import json
import os
from pathlib import Path
import tempfile
from typing import Any, Literal, Self
import uuid

from beartype import beartype
import pandas as pd
Expand All @@ -37,6 +32,11 @@
AsyncExecutableMixin,
SyncExecutableMixin,
)
from deeporigin.drug_discovery.ligand_list_file import (
ligand_rows_from_inputs,
ligands_from_rows,
upload_ligand_list,
)
from deeporigin.drug_discovery.metabolism import _ligand_payloads
from deeporigin.drug_discovery.notebook_watch_mixin import NotebookWatchMixin
from deeporigin.drug_discovery.structures.ligand import Ligand, LigandSet
Expand All @@ -47,10 +47,11 @@
from deeporigin.platform.project_scope import require_client_project_id
from deeporigin.utils.constants import (
ADMET_EXECUTION_TIMEOUT_SECONDS,
ADMET_INLINE_LIGAND_CAP,
ADMET_LIGAND_LIST_UPLOAD_PREFIX,
ADMET_MIN_BATCH_SIZE,
ADMET_RESULT_EXPLORER_PAGE_SIZE,
ADMET_WORKFLOW_LIGAND_THRESHOLD,
INLINE_LIGAND_CAP,
QUOTE_APPROVE_AMOUNT,
)

Expand Down Expand Up @@ -155,47 +156,6 @@ def _expand_admet_payload(payload: dict[str, Any]) -> list[dict[str, Any]]:
return []


def _ligands_from_payload_rows(raw: list[Any]) -> list[Ligand]:
"""Rebuild ligands from inline or file JSON ligand rows."""

ligands: list[Ligand] = []
for idx, row in enumerate(raw):
if not isinstance(row, dict):
raise ValueError(
f"Cannot rehydrate Admet: ligands[{idx}] is not an object."
)
smiles = row.get("smiles")
if not smiles or not isinstance(smiles, str):
raise ValueError(f"Cannot rehydrate Admet: ligands[{idx}] has no SMILES.")
ligand = Ligand.from_smiles(smiles)
if row.get("id") is not None:
ligand.id = str(row["id"])
ligands.append(ligand)
return ligands


def _ligands_from_list_file_bytes(payload: bytes) -> list[Ligand]:
"""Parse a Ligand list file body into ligands."""

try:
text = payload.decode("utf-8")
except UnicodeDecodeError as exc:
raise ValueError(
f"Cannot rehydrate Admet: ligands_file is not valid UTF-8: {exc}"
) from exc
try:
parsed = json.loads(text)
except json.JSONDecodeError as exc:
raise ValueError(
f"Cannot rehydrate Admet: ligands_file is not valid JSON: {exc.msg}"
) from exc
if not isinstance(parsed, list) or not parsed:
raise ValueError(
"Cannot rehydrate Admet: ligands_file must be a non-empty JSON array."
)
return _ligands_from_payload_rows(parsed)


def _properties_from_inputs(
inputs: dict[str, Any],
) -> tuple[str, ...] | None:
Expand Down Expand Up @@ -233,31 +193,13 @@ def _ligands_from_inputs(
if isinstance(project, dict) and project.get("id"):
return []

raw = inputs.get("ligands")
if isinstance(raw, list) and raw:
return _ligands_from_payload_rows(raw)

remote = inputs.get("ligands_file")
if isinstance(remote, str) and remote.strip():
if client is None or client.files is None:
raise ValueError(
"Cannot rehydrate Admet: client with files is required "
"to download ligands_file."
)
try:
local_path = client.files.download(remote.strip(), direct=True)
payload = Path(local_path).read_bytes()
except Exception as exc:
raise ValueError(
f"Cannot rehydrate Admet: failed to download ligands_file "
f"{remote!r}: {exc}"
) from exc
return _ligands_from_list_file_bytes(payload)

raise ValueError(
"Cannot rehydrate Admet: stored inputs have no ligands, ligands_file, "
"or project."
)
rows = ligand_rows_from_inputs(inputs, client=client, label="Admet")
if not rows:
raise ValueError(
"Cannot rehydrate Admet: stored inputs have no ligands, ligands_file, "
"or project."
)
return ligands_from_rows(rows, label="Admet")


class Admet(
Expand All @@ -275,7 +217,7 @@ class Admet(
``run(quote=True)`` which sends SMILES without syncing first.

Use :meth:`run` for at most
:data:`~deeporigin.utils.constants.ADMET_INLINE_LIGAND_CAP` ligands
:data:`~deeporigin.utils.constants.INLINE_LIGAND_CAP` ligands
(blocking served path). For larger batches, call :meth:`start`, then
:meth:`wait` or :meth:`watch`, then :meth:`get_results`. Project-wide runs
use :class:`Admet` with ``ligands=[]`` and ``client.project_id`` set; call
Expand All @@ -285,6 +227,8 @@ class Admet(
ligands: Ligands whose SMILES are sent to the tool (empty for project runs).
properties: Endpoint names for this run.
method: Inference path — ``togo`` (default) or ``maplight``.
batch_size: Ligands per workflow pod on file and project runs, or
``None`` for the tool default.
"""

tool_key: str = TOOL_KEYS_AND_VERSIONS["admet"]["tool_key"]
Expand All @@ -296,14 +240,26 @@ def __init__(
*,
ligands: list[Ligand] | LigandSet,
method: Literal["maplight", "togo"] = "togo",
batch_size: int | None = None,
client: DeepOriginClient | None = None,
) -> None:
"""Configure an ADMET prediction run.

Fetches the live tool definition and fills :attr:`properties` with its
endpoint enum. Pass ``ligands=[]`` for a project-wide workflow run
(requires ``client.project_id``).

``batch_size`` sets ligands per workflow pod (``batchSize``) on file
and project runs; a smaller value fans out to more pods. It is ignored
on inline runs, which never fan out. ``None`` uses the tool default.

Raises:
ValueError: If ``batch_size`` is below the tool minimum (50).
"""
if batch_size is not None and batch_size < ADMET_MIN_BATCH_SIZE:
raise ValueError(
f"batch_size must be at least {ADMET_MIN_BATCH_SIZE} (got {batch_size})."
)
super().__init__(client=client)
if isinstance(ligands, LigandSet):
self._ligands: list[Ligand] = list(ligands.ligands)
Expand All @@ -313,6 +269,7 @@ def __init__(
self._allowed_endpoints: frozenset[str] | None = frozenset(endpoints)
self._properties: list[str] | tuple[str, ...] | None = list(endpoints)
self._method = method
self._batch_size = batch_size
self._remote_ligands_file: str | None = None

@property
Expand All @@ -325,6 +282,11 @@ def method(self) -> str:
"""Selected admet-now inference method."""
return self._method

@property
def batch_size(self) -> int | None:
"""Ligands per workflow pod, or ``None`` for the tool default."""
return self._batch_size

def _is_project_run(self) -> bool:
"""True when this instance targets all ligands in ``client.project_id``."""

Expand Down Expand Up @@ -403,30 +365,15 @@ def _ensure_platform_inputs(self) -> None:
LigandSet(ligands=self._ligands).sync(lazy=True, client=self.client)

def _ensure_ligands_file_uploaded(self) -> str:
"""Upload ligands JSON to UFA and return the remote path."""
"""Upload ligands JSON to UFA once and return the remote path."""

if self._remote_ligands_file is not None:
return self._remote_ligands_file
if self.client.files is None:
raise ValueError(
"Cannot upload Ligand list file: client.files is not available."
if self._remote_ligands_file is None:
self._remote_ligands_file = upload_ligand_list(
_ligand_payloads(self._ligands),
client=self.client,
prefix=ADMET_LIGAND_LIST_UPLOAD_PREFIX,
)

payloads = _ligand_payloads(self._ligands)
remote_path = f"{ADMET_LIGAND_LIST_UPLOAD_PREFIX}{uuid.uuid4().hex}.json"
fd, tmp_name = tempfile.mkstemp(suffix=".json", prefix="admet-ligands-")
try:
with os.fdopen(fd, "w", encoding="utf-8") as handle:
json.dump(payloads, handle, allow_nan=False)
self.client.files.upload(tmp_name, remote_path)
finally:
try:
os.unlink(tmp_name)
except OSError:
pass

self._remote_ligands_file = remote_path
return remote_path
return self._remote_ligands_file

def _make_inputs(self) -> dict[str, Any]:
"""Build tool ``inputs`` matching the admet-properties schema."""
Expand All @@ -436,14 +383,17 @@ def _make_inputs(self) -> dict[str, Any]:
inputs["method"] = self._method
if self._properties is not None:
inputs["properties"] = list(self._properties)
workflow_run = self._is_project_run() or len(self._ligands) > INLINE_LIGAND_CAP
if workflow_run and self._batch_size is not None:
inputs["batchSize"] = self._batch_size

if self._is_project_run():
project_id = require_client_project_id(self.client)
inputs["project"] = {"id": project_id}
return inputs

n = len(self._ligands)
if n > ADMET_INLINE_LIGAND_CAP:
if n > INLINE_LIGAND_CAP:
remote = self._ensure_ligands_file_uploaded()
inputs["ligands_file"] = remote
inputs["ligands_count"] = n
Expand Down Expand Up @@ -499,7 +449,7 @@ def _ensure_run_ligand_count(self) -> None:
n = len(self._ligands)
if n >= ADMET_WORKFLOW_LIGAND_THRESHOLD:
raise ValueError(
f"run() supports at most {ADMET_INLINE_LIGAND_CAP} ligands "
f"run() supports at most {INLINE_LIGAND_CAP} ligands "
f"(got {n}). Use start() then wait() or watch()."
)

Expand Down Expand Up @@ -660,6 +610,8 @@ def from_dto(
instance._allowed_endpoints = None
method = inputs.get("method")
instance._method = method if method in ("maplight", "togo") else "togo"
raw_batch = inputs.get("batchSize")
instance._batch_size = raw_batch if isinstance(raw_batch, int) else None
remote = inputs.get("ligands_file")
if isinstance(remote, str) and remote.strip():
instance._remote_ligands_file = remote.strip()
Expand Down
28 changes: 20 additions & 8 deletions src/drug_discovery/docking.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,10 @@
AsyncExecutableMixin,
SyncExecutableMixin,
)
from deeporigin.drug_discovery.ligand_list_file import (
ligand_rows_from_inputs,
ligands_input,
)
from deeporigin.drug_discovery.notebook_watch_mixin import NotebookWatchMixin
from deeporigin.drug_discovery.structures.ligand import Ligand, LigandSet
from deeporigin.drug_discovery.structures.pocket import Pocket
Expand Down Expand Up @@ -398,8 +402,10 @@ def _build_tool_inputs(
) -> tuple[dict, dict]:
"""Build params and metadata for ``client.executions.create``.

Does not sync or upload; call :meth:`_ensure_platform_inputs` first when
inputs may not yet exist on the platform.
Does not sync; call :meth:`_ensure_platform_inputs` first when inputs may
not yet exist on the platform. Above
:data:`~deeporigin.utils.constants.INLINE_LIGAND_CAP` ligands, uploads a
Ligand list file.

Args:
ligand_set: Ligands to include in tool ``inputs`` (default: all
Expand Down Expand Up @@ -427,7 +433,11 @@ def _build_tool_inputs(
"id": self.protein.id,
"file_path": self.protein.remote_path,
},
"ligands": [_ligand_tool_input_row(lig) for lig in ligands],
**ligands_input(
[_ligand_tool_input_row(lig) for lig in ligands],
client=self.client,
prefix="docking/ligand-lists/",
),
}

return params, metadata
Expand Down Expand Up @@ -513,10 +523,12 @@ def from_dto(
"this execution may have been created with an older input schema."
)

ligands_input = inputs.get("ligands", [])
if not ligands_input:
ligand_rows = ligand_rows_from_inputs(
inputs, client=instance.client, label="Docking"
)
if not ligand_rows:
raise ValueError(
"Missing 'ligands' in execution userInputs; "
"Missing 'ligands' or 'ligands_file' in execution userInputs; "
"this execution may have been created with an older input schema."
)

Expand All @@ -530,10 +542,10 @@ def from_dto(
)
fut_ligands = executor.submit(
LigandSet.from_ids,
[lig["id"] for lig in ligands_input],
[lig["id"] for lig in ligand_rows],
client=instance.client,
download=False,
ligand_inputs=ligands_input,
ligand_inputs=ligand_rows,
)
if pocket_id is not None:
fut_pocket = executor.submit(
Expand Down
Loading
Loading