107 lines
3.3 KiB
Python
107 lines
3.3 KiB
Python
|
|
"""
|
||
|
|
Policy Rollback — auto-rollback if shadow performance degrades.
|
||
|
|
|
||
|
|
Enables:
|
||
|
|
- Detecting when a promoted policy performs worse than baseline
|
||
|
|
- Automatically reverting to previous best
|
||
|
|
- Preventing bad policies from reaching live
|
||
|
|
"""
|
||
|
|
from __future__ import annotations
|
||
|
|
|
||
|
|
import time
|
||
|
|
from dataclasses import dataclass, field
|
||
|
|
from typing import Any, List, Mapping, Optional, Tuple
|
||
|
|
|
||
|
|
from malkhut.state import FulfilmentPolicyParams
|
||
|
|
from malkhut.training.registry import PolicyRegistry, PolicyStage, PolicyRecord
|
||
|
|
|
||
|
|
|
||
|
|
@dataclass(frozen=True, slots=True)
|
||
|
|
class RollbackEvent:
|
||
|
|
"""Record of a rollback event."""
|
||
|
|
ts_ns: int
|
||
|
|
rolled_back_version: str
|
||
|
|
rolled_back_to: str
|
||
|
|
reason: str
|
||
|
|
performance_drop: float
|
||
|
|
|
||
|
|
|
||
|
|
class PolicyRollback:
|
||
|
|
"""
|
||
|
|
Auto-rollback if shadow performance degrades.
|
||
|
|
|
||
|
|
Monitors active policy performance and reverts to previous best
|
||
|
|
if performance drops below threshold.
|
||
|
|
"""
|
||
|
|
|
||
|
|
def __init__(
|
||
|
|
self,
|
||
|
|
registry: PolicyRegistry,
|
||
|
|
degradation_threshold: float = -5.0, # bps
|
||
|
|
min_shadow_steps: int = 10,
|
||
|
|
) -> None:
|
||
|
|
self._registry = registry
|
||
|
|
self._degradation_threshold = degradation_threshold
|
||
|
|
self._min_shadow_steps = min_shadow_steps
|
||
|
|
self._shadow_scores: List[float] = []
|
||
|
|
self._rollback_events: List[RollbackEvent] = []
|
||
|
|
|
||
|
|
def record_shadow_score(self, score: float) -> None:
|
||
|
|
"""Record a shadow performance score."""
|
||
|
|
self._shadow_scores.append(score)
|
||
|
|
|
||
|
|
def check_rollback(self) -> Optional[RollbackEvent]:
|
||
|
|
"""
|
||
|
|
Check if rollback is needed.
|
||
|
|
|
||
|
|
Returns RollbackEvent if rollback should happen, None otherwise.
|
||
|
|
"""
|
||
|
|
if len(self._shadow_scores) < self._min_shadow_steps:
|
||
|
|
return None
|
||
|
|
|
||
|
|
# Compare recent average to baseline
|
||
|
|
recent = self._shadow_scores[-self._min_shadow_steps:]
|
||
|
|
avg_recent = sum(recent) / len(recent)
|
||
|
|
|
||
|
|
# Get baseline score (first policy in registry)
|
||
|
|
baseline = self._registry.get_by_stage(PolicyStage.ACTIVE)
|
||
|
|
if not baseline:
|
||
|
|
return None
|
||
|
|
|
||
|
|
baseline_record = baseline[0] if baseline else None
|
||
|
|
if not baseline_record:
|
||
|
|
return None
|
||
|
|
|
||
|
|
# Check degradation
|
||
|
|
if avg_recent < self._degradation_threshold:
|
||
|
|
# Find previous best to rollback to
|
||
|
|
previous = self._registry.get_by_stage(PolicyStage.RETIRED)
|
||
|
|
if previous:
|
||
|
|
rollback_to = previous[0].version
|
||
|
|
else:
|
||
|
|
rollback_to = "baseline"
|
||
|
|
|
||
|
|
# Perform rollback
|
||
|
|
current = baseline_record[0] if isinstance(baseline_record, list) else baseline_record
|
||
|
|
self._registry.retire(current.version, "auto_rollback_degradation")
|
||
|
|
|
||
|
|
event = RollbackEvent(
|
||
|
|
ts_ns=time.time_ns(),
|
||
|
|
rolled_back_version=current.version,
|
||
|
|
rolled_back_to=rollback_to,
|
||
|
|
reason=f"performance_drop_{avg_recent:.2f}",
|
||
|
|
performance_drop=avg_recent - self._degradation_threshold,
|
||
|
|
)
|
||
|
|
self._rollback_events.append(event)
|
||
|
|
return event
|
||
|
|
|
||
|
|
return None
|
||
|
|
|
||
|
|
@property
|
||
|
|
def rollback_events(self) -> List[RollbackEvent]:
|
||
|
|
return list(self._rollback_events)
|
||
|
|
|
||
|
|
@property
|
||
|
|
def shadow_score_count(self) -> int:
|
||
|
|
return len(self._shadow_scores)
|