Add descriptor-batch SDK reference workload
This commit is contained in:
@@ -48,7 +48,9 @@ def current_environment_digest() -> str:
|
||||
return "sha256:" + hashlib.sha256(payload.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def similarity_search_sdk_adapter(*, shard_rows: int = 10_000) -> LegacyDistributedWorkloadAdapter:
|
||||
def similarity_search_sdk_adapter(
|
||||
*, shard_rows: int = 10_000
|
||||
) -> LegacyDistributedWorkloadAdapter:
|
||||
dataset_schema = ArtifactSchema(
|
||||
SchemaRef("molecule-table", 1),
|
||||
"text/tab-separated-values",
|
||||
@@ -119,7 +121,9 @@ def similarity_search_sdk_adapter(*, shard_rows: int = 10_000) -> LegacyDistribu
|
||||
|
||||
def default_sdk_registry(*, shard_rows: int = 10_000) -> WorkloadRegistry:
|
||||
registry = WorkloadRegistry()
|
||||
registry.register(similarity_search_sdk_adapter(shard_rows=shard_rows).definition(), enabled=True)
|
||||
registry.register(
|
||||
similarity_search_sdk_adapter(shard_rows=shard_rows).definition(), enabled=True
|
||||
)
|
||||
return registry
|
||||
|
||||
|
||||
@@ -135,7 +139,7 @@ def default_sdk_runtime() -> RuntimeCapabilities:
|
||||
protocol_version="1.0.0",
|
||||
profiles=("core-batch-v1",),
|
||||
features={"artifact-collections": "1.0.0", "exact-verifier": "1.0.0"},
|
||||
workload_capabilities=("similarity-search",),
|
||||
workload_capabilities=("similarity-search", "descriptor-batch"),
|
||||
inventory=ResourceInventory(
|
||||
cpu_cores=max(os.cpu_count() or 1, 1),
|
||||
memory_mb=4096,
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
"""SDK-native ``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-native ``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 ..artifacts import (
|
||||
ArtifactCollection,
|
||||
ArtifactItem,
|
||||
ArtifactRef,
|
||||
ArtifactSchema,
|
||||
Cardinality,
|
||||
CollectionKind,
|
||||
OutputManifest,
|
||||
PortSpec,
|
||||
)
|
||||
from ..builtins import current_environment_digest, current_scimesh_package_digest
|
||||
from ..execution import (
|
||||
CheckpointPolicy,
|
||||
ExecutionProfile,
|
||||
NetworkPolicy,
|
||||
RetryPolicy,
|
||||
)
|
||||
from ..identity import ComponentRef, SchemaRef, VersionRange, WorkloadId
|
||||
from ..manifest import (
|
||||
DeterminismProfile,
|
||||
EnvironmentSpec,
|
||||
PackageSpec,
|
||||
TrustMode,
|
||||
VerifierSpec,
|
||||
WorkloadLimits,
|
||||
WorkloadManifest,
|
||||
)
|
||||
from ..plans import JobRequest, TaskSpec, ValidatedJob, WorkflowPlan
|
||||
from ..protocols import PlanningContext, ReduceContext, TaskContext
|
||||
from ..registry import WorkloadDefinition
|
||||
from ..resources import ResourceRequirements
|
||||
from ..verification import ExactArtifactVerifier
|
||||
from ..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.sdk.descriptors.definition:map_descriptors@v1"
|
||||
REDUCE_ENTRY_POINT = "scimesh.sdk.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-native:
|
||||
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()
|
||||
Reference in New Issue
Block a user