Skip to main content

Model Interpretation & Explainability: Understanding the Black Box

๐Ÿ“š What You'll Learn

By the end of this lesson, you will be able to:

โฑ๏ธ Estimated Time: 45โ€“60 minutes

๐ŸŽฏ Project: Take a trained black-box model (random forest or gradient boosting) and produce a full explainability report using permutation importance, SHAP, and partial dependence plots.

Opening the Black Box: Why Explainability Matters ๐Ÿ”

"With great power comes great responsibility." As machine learning models become more complex and powerful, understanding their decisions becomes crucial. From regulatory compliance to debugging models, from building trust to discovering insights, interpretability is no longer optionalโ€”it's essential for responsible AI deployment.

The Landscape of Model Interpretability

graph TD A[Model Interpretability] --> B[Global Interpretability] A --> C[Local Interpretability] B --> D[Feature Importance] B --> E[Partial Dependence] B --> F[Model-Specific Methods] C --> G[LIME] C --> H[SHAP] C --> I[Counterfactuals] D --> J[Permutation Importance] D --> K[Drop-Column Importance] E --> L[PDP Plots] E --> M[ICE Plots] F --> N[Tree Visualization] F --> O[Linear Coefficients] style A fill:#667eea,color:#fff style B fill:#51cf66 style C fill:#4ecdc4
# Essential imports for model interpretation
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.model_selection import train_test_split
from sklearn.ensemble import RandomForestClassifier, GradientBoostingRegressor
from sklearn.linear_model import LogisticRegression, LinearRegression
from sklearn.preprocessing import StandardScaler
from sklearn.datasets import load_breast_cancer, load_boston, make_classification
import warnings
warnings.filterwarnings('ignore')

# Set visualization style
plt.style.use('seaborn-v0_8-whitegrid')
sns.set_palette("husl")

print("="*60)
print("MODEL INTERPRETABILITY: UNDERSTANDING PREDICTIONS")
print("="*60)

# Load dataset for demonstrations
data = load_breast_cancer()
X = pd.DataFrame(data.data, columns=data.feature_names)
y = data.target

# Create meaningful feature names (shortened for visualization)
feature_names_short = [name.split(' ')[0] for name in data.feature_names]
X.columns = feature_names_short

X_train, X_test, y_train, y_test = train_test_split(
    X, y, test_size=0.2, random_state=42, stratify=y
)

print(f"Dataset: Breast Cancer Classification")
print(f"Shape: {X.shape}")
print(f"Features: {X.columns.tolist()[:5]}...")
print(f"Classes: {data.target_names}")

# Train multiple models for comparison
models = {
    'Logistic Regression': LogisticRegression(max_iter=1000, random_state=42),
    'Random Forest': RandomForestClassifier(n_estimators=100, random_state=42),
    'Gradient Boosting': GradientBoostingRegressor(n_estimators=100, random_state=42)
}

fitted_models = {}
for name, model in models.items():
    if name == 'Gradient Boosting':
        model.fit(X_train, y_train)
    else:
        model.fit(X_train, y_train)
    fitted_models[name] = model
    print(f"โœ“ Trained {name}")

Global Interpretability: Understanding Overall Model Behavior

1. Permutation Importance

from sklearn.inspection import permutation_importance
from sklearn.metrics import accuracy_score

print("\n" + "="*60)
print("1. PERMUTATION IMPORTANCE")
print("="*60)

def calculate_permutation_importance(model, X, y, n_repeats=10):
    """
    Calculate feature importance by permuting each feature
    and measuring the decrease in model performance
    """
    
    # Baseline score
    baseline_score = model.score(X, y) if hasattr(model, 'score') else \
                    accuracy_score(y, model.predict(X))
    
    importances = []
    
    for col_idx, col in enumerate(X.columns):
        scores = []
        
        for _ in range(n_repeats):
            # Create a copy and permute the column
            X_permuted = X.copy()
            X_permuted[col] = np.random.permutation(X_permuted[col].values)
            
            # Calculate score with permuted feature
            if hasattr(model, 'score'):
                permuted_score = model.score(X_permuted, y)
            else:
                permuted_score = accuracy_score(y, model.predict(X_permuted))
            
            # Calculate importance as decrease in performance
            importance = baseline_score - permuted_score
            scores.append(importance)
        
        importances.append({
            'feature': col,
            'importance': np.mean(scores),
            'std': np.std(scores)
        })
    
    return pd.DataFrame(importances).sort_values('importance', ascending=False)

# Calculate permutation importance for Random Forest
rf_model = fitted_models['Random Forest']

print("Calculating permutation importance...")
perm_importance = calculate_permutation_importance(rf_model, X_test, y_test)

print("\nTop 10 Most Important Features (Permutation):")
print(perm_importance.head(10).to_string(index=False))

# Visualize permutation importance
fig, axes = plt.subplots(1, 2, figsize=(14, 6))

# Plot 1: Bar plot of importance
top_features = perm_importance.head(10)
axes[0].barh(range(len(top_features)), top_features['importance'].values)
axes[0].set_yticks(range(len(top_features)))
axes[0].set_yticklabels(top_features['feature'].values)
axes[0].set_xlabel('Importance (Performance Decrease)')
axes[0].set_title('Permutation Feature Importance')
axes[0].invert_yaxis()

# Add error bars
axes[0].errorbar(top_features['importance'].values, 
                range(len(top_features)),
                xerr=top_features['std'].values,
                fmt='none', color='black', alpha=0.5)

# Plot 2: Compare with built-in feature importance
if hasattr(rf_model, 'feature_importances_'):
    builtin_importance = pd.DataFrame({
        'feature': X.columns,
        'importance': rf_model.feature_importances_
    }).sort_values('importance', ascending=False).head(10)
    
    axes[1].barh(range(len(builtin_importance)), builtin_importance['importance'].values)
    axes[1].set_yticks(range(len(builtin_importance)))
    axes[1].set_yticklabels(builtin_importance['feature'].values)
    axes[1].set_xlabel('Importance (Gini/Entropy)')
    axes[1].set_title('Built-in Feature Importance (Tree-based)')
    axes[1].invert_yaxis()

