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}")

---

Go deeper with CFSGPT

Get AI-powered deep-dives, save terms, and run advanced simulations — free account.

Create Free Account