lime-audit 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- lime_audit/__init__.py +3 -0
- lime_audit/analyse.py +290 -0
- lime_audit/charts.py +179 -0
- lime_audit/cli.py +75 -0
- lime_audit/config.py +14 -0
- lime_audit/default_test_set.json +36 -0
- lime_audit/metrics.py +96 -0
- lime_audit/runner.py +297 -0
- lime_audit-0.1.0.dist-info/METADATA +248 -0
- lime_audit-0.1.0.dist-info/RECORD +13 -0
- lime_audit-0.1.0.dist-info/WHEEL +4 -0
- lime_audit-0.1.0.dist-info/entry_points.txt +2 -0
- lime_audit-0.1.0.dist-info/licenses/LICENSE +21 -0
lime_audit/__init__.py
ADDED
lime_audit/analyse.py
ADDED
|
@@ -0,0 +1,290 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Analyse LIME audit results - compute stability and faithfulness aggregates.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
import csv
|
|
6
|
+
import json
|
|
7
|
+
import os
|
|
8
|
+
from collections import defaultdict
|
|
9
|
+
|
|
10
|
+
import numpy as np
|
|
11
|
+
|
|
12
|
+
from lime_audit.config import (
|
|
13
|
+
FAITHFULNESS_THRESHOLD_DIRECTION,
|
|
14
|
+
LIME_NUM_FEATURES,
|
|
15
|
+
STABILITY_THRESHOLD_JACCARD,
|
|
16
|
+
STABILITY_THRESHOLD_KENDALL,
|
|
17
|
+
TOP_K_FOR_JACCARD,
|
|
18
|
+
)
|
|
19
|
+
from lime_audit.metrics import bootstrap_ci, compute_pairwise_stability, count_tokenizer_mismatch
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def load_raw_attributions(path, num_features=LIME_NUM_FEATURES):
|
|
23
|
+
by_input = defaultdict(dict)
|
|
24
|
+
with open(path, "r", encoding="utf-8") as f:
|
|
25
|
+
reader = csv.DictReader(f)
|
|
26
|
+
for row in reader:
|
|
27
|
+
if not row.get("input_id") or not row.get("seed"):
|
|
28
|
+
continue
|
|
29
|
+
input_id = int(row["input_id"])
|
|
30
|
+
seed = int(row["seed"])
|
|
31
|
+
tokens = []
|
|
32
|
+
for i in range(1, num_features + 1):
|
|
33
|
+
t = row.get(f"token_{i}", "")
|
|
34
|
+
w = float(row.get(f"weight_{i}", 0))
|
|
35
|
+
if t:
|
|
36
|
+
tokens.append((t, w))
|
|
37
|
+
by_input[input_id][seed] = {
|
|
38
|
+
"tokens": tokens,
|
|
39
|
+
"base_label": row["base_label"],
|
|
40
|
+
"base_confidence": float(row["base_confidence"]),
|
|
41
|
+
"base_positive_score": float(row["base_positive_score"]),
|
|
42
|
+
"lime_score": float(row["lime_score"]),
|
|
43
|
+
"duration_ms": int(row["duration_ms"]),
|
|
44
|
+
"category": row["category"],
|
|
45
|
+
"text": row["text"],
|
|
46
|
+
}
|
|
47
|
+
return dict(by_input)
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def load_deletion_results(path):
|
|
51
|
+
by_input = {}
|
|
52
|
+
with open(path, "r", encoding="utf-8") as f:
|
|
53
|
+
reader = csv.DictReader(f)
|
|
54
|
+
for row in reader:
|
|
55
|
+
status = row.get("status", "tested")
|
|
56
|
+
if status == "skipped_empty":
|
|
57
|
+
by_input[int(row["input_id"])] = {
|
|
58
|
+
"status": "skipped_empty",
|
|
59
|
+
"category": row["category"],
|
|
60
|
+
"text": row["text"],
|
|
61
|
+
}
|
|
62
|
+
continue
|
|
63
|
+
entry = {
|
|
64
|
+
"status": "tested",
|
|
65
|
+
"removed_token": row["removed_token"],
|
|
66
|
+
"token_lime_weight": float(row["token_lime_weight"]),
|
|
67
|
+
"original_label": row["original_label"],
|
|
68
|
+
"original_positive_score": float(row["original_positive_score"]),
|
|
69
|
+
"modified_label": row["modified_label"],
|
|
70
|
+
"modified_positive_score": float(row["modified_positive_score"]),
|
|
71
|
+
"confidence_delta": float(row["confidence_delta"]),
|
|
72
|
+
"label_flipped": row["label_flipped"] == "True",
|
|
73
|
+
"direction_correct": row["direction_correct"] == "True",
|
|
74
|
+
"category": row["category"],
|
|
75
|
+
"text": row["text"],
|
|
76
|
+
}
|
|
77
|
+
if row.get("delta_top3") and row["delta_top3"] != "":
|
|
78
|
+
entry["delta_top3"] = float(row["delta_top3"])
|
|
79
|
+
entry["flipped_top3"] = row["flipped_top3"] == "True"
|
|
80
|
+
if row.get("delta_top5") and row["delta_top5"] != "":
|
|
81
|
+
entry["delta_top5"] = float(row["delta_top5"])
|
|
82
|
+
entry["flipped_top5"] = row["flipped_top5"] == "True"
|
|
83
|
+
by_input[int(row["input_id"])] = entry
|
|
84
|
+
return by_input
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def analyse(output_dir, model_name):
|
|
88
|
+
raw_csv = os.path.join(output_dir, "raw_attributions.csv")
|
|
89
|
+
del_csv = os.path.join(output_dir, "deletion_faithfulness.csv")
|
|
90
|
+
stability_csv = os.path.join(output_dir, "stability_metrics.csv")
|
|
91
|
+
summary_path = os.path.join(output_dir, "summary.json")
|
|
92
|
+
|
|
93
|
+
print("[1/4] Loading raw attributions...")
|
|
94
|
+
raw = load_raw_attributions(raw_csv)
|
|
95
|
+
print(f" {len(raw)} inputs loaded")
|
|
96
|
+
|
|
97
|
+
print("[2/4] Loading deletion results...")
|
|
98
|
+
deletions = load_deletion_results(del_csv)
|
|
99
|
+
print(f" {len(deletions)} deletion tests loaded")
|
|
100
|
+
|
|
101
|
+
print("[3/4] Computing stability and tokenizer mismatch...")
|
|
102
|
+
stability_rows = []
|
|
103
|
+
|
|
104
|
+
for input_id in sorted(raw.keys()):
|
|
105
|
+
seeds_data = raw[input_id]
|
|
106
|
+
first = next(iter(seeds_data.values()))
|
|
107
|
+
category = first["category"]
|
|
108
|
+
text = first["text"]
|
|
109
|
+
|
|
110
|
+
seed_tokens = {s: d["tokens"] for s, d in seeds_data.items()}
|
|
111
|
+
pairwise = compute_pairwise_stability(seed_tokens, k=TOP_K_FOR_JACCARD)
|
|
112
|
+
|
|
113
|
+
labels = [d["base_label"] for d in seeds_data.values()]
|
|
114
|
+
confidences = [d["base_confidence"] for d in seeds_data.values()]
|
|
115
|
+
|
|
116
|
+
mean_j = round(pairwise["mean_jaccard"], 4)
|
|
117
|
+
jaccard_ci = bootstrap_ci(pairwise["all_jaccards"])
|
|
118
|
+
tok_mismatch = count_tokenizer_mismatch(text, model_name)
|
|
119
|
+
|
|
120
|
+
stability_rows.append({
|
|
121
|
+
"input_id": input_id,
|
|
122
|
+
"category": category,
|
|
123
|
+
"text": text[:80],
|
|
124
|
+
"mean_jaccard_top5": mean_j,
|
|
125
|
+
"jaccard_ci_lower": round(jaccard_ci[0], 4),
|
|
126
|
+
"jaccard_ci_upper": round(jaccard_ci[1], 4),
|
|
127
|
+
"min_jaccard_top5": round(pairwise["min_jaccard"], 4),
|
|
128
|
+
"mean_kendall_tau": round(pairwise["mean_kendall_tau"], 4),
|
|
129
|
+
"min_kendall_tau": round(pairwise["min_kendall_tau"], 4) if not np.isnan(pairwise["min_kendall_tau"]) else "NaN",
|
|
130
|
+
"label_stable": len(set(labels)) == 1,
|
|
131
|
+
"confidence_std": round(float(np.std(confidences)), 6),
|
|
132
|
+
"top1_unanimous": pairwise["top1_unanimous"],
|
|
133
|
+
"lime_token_count": tok_mismatch["lime_token_count"],
|
|
134
|
+
"wordpiece_token_count": tok_mismatch["wordpiece_token_count"],
|
|
135
|
+
"token_mismatch": tok_mismatch["token_mismatch"],
|
|
136
|
+
})
|
|
137
|
+
|
|
138
|
+
fieldnames = [
|
|
139
|
+
"input_id", "category", "text", "mean_jaccard_top5",
|
|
140
|
+
"jaccard_ci_lower", "jaccard_ci_upper",
|
|
141
|
+
"min_jaccard_top5", "mean_kendall_tau", "min_kendall_tau",
|
|
142
|
+
"label_stable", "confidence_std", "top1_unanimous",
|
|
143
|
+
"lime_token_count", "wordpiece_token_count", "token_mismatch",
|
|
144
|
+
]
|
|
145
|
+
with open(stability_csv, "w", newline="", encoding="utf-8") as f:
|
|
146
|
+
writer = csv.DictWriter(f, fieldnames=fieldnames)
|
|
147
|
+
writer.writeheader()
|
|
148
|
+
writer.writerows(stability_rows)
|
|
149
|
+
print(f" Saved stability CSV: {stability_csv}")
|
|
150
|
+
|
|
151
|
+
print("[4/4] Computing summary...")
|
|
152
|
+
jaccards = [r["mean_jaccard_top5"] for r in stability_rows]
|
|
153
|
+
taus = [r["mean_kendall_tau"] for r in stability_rows if not isinstance(r["mean_kendall_tau"], str)]
|
|
154
|
+
overall_jaccard_ci = bootstrap_ci(jaccards) if jaccards else (None, None)
|
|
155
|
+
mismatches = [r["token_mismatch"] for r in stability_rows]
|
|
156
|
+
|
|
157
|
+
cat_stability = defaultdict(lambda: {"jaccards": [], "taus": []})
|
|
158
|
+
for r in stability_rows:
|
|
159
|
+
cat = r["category"]
|
|
160
|
+
cat_stability[cat]["jaccards"].append(r["mean_jaccard_top5"])
|
|
161
|
+
if not isinstance(r["mean_kendall_tau"], str):
|
|
162
|
+
cat_stability[cat]["taus"].append(r["mean_kendall_tau"])
|
|
163
|
+
|
|
164
|
+
tested = {k: v for k, v in deletions.items() if v.get("status") != "skipped_empty"}
|
|
165
|
+
skipped = {k: v for k, v in deletions.items() if v.get("status") == "skipped_empty"}
|
|
166
|
+
|
|
167
|
+
del_correct = [d for d in tested.values() if d["direction_correct"]]
|
|
168
|
+
del_flipped = [d for d in tested.values() if d["label_flipped"]]
|
|
169
|
+
del_deltas = [abs(d["confidence_delta"]) for d in tested.values() if not np.isnan(d["confidence_delta"])]
|
|
170
|
+
|
|
171
|
+
del3_deltas = [abs(d["delta_top3"]) for d in tested.values() if "delta_top3" in d]
|
|
172
|
+
del3_flipped = [d for d in tested.values() if d.get("flipped_top3")]
|
|
173
|
+
del5_deltas = [abs(d["delta_top5"]) for d in tested.values() if "delta_top5" in d]
|
|
174
|
+
del5_flipped = [d for d in tested.values() if d.get("flipped_top5")]
|
|
175
|
+
|
|
176
|
+
env_path = os.path.join(output_dir, "environment.json")
|
|
177
|
+
env_info = {}
|
|
178
|
+
if os.path.exists(env_path):
|
|
179
|
+
with open(env_path, "r") as f:
|
|
180
|
+
env_info = json.load(f)
|
|
181
|
+
|
|
182
|
+
cat_faith = defaultdict(lambda: {"correct": 0, "flipped": 0, "total": 0, "deltas": [],
|
|
183
|
+
"deltas_top3": [], "flipped_top3": 0,
|
|
184
|
+
"deltas_top5": [], "flipped_top5": 0})
|
|
185
|
+
for d in tested.values():
|
|
186
|
+
cat = d["category"]
|
|
187
|
+
cat_faith[cat]["total"] += 1
|
|
188
|
+
if d["direction_correct"]:
|
|
189
|
+
cat_faith[cat]["correct"] += 1
|
|
190
|
+
if d["label_flipped"]:
|
|
191
|
+
cat_faith[cat]["flipped"] += 1
|
|
192
|
+
if not np.isnan(d["confidence_delta"]):
|
|
193
|
+
cat_faith[cat]["deltas"].append(abs(d["confidence_delta"]))
|
|
194
|
+
if "delta_top3" in d:
|
|
195
|
+
cat_faith[cat]["deltas_top3"].append(abs(d["delta_top3"]))
|
|
196
|
+
if d.get("flipped_top3"):
|
|
197
|
+
cat_faith[cat]["flipped_top3"] += 1
|
|
198
|
+
if "delta_top5" in d:
|
|
199
|
+
cat_faith[cat]["deltas_top5"].append(abs(d["delta_top5"]))
|
|
200
|
+
if d.get("flipped_top5"):
|
|
201
|
+
cat_faith[cat]["flipped_top5"] += 1
|
|
202
|
+
|
|
203
|
+
summary = {
|
|
204
|
+
"experiment_metadata": env_info,
|
|
205
|
+
"stability_summary": {
|
|
206
|
+
"overall_mean_jaccard_top5": round(float(np.mean(jaccards)), 4) if jaccards else None,
|
|
207
|
+
"overall_jaccard_ci_95": [round(overall_jaccard_ci[0], 4), round(overall_jaccard_ci[1], 4)] if overall_jaccard_ci[0] is not None else None,
|
|
208
|
+
"overall_mean_kendall_tau": round(float(np.mean(taus)), 4) if taus else None,
|
|
209
|
+
"inputs_with_perfect_jaccard": sum(1 for j in jaccards if j >= 1.0),
|
|
210
|
+
"inputs_with_jaccard_below_threshold": sum(1 for j in jaccards if j < STABILITY_THRESHOLD_JACCARD),
|
|
211
|
+
"inputs_with_label_instability": sum(1 for r in stability_rows if not r["label_stable"]),
|
|
212
|
+
"inputs_with_unanimous_top1": sum(1 for r in stability_rows if r["top1_unanimous"]),
|
|
213
|
+
"total_inputs": len(stability_rows),
|
|
214
|
+
"threshold_jaccard": STABILITY_THRESHOLD_JACCARD,
|
|
215
|
+
"threshold_kendall": STABILITY_THRESHOLD_KENDALL,
|
|
216
|
+
"per_category": {
|
|
217
|
+
cat: {
|
|
218
|
+
"mean_jaccard": round(float(np.mean(v["jaccards"])), 4),
|
|
219
|
+
"jaccard_ci_95": [round(x, 4) for x in bootstrap_ci(v["jaccards"])] if len(v["jaccards"]) > 1 else None,
|
|
220
|
+
"mean_kendall_tau": round(float(np.mean(v["taus"])), 4) if v["taus"] else None,
|
|
221
|
+
}
|
|
222
|
+
for cat, v in cat_stability.items()
|
|
223
|
+
},
|
|
224
|
+
},
|
|
225
|
+
"tokenizer_mismatch_summary": {
|
|
226
|
+
"mean_mismatch": round(float(np.mean(mismatches)), 2),
|
|
227
|
+
"max_mismatch": int(np.max(mismatches)),
|
|
228
|
+
"inputs_with_mismatch": sum(1 for m in mismatches if m != 0),
|
|
229
|
+
"total_inputs": len(mismatches),
|
|
230
|
+
},
|
|
231
|
+
"faithfulness_summary": {
|
|
232
|
+
"total_deletion_tests": len(tested),
|
|
233
|
+
"skipped_inputs": len(skipped),
|
|
234
|
+
"direction_correct_count": len(del_correct),
|
|
235
|
+
"direction_correct_rate": round(len(del_correct) / len(tested), 4) if tested else None,
|
|
236
|
+
"label_flip_count": len(del_flipped),
|
|
237
|
+
"label_flip_rate": round(len(del_flipped) / len(tested), 4) if tested else None,
|
|
238
|
+
"mean_abs_confidence_delta": round(float(np.mean(del_deltas)), 4) if del_deltas else None,
|
|
239
|
+
"threshold_direction_correct": FAITHFULNESS_THRESHOLD_DIRECTION,
|
|
240
|
+
"top3_deletion": {
|
|
241
|
+
"mean_abs_delta": round(float(np.mean(del3_deltas)), 4) if del3_deltas else None,
|
|
242
|
+
"flip_count": len(del3_flipped),
|
|
243
|
+
"flip_rate": round(len(del3_flipped) / len(del3_deltas), 4) if del3_deltas else None,
|
|
244
|
+
},
|
|
245
|
+
"top5_deletion": {
|
|
246
|
+
"mean_abs_delta": round(float(np.mean(del5_deltas)), 4) if del5_deltas else None,
|
|
247
|
+
"flip_count": len(del5_flipped),
|
|
248
|
+
"flip_rate": round(len(del5_flipped) / len(del5_deltas), 4) if del5_deltas else None,
|
|
249
|
+
},
|
|
250
|
+
"per_category": {
|
|
251
|
+
cat: {
|
|
252
|
+
"direction_correct_rate": round(v["correct"] / v["total"], 4) if v["total"] else None,
|
|
253
|
+
"label_flip_rate": round(v["flipped"] / v["total"], 4) if v["total"] else None,
|
|
254
|
+
"mean_abs_delta": round(float(np.mean(v["deltas"])), 4) if v["deltas"] else None,
|
|
255
|
+
"mean_abs_delta_top3": round(float(np.mean(v["deltas_top3"])), 4) if v["deltas_top3"] else None,
|
|
256
|
+
"flip_rate_top3": round(v["flipped_top3"] / len(v["deltas_top3"]), 4) if v["deltas_top3"] else None,
|
|
257
|
+
"mean_abs_delta_top5": round(float(np.mean(v["deltas_top5"])), 4) if v["deltas_top5"] else None,
|
|
258
|
+
"flip_rate_top5": round(v["flipped_top5"] / len(v["deltas_top5"]), 4) if v["deltas_top5"] else None,
|
|
259
|
+
}
|
|
260
|
+
for cat, v in cat_faith.items()
|
|
261
|
+
},
|
|
262
|
+
},
|
|
263
|
+
}
|
|
264
|
+
|
|
265
|
+
with open(summary_path, "w", encoding="utf-8") as f:
|
|
266
|
+
json.dump(summary, f, indent=2)
|
|
267
|
+
|
|
268
|
+
ss = summary["stability_summary"]
|
|
269
|
+
fs = summary["faithfulness_summary"]
|
|
270
|
+
ts = summary["tokenizer_mismatch_summary"]
|
|
271
|
+
|
|
272
|
+
print(f"\n--- STABILITY ---")
|
|
273
|
+
ci = ss.get("overall_jaccard_ci_95")
|
|
274
|
+
ci_str = f" [{ci[0]}, {ci[1]}]" if ci else ""
|
|
275
|
+
print(f" Mean Jaccard top-5: {ss['overall_mean_jaccard_top5']}{ci_str}")
|
|
276
|
+
print(f" Mean Kendall tau: {ss['overall_mean_kendall_tau']}")
|
|
277
|
+
|
|
278
|
+
print(f"\n--- TOKENIZER MISMATCH ---")
|
|
279
|
+
print(f" Mean mismatch: {ts['mean_mismatch']} extra WordPiece tokens")
|
|
280
|
+
|
|
281
|
+
print(f"\n--- FAITHFULNESS ---")
|
|
282
|
+
print(f" Direction correct: {fs['direction_correct_rate']}")
|
|
283
|
+
print(f" Flip rate (top-1): {fs['label_flip_rate']}")
|
|
284
|
+
t3 = fs["top3_deletion"]
|
|
285
|
+
t5 = fs["top5_deletion"]
|
|
286
|
+
print(f" Mean |delta| top-3: {t3['mean_abs_delta']}, flip rate: {t3['flip_rate']}")
|
|
287
|
+
print(f" Mean |delta| top-5: {t5['mean_abs_delta']}, flip rate: {t5['flip_rate']}")
|
|
288
|
+
|
|
289
|
+
print(f"\nSaved: {summary_path}")
|
|
290
|
+
return summary_path
|
lime_audit/charts.py
ADDED
|
@@ -0,0 +1,179 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Generate charts for LIME audit results.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
import csv
|
|
6
|
+
import os
|
|
7
|
+
|
|
8
|
+
import matplotlib.pyplot as plt
|
|
9
|
+
import numpy as np
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
CATEGORY_COLORS = {
|
|
13
|
+
"negation_minimal_pairs": "#e63946",
|
|
14
|
+
"lexical_shortcuts": "#457b9d",
|
|
15
|
+
"ambiguity": "#2a9d8f",
|
|
16
|
+
"distribution_style_shift": "#e9c46a",
|
|
17
|
+
"strong_baselines": "#f4a261",
|
|
18
|
+
"edge_cases": "#264653",
|
|
19
|
+
}
|
|
20
|
+
|
|
21
|
+
CATEGORY_LABELS = {
|
|
22
|
+
"negation_minimal_pairs": "Negation",
|
|
23
|
+
"lexical_shortcuts": "Lexical Shortcuts",
|
|
24
|
+
"ambiguity": "Ambiguity",
|
|
25
|
+
"distribution_style_shift": "Distribution Shift",
|
|
26
|
+
"strong_baselines": "Strong Baselines",
|
|
27
|
+
"edge_cases": "Edge Cases",
|
|
28
|
+
}
|
|
29
|
+
|
|
30
|
+
STABILITY_THRESHOLD = 0.6
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def load_stability(output_dir):
|
|
34
|
+
rows = []
|
|
35
|
+
path = os.path.join(output_dir, "stability_metrics.csv")
|
|
36
|
+
with open(path, "r", encoding="utf-8") as f:
|
|
37
|
+
reader = csv.DictReader(f)
|
|
38
|
+
for row in reader:
|
|
39
|
+
rows.append({
|
|
40
|
+
"input_id": int(row["input_id"]),
|
|
41
|
+
"category": row["category"],
|
|
42
|
+
"text": row["text"],
|
|
43
|
+
"mean_jaccard": float(row["mean_jaccard_top5"]),
|
|
44
|
+
"ci_lower": float(row.get("jaccard_ci_lower", row["mean_jaccard_top5"])),
|
|
45
|
+
"ci_upper": float(row.get("jaccard_ci_upper", row["mean_jaccard_top5"])),
|
|
46
|
+
})
|
|
47
|
+
return rows
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def load_deletion(output_dir):
|
|
51
|
+
rows = []
|
|
52
|
+
path = os.path.join(output_dir, "deletion_faithfulness.csv")
|
|
53
|
+
with open(path, "r", encoding="utf-8") as f:
|
|
54
|
+
reader = csv.DictReader(f)
|
|
55
|
+
for row in reader:
|
|
56
|
+
if row.get("status") == "skipped_empty":
|
|
57
|
+
continue
|
|
58
|
+
entry = {
|
|
59
|
+
"input_id": int(row["input_id"]),
|
|
60
|
+
"category": row["category"],
|
|
61
|
+
"confidence_delta": abs(float(row["confidence_delta"])),
|
|
62
|
+
"label_flipped": row["label_flipped"] == "True",
|
|
63
|
+
}
|
|
64
|
+
if row.get("delta_top3") and row["delta_top3"] != "":
|
|
65
|
+
entry["delta_top3"] = abs(float(row["delta_top3"]))
|
|
66
|
+
entry["flipped_top3"] = row["flipped_top3"] == "True"
|
|
67
|
+
if row.get("delta_top5") and row["delta_top5"] != "":
|
|
68
|
+
entry["delta_top5"] = abs(float(row["delta_top5"]))
|
|
69
|
+
entry["flipped_top5"] = row["flipped_top5"] == "True"
|
|
70
|
+
rows.append(entry)
|
|
71
|
+
return rows
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def chart_jaccard(stability_rows, output_dir):
|
|
75
|
+
sorted_rows = sorted(stability_rows, key=lambda r: r["mean_jaccard"])
|
|
76
|
+
fig, ax = plt.subplots(figsize=(12, 6))
|
|
77
|
+
|
|
78
|
+
x = np.arange(len(sorted_rows))
|
|
79
|
+
colors = [CATEGORY_COLORS.get(r["category"], "#888") for r in sorted_rows]
|
|
80
|
+
means = [r["mean_jaccard"] for r in sorted_rows]
|
|
81
|
+
ci_lower = [r["mean_jaccard"] - r["ci_lower"] for r in sorted_rows]
|
|
82
|
+
ci_upper = [r["ci_upper"] - r["mean_jaccard"] for r in sorted_rows]
|
|
83
|
+
|
|
84
|
+
ax.bar(x, means, color=colors, edgecolor="white", linewidth=0.5)
|
|
85
|
+
ax.errorbar(x, means, yerr=[ci_lower, ci_upper], fmt="none", ecolor="#333333",
|
|
86
|
+
capsize=3, linewidth=1)
|
|
87
|
+
ax.axhline(y=STABILITY_THRESHOLD, color="#cc0000", linestyle="--",
|
|
88
|
+
linewidth=1.5, label=f"Threshold ({STABILITY_THRESHOLD})")
|
|
89
|
+
|
|
90
|
+
handles = [plt.Rectangle((0, 0), 1, 1, color=c) for c in CATEGORY_COLORS.values()]
|
|
91
|
+
labels = list(CATEGORY_LABELS.values())
|
|
92
|
+
ax.legend(handles, labels, loc="upper left", fontsize=8, ncol=2)
|
|
93
|
+
|
|
94
|
+
ax.set_xlabel("Inputs (sorted by stability)", fontsize=11)
|
|
95
|
+
ax.set_ylabel("Mean Jaccard Top-5 Overlap", fontsize=11)
|
|
96
|
+
ax.set_title("LIME Attribution Stability Across 5 Random Seeds", fontsize=13, fontweight="bold")
|
|
97
|
+
ax.set_ylim(0, 1.1)
|
|
98
|
+
ax.set_xticks([])
|
|
99
|
+
|
|
100
|
+
fig.tight_layout()
|
|
101
|
+
path = os.path.join(output_dir, "chart_jaccard_by_input.png")
|
|
102
|
+
fig.savefig(path, dpi=150)
|
|
103
|
+
plt.close(fig)
|
|
104
|
+
print(f"Saved: {path}")
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def chart_faithfulness(deletion_rows, output_dir):
|
|
108
|
+
cat_deltas = {}
|
|
109
|
+
for row in deletion_rows:
|
|
110
|
+
cat = row["category"]
|
|
111
|
+
if cat not in cat_deltas:
|
|
112
|
+
cat_deltas[cat] = {"top1": [], "top3": [], "top5": [], "flip1": [], "flip3": [], "flip5": []}
|
|
113
|
+
cat_deltas[cat]["top1"].append(row["confidence_delta"])
|
|
114
|
+
cat_deltas[cat]["flip1"].append(row["label_flipped"])
|
|
115
|
+
if "delta_top3" in row:
|
|
116
|
+
cat_deltas[cat]["top3"].append(row["delta_top3"])
|
|
117
|
+
cat_deltas[cat]["flip3"].append(row.get("flipped_top3", False))
|
|
118
|
+
if "delta_top5" in row:
|
|
119
|
+
cat_deltas[cat]["top5"].append(row["delta_top5"])
|
|
120
|
+
cat_deltas[cat]["flip5"].append(row.get("flipped_top5", False))
|
|
121
|
+
|
|
122
|
+
display_order = [
|
|
123
|
+
"strong_baselines", "distribution_style_shift", "edge_cases",
|
|
124
|
+
"ambiguity", "negation_minimal_pairs", "lexical_shortcuts",
|
|
125
|
+
]
|
|
126
|
+
cats_present = [c for c in display_order if c in cat_deltas]
|
|
127
|
+
if not cats_present:
|
|
128
|
+
cats_present = list(cat_deltas.keys())
|
|
129
|
+
|
|
130
|
+
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5.5))
|
|
131
|
+
x = np.arange(len(cats_present))
|
|
132
|
+
width = 0.25
|
|
133
|
+
|
|
134
|
+
means_1 = [np.mean(cat_deltas[c]["top1"]) if cat_deltas[c]["top1"] else 0 for c in cats_present]
|
|
135
|
+
means_3 = [np.mean(cat_deltas[c]["top3"]) if cat_deltas[c]["top3"] else 0 for c in cats_present]
|
|
136
|
+
means_5 = [np.mean(cat_deltas[c]["top5"]) if cat_deltas[c]["top5"] else 0 for c in cats_present]
|
|
137
|
+
|
|
138
|
+
ax1.bar(x - width, means_1, width, label="Top-1 removed", color="#264653")
|
|
139
|
+
ax1.bar(x, means_3, width, label="Top-3 removed", color="#2a9d8f")
|
|
140
|
+
ax1.bar(x + width, means_5, width, label="Top-5 removed", color="#e9c46a")
|
|
141
|
+
|
|
142
|
+
ax1.set_xlabel("Category", fontsize=10)
|
|
143
|
+
ax1.set_ylabel("Mean |Confidence Delta|", fontsize=10)
|
|
144
|
+
ax1.set_title("Impact of Removing Top-K Tokens", fontsize=12, fontweight="bold")
|
|
145
|
+
ax1.set_xticks(x)
|
|
146
|
+
ax1.set_xticklabels([CATEGORY_LABELS.get(c, c) for c in cats_present], rotation=30, ha="right", fontsize=8)
|
|
147
|
+
ax1.legend(fontsize=8)
|
|
148
|
+
ax1.set_ylim(0, 1.1)
|
|
149
|
+
|
|
150
|
+
flip1 = [np.mean(cat_deltas[c]["flip1"]) * 100 if cat_deltas[c]["flip1"] else 0 for c in cats_present]
|
|
151
|
+
flip3 = [np.mean(cat_deltas[c]["flip3"]) * 100 if cat_deltas[c]["flip3"] else 0 for c in cats_present]
|
|
152
|
+
flip5 = [np.mean(cat_deltas[c]["flip5"]) * 100 if cat_deltas[c]["flip5"] else 0 for c in cats_present]
|
|
153
|
+
|
|
154
|
+
ax2.bar(x - width, flip1, width, label="Top-1 removed", color="#264653")
|
|
155
|
+
ax2.bar(x, flip3, width, label="Top-3 removed", color="#2a9d8f")
|
|
156
|
+
ax2.bar(x + width, flip5, width, label="Top-5 removed", color="#e9c46a")
|
|
157
|
+
|
|
158
|
+
ax2.set_xlabel("Category", fontsize=10)
|
|
159
|
+
ax2.set_ylabel("Label Flip Rate (%)", fontsize=10)
|
|
160
|
+
ax2.set_title("Label Flip Rate by Deletion Depth", fontsize=12, fontweight="bold")
|
|
161
|
+
ax2.set_xticks(x)
|
|
162
|
+
ax2.set_xticklabels([CATEGORY_LABELS.get(c, c) for c in cats_present], rotation=30, ha="right", fontsize=8)
|
|
163
|
+
ax2.legend(fontsize=8)
|
|
164
|
+
ax2.set_ylim(0, 110)
|
|
165
|
+
|
|
166
|
+
fig.tight_layout()
|
|
167
|
+
path = os.path.join(output_dir, "chart_faithfulness_bars.png")
|
|
168
|
+
fig.savefig(path, dpi=150)
|
|
169
|
+
plt.close(fig)
|
|
170
|
+
print(f"Saved: {path}")
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
def generate_charts(output_dir):
|
|
174
|
+
print("Generating charts...")
|
|
175
|
+
stability = load_stability(output_dir)
|
|
176
|
+
deletion = load_deletion(output_dir)
|
|
177
|
+
chart_jaccard(stability, output_dir)
|
|
178
|
+
chart_faithfulness(deletion, output_dir)
|
|
179
|
+
print("Done.")
|
lime_audit/cli.py
ADDED
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
"""
|
|
2
|
+
CLI entry point for lime-audit.
|
|
3
|
+
|
|
4
|
+
Usage:
|
|
5
|
+
lime-audit run --model distilbert/distilbert-base-uncased-finetuned-sst-2-english
|
|
6
|
+
lime-audit run --model <name> --test-set my_inputs.json --output results/
|
|
7
|
+
lime-audit report --output results/
|
|
8
|
+
lime-audit charts --output results/
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
import argparse
|
|
12
|
+
import sys
|
|
13
|
+
|
|
14
|
+
from lime_audit import __version__
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
DEFAULT_MODEL = "distilbert/distilbert-base-uncased-finetuned-sst-2-english"
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def cmd_run(args):
|
|
21
|
+
from lime_audit.runner import run_audit
|
|
22
|
+
output_dir = run_audit(
|
|
23
|
+
model_name=args.model,
|
|
24
|
+
test_set_path=args.test_set if args.test_set else None,
|
|
25
|
+
output_dir=args.output if args.output else None,
|
|
26
|
+
)
|
|
27
|
+
if not args.skip_report:
|
|
28
|
+
from lime_audit.analyse import analyse
|
|
29
|
+
analyse(output_dir, args.model)
|
|
30
|
+
if not args.skip_charts:
|
|
31
|
+
from lime_audit.charts import generate_charts
|
|
32
|
+
generate_charts(output_dir)
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def cmd_report(args):
|
|
36
|
+
from lime_audit.analyse import analyse
|
|
37
|
+
analyse(args.output, args.model)
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def cmd_charts(args):
|
|
41
|
+
from lime_audit.charts import generate_charts
|
|
42
|
+
generate_charts(args.output)
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def main():
|
|
46
|
+
parser = argparse.ArgumentParser(
|
|
47
|
+
prog="lime-audit",
|
|
48
|
+
description="Audit LIME explanations for stability and faithfulness",
|
|
49
|
+
)
|
|
50
|
+
parser.add_argument("--version", action="version", version=f"%(prog)s {__version__}")
|
|
51
|
+
subparsers = parser.add_subparsers(dest="command", required=True)
|
|
52
|
+
|
|
53
|
+
run_parser = subparsers.add_parser("run", help="Run the full LIME audit")
|
|
54
|
+
run_parser.add_argument("--model", default=DEFAULT_MODEL, help="HuggingFace model name (default: distilbert-sst2)")
|
|
55
|
+
run_parser.add_argument("--test-set", default=None, help="Path to test set JSON (default: bundled 30 inputs)")
|
|
56
|
+
run_parser.add_argument("--output", default=None, help="Output directory (default: ./lime_audit_results)")
|
|
57
|
+
run_parser.add_argument("--skip-report", action="store_true", help="Skip analysis report after run")
|
|
58
|
+
run_parser.add_argument("--skip-charts", action="store_true", help="Skip chart generation after run")
|
|
59
|
+
run_parser.set_defaults(func=cmd_run)
|
|
60
|
+
|
|
61
|
+
report_parser = subparsers.add_parser("report", help="Generate analysis report from existing results")
|
|
62
|
+
report_parser.add_argument("--output", default="lime_audit_results", help="Results directory")
|
|
63
|
+
report_parser.add_argument("--model", default=DEFAULT_MODEL, help="Model name (for tokenizer mismatch)")
|
|
64
|
+
report_parser.set_defaults(func=cmd_report)
|
|
65
|
+
|
|
66
|
+
charts_parser = subparsers.add_parser("charts", help="Generate charts from existing results")
|
|
67
|
+
charts_parser.add_argument("--output", default="lime_audit_results", help="Results directory")
|
|
68
|
+
charts_parser.set_defaults(func=cmd_charts)
|
|
69
|
+
|
|
70
|
+
args = parser.parse_args()
|
|
71
|
+
args.func(args)
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
if __name__ == "__main__":
|
|
75
|
+
main()
|
lime_audit/config.py
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
"""Default configuration for lime-audit runs."""
|
|
2
|
+
|
|
3
|
+
LIME_NUM_SAMPLES = 300
|
|
4
|
+
LIME_NUM_FEATURES = 10
|
|
5
|
+
LIME_CLASS_NAMES = ["negative", "positive"]
|
|
6
|
+
|
|
7
|
+
RANDOM_SEEDS = [42, 123, 456, 789, 1024]
|
|
8
|
+
CANONICAL_SEED = 42
|
|
9
|
+
|
|
10
|
+
TOP_K_FOR_JACCARD = 5
|
|
11
|
+
|
|
12
|
+
STABILITY_THRESHOLD_JACCARD = 0.6
|
|
13
|
+
STABILITY_THRESHOLD_KENDALL = 0.5
|
|
14
|
+
FAITHFULNESS_THRESHOLD_DIRECTION = 0.7
|
|
@@ -0,0 +1,36 @@
|
|
|
1
|
+
{
|
|
2
|
+
"version": "1.0",
|
|
3
|
+
"description": "Default test set for LIME attribution audit. 30 inputs across 6 categories.",
|
|
4
|
+
"inputs": [
|
|
5
|
+
{"id": 1, "category": "negation_minimal_pairs", "text": "I am not entirely unhappy with this result."},
|
|
6
|
+
{"id": 2, "category": "negation_minimal_pairs", "text": "This is not good."},
|
|
7
|
+
{"id": 3, "category": "negation_minimal_pairs", "text": "This is good."},
|
|
8
|
+
{"id": 4, "category": "negation_minimal_pairs", "text": "I would not recommend this to anyone."},
|
|
9
|
+
{"id": 5, "category": "negation_minimal_pairs", "text": "Nothing about this experience was disappointing."},
|
|
10
|
+
{"id": 6, "category": "lexical_shortcuts", "text": "The movie was terrible but I loved every minute of it."},
|
|
11
|
+
{"id": 7, "category": "lexical_shortcuts", "text": "Excellent packaging, but the product itself is useless."},
|
|
12
|
+
{"id": 8, "category": "lexical_shortcuts", "text": "I love how badly this was designed."},
|
|
13
|
+
{"id": 9, "category": "lexical_shortcuts", "text": "Fine."},
|
|
14
|
+
{"id": 10, "category": "lexical_shortcuts", "text": "Absolutely phenomenal waste of my time."},
|
|
15
|
+
{"id": 11, "category": "ambiguity", "text": "It was okay, I guess."},
|
|
16
|
+
{"id": 12, "category": "ambiguity", "text": "That is one way to do it."},
|
|
17
|
+
{"id": 13, "category": "ambiguity", "text": "I have seen worse."},
|
|
18
|
+
{"id": 14, "category": "ambiguity", "text": "The service was exactly what I expected."},
|
|
19
|
+
{"id": 15, "category": "ambiguity", "text": "It is what it is."},
|
|
20
|
+
{"id": 16, "category": "distribution_style_shift", "text": "ngl this slaps fr fr no cap"},
|
|
21
|
+
{"id": 17, "category": "distribution_style_shift", "text": "The patient presents with acute exacerbation of chronic symptoms."},
|
|
22
|
+
{"id": 18, "category": "distribution_style_shift", "text": "Revenue increased 12% YoY driven by strong Q4 performance."},
|
|
23
|
+
{"id": 19, "category": "distribution_style_shift", "text": "lmaooo this is so bad its good"},
|
|
24
|
+
{"id": 20, "category": "distribution_style_shift", "text": "Per the attached memo, please advise on next steps."},
|
|
25
|
+
{"id": 21, "category": "strong_baselines", "text": "This is the best product I have ever purchased."},
|
|
26
|
+
{"id": 22, "category": "strong_baselines", "text": "Terrible experience, complete waste of money."},
|
|
27
|
+
{"id": 23, "category": "strong_baselines", "text": "I absolutely love everything about this."},
|
|
28
|
+
{"id": 24, "category": "strong_baselines", "text": "This is awful and I regret buying it."},
|
|
29
|
+
{"id": 25, "category": "strong_baselines", "text": "The quality exceeded all my expectations and I am thrilled."},
|
|
30
|
+
{"id": 26, "category": "edge_cases", "text": "good good good good good"},
|
|
31
|
+
{"id": 27, "category": "edge_cases", "text": "The the the the movie was great."},
|
|
32
|
+
{"id": 28, "category": "edge_cases", "text": "I think that maybe it could possibly be somewhat decent."},
|
|
33
|
+
{"id": 29, "category": "edge_cases", "text": "Amazing! Horrible! Amazing! Horrible!"},
|
|
34
|
+
{"id": 30, "category": "edge_cases", "text": " "}
|
|
35
|
+
]
|
|
36
|
+
}
|
lime_audit/metrics.py
ADDED
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
"""Stability and faithfulness metrics for LIME audit."""
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
from scipy import stats
|
|
5
|
+
from itertools import combinations
|
|
6
|
+
from transformers import AutoTokenizer
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def jaccard_top_k(tokens_a: list[str], tokens_b: list[str], k: int = 5) -> float:
|
|
10
|
+
set_a = set(tokens_a[:k])
|
|
11
|
+
set_b = set(tokens_b[:k])
|
|
12
|
+
if not set_a and not set_b:
|
|
13
|
+
return 1.0
|
|
14
|
+
if not set_a or not set_b:
|
|
15
|
+
return 0.0
|
|
16
|
+
return len(set_a & set_b) / len(set_a | set_b)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def kendall_tau_top_k(
|
|
20
|
+
ranking_a: list[tuple[str, float]],
|
|
21
|
+
ranking_b: list[tuple[str, float]],
|
|
22
|
+
) -> float:
|
|
23
|
+
tokens_a = {t: i for i, (t, _) in enumerate(ranking_a)}
|
|
24
|
+
tokens_b = {t: i for i, (t, _) in enumerate(ranking_b)}
|
|
25
|
+
common = sorted(set(tokens_a.keys()) & set(tokens_b.keys()))
|
|
26
|
+
if len(common) < 2:
|
|
27
|
+
return float("nan")
|
|
28
|
+
ranks_a = [tokens_a[t] for t in common]
|
|
29
|
+
ranks_b = [tokens_b[t] for t in common]
|
|
30
|
+
tau, _ = stats.kendalltau(ranks_a, ranks_b)
|
|
31
|
+
return tau
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def confidence_delta(original_pos: float, modified_pos: float) -> float:
|
|
35
|
+
return modified_pos - original_pos
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def faithfulness_direction_correct(token_weight: float, conf_delta: float) -> bool:
|
|
39
|
+
if abs(token_weight) < 1e-6 or abs(conf_delta) < 1e-6:
|
|
40
|
+
return True
|
|
41
|
+
return (token_weight > 0 and conf_delta < 0) or (token_weight < 0 and conf_delta > 0)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def bootstrap_ci(
|
|
45
|
+
values: list[float],
|
|
46
|
+
n_bootstrap: int = 10000,
|
|
47
|
+
alpha: float = 0.05,
|
|
48
|
+
rng_seed: int = 42,
|
|
49
|
+
) -> tuple[float, float]:
|
|
50
|
+
rng = np.random.RandomState(rng_seed)
|
|
51
|
+
arr = np.array(values)
|
|
52
|
+
boot_means = np.array([
|
|
53
|
+
np.mean(rng.choice(arr, size=len(arr), replace=True))
|
|
54
|
+
for _ in range(n_bootstrap)
|
|
55
|
+
])
|
|
56
|
+
lower = float(np.percentile(boot_means, 100 * alpha / 2))
|
|
57
|
+
upper = float(np.percentile(boot_means, 100 * (1 - alpha / 2)))
|
|
58
|
+
return lower, upper
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def count_tokenizer_mismatch(text: str, model_name: str) -> dict:
|
|
62
|
+
lime_tokens = text.split()
|
|
63
|
+
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
|
64
|
+
wp_tokens = tokenizer.tokenize(text)
|
|
65
|
+
return {
|
|
66
|
+
"lime_token_count": len(lime_tokens),
|
|
67
|
+
"wordpiece_token_count": len(wp_tokens),
|
|
68
|
+
"token_mismatch": len(wp_tokens) - len(lime_tokens),
|
|
69
|
+
}
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def compute_pairwise_stability(
|
|
73
|
+
seed_results: dict[int, list[tuple[str, float]]],
|
|
74
|
+
k: int = 5,
|
|
75
|
+
) -> dict:
|
|
76
|
+
seeds = sorted(seed_results.keys())
|
|
77
|
+
jaccards = []
|
|
78
|
+
taus = []
|
|
79
|
+
|
|
80
|
+
for s1, s2 in combinations(seeds, 2):
|
|
81
|
+
tokens_a = [t for t, _ in seed_results[s1]]
|
|
82
|
+
tokens_b = [t for t, _ in seed_results[s2]]
|
|
83
|
+
jaccards.append(jaccard_top_k(tokens_a, tokens_b, k))
|
|
84
|
+
taus.append(kendall_tau_top_k(seed_results[s1], seed_results[s2]))
|
|
85
|
+
|
|
86
|
+
top1_tokens = [seed_results[s][0][0] if seed_results[s] else "" for s in seeds]
|
|
87
|
+
|
|
88
|
+
return {
|
|
89
|
+
"mean_jaccard": float(np.mean(jaccards)),
|
|
90
|
+
"min_jaccard": float(np.min(jaccards)),
|
|
91
|
+
"mean_kendall_tau": float(np.nanmean(taus)),
|
|
92
|
+
"min_kendall_tau": float(np.nanmin(taus)) if not all(np.isnan(taus)) else float("nan"),
|
|
93
|
+
"top1_unanimous": len(set(top1_tokens)) == 1 and top1_tokens[0] != "",
|
|
94
|
+
"all_jaccards": jaccards,
|
|
95
|
+
"all_taus": taus,
|
|
96
|
+
}
|
lime_audit/runner.py
ADDED
|
@@ -0,0 +1,297 @@
|
|
|
1
|
+
"""
|
|
2
|
+
LIME Audit Runner - runs LIME on test inputs, measures stability and faithfulness.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
import csv
|
|
6
|
+
import json
|
|
7
|
+
import os
|
|
8
|
+
import platform
|
|
9
|
+
import sys
|
|
10
|
+
import time
|
|
11
|
+
from datetime import datetime, timezone
|
|
12
|
+
|
|
13
|
+
import numpy as np
|
|
14
|
+
from lime.lime_text import LimeTextExplainer
|
|
15
|
+
from transformers import pipeline as hf_pipeline
|
|
16
|
+
|
|
17
|
+
from lime_audit.config import (
|
|
18
|
+
CANONICAL_SEED,
|
|
19
|
+
LIME_CLASS_NAMES,
|
|
20
|
+
LIME_NUM_FEATURES,
|
|
21
|
+
LIME_NUM_SAMPLES,
|
|
22
|
+
RANDOM_SEEDS,
|
|
23
|
+
)
|
|
24
|
+
from lime_audit.metrics import faithfulness_direction_correct
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def _get_default_test_set_path():
|
|
28
|
+
return os.path.join(os.path.dirname(__file__), "default_test_set.json")
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def load_test_set(path: str) -> list[dict]:
|
|
32
|
+
with open(path, "r", encoding="utf-8") as f:
|
|
33
|
+
data = json.load(f)
|
|
34
|
+
return data["inputs"]
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def load_model(model_name: str):
|
|
38
|
+
print(f"Loading model: {model_name}")
|
|
39
|
+
model = hf_pipeline("text-classification", model=model_name, top_k=None)
|
|
40
|
+
print("Model loaded.")
|
|
41
|
+
return model
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def get_positive_score(pipeline_output: list) -> float:
|
|
45
|
+
for item in pipeline_output[0]:
|
|
46
|
+
if "pos" in item["label"].lower():
|
|
47
|
+
return item["score"]
|
|
48
|
+
return 1.0 - max(item["score"] for item in pipeline_output[0])
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def make_predict_fn(model):
|
|
52
|
+
def predict_proba(texts: list[str]) -> np.ndarray:
|
|
53
|
+
results = []
|
|
54
|
+
for text in texts:
|
|
55
|
+
output = model(text, truncation=True, max_length=512)
|
|
56
|
+
pos = get_positive_score(output)
|
|
57
|
+
results.append([1.0 - pos, pos])
|
|
58
|
+
return np.array(results)
|
|
59
|
+
return predict_proba
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def run_lime_single(text, predict_fn, seed, num_samples=LIME_NUM_SAMPLES, num_features=LIME_NUM_FEATURES):
|
|
63
|
+
np.random.seed(seed)
|
|
64
|
+
explainer = LimeTextExplainer(class_names=LIME_CLASS_NAMES, random_state=seed)
|
|
65
|
+
explanation = explainer.explain_instance(text, predict_fn, num_features=num_features, num_samples=num_samples)
|
|
66
|
+
token_weights = explanation.as_list()
|
|
67
|
+
sorted_by_abs = sorted(token_weights, key=lambda x: abs(x[1]), reverse=True)
|
|
68
|
+
return {"tokens": sorted_by_abs, "lime_score": explanation.score, "intercept": explanation.intercept.get(1, 0.0)}
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def run_deletion_test(text, top_token, predict_fn):
|
|
72
|
+
words = text.split()
|
|
73
|
+
masked_words = [w for w in words if w != top_token]
|
|
74
|
+
if not masked_words or masked_words == words:
|
|
75
|
+
for i, w in enumerate(words):
|
|
76
|
+
if top_token.lower() in w.lower():
|
|
77
|
+
masked_words = words[:i] + words[i + 1:]
|
|
78
|
+
break
|
|
79
|
+
if not masked_words:
|
|
80
|
+
return {"modified_text": "", "modified_label": "error", "modified_positive_score": float("nan")}
|
|
81
|
+
masked_text = " ".join(masked_words)
|
|
82
|
+
proba = predict_fn([masked_text])
|
|
83
|
+
modified_pos = float(proba[0][1])
|
|
84
|
+
return {"modified_text": masked_text, "modified_label": "positive" if modified_pos >= 0.5 else "negative", "modified_positive_score": modified_pos}
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def run_deletion_topk(text, top_tokens, predict_fn):
|
|
88
|
+
words = text.split()
|
|
89
|
+
token_set = set(top_tokens)
|
|
90
|
+
masked_words = [w for w in words if w not in token_set]
|
|
91
|
+
if not masked_words:
|
|
92
|
+
return {"modified_text": "", "modified_label": "error", "modified_positive_score": float("nan")}
|
|
93
|
+
masked_text = " ".join(masked_words)
|
|
94
|
+
proba = predict_fn([masked_text])
|
|
95
|
+
modified_pos = float(proba[0][1])
|
|
96
|
+
return {"modified_text": masked_text, "modified_label": "positive" if modified_pos >= 0.5 else "negative", "modified_positive_score": modified_pos}
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def get_environment_info(model_name):
|
|
100
|
+
import lime
|
|
101
|
+
import scipy
|
|
102
|
+
import sklearn
|
|
103
|
+
import torch
|
|
104
|
+
import transformers
|
|
105
|
+
return {
|
|
106
|
+
"model": model_name,
|
|
107
|
+
"lime_version": getattr(lime, "__version__", "0.2.0.1"),
|
|
108
|
+
"torch_version": torch.__version__,
|
|
109
|
+
"transformers_version": transformers.__version__,
|
|
110
|
+
"numpy_version": np.__version__,
|
|
111
|
+
"scipy_version": scipy.__version__,
|
|
112
|
+
"sklearn_version": sklearn.__version__,
|
|
113
|
+
"python_version": sys.version,
|
|
114
|
+
"platform": platform.platform(),
|
|
115
|
+
"seeds": RANDOM_SEEDS,
|
|
116
|
+
"num_samples": LIME_NUM_SAMPLES,
|
|
117
|
+
"num_features": LIME_NUM_FEATURES,
|
|
118
|
+
}
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
RAW_FIELDNAMES = (
|
|
122
|
+
["input_id", "category", "text", "seed", "base_label", "base_confidence", "base_positive_score"]
|
|
123
|
+
+ [f"token_{i}" for i in range(1, LIME_NUM_FEATURES + 1)]
|
|
124
|
+
+ [f"weight_{i}" for i in range(1, LIME_NUM_FEATURES + 1)]
|
|
125
|
+
+ ["lime_score", "duration_ms"]
|
|
126
|
+
)
|
|
127
|
+
|
|
128
|
+
DELETION_FIELDNAMES = [
|
|
129
|
+
"input_id", "category", "text", "status",
|
|
130
|
+
"removed_token", "token_lime_weight",
|
|
131
|
+
"original_label", "original_positive_score",
|
|
132
|
+
"modified_label", "modified_positive_score",
|
|
133
|
+
"confidence_delta", "label_flipped", "direction_correct",
|
|
134
|
+
"removed_top3", "delta_top3", "flipped_top3",
|
|
135
|
+
"removed_top5", "delta_top5", "flipped_top5",
|
|
136
|
+
]
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
def run_audit(model_name: str, test_set_path: str = None, output_dir: str = None):
|
|
140
|
+
if test_set_path is None:
|
|
141
|
+
test_set_path = _get_default_test_set_path()
|
|
142
|
+
if output_dir is None:
|
|
143
|
+
output_dir = os.path.join(os.getcwd(), "lime_audit_results")
|
|
144
|
+
|
|
145
|
+
os.makedirs(output_dir, exist_ok=True)
|
|
146
|
+
raw_csv = os.path.join(output_dir, "raw_attributions.csv")
|
|
147
|
+
del_csv = os.path.join(output_dir, "deletion_faithfulness.csv")
|
|
148
|
+
|
|
149
|
+
inputs = load_test_set(test_set_path)
|
|
150
|
+
print(f"Loaded {len(inputs)} inputs")
|
|
151
|
+
|
|
152
|
+
model = load_model(model_name)
|
|
153
|
+
predict_fn = make_predict_fn(model)
|
|
154
|
+
env_info = get_environment_info(model_name)
|
|
155
|
+
|
|
156
|
+
start_total = time.perf_counter()
|
|
157
|
+
skipped_inputs = []
|
|
158
|
+
canonical_results = {}
|
|
159
|
+
|
|
160
|
+
with open(raw_csv, "w", newline="", encoding="utf-8") as raw_f:
|
|
161
|
+
raw_writer = csv.DictWriter(raw_f, fieldnames=RAW_FIELDNAMES)
|
|
162
|
+
raw_writer.writeheader()
|
|
163
|
+
|
|
164
|
+
for inp in inputs:
|
|
165
|
+
input_id = inp["id"]
|
|
166
|
+
text = inp["text"]
|
|
167
|
+
category = inp.get("category", "uncategorized")
|
|
168
|
+
|
|
169
|
+
if not text.strip():
|
|
170
|
+
skipped_inputs.append({"input_id": input_id, "category": category, "reason": "empty_or_whitespace"})
|
|
171
|
+
print(f" [{input_id}] SKIP empty input")
|
|
172
|
+
continue
|
|
173
|
+
|
|
174
|
+
base_proba = predict_fn([text])
|
|
175
|
+
base_pos = float(base_proba[0][1])
|
|
176
|
+
base_label = "positive" if base_pos >= 0.5 else "negative"
|
|
177
|
+
base_conf = base_pos if base_label == "positive" else 1.0 - base_pos
|
|
178
|
+
print(f" [{input_id}] {base_label} ({base_conf:.4f})")
|
|
179
|
+
|
|
180
|
+
for seed in RANDOM_SEEDS:
|
|
181
|
+
start_call = time.perf_counter()
|
|
182
|
+
try:
|
|
183
|
+
result = run_lime_single(text, predict_fn, seed)
|
|
184
|
+
except Exception as e:
|
|
185
|
+
print(f" Seed {seed}: ERROR {e}")
|
|
186
|
+
continue
|
|
187
|
+
|
|
188
|
+
duration_ms = round((time.perf_counter() - start_call) * 1000)
|
|
189
|
+
tokens = result["tokens"]
|
|
190
|
+
|
|
191
|
+
row = {
|
|
192
|
+
"input_id": input_id, "category": category, "text": text,
|
|
193
|
+
"seed": seed, "base_label": base_label,
|
|
194
|
+
"base_confidence": round(base_conf, 6),
|
|
195
|
+
"base_positive_score": round(base_pos, 6),
|
|
196
|
+
"lime_score": round(result["lime_score"], 6),
|
|
197
|
+
"duration_ms": duration_ms,
|
|
198
|
+
}
|
|
199
|
+
for i in range(LIME_NUM_FEATURES):
|
|
200
|
+
if i < len(tokens):
|
|
201
|
+
row[f"token_{i + 1}"] = tokens[i][0]
|
|
202
|
+
row[f"weight_{i + 1}"] = round(tokens[i][1], 6)
|
|
203
|
+
else:
|
|
204
|
+
row[f"token_{i + 1}"] = ""
|
|
205
|
+
row[f"weight_{i + 1}"] = 0.0
|
|
206
|
+
raw_writer.writerow(row)
|
|
207
|
+
raw_f.flush()
|
|
208
|
+
|
|
209
|
+
if seed == CANONICAL_SEED:
|
|
210
|
+
canonical_results[input_id] = {
|
|
211
|
+
"tokens": tokens, "base_label": base_label,
|
|
212
|
+
"base_pos": base_pos, "base_conf": base_conf,
|
|
213
|
+
}
|
|
214
|
+
|
|
215
|
+
print(f"\nRunning deletion faithfulness tests...")
|
|
216
|
+
|
|
217
|
+
with open(del_csv, "w", newline="", encoding="utf-8") as del_f:
|
|
218
|
+
del_writer = csv.DictWriter(del_f, fieldnames=DELETION_FIELDNAMES)
|
|
219
|
+
del_writer.writeheader()
|
|
220
|
+
|
|
221
|
+
for inp in inputs:
|
|
222
|
+
input_id = inp["id"]
|
|
223
|
+
text = inp["text"]
|
|
224
|
+
category = inp.get("category", "uncategorized")
|
|
225
|
+
|
|
226
|
+
if not text.strip():
|
|
227
|
+
row = {"input_id": input_id, "category": category, "text": repr(text), "status": "skipped_empty"}
|
|
228
|
+
for field in DELETION_FIELDNAMES:
|
|
229
|
+
row.setdefault(field, "")
|
|
230
|
+
del_writer.writerow(row)
|
|
231
|
+
del_f.flush()
|
|
232
|
+
continue
|
|
233
|
+
|
|
234
|
+
if input_id not in canonical_results:
|
|
235
|
+
continue
|
|
236
|
+
|
|
237
|
+
canon = canonical_results[input_id]
|
|
238
|
+
if not canon["tokens"]:
|
|
239
|
+
continue
|
|
240
|
+
|
|
241
|
+
top_token = canon["tokens"][0][0]
|
|
242
|
+
top_weight = canon["tokens"][0][1]
|
|
243
|
+
deletion = run_deletion_test(text, top_token, predict_fn)
|
|
244
|
+
|
|
245
|
+
if deletion["modified_label"] == "error":
|
|
246
|
+
continue
|
|
247
|
+
|
|
248
|
+
delta = deletion["modified_positive_score"] - canon["base_pos"]
|
|
249
|
+
flipped = deletion["modified_label"] != canon["base_label"]
|
|
250
|
+
dir_correct = faithfulness_direction_correct(top_weight, delta)
|
|
251
|
+
|
|
252
|
+
top3_tokens = [t for t, _ in canon["tokens"][:3]]
|
|
253
|
+
del3 = run_deletion_topk(text, top3_tokens, predict_fn)
|
|
254
|
+
delta3 = flipped3 = ""
|
|
255
|
+
if del3["modified_label"] != "error":
|
|
256
|
+
delta3 = round(del3["modified_positive_score"] - canon["base_pos"], 6)
|
|
257
|
+
flipped3 = del3["modified_label"] != canon["base_label"]
|
|
258
|
+
|
|
259
|
+
top5_tokens = [t for t, _ in canon["tokens"][:5]]
|
|
260
|
+
del5 = run_deletion_topk(text, top5_tokens, predict_fn)
|
|
261
|
+
delta5 = flipped5 = ""
|
|
262
|
+
if del5["modified_label"] != "error":
|
|
263
|
+
delta5 = round(del5["modified_positive_score"] - canon["base_pos"], 6)
|
|
264
|
+
flipped5 = del5["modified_label"] != canon["base_label"]
|
|
265
|
+
|
|
266
|
+
row = {
|
|
267
|
+
"input_id": input_id, "category": category, "text": text,
|
|
268
|
+
"status": "tested",
|
|
269
|
+
"removed_token": top_token,
|
|
270
|
+
"token_lime_weight": round(top_weight, 6),
|
|
271
|
+
"original_label": canon["base_label"],
|
|
272
|
+
"original_positive_score": round(canon["base_pos"], 6),
|
|
273
|
+
"modified_label": deletion["modified_label"],
|
|
274
|
+
"modified_positive_score": round(deletion["modified_positive_score"], 6),
|
|
275
|
+
"confidence_delta": round(delta, 6),
|
|
276
|
+
"label_flipped": flipped,
|
|
277
|
+
"direction_correct": dir_correct,
|
|
278
|
+
"removed_top3": "|".join(top3_tokens),
|
|
279
|
+
"delta_top3": delta3, "flipped_top3": flipped3,
|
|
280
|
+
"removed_top5": "|".join(top5_tokens),
|
|
281
|
+
"delta_top5": delta5, "flipped_top5": flipped5,
|
|
282
|
+
}
|
|
283
|
+
del_writer.writerow(row)
|
|
284
|
+
del_f.flush()
|
|
285
|
+
|
|
286
|
+
total_time = time.perf_counter() - start_total
|
|
287
|
+
|
|
288
|
+
env_info["run_timestamp"] = datetime.now(timezone.utc).isoformat()
|
|
289
|
+
env_info["total_duration_seconds"] = round(total_time, 1)
|
|
290
|
+
env_info["skipped_inputs"] = skipped_inputs
|
|
291
|
+
env_path = os.path.join(output_dir, "environment.json")
|
|
292
|
+
with open(env_path, "w", encoding="utf-8") as f:
|
|
293
|
+
json.dump(env_info, f, indent=2)
|
|
294
|
+
|
|
295
|
+
print(f"\nAudit complete in {total_time:.1f}s ({total_time / 60:.1f} min)")
|
|
296
|
+
print(f"Results in: {output_dir}")
|
|
297
|
+
return output_dir
|
|
@@ -0,0 +1,248 @@
|
|
|
1
|
+
Metadata-Version: 2.5
|
|
2
|
+
Name: lime-audit
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: Audit LIME explanations for stability and faithfulness on any HuggingFace text classifier
|
|
5
|
+
Project-URL: Homepage, https://github.com/parshvi1508/XAI_Forensic
|
|
6
|
+
Project-URL: Documentation, https://github.com/parshvi1508/XAI_Forensic#readme
|
|
7
|
+
Project-URL: Issues, https://github.com/parshvi1508/XAI_Forensic/issues
|
|
8
|
+
Author-email: Parshvi Jain <parshvijain1508@gmail.com>
|
|
9
|
+
License: MIT
|
|
10
|
+
License-File: LICENSE
|
|
11
|
+
Keywords: LIME,NLP,XAI,audit,explainability,transformers
|
|
12
|
+
Classifier: Development Status :: 3 - Alpha
|
|
13
|
+
Classifier: Intended Audience :: Science/Research
|
|
14
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
15
|
+
Classifier: Programming Language :: Python :: 3
|
|
16
|
+
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
|
17
|
+
Requires-Python: >=3.10
|
|
18
|
+
Requires-Dist: lime>=0.2
|
|
19
|
+
Requires-Dist: matplotlib>=3.7
|
|
20
|
+
Requires-Dist: numpy>=1.24
|
|
21
|
+
Requires-Dist: scikit-learn>=1.2
|
|
22
|
+
Requires-Dist: scipy>=1.10
|
|
23
|
+
Requires-Dist: torch>=2.0
|
|
24
|
+
Requires-Dist: transformers>=4.30
|
|
25
|
+
Provides-Extra: dev
|
|
26
|
+
Requires-Dist: build; extra == 'dev'
|
|
27
|
+
Requires-Dist: pytest; extra == 'dev'
|
|
28
|
+
Requires-Dist: twine; extra == 'dev'
|
|
29
|
+
Description-Content-Type: text/markdown
|
|
30
|
+
|
|
31
|
+
# XAI Forensics
|
|
32
|
+
|
|
33
|
+
A diagnostic tool for evaluating when LIME token attributions are trustworthy on transformer sentiment classifiers. Runs three independent forensic checks (attribution stability, counterfactual faithfulness, cross-model agreement) and validates results against a 30-input pre-registered audit.
|
|
34
|
+
|
|
35
|
+
**Key finding:** high-confidence predictions (>95%) produce the least stable attributions (mean Jaccard 0.46 on strong baselines), while lexical-shortcut inputs are perfectly faithful (Jaccard 1.0, 100% flip rate). This tool shows you when LIME is signal vs. noise.
|
|
36
|
+
|
|
37
|
+
Sentiment classification is the controlled test task. The project evaluates the explanation method, not the classifier output.
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
## Live Demo
|
|
42
|
+
|
|
43
|
+
| Component | Link |
|
|
44
|
+
|-----------|------|
|
|
45
|
+
| Frontend | [xai-forensic.vercel.app](https://xai-forensic.vercel.app/) |
|
|
46
|
+
| Backend API | [jainparshvi-xai-forensics-backend.hf.space](https://jainparshvi-xai-forensics-backend.hf.space) |
|
|
47
|
+
| API Docs (Swagger) | [/docs](https://jainparshvi-xai-forensics-backend.hf.space/docs) |
|
|
48
|
+
|
|
49
|
+
**Demo input:** `I am not entirely unhappy with this result.`
|
|
50
|
+
|
|
51
|
+
This sentence contains a double negation that causes the two models to genuinely disagree.
|
|
52
|
+
|
|
53
|
+
> The backend runs on Hugging Face Spaces free tier. If the Space has been idle, the first request triggers a cold start (30-60 seconds) while models download. Subsequent requests are faster.
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+

|
|
58
|
+

|
|
59
|
+

|
|
60
|
+

|
|
61
|
+

|
|
62
|
+

|
|
63
|
+
|
|
64
|
+
## What This Does
|
|
65
|
+
|
|
66
|
+
XAI Forensics runs three independent diagnostic checks on any short English text:
|
|
67
|
+
|
|
68
|
+
1. **WHY (Attribution Stability)** - Tests whether LIME can produce a stable, reproducible token ranking for a given prediction. Seeds the random state for deterministic output.
|
|
69
|
+
2. **FLIP (Counterfactual Faithfulness)** - Tests whether the tokens LIME identifies as important are actually causally influential. Removes the top-attributed word and measures the real confidence shift.
|
|
70
|
+
3. **DISAGREE (Cross-Model Consistency)** - Tests whether the prediction itself is domain-stable enough to warrant attribution analysis. Compares DistilBERT-SST2 against Twitter-RoBERTa.
|
|
71
|
+
|
|
72
|
+
Each check targets a different failure mode of LIME: instability under re-sampling, unfaithfulness to the model's actual reasoning, and domain sensitivity of the underlying prediction.
|
|
73
|
+
|
|
74
|
+
### Audit Results
|
|
75
|
+
|
|
76
|
+
A pre-registered 30-input audit (5 seeds each, 150 total LIME runs) found:
|
|
77
|
+
- Overall mean Jaccard stability: **0.81** (passes 0.6 threshold)
|
|
78
|
+
- Deletion faithfulness direction correct: **89.3%** (passes 70% threshold)
|
|
79
|
+
- Strong baselines category (high confidence): **Jaccard 0.67** (lowest category, driven by redundant evidence)
|
|
80
|
+
- Lexical shortcuts category: **100% label flip rate** (highest faithfulness)
|
|
81
|
+
|
|
82
|
+
Full interactive results available at [/audit](https://xai-forensic.vercel.app/audit).
|
|
83
|
+
|
|
84
|
+
## How It Works
|
|
85
|
+
|
|
86
|
+
### WHY: LIME Token Attribution
|
|
87
|
+
|
|
88
|
+
Uses [LIME](https://arxiv.org/abs/1602.04938) (Ribeiro et al., 2016) for model-agnostic local explanations. LIME generates perturbed versions of the input text, reruns the classifier on each perturbation, and fits a local linear model to estimate which tokens most influenced the prediction.
|
|
89
|
+
|
|
90
|
+
- Runs 300 perturbation samples per explanation
|
|
91
|
+
- Returns the top 10 tokens with signed weights (positive = pushes toward positive class)
|
|
92
|
+
- SHAP was considered and rejected: slower for transformers, expensive on free-tier CPU
|
|
93
|
+
- Attention weights were rejected as explanations per Jain and Wallace (2019)
|
|
94
|
+
|
|
95
|
+
### FLIP: Counterfactual Word Removal
|
|
96
|
+
|
|
97
|
+
Removes each word one at a time, reruns inference, and finds the word whose removal causes the largest confidence shift. Shows the full before-and-after comparison: original label, modified label, confidence delta, and whether the verdict changed.
|
|
98
|
+
|
|
99
|
+
- Greedy O(n) search over words in the input
|
|
100
|
+
- Deterministic, no generation model required
|
|
101
|
+
- Word removal can create ungrammatical text (documented limitation)
|
|
102
|
+
- Does not always flip the label on highly confident predictions
|
|
103
|
+
|
|
104
|
+
### DISAGREE: Dual-Model Divergence
|
|
105
|
+
|
|
106
|
+
Runs the same text through both models and computes the absolute difference in their positive-class confidence scores. A high divergence score means the models have different views on the same text, which reveals linguistic ambiguity across training domains.
|
|
107
|
+
|
|
108
|
+
- Divergence = abs(positive_score_A - positive_score_B)
|
|
109
|
+
- This is an interpretable confidence delta, not a formal divergence metric like KL divergence
|
|
110
|
+
- KL divergence was considered and rejected: harder to interpret for non-technical audiences
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
## Models Used
|
|
114
|
+
|
|
115
|
+
| Model | Training Data | Strength |
|
|
116
|
+
|-------|---------------|----------|
|
|
117
|
+
| [distilbert-base-uncased-finetuned-sst-2-english](https://huggingface.co/distilbert-base-uncased-finetuned-sst-2-english) | SST-2 movie reviews | Formal, structured English |
|
|
118
|
+
| [cardiffnlp/twitter-roberta-base-sentiment-latest](https://huggingface.co/cardiffnlp/twitter-roberta-base-sentiment-latest) | 124M tweets | Sarcasm, slang, informal tone |
|
|
119
|
+
|
|
120
|
+
These two models are chosen because they genuinely disagree on informal or ambiguous language. DistilBERT expects clean, formal text. Twitter-RoBERTa handles internet language better. This domain mismatch makes the DISAGREE panel meaningful rather than artificial.
|
|
121
|
+
|
|
122
|
+
## Architecture
|
|
123
|
+
|
|
124
|
+
```
|
|
125
|
+
Frontend (Vercel) Backend (HF Spaces Docker)
|
|
126
|
+
Next.js + Tailwind FastAPI + PyTorch
|
|
127
|
+
| |
|
|
128
|
+
|--- POST /why ------------->|--- LIME (300 perturbations)
|
|
129
|
+
|--- POST /flip ------------->|--- Greedy word removal
|
|
130
|
+
|--- POST /disagree ---------->|--- Dual model inference
|
|
131
|
+
| |
|
|
132
|
+
|<---- JSON responses ---------|
|
|
133
|
+
```
|
|
134
|
+
|
|
135
|
+
- Frontend and backend are fully decoupled
|
|
136
|
+
- All endpoints accept `{"text": "..."}` and return structured JSON
|
|
137
|
+
- Frontend calls all three endpoints in parallel using `Promise.allSettled`
|
|
138
|
+
- Models load once at container startup, not per request
|
|
139
|
+
|
|
140
|
+
## Run it Locally
|
|
141
|
+
|
|
142
|
+
### Backend
|
|
143
|
+
|
|
144
|
+
```bash
|
|
145
|
+
cd backend
|
|
146
|
+
python -m venv venv
|
|
147
|
+
source venv/bin/activate # Windows: venv\Scripts\activate
|
|
148
|
+
pip install -r requirements.txt
|
|
149
|
+
uvicorn main:app --reload --port 8000
|
|
150
|
+
```
|
|
151
|
+
|
|
152
|
+
First run downloads models (~500MB total). Subsequent starts use cached weights.
|
|
153
|
+
|
|
154
|
+
### Frontend
|
|
155
|
+
|
|
156
|
+
```bash
|
|
157
|
+
cd frontend
|
|
158
|
+
npm install
|
|
159
|
+
npm run dev
|
|
160
|
+
```
|
|
161
|
+
|
|
162
|
+
Frontend runs at `http://localhost:3000` and calls the backend at `http://localhost:8000` by default. To change the backend URL, set `NEXT_PUBLIC_API_URL` in `frontend/.env.local`.
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
## API Endpoints
|
|
166
|
+
|
|
167
|
+
All endpoints accept POST with `Content-Type: application/json` and a body of `{"text": "your input"}`.
|
|
168
|
+
|
|
169
|
+
| Endpoint | Returns |
|
|
170
|
+
|----------|---------|
|
|
171
|
+
| `POST /why` | Predicted label, confidence, and top 10 token weights from LIME |
|
|
172
|
+
| `POST /flip` | Original and modified labels, removed word, confidence delta, verdict changed status |
|
|
173
|
+
| `POST /disagree` | Both model predictions, divergence score, agreement status |
|
|
174
|
+
| `GET /` | Health check: `{"status": "ok"}` |
|
|
175
|
+
| `GET /docs` | Interactive Swagger API documentation |
|
|
176
|
+
|
|
177
|
+
Input is limited to 1000 characters. Requests with longer text return HTTP 400.
|
|
178
|
+
|
|
179
|
+
|
|
180
|
+
## Latency Notes
|
|
181
|
+
|
|
182
|
+
| Endpoint | Approx. CPU time | Why |
|
|
183
|
+
|----------|-------------------|-----|
|
|
184
|
+
| `/why` | 15-45 seconds | LIME runs 300 model inference calls per explanation |
|
|
185
|
+
| `/flip` | 2-10 seconds | One model call per word in the input |
|
|
186
|
+
| `/disagree` | Under 1 second | Two forward passes |
|
|
187
|
+
|
|
188
|
+
LIME is the main latency source. CPU inference is used because deployment is on the Hugging Face free tier, which does not guarantee GPU availability. This is acceptable for a demo tool with short text inputs.
|
|
189
|
+
|
|
190
|
+
|
|
191
|
+
## Latency Benchmark
|
|
192
|
+
|
|
193
|
+
The backend runs on Hugging Face Spaces free-tier CPU. The table below reports median endpoint runtime across 3 runs after one warmup request. Timings are approximate because free-tier CPU performance varies.
|
|
194
|
+
|
|
195
|
+
| Words | WHY median | FLIP median | DISAGREE median |
|
|
196
|
+
|---:|---:|---:|---:|
|
|
197
|
+
| 5 | 4.5s | 414ms | 98ms |
|
|
198
|
+
| 10 | 5.8s | 270ms | 75ms |
|
|
199
|
+
| 20 | 6.1s | 475ms | 72ms |
|
|
200
|
+
| 50 | 8.3s | 4.2s | 104ms |
|
|
201
|
+
|
|
202
|
+
WHY is slowest because LIME generates perturbed versions of the input and reruns model inference many times. FLIP scales with word count because it removes candidate words one at a time and reruns inference. DISAGREE is fastest because it only runs two model forward passes.
|
|
203
|
+
|
|
204
|
+
The benchmark can be reproduced with `scripts/benchmark_latency.py`.
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
## Security and Cost
|
|
208
|
+
|
|
209
|
+
- No paid APIs. Both models are public Hugging Face models.
|
|
210
|
+
- No database. Nothing is stored.
|
|
211
|
+
- No authentication. This is a public demo tool.
|
|
212
|
+
- No user data is collected, logged, or persisted.
|
|
213
|
+
- CORS allows all origins (appropriate for a public demo with no sensitive operations).
|
|
214
|
+
- Input capped at 1000 characters to prevent LIME timeouts on free-tier CPU.
|
|
215
|
+
|
|
216
|
+
|
|
217
|
+
## Known Behavior and Limitations
|
|
218
|
+
|
|
219
|
+
Single-word inputs are not ideal for this tool. LIME works by perturbing parts of the input and observing prediction changes. With only one word, there is very little structure to perturb, so the explanation can be unstable or uninformative. In testing, short inputs such as "Fine." can produce domain-sensitive behavior because different models interpret minimal context differently.
|
|
220
|
+
|
|
221
|
+
Highly confident predictions may not flip after one-word removal. For example, strongly positive sentences such as "I absolutely love this, it is the best thing ever." often remain positive after removing one word. This does not mean every word is irrelevant. It means the model found enough evidence across the sentence that removing one token did not change the final verdict.
|
|
222
|
+
|
|
223
|
+
Counterfactual removal can create ungrammatical text because the method deletes a word rather than rewriting the sentence. This is a deliberate MVP tradeoff. The FLIP panel is a fragility test, not a full natural-language counterfactual generator. For example, removing a key word from "I am not entirely unhappy with this result." can flip the verdict, but the modified sentence may not always be natural English.
|
|
224
|
+
|
|
225
|
+
Additional limitations:
|
|
226
|
+
|
|
227
|
+
- LIME explanation takes 15-45 seconds on CPU for short text
|
|
228
|
+
- Input limited to 1000 characters
|
|
229
|
+
- Hugging Face free tier sleeps after inactivity; first request after sleep has a 30-60 second cold start
|
|
230
|
+
- Only tested on English text
|
|
231
|
+
- Sentiment-specific; adapting to other tasks would require different models and possibly different XAI methods
|
|
232
|
+
|
|
233
|
+
|
|
234
|
+
## Tech Stack
|
|
235
|
+
|
|
236
|
+
| Layer | Technology |
|
|
237
|
+
|-------|-----------|
|
|
238
|
+
| Backend | Python, FastAPI, PyTorch, Hugging Face Transformers, LIME |
|
|
239
|
+
| Frontend | Next.js, React, Tailwind CSS |
|
|
240
|
+
| Backend deployment | Hugging Face Spaces (Docker) |
|
|
241
|
+
| Frontend deployment | Vercel |
|
|
242
|
+
| ML models | DistilBERT-SST2, Twitter-RoBERTa |
|
|
243
|
+
| Infrastructure cost | Zero (free tier only) |
|
|
244
|
+
|
|
245
|
+
|
|
246
|
+
## License
|
|
247
|
+
|
|
248
|
+
MIT
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
lime_audit/__init__.py,sha256=9dECzSKe22L8jLzar7wgc03KG7cIKsnTvmmUDL3LNbo,97
|
|
2
|
+
lime_audit/analyse.py,sha256=RxBNDNkq6uV8mUj1wJzNtsghSViYIpE9LZtNGXjNeDw,13700
|
|
3
|
+
lime_audit/charts.py,sha256=Oap1ERywe_w47cMYqtBoLQzG70fJXg6TnN7aPmv_azk,7207
|
|
4
|
+
lime_audit/cli.py,sha256=ATwCQTTB_XdwAQ_7JiaOgz2YZyS2w9E14axIaMIqZyQ,2835
|
|
5
|
+
lime_audit/config.py,sha256=JudPbPku6R6eUuMGPcn4yGSzCwhQFHVEbwBsnZinHeo,333
|
|
6
|
+
lime_audit/metrics.py,sha256=YuVBzuMuY1CO1HEXW6YWIlzQchUQ61KM24OOwLpS7Ec,3186
|
|
7
|
+
lime_audit/runner.py,sha256=Q1KII97vRCNc0-N_oWN41jwHUdZGFn4Ed3U7t3-6M_s,11558
|
|
8
|
+
lime_audit/default_test_set.json,sha256=MihT_tWrudFytT1AuU94s4_ljZ0lpu7UzhgakKXJYGQ,3014
|
|
9
|
+
lime_audit-0.1.0.dist-info/METADATA,sha256=e5ZrnkyDt6vOoiQ6_8A3UP36nyKnmdmlyNMDP-ync3M,12035
|
|
10
|
+
lime_audit-0.1.0.dist-info/WHEEL,sha256=zOwg4jB6zX2kU910N-cMawjivD6tO8NEWvE12je1bVk,87
|
|
11
|
+
lime_audit-0.1.0.dist-info/entry_points.txt,sha256=_LAbstv2k2DyUBcMmP4CM_wi39PybvK4V-r7gHJLCuU,51
|
|
12
|
+
lime_audit-0.1.0.dist-info/licenses/LICENSE,sha256=yKv5VwH3OO6SDwLig_aibiZylHT-PATHcRtxED9_Mz8,1069
|
|
13
|
+
lime_audit-0.1.0.dist-info/RECORD,,
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 Parshvi Jain
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|