File size: 3,934 Bytes
66d7c1e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Standalone copy of LIME-based explainability used by the dashboard."""

import logging
from typing import Dict, List, Any, Optional, Tuple

import numpy as np

try:
    from lime.lime_text import LimeTextExplainer
    LIME_AVAILABLE = True
except ImportError:
    LIME_AVAILABLE = False

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)


class ContractExplainer:
    def __init__(self, model, vectorizer, class_names: List[str], feature_selector=None, random_state: int = 42):
        if not LIME_AVAILABLE:
            raise ImportError(
                "LIME not available. Install with: pip install lime")
        self.model = model
        self.vectorizer = vectorizer
        self.feature_selector = feature_selector
        self.class_names = class_names
        self.random_state = random_state
        self.explainer = LimeTextExplainer(
            class_names=class_names, random_state=random_state)

    def explain_prediction(self, text: str, num_features: int = 10, num_samples: int = 500) -> Dict[str, Any]:
        try:
            def predict_proba_wrapper(texts):
                features = self.vectorizer.transform(texts)
                if self.feature_selector is not None:
                    features = self.feature_selector.transform(features)
                return self.model.predict_proba(features)

            exp = self.explainer.explain_instance(
                text,
                predict_proba_wrapper,
                num_features=num_features,
                num_samples=num_samples,
                top_labels=1,
            )

            all_probs = predict_proba_wrapper([text])[0]
            predicted_index = int(np.argmax(all_probs))
            predicted_class = self.class_names[predicted_index]
            confidence = float(all_probs[predicted_index])

            important_features = exp.as_list(label=predicted_index)
            processed_features = self._get_best_phrase_feature(
                important_features, text)

            return {
                "text": text[:200] + "..." if len(text) > 200 else text,
                "full_text": text,
                "prediction": predicted_class,
                "confidence": confidence,
                "important_features": processed_features,
                "explanation_html": exp.as_html(),
                "num_features": num_features,
                "success": True,
                "explanation_object": exp,
            }
        except Exception as e:
            logger.exception("Explain failed")
            return {"success": False, "error": str(e), "text": text[:200] + "..." if len(text) > 200 else text, "full_text": text}

    def _get_best_phrase_feature(self, important_features: List[Tuple[str, float]], text: str) -> List[Tuple[str, float]]:
        text_lower = text.lower()
        candidate_phrases: List[Tuple[str, float]] = []

        for feature, score in important_features:
            if " " in feature and len(feature.split()) >= 3:
                candidate_phrases.append((feature, abs(float(score))))
            else:
                feature_lower = feature.lower()
                words = text_lower.split()
                for i, word in enumerate(words):
                    if feature_lower in word.lower():
                        start_idx = max(0, i - 2)
                        end_idx = min(len(words), i + 4)
                        context_phrase = " ".join(
                            words[start_idx:end_idx]).strip('.,!?;:"()[]{}')
                        if len(context_phrase.split()) >= 3:
                            candidate_phrases.append(
                                (context_phrase, abs(float(score))))
                            break

        if candidate_phrases:
            best = max(candidate_phrases, key=lambda x: x[1])
            return [best]
        return [important_features[0]] if important_features else []