Replace legacy distributed protocol with SDK-built workloads

This commit is contained in:
Emil
2026-08-01 23:57:41 +03:00
parent 96169086f0
commit 19fbb8e926
34 changed files with 3009 additions and 1986 deletions
+41
View File
@@ -0,0 +1,41 @@
"""SDK-built ``descriptor-batch`` reference workload.
See ``core.py`` for the pinned scientific contract and ``definition.py`` for
the manifest-backed planner/runner/reducer handlers.
"""
from .core import (
DESCRIPTOR_COLUMNS,
DESCRIPTOR_NAMES,
DescriptorRow,
compute_descriptor_batch,
concatenate_descriptor_shards,
descriptor_calculator,
validate_descriptor_names,
write_descriptor_rows,
write_descriptor_shards,
)
from .definition import (
MAP_ENTRY_POINT,
REDUCE_ENTRY_POINT,
DescriptorBatchWorkload,
descriptor_batch_sdk_definition,
workload_definition,
)
__all__ = [
"DESCRIPTOR_COLUMNS",
"DESCRIPTOR_NAMES",
"MAP_ENTRY_POINT",
"REDUCE_ENTRY_POINT",
"DescriptorBatchWorkload",
"DescriptorRow",
"compute_descriptor_batch",
"concatenate_descriptor_shards",
"descriptor_batch_sdk_definition",
"descriptor_calculator",
"validate_descriptor_names",
"workload_definition",
"write_descriptor_rows",
"write_descriptor_shards",
]
+314
View File
@@ -0,0 +1,314 @@
"""Pinned RDKit 2D descriptor computation for the descriptor-batch workload.
The scientific contract of ``descriptor-batch`` is deliberately small and
fully pinned:
- exactly one output row per valid input molecule, in input order;
- RDKit canonical SMILES recomputed with ``MolToSmiles(..., canonical=True)``;
- the descriptor set is an explicit, versioned tuple of RDKit
``Descriptors.descList`` names (2D only), not a scan of installed names;
- float values are serialized with fixed ``%.6f`` formatting so that the
output is byte-identical for identical inputs and a pinned environment;
- invalid SMILES rows are either skipped (counted) or fail the run, selected
by the explicit ``skip_invalid`` parameter.
"""
from __future__ import annotations
import csv
from dataclasses import dataclass
from functools import lru_cache
from pathlib import Path
from typing import Any, Iterator, Mapping, Sequence
from rdkit import Chem
from rdkit.ML.Descriptors.MoleculeDescriptors import MolecularDescriptorCalculator
from scimesh.chemistry.dataset import iter_rows
# Explicit pinned list. Names must exist in the installed RDKit ``descList``;
# the list itself is the reproducibility contract and must change version
# together with the workload (descriptor-batch@1.0.0).
DESCRIPTOR_NAMES: tuple[str, ...] = (
"ExactMolWt",
"MolWt",
"HeavyAtomMolWt",
"HeavyAtomCount",
"NumHDonors",
"NumHAcceptors",
"NumRotatableBonds",
"NumHeteroatoms",
"NumRadicalElectrons",
"NumValenceElectrons",
"FractionCSP3",
"RingCount",
"NumAromaticRings",
"NumSaturatedRings",
"NumAliphaticRings",
"NumAromaticHeterocycles",
"NumSaturatedHeterocycles",
"NumAliphaticHeterocycles",
"NumAromaticCarbocycles",
"NumSaturatedCarbocycles",
"NumAliphaticCarbocycles",
"TPSA",
"LabuteASA",
"MolLogP",
"MolMR",
"BalabanJ",
"BertzCT",
"HallKierAlpha",
"Kappa1",
"Kappa2",
"Kappa3",
"Chi0",
"Chi1",
"Chi0n",
"Chi1n",
"Chi2n",
"Chi3n",
"Chi4n",
"Chi0v",
"Chi1v",
"Chi2v",
"Chi3v",
"Chi4v",
"PEOE_VSA1",
"PEOE_VSA2",
"PEOE_VSA3",
"PEOE_VSA4",
"PEOE_VSA5",
"PEOE_VSA6",
"PEOE_VSA7",
"PEOE_VSA8",
"PEOE_VSA9",
"PEOE_VSA10",
"PEOE_VSA11",
"PEOE_VSA12",
"PEOE_VSA13",
"PEOE_VSA14",
"SMR_VSA1",
"SMR_VSA2",
"SMR_VSA3",
"SMR_VSA4",
"SMR_VSA5",
"SMR_VSA6",
"SMR_VSA7",
"SMR_VSA8",
"SMR_VSA9",
"SMR_VSA10",
"SlogP_VSA1",
"SlogP_VSA2",
"SlogP_VSA3",
"SlogP_VSA4",
"SlogP_VSA5",
"SlogP_VSA6",
"SlogP_VSA7",
"SlogP_VSA8",
"SlogP_VSA9",
"SlogP_VSA10",
"SlogP_VSA11",
"SlogP_VSA12",
"NHOHCount",
"NOCount",
)
DESCRIPTOR_COLUMNS: tuple[str, ...] = (
"chembl_id",
"canonical_smiles",
) + DESCRIPTOR_NAMES
@lru_cache(maxsize=1)
def descriptor_calculator() -> MolecularDescriptorCalculator:
"""Build the pinned calculator once per process."""
return MolecularDescriptorCalculator(DESCRIPTOR_NAMES)
def validate_descriptor_names() -> None:
"""Fail fast when the pinned list is unavailable in the installed RDKit."""
from rdkit.Chem import Descriptors
available = {name for name, _ in Descriptors.descList}
missing = [name for name in DESCRIPTOR_NAMES if name not in available]
if missing:
raise ValueError(
"pinned descriptor-batch descriptors are missing from RDKit: "
+ ", ".join(missing)
)
@dataclass(frozen=True)
class DescriptorRow:
"""One canonical descriptor row for a valid input molecule."""
molecule_id: str
canonical_smiles: str
values: tuple[float, ...]
class DescriptorStats:
"""Row counters collected while computing a descriptor batch."""
def __init__(self) -> None:
self.scanned = 0
self.invalid = 0
self.emitted = 0
def as_metrics(self) -> dict[str, int]:
return {
"rows_scanned": self.scanned,
"invalid_rows": self.invalid,
"rows_emitted": self.emitted,
}
def iter_descriptor_rows(
input_path: Path,
*,
skip_invalid: bool = True,
) -> tuple[Iterator[DescriptorRow], DescriptorStats]:
"""Yield canonical descriptor rows in input order with streaming stats."""
calculator = descriptor_calculator()
stats = DescriptorStats()
def generate() -> Iterator[DescriptorRow]:
for row in iter_rows(input_path):
stats.scanned += 1
smiles = row.get("canonical_smiles", "")
molecule = Chem.MolFromSmiles(smiles)
if molecule is None:
stats.invalid += 1
if not skip_invalid:
raise ValueError(
f"row {stats.scanned} has an invalid canonical_smiles"
)
continue
canonical = Chem.MolToSmiles(molecule, canonical=True)
values = tuple(
float(value) for value in calculator.CalcDescriptors(molecule)
)
stats.emitted += 1
yield DescriptorRow(row.get("chembl_id", ""), canonical, values)
return generate(), stats
def write_descriptor_rows(
output_path: Path,
rows: Sequence[DescriptorRow],
) -> None:
"""Write a canonical one-row-per-input descriptor CSV with one header."""
output_path.parent.mkdir(parents=True, exist_ok=True)
with output_path.open("w", encoding="utf-8", newline="") as destination:
writer = csv.DictWriter(
destination, fieldnames=list(DESCRIPTOR_COLUMNS), lineterminator="\n"
)
writer.writeheader()
for row in rows:
writer.writerow(
{
"chembl_id": row.molecule_id,
"canonical_smiles": row.canonical_smiles,
**{
name: f"{value:.6f}"
for name, value in zip(DESCRIPTOR_NAMES, row.values)
},
}
)
def compute_descriptor_batch(
input_path: Path,
output_path: Path,
*,
skip_invalid: bool = True,
) -> dict[str, int]:
"""Single-process reference: read the whole input and write the CSV."""
rows, stats = iter_descriptor_rows(input_path, skip_invalid=skip_invalid)
materialized = list(rows)
write_descriptor_rows(output_path, materialized)
return stats.as_metrics()
def write_descriptor_shards(
input_path: Path,
workspace: Path,
shard_rows: int,
) -> list[Path]:
"""Split the input TSV into deterministic row-bounded shards with headers."""
if (
isinstance(shard_rows, bool)
or not isinstance(shard_rows, int)
or shard_rows < 1
):
raise ValueError("shard_rows must be a positive integer")
paths: list[Path] = []
current: Path | None = None
destination = None
writer = None
rows_in_shard = 0
try:
with input_path.open("r", encoding="utf-8", newline="") as source:
reader = csv.DictReader(source, delimiter="\t")
fieldnames = tuple(reader.fieldnames or ())
if not {"chembl_id", "canonical_smiles"}.issubset(set(fieldnames)):
raise ValueError(
"dataset is missing required columns: chembl_id, canonical_smiles"
)
for row in reader:
if destination is None or rows_in_shard == shard_rows:
if destination is not None:
destination.close()
current = workspace / f"shard-{len(paths)}.tsv"
destination = current.open("w", encoding="utf-8", newline="")
writer = csv.DictWriter(
destination,
fieldnames=list(fieldnames),
delimiter="\t",
lineterminator="\n",
)
writer.writeheader()
paths.append(current)
rows_in_shard = 0
assert writer is not None
writer.writerow(row)
rows_in_shard += 1
finally:
if destination is not None:
destination.close()
if not paths:
raise ValueError("dataset has no data rows")
return paths
def concatenate_descriptor_shards(
partial_paths: Sequence[Path],
output_path: Path,
) -> dict[str, int]:
"""Merge shard partial CSVs by shard index with exactly one header.
Every partial is a full CSV with the same header. The first partial is
copied verbatim; each later partial contributes only its data rows, so the
merged file is byte-identical to the single-process reference for the same
input rows.
"""
if not partial_paths:
raise ValueError("descriptor reducer requires at least one partial")
output_path.parent.mkdir(parents=True, exist_ok=True)
rows_emitted = 0
with output_path.open("w", encoding="utf-8", newline="") as destination:
for index, partial in enumerate(partial_paths):
with partial.open("r", encoding="utf-8", newline="") as source:
for line_index, line in enumerate(source):
if line_index == 0:
if index > 0:
continue
if line.rstrip("\r\n") != ",".join(DESCRIPTOR_COLUMNS):
raise ValueError(
"partial descriptor CSV has an invalid header"
)
destination.write(line)
if line_index > 0:
rows_emitted += 1
return {"partial_count": len(partial_paths), "rows_emitted": rows_emitted}
+443
View File
@@ -0,0 +1,443 @@
"""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.
"""
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 ..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,
concatenate_descriptor_shards,
validate_descriptor_names,
write_descriptor_shards,
)
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 {
"type": "object",
"additionalProperties": False,
"properties": {
"skip_invalid": {
"type": "boolean",
"default": True,
"description": "Skip rows with invalid SMILES instead of failing",
},
},
}
def _input_schema() -> ArtifactSchema:
return ArtifactSchema(
SchemaRef("molecule-table", 1),
"text/tab-separated-values",
"utf-8",
max_bytes=10 * 1024 * 1024 * 1024,
validator=ComponentRef("delimited-table", 1),
validator_configuration={
"required_columns": ["canonical_smiles", "chembl_id"],
},
max_records=100_000_000,
canonicalizer="scimesh-tsv-v1",
)
def _descriptor_schema() -> ArtifactSchema:
return ArtifactSchema(
SchemaRef("descriptor-table", 1),
"text/csv",
"utf-8",
max_bytes=100 * 1024 * 1024 * 1024,
validator=ComponentRef("delimited-table", 1),
validator_configuration={
"columns": list(DESCRIPTOR_COLUMNS),
},
max_records=100_000_000,
canonicalizer="descriptor-table-v1",
)
class DescriptorBatchWorkload:
"""Manifest-backed planner, runner, and reducer for descriptor-batch.
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.
"""
def __init__(
self,
*,
shard_rows: int,
package_digest: str,
environment_digest: str,
) -> None:
if (
isinstance(shard_rows, bool)
or not isinstance(shard_rows, int)
or shard_rows < 1
):
raise ValueError("shard_rows must be a positive integer")
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,
)
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},
)
@staticmethod
def _skip_invalid(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 descriptor_batch_sdk_definition(
*,
shard_rows: int = 10_000,
package_digest: str | None = None,
environment_digest: str | None = None,
) -> DescriptorBatchWorkload:
"""Build the default local descriptor-batch definition for tests."""
return DescriptorBatchWorkload(
shard_rows=shard_rows,
package_digest=package_digest or current_scimesh_package_digest(),
environment_digest=environment_digest or current_environment_digest(),
)
def workload_definition() -> WorkloadDefinition:
"""Installed entry-point factory for the default descriptor-batch definition."""
return descriptor_batch_sdk_definition().definition()
+42
View File
@@ -0,0 +1,42 @@
"""Local development digest helpers for built-in SDK workload definitions.
This module is a leaf: it must not import other SDK workload modules, so the
workload definition packages (``search``, ``graph``, ``descriptors``) and the
``builtins`` registry wiring can import it without creating cycles.
"""
from __future__ import annotations
import hashlib
import os
import platform
import sys
from rdkit import rdBase
from scimesh.sdk.integrity import installed_distribution_digest
def current_scimesh_package_digest() -> str:
"""Hash installed SciMesh Python sources for the built-in trusted adapters.
This is a local immutable-code pin, not a package signature or container
attestation. Consequently the built-in manifests are trusted only; an
administrator must supply signed image metadata before enabling an
untrusted quorum policy.
"""
# Source/editable installs are allowed only for this explicit local
# development helper. Registry discovery keeps the secure default.
return installed_distribution_digest("scimesh", allow_editable=True)
def current_environment_digest() -> str:
payload = "\n".join(
(
current_scimesh_package_digest(),
f"python={sys.implementation.name}-{platform.python_version()}",
f"rdkit={rdBase.rdkitVersion}",
f"platform={sys.platform}-{platform.machine().lower()}",
)
)
return "sha256:" + hashlib.sha256(payload.encode("utf-8")).hexdigest()
+40
View File
@@ -0,0 +1,40 @@
"""SDK-built ``similarity-graph`` workload.
A user workload script built on the SciMesh Workload SDK. See ``core.py`` for
the block/pair scientific core and ``definition.py`` for the manifest-backed
planner/runner/reducer handlers.
"""
from .core import (
block_pair_from_key,
check_pair_coverage,
compute_block_edges,
merge_edge_partials,
parse_molecule_blocks,
read_block_rows,
write_block_tsv,
write_edge_csv,
)
from .definition import (
MAP_ENTRY_POINT,
REDUCE_ENTRY_POINT,
SimilarityGraphSDKWorkload,
similarity_graph_sdk_definition,
workload_definition,
)
__all__ = [
"MAP_ENTRY_POINT",
"REDUCE_ENTRY_POINT",
"SimilarityGraphSDKWorkload",
"block_pair_from_key",
"check_pair_coverage",
"compute_block_edges",
"merge_edge_partials",
"parse_molecule_blocks",
"read_block_rows",
"similarity_graph_sdk_definition",
"workload_definition",
"write_block_tsv",
"write_edge_csv",
]
+249
View File
@@ -0,0 +1,249 @@
"""Scientific core for the SDK-built ``similarity-graph`` workload.
Molecules are parsed once into deterministic row-ordered blocks; every block
pair ``(i, j)`` with ``i <= j`` becomes one map task (diagonal tasks compare
pairs ``a < b`` inside a block, off-diagonal tasks compare every molecule
across two blocks). The reducer enforces the CTX-10 pair-coverage invariant:
the union of task pair sets must equal all unordered molecule pairs exactly
once, and the merged edge set must contain no duplicate unordered pair.
"""
from __future__ import annotations
import csv
from pathlib import Path
from typing import Iterable, Sequence
from rdkit import Chem, DataStructs
from scimesh.chemistry.dataset import iter_valid_molecules
from scimesh.chemistry.fingerprints import fingerprint
EDGE_COLUMNS = ("source_id", "target_id", "similarity")
MoleculeBlock = list[tuple[str, str]] # (chembl_id, smiles), row-ordered
def parse_molecule_blocks(
input_path: Path,
block_size: int,
max_rows: int | None = None,
) -> tuple[list[MoleculeBlock], dict[str, int]]:
"""Parse valid molecules into deterministic row-ordered blocks.
Mirrors the local reference's strictness: an empty or duplicate
``chembl_id`` fails the run, because the edge identity is the molecule id.
Invalid SMILES rows are skipped and counted.
"""
if (
isinstance(block_size, bool)
or not isinstance(block_size, int)
or block_size < 1
):
raise ValueError("block_size must be a positive integer")
from scimesh.chemistry.dataset import DatasetStats
stats = DatasetStats()
blocks: list[MoleculeBlock] = []
current: MoleculeBlock = []
seen_ids: set[str] = set()
for record in iter_valid_molecules(input_path, stats, max_rows=max_rows):
if not record.molecule_id:
raise ValueError("dataset contains an empty chembl_id")
if record.molecule_id in seen_ids:
raise ValueError(
f"dataset contains a duplicate chembl_id: {record.molecule_id}"
)
seen_ids.add(record.molecule_id)
current.append((record.molecule_id, record.smiles))
if len(current) == block_size:
blocks.append(current)
current = []
if current:
blocks.append(current)
if not blocks:
raise ValueError("dataset has no valid molecules")
return blocks, {
"rows_scanned": stats.scanned,
"valid_molecules": stats.valid,
"invalid_smiles": stats.invalid,
"block_count": len(blocks),
}
def write_block_tsv(rows: MoleculeBlock, path: Path) -> None:
"""Write one molecule block as a header TSV with the input column names."""
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8", newline="") as destination:
writer = csv.DictWriter(
destination,
fieldnames=["chembl_id", "canonical_smiles"],
delimiter="\t",
lineterminator="\n",
)
writer.writeheader()
for molecule_id, smiles in rows:
writer.writerow({"chembl_id": molecule_id, "canonical_smiles": smiles})
def read_block_rows(path: Path) -> MoleculeBlock:
rows: MoleculeBlock = []
with path.open("r", encoding="utf-8", newline="") as source:
reader = csv.DictReader(source, delimiter="\t")
for row in reader:
molecule_id = row.get("chembl_id", "")
smiles = row.get("canonical_smiles", "")
if not molecule_id or not smiles:
raise ValueError("block artifact contains an invalid row")
rows.append((molecule_id, smiles))
return rows
def compute_block_edges(
left: MoleculeBlock,
right: MoleculeBlock,
threshold: float,
threshold_direction: str,
) -> list[tuple[str, str, float]]:
"""Compare one planned block pair and emit only thresholded edges."""
if not 0.0 <= threshold <= 1.0:
raise ValueError("threshold must be between 0 and 1")
if threshold_direction not in {"greater", "less"}:
raise ValueError("threshold_direction must be 'greater' or 'less'")
left_fingerprints = [
(molecule_id, fingerprint(Chem.MolFromSmiles(smiles)))
for molecule_id, smiles in left
]
right_fingerprints = [
(molecule_id, fingerprint(Chem.MolFromSmiles(smiles)))
for molecule_id, smiles in right
]
diagonal = left is right or left == right
edges: list[tuple[str, str, float]] = []
for left_index, (left_id, left_fp) in enumerate(left_fingerprints):
right_start = left_index + 1 if diagonal else 0
for right_index in range(right_start, len(right_fingerprints)):
right_id, right_fp = right_fingerprints[right_index]
similarity = DataStructs.TanimotoSimilarity(left_fp, right_fp)
matches_threshold = (
similarity >= threshold
if threshold_direction == "greater"
else similarity <= threshold
)
if matches_threshold:
edges.append((left_id, right_id, similarity))
return edges
def write_edge_csv(output_path: Path, edges: Iterable[tuple[str, str, float]]) -> None:
"""Write an edge table CSV with six-decimal similarity values.
Uses the CSV module's default ``\\r\\n`` line terminator so the bytes match
the local ``write_graph_edges`` reference exactly.
"""
output_path.parent.mkdir(parents=True, exist_ok=True)
with output_path.open("w", encoding="utf-8", newline="") as destination:
writer = csv.DictWriter(destination, fieldnames=list(EDGE_COLUMNS))
writer.writeheader()
for source_id, target_id, similarity in edges:
writer.writerow(
{
"source_id": source_id,
"target_id": target_id,
"similarity": f"{similarity:.6f}",
}
)
def read_edge_csv(path: Path) -> list[tuple[str, str, float]]:
"""Read a materialized edge CSV with strict row validation."""
if not path.is_file():
raise ValueError("materialized edge partial is missing")
edges: list[tuple[str, str, float]] = []
with path.open("r", encoding="utf-8", newline="") as source:
reader = csv.DictReader(source)
if tuple(reader.fieldnames or ()) != EDGE_COLUMNS:
raise ValueError("edge partial has an invalid CSV header")
for row in reader:
if set(row) != set(EDGE_COLUMNS):
raise ValueError("edge partial has an invalid row")
source_id = row["source_id"]
target_id = row["target_id"]
if not source_id or not target_id:
raise ValueError("edge partial contains an empty molecule id")
try:
similarity = float(row["similarity"])
except (TypeError, ValueError) as error:
raise ValueError("edge partial has an invalid similarity") from error
if not 0 <= similarity <= 1:
raise ValueError("edge partial has an invalid similarity")
edges.append((source_id, target_id, similarity))
return edges
def block_pair_from_key(key: str) -> tuple[int, int]:
"""Parse ``map.<i>x<j>`` partial keys into block indices."""
prefix = "map."
if not key.startswith(prefix):
raise ValueError("graph partial key must use map.<left>x<right>")
raw = key[len(prefix) :]
left_raw, separator, right_raw = raw.partition("x")
if not separator or not left_raw.isdigit() or not right_raw.isdigit():
raise ValueError("graph partial key must use map.<left>x<right>")
return int(left_raw), int(right_raw)
def check_pair_coverage(pairs: Sequence[tuple[int, int]]) -> None:
"""Enforce the pair-coverage invariant: every block pair exactly once.
``pairs`` must contain every ``(i, j)`` with ``0 <= i <= j < n`` exactly
once, where ``n`` is derived from the largest referenced block index.
"""
if not pairs:
raise ValueError("graph partial keys cover no block pairs")
unique = set(pairs)
if len(unique) != len(pairs):
raise ValueError("graph partial keys contain a duplicate block pair")
if any(left > right or left < 0 or right < 0 for left, right in unique):
raise ValueError("graph partial keys reference an invalid block pair")
n = max(right for _, right in unique) + 1
expected = {(left, right) for left in range(n) for right in range(left, n)}
missing = sorted(expected - unique)
if missing:
raise ValueError(
"graph partial keys do not cover the full block pair set: "
+ ", ".join(f"{left}x{right}" for left, right in missing)
)
unexpected = sorted(unique - expected)
if unexpected:
raise ValueError(
"graph partial keys cover pairs outside the block pair set: "
+ ", ".join(f"{left}x{right}" for left, right in unexpected)
)
def merge_edge_partials(
partial_paths: Sequence[Path],
output_path: Path,
) -> dict[str, int]:
"""Merge edge partials with duplicate detection and deterministic sort.
The merged edge list is sorted by ``(source_id, target_id, -similarity)``,
matching the local brute-force reference exactly.
"""
if not partial_paths:
raise ValueError("graph reducer requires at least one edge partial")
edges: list[tuple[str, str, float]] = []
seen_pairs: set[tuple[str, str]] = set()
for path in partial_paths:
for source_id, target_id, similarity in read_edge_csv(path):
unordered = (min(source_id, target_id), max(source_id, target_id))
if unordered in seen_pairs:
raise ValueError(
f"graph partials contain a duplicate unordered pair: {unordered[0]}, {unordered[1]}"
)
seen_pairs.add(unordered)
edges.append((source_id, target_id, similarity))
edges.sort(key=lambda edge: (edge[0], edge[1], -edge[2]))
write_edge_csv(output_path, edges)
return {"partial_count": len(partial_paths), "edges_emitted": len(edges)}
+500
View File
@@ -0,0 +1,500 @@
"""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.
"""
from __future__ import annotations
import hashlib
import shutil
from pathlib import Path
from typing import Any, Mapping
from ...sdk.artifacts import (
ArtifactCollection,
ArtifactItem,
ArtifactRef,
ArtifactSchema,
Cardinality,
CollectionKind,
OutputManifest,
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 ..environment import current_environment_digest, current_scimesh_package_digest
from .core import (
block_pair_from_key,
check_pair_coverage,
compute_block_edges,
merge_edge_partials,
parse_molecule_blocks,
read_block_rows,
write_block_tsv,
write_edge_csv,
)
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]:
return {
"type": "object",
"additionalProperties": False,
"properties": {
"threshold": {"type": "number", "minimum": 0, "maximum": 1},
"threshold_direction": {"enum": ["greater", "less"]},
"block_size": {"type": "integer", "minimum": 1},
"max_rows": {"type": "integer", "minimum": 1},
},
"required": ["threshold"],
}
def _molecule_schema() -> ArtifactSchema:
return ArtifactSchema(
SchemaRef("molecule-table", 1),
"text/tab-separated-values",
"utf-8",
max_bytes=10 * 1024 * 1024 * 1024,
validator=ComponentRef("delimited-table", 1),
validator_configuration={
"required_columns": ["canonical_smiles", "chembl_id"],
},
max_records=100_000_000,
canonicalizer="scimesh-tsv-v1",
)
def _edge_schema() -> ArtifactSchema:
return ArtifactSchema(
SchemaRef("similarity-edge-table", 1),
"text/csv",
"utf-8",
max_bytes=100 * 1024 * 1024 * 1024,
validator=ComponentRef("delimited-table", 1),
validator_configuration={
"columns": ["source_id", "target_id", "similarity"],
},
max_records=1_000_000_000,
canonicalizer="similarity-edge-table-v1",
)
class SimilarityGraphSDKWorkload:
"""Manifest-backed planner, runner, and reducer for similarity-graph."""
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()
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
unknown = set(parameters) - {
"threshold",
"threshold_direction",
"block_size",
"max_rows",
}
if unknown:
raise ValueError(
"unsupported similarity-graph parameters: " + ", ".join(sorted(unknown))
)
threshold = parameters.get("threshold")
if threshold is None:
raise ValueError("threshold is required")
self._unit_interval(threshold, "threshold")
if "threshold_direction" in parameters and parameters[
"threshold_direction"
] not in {"greater", "less"}:
raise ValueError("threshold_direction must be 'greater' or 'less'")
if "block_size" in parameters:
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'")
block_size = int(parameters.get("block_size", 1_000))
max_rows = parameters.get("max_rows")
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] = []
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,
)
)
tasks: list[TaskSpec] = []
for left in range(len(blocks)):
for right in range(left, len(blocks)):
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={
**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),
)
def run(self, context: TaskContext) -> OutputManifest:
context.cancellation.raise_if_cancelled()
parameters = context.task.parameters
left_block = parameters.get("left_block")
right_block = parameters.get("right_block")
if (
isinstance(left_block, bool)
or not isinstance(left_block, int)
or isinstance(right_block, bool)
or not isinstance(right_block, int)
):
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'")
checked_pairs = (
len(left_rows) * (len(left_rows) - 1) // 2
if diagonal
else len(left_rows) * len(right_rows)
)
edges = compute_block_edges(
left_rows,
right_rows,
threshold,
direction,
)
output_path = workspace / "result.csv"
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,
)
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 similarity_graph_sdk_definition(
*,
package_digest: str | None = None,
environment_digest: str | None = None,
) -> SimilarityGraphSDKWorkload:
"""Build the SDK-based similarity-graph definition for tests."""
return SimilarityGraphSDKWorkload(
package_digest=package_digest or current_scimesh_package_digest(),
environment_digest=environment_digest or current_environment_digest(),
)
def workload_definition() -> WorkloadDefinition:
"""Installed entry-point factory for the SDK-based similarity-graph."""
return similarity_graph_sdk_definition().definition()
+68
View File
@@ -0,0 +1,68 @@
"""Built-in workload library: default registry and runtime.
This module is workload code, not SDK framework code. It composes the
installed SciMesh workload packages (``search``, ``graph``, ``descriptors``)
with the SDK registry and runtime. External workload libraries can follow the
same pattern with their own packages and entry points.
"""
from __future__ import annotations
import os
import platform
from scimesh.sdk.identity import SDK_API_VERSION
from scimesh.sdk.registry import WorkloadRegistry
from scimesh.sdk.resources import ResourceInventory
from scimesh.sdk.runtime import RuntimeCapabilities
from .descriptors import descriptor_batch_sdk_definition
from .environment import current_environment_digest
from .graph import similarity_graph_sdk_definition
from .search import similarity_search_sdk_definition
__all__ = [
"default_sdk_registry",
"default_sdk_runtime",
]
def default_sdk_registry(*, shard_rows: int = 10_000) -> WorkloadRegistry:
"""Registry of every built-in SDK-built workload, all enabled."""
registry = WorkloadRegistry()
registry.register(
similarity_search_sdk_definition(shard_rows=shard_rows).definition(),
enabled=True,
)
registry.register(
similarity_graph_sdk_definition().definition(),
enabled=True,
)
registry.register(
descriptor_batch_sdk_definition(shard_rows=shard_rows).definition(),
enabled=True,
)
return registry
def default_sdk_runtime() -> RuntimeCapabilities:
"""Runtime advertising the built-in workloads' capabilities and inventory."""
architecture = platform.machine().lower() or "unknown"
return RuntimeCapabilities(
sdk_api_version=SDK_API_VERSION,
protocol_version="1.0.0",
profiles=("core-batch-v1",),
features={"artifact-collections": "1.0.0", "exact-verifier": "1.0.0"},
workload_capabilities=(
"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(),),
),
)
+31
View File
@@ -0,0 +1,31 @@
"""SDK-built ``similarity-search`` reference workload.
See ``core.py`` for the sharding/merge scientific core and ``definition.py``
for the manifest-backed planner/runner/reducer handlers.
"""
from .core import (
merge_search_partials,
run_search_shard,
write_search_partial,
write_search_shards,
)
from .definition import (
MAP_ENTRY_POINT,
REDUCE_ENTRY_POINT,
SimilaritySearchSDKWorkload,
similarity_search_sdk_definition,
workload_definition,
)
__all__ = [
"MAP_ENTRY_POINT",
"REDUCE_ENTRY_POINT",
"SimilaritySearchSDKWorkload",
"merge_search_partials",
"run_search_shard",
"similarity_search_sdk_definition",
"workload_definition",
"write_search_partial",
"write_search_shards",
]
+261
View File
@@ -0,0 +1,261 @@
"""Scientific core for the SDK-built ``similarity-search`` workload.
Reuses the local reference implementation (``search_similar``) and the shared
CTX-08 partial format (full-precision ``repr`` scores) so that shard outputs
and the merged final CSV are byte-identical to the legacy distributed path and
to the single-process reference.
"""
from __future__ import annotations
import csv
import heapq
from pathlib import Path
from typing import Any, Iterator, Mapping, Sequence
from scimesh.chemistry.dataset import MoleculeRecord, parse_smiles
from scimesh.workloads.similarity_search import (
SimilarityMatch,
_HeapEntry,
search_similar,
write_search_results,
)
SEARCH_COLUMNS = ("rank", "chembl_id", "canonical_smiles", "similarity")
REQUIRED_COLUMNS = {"chembl_id", "canonical_smiles"}
def write_search_partial(
output_path: Path,
matches: Sequence[SimilarityMatch],
) -> None:
"""Write a worker partial with a round-trip score, not display rounding.
The public final CSV continues to use the local CLI's six-decimal display.
A reducer needs the full binary float representation to rank candidates
from separate shards exactly as the single-process reference does.
"""
output_path.parent.mkdir(parents=True, exist_ok=True)
with output_path.open("w", encoding="utf-8", newline="") as destination:
writer = csv.DictWriter(destination, fieldnames=list(SEARCH_COLUMNS))
writer.writeheader()
for rank, match in enumerate(matches, start=1):
writer.writerow(
{
"rank": rank,
"chembl_id": match.molecule_id,
"canonical_smiles": match.smiles,
"similarity": repr(match.similarity),
}
)
def run_search_shard(
input_path: Path,
parameters: Mapping[str, object],
output_path: Path,
) -> 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``).
"""
allowed = {
"query_smiles",
"top_k",
"threshold",
"threshold_direction",
"progress_every",
}
unknown = set(parameters) - allowed
if unknown:
raise ValueError(
f"unsupported similarity-search parameters: {', '.join(sorted(unknown))}"
)
query_smiles = parameters.get("query_smiles")
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)
if molecule is None:
raise ValueError("query_smiles is invalid")
top_k = _positive_int(parameters.get("top_k", 20), "top_k")
threshold = None
if "threshold" in parameters:
threshold = _unit_interval(parameters["threshold"], "threshold")
direction = parameters.get("threshold_direction", "greater")
if direction not in {"greater", "less"}:
raise ValueError("threshold_direction must be 'greater' or 'less'")
assert isinstance(direction, str)
progress_every = 0
if "progress_every" in parameters:
progress_every = _nonnegative_int(
parameters["progress_every"], "progress_every"
)
result = search_similar(
input_path,
MoleculeRecord("query", query_smiles, molecule),
top_k=top_k,
progress_every=progress_every,
threshold=threshold,
threshold_direction=direction,
)
write_search_partial(output_path, result.matches)
return {
"scanned_rows": result.stats.scanned,
"valid_molecules": result.stats.valid,
"invalid_smiles": result.stats.invalid,
"matches_emitted": len(result.matches),
}
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 _nonnegative_int(value: object, name: str) -> int:
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
raise ValueError(f"{name} must be a non-negative integer")
return value
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)
def write_search_shards(
input_path: Path,
workspace: Path,
shard_rows: int,
max_rows: int | None = None,
) -> list[Path]:
"""Split the input TSV into deterministic row-bounded shards with headers."""
if (
isinstance(shard_rows, bool)
or not isinstance(shard_rows, int)
or shard_rows < 1
):
raise ValueError("shard_rows must be a positive integer")
if max_rows is not None and (
isinstance(max_rows, bool) or not isinstance(max_rows, int) or max_rows < 1
):
raise ValueError("max_rows must be a positive integer")
paths: list[Path] = []
current: Path | None = None
destination = None
writer = None
rows_in_shard = 0
seen_rows = 0
try:
with input_path.open("r", encoding="utf-8", newline="") as source:
reader = csv.DictReader(source, delimiter="\t")
fieldnames = tuple(reader.fieldnames or ())
if not REQUIRED_COLUMNS.issubset(set(fieldnames)):
missing = sorted(REQUIRED_COLUMNS - set(fieldnames))
raise ValueError(
f"dataset is missing required columns: {', '.join(missing)}"
)
for row in reader:
if max_rows is not None and seen_rows >= max_rows:
break
if destination is None or rows_in_shard == shard_rows:
if destination is not None:
destination.close()
current = workspace / f"shard-{len(paths)}.tsv"
destination = current.open("w", encoding="utf-8", newline="")
writer = csv.DictWriter(
destination,
fieldnames=list(fieldnames),
delimiter="\t",
lineterminator="\n",
)
writer.writeheader()
paths.append(current)
rows_in_shard = 0
assert writer is not None
writer.writerow(row)
rows_in_shard += 1
seen_rows += 1
finally:
if destination is not None:
destination.close()
if not paths:
raise ValueError("dataset has no data rows")
return paths
def iter_search_partial(
path: Path,
threshold_direction: str,
) -> Iterator[SimilarityMatch]:
"""Yield strictly ordered partial matches with full-precision scores."""
if not path.is_file():
raise ValueError("materialized partial result is missing")
if threshold_direction not in {"greater", "less"}:
raise ValueError("threshold_direction must be 'greater' or 'less'")
previous_key: tuple[float, str, str] | None = None
with path.open("r", encoding="utf-8", newline="") as source:
reader = csv.DictReader(source)
if tuple(reader.fieldnames or ()) != SEARCH_COLUMNS:
raise ValueError("partial result has an invalid CSV header")
for expected_rank, row in enumerate(reader, start=1):
if set(row) != set(SEARCH_COLUMNS) or row["rank"] != str(expected_rank):
raise ValueError("partial result has an invalid rank")
try:
similarity = float(row["similarity"])
except (TypeError, ValueError) as error:
raise ValueError("partial result has an invalid similarity") from error
if not 0 <= similarity <= 1:
raise ValueError("partial result has an invalid similarity")
match = SimilarityMatch(
similarity, row["chembl_id"], row["canonical_smiles"]
)
key = match.sort_key(threshold_direction)
if previous_key is not None and key < previous_key:
raise ValueError("partial result is not sorted deterministically")
previous_key = key
yield match
def merge_search_partials(
partial_paths: Sequence[Path],
parameters: Mapping[str, Any],
output_path: Path,
) -> dict[str, int]:
"""Merge sorted shard partials into one deterministic final top-k CSV.
Mirrors the CTX-08/CTX-09 reducer: a bounded heap with the local
tie-breaker, so the merged file equals the single-process reference
byte-for-byte for the same input and options.
"""
if not partial_paths:
raise ValueError("at least one partial result is required")
raw_top_k = parameters.get("top_k", 20)
if isinstance(raw_top_k, bool) or not isinstance(raw_top_k, int) or raw_top_k < 1:
raise ValueError("top_k must be a positive integer")
top_k = raw_top_k
direction = parameters.get("threshold_direction", "greater")
if direction not in {"greater", "less"}:
raise ValueError("threshold_direction must be 'greater' or 'less'")
heap: list[_HeapEntry] = []
for path in partial_paths:
for match in iter_search_partial(path, direction):
rank_key = match.sort_key(direction)
entry = _HeapEntry(match, rank_key)
if len(heap) < top_k:
heapq.heappush(heap, entry)
elif rank_key < heap[0].rank_key:
heapq.heapreplace(heap, entry)
matches = sorted(
(entry.match for entry in heap),
key=lambda match: match.sort_key(direction),
)
write_search_results(output_path, matches)
return {"matches_emitted": len(matches), "partial_count": len(partial_paths)}
+565
View File
@@ -0,0 +1,565 @@
"""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.
"""
from __future__ import annotations
import hashlib
import shutil
from pathlib import Path
from typing import Any, Mapping
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 ...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_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]:
return {
"type": "object",
"additionalProperties": False,
"properties": {
"query_id": {"type": "string", "minLength": 1, "maxLength": 200},
"query_smiles": {"type": "string", "minLength": 1, "maxLength": 200},
"top_k": {"type": "integer", "minimum": 1},
"threshold": {"type": "number", "minimum": 0, "maximum": 1},
"threshold_direction": {"enum": ["greater", "less"]},
"max_rows": {"type": "integer", "minimum": 1},
"progress_every": {"type": "integer", "minimum": 0},
},
"oneOf": [
{"required": ["query_id"], "not": {"required": ["query_smiles"]}},
{"required": ["query_smiles"], "not": {"required": ["query_id"]}},
],
}
def _dataset_schema() -> ArtifactSchema:
return ArtifactSchema(
SchemaRef("molecule-table", 1),
"text/tab-separated-values",
"utf-8",
max_bytes=10 * 1024 * 1024 * 1024,
validator=ComponentRef("delimited-table", 1),
validator_configuration={
"required_columns": ["canonical_smiles", "chembl_id"],
},
max_records=100_000_000,
canonicalizer="scimesh-tsv-v1",
)
def _search_table_schema(ref: SchemaRef, canonicalizer: str) -> ArtifactSchema:
return ArtifactSchema(
ref,
"text/csv",
"utf-8",
max_bytes=1024 * 1024 * 1024,
validator=ComponentRef("delimited-table", 1),
validator_configuration={
"columns": ["rank", "chembl_id", "canonical_smiles", "similarity"],
},
max_records=100_000,
canonicalizer=canonicalizer,
)
class SimilaritySearchSDKWorkload:
"""Manifest-backed planner, runner, and reducer for similarity-search."""
def __init__(
self,
*,
shard_rows: int,
package_digest: str,
environment_digest: str,
) -> None:
if (
isinstance(shard_rows, bool)
or not isinstance(shard_rows, int)
or shard_rows < 1
):
raise ValueError("shard_rows must be a positive integer")
self.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"
)
)
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},
)
@staticmethod
def _string(value: object, name: str) -> str:
if not isinstance(value, str) or not value.strip() or len(value) > 200:
raise ValueError(f"{name} must be a non-empty string")
return 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
@staticmethod
def _nonnegative_int(value: object, name: str) -> int:
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
raise ValueError(f"{name} must be a non-negative integer")
return value
@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)
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
unknown = set(parameters) - {
"query_id",
"query_smiles",
"top_k",
"threshold",
"threshold_direction",
"max_rows",
"progress_every",
}
if unknown:
raise ValueError(
"unsupported similarity-search parameters: "
+ ", ".join(sorted(unknown))
)
query_id = parameters.get("query_id")
query_smiles = parameters.get("query_smiles")
if (query_id is None) == (query_smiles is None):
raise ValueError("exactly one of query_id or query_smiles is required")
if query_id is not None:
self._string(query_id, "query_id")
if query_smiles is not None:
self._string(query_smiles, "query_smiles")
self._positive_int(parameters.get("top_k", 20), "top_k")
if "max_rows" in parameters:
self._positive_int(parameters["max_rows"], "max_rows")
if "progress_every" in parameters:
self._nonnegative_int(parameters["progress_every"], "progress_every")
if "threshold" in parameters:
self._unit_interval(parameters["threshold"], "threshold")
if "threshold_direction" in parameters and parameters[
"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]:
parameters = request.parameters
query_id = parameters.get("query_id")
if isinstance(query_id, str):
query_source: dict[str, str] = {"kind": "chembl_id", "value": query_id}
else:
query_source = {
"kind": "smiles",
"value": self._string(parameters.get("query_smiles"), "query_smiles"),
}
resolved: dict[str, object] = {
"query_source": query_source,
"top_k": self._positive_int(parameters.get("top_k", 20), "top_k"),
"threshold_direction": parameters.get("threshold_direction", "greater"),
"fingerprint": {
"algorithm": "morgan",
"radius": FP_RADIUS,
"fp_size": FP_SIZE,
},
}
if "threshold" in parameters:
resolved["threshold"] = self._unit_interval(
parameters["threshold"], "threshold"
)
if "max_rows" in parameters:
resolved["max_rows"] = self._positive_int(
parameters["max_rows"], "max_rows"
)
if "progress_every" in parameters:
resolved["progress_every"] = self._nonnegative_int(
parameters["progress_every"], "progress_every"
)
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(
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 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 similarity_search_sdk_definition(
*,
shard_rows: int = 10_000,
package_digest: str | None = None,
environment_digest: str | None = None,
) -> SimilaritySearchSDKWorkload:
"""Build the SDK-built similarity-search definition for tests."""
return SimilaritySearchSDKWorkload(
shard_rows=shard_rows,
package_digest=package_digest or current_scimesh_package_digest(),
environment_digest=environment_digest or current_environment_digest(),
)
def workload_definition() -> WorkloadDefinition:
"""Installed entry-point factory for the SDK-built similarity-search."""
return similarity_search_sdk_definition().definition()