plt.tight_layout()
plt.show()

print("\n๐Ÿ’ก Insight: Permutation importance is model-agnostic and captures")
print("   feature interactions, while built-in importance may miss them.")

2. Partial Dependence Plots (PDP)

from sklearn.inspection import PartialDependenceDisplay
import matplotlib.pyplot as plt

print("\n" + "="*60)
print("2. PARTIAL DEPENDENCE PLOTS")
print("="*60)

def create_partial_dependence_plots(model, X, features, feature_names=None):
    """
    Create partial dependence plots to show the marginal effect
    of features on the predicted outcome
    """
    
    if feature_names is None:
        feature_names = X.columns
    
    # Create figure
    fig, axes = plt.subplots(2, 2, figsize=(12, 10))
    axes = axes.flatten()
    
    for idx, feature in enumerate(features[:4]):
        ax = axes[idx]
        
        # Calculate partial dependence manually
        feature_idx = list(X.columns).index(feature)
        x_values = np.linspace(X[feature].min(), X[feature].max(), 100)
        
        # Create dataset with feature values varying
        pd_values = []
        for val in x_values:
            X_temp = X.copy()
            X_temp[feature] = val
            
            # Get average prediction
            if hasattr(model, 'predict_proba'):
                preds = model.predict_proba(X_temp)[:, 1]
            else:
                preds = model.predict(X_temp)
            
            pd_values.append(preds.mean())
        
        # Plot
        ax.plot(x_values, pd_values, 'b-', linewidth=2)
        ax.set_xlabel(feature)
        ax.set_ylabel('Partial Dependence')
        ax.set_title(f'Partial Dependence of {feature}')
        ax.grid(True, alpha=0.3)
        
        # Add rug plot for data distribution
        ax.plot(X[feature].values, [ax.get_ylim()[0]] * len(X), '|', 
               alpha=0.1, color='gray')
    
    plt.suptitle('Partial Dependence Plots - Random Forest', fontsize=14)
    plt.tight_layout()
    plt.show()

# Select important features for PDP
top_features_for_pdp = perm_importance.head(4)['feature'].tolist()

print(f"Creating PDPs for features: {top_features_for_pdp}")
create_partial_dependence_plots(rf_model, X_test, top_features_for_pdp)

# Using sklearn's built-in PDP (if available)
try:
    from sklearn.inspection import plot_partial_dependence
    
    fig, ax = plt.subplots(figsize=(12, 8))
    display = plot_partial_dependence(
        rf_model, X_test, features=top_features_for_pdp[:4],
        kind="both", subsample=50, n_jobs=-1,
        grid_resolution=50, ax=ax
    )
    plt.suptitle('Partial Dependence and Individual Conditional Expectation')
    plt.tight_layout()
    plt.show()
    
    print("\n๐Ÿ’ก ICE plots (light lines) show individual predictions,")
    print("   PDP (dark line) shows the average effect.")
except:
    print("\nNote: Built-in PDP not available in this sklearn version")

Local Interpretability: Understanding Individual Predictions

3. LIME (Local Interpretable Model-agnostic Explanations)

# Note: Requires installation: pip install lime
print("\n" + "="*60)
print("3. LIME - LOCAL EXPLANATIONS")
print("="*60)

# Simplified LIME implementation for demonstration
class SimpleLIME:
    """
    Simplified LIME implementation for educational purposes
    """
    
    def __init__(self, model, num_features=10, num_samples=5000):
        self.model = model
        self.num_features = num_features
        self.num_samples = num_samples
    
    def explain_instance(self, instance, X_train, feature_names=None):
        """
        Explain a single prediction using local linear approximation
        """
        if feature_names is None:
            feature_names = [f'Feature_{i}' for i in range(len(instance))]
        
        # Generate perturbed samples around the instance
        np.random.seed(42)
        perturbations = []
        predictions = []
        weights = []
        
        for _ in range(self.num_samples):
            # Create perturbation
            perturbation = instance.copy()
            
            # Randomly change some features
            num_changes = np.random.randint(1, len(instance))
            features_to_change = np.random.choice(len(instance), num_changes, replace=False)
            
            for feat_idx in features_to_change:
                # Sample from training data distribution
                perturbation[feat_idx] = np.random.choice(X_train.iloc[:, feat_idx])
            
            perturbations.append(perturbation)
            
            # Get prediction for perturbation
            if hasattr(self.model, 'predict_proba'):
                pred = self.model.predict_proba(perturbation.values.reshape(1, -1))[0, 1]
            else:
                pred = self.model.predict(perturbation.values.reshape(1, -1))[0]
            predictions.append(pred)
            
            # Calculate weight based on distance to original instance
            distance = np.linalg.norm(perturbation - instance)
            weight = np.exp(-distance**2 / 2)  # Gaussian kernel
            weights.append(weight)
        
        # Convert to arrays
        X_perturbed = np.array(perturbations)
        y_perturbed = np.array(predictions)
        weights = np.array(weights)
        
        # Fit weighted linear model
        from sklearn.linear_model import Ridge
        
        # Create binary features (different from original or not)
        X_binary = (X_perturbed != instance.values).astype(int)
        
        # Fit weighted Ridge regression
        linear_model = Ridge(alpha=1.0)
        linear_model.fit(X_binary, y_perturbed, sample_weight=weights)
        
        # Get feature importance from linear model
        importances = pd.DataFrame({
            'feature': feature_names,
            'importance': linear_model.coef_
        }).sort_values('importance', key=abs, ascending=False)
        
        # Get prediction for original instance
        if hasattr(self.model, 'predict_proba'):
            original_pred = self.model.predict_proba(instance.values.reshape(1, -1))[0, 1]
        else:
            original_pred = self.model.predict(instance.values.reshape(1, -1))[0]
        
        return {
            'prediction': original_pred,
            'intercept': linear_model.intercept_,
            'importances': importances.head(self.num_features)
        }

