Add molwt-filter workload with default scaffold hooks
This commit is contained in:
@@ -9,8 +9,7 @@ partition, compute, and merge.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any, Mapping, Sequence
|
||||
from typing import Any, Mapping
|
||||
|
||||
from scimesh.sdk.artifacts import ArtifactSchema, ComponentRef, PortSpec
|
||||
from scimesh.sdk.batch import MapReduceWorkload
|
||||
@@ -21,9 +20,7 @@ from ..environment import current_environment_digest, current_scimesh_package_di
|
||||
from .core import (
|
||||
DESCRIPTOR_COLUMNS,
|
||||
compute_descriptor_batch,
|
||||
concatenate_descriptor_shards,
|
||||
validate_descriptor_names,
|
||||
write_descriptor_shards,
|
||||
)
|
||||
|
||||
MAP_ENTRY_POINT = "scimesh.workloads.descriptors.definition:map_descriptors@v1"
|
||||
@@ -116,14 +113,6 @@ class DescriptorBatchWorkload(MapReduceWorkload):
|
||||
if not isinstance(value, bool):
|
||||
raise ValueError("skip_invalid must be a boolean")
|
||||
|
||||
def partition_input(
|
||||
self,
|
||||
input_path: Path,
|
||||
parameters: Mapping[str, Any],
|
||||
workspace: Path,
|
||||
) -> list[Path]:
|
||||
return write_descriptor_shards(input_path, workspace, self.shard_rows)
|
||||
|
||||
def compute_shard(
|
||||
self,
|
||||
inputs: Mapping[str, Path],
|
||||
@@ -143,15 +132,6 @@ class DescriptorBatchWorkload(MapReduceWorkload):
|
||||
raise ValueError("skip_invalid must be a boolean")
|
||||
return value
|
||||
|
||||
def reduce_partials(
|
||||
self,
|
||||
partial_paths: Sequence[Path],
|
||||
parameters: Mapping[str, Any],
|
||||
output_path: Path,
|
||||
) -> Mapping[str, int | float]:
|
||||
return concatenate_descriptor_shards(partial_paths, output_path)
|
||||
|
||||
|
||||
def descriptor_batch_sdk_definition(
|
||||
*,
|
||||
shard_rows: int = 10_000,
|
||||
|
||||
@@ -19,6 +19,7 @@ from scimesh.sdk.runtime import RuntimeCapabilities
|
||||
from .descriptors import descriptor_batch_sdk_definition
|
||||
from .environment import current_environment_digest
|
||||
from .graph import similarity_graph_sdk_definition
|
||||
from .molwt_filter import molwt_filter_sdk_definition
|
||||
from .search import similarity_search_sdk_definition
|
||||
|
||||
__all__ = [
|
||||
@@ -53,6 +54,10 @@ def default_sdk_registry(
|
||||
descriptor_batch_sdk_definition(shard_rows=shard_rows).definition(),
|
||||
enabled=True,
|
||||
)
|
||||
registry.register(
|
||||
molwt_filter_sdk_definition(shard_rows=shard_rows).definition(),
|
||||
enabled=True,
|
||||
)
|
||||
return registry
|
||||
|
||||
|
||||
@@ -75,7 +80,7 @@ def default_sdk_runtime(
|
||||
features={"artifact-collections": "1.0.0", "exact-verifier": "1.0.0"},
|
||||
workload_capabilities=(
|
||||
workload_capabilities
|
||||
or ("similarity-search", "similarity-graph", "descriptor-batch")
|
||||
or ("similarity-search", "similarity-graph", "descriptor-batch", "molwt-filter")
|
||||
),
|
||||
inventory=ResourceInventory(
|
||||
cpu_cores=max(os.cpu_count() or 1, 1),
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
"""SDK-built ``molwt-filter`` workload.
|
||||
|
||||
A minimal authoring example built on ``MapReduceWorkload``: the scaffold's
|
||||
default sharding and concatenation hooks are used unchanged, so the workload
|
||||
only declares identity, parameters, ports, and the single scientific hook.
|
||||
"""
|
||||
|
||||
from .core import MOLWT_COLUMNS, filter_molecules_by_molwt
|
||||
from .definition import (
|
||||
MAP_ENTRY_POINT,
|
||||
REDUCE_ENTRY_POINT,
|
||||
MolwtFilterWorkload,
|
||||
molwt_filter_sdk_definition,
|
||||
workload_definition,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"MAP_ENTRY_POINT",
|
||||
"MOLWT_COLUMNS",
|
||||
"REDUCE_ENTRY_POINT",
|
||||
"MolwtFilterWorkload",
|
||||
"filter_molecules_by_molwt",
|
||||
"molwt_filter_sdk_definition",
|
||||
"workload_definition",
|
||||
]
|
||||
@@ -0,0 +1,89 @@
|
||||
"""Scientific core for the ``molwt-filter`` workload.
|
||||
|
||||
Keeps one canonical row per input molecule whose exact RDKit molecular weight
|
||||
falls inside the requested bounds; rows keep input order, invalid SMILES are
|
||||
skipped or fail the run, and molecular weights are serialized with fixed
|
||||
``%.6f`` formatting so the output is byte-identical across workers.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
from pathlib import Path
|
||||
from typing import Mapping
|
||||
|
||||
from rdkit import Chem
|
||||
from rdkit.Chem import Descriptors
|
||||
|
||||
from scimesh.chemistry.dataset import iter_rows
|
||||
|
||||
MOLWT_COLUMNS = ("chembl_id", "canonical_smiles", "molwt")
|
||||
|
||||
|
||||
def _bound(value: object, name: str) -> float:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
raise ValueError(f"{name} must be a number")
|
||||
return float(value)
|
||||
|
||||
|
||||
def filter_molecules_by_molwt(
|
||||
input_path: Path,
|
||||
output_path: Path,
|
||||
*,
|
||||
min_molwt: object,
|
||||
max_molwt: object,
|
||||
skip_invalid: bool,
|
||||
) -> dict[str, int]:
|
||||
"""Write one CSV row per molecule whose MolWt is within [min, max]."""
|
||||
minimum = _bound(min_molwt, "min_molwt") if min_molwt is not None else None
|
||||
maximum = _bound(max_molwt, "max_molwt") if max_molwt is not None else None
|
||||
if minimum is None and maximum is None:
|
||||
raise ValueError("at least one of min_molwt or max_molwt is required")
|
||||
if minimum is not None and minimum < 0:
|
||||
raise ValueError("min_molwt must be non-negative")
|
||||
if maximum is not None and maximum < 0:
|
||||
raise ValueError("max_molwt must be non-negative")
|
||||
if minimum is not None and maximum is not None and minimum > maximum:
|
||||
raise ValueError("min_molwt must not exceed max_molwt")
|
||||
if not isinstance(skip_invalid, bool):
|
||||
raise ValueError("skip_invalid must be a boolean")
|
||||
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
scanned = 0
|
||||
invalid = 0
|
||||
emitted = 0
|
||||
with output_path.open("w", encoding="utf-8", newline="") as destination:
|
||||
writer = csv.DictWriter(
|
||||
destination,
|
||||
fieldnames=list(MOLWT_COLUMNS),
|
||||
lineterminator="\n",
|
||||
)
|
||||
writer.writeheader()
|
||||
for row in iter_rows(input_path):
|
||||
scanned += 1
|
||||
smiles = row.get("canonical_smiles", "")
|
||||
molecule = Chem.MolFromSmiles(smiles)
|
||||
if molecule is None:
|
||||
invalid += 1
|
||||
if not skip_invalid:
|
||||
raise ValueError(f"row {scanned} has an invalid canonical_smiles")
|
||||
continue
|
||||
molwt = Descriptors.MolWt(molecule)
|
||||
if minimum is not None and molwt < minimum:
|
||||
continue
|
||||
if maximum is not None and molwt > maximum:
|
||||
continue
|
||||
canonical = Chem.MolToSmiles(molecule, canonical=True)
|
||||
writer.writerow(
|
||||
{
|
||||
"chembl_id": row.get("chembl_id", ""),
|
||||
"canonical_smiles": canonical,
|
||||
"molwt": f"{molwt:.6f}",
|
||||
}
|
||||
)
|
||||
emitted += 1
|
||||
return {
|
||||
"rows_scanned": scanned,
|
||||
"invalid_rows": invalid,
|
||||
"rows_emitted": emitted,
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
"""SDK-built ``molwt-filter`` workload definition.
|
||||
|
||||
The minimal authoring example: a subclass of ``MapReduceWorkload`` that only
|
||||
declares identity, parameters, ports, and one scientific hook. The default
|
||||
hooks of the scaffold provide deterministic row-bounded sharding and
|
||||
header-preserving concatenation, so nothing else is needed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any, Mapping
|
||||
|
||||
from scimesh.sdk.artifacts import ArtifactSchema, ComponentRef, PortSpec
|
||||
from scimesh.sdk.batch import MapReduceWorkload
|
||||
from scimesh.sdk.identity import SchemaRef, WorkloadId
|
||||
from scimesh.sdk.registry import WorkloadDefinition
|
||||
|
||||
from ..environment import current_environment_digest, current_scimesh_package_digest
|
||||
from .core import MOLWT_COLUMNS, filter_molecules_by_molwt
|
||||
|
||||
MAP_ENTRY_POINT = "scimesh.workloads.molwt_filter.definition:map_molwt_filter@v1"
|
||||
REDUCE_ENTRY_POINT = "scimesh.workloads.molwt_filter.definition:reduce_molwt_filter@v1"
|
||||
|
||||
|
||||
def _parameters_schema() -> dict[str, Any]:
|
||||
return {
|
||||
"type": "object",
|
||||
"additionalProperties": False,
|
||||
"properties": {
|
||||
"min_molwt": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Keep molecules with MolWt >= this value",
|
||||
},
|
||||
"max_molwt": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Keep molecules with MolWt <= this value",
|
||||
},
|
||||
"skip_invalid": {
|
||||
"type": "boolean",
|
||||
"default": True,
|
||||
"description": "Skip rows with invalid SMILES instead of failing",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _molecule_schema() -> ArtifactSchema:
|
||||
return ArtifactSchema(
|
||||
SchemaRef("molecule-table", 1),
|
||||
"text/tab-separated-values",
|
||||
"utf-8",
|
||||
max_bytes=10 * 1024 * 1024 * 1024,
|
||||
validator=ComponentRef("delimited-table", 1),
|
||||
validator_configuration={
|
||||
"required_columns": ["canonical_smiles", "chembl_id"],
|
||||
},
|
||||
max_records=100_000_000,
|
||||
canonicalizer="scimesh-tsv-v1",
|
||||
)
|
||||
|
||||
|
||||
def _filtered_schema() -> ArtifactSchema:
|
||||
return ArtifactSchema(
|
||||
SchemaRef("molwt-filtered-table", 1),
|
||||
"text/csv",
|
||||
"utf-8",
|
||||
max_bytes=100 * 1024 * 1024 * 1024,
|
||||
validator=ComponentRef("delimited-table", 1),
|
||||
validator_configuration={
|
||||
"columns": list(MOLWT_COLUMNS),
|
||||
},
|
||||
max_records=100_000_000,
|
||||
canonicalizer="molwt-filtered-table-v1",
|
||||
)
|
||||
|
||||
|
||||
class MolwtFilterWorkload(MapReduceWorkload):
|
||||
"""Keep molecules whose exact RDKit molecular weight is within bounds."""
|
||||
|
||||
workload_id = WorkloadId("molwt-filter", "1.0.0")
|
||||
description = (
|
||||
"Filter molecules by exact RDKit molecular weight, one canonical "
|
||||
"CSV row per kept input molecule, in deterministic input order."
|
||||
)
|
||||
parameters_schema = _parameters_schema()
|
||||
input_port = PortSpec(_molecule_schema())
|
||||
partial_port = PortSpec(_filtered_schema())
|
||||
output_port = PortSpec(_filtered_schema())
|
||||
map_parameter_names = ("min_molwt", "max_molwt", "skip_invalid")
|
||||
reduce_parameter_names = ("min_molwt", "max_molwt", "skip_invalid")
|
||||
map_entry_point = MAP_ENTRY_POINT
|
||||
reduce_entry_point = REDUCE_ENTRY_POINT
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
shard_rows: int = 10_000,
|
||||
package_digest: str,
|
||||
environment_digest: str,
|
||||
) -> None:
|
||||
if (
|
||||
isinstance(shard_rows, bool)
|
||||
or not isinstance(shard_rows, int)
|
||||
or shard_rows < 1
|
||||
):
|
||||
raise ValueError("shard_rows must be a positive integer")
|
||||
self.shard_rows = shard_rows
|
||||
super().__init__(
|
||||
package_digest=package_digest,
|
||||
environment_digest=environment_digest,
|
||||
)
|
||||
|
||||
def domain_validate(self, parameters: Mapping[str, Any]) -> None:
|
||||
unknown = set(parameters) - {"min_molwt", "max_molwt", "skip_invalid"}
|
||||
if unknown:
|
||||
raise ValueError(
|
||||
"unsupported molwt-filter parameters: " + ", ".join(sorted(unknown))
|
||||
)
|
||||
minimum = parameters.get("min_molwt")
|
||||
maximum = parameters.get("max_molwt")
|
||||
if minimum is None and maximum is None:
|
||||
raise ValueError("at least one of min_molwt or max_molwt is required")
|
||||
for value, name in ((minimum, "min_molwt"), (maximum, "max_molwt")):
|
||||
if value is not None and (
|
||||
isinstance(value, bool) or not isinstance(value, (int, float))
|
||||
):
|
||||
raise ValueError(f"{name} must be a number")
|
||||
if (
|
||||
minimum is not None
|
||||
and maximum is not None
|
||||
and float(minimum) > float(maximum)
|
||||
):
|
||||
raise ValueError("min_molwt must not exceed max_molwt")
|
||||
skip_invalid = parameters.get("skip_invalid", True)
|
||||
if not isinstance(skip_invalid, bool):
|
||||
raise ValueError("skip_invalid must be a boolean")
|
||||
|
||||
def compute_shard(
|
||||
self,
|
||||
inputs: Mapping[str, Path],
|
||||
parameters: Mapping[str, Any],
|
||||
output_path: Path,
|
||||
) -> Mapping[str, int | float]:
|
||||
return filter_molecules_by_molwt(
|
||||
inputs["input"],
|
||||
output_path,
|
||||
min_molwt=parameters.get("min_molwt"),
|
||||
max_molwt=parameters.get("max_molwt"),
|
||||
skip_invalid=parameters.get("skip_invalid", True),
|
||||
)
|
||||
|
||||
|
||||
def molwt_filter_sdk_definition(
|
||||
*,
|
||||
shard_rows: int = 10_000,
|
||||
package_digest: str | None = None,
|
||||
environment_digest: str | None = None,
|
||||
) -> MolwtFilterWorkload:
|
||||
"""Build the molwt-filter definition for tests."""
|
||||
return MolwtFilterWorkload(
|
||||
shard_rows=shard_rows,
|
||||
package_digest=package_digest or current_scimesh_package_digest(),
|
||||
environment_digest=environment_digest or current_environment_digest(),
|
||||
)
|
||||
|
||||
|
||||
def workload_definition() -> WorkloadDefinition:
|
||||
"""Installed entry-point factory for molwt-filter."""
|
||||
return molwt_filter_sdk_definition().definition()
|
||||
Reference in New Issue
Block a user