Add SDK-declared workload UI elements and reduction metadata
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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) - {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user