Files
sentiment-engine/MALKHUT/malkhut/training/rollback.py

107 lines
3.3 KiB
Python
Raw Normal View History

"""
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)