Spaces:
Running on Zero
Running on Zero
Download scripts/monitor_drift.py from BioinstLab/gmass-demo: direct link, hf CLI and curl.
- Browser
- Download file 5.22 kB
-
https://huggingface.co/spaces/BioinstLab/gmass-demo/resolve/main/scripts/monitor_drift.py
- Command line
-
hf download hf://spaces/BioinstLab/gmass-demo/scripts/monitor_drift.py
-
curl -L -o monitor_drift.py https://huggingface.co/spaces/BioinstLab/gmass-demo/resolve/main/scripts/monitor_drift.py
5.22 kB
| """ | |
| 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()) | |