File size: 14,217 Bytes
09801ca
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
"""
πŸ“Š SHAP EXPLAINABILITY ENGINE
=============================

Provides model explanations using SHAP (SHapley Additive exPlanations):
- Global feature importance (summary plot)
- Local explanations (waterfall for single predictions)
- Force plots for decision paths

Uses TreeExplainer for tree-based models, KernelExplainer as fallback.
"""

import numpy as np
import pandas as pd
import base64
import io
import logging
from typing import Dict, List, Any, Optional, Tuple

import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt

logger = logging.getLogger(__name__)

# Check if SHAP is available
try:
    import shap
    HAS_SHAP = True
    logger.info("βœ… SHAP library available for explainability")
except ImportError:
    HAS_SHAP = False
    logger.warning("⚠️ SHAP not installed. Run: pip install shap")


def _fig_to_base64(fig) -> str:
    """Convert matplotlib figure to base64 string"""
    buf = io.BytesIO()
    fig.savefig(buf, format='png', dpi=150, bbox_inches='tight', 
                facecolor='white', edgecolor='none')
    buf.seek(0)
    img_base64 = base64.b64encode(buf.read()).decode('utf-8')
    plt.close(fig)
    return f"data:image/png;base64,{img_base64}"


class SHAPExplainer:
    """
    πŸ” SHAP-based model explainer.
    
    Creates explanations for:
    - Why did the model make this prediction?
    - Which features are most important globally?
    - How does each feature push the prediction?
    """
    
    def __init__(self, model, X_background: np.ndarray, feature_names: List[str] = None):
        """
        Initialize explainer with model and background data.
        
        Args:
            model: Trained sklearn model
            X_background: Background dataset for SHAP (subset of training data)
            feature_names: Optional list of feature names
        """
        self.model = model
        self.feature_names = feature_names or [f"Feature {i}" for i in range(X_background.shape[1])]
        self.explainer = None
        self.explainer_type = None
        
        if not HAS_SHAP:
            logger.warning("SHAP not available")
            return
        
        # Sample background data for efficiency (max 100 samples)
        if len(X_background) > 100:
            indices = np.random.choice(len(X_background), 100, replace=False)
            X_background = X_background[indices]
        
        self.X_background = X_background
        
        # Choose appropriate explainer
        model_name = type(model).__name__
        tree_models = ['RandomForest', 'GradientBoosting', 'XGB', 'LGBM', 
                       'CatBoost', 'ExtraTrees', 'DecisionTree', 'HistGradient']
        
        try:
            if any(name in model_name for name in tree_models):
                # TreeExplainer is fast and exact for tree models
                self.explainer = shap.TreeExplainer(model)
                self.explainer_type = "TreeExplainer"
                logger.info(f"βœ… Using TreeExplainer for {model_name}")
            else:
                # KernelExplainer works for any model but is slower
                if hasattr(model, 'predict_proba'):
                    self.explainer = shap.KernelExplainer(model.predict_proba, X_background)
                else:
                    self.explainer = shap.KernelExplainer(model.predict, X_background)
                self.explainer_type = "KernelExplainer"
                logger.info(f"βœ… Using KernelExplainer for {model_name}")
        except Exception as e:
            logger.warning(f"⚠️ Could not create SHAP explainer: {e}")
            self.explainer = None
    
    def explain_prediction(self, X_single: np.ndarray) -> Dict[str, Any]:
        """
        Explain a single prediction.
        
        Args:
            X_single: Single sample to explain (1D or 2D array)
            
        Returns:
            Dictionary with SHAP values, base value, and explanation chart
        """
        if not HAS_SHAP or self.explainer is None:
            return {"error": "SHAP not available"}
        
        try:
            # Ensure 2D
            if X_single.ndim == 1:
                X_single = X_single.reshape(1, -1)
            
            # Calculate SHAP values
            shap_values = self.explainer.shap_values(X_single)
            
            # Handle multi-class (use positive class for binary)
            if isinstance(shap_values, list):
                # Classification with multiple classes
                shap_values = shap_values[1] if len(shap_values) == 2 else shap_values[0]
            
            shap_values = shap_values.flatten()
            
            # Get base value
            if hasattr(self.explainer, 'expected_value'):
                base_value = self.explainer.expected_value
                if isinstance(base_value, np.ndarray):
                    base_value = base_value[1] if len(base_value) == 2 else base_value[0]
            else:
                base_value = 0
            
            # Create feature contributions dict
            contributions = []
            for i, (name, value, shap_val) in enumerate(zip(
                self.feature_names, X_single.flatten(), shap_values
            )):
                contributions.append({
                    "feature": name,
                    "value": float(value),
                    "shap_value": float(shap_val),
                    "direction": "positive" if shap_val > 0 else "negative"
                })
            
            # Sort by absolute SHAP value
            contributions.sort(key=lambda x: abs(x["shap_value"]), reverse=True)
            
            # Generate waterfall chart
            waterfall_chart = self._generate_waterfall_chart(
                shap_values, X_single.flatten(), base_value
            )
            
            return {
                "success": True,
                "base_value": float(base_value),
                "prediction_contribution": float(sum(shap_values)),
                "contributions": contributions[:50],  # Top 50 features (increased from 15)
                "waterfall_chart": waterfall_chart
            }
            
        except Exception as e:
            logger.error(f"SHAP explanation error: {e}")
            return {"error": str(e)}
    
    def get_global_importance(self, X_sample: np.ndarray = None) -> Dict[str, Any]:
        """
        Get global feature importance using SHAP.
        
        Args:
            X_sample: Sample of data to calculate importance (uses background if None)
            
        Returns:
            Dictionary with importance values and summary chart
        """
        if not HAS_SHAP or self.explainer is None:
            return {"error": "SHAP not available"}
        
        try:
            X_sample = X_sample if X_sample is not None else self.X_background
            
            # Limit samples for speed
            if len(X_sample) > 200:
                indices = np.random.choice(len(X_sample), 200, replace=False)
                X_sample = X_sample[indices]
            
            # Calculate SHAP values
            shap_values = self.explainer.shap_values(X_sample)
            
            # Handle multi-class
            if isinstance(shap_values, list):
                shap_values = shap_values[1] if len(shap_values) == 2 else shap_values[0]
            
            # Calculate mean absolute SHAP values for importance
            importance = np.abs(shap_values).mean(axis=0)
            
            # Create importance dict
            feature_importance = []
            for name, imp in zip(self.feature_names, importance):
                feature_importance.append({
                    "feature": name,
                    "importance": float(imp)
                })
            
            # Sort by importance
            feature_importance.sort(key=lambda x: x["importance"], reverse=True)
            
            # Generate summary plot
            summary_chart = self._generate_summary_chart(shap_values, X_sample)
            beeswarm_chart = self._generate_beeswarm_chart(shap_values, X_sample)
            
            return {
                "success": True,
                "feature_importance": feature_importance[:20],
                "summary_chart": summary_chart,
                "beeswarm_chart": beeswarm_chart
            }
            
        except Exception as e:
            logger.error(f"Global importance error: {e}")
            return {"error": str(e)}
    
    def _generate_waterfall_chart(self, shap_values: np.ndarray, 
                                   feature_values: np.ndarray, 
                                   base_value: float) -> Optional[str]:
        """Generate SHAP waterfall chart for single prediction"""
        try:
            fig, ax = plt.subplots(figsize=(10, 8))
            
            # Sort by absolute value
            indices = np.argsort(np.abs(shap_values))[::-1][:12]  # Top 12
            
            sorted_shap = shap_values[indices]
            sorted_names = [f"{self.feature_names[i][:20]} = {feature_values[i]:.2g}" 
                           for i in indices]
            
            # Calculate cumulative values
            cumsum = np.cumsum(sorted_shap)
            
            # Colors: green for positive, red for negative
            colors = ['#10B981' if v > 0 else '#EF4444' for v in sorted_shap]
            
            # Create waterfall
            y_pos = np.arange(len(sorted_shap))
            bars = ax.barh(y_pos, sorted_shap, color=colors, edgecolor='white', height=0.7)
            
            # Add base value annotation
            final_value = base_value + sum(shap_values)
            ax.axvline(x=0, color='gray', linestyle='--', alpha=0.5)
            
            ax.set_yticks(y_pos)
            ax.set_yticklabels(sorted_names, fontsize=9)
            ax.set_xlabel('SHAP Value (impact on prediction)', fontsize=11)
            ax.set_title(f'πŸ” Prediction Explanation\nBase: {base_value:.3f} β†’ Final: {final_value:.3f}', 
                        fontsize=12, pad=10)
            
            # Add value labels
            for bar, val in zip(bars, sorted_shap):
                width = bar.get_width()
                ax.text(width + 0.01, bar.get_y() + bar.get_height()/2,
                       f'{val:+.3f}', va='center', fontsize=8,
                       color='#10B981' if val > 0 else '#EF4444')
            
            ax.invert_yaxis()
            plt.tight_layout()
            
            return _fig_to_base64(fig)
            
        except Exception as e:
            logger.warning(f"Waterfall chart error: {e}")
            return None
    
    def _generate_summary_chart(self, shap_values: np.ndarray, X: np.ndarray) -> Optional[str]:
        """Generate SHAP summary bar chart"""
        try:
            fig, ax = plt.subplots(figsize=(10, 8))
            
            # Mean absolute SHAP values
            importance = np.abs(shap_values).mean(axis=0)
            indices = np.argsort(importance)[::-1][:15]  # Top 15
            
            sorted_importance = importance[indices]
            sorted_names = [self.feature_names[i][:25] for i in indices]
            
            # Gradient colors
            colors = plt.cm.viridis(np.linspace(0.3, 0.9, len(sorted_importance)))[::-1]
            
            y_pos = np.arange(len(sorted_importance))
            bars = ax.barh(y_pos, sorted_importance, color=colors, edgecolor='white', height=0.7)
            
            ax.set_yticks(y_pos)
            ax.set_yticklabels(sorted_names, fontsize=10)
            ax.set_xlabel('Mean |SHAP Value|', fontsize=11)
            ax.set_title('πŸ“Š Global Feature Importance (SHAP)', fontsize=13, pad=10)
            
            # Add value labels
            for bar, val in zip(bars, sorted_importance):
                ax.text(bar.get_width() + 0.002, bar.get_y() + bar.get_height()/2,
                       f'{val:.3f}', va='center', fontsize=9)
            
            ax.invert_yaxis()
            plt.tight_layout()
            
            return _fig_to_base64(fig)
            
        except Exception as e:
            logger.warning(f"Summary chart error: {e}")
            return None
    
    def _generate_beeswarm_chart(self, shap_values: np.ndarray, X: np.ndarray) -> Optional[str]:
        """Generate SHAP beeswarm plot showing feature value distributions"""
        try:
            if not HAS_SHAP:
                return None
                
            fig, ax = plt.subplots(figsize=(10, 8))
            
            # Use shap's built-in summary plot
            shap.summary_plot(shap_values, X, 
                             feature_names=self.feature_names,
                             max_display=15,
                             show=False,
                             plot_size=None)
            
            plt.title('🐝 SHAP Beeswarm (Feature Impact Distribution)', fontsize=12, pad=10)
            plt.tight_layout()
            
            # Get current figure
            fig = plt.gcf()
            return _fig_to_base64(fig)
            
        except Exception as e:
            logger.warning(f"Beeswarm chart error: {e}")
            return None


def explain_model(model, X_train: np.ndarray, X_explain: np.ndarray, 
                  feature_names: List[str] = None) -> Dict[str, Any]:
    """
    Convenience function to explain a model.
    
    Args:
        model: Trained model
        X_train: Training data (for background)
        X_explain: Data to explain (single sample or multiple)
        feature_names: Optional feature names
        
    Returns:
        Dictionary with explanations and charts
    """
    explainer = SHAPExplainer(model, X_train, feature_names)
    
    result = {}
    
    # Global importance
    global_exp = explainer.get_global_importance()
    if global_exp.get("success"):
        result["global"] = global_exp
    
    # Single prediction explanation
    if X_explain is not None:
        local_exp = explainer.explain_prediction(X_explain)
        if local_exp.get("success"):
            result["local"] = local_exp
    
    return result