Add SDK-declared workload UI elements and reduction metadata

This commit is contained in:
Emil
2026-08-02 17:47:43 +03:00
parent a18b8b8ae4
commit 5d738e0a14
10 changed files with 397 additions and 18 deletions
+13 -1
View File
@@ -9,13 +9,15 @@ partition, compute, and merge.
from __future__ import annotations
from collections.abc import Mapping
from pathlib import Path
from typing import Any, Mapping
from typing import Any
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 scimesh.sdk.ui import UIElement
from ..environment import current_environment_digest, current_scimesh_package_digest
from .core import (
@@ -88,6 +90,16 @@ class DescriptorBatchWorkload(MapReduceWorkload):
reduce_parameter_names = ("skip_invalid",)
map_entry_point = MAP_ENTRY_POINT
reduce_entry_point = REDUCE_ENTRY_POINT
ui_elements = (
UIElement(
"skip_invalid",
"checkbox",
"Skip invalid molecules",
help="Skip rows with invalid SMILES instead of failing the shard.",
default=True,
order=1,
),
)
def __init__(
self,
+30 -1
View File
@@ -9,8 +9,9 @@ byte-identical to the local brute-force reference.
from __future__ import annotations
from collections.abc import Mapping, Sequence
from pathlib import Path
from typing import Any, Mapping, Sequence
from typing import Any
from scimesh.sdk.artifacts import (
ArtifactCollection,
@@ -23,6 +24,7 @@ from scimesh.sdk.identity import SchemaRef, WorkloadId
from scimesh.sdk.plans import TaskSpec, ValidatedJob
from scimesh.sdk.protocols import PlanningContext
from scimesh.sdk.registry import WorkloadDefinition
from scimesh.sdk.ui import UIElement
from scimesh.sdk.workflow import StageSpec
from ..environment import current_environment_digest, current_scimesh_package_digest
@@ -109,8 +111,35 @@ class SimilarityGraphSDKWorkload(MapReduceWorkload):
"max_rows",
)
workflow_id = "graph-block-pairs-v1"
upload_ready = False
map_entry_point = MAP_ENTRY_POINT
reduce_entry_point = REDUCE_ENTRY_POINT
ui_elements = (
UIElement(
"threshold",
"number",
"Similarity threshold",
help="Minimum (greater) or maximum (less) edge similarity. Required.",
order=1,
),
UIElement(
"threshold_direction",
"select",
"Direction",
help="Whether to keep edges above (greater) or below (less) the threshold.",
options=("greater", "less"),
default="greater",
order=2,
),
UIElement(
"block_size",
"number",
"Block size",
help="Deterministic block size for pair sharding.",
default=100,
order=3,
),
)
def domain_validate(self, parameters: Mapping[str, Any]) -> None:
unknown = set(parameters) - {
+29 -1
View File
@@ -8,13 +8,15 @@ header-preserving concatenation, so nothing else is needed.
from __future__ import annotations
from collections.abc import Mapping
from pathlib import Path
from typing import Any, Mapping
from typing import Any
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 scimesh.sdk.ui import UIElement
from ..environment import current_environment_digest, current_scimesh_package_digest
from .core import MOLWT_COLUMNS, filter_molecules_by_molwt
@@ -93,6 +95,32 @@ class MolwtFilterWorkload(MapReduceWorkload):
reduce_parameter_names = ("min_molwt", "max_molwt", "skip_invalid")
map_entry_point = MAP_ENTRY_POINT
reduce_entry_point = REDUCE_ENTRY_POINT
ui_elements = (
UIElement(
"min_molwt",
"number",
"Minimum molecular weight",
help="Keep molecules with MolWt at least this value. Optional.",
placeholder="e.g. 100",
order=1,
),
UIElement(
"max_molwt",
"number",
"Maximum molecular weight",
help="Keep molecules with MolWt at most this value. Optional.",
placeholder="e.g. 600",
order=2,
),
UIElement(
"skip_invalid",
"checkbox",
"Skip invalid molecules",
help="Skip rows with invalid SMILES instead of failing the shard.",
default=True,
order=3,
),
)
def __init__(
self,
+45 -1
View File
@@ -9,8 +9,9 @@ and merge.
from __future__ import annotations
from collections.abc import Mapping, Sequence
from pathlib import Path
from typing import Any, Mapping, Sequence
from typing import Any
from rdkit import Chem
@@ -21,6 +22,7 @@ from scimesh.sdk.batch import MapReduceWorkload
from scimesh.sdk.identity import SchemaRef, WorkloadId
from scimesh.sdk.plans import JobRequest, ValidatedJob
from scimesh.sdk.registry import WorkloadDefinition
from scimesh.sdk.ui import UIElement
from ..environment import current_environment_digest, current_scimesh_package_digest
from .core import merge_search_partials, run_search_shard, write_search_shards
@@ -112,6 +114,48 @@ class SimilaritySearchSDKWorkload(MapReduceWorkload):
reduce_parameter_names = _MAP_PARAMETERS + ("query_source", "fingerprint")
map_entry_point = MAP_ENTRY_POINT
reduce_entry_point = REDUCE_ENTRY_POINT
reduction = "top-k"
ui_elements = (
UIElement(
"query_id",
"text",
"Query molecule id",
help="ChEMBL id of the query molecule. Provide exactly one of id or SMILES.",
order=1,
),
UIElement(
"query_smiles",
"text",
"Query molecule SMILES",
help="SMILES of the query molecule. Provide exactly one of id or SMILES.",
order=2,
),
UIElement(
"top_k",
"number",
"Top k",
help="Number of most similar molecules to keep per shard (global merge keeps the best of these).",
default=20,
order=3,
),
UIElement(
"threshold_direction",
"select",
"Direction",
help="Keep molecules with similarity greater or less than the threshold.",
options=("greater", "less"),
default="greater",
order=4,
),
UIElement(
"threshold",
"number",
"Similarity threshold",
help="Optional similarity bound: results are filtered to this direction.",
placeholder="e.g. 0.8",
order=5,
),
)
def __init__(
self,
+44 -1
View File
@@ -50,6 +50,12 @@ class WorkloadCLI:
)
export_parser.set_defaults(workload_handler=self.export_workloads)
allowlist_parser = subparsers.add_parser(
"allowlist",
help="Print the installed-package workload allowlist for worker configuration.",
)
allowlist_parser.set_defaults(workload_handler=self.export_allowlist)
run_parser = subparsers.add_parser(
"run", help="Run one SDK workload locally against an input file."
)
@@ -147,6 +153,11 @@ class WorkloadCLI:
"verifier": manifest.verifier.verifier.canonical,
"enabled": item.enabled,
"parameters_schema": thaw_json(manifest.parameters_schema),
"ui_elements": [
element.to_dict() for element in manifest.ui_elements
],
"reduction": manifest.reduction,
"upload_ready": manifest.upload_ready,
"inputs": {
name: port.schema.to_dict()
for name, port in manifest.inputs.items()
@@ -158,7 +169,7 @@ class WorkloadCLI:
}
)
payload: dict[str, object] = {
"schema_version": 1,
"schema_version": 2,
"generated_by": "scimesh workload export",
"workloads": workloads,
}
@@ -169,6 +180,38 @@ class WorkloadCLI:
print(f"Exported {len(workloads)} workloads to {args.output}")
return 0
def export_allowlist(self, args: argparse.Namespace) -> int:
"""Print the allowlist JSON that worker environments consume.
The printed array feeds ``SCIMESH_WORKLOAD_ALLOWLIST`` on workers and
mirrors the digest pins of the installed distribution.
"""
import json
registry = self._registry(args)
payload = []
for item in sorted(
registry.descriptions(), key=lambda value: value.workload.name
):
if not item.enabled:
continue
definition, _ = registry.require(
item.workload.name,
item.workload.version,
item.package_digest,
)
manifest = definition.manifest
payload.append(
{
"distribution": manifest.package.distribution,
"name": manifest.workload.name,
"version": manifest.workload.version,
"digest": manifest.package.digest,
}
)
print(json.dumps(payload, indent=2, sort_keys=True))
return 0
def run_workload(self, args: argparse.Namespace) -> int:
registry = self._registry(args)
descriptions = registry.descriptions()