278 lines
11 KiB
Python
278 lines
11 KiB
Python
"""
|
||
Thompson Sampling Contextual Bandit — Strategy Selector
|
||
=========================================================
|
||
Replaces the greedy SQL auto-tuner with proper online RL.
|
||
|
||
Problem
|
||
-------
|
||
The old greedy tuner picks the strategy with the highest *average* quality
|
||
across all history. It never explores alternatives once one strategy leads,
|
||
and it ignores context (the same strategy isn't best at all times of day
|
||
and all load levels).
|
||
|
||
Solution
|
||
--------
|
||
A contextual multi-armed bandit with Thompson Sampling:
|
||
|
||
Context = (time_band, load_band) — up to 7 × 4 = 28 states
|
||
Arms = 4 strategies — balanced / fuel_saver / aggressive_speed / zone_strict
|
||
Reward = quality_score / 100 — 0..1 float, already logged by MLDataCollector
|
||
|
||
How Thompson Sampling works
|
||
---------------------------
|
||
Each (context, arm) pair has a Beta(α, β) posterior where:
|
||
α = sum of rewards seen so far (high quality calls push α up)
|
||
β = sum of "anti-rewards" (low quality calls push β up)
|
||
|
||
To SELECT a strategy:
|
||
1. For each arm, sample θ ~ Beta(α, β)
|
||
2. Pick arm with highest θ
|
||
→ Naturally balances exploration (uncertain arms get sampled often)
|
||
with exploitation (well-known good arms dominate when confident)
|
||
|
||
To UPDATE after an assignment:
|
||
reward = quality_score / 100
|
||
α += reward
|
||
β += (1 - reward)
|
||
|
||
Bootstrap
|
||
---------
|
||
On first startup the entire historical SQLite DB is replayed to warm up
|
||
posteriors so the bandit starts informed, not blank.
|
||
If saved state already exists it is loaded from disk instead.
|
||
|
||
Persistence
|
||
-----------
|
||
ml_data/strategy_bandit.json — saved every 10 updates.
|
||
"""
|
||
|
||
import json
|
||
import logging
|
||
import os
|
||
import threading
|
||
from typing import Dict, List, Optional, Tuple
|
||
|
||
import numpy as np
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
_SAVE_PATH = os.getenv("BANDIT_PATH", "ml_data/strategy_bandit.json")
|
||
_DB_PATH = os.getenv("ML_DB_PATH", "ml_data/ml_store.db")
|
||
|
||
ARMS = ["balanced", "fuel_saver", "aggressive_speed", "zone_strict"]
|
||
|
||
|
||
class ContextualBandit:
|
||
"""
|
||
Thompson Sampling bandit for ml_strategy selection.
|
||
|
||
Context key : "{time_band}|{load_band}"
|
||
Arms : ARMS list (4 strategies)
|
||
Prior : Beta(1, 1) — uniform, no initial preference
|
||
"""
|
||
|
||
def __init__(self):
|
||
self._lock = threading.Lock()
|
||
# _posteriors[ctx_key][arm] = [alpha, beta]
|
||
self._posteriors: Dict[str, Dict[str, List[float]]] = {}
|
||
self._update_count = 0
|
||
self._total_pulls = 0
|
||
self._load()
|
||
self._bootstrap_from_db()
|
||
|
||
# ── Public API ───────────────────────────────────────────────────────────
|
||
|
||
def select(self, time_band: str, load_band: str) -> str:
|
||
"""
|
||
Sample from Beta posteriors and return the strategy with the highest
|
||
sample. Explores uncertain arms naturally; exploits known-good arms
|
||
as confidence grows.
|
||
"""
|
||
ctx = self._ctx(time_band, load_band)
|
||
with self._lock:
|
||
samples = {
|
||
arm: float(np.random.beta(*self._ab(ctx, arm)))
|
||
for arm in ARMS
|
||
}
|
||
self._total_pulls += 1
|
||
chosen = max(samples, key=samples.__getitem__)
|
||
logger.debug(
|
||
f"[Bandit] ctx={ctx} "
|
||
f"samples={{{', '.join(f'{k}:{v:.3f}' for k, v in samples.items())}}} "
|
||
f"→ {chosen}"
|
||
)
|
||
return chosen
|
||
|
||
def update(self, time_band: str, load_band: str,
|
||
strategy: str, quality_score: float) -> None:
|
||
"""Update the Beta posterior for (context, arm) after observing quality."""
|
||
if strategy not in ARMS:
|
||
return
|
||
ctx = self._ctx(time_band, load_band)
|
||
reward = min(1.0, max(0.0, float(quality_score) / 100.0))
|
||
with self._lock:
|
||
ab = self._ab(ctx, strategy)
|
||
ab[0] += reward # alpha ← quality adds to success mass
|
||
ab[1] += (1.0 - reward) # beta ← (1-quality) adds to failure mass
|
||
self._update_count += 1
|
||
should_save = self._update_count % 10 == 0
|
||
if should_save:
|
||
self._save()
|
||
|
||
def best_arm(self, time_band: str, load_band: str) -> Tuple[str, float]:
|
||
"""
|
||
Return the arm with the highest posterior mean (pure exploitation, no
|
||
sampling noise). Used for logging and dashboard, not for live selection.
|
||
"""
|
||
ctx = self._ctx(time_band, load_band)
|
||
with self._lock:
|
||
means = {
|
||
arm: self._ab(ctx, arm)[0] / sum(self._ab(ctx, arm))
|
||
for arm in ARMS
|
||
}
|
||
best = max(means, key=means.__getitem__)
|
||
return best, round(means[best], 4)
|
||
|
||
def get_stats(self) -> Dict:
|
||
"""Return per-context arm statistics for the ML admin dashboard."""
|
||
with self._lock:
|
||
stats: Dict[str, Dict] = {}
|
||
for ctx, arms in sorted(self._posteriors.items()):
|
||
stats[ctx] = {}
|
||
for arm in ARMS:
|
||
if arm not in arms:
|
||
stats[ctx][arm] = {"mean_reward": 0.5, "observations": 0,
|
||
"alpha": 1.0, "beta": 1.0}
|
||
continue
|
||
alpha, beta = arms[arm]
|
||
n = alpha + beta - 2.0 # subtract the Beta(1,1) prior mass
|
||
mean = alpha / (alpha + beta)
|
||
stats[ctx][arm] = {
|
||
"mean_reward": round(mean, 4),
|
||
"observations": round(max(0, n), 1),
|
||
"alpha": round(alpha, 2),
|
||
"beta": round(beta, 2),
|
||
}
|
||
return {
|
||
"total_pulls": self._total_pulls,
|
||
"total_updates": self._update_count,
|
||
"context_count": len(self._posteriors),
|
||
"arms": ARMS,
|
||
"contexts": stats,
|
||
}
|
||
|
||
# ── Persistence ──────────────────────────────────────────────────────────
|
||
|
||
def _save(self) -> None:
|
||
try:
|
||
with self._lock:
|
||
snapshot = {
|
||
"posteriors": {
|
||
ctx: {arm: list(ab) for arm, ab in arms.items()}
|
||
for ctx, arms in self._posteriors.items()
|
||
},
|
||
"update_count": self._update_count,
|
||
"total_pulls": self._total_pulls,
|
||
}
|
||
os.makedirs(os.path.dirname(_SAVE_PATH) or ".", exist_ok=True)
|
||
with open(_SAVE_PATH, "w", encoding="utf-8") as f:
|
||
json.dump(snapshot, f, indent=2)
|
||
except Exception as e:
|
||
logger.warning(f"[Bandit] Save failed: {e}")
|
||
|
||
def _load(self) -> None:
|
||
try:
|
||
if not os.path.exists(_SAVE_PATH):
|
||
logger.info("[Bandit] No saved state — will bootstrap from DB.")
|
||
return
|
||
with open(_SAVE_PATH, "r", encoding="utf-8") as f:
|
||
data = json.load(f)
|
||
self._posteriors = data.get("posteriors", {})
|
||
self._update_count = data.get("update_count", 0)
|
||
self._total_pulls = data.get("total_pulls", 0)
|
||
logger.info(
|
||
f"[Bandit] Loaded from disk — "
|
||
f"{len(self._posteriors)} contexts, {self._update_count} updates"
|
||
)
|
||
except Exception as e:
|
||
logger.warning(f"[Bandit] Load failed (starting fresh): {e}")
|
||
|
||
# ── Bootstrap ────────────────────────────────────────────────────────────
|
||
|
||
def _bootstrap_from_db(self) -> None:
|
||
"""
|
||
Warm-up posteriors by replaying every historical assignment event.
|
||
Skipped if the saved JSON already reflects current DB data.
|
||
"""
|
||
if self._update_count > 0:
|
||
logger.info(
|
||
f"[Bandit] Already warmed ({self._update_count} updates) — "
|
||
"skipping DB bootstrap."
|
||
)
|
||
return
|
||
try:
|
||
import sqlite3
|
||
from app.services.ml.behavior_analyzer import (
|
||
time_band as _tb,
|
||
load_band as _lb,
|
||
)
|
||
conn = sqlite3.connect(_DB_PATH)
|
||
rows = conn.execute(
|
||
"SELECT timestamp, avg_load, ml_strategy, quality_score "
|
||
"FROM assignment_ml_log "
|
||
"WHERE ml_strategy IS NOT NULL AND quality_score IS NOT NULL "
|
||
"ORDER BY id ASC"
|
||
).fetchall()
|
||
conn.close()
|
||
|
||
if not rows:
|
||
logger.info("[Bandit] No historical data — starting with uniform priors.")
|
||
return
|
||
|
||
for ts, avg_load, strategy, quality in rows:
|
||
t_band = _tb(str(ts)) if ts else "UNKNOWN"
|
||
l_band = _lb(float(avg_load or 0))
|
||
self.update(t_band, l_band, strategy, float(quality or 50))
|
||
|
||
# Reset counter so we don't confuse bootstrap updates with live ones
|
||
with self._lock:
|
||
self._update_count = len(rows)
|
||
|
||
logger.info(
|
||
f"[Bandit] Bootstrapped from {len(rows)} historical events — "
|
||
f"{len(self._posteriors)} contexts, {len(ARMS)} arms."
|
||
)
|
||
self._save()
|
||
except Exception as e:
|
||
logger.warning(f"[Bandit] Bootstrap failed (non-fatal): {e}")
|
||
|
||
# ── Helpers ──────────────────────────────────────────────────────────────
|
||
|
||
@staticmethod
|
||
def _ctx(time_band: str, load_band: str) -> str:
|
||
return f"{time_band}|{load_band}"
|
||
|
||
def _ab(self, ctx: str, arm: str) -> List[float]:
|
||
"""Get or create Beta(1, 1) uniform prior for (ctx, arm). NOT thread-safe alone."""
|
||
if ctx not in self._posteriors:
|
||
self._posteriors[ctx] = {}
|
||
if arm not in self._posteriors[ctx]:
|
||
self._posteriors[ctx][arm] = [1.0, 1.0]
|
||
return self._posteriors[ctx][arm]
|
||
|
||
|
||
# ── Singleton ────────────────────────────────────────────────────────────────
|
||
|
||
_bandit_instance: Optional[ContextualBandit] = None
|
||
_bandit_lock = threading.Lock()
|
||
|
||
|
||
def get_bandit() -> ContextualBandit:
|
||
"""Return the process-level singleton ContextualBandit."""
|
||
global _bandit_instance
|
||
if _bandit_instance is None:
|
||
with _bandit_lock:
|
||
if _bandit_instance is None:
|
||
_bandit_instance = ContextualBandit()
|
||
return _bandit_instance
|