Add MapReduceWorkload scaffold and generic workload execution

This commit is contained in:
Emil
2026-08-02 01:02:24 +03:00
parent 19fbb8e926
commit bc76f386e5
21 changed files with 2241 additions and 1176 deletions
+2
View File
@@ -6,6 +6,7 @@ from scimesh.core.registry import WorkloadRegistry
from scimesh.workloads.help import HelpWorkload
from scimesh.workloads.similarity_graph import SimilarityGraphWorkload
from scimesh.workloads.similarity_search import SimilaritySearchWorkload
from scimesh.workloads.workload_cli import WorkloadCLI
def register_workloads(registry: WorkloadRegistry) -> None:
@@ -13,3 +14,4 @@ def register_workloads(registry: WorkloadRegistry) -> None:
registry.register(HelpWorkload())
registry.register(SimilaritySearchWorkload())
registry.register(SimilarityGraphWorkload())
registry.register(WorkloadCLI())
+61 -333
View File
@@ -1,53 +1,23 @@
"""SDK-built ``descriptor-batch`` workload definition and handlers.
This module is the first non-adapter reference workload built directly on the
``core-batch-v1`` profile: an explicit immutable manifest, a static map/reduce
workflow, a row-bounded planner, pinned descriptor computation, deterministic
shard concatenation, and the exact-artifact verifier. It is the intended
first ``untrusted_quorum`` candidate: ``byte_exact`` determinism with whole
file SHA-256 agreement from distinct owners.
A thin subclass of ``MapReduceWorkload``: the SDK assembles the manifest, the
map/reduce stages, the workflow, and the digest-pinned handlers; this module
only declares the scientific contract (pinned descriptors, canonical CSV,
row-bounded shards, header-preserving concatenation) and the three hooks that
partition, compute, and merge.
"""
from __future__ import annotations
import hashlib
import shutil
from pathlib import Path
from typing import Any, Mapping, Sequence
from ...sdk.artifacts import (
ArtifactCollection,
ArtifactItem,
ArtifactRef,
ArtifactSchema,
Cardinality,
CollectionKind,
OutputManifest,
PortSpec,
)
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 ...sdk.execution import (
CheckpointPolicy,
ExecutionProfile,
NetworkPolicy,
RetryPolicy,
)
from ...sdk.identity import ComponentRef, SchemaRef, VersionRange, WorkloadId
from ...sdk.manifest import (
DeterminismProfile,
EnvironmentSpec,
PackageSpec,
TrustMode,
VerifierSpec,
WorkloadLimits,
WorkloadManifest,
)
from ...sdk.plans import JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
from ...sdk.protocols import PlanningContext, ReduceContext, TaskContext
from ...sdk.registry import WorkloadDefinition
from ...sdk.resources import ResourceRequirements
from ...sdk.verification import ExactArtifactVerifier
from ...sdk.workflow import ArtifactEdge, PortRef, StageKind, StageSpec, WorkflowSpec
from .core import (
DESCRIPTOR_COLUMNS,
compute_descriptor_batch,
@@ -59,16 +29,6 @@ from .core import (
MAP_ENTRY_POINT = "scimesh.workloads.descriptors.definition:map_descriptors@v1"
REDUCE_ENTRY_POINT = "scimesh.workloads.descriptors.definition:reduce_descriptors@v1"
_DESCRIPTOR_PARAMETERS = ("skip_invalid",)
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as source:
for block in iter(lambda: source.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def _parameters_schema() -> dict[str, Any]:
return {
@@ -114,14 +74,22 @@ def _descriptor_schema() -> ArtifactSchema:
)
class DescriptorBatchWorkload:
"""Manifest-backed planner, runner, and reducer for descriptor-batch.
class DescriptorBatchWorkload(MapReduceWorkload):
"""Pinned RDKit 2D descriptor computation, one canonical row per input."""
The class follows the legacy adapter's structural pattern (one object
registered under each stage entry point) while remaining fully SDK-built:
sharding is explicit and deterministic, every artifact is sealed through
the bridge-owned sink, and no filesystem path ever enters a plan or task.
"""
workload_id = WorkloadId("descriptor-batch", "1.0.0")
description = (
"Compute a pinned set of RDKit 2D descriptors, one canonical "
"CSV row per input molecule, in deterministic input order."
)
parameters_schema = _parameters_schema()
input_port = PortSpec(_input_schema())
partial_port = PortSpec(_descriptor_schema())
output_port = PortSpec(_descriptor_schema())
map_parameter_names = ("skip_invalid",)
reduce_parameter_names = ("skip_invalid",)
map_entry_point = MAP_ENTRY_POINT
reduce_entry_point = REDUCE_ENTRY_POINT
def __init__(
self,
@@ -137,291 +105,51 @@ class DescriptorBatchWorkload:
):
raise ValueError("shard_rows must be a positive integer")
validate_descriptor_names()
self.entry_point = MAP_ENTRY_POINT
self.shard_rows = shard_rows
self.input_port = PortSpec(_input_schema())
self.partial_port = PortSpec(_descriptor_schema())
self.output_port = PortSpec(_descriptor_schema())
resources = ResourceRequirements(
profile="descriptor-cpu-v1",
cpu_cores=1,
memory_mb=1024,
scratch_mb=1024,
max_duration_seconds=3600,
super().__init__(
package_digest=package_digest,
environment_digest=environment_digest,
)
execution = ExecutionProfile(
profile="descriptor-python-process-v1",
network=NetworkPolicy.TRUSTED,
timeout_seconds=3600,
checkpoint=CheckpointPolicy(),
)
limits = WorkloadLimits(
max_input_bytes=self.input_port.schema.max_bytes,
max_tasks=10_000,
max_output_bytes=self.output_port.schema.max_bytes,
)
trust_modes = ("trusted", "untrusted_quorum")
map_stage = StageSpec(
stage_id="map",
kind=StageKind.MAP,
entry_point=MAP_ENTRY_POINT,
needs=(),
inputs={"input": self.input_port},
outputs={"partial": self.partial_port},
parameter_names=_DESCRIPTOR_PARAMETERS,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=ComponentRef("exact-artifact", 1),
trust_modes=trust_modes,
max_fan_out=limits.max_tasks,
cacheable=True,
)
reduce_input = PortSpec(
schema=self.partial_port.schema,
cardinality=Cardinality.MANY,
collection=CollectionKind.KEYED,
)
reduce_stage = StageSpec(
stage_id="reduce",
kind=StageKind.REDUCE,
entry_point=REDUCE_ENTRY_POINT,
needs=("map",),
inputs={"partials": reduce_input},
outputs={"result": self.output_port},
parameter_names=_DESCRIPTOR_PARAMETERS,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=ComponentRef("exact-artifact", 1),
trust_modes=trust_modes,
max_fan_out=1,
cacheable=True,
)
workflow = WorkflowSpec(
workflow_id="descriptor-map-reduce-v1",
inputs={"input": self.input_port},
stages=(map_stage, reduce_stage),
edges=(
ArtifactEdge(PortRef("input"), PortRef("input", "map")),
ArtifactEdge(PortRef("partial", "map"), PortRef("partials", "reduce")),
),
outputs={"result": PortRef("result", "reduce")},
max_tasks=limits.max_tasks,
max_output_bytes=limits.max_output_bytes,
)
self.manifest = WorkloadManifest(
sdk_api=VersionRange(">=1.0,<2.0"),
protocol=VersionRange(">=1,<2"),
workload=WorkloadId("descriptor-batch", "1.0.0"),
description=(
"Compute a pinned set of RDKit 2D descriptors, one canonical "
"CSV row per input molecule, in deterministic input order."
),
package=PackageSpec("scimesh", package_digest),
environment=EnvironmentSpec(
"python-process",
environment_digest,
{"adapter": "sdk-native"},
),
parameters_schema=_parameters_schema(),
workflow=workflow,
inputs={"input": self.input_port},
outputs={"result": self.output_port},
determinism=DeterminismProfile.BYTE_EXACT,
trust_modes=(TrustMode.TRUSTED, TrustMode.UNTRUSTED_QUORUM),
verifier=VerifierSpec(ComponentRef("exact-artifact", 1), {}),
limits=limits,
capabilities=("descriptor-batch",),
conformance_profiles=("core-batch-v1",),
)
self._exact_verifier = ExactArtifactVerifier()
def definition(self) -> WorkloadDefinition:
return WorkloadDefinition(
manifest=self.manifest,
planner=self,
runners={MAP_ENTRY_POINT: self},
reducers={REDUCE_ENTRY_POINT: self},
verifiers={self._exact_verifier.identity.canonical: self._exact_verifier},
def domain_validate(self, parameters: Mapping[str, Any]) -> None:
value = parameters.get("skip_invalid", True)
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],
parameters: Mapping[str, Any],
output_path: Path,
) -> Mapping[str, int | float]:
return compute_descriptor_batch(
inputs["input"],
output_path,
skip_invalid=self.domain_validate_check(parameters),
)
@staticmethod
def _skip_invalid(parameters: Mapping[str, Any]) -> bool:
def domain_validate_check(parameters: Mapping[str, Any]) -> bool:
value = parameters.get("skip_invalid", True)
if not isinstance(value, bool):
raise ValueError("skip_invalid must be a boolean")
return value
def validate(self, request: JobRequest) -> ValidatedJob:
if request.workload != self.manifest.workload:
raise ValueError("descriptor-batch received a request for another workload")
self._skip_invalid(request.parameters)
return ValidatedJob(request, request.parameters)
def plan(self, job: ValidatedJob, context: PlanningContext) -> WorkflowPlan:
if not isinstance(job, ValidatedJob):
raise ValueError("job must be a ValidatedJob")
collection = job.request.inputs.get("input")
if collection is None:
raise ValueError("descriptor-batch requires the input port")
self.input_port.validate_collection(collection, "job input")
input_artifact = collection.items[0].artifact
input_path = context.catalog.materialize(input_artifact)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
shard_paths = write_descriptor_shards(
input_path,
workspace,
self.shard_rows,
)
negotiated = context.negotiated
map_stage = self.manifest.workflow.stages[0]
assert map_stage.verifier is not None
tasks: list[TaskSpec] = []
for index, path in enumerate(shard_paths):
sealed = context.sink.seal(
path,
declaration=self.input_port.schema,
)
tasks.append(
TaskSpec(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=map_stage.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
task_key=f"map/{index:08d}",
stage_id="map",
parameters=job.resolved_parameters,
inputs={"input": ArtifactCollection.single(sealed)},
expected_outputs={"partial": self.partial_port},
resources=map_stage.resources,
execution=map_stage.execution,
)
)
return WorkflowPlan(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=self.manifest.verifier.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
workflow_id=self.manifest.workflow.workflow_id,
resolved_parameters=job.resolved_parameters,
tasks=tuple(tasks),
)
def run(self, context: TaskContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
collection = context.task.inputs.get("input")
if collection is None:
raise ValueError("descriptor map task requires one input collection")
self.input_port.validate_collection(collection, "descriptor map input")
source = context.catalog.materialize(collection.items[0].artifact)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
input_path = workspace / "input"
output_path = workspace / "result.csv"
if source.resolve() != input_path.resolve():
shutil.copyfile(source, input_path)
metrics = compute_descriptor_batch(
input_path,
output_path,
skip_invalid=self._skip_invalid(context.task.parameters),
)
context.cancellation.raise_if_cancelled()
sealed = context.sink.seal(
output_path,
declaration=self.partial_port.schema,
)
return OutputManifest(
context.task.task_key,
{"partial": ArtifactCollection.single(sealed)},
metrics,
context.provenance,
).validate_against(
context.task.expected_outputs,
max_output_bytes=self.manifest.limits.max_output_bytes,
)
def reduce(self, context: ReduceContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
collection = context.accepted_inputs.get("partials")
if (
collection is None
or collection.kind is not CollectionKind.KEYED
or not collection.items
):
raise ValueError(
"descriptor reducer requires a non-empty keyed partial collection"
)
self.manifest.workflow.stages[1].inputs["partials"].validate_collection(
collection,
"descriptor reducer partials",
)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
indexed_items: list[tuple[int, ArtifactItem]] = []
for item in collection.items:
key = item.key or ""
prefix = "map."
raw_index = key[len(prefix) :] if key.startswith(prefix) else ""
if len(raw_index) != 8 or not raw_index.isdigit():
raise ValueError(
"descriptor partial key must use map.<eight-digit-index>"
)
indexed_items.append((int(raw_index), item))
expected_keys = context.task.expected_input_keys.get("partials")
if expected_keys is None or {item.key for item in collection.items} != set(
expected_keys
):
raise ValueError(
"descriptor partial keys do not match the coordinator expected set"
)
if sorted(index for index, _ in indexed_items) != list(
range(len(indexed_items))
):
raise ValueError("descriptor partial keys must be complete and contiguous")
partial_paths: list[Path] = []
for index, item in sorted(indexed_items):
artifact: ArtifactRef = item.artifact
source = context.catalog.materialize(artifact)
target = workspace / artifact.artifact_id
if source.resolve() != target.resolve():
shutil.copyfile(source, target)
if _sha256_file(target) != artifact.sha256:
raise ValueError("materialized partial checksum does not match")
partial_paths.append(target)
result_path = workspace / "result.csv"
metrics = concatenate_descriptor_shards(partial_paths, result_path)
context.cancellation.raise_if_cancelled()
sealed = context.sink.seal(
result_path,
declaration=self.output_port.schema,
)
return OutputManifest(
context.task.task_key,
{"result": ArtifactCollection.single(sealed)},
metrics,
context.provenance,
).validate_against(
context.task.expected_outputs,
max_output_bytes=self.manifest.limits.max_output_bytes,
)
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(
@@ -430,7 +158,7 @@ def descriptor_batch_sdk_definition(
package_digest: str | None = None,
environment_digest: str | None = None,
) -> DescriptorBatchWorkload:
"""Build the default local descriptor-batch definition for tests."""
"""Build the default descriptor-batch definition for tests."""
return DescriptorBatchWorkload(
shard_rows=shard_rows,
package_digest=package_digest or current_scimesh_package_digest(),
@@ -439,5 +167,5 @@ def descriptor_batch_sdk_definition(
def workload_definition() -> WorkloadDefinition:
"""Installed entry-point factory for the default descriptor-batch definition."""
"""Installed entry-point factory for the descriptor-batch workload."""
return descriptor_batch_sdk_definition().definition()
+115 -344
View File
@@ -1,52 +1,30 @@
"""SDK-built ``similarity-graph`` workload definition and handlers.
Built directly on the ``core-batch-v1`` profile: molecules are parsed once
into deterministic row-ordered blocks, every block pair ``(i, j)`` with
``i <= j`` becomes one map task, and the reducer enforces the CTX-10
pair-coverage invariant (every unordered molecule pair compared exactly once)
before emitting a deterministically sorted edge list that is byte-identical
to the local brute-force reference.
A ``MapReduceWorkload`` subclass with two non-default hooks: ``plan_tasks``
builds one task per block pair ``(i, j)`` with ``i <= j`` (each task receives
two block inputs), and the partial-key hooks parse ``map.<i>x<j>`` keys and
enforce the CTX-10 pair-coverage invariant. The final edge list is
byte-identical to the local brute-force reference.
"""
from __future__ import annotations
import hashlib
import shutil
from pathlib import Path
from typing import Any, Mapping
from typing import Any, Mapping, Sequence
from ...sdk.artifacts import (
from scimesh.sdk.artifacts import (
ArtifactCollection,
ArtifactItem,
ArtifactRef,
ArtifactSchema,
Cardinality,
CollectionKind,
OutputManifest,
ComponentRef,
PortSpec,
)
from ...sdk.execution import (
CheckpointPolicy,
ExecutionProfile,
NetworkPolicy,
RetryPolicy,
)
from ...sdk.identity import ComponentRef, SchemaRef, VersionRange, WorkloadId
from ...sdk.manifest import (
DeterminismProfile,
EnvironmentSpec,
PackageSpec,
TrustMode,
VerifierSpec,
WorkloadLimits,
WorkloadManifest,
)
from ...sdk.plans import JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
from ...sdk.protocols import PlanningContext, ReduceContext, TaskContext
from ...sdk.registry import WorkloadDefinition
from ...sdk.resources import ResourceRequirements
from ...sdk.verification import ExactArtifactVerifier
from ...sdk.workflow import ArtifactEdge, PortRef, StageKind, StageSpec, WorkflowSpec
from scimesh.sdk.batch import MapReduceWorkload
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.workflow import StageSpec
from ..environment import current_environment_digest, current_scimesh_package_digest
from .core import (
block_pair_from_key,
@@ -63,15 +41,6 @@ MAP_ENTRY_POINT = "scimesh.workloads.graph.definition:map_graph@v1"
REDUCE_ENTRY_POINT = "scimesh.workloads.graph.definition:reduce_graph@v1"
_MAP_PARAMETERS = ("left_block", "right_block", "threshold", "threshold_direction")
_REDUCE_PARAMETERS = ("threshold", "threshold_direction", "block_size", "max_rows")
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as source:
for block in iter(lambda: source.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def _parameters_schema() -> dict[str, Any]:
@@ -118,141 +87,32 @@ def _edge_schema() -> ArtifactSchema:
)
class SimilarityGraphSDKWorkload:
"""Manifest-backed planner, runner, and reducer for similarity-graph."""
class SimilarityGraphSDKWorkload(MapReduceWorkload):
"""Exact sparse Tanimoto graph over deterministic block pairs."""
def __init__(
self,
*,
package_digest: str,
environment_digest: str,
) -> None:
self.entry_point = MAP_ENTRY_POINT
self.input_port = PortSpec(_molecule_schema())
self.block_port = PortSpec(_molecule_schema())
self.partial_port = PortSpec(_edge_schema())
self.output_port = PortSpec(_edge_schema())
resources = ResourceRequirements(
profile="graph-cpu-v1",
cpu_cores=1,
memory_mb=1024,
scratch_mb=1024,
max_duration_seconds=3600,
)
execution = ExecutionProfile(
profile="graph-python-process-v1",
network=NetworkPolicy.TRUSTED,
timeout_seconds=3600,
checkpoint=CheckpointPolicy(),
)
limits = WorkloadLimits(
max_input_bytes=self.input_port.schema.max_bytes,
max_tasks=10_000,
max_output_bytes=self.output_port.schema.max_bytes,
)
trust_modes = ("trusted", "untrusted_quorum")
map_stage = StageSpec(
stage_id="map",
kind=StageKind.MAP,
entry_point=MAP_ENTRY_POINT,
needs=(),
inputs={"left": self.block_port, "right": self.block_port},
outputs={"partial": self.partial_port},
parameter_names=_MAP_PARAMETERS,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=ComponentRef("exact-artifact", 1),
trust_modes=trust_modes,
max_fan_out=limits.max_tasks,
cacheable=True,
)
reduce_input = PortSpec(
schema=self.partial_port.schema,
cardinality=Cardinality.MANY,
collection=CollectionKind.KEYED,
)
reduce_stage = StageSpec(
stage_id="reduce",
kind=StageKind.REDUCE,
entry_point=REDUCE_ENTRY_POINT,
needs=("map",),
inputs={"partials": reduce_input},
outputs={"result": self.output_port},
parameter_names=_REDUCE_PARAMETERS,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=ComponentRef("exact-artifact", 1),
trust_modes=trust_modes,
max_fan_out=1,
cacheable=True,
)
workflow = WorkflowSpec(
workflow_id="graph-block-pairs-v1",
inputs={"input": self.input_port},
stages=(map_stage, reduce_stage),
edges=(
ArtifactEdge(PortRef("input"), PortRef("left", "map")),
ArtifactEdge(PortRef("input"), PortRef("right", "map")),
ArtifactEdge(PortRef("partial", "map"), PortRef("partials", "reduce")),
),
outputs={"result": PortRef("result", "reduce")},
max_tasks=limits.max_tasks,
max_output_bytes=limits.max_output_bytes,
)
self.manifest = WorkloadManifest(
sdk_api=VersionRange(">=1.0,<2.0"),
protocol=VersionRange(">=1,<2"),
workload=WorkloadId("similarity-graph", "1.0.0"),
description=(
"Exact sparse Tanimoto similarity graph over deterministic "
"block pairs with a duplicate-safe, coverage-checked merge."
),
package=PackageSpec("scimesh", package_digest),
environment=EnvironmentSpec(
"python-process",
environment_digest,
{"adapter": "sdk-native"},
),
parameters_schema=_parameters_schema(),
workflow=workflow,
inputs={"input": self.input_port},
outputs={"result": self.output_port},
determinism=DeterminismProfile.BYTE_EXACT,
trust_modes=(TrustMode.TRUSTED, TrustMode.UNTRUSTED_QUORUM),
verifier=VerifierSpec(ComponentRef("exact-artifact", 1), {}),
limits=limits,
capabilities=("similarity-graph",),
conformance_profiles=("core-batch-v1",),
)
self._exact_verifier = ExactArtifactVerifier()
workload_id = WorkloadId("similarity-graph", "1.0.0")
description = (
"Exact sparse Tanimoto similarity graph over deterministic "
"block pairs with a duplicate-safe, coverage-checked merge."
)
parameters_schema = _parameters_schema()
input_port = PortSpec(_molecule_schema())
block_port = PortSpec(_molecule_schema())
partial_port = PortSpec(_edge_schema())
output_port = PortSpec(_edge_schema())
map_stage_inputs = {"left": block_port, "right": block_port}
map_parameter_names = _MAP_PARAMETERS
reduce_parameter_names = (
"threshold",
"threshold_direction",
"block_size",
"max_rows",
)
workflow_id = "graph-block-pairs-v1"
map_entry_point = MAP_ENTRY_POINT
reduce_entry_point = REDUCE_ENTRY_POINT
def definition(self) -> WorkloadDefinition:
return WorkloadDefinition(
manifest=self.manifest,
planner=self,
runners={MAP_ENTRY_POINT: self},
reducers={REDUCE_ENTRY_POINT: self},
verifiers={self._exact_verifier.identity.canonical: self._exact_verifier},
)
@staticmethod
def _unit_interval(value: object, name: str) -> float:
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValueError(f"{name} must be a number between 0 and 1")
return float(value)
@staticmethod
def _positive_int(value: object, name: str) -> int:
if isinstance(value, bool) or not isinstance(value, int) or value < 1:
raise ValueError(f"{name} must be a positive integer")
return value
def validate(self, request: JobRequest) -> ValidatedJob:
if request.workload != self.manifest.workload:
raise ValueError("similarity-graph received a request for another workload")
parameters = request.parameters
def domain_validate(self, parameters: Mapping[str, Any]) -> None:
unknown = set(parameters) - {
"threshold",
"threshold_direction",
@@ -275,102 +135,87 @@ class SimilarityGraphSDKWorkload:
self._positive_int(parameters["block_size"], "block_size")
if "max_rows" in parameters:
self._positive_int(parameters["max_rows"], "max_rows")
return ValidatedJob(request, request.parameters)
def plan(self, job: ValidatedJob, context: PlanningContext) -> WorkflowPlan:
if not isinstance(job, ValidatedJob):
raise ValueError("job must be a ValidatedJob")
collection = job.request.inputs.get("input")
if collection is None:
raise ValueError("similarity-graph requires the input port")
self.input_port.validate_collection(collection, "job input")
input_artifact = collection.items[0].artifact
input_path = context.catalog.materialize(input_artifact)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
parameters = job.resolved_parameters
threshold = self._unit_interval(parameters.get("threshold"), "threshold")
direction = parameters.get("threshold_direction", "greater")
if direction not in {"greater", "less"}:
raise ValueError("threshold_direction must be 'greater' or 'less'")
@staticmethod
def _unit_interval(value: object, name: str) -> float:
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValueError(f"{name} must be a number between 0 and 1")
return float(value)
@staticmethod
def _positive_int(value: object, name: str) -> int:
if isinstance(value, bool) or not isinstance(value, int) or value < 1:
raise ValueError(f"{name} must be a positive integer")
return value
def partition_input(
self,
input_path: Path,
parameters: Mapping[str, Any],
workspace: Path,
) -> list[Path]:
block_size = int(parameters.get("block_size", 1_000))
max_rows = parameters.get("max_rows")
blocks, stats = parse_molecule_blocks(
blocks, _stats = parse_molecule_blocks(
input_path,
block_size,
int(max_rows) if isinstance(max_rows, int) else None,
)
task_parameters = {
"threshold": threshold,
"threshold_direction": direction,
}
negotiated = context.negotiated
map_stage = self.manifest.workflow.stages[0]
assert map_stage.verifier is not None
block_refs: list[ArtifactRef] = []
paths: list[Path] = []
for index, block in enumerate(blocks):
path = workspace / f"block-{index:04d}.tsv"
write_block_tsv(block, path)
block_refs.append(
context.sink.seal(
path,
declaration=self.block_port.schema,
)
paths.append(path)
return paths
def plan_tasks(
self,
shard_paths: Sequence[Path],
resolved: Mapping[str, Any],
job: ValidatedJob,
negotiated: Any,
map_stage: StageSpec,
context: PlanningContext,
) -> list[TaskSpec]:
task_parameters = {
"threshold": self._unit_interval(resolved.get("threshold"), "threshold"),
"threshold_direction": resolved.get("threshold_direction", "greater"),
}
block_refs = [
context.sink.seal(
path,
declaration=self.input_port.schema,
)
for path in shard_paths
]
tasks: list[TaskSpec] = []
for left in range(len(blocks)):
for right in range(left, len(blocks)):
for left in range(len(block_refs)):
for right in range(left, len(block_refs)):
tasks.append(
TaskSpec(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=map_stage.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
task_key=f"map/{left:04d}x{right:04d}",
stage_id="map",
parameters={
self.task_spec(
map_stage,
job,
negotiated,
f"map/{left:04d}x{right:04d}",
{
**task_parameters,
"left_block": left,
"right_block": right,
},
inputs={
{
"left": ArtifactCollection.single(block_refs[left]),
"right": ArtifactCollection.single(block_refs[right]),
},
expected_outputs={"partial": self.partial_port},
resources=map_stage.resources,
execution=map_stage.execution,
)
)
return WorkflowPlan(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=self.manifest.verifier.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
workflow_id=self.manifest.workflow.workflow_id,
resolved_parameters=dict(parameters),
tasks=tuple(tasks),
)
return tasks
def run(self, context: TaskContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
parameters = context.task.parameters
def compute_shard(
self,
inputs: Mapping[str, Path],
parameters: Mapping[str, Any],
output_path: Path,
) -> Mapping[str, int | float]:
left_block = parameters.get("left_block")
right_block = parameters.get("right_block")
if (
@@ -381,106 +226,32 @@ class SimilarityGraphSDKWorkload:
):
raise ValueError("graph map task requires block indices")
diagonal = left_block == right_block
left_collection = context.task.inputs.get("left")
right_collection = context.task.inputs.get("right")
if left_collection is None or right_collection is None:
raise ValueError("graph map task requires left and right block inputs")
self.block_port.validate_collection(left_collection, "graph map left input")
self.block_port.validate_collection(right_collection, "graph map right input")
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
left_path = context.catalog.materialize(left_collection.items[0].artifact)
right_path = context.catalog.materialize(right_collection.items[0].artifact)
left_rows = read_block_rows(left_path)
right_rows = (
left_rows
if diagonal and left_path.resolve() == right_path.resolve()
else read_block_rows(right_path)
)
threshold = self._unit_interval(parameters.get("threshold"), "threshold")
direction = parameters.get("threshold_direction", "greater")
if direction not in {"greater", "less"}:
raise ValueError("threshold_direction must be 'greater' or 'less'")
left = read_block_rows(inputs["left"])
right = left if diagonal else read_block_rows(inputs["right"])
checked_pairs = (
len(left_rows) * (len(left_rows) - 1) // 2
if diagonal
else len(left_rows) * len(right_rows)
len(left) * (len(left) - 1) // 2 if diagonal else len(left) * len(right)
)
edges = compute_block_edges(
left_rows,
right_rows,
threshold,
direction,
)
output_path = workspace / "result.csv"
edges = compute_block_edges(left, right, threshold, direction)
write_edge_csv(output_path, edges)
context.cancellation.raise_if_cancelled()
sealed = context.sink.seal(
output_path,
declaration=self.partial_port.schema,
)
return OutputManifest(
context.task.task_key,
{"partial": ArtifactCollection.single(sealed)},
{"checked_pairs": checked_pairs, "edges_emitted": len(edges)},
context.provenance,
).validate_against(
context.task.expected_outputs,
max_output_bytes=self.manifest.limits.max_output_bytes,
)
return {"checked_pairs": checked_pairs, "edges_emitted": len(edges)}
def reduce(self, context: ReduceContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
collection = context.accepted_inputs.get("partials")
if (
collection is None
or collection.kind is not CollectionKind.KEYED
or not collection.items
):
raise ValueError(
"graph reducer requires a non-empty keyed partial collection"
)
self.manifest.workflow.stages[1].inputs["partials"].validate_collection(
collection,
"graph reducer partials",
)
pairs = [block_pair_from_key(item.key or "") for item in collection.items]
check_pair_coverage(pairs)
expected_keys = context.task.expected_input_keys.get("partials")
if expected_keys is None or {item.key for item in collection.items} != set(
expected_keys
):
raise ValueError(
"graph partial keys do not match the coordinator expected set"
)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
partial_paths: list[Path] = []
for item in sorted(collection.items, key=lambda value: value.key or ""):
artifact: ArtifactRef = item.artifact
source = context.catalog.materialize(artifact)
target = workspace / artifact.artifact_id
if source.resolve() != target.resolve():
shutil.copyfile(source, target)
if _sha256_file(target) != artifact.sha256:
raise ValueError("materialized partial checksum does not match")
partial_paths.append(target)
result_path = workspace / "result.csv"
metrics = merge_edge_partials(partial_paths, result_path)
context.cancellation.raise_if_cancelled()
sealed = context.sink.seal(
result_path,
declaration=self.output_port.schema,
)
return OutputManifest(
context.task.task_key,
{"result": ArtifactCollection.single(sealed)},
metrics,
context.provenance,
).validate_against(
context.task.expected_outputs,
max_output_bytes=self.manifest.limits.max_output_bytes,
)
def parse_partial_key(self, key: str) -> Any:
return block_pair_from_key(key)
def validate_partial_keys(self, parsed: Sequence[Any]) -> None:
check_pair_coverage(tuple(parsed))
def reduce_partials(
self,
partial_paths: Sequence[Path],
parameters: Mapping[str, Any],
output_path: Path,
) -> Mapping[str, int | float]:
return merge_edge_partials(partial_paths, output_path)
def similarity_graph_sdk_definition(
+27 -8
View File
@@ -27,9 +27,20 @@ __all__ = [
]
def default_sdk_registry(*, shard_rows: int = 10_000) -> WorkloadRegistry:
"""Registry of every built-in SDK-built workload, all enabled."""
def default_sdk_registry(
*,
shard_rows: int = 10_000,
allowlist: tuple[AllowedPackage, ...] | None = None,
) -> WorkloadRegistry:
"""Registry of every built-in SDK-built workload, all enabled.
When ``allowlist`` is provided, installed workloads are discovered through
the ``scimesh.workloads`` entry points instead of the built-ins.
"""
registry = WorkloadRegistry()
if allowlist:
registry.discover_installed(allowlist)
return registry
registry.register(
similarity_search_sdk_definition(shard_rows=shard_rows).definition(),
enabled=True,
@@ -45,8 +56,17 @@ def default_sdk_registry(*, shard_rows: int = 10_000) -> WorkloadRegistry:
return registry
def default_sdk_runtime() -> RuntimeCapabilities:
"""Runtime advertising the built-in workloads' capabilities and inventory."""
def default_sdk_runtime(
*,
workload_capabilities: tuple[str, ...] | None = None,
environment_digests: tuple[str, ...] | None = None,
) -> RuntimeCapabilities:
"""Runtime advertising the built-in workloads' capabilities and inventory.
``workload_capabilities`` and ``environment_digests`` override the built-in
defaults, for example when a registry was populated from an allowlist of
installed user workloads instead of the built-ins.
"""
architecture = platform.machine().lower() or "unknown"
return RuntimeCapabilities(
sdk_api_version=SDK_API_VERSION,
@@ -54,15 +74,14 @@ def default_sdk_runtime() -> RuntimeCapabilities:
profiles=("core-batch-v1",),
features={"artifact-collections": "1.0.0", "exact-verifier": "1.0.0"},
workload_capabilities=(
"similarity-search",
"similarity-graph",
"descriptor-batch",
workload_capabilities
or ("similarity-search", "similarity-graph", "descriptor-batch")
),
inventory=ResourceInventory(
cpu_cores=max(os.cpu_count() or 1, 1),
memory_mb=4096,
scratch_mb=4096,
architecture=architecture,
environment_digests=(current_environment_digest(),),
environment_digests=environment_digests or (current_environment_digest(),),
),
)
+12 -6
View File
@@ -57,14 +57,12 @@ def run_search_shard(
) -> dict[str, int]:
"""Run one planned shard with the local reference implementation.
This is the worker entry used by the SDK-built runner. It deliberately
accepts only resolved ``query_smiles``: resolving an identifier
independently in each shard would make the distributed search
scientifically invalid, so identifier resolution happens once in the
planner (or at the worker bridge for v1-wire tasks that still carry
``query_id``).
Accepts either a resolved ``query_smiles`` or a raw ``query_id`` that is
resolved against the shard (worker tasks on the v1 wire may still carry
the identifier; the SDK planner always resolves it once at plan time).
"""
allowed = {
"query_id",
"query_smiles",
"top_k",
"threshold",
@@ -77,6 +75,14 @@ def run_search_shard(
f"unsupported similarity-search parameters: {', '.join(sorted(unknown))}"
)
query_smiles = parameters.get("query_smiles")
query_id = parameters.get("query_id")
if isinstance(query_id, str) and not isinstance(query_smiles, str):
from rdkit import Chem
from scimesh.chemistry.dataset import find_molecule_by_id
record = find_molecule_by_id(input_path, query_id)
query_smiles = Chem.MolToSmiles(record.molecule, canonical=True)
if not isinstance(query_smiles, str) or not query_smiles.strip():
raise ValueError("query_smiles is required for a distributed shard")
molecule = parse_smiles(query_smiles)
+86 -353
View File
@@ -1,79 +1,41 @@
"""SDK-built ``similarity-search`` workload definition and handlers.
A direct ``core-batch-v1`` definition (not the legacy adapter): the planner
resolves the query once, shards deterministically, each map task computes the
local top-k with the reference implementation, and the reducer merges the
sorted partials with the same bounded heap and tie-breakers as the local CLI.
The manifest declares ``byte_exact`` with the exact-artifact verifier and both
``trusted`` and ``untrusted_quorum`` trust modes.
A thin subclass of ``MapReduceWorkload``: the SDK assembles the manifest,
stages, workflow, and digest-pinned handlers. This module declares the search
scientific contract (plan-time query resolution, deterministic sharding, local
top-k per shard, bounded heap merge) and the hooks that partition, compute,
and merge.
"""
from __future__ import annotations
import hashlib
import shutil
from pathlib import Path
from typing import Any, Mapping
from typing import Any, Mapping, Sequence
from rdkit import Chem
from scimesh.chemistry.dataset import find_molecule_by_id, parse_smiles
from scimesh.chemistry.fingerprints import FP_RADIUS, FP_SIZE
from scimesh.sdk.artifacts import ArtifactSchema, ComponentRef, PortSpec
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 ...sdk.artifacts import (
ArtifactCollection,
ArtifactItem,
ArtifactRef,
ArtifactSchema,
Cardinality,
CollectionKind,
OutputManifest,
PortSpec,
)
from ..environment import current_environment_digest, current_scimesh_package_digest
from ...sdk.execution import (
CheckpointPolicy,
ExecutionProfile,
NetworkPolicy,
RetryPolicy,
)
from ...sdk.identity import ComponentRef, SchemaRef, VersionRange, WorkloadId
from ...sdk.manifest import (
DeterminismProfile,
EnvironmentSpec,
PackageSpec,
TrustMode,
VerifierSpec,
WorkloadLimits,
WorkloadManifest,
)
from ...sdk.plans import JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
from ...sdk.protocols import PlanningContext, ReduceContext, TaskContext
from ...sdk.registry import WorkloadDefinition
from ...sdk.resources import ResourceRequirements
from ...sdk.verification import ExactArtifactVerifier
from ...sdk.workflow import ArtifactEdge, PortRef, StageKind, StageSpec, WorkflowSpec
from .core import merge_search_partials, run_search_shard, write_search_shards
MAP_ENTRY_POINT = "scimesh.workloads.search.definition:map_search@v1"
REDUCE_ENTRY_POINT = "scimesh.workloads.search.definition:reduce_search@v1"
_MAP_PARAMETERS = (
"query_id",
"query_smiles",
"top_k",
"threshold",
"threshold_direction",
"progress_every",
)
_REDUCE_PARAMETERS = _MAP_PARAMETERS + ("query_source", "fingerprint")
def _sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as source:
for block in iter(lambda: source.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def _parameters_schema() -> dict[str, Any]:
@@ -126,8 +88,30 @@ def _search_table_schema(ref: SchemaRef, canonicalizer: str) -> ArtifactSchema:
)
class SimilaritySearchSDKWorkload:
"""Manifest-backed planner, runner, and reducer for similarity-search."""
class SimilaritySearchSDKWorkload(MapReduceWorkload):
"""Exact top-k Tanimoto search over deterministic shards with a bounded merge."""
workload_id = WorkloadId("similarity-search", "1.0.0")
description = (
"Exact top-k Tanimoto molecular similarity search over "
"deterministic TSV shards with a bounded merge."
)
parameters_schema = _parameters_schema()
input_port = PortSpec(_dataset_schema())
partial_port = PortSpec(
_search_table_schema(
SchemaRef("similarity-search-partial", 1), "scimesh-search-partial-v1"
)
)
output_port = PortSpec(
_search_table_schema(
SchemaRef("similarity-search-result", 1), "scimesh-search-result-v1"
)
)
map_parameter_names = _MAP_PARAMETERS
reduce_parameter_names = _MAP_PARAMETERS + ("query_source", "fingerprint")
map_entry_point = MAP_ENTRY_POINT
reduce_entry_point = REDUCE_ENTRY_POINT
def __init__(
self,
@@ -142,122 +126,15 @@ class SimilaritySearchSDKWorkload:
or shard_rows < 1
):
raise ValueError("shard_rows must be a positive integer")
self.entry_point = MAP_ENTRY_POINT
self.shard_rows = shard_rows
self.input_port = PortSpec(_dataset_schema())
self.partial_port = PortSpec(
_search_table_schema(
SchemaRef("similarity-search-partial", 1), "scimesh-search-partial-v1"
)
super().__init__(
package_digest=package_digest,
environment_digest=environment_digest,
)
self.output_port = PortSpec(
_search_table_schema(
SchemaRef("similarity-search-result", 1), "scimesh-search-result-v1"
)
)
resources = ResourceRequirements(
profile="search-cpu-v1",
cpu_cores=1,
memory_mb=1024,
scratch_mb=1024,
max_duration_seconds=3600,
)
execution = ExecutionProfile(
profile="search-python-process-v1",
network=NetworkPolicy.TRUSTED,
timeout_seconds=3600,
checkpoint=CheckpointPolicy(),
)
limits = WorkloadLimits(
max_input_bytes=self.input_port.schema.max_bytes,
max_tasks=10_000,
max_output_bytes=self.output_port.schema.max_bytes,
)
trust_modes = ("trusted", "untrusted_quorum")
map_stage = StageSpec(
stage_id="map",
kind=StageKind.MAP,
entry_point=MAP_ENTRY_POINT,
needs=(),
inputs={"input": self.input_port},
outputs={"partial": self.partial_port},
parameter_names=_MAP_PARAMETERS,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=ComponentRef("exact-artifact", 1),
trust_modes=trust_modes,
max_fan_out=limits.max_tasks,
cacheable=True,
)
reduce_input = PortSpec(
schema=self.partial_port.schema,
cardinality=Cardinality.MANY,
collection=CollectionKind.KEYED,
)
reduce_stage = StageSpec(
stage_id="reduce",
kind=StageKind.REDUCE,
entry_point=REDUCE_ENTRY_POINT,
needs=("map",),
inputs={"partials": reduce_input},
outputs={"result": self.output_port},
parameter_names=_REDUCE_PARAMETERS,
resources=resources,
execution=execution,
retry=RetryPolicy(),
verifier=ComponentRef("exact-artifact", 1),
trust_modes=trust_modes,
max_fan_out=1,
cacheable=True,
)
workflow = WorkflowSpec(
workflow_id="search-map-reduce-v1",
inputs={"input": self.input_port},
stages=(map_stage, reduce_stage),
edges=(
ArtifactEdge(PortRef("input"), PortRef("input", "map")),
ArtifactEdge(PortRef("partial", "map"), PortRef("partials", "reduce")),
),
outputs={"result": PortRef("result", "reduce")},
max_tasks=limits.max_tasks,
max_output_bytes=limits.max_output_bytes,
)
self.manifest = WorkloadManifest(
sdk_api=VersionRange(">=1.0,<2.0"),
protocol=VersionRange(">=1,<2"),
workload=WorkloadId("similarity-search", "1.0.0"),
description=(
"Exact top-k Tanimoto molecular similarity search over "
"deterministic TSV shards with a bounded merge."
),
package=PackageSpec("scimesh", package_digest),
environment=EnvironmentSpec(
"python-process",
environment_digest,
{"adapter": "sdk-native"},
),
parameters_schema=_parameters_schema(),
workflow=workflow,
inputs={"input": self.input_port},
outputs={"result": self.output_port},
determinism=DeterminismProfile.BYTE_EXACT,
trust_modes=(TrustMode.TRUSTED, TrustMode.UNTRUSTED_QUORUM),
verifier=VerifierSpec(ComponentRef("exact-artifact", 1), {}),
limits=limits,
capabilities=("similarity-search",),
conformance_profiles=("core-batch-v1",),
)
self._exact_verifier = ExactArtifactVerifier()
def definition(self) -> WorkloadDefinition:
return WorkloadDefinition(
manifest=self.manifest,
planner=self,
runners={MAP_ENTRY_POINT: self},
reducers={REDUCE_ENTRY_POINT: self},
verifiers={self._exact_verifier.identity.canonical: self._exact_verifier},
)
# ------------------------------------------------------------------
# Scientific hooks
# ------------------------------------------------------------------
@staticmethod
def _string(value: object, name: str) -> str:
@@ -283,12 +160,7 @@ class SimilaritySearchSDKWorkload:
raise ValueError(f"{name} must be a number between 0 and 1")
return float(value)
def validate(self, request: JobRequest) -> ValidatedJob:
if request.workload != self.manifest.workload:
raise ValueError(
"similarity-search received a request for another workload"
)
parameters = request.parameters
def domain_validate(self, parameters: Mapping[str, Any]) -> None:
unknown = set(parameters) - {
"query_id",
"query_smiles",
@@ -322,9 +194,8 @@ class SimilaritySearchSDKWorkload:
"threshold_direction"
] not in {"greater", "less"}:
raise ValueError("threshold_direction must be 'greater' or 'less'")
return ValidatedJob(request, self._resolved_parameters(request))
def _resolved_parameters(self, request: JobRequest) -> dict[str, object]:
def resolved_parameters(self, request: JobRequest) -> dict[str, Any]:
parameters = request.parameters
query_id = parameters.get("query_id")
if isinstance(query_id, str):
@@ -334,7 +205,7 @@ class SimilaritySearchSDKWorkload:
"kind": "smiles",
"value": self._string(parameters.get("query_smiles"), "query_smiles"),
}
resolved: dict[str, object] = {
resolved: dict[str, Any] = {
"query_source": query_source,
"top_k": self._positive_int(parameters.get("top_k", 20), "top_k"),
"threshold_direction": parameters.get("threshold_direction", "greater"),
@@ -358,192 +229,54 @@ class SimilaritySearchSDKWorkload:
)
return resolved
def plan(self, job: ValidatedJob, context: PlanningContext) -> WorkflowPlan:
if not isinstance(job, ValidatedJob):
raise ValueError("job must be a ValidatedJob")
collection = job.request.inputs.get("input")
if collection is None:
raise ValueError("similarity-search requires the input port")
self.input_port.validate_collection(collection, "job input")
input_artifact = collection.items[0].artifact
input_path = context.catalog.materialize(input_artifact)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
resolved = dict(job.resolved_parameters)
query_smiles = self._resolve_query(input_path, job.request.parameters)
resolved["query_smiles"] = query_smiles
max_rows = resolved.get("max_rows")
shard_paths = write_search_shards(
def resolved_parameters_for_plan(
self,
job: ValidatedJob,
input_path: Path,
resolved: dict[str, Any],
) -> dict[str, Any]:
query_id = job.request.parameters.get("query_id")
if isinstance(query_id, str):
record = find_molecule_by_id(input_path, query_id)
resolved["query_smiles"] = Chem.MolToSmiles(record.molecule, canonical=True)
else:
supplied = job.request.parameters["query_smiles"]
assert isinstance(supplied, str)
molecule = parse_smiles(supplied)
if molecule is None:
raise ValueError("query_smiles is invalid")
resolved["query_smiles"] = Chem.MolToSmiles(molecule, canonical=True)
return resolved
def partition_input(
self,
input_path: Path,
parameters: Mapping[str, Any],
workspace: Path,
) -> list[Path]:
max_rows = parameters.get("max_rows")
return write_search_shards(
input_path,
workspace,
self.shard_rows,
int(max_rows) if isinstance(max_rows, int) else None,
)
task_parameters = {
key: value for key, value in resolved.items() if key in set(_MAP_PARAMETERS)
}
negotiated = context.negotiated
map_stage = self.manifest.workflow.stages[0]
assert map_stage.verifier is not None
tasks: list[TaskSpec] = []
for index, path in enumerate(shard_paths):
sealed = context.sink.seal(
path,
declaration=self.input_port.schema,
)
tasks.append(
TaskSpec(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=map_stage.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
task_key=f"map/{index:08d}",
stage_id="map",
parameters=task_parameters,
inputs={"input": ArtifactCollection.single(sealed)},
expected_outputs={"partial": self.partial_port},
resources=map_stage.resources,
execution=map_stage.execution,
)
)
return WorkflowPlan(
workload=self.manifest.workload,
package_digest=self.manifest.package.digest,
manifest_digest=self.manifest.digest,
trust_mode=job.request.trust_mode,
sdk_api_version=negotiated.sdk_api_version,
protocol_version=negotiated.protocol_version,
manifest_schema_version=self.manifest.manifest_schema_version,
workflow_schema_version=self.manifest.workflow.schema_version,
environment_digest=self.manifest.environment.digest,
verifier=self.manifest.verifier.verifier,
selected_features=negotiated.selected_features,
optional_fallbacks=negotiated.optional_fallbacks,
workflow_id=self.manifest.workflow.workflow_id,
resolved_parameters=resolved,
tasks=tuple(tasks),
)
@staticmethod
def _resolve_query(input_path: Path, parameters: Mapping[str, object]) -> str:
query_id = parameters.get("query_id")
if isinstance(query_id, str):
record = find_molecule_by_id(input_path, query_id)
return Chem.MolToSmiles(record.molecule, canonical=True)
supplied = parameters["query_smiles"]
assert isinstance(supplied, str)
molecule = parse_smiles(supplied)
if molecule is None:
raise ValueError("query_smiles is invalid")
return Chem.MolToSmiles(molecule, canonical=True)
def compute_shard(
self,
inputs: Mapping[str, Path],
parameters: Mapping[str, Any],
output_path: Path,
) -> Mapping[str, int | float]:
return run_search_shard(inputs["input"], parameters, output_path)
def run(self, context: TaskContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
collection = context.task.inputs.get("input")
if collection is None:
raise ValueError("search map task requires one input collection")
self.input_port.validate_collection(collection, "search map input")
source = context.catalog.materialize(collection.items[0].artifact)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
input_path = workspace / "input"
output_path = workspace / "result.csv"
if source.resolve() != input_path.resolve():
shutil.copyfile(source, input_path)
metrics = run_search_shard(
input_path,
context.task.parameters,
output_path,
)
context.cancellation.raise_if_cancelled()
sealed = context.sink.seal(
output_path,
declaration=self.partial_port.schema,
)
return OutputManifest(
context.task.task_key,
{"partial": ArtifactCollection.single(sealed)},
metrics,
context.provenance,
).validate_against(
context.task.expected_outputs,
max_output_bytes=self.manifest.limits.max_output_bytes,
)
def reduce(self, context: ReduceContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
collection = context.accepted_inputs.get("partials")
if (
collection is None
or collection.kind is not CollectionKind.KEYED
or not collection.items
):
raise ValueError(
"search reducer requires a non-empty keyed partial collection"
)
self.manifest.workflow.stages[1].inputs["partials"].validate_collection(
collection,
"search reducer partials",
)
workspace = context.workspace
workspace.mkdir(parents=True, exist_ok=True)
indexed_items: list[tuple[int, ArtifactItem]] = []
for item in collection.items:
key = item.key or ""
prefix = "map."
raw_index = key[len(prefix) :] if key.startswith(prefix) else ""
if len(raw_index) != 8 or not raw_index.isdigit():
raise ValueError("search partial key must use map.<eight-digit-index>")
indexed_items.append((int(raw_index), item))
expected_keys = context.task.expected_input_keys.get("partials")
if expected_keys is None or {item.key for item in collection.items} != set(
expected_keys
):
raise ValueError(
"search partial keys do not match the coordinator expected set"
)
if sorted(index for index, _ in indexed_items) != list(
range(len(indexed_items))
):
raise ValueError("search partial keys must be complete and contiguous")
partial_paths: list[Path] = []
for index, item in sorted(indexed_items):
artifact: ArtifactRef = item.artifact
source = context.catalog.materialize(artifact)
target = workspace / artifact.artifact_id
if source.resolve() != target.resolve():
shutil.copyfile(source, target)
if _sha256_file(target) != artifact.sha256:
raise ValueError("materialized partial checksum does not match")
partial_paths.append(target)
result_path = workspace / "result.csv"
metrics = merge_search_partials(
partial_paths,
context.task.parameters,
result_path,
)
context.cancellation.raise_if_cancelled()
sealed = context.sink.seal(
result_path,
declaration=self.output_port.schema,
)
return OutputManifest(
context.task.task_key,
{"result": ArtifactCollection.single(sealed)},
metrics,
context.provenance,
).validate_against(
context.task.expected_outputs,
max_output_bytes=self.manifest.limits.max_output_bytes,
)
def reduce_partials(
self,
partial_paths: Sequence[Path],
parameters: Mapping[str, Any],
output_path: Path,
) -> Mapping[str, int | float]:
return merge_search_partials(partial_paths, parameters, output_path)
def similarity_search_sdk_definition(
+167
View File
@@ -0,0 +1,167 @@
"""Generic SDK workload runner CLI: ``scimesh workload list|run``.
This is a generic SDK tool, not a workload. It lists the enabled SDK-built
workloads and executes any of them locally through ``LocalCoreBatchExecutor``,
so a user-written workload package can be inspected and verified without
touching any other part of the program.
"""
from __future__ import annotations
import argparse
import json
import shutil
import tempfile
from pathlib import Path
from scimesh.sdk import (
ArtifactCollection,
JobRequest,
LocalArtifactStore,
LocalCoreBatchExecutor,
)
from scimesh.sdk.registry import workload_allowlist_from_json
from scimesh.workloads.library import default_sdk_registry, default_sdk_runtime
class WorkloadCLI:
"""Inspect and run SDK-built workloads from the command line."""
name = "workload"
help = "List and run SDK-built workloads locally."
def configure_parser(self, parser: argparse.ArgumentParser) -> None:
subparsers = parser.add_subparsers(dest="workload_command", required=True)
list_parser = subparsers.add_parser(
"list", help="List installed and enabled SDK workloads."
)
list_parser.set_defaults(workload_handler=self.list_workloads)
run_parser = subparsers.add_parser(
"run", help="Run one SDK workload locally against an input file."
)
run_parser.add_argument(
"name", help="Workload name, for example descriptor-batch"
)
run_parser.add_argument(
"--version", help="Exact workload version (default: the enabled one)"
)
run_parser.add_argument(
"--input", required=True, type=Path, help="Input dataset file"
)
run_parser.add_argument(
"--params", default="{}", help="Job parameters as a JSON object"
)
run_parser.add_argument(
"--shard-rows",
type=int,
default=10_000,
help="Rows per planned shard for workloads that shard by rows",
)
run_parser.add_argument(
"-o",
"--output",
type=Path,
default=Path("workload_result.csv"),
help="Output path for the final artifact",
)
run_parser.add_argument(
"--work-dir",
type=Path,
help="Temporary working directory (default: a fresh temporary directory)",
)
run_parser.set_defaults(workload_handler=self.run_workload)
def run(self, args: argparse.Namespace) -> int:
handler = getattr(args, "workload_handler", None)
if handler is None:
raise ValueError("select a workload subcommand: list or run")
return handler(args)
@staticmethod
def _registry(args: argparse.Namespace):
import os
allowlist = workload_allowlist_from_json(
os.getenv("SCIMESH_WORKLOAD_ALLOWLIST")
)
return default_sdk_registry(
shard_rows=getattr(args, "shard_rows", 10_000),
allowlist=allowlist,
)
def list_workloads(self, args: argparse.Namespace) -> int:
registry = self._registry(args)
descriptions = registry.descriptions()
if not descriptions:
print("No SDK workloads are installed or enabled.")
return 0
width = max(len(item.workload.name) for item in descriptions)
for item in sorted(descriptions, key=lambda value: value.workload.name):
digest = item.package_digest.removeprefix("sha256:")[:12]
state = "enabled" if item.enabled else "disabled"
print(
f"{item.workload.name:<{width}} {item.workload.version} "
f"{item.description} [{state} {digest}]"
)
return 0
def run_workload(self, args: argparse.Namespace) -> int:
registry = self._registry(args)
descriptions = registry.descriptions()
try:
parameters = json.loads(args.params)
except (TypeError, json.JSONDecodeError, RecursionError) as error:
raise ValueError("--params must be a valid JSON object") from error
if not isinstance(parameters, dict):
raise ValueError("--params must be a JSON object")
description = next(
(
item
for item in descriptions
if item.workload.name == args.name
and (args.version is None or item.workload.version == args.version)
),
None,
)
if description is None:
raise ValueError(f"unknown or disabled SDK workload: {args.name}")
definition, _ = registry.require(
description.workload.name,
description.workload.version,
description.package_digest,
)
runtime = default_sdk_runtime(
workload_capabilities=tuple(item.workload.name for item in descriptions),
environment_digests=(definition.manifest.environment.digest,),
)
if not args.input.is_file():
raise ValueError(f"input file does not exist: {args.input}")
with tempfile.TemporaryDirectory(prefix="scimesh-workload-") as temporary:
root = Path(temporary)
store = LocalArtifactStore(root / "artifacts")
artifact = store.import_file(
args.input,
declaration=definition.manifest.inputs["input"].schema,
)
request = JobRequest(
workload=definition.manifest.workload,
parameters=parameters,
inputs={"input": ArtifactCollection.single(artifact)},
)
result = LocalCoreBatchExecutor(
registry,
runtime,
store,
args.work_dir or root / "attempts",
).execute(request, description.package_digest)
result_artifact = result.outputs["result"].items[0].artifact
source = store.materialize(result_artifact)
args.output.parent.mkdir(parents=True, exist_ok=True)
shutil.copyfile(source, args.output)
print(
f"Saved {description.workload.name} result to {args.output} "
f"(metrics: {dict(result.metrics)})"
)
return 0