contract-classifier / explainability.py
pryyyynz's picture
Basic files, no model
66d7c1e verified
Raw
History Blame Contribute Delete
3.93 kB
"""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 []