# Create LIME explainer
lime_explainer = SimpleLIME(rf_model, num_features=10)

# Explain a specific instance
instance_idx = 0
instance = X_test.iloc[instance_idx]
true_label = y_test[instance_idx]

print(f"Explaining prediction for instance {instance_idx}")
print(f"True label: {true_label} ({'benign' if true_label == 1 else 'malignant'})")

explanation = lime_explainer.explain_instance(instance, X_train, X.columns)

print(f"Predicted probability: {explanation['prediction']:.3f}")
print(f"\nTop features influencing this prediction:")
print(explanation['importances'].to_string(index=False))

# Visualize LIME explanation
fig, axes = plt.subplots(1, 2, figsize=(14, 6))

# Plot 1: Feature contributions
top_lime = explanation['importances'].head(10)
colors = ['green' if x > 0 else 'red' for x in top_lime['importance']]

axes[0].barh(range(len(top_lime)), top_lime['importance'].values, color=colors)
axes[0].set_yticks(range(len(top_lime)))
axes[0].set_yticklabels(top_lime['feature'].values)
axes[0].set_xlabel('Contribution to Prediction')
axes[0].set_title(f'LIME Explanation for Instance {instance_idx}')
axes[0].axvline(x=0, color='black', linestyle='-', linewidth=0.5)
axes[0].invert_yaxis()

# Plot 2: Feature values
feature_values = pd.DataFrame({
    'feature': top_lime['feature'].values,
    'value': [instance[feat] for feat in top_lime['feature']],
    'percentile': [(instance[feat] > X_train[feat]).mean() * 100 
                  for feat in top_lime['feature']]
})

axes[1].barh(range(len(feature_values)), feature_values['percentile'].values)
axes[1].set_yticks(range(len(feature_values)))
axes[1].set_yticklabels([f"{feat}\n({val:.2f})" 
                         for feat, val in zip(feature_values['feature'], 
                                             feature_values['value'])])
axes[1].set_xlabel('Percentile in Training Data')
axes[1].set_title('Feature Values (Percentiles)')
axes[1].invert_yaxis()

plt.tight_layout()
plt.show()

print("\n๐Ÿ’ก Green bars increase prediction probability,")
print("   red bars decrease it.")

4. SHAP (SHapley Additive exPlanations)

# Note: Requires installation: pip install shap
print("\n" + "="*60)
print("4. SHAP VALUES - GAME THEORY FOR ML")
print("="*60)

# Simplified SHAP implementation
class SimplifiedSHAP:
    """
    Simplified SHAP implementation for demonstration
    Shows the core concept of Shapley values
    """
    
    def __init__(self, model, background_data, n_samples=100):
        self.model = model
        self.background_data = background_data
        self.n_samples = n_samples
    
    def calculate_shap_values(self, instance):
        """
        Calculate approximate SHAP values for a single instance
        """
        n_features = len(instance)
        shap_values = np.zeros(n_features)
        
        # For each feature
        for feature_idx in range(n_features):
            marginal_contributions = []
            
            # Sample random coalitions
            for _ in range(self.n_samples):
                # Random coalition (subset of features)
                coalition = np.random.choice([True, False], n_features)
                coalition[feature_idx] = False  # Feature not in coalition
                
                # Prediction without feature
                X_without = self._create_instance(instance, coalition, feature_idx, False)
                if hasattr(self.model, 'predict_proba'):
                    pred_without = self.model.predict_proba(X_without)[0, 1]
                else:
                    pred_without = self.model.predict(X_without)[0]
                
                # Prediction with feature
                coalition[feature_idx] = True
                X_with = self._create_instance(instance, coalition, feature_idx, True)
                if hasattr(self.model, 'predict_proba'):
                    pred_with = self.model.predict_proba(X_with)[0, 1]
                else:
                    pred_with = self.model.predict(X_with)[0]
                
                # Marginal contribution
                marginal_contribution = pred_with - pred_without
                marginal_contributions.append(marginal_contribution)
            
            # Average marginal contribution is the SHAP value
            shap_values[feature_idx] = np.mean(marginal_contributions)
        
        return shap_values
    
    def _create_instance(self, instance, coalition, feature_idx, include_feature):
        """
        Create an instance with features from coalition
        """
        X = instance.copy()
        
        # Replace non-coalition features with background values
        for idx in range(len(instance)):
            if not coalition[idx]:
                # Use mean from background data
                X[idx] = self.background_data.iloc[:, idx].mean()
        
        if include_feature:
            X[feature_idx] = instance[feature_idx]
        else:
            X[feature_idx] = self.background_data.iloc[:, feature_idx].mean()
        
        return X.values.reshape(1, -1)

# Calculate SHAP values
shap_explainer = SimplifiedSHAP(rf_model, X_train, n_samples=50)

# Explain multiple instances
n_explain = 5
shap_values_list = []
predictions_list = []

print(f"Calculating SHAP values for {n_explain} instances...")

for i in range(n_explain):
    instance = X_test.iloc[i]
    shap_vals = shap_explainer.calculate_shap_values(instance)
    shap_values_list.append(shap_vals)
    
    if hasattr(rf_model, 'predict_proba'):
        pred = rf_model.predict_proba(instance.values.reshape(1, -1))[0, 1]
    else:
        pred = rf_model.predict(instance.values.reshape(1, -1))[0]
    predictions_list.append(pred)

shap_values_array = np.array(shap_values_list)

# Visualize SHAP values
fig, axes = plt.subplots(2, 2, figsize=(14, 10))

# Plot 1: Waterfall plot for single instance
instance_idx = 0
shap_vals = shap_values_array[instance_idx]
feature_order = np.argsort(np.abs(shap_vals))[::-1][:10]

