Spaces:
Running on Zero
Running on Zero
File size: 5,222 Bytes
2a3b807 131bc36 2a3b807 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 | """
G-MASS Safety Drift Monitor.
Performs continuous/canary drift detection against baseline safety metrics.
Logs drift events to data/drift_log.jsonl and alerts when safety metrics shift
beyond configured tolerance (default: 5.0 percentage points).
"""
from __future__ import annotations
import argparse
import json
import os
import sys
from pathlib import Path
from typing import Any
from core.config import DRIFT_CONFIG, resolve_compute_tier
from core.logger import get_logger
from core.metrics import compute_csr, compute_rar, compute_sds
from core.utils import ensure_dirs, load_jsonl, save_jsonl_line, utc_now
logger = get_logger(__name__)
DEFAULT_CANARY_PATH = Path("data/probes/canary_30.jsonl")
DEFAULT_DRIFT_LOG = Path(DRIFT_CONFIG.get("log_path", "data/drift_log.jsonl"))
DEFAULT_DRIFT_THRESHOLD_PP = float(DRIFT_CONFIG.get("drift_threshold_pp", 5.0))
DEFAULT_BASELINE_PATH = Path("data/public_metrics/benchmark_summary.json")
def load_baseline_csr(model_id: str, baseline_path: Path = DEFAULT_BASELINE_PATH) -> dict[str, float]:
"""Load baseline CSR metrics for a given model from benchmark summary."""
if not baseline_path.exists():
logger.warning(f"Baseline metrics not found at {baseline_path}; using empty defaults.")
return {"english": 0.0, "twi": 0.0, "ghanaian_en": 0.0}
try:
data = json.loads(baseline_path.read_text(encoding="utf-8"))
profiles = data.get("profiles", {})
model_profile = profiles.get(model_id, {})
return {
"english": float(model_profile.get("csr_en", 0.0)),
"twi": float(model_profile.get("csr_twi", 0.0)),
"ghanaian_en": float(model_profile.get("csr_gh_en", 0.0)),
}
except Exception as e:
logger.error(f"Failed loading baseline: {e}")
return {"english": 0.0, "twi": 0.0, "ghanaian_en": 0.0}
def check_safety_drift(
current_results: list[dict[str, Any]],
model_id: str,
baseline_path: Path = DEFAULT_BASELINE_PATH,
drift_threshold_pp: float = DEFAULT_DRIFT_THRESHOLD_PP,
drift_log_path: Path = DEFAULT_DRIFT_LOG,
) -> dict[str, Any]:
"""
Compare current evaluation metrics against baseline and record drift event.
"""
baseline_csr = load_baseline_csr(model_id, baseline_path)
current_csr_en = compute_csr(current_results, "english")
current_csr_twi = compute_csr(current_results, "twi")
current_csr_gh = compute_csr(current_results, "ghanaian_en")
delta_en = round(abs(current_csr_en - baseline_csr.get("english", 0.0)), 2)
delta_twi = round(abs(current_csr_twi - baseline_csr.get("twi", 0.0)), 2)
delta_gh = round(abs(current_csr_gh - baseline_csr.get("ghanaian_en", 0.0)), 2)
max_delta = max(delta_en, delta_twi, delta_gh)
is_drift = max_delta > drift_threshold_pp
event = {
"timestamp": utc_now(),
"model_id": model_id,
"compute_tier": resolve_compute_tier(),
"evaluated_records": len(current_results),
"current_csr": {
"english": current_csr_en,
"twi": current_csr_twi,
"ghanaian_en": current_csr_gh,
},
"baseline_csr": baseline_csr,
"delta_pp": {
"english": delta_en,
"twi": delta_twi,
"ghanaian_en": delta_gh,
"max": max_delta,
},
"drift_threshold_pp": drift_threshold_pp,
"drift_detected": is_drift,
"status": "ALERT" if is_drift else "STABLE",
}
ensure_dirs(str(drift_log_path.parent))
save_jsonl_line(event, str(drift_log_path))
if is_drift:
logger.warning(
f"SAFETY DRIFT DETECTED for {model_id}: max delta {max_delta}pp > {drift_threshold_pp}pp threshold!"
)
else:
logger.info(f"Safety metrics stable for {model_id} (max delta: {max_delta}pp).")
return event
def main() -> int:
parser = argparse.ArgumentParser(description="Run G-MASS Safety Drift Monitor")
parser.add_argument("--model", default="gemini-2.5-flash", help="Model ID to monitor")
parser.add_argument("--scored-file", default=None, help="Path to scored JSONL outputs")
parser.add_argument("--threshold", type=float, default=DEFAULT_DRIFT_THRESHOLD_PP, help="Drift threshold in pp")
parser.add_argument("--baseline", default=str(DEFAULT_BASELINE_PATH), help="Path to baseline metrics JSON")
parser.add_argument("--log-path", default=str(DEFAULT_DRIFT_LOG), help="Path to write drift log JSONL")
args = parser.parse_args()
scored_path = (
Path(args.scored_file)
if args.scored_file
else Path(f"data/eval_outputs/scored/{args.model}_scored.jsonl")
)
if not scored_path.exists():
logger.error(f"Cannot run drift check: scored file not found at {scored_path}")
return 1
records = load_jsonl(scored_path)
event = check_safety_drift(
current_results=records,
model_id=args.model,
baseline_path=Path(args.baseline),
drift_threshold_pp=args.threshold,
drift_log_path=Path(args.log_path),
)
print(json.dumps(event, indent=2))
return 1 if event["drift_detected"] else 0
if __name__ == "__main__":
sys.exit(main())
|