tostido's picture
Build MeshScale CPU worker template
96ef23c verified
Raw History Blame Contribute Delete
6.54 kB
from __future__ import annotations
import argparse
import os
import time
from .contracts import (
ExpansionRequest,
FarmPolicy,
RegulationIntent,
ShardResult,
WorkerPulse,
utc_now,
)
from .evaluator import evaluate_job
from .regulation import create_pulse, tick_pulse
from .store import FarmStore
def maybe_request_expansion(
store: FarmStore,
job_run_id: str,
generation: int,
policy: FarmPolicy,
reason: str,
formation: str,
) -> None:
pending = len(store.pending_jobs(formation))
active = store.active_worker_count(policy.pulse_ttl_seconds, formation)
allowed = policy.effective_max_workers()
if active >= allowed:
return
if pending <= max(1, active * 2):
return
requested = min(allowed - active, max(1, pending // 4))
request = ExpansionRequest(
run_id=job_run_id,
generation=generation,
reason=reason,
requested_workers=requested,
active_workers=active,
pending_jobs=pending,
policy=policy.to_dict(),
formation=formation,
)
store.write_expansion_request(request)
def primary_formation(policy: FarmPolicy) -> str:
formations = tuple(policy.allowed_formations or ("cpu-worker",))
return formations[0] if formations else "cpu-worker"
def worker_tick(
store: FarmStore,
worker_id: str,
policy: FarmPolicy,
pulse_state: dict[str, float],
) -> bool:
formation = primary_formation(policy)
claimed = store.claim_next(worker_id, tuple(policy.allowed_formations))
if claimed is None:
store.write_worker_pulse(
WorkerPulse(worker_id, None, None, "idle", formation=formation, pulse=pulse_state)
)
return False
job, claimed_path = claimed
store.write_worker_pulse(
WorkerPulse(
worker_id,
job.run_id,
job.generation,
"claimed",
job.shard_id,
formation=formation,
pulse=pulse_state,
)
)
children = job.split(policy)
if len(children) > 1:
store.put_jobs(children)
result = ShardResult(
run_id=job.run_id,
generation=job.generation,
shard_id=job.shard_id,
worker_id=worker_id,
status="refracted",
node_results=[],
started_at=utc_now(),
message=f"split into {len(children)} child shards",
)
store.complete_job(job, result, claimed_path)
maybe_request_expansion(
store, job.run_id, job.generation, policy, "job_refraction", job.formation
)
return True
result = evaluate_job(job, worker_id)
store.complete_job(job, result, claimed_path)
maybe_request_expansion(store, job.run_id, job.generation, policy, "pending_backlog", job.formation)
return True
def _drain_applies_to_worker(
store: FarmStore,
worker_id: str,
control: RegulationIntent,
policy: FarmPolicy,
worker_formation: str,
) -> bool:
if control.target_workers is None:
return False
if control.target_workers <= 0:
return True
active = store.active_worker_ids(policy.pulse_ttl_seconds, worker_formation)
if len(active) <= control.target_workers:
return False
keepers = set(active[: control.target_workers])
return worker_id not in keepers
def worker_exit_reason(store: FarmStore, worker_id: str, policy: FarmPolicy) -> str | None:
control = store.read_regulation()
worker_formation = primary_formation(policy)
if control.formation is not None and control.formation != worker_formation:
return None
if control.mode == "shutdown":
return f"shutdown:{control.reason}"
if control.mode == "drain" and _drain_applies_to_worker(
store, worker_id, control, policy, worker_formation
):
return f"drain:{control.reason}"
if policy.exit_when_pulse_lost and not store.controller_pulse_alive(
policy.controller_pulse_ttl_seconds
):
return "controller_pulse_lost"
return None
def worker_loop(
store: FarmStore,
worker_id: str,
policy: FarmPolicy,
*,
poll_seconds: float = 2.0,
once: bool = False,
) -> None:
pulse = create_pulse(policy.pulse_rate)
formation = primary_formation(policy)
while True:
pulse_state = tick_pulse(pulse)
reason = worker_exit_reason(store, worker_id, policy)
if reason is not None:
store.write_worker_pulse(
WorkerPulse(
worker_id,
None,
None,
f"exiting:{reason}",
formation=formation,
pulse=pulse_state,
)
)
store.append_event("farm_worker_exiting", {"worker_id": worker_id, "reason": reason})
return
store.requeue_expired_claims(policy)
did_work = worker_tick(store, worker_id, policy, pulse_state)
if once:
return
if not did_work:
time.sleep(poll_seconds)
def main() -> None:
parser = argparse.ArgumentParser(description="KEY CPU farm worker")
parser.add_argument("--root", default=os.environ.get("KEY_FARM_ROOT", "/data/farm"))
parser.add_argument("--worker-id", default=os.environ.get("KEY_FARM_WORKER_ID"))
parser.add_argument("--once", action="store_true")
parser.add_argument("--poll-seconds", type=float, default=2.0)
parser.add_argument("--target-batch-size", type=int, default=None)
parser.add_argument("--max-workers", type=int, default=None)
parser.add_argument("--budget-hourly-usd", type=float, default=None)
parser.add_argument("--formation", default=os.environ.get("KEY_FARM_FORMATION", "cpu-worker"))
args = parser.parse_args()
worker_id = args.worker_id or f"worker_{os.getpid()}"
policy_kwargs = {}
if args.target_batch_size is not None:
policy_kwargs["target_batch_size"] = args.target_batch_size
if args.max_workers is not None:
policy_kwargs["max_workers"] = args.max_workers
if args.budget_hourly_usd is not None:
policy_kwargs["budget_hourly_usd"] = args.budget_hourly_usd
policy_kwargs["allowed_formations"] = (args.formation,)
worker_loop(
FarmStore(args.root),
worker_id,
FarmPolicy(**policy_kwargs),
poll_seconds=args.poll_seconds,
once=args.once,
)
if __name__ == "__main__":
main()