axes[0, 0].barh(range(len(feature_order)), shap_vals[feature_order])
axes[0, 0].set_yticks(range(len(feature_order)))
axes[0, 0].set_yticklabels([X.columns[i] for i in feature_order])
axes[0, 0].set_xlabel('SHAP Value')
axes[0, 0].set_title(f'SHAP Values for Instance {instance_idx}')
axes[0, 0].axvline(x=0, color='black', linestyle='-', linewidth=0.5)
axes[0, 0].invert_yaxis()

# Plot 2: Global feature importance (mean absolute SHAP)
mean_abs_shap = np.mean(np.abs(shap_values_array), axis=0)
feature_importance_shap = pd.DataFrame({
    'feature': X.columns,
    'importance': mean_abs_shap
}).sort_values('importance', ascending=False).head(10)

axes[0, 1].barh(range(len(feature_importance_shap)), 
               feature_importance_shap['importance'].values)
axes[0, 1].set_yticks(range(len(feature_importance_shap)))
axes[0, 1].set_yticklabels(feature_importance_shap['feature'].values)
axes[0, 1].set_xlabel('Mean |SHAP Value|')
axes[0, 1].set_title('Global Feature Importance (SHAP)')
axes[0, 1].invert_yaxis()

# Plot 3: SHAP summary plot (simplified)
top_features_idx = feature_importance_shap['feature'].apply(
    lambda x: list(X.columns).index(x)).values[:5]

for idx, feat_idx in enumerate(top_features_idx):
    y_pos = np.ones(n_explain) * idx + np.random.normal(0, 0.02, n_explain)
    colors = X_test.iloc[:n_explain, feat_idx].values
    
    scatter = axes[1, 0].scatter(shap_values_array[:, feat_idx], y_pos, 
                                 c=colors, cmap='coolwarm', alpha=0.7, s=50)

axes[1, 0].set_yticks(range(len(top_features_idx)))
axes[1, 0].set_yticklabels([X.columns[i] for i in top_features_idx])
axes[1, 0].set_xlabel('SHAP Value')
axes[1, 0].set_title('SHAP Summary Plot')
axes[1, 0].axvline(x=0, color='black', linestyle='-', linewidth=0.5)
plt.colorbar(scatter, ax=axes[1, 0], label='Feature Value')

# Plot 4: Force plot visualization (simplified)
instance_idx = 0
shap_vals = shap_values_array[instance_idx]
base_value = X_train.mean().values  # Simplified base value

# Sort features by SHAP value
sorted_idx = np.argsort(shap_vals)[::-1]
cumsum = np.cumsum(shap_vals[sorted_idx])

axes[1, 1].bar(range(len(sorted_idx[:10])), shap_vals[sorted_idx[:10]])
axes[1, 1].set_xticks(range(len(sorted_idx[:10])))
axes[1, 1].set_xticklabels([X.columns[i] for i in sorted_idx[:10]], rotation=45, ha='right')
axes[1, 1].set_ylabel('SHAP Value')
axes[1, 1].set_title(f'Force Plot - Instance {instance_idx}')
axes[1, 1].axhline(y=0, color='black', linestyle='-', linewidth=0.5)

plt.tight_layout()
plt.show()

print("\n๐Ÿ’ก SHAP values provide consistent, theoretically grounded")
print("   feature attributions based on game theory.")

Model-Specific Interpretation Methods

print("\n" + "="*60)
print("5. MODEL-SPECIFIC INTERPRETATION")
print("="*60)

# 1. Linear Model Coefficients
print("\n1. Linear Model Interpretation")
print("-" * 40)

# Scale data for linear model
scaler = StandardScaler()
X_train_scaled = scaler.fit_transform(X_train)
X_test_scaled = scaler.transform(X_test)

# Train logistic regression
lr_model = LogisticRegression(max_iter=1000, random_state=42)
lr_model.fit(X_train_scaled, y_train)

# Get coefficients
coefficients = pd.DataFrame({
    'feature': X.columns,
    'coefficient': lr_model.coef_[0],
    'abs_coefficient': np.abs(lr_model.coef_[0])
}).sort_values('abs_coefficient', ascending=False)

print("Top 10 Most Important Features (Linear Model):")
print(coefficients.head(10)[['feature', 'coefficient']].to_string(index=False))

# Visualize coefficients
fig, axes = plt.subplots(1, 2, figsize=(14, 6))

# Plot 1: Coefficients
top_coef = coefficients.head(15)
colors = ['green' if x > 0 else 'red' for x in top_coef['coefficient']]

axes[0].barh(range(len(top_coef)), top_coef['coefficient'].values, color=colors)
axes[0].set_yticks(range(len(top_coef)))
axes[0].set_yticklabels(top_coef['feature'].values)
axes[0].set_xlabel('Coefficient Value')
axes[0].set_title('Logistic Regression Coefficients')
axes[0].axvline(x=0, color='black', linestyle='-', linewidth=0.5)
axes[0].invert_yaxis()

# Plot 2: Odds ratios
top_coef['odds_ratio'] = np.exp(top_coef['coefficient'])
axes[1].barh(range(len(top_coef)), top_coef['odds_ratio'].values)
axes[1].set_yticks(range(len(top_coef)))
axes[1].set_yticklabels(top_coef['feature'].values)
axes[1].set_xlabel('Odds Ratio')
axes[1].set_title('Odds Ratios (exp(coefficient))')
axes[1].axvline(x=1, color='black', linestyle='-', linewidth=0.5)
axes[1].invert_yaxis()

plt.tight_layout()
plt.show()

print("\n๐Ÿ’ก Positive coefficients increase the probability of positive class,")
print("   Odds ratio > 1 means feature increases odds of positive class.")

# 2. Decision Tree Visualization
print("\n2. Decision Tree Interpretation")
print("-" * 40)

