Add MapReduceWorkload scaffold and generic workload execution
This commit is contained in:
@@ -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())
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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,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(),),
|
||||
),
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user