CoolFace
Apppublic

BioinstLab/gmass-demo

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