from sklearn.tree import DecisionTreeClassifier, plot_tree

# Train a simple decision tree
dt_model = DecisionTreeClassifier(max_depth=3, random_state=42)
dt_model.fit(X_train, y_train)

# Visualize tree
fig, ax = plt.subplots(figsize=(20, 10))
plot_tree(dt_model, 
          feature_names=X.columns,
          class_names=['Malignant', 'Benign'],
          filled=True,
          rounded=True,
          fontsize=10,
          ax=ax)
plt.title('Decision Tree Visualization (Max Depth = 3)')
plt.show()

# Extract decision rules
def extract_rules(tree, feature_names):
    """Extract decision rules from tree"""
    tree_ = tree.tree_
    feature_name = [
        feature_names[i] if i != -2 else "undefined!"
        for i in tree_.feature
    ]
    
    rules = []
    
    def recurse(node, depth, parent_rule=""):
        if tree_.feature[node] != -2:
            name = feature_name[node]
            threshold = tree_.threshold[node]
            
            # Left child (<=)
            left_rule = f"{parent_rule} & " if parent_rule else ""
            left_rule += f"({name} <= {threshold:.2f})"
            recurse(tree_.children_left[node], depth + 1, left_rule)
            
            # Right child (>)
            right_rule = f"{parent_rule} & " if parent_rule else ""
            right_rule += f"({name} > {threshold:.2f})"
            recurse(tree_.children_right[node], depth + 1, right_rule)
        else:
            # Leaf node
            value = tree_.value[node][0]
            class_prediction = np.argmax(value)
            confidence = value[class_prediction] / value.sum()
            
            rules.append({
                'rule': parent_rule,
                'class': class_prediction,
                'samples': int(tree_.n_node_samples[node]),
                'confidence': confidence
            })
    
    recurse(0, 1)
    return rules

# Extract and display rules
rules = extract_rules(dt_model, X.columns)
rules_df = pd.DataFrame(rules).sort_values('samples', ascending=False).head(5)

print("\nTop 5 Decision Rules by Sample Coverage:")
for idx, row in rules_df.iterrows():
    class_name = 'Benign' if row['class'] == 1 else 'Malignant'
    print(f"\nRule: {row['rule']}")
    print(f"  โ†’ Predicts: {class_name}")
    print(f"  โ†’ Samples: {row['samples']}")
    print(f"  โ†’ Confidence: {row['confidence']:.2%}")

Advanced Interpretation Techniques

print("\n" + "="*60)
print("6. ADVANCED INTERPRETATION TECHNIQUES")
print("="*60)

# 1. Anchors: High-Precision Rules
print("\n1. Anchor Explanations")
print("-" * 40)

class SimpleAnchorExplainer:
    """
    Find minimal sufficient conditions (anchors) for predictions
    """
    
    def __init__(self, model, precision_threshold=0.95):
        self.model = model
        self.precision_threshold = precision_threshold
    
    def find_anchor(self, instance, X_data, max_features=5):
        """
        Find minimal set of features that anchor the prediction
        """
        # Get original prediction
        original_pred = self.model.predict(instance.values.reshape(1, -1))[0]
        
        # Start with empty anchor
        anchor_features = []
        anchor_conditions = {}
        
        # Greedily add features
        remaining_features = list(X_data.columns)
        
        for _ in range(max_features):
            best_feature = None
            best_precision = 0
            best_condition = None
            
            for feature in remaining_features:
                # Create condition based on instance value
                feature_value = instance[feature]
                
                # For numerical features, use range
                if X_data[feature].dtype in ['float64', 'int64']:
                    std = X_data[feature].std()
                    condition = (X_data[feature] >= feature_value - 0.1*std) & \
                               (X_data[feature] <= feature_value + 0.1*std)
                else:
                    condition = X_data[feature] == feature_value
                
                # Combine with existing anchor conditions
                combined_condition = condition
                for anchor_feat, anchor_cond in anchor_conditions.items():
                    combined_condition = combined_condition & anchor_cond
                
                # Test precision
                matching_samples = X_data[combined_condition]
                if len(matching_samples) > 0:
                    predictions = self.model.predict(matching_samples)
                    precision = (predictions == original_pred).mean()
                    
                    if precision > best_precision:
                        best_precision = precision
                        best_feature = feature
                        best_condition = condition
            
            # Add best feature to anchor
            if best_precision >= self.precision_threshold:
                anchor_features.append(best_feature)
                anchor_conditions[best_feature] = best_condition
                break
            elif best_feature:
                anchor_features.append(best_feature)
                anchor_conditions[best_feature] = best_condition
                remaining_features.remove(best_feature)
        
        return {
            'features': anchor_features,
            'precision': best_precision,
            'coverage': len(X_data[combined_condition]) / len(X_data),
            'prediction': original_pred
        }

# Find anchors for an instance
anchor_explainer = SimpleAnchorExplainer(rf_model, precision_threshold=0.9)
instance_idx = 0
instance = X_test.iloc[instance_idx]

anchor_result = anchor_explainer.find_anchor(instance, X_train, max_features=3)

print(f"Anchor Explanation for Instance {instance_idx}:")
print(f"Prediction: {anchor_result['prediction']}")
print(f"Anchor features: {anchor_result['features']}")
print(f"Precision: {anchor_result['precision']:.2%}")
print(f"Coverage: {anchor_result['coverage']:.2%}")

# 2. Counterfactual Explanations
print("\n2. Counterfactual Explanations")
print("-" * 40)

