Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -126,7 +126,9 @@ explain(model, X, instance=0,
explainer_kwargs={"surrogate": {"n_samples": 2000},
"coalition": {"n_background": 100}})
```
### Examples

- [Consensus Explanation Walkthrough](examples/01_consensus_walkthrough.ipynb) — demonstrates local consensus explanations, agreement between explanation methods, disagreement cases, and global explanations.
### Use the explainer classes directly

```python
Expand Down
285 changes: 285 additions & 0 deletions examples/01_consensus_walkthrough.ipynb
Original file line number Diff line number Diff line change
@@ -0,0 +1,285 @@
{
"cells": [
{
"cell_type": "markdown",
"id": "9a5ae9ca",
"metadata": {},
"source": [
"# Consensus Explanation Walkthrough\n",
"\n",
"This notebook demonstrates how XAI-Framework combines multiple\n",
"explanation methods and reports how much they agree.\n",
"\n",
"We will cover:\n",
"\n",
"1. A local explanation using a Random Forest.\n",
"2. A case where explanation methods disagree.\n",
"3. Global explanations with and without ground-truth labels."
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "664e2b62",
"metadata": {},
"outputs": [],
"source": [
"from sklearn.datasets import load_breast_cancer\n",
"from sklearn.ensemble import RandomForestClassifier\n",
"from sklearn.model_selection import train_test_split\n",
"\n",
"from xai_framework import explain, explain_global\n",
"\n",
"\n",
"# Load a small scikit-learn dataset\n",
"data = load_breast_cancer(as_frame=True)\n",
"\n",
"X = data.data\n",
"y = data.target\n",
"\n",
"# Split into training and test data\n",
"X_train, X_test, y_train, y_test = train_test_split(\n",
" X,\n",
" y,\n",
" test_size=0.2,\n",
" random_state=42,\n",
" stratify=y,\n",
")\n",
"\n",
"# Train a Random Forest classifier\n",
"model = RandomForestClassifier(\n",
" n_estimators=100,\n",
" random_state=42,\n",
")\n",
"\n",
"model.fit(X_train, y_train)\n",
"\n",
"print(\"Model trained successfully.\")\n",
"print(\"Training samples:\", len(X_train))\n",
"print(\"Test samples:\", len(X_test))"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "848db025",
"metadata": {},
"outputs": [],
"source": [
"# Explain one test instance using multiple explanation methods\n",
"instance = X_test.iloc[0]\n",
"\n",
"result = explain(\n",
" model,\n",
" X_train,\n",
" instance=instance,\n",
" random_state=42,\n",
")\n",
"\n",
"result"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "3b21c2e9",
"metadata": {},
"outputs": [],
"source": [
"# Display the local explanation plot\n",
"result.plot()"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "a09a6fb3",
"metadata": {},
"outputs": [],
"source": [
"# Show the agreement between the explanation methods\n",
"result.agreement_matrix"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "1c9233a1",
"metadata": {},
"outputs": [],
"source": [
"# Plot the consensus explanation\n",
"result.plot()"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "0d01b298",
"metadata": {},
"outputs": [],
"source": [
"from sklearn.datasets import make_classification\n",
"from sklearn.neural_network import MLPClassifier\n",
"\n",
"# Create a noisy dataset\n",
"X_noisy, y_noisy = make_classification(\n",
" n_samples=500,\n",
" n_features=10,\n",
" n_informative=3,\n",
" n_redundant=2,\n",
" flip_y=0.25,\n",
" random_state=42,\n",
")\n",
"\n",
"# Train an MLP classifier\n",
"mlp = MLPClassifier(\n",
" hidden_layer_sizes=(20,),\n",
" max_iter=1000,\n",
" random_state=42,\n",
")\n",
"\n",
"mlp.fit(X_noisy, y_noisy)\n",
"\n",
"print(\"MLP model trained successfully.\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "4c8612c0",
"metadata": {},
"outputs": [],
"source": [
"# Explain one instance using the MLP model\n",
"noisy_instance = X_noisy[1]\n",
"\n",
"mlp_result = explain(\n",
" mlp,\n",
" X_noisy,\n",
" instance=noisy_instance,\n",
" random_state=42,\n",
")\n",
"\n",
"mlp_result"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "f22dc92d",
"metadata": {},
"outputs": [],
"source": [
"# Show the explanation plot\n",
"mlp_result.plot()\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "3e99fb65",
"metadata": {},
"outputs": [],
"source": [
"# Show the agreement matrix\n",
"mlp_result.agreement_matrix"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "faad45f3",
"metadata": {},
"outputs": [],
"source": [
"# Increase the number of samples used by the surrogate explainer\n",
"mlp_result_more_samples = explain(\n",
" mlp,\n",
" X_noisy,\n",
" instance=noisy_instance,\n",
" random_state=42,\n",
" explainer_kwargs={\n",
" \"surrogate\": {\n",
" \"n_samples\": 100000\n",
" }\n",
" },\n",
")\n",
"\n",
"mlp_result_more_samples"
]
},
{
"cell_type": "markdown",
"id": "49b1c417",
"metadata": {},
"source": [
"### Effect of increasing surrogate samples\n",
"\n",
"For this noisy instance, the default surrogate explanation shows substantial disagreement\n",
"between the coalition and surrogate explanations, with an agreement of about 0.24.\n",
"\n",
"After increasing the surrogate sample size to 100,000, the agreement changes to about 0.18.\n",
"The agreement does not necessarily improve when the surrogate sample size is increased,\n",
"but the change shows that surrogate sampling settings can affect the explanation.\n",
"\n",
"For this instance, f8 remains the strongest positive feature, while f7 has a negative\n",
"contribution. Both explanations target class 0."
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "dbedd824",
"metadata": {},
"outputs": [],
"source": [
"# Global explanation without ground-truth labels\n",
"global_without_y = explain_global(\n",
" mlp,\n",
" X_noisy,\n",
")\n",
"\n",
"global_without_y"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "4d375c06",
"metadata": {},
"outputs": [],
"source": [
"# Global explanation with ground-truth labels\n",
"global_with_y = explain_global(\n",
" mlp,\n",
" X_noisy,\n",
" y=y_noisy,\n",
")\n",
"\n",
"global_with_y"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.14.7"
}
},
"nbformat": 4,
"nbformat_minor": 5
}
Loading