Add molwt-filter workload with default scaffold hooks

This commit is contained in:
Emil
2026-08-02 01:09:13 +03:00
parent bc76f386e5
commit 5c5a2af0a1
13 changed files with 661 additions and 37 deletions
+1 -21
View File
@@ -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,
+6 -1
View File
@@ -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",
]
+89
View File
@@ -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()