Spaces:
Sleeping
Sleeping
| """End-to-End Pipeline Evaluation Script. | |
| Evaluates the full fused pipeline (L1 + L2 + fusion) against labelled datasets. | |
| Performs threshold tuning (Grid Search over tau_m and tau_s) and computes metrics. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import asyncio | |
| import json | |
| import logging | |
| import random | |
| import sys | |
| import time | |
| from collections import defaultdict | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from typing import Any | |
| import numpy as np | |
| REPO_ROOT = Path(__file__).resolve().parents[1] | |
| if str(REPO_ROOT) not in sys.path: | |
| sys.path.insert(0, str(REPO_ROOT)) | |
| from app.engines.adapters import L1HeuristicsAnalyzer, L2SemanticAnalyzer | |
| from app.engines.contracts import EmailPayload, Verdict, AnalyzerLayer | |
| from app.engines.manager import AnalysisManager, _MALICIOUS_THRESHOLD, _SUSPICIOUS_THRESHOLD | |
| from app.engines.semantic.orchestrator import load_models, shutdown_models | |
| logger = logging.getLogger("evaluate_pipeline") | |
| DEFAULT_DATA_DIR = Path("data/l2_datasets") | |
| DEFAULT_OUTPUT_DIR = Path("evaluation_e2e") | |
| class CachedSample: | |
| sample_id: str | |
| task: str | |
| label: int | |
| fused_score: float | |
| verdict: Verdict | |
| l1_score: float | |
| l2_score: float | |
| l2_confidence: float | |
| winner_layer: str | |
| short_circuited: bool | |
| elapsed_ms: float | |
| fusion_metadata: dict[str, Any] | |
| def load_jsonl(path: Path, task_name: str) -> list[dict[str, Any]]: | |
| if not path.exists(): | |
| logger.warning(f"File not found1: {path}") | |
| return [] | |
| samples = [] | |
| with path.open("r", encoding="utf-8") as handle: | |
| for line_num, line in enumerate(handle, start=1): | |
| line = line.strip() | |
| if not line: | |
| continue | |
| try: | |
| row = json.loads(line) | |
| except json.JSONDecodeError: | |
| continue | |
| row["task"] = task_name | |
| samples.append(row) | |
| return samples | |
| def stratified_sample(samples: list[dict[str, Any]], limit: int, seed: int) -> list[dict[str, Any]]: | |
| if not samples or limit >= len(samples): | |
| return samples | |
| positives = [s for s in samples if int(s.get("label", 0)) == 1] | |
| negatives = [s for s in samples if int(s.get("label", 0)) == 0] | |
| pos_ratio = len(positives) / len(samples) if samples else 0 | |
| n_pos = int(round(limit * pos_ratio)) | |
| n_neg = limit - n_pos | |
| rng = random.Random(seed) | |
| sampled_pos = rng.sample(positives, min(n_pos, len(positives))) | |
| sampled_neg = rng.sample(negatives, min(n_neg, len(negatives))) | |
| combined = sampled_pos + sampled_neg | |
| rng.shuffle(combined) | |
| return combined | |
| async def run_inference(samples: list[dict[str, Any]], manager: AnalysisManager) -> list[CachedSample]: | |
| results = [] | |
| for idx, sample in enumerate(samples): | |
| text = str(sample.get("text_body", "")) | |
| html = str(sample.get("html_body", "")) | |
| label = int(sample.get("label", 0)) | |
| sample_id = sample.get("id", f"sample-{idx}") | |
| task = sample.get("task", "unknown") | |
| payload = EmailPayload(body_html=html, body_text=text) | |
| t0 = time.perf_counter() | |
| verdict_res = await manager.run(payload) | |
| elapsed_ms = (time.perf_counter() - t0) * 1000.0 | |
| l1_score = 0.0 | |
| l2_score = 0.0 | |
| l2_confidence = 0.0 | |
| for r in verdict_res.layer_results: | |
| if r.layer == AnalyzerLayer.HEURISTICS: | |
| l1_score = r.score | |
| elif r.layer == AnalyzerLayer.SEMANTIC: | |
| l2_score = r.score | |
| l2_confidence = r.confidence | |
| results.append(CachedSample( | |
| sample_id=sample_id, | |
| task=task, | |
| label=label, | |
| fused_score=verdict_res.fused_score, | |
| verdict=verdict_res.verdict, | |
| l1_score=l1_score, | |
| l2_score=l2_score, | |
| l2_confidence=l2_confidence, | |
| winner_layer=str(verdict_res.fusion_metadata.get("winner") or "none"), | |
| short_circuited=verdict_res.short_circuited, | |
| elapsed_ms=elapsed_ms, | |
| fusion_metadata=verdict_res.fusion_metadata | |
| )) | |
| if (idx + 1) % 100 == 0: | |
| logger.info(f"Processed {idx + 1} / {len(samples)} samples") | |
| return results | |
| def compute_metrics(labels: np.ndarray, preds: np.ndarray) -> dict[str, float]: | |
| tp = int(np.sum((labels == 1) & (preds == 1))) | |
| fp = int(np.sum((labels == 0) & (preds == 1))) | |
| fn = int(np.sum((labels == 1) & (preds == 0))) | |
| tn = int(np.sum((labels == 0) & (preds == 0))) | |
| precision = tp / (tp + fp) if (tp + fp) > 0 else 0.0 | |
| recall = tp / (tp + fn) if (tp + fn) > 0 else 0.0 | |
| f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0.0 | |
| return { | |
| "tp": float(tp), "fp": float(fp), "tn": float(tn), "fn": float(fn), | |
| "precision": precision, "recall": recall, "f1": f1 | |
| } | |
| def print_metrics(m: dict[str, float], label: str): | |
| print(f"\n{label}") | |
| print(f" Precision: {m['precision']:.4f}") | |
| print(f" Recall: {m['recall']:.4f}") | |
| print(f" F1: {m['f1']:.4f}") | |
| def phase_grid_search(cached: list[CachedSample], recall_floor: float = 0.85) -> dict[str, Any]: | |
| labels = np.array([s.label for s in cached]) | |
| scores = np.array([s.fused_score for s in cached]) | |
| best_f1 = -1.0 | |
| best_tm = 70.0 | |
| best_ts = 30.0 | |
| best_m = {} | |
| grid_results = [] | |
| for tm in np.arange(30.0, 97.5, 2.5): | |
| for ts in np.arange(10.0, tm, 2.5): | |
| preds = (scores >= tm).astype(int) | |
| m = compute_metrics(labels, preds) | |
| # soft detection (verdict >= suspicious) | |
| soft_preds = (scores >= ts).astype(int) | |
| soft_m = compute_metrics(labels, soft_preds) | |
| soft_detection_rate = soft_m["recall"] | |
| res = { | |
| "tau_m": float(tm), "tau_s": float(ts), | |
| "f1": m["f1"], "precision": m["precision"], "recall": m["recall"], | |
| "soft_detection_rate": soft_detection_rate | |
| } | |
| grid_results.append(res) | |
| if m["recall"] >= recall_floor: | |
| if m["f1"] > best_f1: | |
| best_f1 = m["f1"] | |
| best_tm = float(tm) | |
| best_ts = float(ts) | |
| best_m = res | |
| # If no config met the recall floor, take the one with highest recall | |
| if best_f1 < 0 and grid_results: | |
| logger.warning(f"No threshold combination met recall_floor={recall_floor}. Taking highest F1 overall.") | |
| best_res = max(grid_results, key=lambda x: x["f1"]) | |
| best_tm = best_res["tau_m"] | |
| best_ts = best_res["tau_s"] | |
| best_m = best_res | |
| # Sort grid results for top 5 | |
| top_5 = sorted([r for r in grid_results if r["recall"] >= recall_floor], key=lambda x: x["f1"], reverse=True)[:5] | |
| if not top_5: | |
| top_5 = sorted(grid_results, key=lambda x: x["f1"], reverse=True)[:5] | |
| return { | |
| "recommended_thresholds": { | |
| "malicious_threshold": best_tm, | |
| "suspicious_threshold": best_ts, | |
| "recall_floor": recall_floor | |
| }, | |
| "metrics_at_recommended": best_m, | |
| "top_5_grid_configurations": top_5 | |
| } | |
| def write_markdown_report(out_dir: Path, def_m: dict, grid_report: dict, test_m: dict, sc_stats: dict): | |
| lines = [ | |
| "# End-to-End Pipeline Evaluation Report", | |
| "", | |
| "## Runtime Summary", | |
| f"- Total Samples: {sc_stats['total']}", | |
| f"- Short-circuited: {sc_stats['short_circuited']} ({sc_stats['sc_pct']:.1f}%)", | |
| f"- Full pipeline invocations: {sc_stats['full_pipeline']} ({sc_stats['full_pct']:.1f}%)", | |
| "", | |
| "## Phase 2: Default Thresholds (Val Set)", | |
| f"- Malicious Threshold: {_MALICIOUS_THRESHOLD}", | |
| f"- Precision: {def_m['precision']:.4f}", | |
| f"- Recall: {def_m['recall']:.4f}", | |
| f"- F1 Score: {def_m['f1']:.4f}", | |
| "", | |
| "## Phase 3: Grid Search Optimization", | |
| ] | |
| if grid_report: | |
| rec = grid_report["recommended_thresholds"] | |
| m = grid_report["metrics_at_recommended"] | |
| lines.extend([ | |
| f"- Recommended Malicious Threshold (`τ_m`): **{rec['malicious_threshold']}**", | |
| f"- Recommended Suspicious Threshold (`τ_s`): **{rec['suspicious_threshold']}**", | |
| f"- Recall Floor applied: {rec['recall_floor']}", | |
| "", | |
| "**Metrics at Recommended:**", | |
| f"- Precision: {m['precision']:.4f}", | |
| f"- Recall: {m['recall']:.4f}", | |
| f"- F1 Score: {m['f1']:.4f}", | |
| f"- Soft Detection Rate (≥ Suspicious): {m['soft_detection_rate']:.4f}", | |
| ]) | |
| else: | |
| lines.append("- *Grid Search Skipped*") | |
| lines.append("") | |
| lines.append("## Phase 4: Test Set Metrics") | |
| if test_m: | |
| lines.extend([ | |
| f"- Precision: {test_m['precision']:.4f}", | |
| f"- Recall: {test_m['recall']:.4f}", | |
| f"- F1 Score: {test_m['f1']:.4f}", | |
| ]) | |
| else: | |
| lines.append("- *Test set evaluation was not run.*") | |
| (out_dir / "evaluation_report.md").write_text("\n".join(lines), encoding="utf-8") | |
| async def async_main(): | |
| parser = argparse.ArgumentParser(description="End-to-End Pipeline Evaluation") | |
| parser.add_argument("--data-dir", type=Path, default=DEFAULT_DATA_DIR) | |
| parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR) | |
| parser.add_argument("--limit", type=int, default=None) | |
| parser.add_argument("--limit-per-task", type=int, default=None) | |
| parser.add_argument("--sample-pct", type=float, default=None) | |
| parser.add_argument("--include-train", action="store_true") | |
| parser.add_argument("--run-test-eval", action="store_true") | |
| parser.add_argument("--skip-grid-search", action="store_true") | |
| parser.add_argument("--recall-floor", type=float, default=0.85) | |
| parser.add_argument("--seed", type=int, default=42) | |
| parser.add_argument("--timeout", type=float, default=5.0) | |
| args = parser.parse_args() | |
| args.output_dir.mkdir(parents=True, exist_ok=True) | |
| # Setup AnalysisManager | |
| manager = AnalysisManager(timeout=args.timeout) | |
| manager.register(L1HeuristicsAnalyzer()) | |
| manager.register(L2SemanticAnalyzer()) | |
| try: | |
| logger.info("Loading L2 models...") | |
| load_models() | |
| # Load samples | |
| tasks = ["prompt_injection", "malicious_intent"] | |
| splits = ["val", "train"] if args.include_train else ["val"] | |
| val_samples = [] | |
| for task in tasks: | |
| task_samples = [] | |
| for split in splits: | |
| p = args.data_dir / task / f"{split}.jsonl" | |
| task_samples.extend(load_jsonl(p, task)) | |
| if args.limit_per_task: | |
| task_samples = stratified_sample(task_samples, args.limit_per_task, args.seed) | |
| elif args.sample_pct: | |
| limit = int(len(task_samples) * args.sample_pct) | |
| task_samples = stratified_sample(task_samples, limit, args.seed) | |
| val_samples.extend(task_samples) | |
| if args.limit: | |
| val_samples = stratified_sample(val_samples, args.limit, args.seed) | |
| logger.info(f"Loaded {len(val_samples)} validation samples") | |
| if not val_samples: | |
| logger.warning("No validation samples loaded! Exiting.") | |
| return | |
| # Phase 1: Inference | |
| logger.info("Phase 1: Running inference on val set...") | |
| cached_val = await run_inference(val_samples, manager) | |
| # Phase 2: Diagnostic metrics (default thresholds) | |
| labels = np.array([s.label for s in cached_val]) | |
| scores = np.array([s.fused_score for s in cached_val]) | |
| preds_def = (scores >= _MALICIOUS_THRESHOLD).astype(int) | |
| def_m = compute_metrics(labels, preds_def) | |
| print_metrics(def_m, f"Default Thresholds (tau_m={_MALICIOUS_THRESHOLD})") | |
| # Short circuit stats | |
| sc_count = sum(1 for s in cached_val if s.short_circuited) | |
| sc_stats = { | |
| "total": len(cached_val), | |
| "short_circuited": sc_count, | |
| "sc_pct": (sc_count / len(cached_val)) * 100 if cached_val else 0.0, | |
| "full_pipeline": len(cached_val) - sc_count, | |
| "full_pct": ((len(cached_val) - sc_count) / len(cached_val)) * 100 if cached_val else 0.0, | |
| } | |
| print(f" Short-circuited: {sc_stats['short_circuited']} / {sc_stats['total']} ({sc_stats['sc_pct']:.1f}%)") | |
| best_tm = _MALICIOUS_THRESHOLD | |
| best_ts = _SUSPICIOUS_THRESHOLD | |
| grid_report = {} | |
| # Phase 3: Grid Search | |
| if not args.skip_grid_search: | |
| logger.info("Phase 3: Running Grid Search...") | |
| grid_report = phase_grid_search(cached_val, args.recall_floor) | |
| best_tm = grid_report["recommended_thresholds"]["malicious_threshold"] | |
| best_ts = grid_report["recommended_thresholds"]["suspicious_threshold"] | |
| print(f"\nGrid Search Recommended thresholds:") | |
| print(f" Malicious: {best_tm}") | |
| print(f" Suspicious: {best_ts}") | |
| print_metrics(grid_report["metrics_at_recommended"], "Metrics at Recommended Thresholds") | |
| else: | |
| logger.info("Skipping Phase 3 Grid Search") | |
| # Phase 4: Test set evaluation | |
| test_m = {} | |
| if args.run_test_eval: | |
| logger.info("Phase 4: Running Test Set Evaluation...") | |
| test_samples = [] | |
| for task in tasks: | |
| p = args.data_dir / task / "test.jsonl" | |
| test_samples.extend(load_jsonl(p, task)) | |
| if args.limit_per_task: | |
| test_samples = stratified_sample(test_samples, args.limit_per_task, args.seed) | |
| elif args.sample_pct: | |
| test_samples = stratified_sample(test_samples, int(len(test_samples) * args.sample_pct), args.seed) | |
| if args.limit: | |
| test_samples = stratified_sample(test_samples, args.limit, args.seed) | |
| logger.info(f"Loaded {len(test_samples)} test samples") | |
| if test_samples: | |
| cached_test = await run_inference(test_samples, manager) | |
| test_labels = np.array([s.label for s in cached_test]) | |
| test_scores = np.array([s.fused_score for s in cached_test]) | |
| test_preds = (test_scores >= best_tm).astype(int) | |
| test_m = compute_metrics(test_labels, test_preds) | |
| print_metrics(test_m, f"Test Set Metrics (tau_m={best_tm})") | |
| else: | |
| logger.warning("No test samples found.") | |
| # Phase 5: Artifacts | |
| out_json = { | |
| "grid_search": grid_report, | |
| "test_metrics": test_m, | |
| "val_default_metrics": def_m, | |
| "runtime_stats": sc_stats | |
| } | |
| (args.output_dir / "metrics.json").write_text(json.dumps(out_json, indent=2)) | |
| write_markdown_report(args.output_dir, def_m, grid_report, test_m, sc_stats) | |
| logger.info(f"Artifacts written to {args.output_dir}") | |
| finally: | |
| shutdown_models() | |
| def main(): | |
| logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s") | |
| asyncio.run(async_main()) | |
| if __name__ == "__main__": | |
| main() | |