Spaces:
Sleeping
Sleeping
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 []
|