class CounterfactualExplainer:
    """
    Find minimal changes to flip the prediction
    """
    
    def __init__(self, model):
        self.model = model
    
    def find_counterfactual(self, instance, X_train, desired_class=None, max_changes=3):
        """
        Find minimal changes to achieve desired prediction
        """
        # Get original prediction
        original_pred = self.model.predict(instance.values.reshape(1, -1))[0]
        
        if desired_class is None:
            desired_class = 1 - original_pred  # Flip the prediction
        
        # Find nearest neighbors with desired class
        from sklearn.neighbors import NearestNeighbors
        
        # Get samples with desired class
        desired_samples = X_train[self.model.predict(X_train) == desired_class]
        
        if len(desired_samples) == 0:
            return None
        
        # Find nearest neighbor
        nn = NearestNeighbors(n_neighbors=1)
        nn.fit(desired_samples)
        
        distances, indices = nn.kneighbors(instance.values.reshape(1, -1))
        nearest = desired_samples.iloc[indices[0, 0]]
        
        # Find minimal changes
        differences = pd.DataFrame({
            'feature': X_train.columns,
            'original': instance.values,
            'counterfactual': nearest.values,
            'change': nearest.values - instance.values,
            'abs_change': np.abs(nearest.values - instance.values)
        }).sort_values('abs_change', ascending=False)
        
        # Select top changes
        top_changes = differences[differences['abs_change'] > 0].head(max_changes)
        
        return {
            'original_prediction': original_pred,
            'desired_prediction': desired_class,
            'changes': top_changes,
            'distance': distances[0, 0]
        }

# Find counterfactual
cf_explainer = CounterfactualExplainer(rf_model)
cf_result = cf_explainer.find_counterfactual(instance, X_train, max_changes=3)

if cf_result:
    print(f"Counterfactual Explanation for Instance {instance_idx}:")
    print(f"Original prediction: {cf_result['original_prediction']}")
    print(f"Desired prediction: {cf_result['desired_prediction']}")
    print(f"\nMinimal changes needed:")
    print(cf_result['changes'][['feature', 'original', 'counterfactual', 'change']].to_string(index=False))

# 3. Prototype and Criticism Selection
print("\n3. Prototypes and Criticisms")
print("-" * 40)

def select_prototypes_criticisms(X, y, model, n_prototypes=5, n_criticisms=5):
    """
    Select representative examples (prototypes) and 
    edge cases (criticisms) for interpretation
    """
    from sklearn.metrics.pairwise import euclidean_distances
    
    # Get model predictions
    predictions = model.predict(X)
    
    # For each class, find prototypes
    prototypes = {}
    criticisms = {}
    
    for class_label in np.unique(y):
        # Get samples of this class
        class_mask = (y == class_label)
        class_samples = X[class_mask]
        class_indices = np.where(class_mask)[0]
        
        if len(class_samples) < n_prototypes + n_criticisms:
            continue
        
        # Calculate centroid
        centroid = class_samples.mean(axis=0).values.reshape(1, -1)
        
        # Find prototypes (closest to centroid)
        distances_to_centroid = euclidean_distances(class_samples, centroid).flatten()
        prototype_indices = class_indices[np.argsort(distances_to_centroid)[:n_prototypes]]
        
        # Find criticisms (furthest from centroid but still correctly classified)
        correct_predictions = predictions[class_mask] == class_label
        correct_indices = class_indices[correct_predictions]
        
        if len(correct_indices) > 0:
            distances_correct = distances_to_centroid[correct_predictions]
            criticism_indices = correct_indices[np.argsort(distances_correct)[-n_criticisms:]]
        else:
            criticism_indices = []
        
        prototypes[class_label] = prototype_indices
        criticisms[class_label] = criticism_indices
    
    return prototypes, criticisms

# Select prototypes and criticisms
prototypes, criticisms = select_prototypes_criticisms(
    X_test, y_test, rf_model, n_prototypes=3, n_criticisms=3
)

print("Prototypes and Criticisms:")
for class_label in [0, 1]:
    class_name = 'Benign' if class_label == 1 else 'Malignant'
    print(f"\n{class_name} Class:")
    
    if class_label in prototypes:
        print(f"  Prototypes (typical examples): {prototypes[class_label]}")
        print(f"  Criticisms (edge cases): {criticisms[class_label]}")

# Visualize prototypes and criticisms
from sklearn.decomposition import PCA

pca = PCA(n_components=2)
X_pca = pca.fit_transform(X_test)

fig, ax = plt.subplots(figsize=(10, 8))

# Plot all points
scatter = ax.scatter(X_pca[:, 0], X_pca[:, 1], 
                    c=y_test, cmap='coolwarm', alpha=0.3, s=20)

# Highlight prototypes
for class_label, proto_indices in prototypes.items():
    if len(proto_indices) > 0:
        ax.scatter(X_pca[proto_indices, 0], X_pca[proto_indices, 1],
                  c='green', s=200, marker='^', 
                  edgecolors='black', linewidth=2,
                  label=f'Prototypes (Class {class_label})')

# Highlight criticisms
for class_label, crit_indices in criticisms.items():
    if len(crit_indices) > 0:
        ax.scatter(X_pca[crit_indices, 0], X_pca[crit_indices, 1],
                  c='red', s=200, marker='s',
                  edgecolors='black', linewidth=2,
                  label=f'Criticisms (Class {class_label})')

ax.set_xlabel('First Principal Component')
ax.set_ylabel('Second Principal Component')
ax.set_title('Prototypes and Criticisms in PCA Space')
ax.legend()
plt.colorbar(scatter, label='Class')
plt.show()

Interactive Interpretation Dashboard

print("\n" + "="*60)
print("7. BUILDING AN INTERPRETATION DASHBOARD")
print("="*60)

