-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdistributed_execution.py
More file actions
143 lines (116 loc) · 5.71 KB
/
Copy pathdistributed_execution.py
File metadata and controls
143 lines (116 loc) · 5.71 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
"""Deterministic parallel execution primitives for QueueCraft Monte Carlo workloads.
This module partitions independent replication indices into reproducible shards,
executes them through a bounded local worker pool, emits progress snapshots, and
supports checkpoint/resume without external infrastructure side effects.
"""
from __future__ import annotations
from concurrent.futures import Future, ThreadPoolExecutor, as_completed
from dataclasses import asdict, dataclass
import hashlib
import json
from pathlib import Path
import threading
from typing import Any, Callable, Iterable, Mapping
def canonical_json(value: Any) -> str:
return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
def fingerprint(value: Any) -> str:
return hashlib.sha256(canonical_json(value).encode("utf-8")).hexdigest()
@dataclass(frozen=True)
class DistributedPlan:
run_id: str
total_tasks: int
worker_count: int = 2
chunk_size: int = 10
seed: int = 42
def validate(self) -> None:
if not self.run_id or self.total_tasks < 1 or self.worker_count < 1 or self.chunk_size < 1:
raise ValueError("run_id, task count, worker count and chunk size must be positive")
@dataclass(frozen=True)
class Checkpoint:
run_id: str
plan_fingerprint: str
completed_tasks: tuple[int, ...]
results: tuple[tuple[int, Any], ...]
def to_dict(self) -> dict[str, Any]:
return asdict(self)
class CheckpointStore:
def __init__(self, path: str | Path) -> None:
self.path = Path(path)
self.path.parent.mkdir(parents=True, exist_ok=True)
self._lock = threading.Lock()
def save(self, checkpoint: Checkpoint) -> None:
payload = checkpoint.to_dict()
with self._lock:
self.path.write_text(json.dumps(payload, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
def load(self) -> Checkpoint | None:
if not self.path.exists():
return None
payload = json.loads(self.path.read_text(encoding="utf-8"))
return Checkpoint(
run_id=str(payload["run_id"]),
plan_fingerprint=str(payload["plan_fingerprint"]),
completed_tasks=tuple(int(item) for item in payload.get("completed_tasks", [])),
results=tuple((int(item[0]), item[1]) for item in payload.get("results", [])),
)
class DistributedExecutor:
"""Bounded local executor with deterministic task identity and resumability."""
def __init__(self, plan: DistributedPlan, checkpoint_store: CheckpointStore | None = None) -> None:
plan.validate()
self.plan = plan
self.store = checkpoint_store
self.plan_fingerprint = fingerprint(asdict(plan))
self._cancel = threading.Event()
def cancel(self) -> None:
self._cancel.set()
def _shards(self, task_ids: Iterable[int]) -> list[list[int]]:
tasks = list(task_ids)
return [tasks[index:index + self.plan.chunk_size] for index in range(0, len(tasks), self.plan.chunk_size)]
def run(
self,
worker: Callable[[int, int], Any],
*,
progress: Callable[[dict[str, Any]], None] | None = None,
resume: bool = True,
) -> dict[str, Any]:
task_ids = list(range(self.plan.total_tasks))
completed: set[int] = set()
results: dict[int, Any] = {}
if resume and self.store:
checkpoint = self.store.load()
if checkpoint:
if checkpoint.run_id != self.plan.run_id or checkpoint.plan_fingerprint != self.plan_fingerprint:
raise ValueError("checkpoint does not match execution plan")
completed.update(checkpoint.completed_tasks)
results.update(dict(checkpoint.results))
pending = [task_id for task_id in task_ids if task_id not in completed]
total = len(task_ids)
if progress:
progress({"status": "started", "completed": len(completed), "total": total, "progress": len(completed) / total})
if not pending:
return {"status": "completed", "run_id": self.plan.run_id, "completed": total, "total": total, "results": results, "resumed": True}
shards = self._shards(pending)
def run_shard(shard: list[int]) -> list[tuple[int, Any]]:
if self._cancel.is_set():
return []
output: list[tuple[int, Any]] = []
for task_id in shard:
if self._cancel.is_set():
break
seed = self.plan.seed + task_id
output.append((task_id, worker(task_id, seed)))
return output
with ThreadPoolExecutor(max_workers=self.plan.worker_count) as executor:
futures: list[Future[list[tuple[int, Any]]]] = [executor.submit(run_shard, shard) for shard in shards]
for future in as_completed(futures):
for task_id, value in future.result():
results[task_id] = value
completed.add(task_id)
if self.store:
checkpoint = Checkpoint(self.plan.run_id, self.plan_fingerprint, tuple(sorted(completed)), tuple(sorted(results.items())))
self.store.save(checkpoint)
if progress:
progress({"status": "cancelled" if self._cancel.is_set() else "running", "completed": len(completed), "total": total, "progress": len(completed) / total})
if self._cancel.is_set():
break
status = "cancelled" if self._cancel.is_set() and len(completed) < total else "completed"
return {"status": status, "run_id": self.plan.run_id, "completed": len(completed), "total": total, "results": results, "resumed": bool(self.store and self.store.load())}