Interpretability and Explainability
# Interpretability and Explainability
## Introduction & Motivation
Interpretability: understanding model decisions. Explainability: communicating model behavior. Critical for trust, debugging, and deployment in high-stakes domains.
Motivation: Build trustworthy AI systems.
Applications: Healthcare, finance, autonomous systems, regulatory compliance.
---
## Core Concepts & Theory
### Feature Importance
Which inputs matter most.
### Attention Visualization
Understanding attention weights.
### Saliency Maps
Spatial importance visualization.
### Model Distillation
Simpler approximations.
---
## Mathematical Formulation
SHAP Values:
$$\phi_i = \sum_{S \subseteq N \setminus \{i\}} \frac{|S|! (|N|-|S|-1)!}{|N|!} [f(S \cup \{i\}) - f(S)]$$
Saliency:
$$ ext{Saliency} = \left| \frac{\partial f(x)}{\partial x}
ight|$$
---
## Advanced Theory & Extensions
### Concept Activation
Semantic concept importance.
### Counterfactual Explanations
What-if scenarios.
### Influence Functions
Training data attribution.
---
## Computational Considerations
SHAP: O(2^D) approximated to O(D²).
Saliency: O(D) backpropagation.
Total: O(D²) to O(D·M) for M explanations.
---
## Practical Implementation Strategies
### Gradient-Based
Use backpropagation for explanations.
### Approximation Methods
LIME, SHAP approximations.
### Visualization
Heatmaps, attention matrices.
---
## Benchmark Datasets & Evaluation
MNIST: Simple interpretability.
ImageNet: Complex visual explanations.
Tabular: Feature importance validation.
---
## Key Challenges & Limitations
### Computational Cost
Expensive for large models.
### Faithfulness
Explanations may not be accurate.
### Completeness
Multiple valid explanations.
---
## Hyperparameter Tuning
SHAP samples: 100-1000.
Saliency smoothing: σ = 0.1-1.0.
Visualization threshold: Percentile 90-95.
---
## Real-World Applications & Case Studies
Medical Diagnosis: Treatment justification.
Credit Decisions: Loan approval explanation.
Autonomous Vehicles: Decision transparency.
---
## Integration with Other Methods
Interpretability + adversarial robustness; + uncertainty estimation; + human feedback.
---
## Summary & Key Takeaways
Interpretability builds trust in AI systems.
Principles:
1. Transparency: Model clarity.
2. Explanation: Communicate decisions.
3. Attribution: Feature importance.
4. Faithfulness: Accurate representations.
5. Actionability: Enable intervention.
---
## Appendix: Practical Labs
### Lab 1: Feature Importance via Gradients
import numpy as np
def compute_gradient_importance(input_x, model, labels):
"""Compute gradient-based feature importance"""
# Simplified: use finite differences
epsilon = 1e-4
gradients = np.zeros_like(input_x)
for i in range(len(input_x)):
x_plus = input_x.copy()
x_plus[i] += epsilon
x_minus = input_x.copy()
x_minus[i] -= epsilon
output_plus = model(x_plus)
output_minus = model(x_minus)
gradients[i] = (output_plus - output_minus) / (2 * epsilon)
return gradients
model = lambda x: (x @ np.random.randn(10)).sum()
x = np.random.randn(10)
importance = compute_gradient_importance(x, model, None)
print(f"✓ Gradient importance: sum={np.sum(np.abs(importance)):.3f}")### Lab 2: Saliency Maps
import numpy as np
def compute_saliency_map(image, model):
"""Compute saliency map via gradient"""
# Image gradients (simplified)
saliency = np.zeros_like(image)
for i in range(image.shape[0]):
for j in range(image.shape[1]):
epsilon = 1e-4
image_plus = image.copy()
image_plus[i, j] += epsilon
pred_plus = model(image_plus)
pred_base = model(image)
gradient = (pred_plus - pred_base) / epsilon
saliency[i, j] = abs(gradient)
return saliency
model = lambda x: x.sum()
image = np.random.randn(5, 5)
saliency = compute_saliency_map(image, model)
print(f"✓ Saliency map: shape={saliency.shape}, max={saliency.max():.3f}")### Lab 3: Attention Visualization
import numpy as np
def visualize_attention(attention_weights, input_tokens=None):
"""Visualize attention patterns"""
# Attention is [seq_len, seq_len] or [heads, seq_len, seq_len]
if len(attention_weights.shape) == 3:
# Multi-head: average over heads
attention_weights = np.mean(attention_weights, axis=0)
# Normalize for visualization
attention_normalized = attention_weights / (np.sum(attention_weights, axis=-1, keepdims=True) + 1e-8)
# Return high-attention pairs
top_attention = []
for i in range(attention_normalized.shape[0]):
top_j = np.argmax(attention_normalized[i])
top_attention.append((i, top_j, attention_normalized[i, top_j]))
return top_attention
attention = np.random.randn(4, 8, 8)
top_pairs = visualize_attention(attention)
print(f"✓ Attention visualization: top pairs = {top_pairs[:2]}")### Lab 4: LIME Approximation
import numpy as np
def lime_local_approximation(instance, model, num_samples=1000, feature_dim=10):
"""LIME: local interpretable model-agnostic explanations"""
# Generate perturbed samples
perturbed_samples = np.random.binomial(1, 0.5, (num_samples, feature_dim))
# Get model predictions
predictions = np.array([model(sample) for sample in perturbed_samples])
# Weight by distance from original
distances = np.linalg.norm(perturbed_samples - instance, axis=1)
weights = np.exp(-distances / distances.std())
# Fit weighted linear model
X = perturbed_samples
y = predictions
W = np.diag(weights)
# Simple weighted regression: (X^T W X)^-1 X^T W y
coefficients = np.zeros(feature_dim)
for i in range(feature_dim):
weighted_sum = np.sum(weights * X[:, i] * y)
weighted_denom = np.sum(weights * X[:, i] ** 2)
coefficients[i] = weighted_sum / (weighted_denom + 1e-8)
return coefficients
model = lambda x: (x @ np.random.randn(10)).sum()
instance = np.random.rand(10)
lime_weights = lime_local_approximation(instance, model)
print(f"✓ LIME weights: sum={np.sum(np.abs(lime_weights)):.3f}")---