class ModelInterpretationDashboard:
    """
    Comprehensive dashboard for model interpretation
    """
    
    def __init__(self, model, X_train, y_train, X_test, y_test, feature_names=None):
        self.model = model
        self.X_train = X_train
        self.y_train = y_train
        self.X_test = X_test
        self.y_test = y_test
        self.feature_names = feature_names or X_train.columns
        
    def generate_report(self, instance_idx=0):
        """
        Generate comprehensive interpretation report
        """
        instance = self.X_test.iloc[instance_idx]
        true_label = self.y_test[instance_idx]
        
        # Prediction
        if hasattr(self.model, 'predict_proba'):
            pred_proba = self.model.predict_proba(instance.values.reshape(1, -1))[0]
            pred_class = np.argmax(pred_proba)
        else:
            pred_class = self.model.predict(instance.values.reshape(1, -1))[0]
            pred_proba = [1-pred_class, pred_class]
        
        report = {
            'instance_id': instance_idx,
            'true_label': true_label,
            'predicted_label': pred_class,
            'prediction_probability': pred_proba,
            'interpretations': {}
        }
        
        # 1. Feature importance (global)
        if hasattr(self.model, 'feature_importances_'):
            report['interpretations']['global_importance'] = pd.DataFrame({
                'feature': self.feature_names,
                'importance': self.model.feature_importances_
            }).sort_values('importance', ascending=False).head(10)
        
        # 2. LIME explanation (local)
        lime_exp = SimpleLIME(self.model, num_features=5)
        lime_result = lime_exp.explain_instance(instance, self.X_train, self.feature_names)
        report['interpretations']['lime'] = lime_result
        
        # 3. Counterfactual
        cf_exp = CounterfactualExplainer(self.model)
        cf_result = cf_exp.find_counterfactual(instance, self.X_train, max_changes=3)
        report['interpretations']['counterfactual'] = cf_result
        
        # 4. Similar examples
        from sklearn.neighbors import NearestNeighbors
        nn = NearestNeighbors(n_neighbors=5)
        nn.fit(self.X_train)
        distances, indices = nn.kneighbors(instance.values.reshape(1, -1))
        
        report['interpretations']['similar_examples'] = {
            'indices': indices[0],
            'distances': distances[0],
            'labels': self.y_train[indices[0]]
        }
        
        return report
    
    def visualize_report(self, report):
        """
        Create visualization of interpretation report
        """
        fig = plt.figure(figsize=(16, 12))
        
        # Title
        fig.suptitle(f"Model Interpretation Dashboard - Instance {report['instance_id']}", 
                    fontsize=16, fontweight='bold')
        
        # Grid layout
        gs = fig.add_gridspec(3, 3, hspace=0.3, wspace=0.3)
        
        # 1. Prediction summary
        ax1 = fig.add_subplot(gs[0, :])
        ax1.axis('off')
        
        summary_text = f"""
        True Label: {report['true_label']} | Predicted: {report['predicted_label']}
        Confidence: {report['prediction_probability'][report['predicted_label']]:.2%}
        Model: {self.model.__class__.__name__}
        """
        ax1.text(0.5, 0.5, summary_text, ha='center', va='center',
                fontsize=14, bbox=dict(boxstyle="round,pad=0.5", facecolor="lightblue"))
        
        # 2. LIME explanation
        ax2 = fig.add_subplot(gs[1, 0])
        lime_imp = report['interpretations']['lime']['importances'].head(5)
        colors = ['green' if x > 0 else 'red' for x in lime_imp['importance']]
        ax2.barh(range(len(lime_imp)), lime_imp['importance'].values, color=colors)
        ax2.set_yticks(range(len(lime_imp)))
        ax2.set_yticklabels(lime_imp['feature'].values)
        ax2.set_xlabel('LIME Importance')
        ax2.set_title('Local Explanation (LIME)')
        ax2.axvline(x=0, color='black', linestyle='-', linewidth=0.5)
        ax2.invert_yaxis()
        
        # 3. Global importance
        if 'global_importance' in report['interpretations']:
            ax3 = fig.add_subplot(gs[1, 1])
            global_imp = report['interpretations']['global_importance'].head(5)
            ax3.barh(range(len(global_imp)), global_imp['importance'].values)
            ax3.set_yticks(range(len(global_imp)))
            ax3.set_yticklabels(global_imp['feature'].values)
            ax3.set_xlabel('Feature Importance')
            ax3.set_title('Global Feature Importance')
            ax3.invert_yaxis()
        
        # 4. Counterfactual changes
        if report['interpretations']['counterfactual']:
            ax4 = fig.add_subplot(gs[1, 2])
            cf_changes = report['interpretations']['counterfactual']['changes'].head(3)
            
            y_pos = range(len(cf_changes))
            ax4.barh(y_pos, cf_changes['change'].values)
            ax4.set_yticks(y_pos)
            ax4.set_yticklabels(cf_changes['feature'].values)
            ax4.set_xlabel('Required Change')
            ax4.set_title('Counterfactual: Changes to Flip Prediction')
            ax4.invert_yaxis()
        
        # 5. Similar examples
        ax5 = fig.add_subplot(gs[2, :])
        similar = report['interpretations']['similar_examples']
        
        similar_text = "Similar Training Examples:\n"
        for i, (idx, dist, label) in enumerate(zip(similar['indices'][:3], 
                                                   similar['distances'][:3],
                                                   similar['labels'][:3])):
            similar_text += f"  {i+1}. Index {idx}: Label={label}, Distance={dist:.2f}\n"
        
        ax5.text(0.1, 0.5, similar_text, va='center', fontsize=11,
                bbox=dict(boxstyle="round,pad=0.5", facecolor="lightyellow"))
        ax5.axis('off')
        ax5.set_title('Nearest Neighbors in Training Set', fontsize=12)
        
        plt.tight_layout()
        plt.show()

# Create and display dashboard
dashboard = ModelInterpretationDashboard(
    rf_model, X_train, y_train, X_test, y_test
)

# Generate report for a specific instance
instance_to_explain = 0
report = dashboard.generate_report(instance_to_explain)

print(f"Generated interpretation report for instance {instance_to_explain}")
print(f"True label: {report['true_label']}")
print(f"Predicted: {report['predicted_label']} (confidence: {report['prediction_probability'][report['predicted_label']]:.2%})")

# Visualize the report
dashboard.visualize_report(report)

Best Practices for Model Interpretation

