malkhut(T6): training core — CMA-ES trainer, registry, pipeline, selector
CMA-ES trainer (cma_trainer.py): self-play pool, bootstrap CI, ScenarioFactory with behavior-driven scenarios, auto-compile, label query interfaces. Policy registry (registry.py): CANDIDATE → ACTIVE lifecycle. Training pipeline (pipeline.py): bounded continuous learning loop + logger. Strategy selector (selector.py): regime → strategy mapping, performance matrix.
This commit is contained in:
165
MALKHUT/malkhut/training/registry.py
Normal file
165
MALKHUT/malkhut/training/registry.py
Normal file
@@ -0,0 +1,165 @@
|
||||
"""
|
||||
Policy Registry — versioned promotion pipeline for trained policies.
|
||||
|
||||
Lifecycle:
|
||||
CANDIDATE → BACKTESTED → SELF_PLAY_CONFIRMED → SHADOW → TINY_LIVE → ACTIVE
|
||||
↓
|
||||
REJECTED
|
||||
|
||||
Promotion gates:
|
||||
1. Unit tests pass
|
||||
2. Replay verification pass
|
||||
3. Self-play robust score pass (bootstrap CI)
|
||||
4. Shadow live discrepancy pass
|
||||
5. Tiny live risk pass
|
||||
6. Manual or rule-based promotion
|
||||
|
||||
Active policy is loaded by the engine via params_provider.
|
||||
Hot-reload via CONTROL_PLANE HOT_RELOAD_POLICY command.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any, Mapping, Optional
|
||||
|
||||
from malkhut.state import FulfilmentPolicyParams
|
||||
from malkhut.storage.ch_store import MalkhutCHStore
|
||||
|
||||
|
||||
class PolicyStage(str, Enum):
|
||||
CANDIDATE = "CANDIDATE"
|
||||
BACKTESTED = "BACKTESTED"
|
||||
SELF_PLAY_CONFIRMED = "SELF_PLAY_CONFIRMED"
|
||||
SHADOW = "SHADOW"
|
||||
TINY_LIVE = "TINY_LIVE"
|
||||
ACTIVE = "ACTIVE"
|
||||
RETIRED = "RETIRED"
|
||||
REJECTED = "REJECTED"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PolicyRecord:
|
||||
version: str
|
||||
stage: PolicyStage
|
||||
params: FulfilmentPolicyParams
|
||||
score: float
|
||||
created_ts_ns: int
|
||||
promoted_ts_ns: int = 0
|
||||
evaluation_summary: Mapping[str, Any] = field(default_factory=dict)
|
||||
promotion_reason: str = ""
|
||||
|
||||
|
||||
class PolicyRegistry:
|
||||
"""
|
||||
Versioned promotion pipeline with CH persistence.
|
||||
|
||||
One ACTIVE policy at a time.
|
||||
Multiple CANDIDATE/SHADOW policies in flight.
|
||||
RETIRED policies kept for audit.
|
||||
"""
|
||||
|
||||
def __init__(self, store: Optional[MalkhutCHStore] = None) -> None:
|
||||
self._store = store
|
||||
self._records: dict[str, PolicyRecord] = {}
|
||||
self._active_version: Optional[str] = None
|
||||
|
||||
# Ensure CH table exists
|
||||
if self._store:
|
||||
self._store.ensure_tables()
|
||||
|
||||
def register_candidate(
|
||||
self,
|
||||
params: FulfilmentPolicyParams,
|
||||
score: float,
|
||||
evaluation_summary: Optional[Mapping[str, Any]] = None,
|
||||
) -> PolicyRecord:
|
||||
"""Register a new candidate policy."""
|
||||
record = PolicyRecord(
|
||||
version=params.version,
|
||||
stage=PolicyStage.CANDIDATE,
|
||||
params=params,
|
||||
score=score,
|
||||
created_ts_ns=time.time_ns(),
|
||||
evaluation_summary=evaluation_summary or {},
|
||||
)
|
||||
self._records[params.version] = record
|
||||
self._persist(record)
|
||||
return record
|
||||
|
||||
def promote(self, version: str, stage: PolicyStage, reason: str = "") -> PolicyRecord:
|
||||
"""Promote a policy to a new stage."""
|
||||
if version not in self._records:
|
||||
raise KeyError(f"Policy {version} not found")
|
||||
|
||||
old = self._records[version]
|
||||
new = PolicyRecord(
|
||||
version=old.version,
|
||||
stage=stage,
|
||||
params=old.params,
|
||||
score=old.score,
|
||||
created_ts_ns=old.created_ts_ns,
|
||||
promoted_ts_ns=time.time_ns(),
|
||||
evaluation_summary=old.evaluation_summary,
|
||||
promotion_reason=reason,
|
||||
)
|
||||
self._records[version] = new
|
||||
self._persist(new)
|
||||
|
||||
if stage == PolicyStage.ACTIVE:
|
||||
self._active_version = version
|
||||
|
||||
return new
|
||||
|
||||
def reject(self, version: str, reason: str = "") -> PolicyRecord:
|
||||
"""Reject a candidate policy."""
|
||||
return self.promote(version, PolicyStage.REJECTED, reason)
|
||||
|
||||
def retire(self, version: str, reason: str = "") -> PolicyRecord:
|
||||
"""Retire the current active policy."""
|
||||
return self.promote(version, PolicyStage.RETIRED, reason)
|
||||
|
||||
def load_active(self) -> Optional[FulfilmentPolicyParams]:
|
||||
"""Load the currently active policy parameters."""
|
||||
if self._active_version and self._active_version in self._records:
|
||||
record = self._records[self._active_version]
|
||||
if record.stage == PolicyStage.ACTIVE:
|
||||
return record.params
|
||||
|
||||
# Fallback: scan records for ACTIVE
|
||||
for record in self._records.values():
|
||||
if record.stage == PolicyStage.ACTIVE:
|
||||
self._active_version = record.version
|
||||
return record.params
|
||||
|
||||
return None
|
||||
|
||||
def get_record(self, version: str) -> Optional[PolicyRecord]:
|
||||
return self._records.get(version)
|
||||
|
||||
def get_by_stage(self, stage: PolicyStage) -> list[PolicyRecord]:
|
||||
return [r for r in self._records.values() if r.stage == stage]
|
||||
|
||||
@property
|
||||
def active_version(self) -> Optional[str]:
|
||||
return self._active_version
|
||||
|
||||
@property
|
||||
def record_count(self) -> int:
|
||||
return len(self._records)
|
||||
|
||||
def _persist(self, record: PolicyRecord) -> None:
|
||||
if self._store:
|
||||
self._store.store_policy_snapshot(
|
||||
version=record.version,
|
||||
score=record.score,
|
||||
params_str=json.dumps({
|
||||
"stage": record.stage.value,
|
||||
"promotion_reason": record.promotion_reason,
|
||||
"created_ts_ns": record.created_ts_ns,
|
||||
"promoted_ts_ns": record.promoted_ts_ns,
|
||||
}),
|
||||
evaluation_summary=json.dumps(dict(record.evaluation_summary)),
|
||||
)
|
||||
Reference in New Issue
Block a user