634 lines
25 KiB
Python
634 lines
25 KiB
Python
from __future__ import annotations
|
|
|
|
import base64
|
|
import hashlib
|
|
import hmac
|
|
import json
|
|
import os
|
|
import re
|
|
import secrets
|
|
import shutil
|
|
import uuid
|
|
from datetime import UTC
|
|
from datetime import datetime
|
|
from datetime import timedelta
|
|
from pathlib import Path
|
|
from threading import Lock
|
|
from typing import Any
|
|
|
|
|
|
ALLOWED_TRAINING_ARTIFACTS = {
|
|
"lstm_forecaster.json",
|
|
"torch_retrain_guard.json",
|
|
"torch_threshold_calibration.json",
|
|
}
|
|
RUNNING_LEASE_TIMEOUT = timedelta(minutes=10)
|
|
ONLINE_WINDOW = timedelta(minutes=3)
|
|
MAX_JOB_ATTEMPTS = 3
|
|
MAX_ARTIFACT_CHUNK_BYTES = 1024 * 1024
|
|
# Keep uploads bounded while leaving room for explicitly requested per-symbol bundles.
|
|
MAX_ARTIFACT_BYTES = 256 * 1024 * 1024
|
|
MAX_ARTIFACT_CHUNKS = 1024
|
|
REQUIRED_MODEL_BUNDLE = set(ALLOWED_TRAINING_ARTIFACTS)
|
|
|
|
|
|
class TrainingCoordinator:
|
|
def __init__(self, runtime_dir: Path) -> None:
|
|
self.runtime_dir = runtime_dir
|
|
self.state_path = runtime_dir / "training_coordination.json"
|
|
self.upload_root = runtime_dir / ".training_uploads"
|
|
self._lock = Lock()
|
|
|
|
def status(self) -> dict[str, Any]:
|
|
with self._lock:
|
|
state = self._load_state()
|
|
self._expire_stale_jobs(state)
|
|
self._save_state(state)
|
|
return self._public_status(state)
|
|
|
|
def request_retrain(self, payload: dict[str, Any] | None = None) -> dict[str, Any]:
|
|
payload = payload or {}
|
|
with self._lock:
|
|
state = self._load_state()
|
|
self._expire_stale_jobs(state)
|
|
existing = self._active_job(state)
|
|
if existing is not None:
|
|
self._save_state(state)
|
|
return {
|
|
"queued": False,
|
|
"reason": "active_job_exists",
|
|
"job": self._public_job(existing),
|
|
"status": self._public_status(state),
|
|
}
|
|
|
|
now = _now()
|
|
job = {
|
|
"id": str(uuid.uuid4()),
|
|
"status": "pending",
|
|
"requested_at": now,
|
|
"requested_by": str(payload.get("source") or "api"),
|
|
"parameters": _safe_parameters(payload.get("parameters")),
|
|
"message": "",
|
|
"artifacts": [],
|
|
"attempts": 0,
|
|
}
|
|
state.setdefault("jobs", []).append(job)
|
|
self._trim_jobs(state)
|
|
self._save_state(state)
|
|
return {
|
|
"queued": True,
|
|
"job": self._public_job(job),
|
|
"status": self._public_status(state),
|
|
}
|
|
|
|
def heartbeat(self, payload: dict[str, Any] | None = None) -> dict[str, Any]:
|
|
payload = payload or {}
|
|
with self._lock:
|
|
state = self._load_state()
|
|
worker = self._worker_from_payload(payload)
|
|
state["worker"] = worker
|
|
self._save_state(state)
|
|
return {"ok": True, "worker": worker, "status": self._public_status(state)}
|
|
|
|
def claim(self, payload: dict[str, Any] | None = None) -> dict[str, Any]:
|
|
payload = payload or {}
|
|
with self._lock:
|
|
state = self._load_state()
|
|
self._expire_stale_jobs(state)
|
|
worker = self._worker_from_payload(payload)
|
|
state["worker"] = worker
|
|
job = self._oldest_pending_job(state)
|
|
if job is None:
|
|
self._save_state(state)
|
|
return {"claimed": False, "job": None, "status": self._public_status(state)}
|
|
|
|
now = _now()
|
|
lease_token = secrets.token_urlsafe(32)
|
|
job["status"] = "running"
|
|
job["claimed_at"] = now
|
|
job["updated_at"] = now
|
|
job["claimed_by"] = worker["id"]
|
|
job["worker"] = worker
|
|
job["lease_token"] = lease_token
|
|
job["attempts"] = int(job.get("attempts", 0)) + 1
|
|
self._save_state(state)
|
|
return {
|
|
"claimed": True,
|
|
"job": self._public_job(job),
|
|
"lease_token": lease_token,
|
|
"status": self._public_status(state),
|
|
}
|
|
|
|
def save_artifact_chunk(self, job_id: str, payload: dict[str, Any]) -> dict[str, Any]:
|
|
job_id = _valid_job_id(job_id)
|
|
name = Path(str(payload.get("name") or "")).name
|
|
if name not in ALLOWED_TRAINING_ARTIFACTS:
|
|
raise ValueError(f"artifact is not allowed: {name}")
|
|
index = int(payload.get("index", -1))
|
|
total = int(payload.get("total", 0))
|
|
sha256 = str(payload.get("sha256") or "").strip().lower()
|
|
if index < 0 or total <= 0 or index >= total:
|
|
raise ValueError("invalid artifact chunk index")
|
|
if total > MAX_ARTIFACT_CHUNKS:
|
|
raise ValueError("artifact has too many chunks")
|
|
if not re.fullmatch(r"[0-9a-f]{64}", sha256):
|
|
raise ValueError("artifact sha256 is invalid")
|
|
try:
|
|
chunk = base64.b64decode(str(payload.get("data_base64") or ""), validate=True)
|
|
except (ValueError, TypeError) as exc:
|
|
raise ValueError("invalid artifact chunk payload") from exc
|
|
if not chunk or len(chunk) > MAX_ARTIFACT_CHUNK_BYTES:
|
|
raise ValueError("artifact chunk size is invalid")
|
|
|
|
with self._lock:
|
|
state = self._load_state()
|
|
job = self._job_by_id(state, job_id)
|
|
if job is None:
|
|
raise ValueError(f"training job not found: {job_id}")
|
|
if job.get("status") != "running" or not job.get("claimed_by"):
|
|
raise ValueError("training job is not claimed and running")
|
|
self._require_lease(job, payload)
|
|
uploads = job.setdefault("uploads", {})
|
|
upload = uploads.setdefault(name, {"sha256": sha256, "total": total})
|
|
if upload.get("sha256") != sha256 or int(upload.get("total", 0)) != total:
|
|
raise ValueError("artifact upload metadata changed during upload")
|
|
|
|
chunk_dir = self.upload_root / job_id / "chunks" / name
|
|
chunk_dir.mkdir(parents=True, exist_ok=True)
|
|
(chunk_dir / f"{index:06d}.part").write_bytes(chunk)
|
|
|
|
received = sum(1 for part in range(total) if (chunk_dir / f"{part:06d}.part").is_file())
|
|
if received < total:
|
|
upload["received"] = received
|
|
job["updated_at"] = _now()
|
|
self._save_state(state)
|
|
return {"complete": False, "received": received, "total": total}
|
|
|
|
ready_dir = self.upload_root / job_id / "ready"
|
|
ready_dir.mkdir(parents=True, exist_ok=True)
|
|
target_tmp = ready_dir / f".{name}.tmp"
|
|
digest = hashlib.sha256()
|
|
size = 0
|
|
with target_tmp.open("wb") as output:
|
|
for part in range(total):
|
|
data = (chunk_dir / f"{part:06d}.part").read_bytes()
|
|
size += len(data)
|
|
if size > MAX_ARTIFACT_BYTES:
|
|
target_tmp.unlink(missing_ok=True)
|
|
raise ValueError("artifact exceeds maximum size")
|
|
digest.update(data)
|
|
output.write(data)
|
|
if digest.hexdigest().lower() != sha256:
|
|
target_tmp.unlink(missing_ok=True)
|
|
raise ValueError("artifact sha256 mismatch")
|
|
|
|
target = ready_dir / name
|
|
os.replace(target_tmp, target)
|
|
_remove_tree(chunk_dir)
|
|
|
|
artifacts = job.setdefault("artifacts", [])
|
|
artifacts = [item for item in artifacts if item.get("name") != name]
|
|
artifacts.append(
|
|
{"name": name, "sha256": sha256, "size": size, "staged_at": _now()}
|
|
)
|
|
job["artifacts"] = artifacts
|
|
job["updated_at"] = _now()
|
|
upload["received"] = total
|
|
upload["complete"] = True
|
|
self._save_state(state)
|
|
return {"complete": True, "staged": True, "name": name, "sha256": sha256}
|
|
|
|
def progress(self, job_id: str, payload: dict[str, Any] | None = None) -> dict[str, Any]:
|
|
payload = payload or {}
|
|
job_id = _valid_job_id(job_id)
|
|
with self._lock:
|
|
state = self._load_state()
|
|
job = self._job_by_id(state, job_id)
|
|
if job is None:
|
|
raise ValueError(f"training job not found: {job_id}")
|
|
if job.get("status") != "running" or not job.get("claimed_by"):
|
|
raise ValueError("training job is not claimed and running")
|
|
self._require_lease(job, payload)
|
|
if isinstance(payload.get("worker"), dict):
|
|
state["worker"] = self._worker_from_payload(payload["worker"])
|
|
job["status"] = "running"
|
|
job["phase"] = str(payload.get("phase") or job.get("phase") or "running")[:80]
|
|
job["message"] = str(payload.get("message") or job.get("message") or "")[:2000]
|
|
job["progress_percent"] = _coerce_percent(payload.get("progress_percent"), job.get("progress_percent", 0))
|
|
job["updated_at"] = _now()
|
|
if isinstance(payload.get("details"), dict):
|
|
job["details"] = payload["details"]
|
|
self._save_state(state)
|
|
return {
|
|
"ok": True,
|
|
"job": self._public_job(job),
|
|
"status": self._public_status(state),
|
|
}
|
|
|
|
def complete(self, job_id: str, payload: dict[str, Any] | None = None) -> dict[str, Any]:
|
|
payload = payload or {}
|
|
job_id = _valid_job_id(job_id)
|
|
with self._lock:
|
|
state = self._load_state()
|
|
job = self._job_by_id(state, job_id)
|
|
if job is None:
|
|
raise ValueError(f"training job not found: {job_id}")
|
|
if job.get("status") != "running" or not job.get("claimed_by"):
|
|
raise ValueError("training job is not claimed and running")
|
|
self._require_lease(job, payload)
|
|
success = bool(payload.get("success", payload.get("status") == "completed"))
|
|
if success and job.get("artifacts"):
|
|
promoted = self._validate_and_promote(job_id, job)
|
|
job["promoted_artifacts"] = promoted
|
|
job["status"] = "completed" if success else "failed"
|
|
job["phase"] = "completed" if success else "failed"
|
|
job["progress_percent"] = 100 if success else _coerce_percent(payload.get("progress_percent"), job.get("progress_percent", 0))
|
|
job["completed_at"] = _now()
|
|
job["message"] = str(payload.get("message") or "")
|
|
if isinstance(payload.get("summary"), dict):
|
|
job["summary"] = payload["summary"]
|
|
if isinstance(payload["summary"].get("accepted"), bool):
|
|
job["model_decision"] = (
|
|
"accepted" if payload["summary"]["accepted"] else "rejected"
|
|
)
|
|
job.pop("lease_token", None)
|
|
self._save_state(state)
|
|
return {
|
|
"ok": True,
|
|
"job": self._public_job(job),
|
|
"status": self._public_status(state),
|
|
}
|
|
|
|
def _validate_and_promote(self, job_id: str, job: dict[str, Any]) -> list[dict[str, Any]]:
|
|
ready_dir = self.upload_root / job_id / "ready"
|
|
staged = {path.name for path in ready_dir.iterdir() if path.is_file()} if ready_dir.is_dir() else set()
|
|
missing = REQUIRED_MODEL_BUNDLE - staged
|
|
if missing:
|
|
raise ValueError("training bundle is incomplete: " + ", ".join(sorted(missing)))
|
|
|
|
model = _read_json(ready_dir / "lstm_forecaster.json")
|
|
guard = _read_json(ready_dir / "torch_retrain_guard.json")
|
|
calibration = _read_json(ready_dir / "torch_threshold_calibration.json")
|
|
if model.get("type") != "pytorch_recurrent_forecaster":
|
|
raise ValueError("candidate model type is invalid")
|
|
symbols = model.get("symbols")
|
|
if not isinstance(symbols, dict) or not symbols:
|
|
raise ValueError("candidate model has no symbol models")
|
|
_validate_symbol_models(symbols)
|
|
model_sha256 = hashlib.sha256((ready_dir / "lstm_forecaster.json").read_bytes()).hexdigest()
|
|
if calibration.get("artifact_sha256") != model_sha256:
|
|
raise ValueError("candidate calibration is not bound to the uploaded model")
|
|
if not bool(guard.get("accepted")):
|
|
raise ValueError("candidate retrain guard did not accept the model")
|
|
if guard.get("candidate_artifact_sha256") != model_sha256:
|
|
raise ValueError("candidate guard is not bound to the uploaded model")
|
|
validation = calibration.get("validation")
|
|
if not isinstance(validation, dict) or not _validation_passed(validation):
|
|
raise ValueError("candidate quality gate did not pass")
|
|
if validation.get("protocol") != "untouched_model_holdout_with_threshold_walk_forward":
|
|
raise ValueError("candidate validation protocol is not an untouched holdout")
|
|
|
|
self.runtime_dir.mkdir(parents=True, exist_ok=True)
|
|
backup_dir = self.runtime_dir / ".model_backups" / f"{_compact_now()}-{job_id}"
|
|
backup_dir.mkdir(parents=True, exist_ok=True)
|
|
for name in sorted(REQUIRED_MODEL_BUNDLE):
|
|
current = self.runtime_dir / name
|
|
if current.is_file():
|
|
shutil.copy2(current, backup_dir / name)
|
|
|
|
promoted: list[dict[str, Any]] = []
|
|
artifact_rows = {
|
|
str(item.get("name")): item
|
|
for item in job.get("artifacts", [])
|
|
if isinstance(item, dict)
|
|
}
|
|
for name in sorted(REQUIRED_MODEL_BUNDLE):
|
|
staged_path = ready_dir / name
|
|
target_tmp = self.runtime_dir / f".{name}.{job_id}.promote"
|
|
shutil.copy2(staged_path, target_tmp)
|
|
os.replace(target_tmp, self.runtime_dir / name)
|
|
row = artifact_rows.get(name, {})
|
|
promoted.append(
|
|
{
|
|
"name": name,
|
|
"sha256": row.get("sha256", ""),
|
|
"promoted_at": _now(),
|
|
}
|
|
)
|
|
_remove_tree(self.upload_root / job_id)
|
|
return promoted
|
|
|
|
def _load_state(self) -> dict[str, Any]:
|
|
try:
|
|
data = json.loads(self.state_path.read_text(encoding="utf-8"))
|
|
except (OSError, json.JSONDecodeError):
|
|
data = {}
|
|
if not isinstance(data, dict):
|
|
data = {}
|
|
data.setdefault("jobs", [])
|
|
return data
|
|
|
|
def _save_state(self, state: dict[str, Any]) -> None:
|
|
self.runtime_dir.mkdir(parents=True, exist_ok=True)
|
|
tmp = self.state_path.with_suffix(".tmp")
|
|
tmp.write_text(json.dumps(state, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
|
|
os.replace(tmp, self.state_path)
|
|
|
|
def _worker_from_payload(self, payload: dict[str, Any]) -> dict[str, Any]:
|
|
worker_id = str(payload.get("worker_id") or payload.get("id") or "windows-training-host").strip()
|
|
return {
|
|
"id": worker_id,
|
|
"name": str(payload.get("name") or worker_id).strip(),
|
|
"path": str(payload.get("path") or "").strip(),
|
|
"version": str(payload.get("version") or "1"),
|
|
"last_seen_at": _now(),
|
|
}
|
|
|
|
def _public_status(self, state: dict[str, Any]) -> dict[str, Any]:
|
|
worker = state.get("worker") if isinstance(state.get("worker"), dict) else {}
|
|
last_seen = _parse_time(str(worker.get("last_seen_at") or ""))
|
|
active = self._active_job(state)
|
|
latest = _latest_job(state)
|
|
recently_seen = bool(last_seen and datetime.now(UTC) - last_seen <= ONLINE_WINDOW)
|
|
agent_busy = bool(
|
|
active
|
|
and active.get("status") == "running"
|
|
and worker
|
|
and active.get("claimed_by") == worker.get("id")
|
|
)
|
|
return {
|
|
"available": True,
|
|
"agent_online": recently_seen or agent_busy,
|
|
"agent_recently_seen": recently_seen,
|
|
"agent_busy": agent_busy,
|
|
"worker": worker,
|
|
"active_job": self._public_job(active),
|
|
"latest_job": self._public_job(latest),
|
|
"pending_jobs": sum(1 for job in state.get("jobs", []) if job.get("status") == "pending"),
|
|
}
|
|
|
|
@staticmethod
|
|
def _public_job(job: dict[str, Any] | None) -> dict[str, Any] | None:
|
|
if job is None:
|
|
return None
|
|
public = dict(job)
|
|
public.pop("lease_token", None)
|
|
public.pop("uploads", None)
|
|
return public
|
|
|
|
@staticmethod
|
|
def _require_lease(job: dict[str, Any], payload: dict[str, Any]) -> None:
|
|
expected = str(job.get("lease_token") or "")
|
|
supplied = str(payload.get("lease_token") or "")
|
|
if not expected or not supplied or not hmac.compare_digest(expected, supplied):
|
|
raise ValueError("training job lease is invalid or expired")
|
|
|
|
def _active_job(self, state: dict[str, Any]) -> dict[str, Any] | None:
|
|
for job in reversed(state.get("jobs", [])):
|
|
if job.get("status") in {"pending", "running"}:
|
|
return job
|
|
return None
|
|
|
|
def _oldest_pending_job(self, state: dict[str, Any]) -> dict[str, Any] | None:
|
|
for job in state.get("jobs", []):
|
|
if job.get("status") == "pending":
|
|
return job
|
|
return None
|
|
|
|
def _job_by_id(self, state: dict[str, Any], job_id: str) -> dict[str, Any] | None:
|
|
for job in state.get("jobs", []):
|
|
if job.get("id") == job_id:
|
|
return job
|
|
return None
|
|
|
|
def _expire_stale_jobs(self, state: dict[str, Any]) -> None:
|
|
now = datetime.now(UTC)
|
|
for job in state.get("jobs", []):
|
|
if job.get("status") != "running":
|
|
continue
|
|
lease_updated_at = _parse_time(
|
|
str(job.get("updated_at") or job.get("claimed_at") or "")
|
|
)
|
|
if not lease_updated_at or now - lease_updated_at <= RUNNING_LEASE_TIMEOUT:
|
|
continue
|
|
job_id = str(job.get("id") or "")
|
|
if job_id:
|
|
_remove_tree(self.upload_root / job_id)
|
|
job.pop("lease_token", None)
|
|
job.pop("uploads", None)
|
|
attempts = int(job.get("attempts", 0))
|
|
if attempts < MAX_JOB_ATTEMPTS:
|
|
job["status"] = "pending"
|
|
job["phase"] = "queued"
|
|
job["progress_percent"] = 0
|
|
job["message"] = "training worker lease expired; queued for retry"
|
|
job["retry_queued_at"] = _now()
|
|
for key in ("claimed_at", "claimed_by", "worker", "updated_at"):
|
|
job.pop(key, None)
|
|
else:
|
|
job["status"] = "failed"
|
|
job["phase"] = "failed"
|
|
job["completed_at"] = _now()
|
|
job["message"] = "training worker lease expired after maximum retries"
|
|
|
|
def _trim_jobs(self, state: dict[str, Any]) -> None:
|
|
jobs = state.get("jobs", [])
|
|
if isinstance(jobs, list) and len(jobs) > 30:
|
|
state["jobs"] = jobs[-30:]
|
|
|
|
|
|
def _safe_parameters(value: Any) -> dict[str, Any]:
|
|
if not isinstance(value, dict):
|
|
return {}
|
|
allowed = {
|
|
"symbols",
|
|
"limit",
|
|
"lookbacks",
|
|
"architectures",
|
|
"hidden_sizes",
|
|
"layers",
|
|
"dropouts",
|
|
"epochs",
|
|
"validation_window",
|
|
"holdout_window",
|
|
"ensemble_seeds",
|
|
"selection_folds",
|
|
"learning_rate",
|
|
"weight_decay",
|
|
"horizon",
|
|
"horizons",
|
|
"patience",
|
|
"context_symbols",
|
|
"features",
|
|
"seed",
|
|
"interval",
|
|
"pooled",
|
|
"resume_candidate",
|
|
}
|
|
result = {key: value[key] for key in allowed if key in value}
|
|
for key, low, high in (
|
|
("limit", 500, 20000),
|
|
("epochs", 1, 200),
|
|
("validation_window", 64, 2000),
|
|
("holdout_window", 64, 1000),
|
|
("selection_folds", 1, 12),
|
|
("horizon", 1, 96),
|
|
("patience", 1, 50),
|
|
("seed", 1, 2_147_483_647),
|
|
):
|
|
if key not in result:
|
|
continue
|
|
try:
|
|
result[key] = max(low, min(high, int(result[key])))
|
|
except (TypeError, ValueError):
|
|
result.pop(key, None)
|
|
if "symbols" in result:
|
|
symbols = [
|
|
item.strip().upper()
|
|
for item in str(result["symbols"]).split(",")
|
|
if re.fullmatch(r"[A-Z0-9]{3,20}", item.strip().upper())
|
|
]
|
|
result["symbols"] = ",".join(symbols[:30])
|
|
if "architectures" in result:
|
|
architectures = [
|
|
item.strip().lower()
|
|
for item in str(result["architectures"]).split(",")
|
|
if item.strip().lower() in {"lstm", "gru"}
|
|
]
|
|
result["architectures"] = ",".join(architectures) or "lstm,gru"
|
|
for key in (
|
|
"lookbacks",
|
|
"hidden_sizes",
|
|
"layers",
|
|
"dropouts",
|
|
"ensemble_seeds",
|
|
"horizons",
|
|
"context_symbols",
|
|
"features",
|
|
"interval",
|
|
):
|
|
if key in result:
|
|
result[key] = str(result[key])[: 4000 if key == "features" else 500]
|
|
for key, low, high in (
|
|
("learning_rate", 0.00001, 0.1),
|
|
("weight_decay", 0.0, 0.1),
|
|
):
|
|
if key not in result:
|
|
continue
|
|
try:
|
|
result[key] = max(low, min(high, float(result[key])))
|
|
except (TypeError, ValueError):
|
|
result.pop(key, None)
|
|
if "pooled" in result:
|
|
result["pooled"] = result["pooled"] is True
|
|
if "resume_candidate" in result:
|
|
result["resume_candidate"] = result["resume_candidate"] is True
|
|
return result
|
|
|
|
|
|
def _valid_job_id(value: str) -> str:
|
|
try:
|
|
return str(uuid.UUID(str(value)))
|
|
except (ValueError, AttributeError, TypeError) as exc:
|
|
raise ValueError("invalid training job id") from exc
|
|
|
|
|
|
def _read_json(path: Path) -> dict[str, Any]:
|
|
try:
|
|
data = json.loads(path.read_text(encoding="utf-8"))
|
|
except (OSError, json.JSONDecodeError) as exc:
|
|
raise ValueError(f"invalid training artifact: {path.name}") from exc
|
|
if not isinstance(data, dict):
|
|
raise ValueError(f"invalid training artifact: {path.name}")
|
|
return data
|
|
|
|
|
|
def _validation_passed(validation: dict[str, Any]) -> bool:
|
|
if "passed" in validation:
|
|
return bool(validation.get("passed"))
|
|
return str(validation.get("status", "")).strip().lower() in {"pass", "passed", "ok"}
|
|
|
|
|
|
def _validate_symbol_models(symbols: dict[str, Any]) -> None:
|
|
for symbol, entry in symbols.items():
|
|
if not isinstance(entry, dict):
|
|
raise ValueError(f"candidate model entry is invalid: {symbol}")
|
|
if entry.get("model") not in {"torch_lstm", "torch_gru"}:
|
|
raise ValueError(f"candidate model architecture is invalid: {symbol}")
|
|
try:
|
|
lookback = int(entry.get("lookback", 0))
|
|
input_size = int(entry.get("input_size", 0))
|
|
hidden_size = int(entry.get("hidden_size", 0))
|
|
except (TypeError, ValueError) as exc:
|
|
raise ValueError(f"candidate model dimensions are invalid: {symbol}") from exc
|
|
if not 4 <= lookback <= 512 or not 1 <= input_size <= 256 or not 1 <= hidden_size <= 1024:
|
|
raise ValueError(f"candidate model dimensions are out of range: {symbol}")
|
|
members = entry.get("ensemble_members")
|
|
payloads = members if isinstance(members, list) and members else [entry]
|
|
for payload in payloads:
|
|
if not isinstance(payload, dict) or not isinstance(payload.get("state_dict"), dict):
|
|
raise ValueError(f"candidate recurrent state is missing: {symbol}")
|
|
merged = {**entry, **payload}
|
|
if merged.get("multitask_head") is True:
|
|
required = (
|
|
"head_hidden_weight",
|
|
"head_hidden_bias",
|
|
"return_head_weight",
|
|
"return_head_bias",
|
|
"event_head_weight",
|
|
"event_head_bias",
|
|
)
|
|
if any(not isinstance(merged.get(name), list) for name in required):
|
|
raise ValueError(f"candidate multitask forecast head is missing: {symbol}")
|
|
elif not isinstance(merged.get("head_weight"), list) or not isinstance(
|
|
merged.get("head_bias"), list
|
|
):
|
|
raise ValueError(f"candidate forecast head is missing: {symbol}")
|
|
|
|
|
|
def _compact_now() -> str:
|
|
return datetime.now(UTC).strftime("%Y%m%dT%H%M%SZ")
|
|
|
|
|
|
def _latest_job(state: dict[str, Any]) -> dict[str, Any] | None:
|
|
jobs = state.get("jobs", [])
|
|
if not jobs:
|
|
return None
|
|
latest = jobs[-1]
|
|
return latest if isinstance(latest, dict) else None
|
|
|
|
|
|
def _now() -> str:
|
|
return datetime.now(UTC).isoformat(timespec="seconds")
|
|
|
|
|
|
def _parse_time(value: str) -> datetime | None:
|
|
if not value:
|
|
return None
|
|
try:
|
|
return datetime.fromisoformat(value.replace("Z", "+00:00"))
|
|
except ValueError:
|
|
return None
|
|
|
|
|
|
def _coerce_percent(value: Any, default: Any = 0) -> int:
|
|
try:
|
|
number = int(float(value))
|
|
except (TypeError, ValueError):
|
|
try:
|
|
number = int(float(default))
|
|
except (TypeError, ValueError):
|
|
number = 0
|
|
return max(0, min(number, 100))
|
|
|
|
|
|
def _remove_tree(path: Path) -> None:
|
|
if not path.exists():
|
|
return
|
|
for child in path.iterdir():
|
|
if child.is_dir():
|
|
_remove_tree(child)
|
|
else:
|
|
child.unlink(missing_ok=True)
|
|
path.rmdir()
|