File size: 6,237 Bytes
622315e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
143
144
145
146
147
148
149
150
151
152
153
154
155
156
"""

MEXAR - Domain Guardrail Threshold Sweep.

Sweeps candidate DOMAIN_SIMILARITY_THRESHOLD values over all domain query sets and agents

to compute out-of-scope rejection Precision, Recall, F1, and In-Domain False Rejection Rate.

Exports results to evaluation_outputs/guardrail_threshold_sweep.json.

"""
import os
import sys
import json
import logging
from typing import Dict, List, Any, Tuple

sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from modules.reasoning_engine import create_reasoning_engine

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

REPO_ROOT = os.path.abspath(os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", ".."))
OUTPUT_DIR_BACKEND = os.path.abspath(os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "evaluation_outputs"))
OUTPUT_DIR_ROOT = os.path.abspath(os.path.join(REPO_ROOT, "evaluation_outputs"))
QUERY_SETS_DIR = os.path.join(REPO_ROOT, "test_data", "query_sets")


def load_all_query_sets() -> Dict[str, List[Dict[str, Any]]]:
    """Load query set JSON files for medical, legal, and financial domains."""
    domains = ["medical", "legal", "financial"]
    query_sets = {}
    for domain in domains:
        filepath = os.path.join(QUERY_SETS_DIR, f"{domain}_queries.json")
        if not os.path.exists(filepath):
            logger.error(f"Query set file not found: {filepath}")
            continue
        with open(filepath, "r", encoding="utf-8") as f:
            query_sets[domain] = json.load(f)
    return query_sets


def run_threshold_sweep() -> Dict[str, Any]:
    """

    Run threshold sweep over candidate threshold values [0.02, 0.05, 0.08, 0.10, 0.15, 0.20, 0.25, 0.30, 0.35, 0.40].

    Measures confusion matrix for out-of-scope rejection across all query x agent combinations.

    """
    engine = create_reasoning_engine()
    query_sets = load_all_query_sets()
    domains = ["medical", "legal", "financial"]

    # Pre-load agents
    agents = {}
    for d in domains:
        agent_name = f"{d}_agent"
        agents[d] = engine._load_agent(agent_name)

    candidate_thresholds = [0.02, 0.05, 0.08, 0.10, 0.15, 0.20, 0.25, 0.30, 0.35, 0.40]

    sweep_results = []
    best_threshold = 0.05
    best_f1 = -1.0
    best_metrics = {}

    for thresh in candidate_thresholds:
        # Override instance threshold
        engine.DOMAIN_SIMILARITY_THRESHOLD = thresh

        tp = 0  # Truly out-of-scope, rejected
        fp = 0  # Truly in-scope, rejected (False Rejection)
        fn = 0  # Truly out-of-scope, accepted (Leak)
        tn = 0  # Truly in-scope, accepted

        for source_domain, queries in query_sets.items():
            for item in queries:
                query_text = item["query"]
                item_is_in_domain = item.get("is_in_domain", True)

                for target_domain in domains:
                    target_agent = agents[target_domain]
                    target_agent_name = f"{target_domain}_agent"

                    # Determine ground truth scope for target agent
                    truly_in_scope = (source_domain == target_domain) and item_is_in_domain
                    truly_out_of_scope = not truly_in_scope

                    # Check guardrail directly
                    in_domain_flag, score = engine._check_guardrail(
                        query_text,
                        target_agent["domain_signature"],
                        target_agent["prompt_analysis"]
                    )
                    rejected = not in_domain_flag

                    if rejected and truly_out_of_scope:
                        tp += 1
                    elif rejected and truly_in_scope:
                        fp += 1
                    elif not rejected and truly_out_of_scope:
                        fn += 1
                    else:
                        tn += 1

        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
        false_rejection_rate = fp / (fp + tn) if (fp + tn) > 0 else 0.0

        res_entry = {
            "threshold": thresh,
            "tp": tp,
            "fp": fp,
            "fn": fn,
            "tn": tn,
            "precision": round(precision, 4),
            "recall": round(recall, 4),
            "f1": round(f1, 4),
            "false_rejection_rate": round(false_rejection_rate, 4)
        }
        sweep_results.append(res_entry)

        if f1 > best_f1:
            best_f1 = f1
            best_threshold = thresh
            best_metrics = res_entry

    # Format table for console output
    print("\n" + "=" * 70)
    print("GUARDRAIL THRESHOLD SWEEP RESULTS")
    print("=" * 70)
    print(f"{'Threshold':<10} | {'Precision':<10} | {'Recall':<10} | {'F1 Score':<10} | {'False Rejection Rate':<20}")
    print("-" * 70)
    for r in sweep_results:
        star = " *" if r["threshold"] == best_threshold else ""
        print(f"{r['threshold']:<10.2f} | {r['precision']:<10.4f} | {r['recall']:<10.4f} | {r['f1']:<10.4f} | {r['false_rejection_rate']:<20.4f}{star}")
    print("=" * 70)
    print(f"Optimal Threshold: {best_threshold} (F1 = {best_f1:.4f}, False Rejection Rate = {best_metrics.get('false_rejection_rate', 0.0):.4f})")
    print("=" * 70 + "\n")

    output_payload = {
        "sweep_results": sweep_results,
        "optimal_threshold": best_threshold,
        "optimal_f1": round(best_f1, 4),
        "optimal_false_rejection_rate": best_metrics.get("false_rejection_rate", 0.0),
        "best_metrics": best_metrics
    }

    for out_dir in [OUTPUT_DIR_BACKEND, OUTPUT_DIR_ROOT]:
        os.makedirs(out_dir, exist_ok=True)
        out_file = os.path.join(out_dir, "guardrail_threshold_sweep.json")
        with open(out_file, "w", encoding="utf-8") as f:
            json.dump(output_payload, f, indent=2)
        logger.info(f"Sweep results written to {out_file}")

    return output_payload


if __name__ == "__main__":
    run_threshold_sweep()