Replace legacy distributed protocol with SDK-built workloads
This commit is contained in:
@@ -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",
|
||||
]
|
||||
@@ -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}
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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)}
|
||||
@@ -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()
|
||||
@@ -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(),),
|
||||
),
|
||||
)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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)}
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user