print("\n" + "="*60)
print("BEST PRACTICES FOR MODEL INTERPRETATION")
print("="*60)

best_practices = """
1. **Choose the Right Method for Your Needs:**
   - Global understanding: Feature importance, PDP plots
   - Individual predictions: LIME, SHAP, Counterfactuals
   - Model debugging: Prototypes, criticisms, anchor explanations
   
2. **Validate Interpretations:**
   - Check consistency across different methods
   - Verify with domain experts
   - Test on known edge cases
   
3. **Consider Your Audience:**
   - Technical stakeholders: SHAP values, feature importance
   - Business users: Simple rules, counterfactuals
   - End users: Similar examples, simple explanations
   
4. **Be Aware of Limitations:**
   - Correlation vs causation
   - Approximation errors in local methods
   - Computational cost for large datasets
   
5. **Combine Multiple Approaches:**
   - No single method tells the complete story
   - Use global + local methods together
   - Cross-validate findings
   
6. **Document and Version:**
   - Track interpretation methods used
   - Version control interpretation code
   - Document assumptions and limitations
   
7. **Interactive Exploration:**
   - Build dashboards for stakeholders
   - Allow exploration of different instances
   - Provide multiple views of the same prediction
"""

print(best_practices)

# Create a comparison table of methods
comparison_df = pd.DataFrame({
    'Method': ['Permutation Importance', 'PDP/ICE', 'LIME', 'SHAP', 
               'Counterfactuals', 'Anchors', 'Prototypes'],
    'Scope': ['Global', 'Global', 'Local', 'Both', 'Local', 'Local', 'Global'],
    'Model-Agnostic': ['Yes', 'Yes', 'Yes', 'Yes*', 'Yes', 'Yes', 'Yes'],
    'Speed': ['Slow', 'Medium', 'Fast', 'Slow', 'Fast', 'Medium', 'Fast'],
    'Interpretability': ['High', 'High', 'High', 'Medium', 'Very High', 'High', 'Very High'],
    'Best For': ['Feature ranking', 'Feature effects', 'Single predictions',
                'Detailed attribution', 'What-if analysis', 'Rule extraction', 'Examples']
})

print("\nInterpretation Methods Comparison:")
print(comparison_df.to_string(index=False))
print("\n* SHAP has model-specific optimizations for tree-based models")

Practice Exercises

Exercise 1: Custom Interpretation Framework

Build a comprehensive interpretation framework that:

  1. Automatically selects appropriate interpretation methods based on model type
  2. Generates both global and local explanations
  3. Provides confidence intervals for interpretations
  4. Detects and warns about potential interpretation issues
  5. Exports interpretations in multiple formats (JSON, HTML, PDF)

Exercise 2: Fairness and Bias Detection

Create tools to detect and explain bias:

  1. Implement disparate impact analysis
  2. Create fairness-aware SHAP values
  3. Build counterfactual fairness explanations
  4. Visualize bias across different demographic groups

Exercise 3: Real-time Interpretation API

Develop a production-ready interpretation service:

  1. REST API for model predictions with explanations
  2. Caching for expensive computations
  3. Batch processing capabilities
  4. Monitoring and logging of interpretation requests
  5. A/B testing different explanation methods

Key Takeaways

Summary

Model interpretation and explainability are critical for deploying machine learning responsibly. From understanding global patterns to explaining individual predictions, the techniques covered here provide the tools to open the black box of complex models. Remember that interpretation is not just about satisfying curiosityโ€”it's about building trust, ensuring fairness, meeting regulatory requirements, and ultimately creating ML systems that work reliably in the real world. As models become more complex, the ability to explain their decisions becomes not just valuable, but essential.

๐Ÿ““ Learning Journal

Keep a learning journal โ€” digital or physical. After this lesson, take a few minutes to write down:

  • Key concepts you learned
  • Techniques that clicked for you
  • Questions or confusion points to revisit
  • Ideas you want to try
  • Your progress and feelings about learning this

โœ๏ธ This lesson's prompt: Imagine a model decision that would affect a real person โ€” a loan approval, a medical flag, a hiring screen. What exactly would you need to be able to explain about that prediction before you would feel comfortable deploying the model?

๐Ÿ“ Lesson Summary

๐ŸŽ“ Key Takeaways

  • Interpretability comes in two scopes: global (how the model behaves overall) and local (why it made one specific prediction).
  • Model-agnostic tools โ€” permutation importance, SHAP, LIME, and partial dependence plots โ€” work on any trained estimator.
  • SHAP values are additive and fairly distribute a prediction's value across its features.
  • Explainability supports trust, debugging, fairness, and regulatory compliance โ€” not just curiosity.

๐ŸŽ‰ What You've Accomplished

You can now open a black-box model and explain both its overall logic and its individual decisions โ€” turning an opaque predictor into something a team, an auditor, and an end user can actually trust.

โ“ Common Questions at This Stage

What's the difference between a global and a local explanation?

A global explanation describes the model's overall behavior across the whole dataset (which features matter most, in general). A local explanation focuses on one specific prediction and shows why that particular case came out the way it did.

Should I use SHAP or LIME?

LIME is fast and intuitive but its local approximations can be unstable. SHAP has stronger theoretical guarantees (it is additive and consistent) but costs more compute. For careful analysis, prefer SHAP; for a quick look, LIME is fine.

Does high feature importance mean that feature causes the outcome?

No. Importance measures predictive contribution and association, not causation. A feature can be important because it correlates with the real cause. Treat importance as a clue to investigate, not proof of cause and effect.

๐Ÿ”ญ Looking Ahead

Once you can explain a model, the next step is deploying and monitoring it responsibly โ€” carrying that same transparency into production so the system stays trustworthy over time.

โœ… Before the Next Lesson

๐ŸŒŸ Encouragement for the Journey

A model you can explain is a model people can trust. Learning to open the black box makes you exactly the kind of practitioner teams and stakeholders rely on. Keep going.