""" Standalone Gradio dashboard for contract classification with LIME explanations. Features: - Upload single or multiple documents (PDF, DOCX, DOC, TXT) - Show prediction and confidence - Show class probability chart - Highlight influential text via LIME HTML - Download CSV for batch results This app loads the same enhanced TF-IDF model used by the API if available. For a fully standalone setup, place the model file under web/models/. """ import os import io import csv import tempfile import shutil import logging from typing import List, Dict, Any, Tuple import mimetypes import numpy as np import pandas as pd # Document processing deps import pdfplumber from docx import Document as DocxDocument from PIL import Image import pytesseract # Optional OCR PDF rasterization try: import fitz # PyMuPDF PYMUPDF_AVAILABLE = True except Exception: PYMUPDF_AVAILABLE = False import gradio as gr from explainability import ContractExplainer import pickle logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) # ------------------------------ # Model loading # ------------------------------ MODEL = None VECTORIZER = None CLASS_NAMES: List[str] = [] FEATURE_SELECTOR = None EXPLAINER: ContractExplainer | None = None def _candidate_model_paths() -> List[str]: return [ os.path.join(os.path.dirname(__file__), "models", "enhanced_tfidf_gradient_boosting_model.pkl"), os.path.join(os.path.dirname(__file__), "..", "enhanced_models_output", "models", "enhanced_tfidf_gradient_boosting_model.pkl"), os.path.join(os.path.dirname(__file__), "..", "models_output", "models", "random_forest_model.pkl"), ] def load_model_if_needed() -> Tuple[bool, str]: global MODEL, VECTORIZER, CLASS_NAMES, FEATURE_SELECTOR, EXPLAINER if EXPLAINER is not None: return True, "Model already loaded" last_error = "" for path in _candidate_model_paths(): try: if not os.path.exists(path): continue with open(path, "rb") as f: data = pickle.load(f) MODEL = data["classifier"] VECTORIZER = data["vectorizer"] CLASS_NAMES = data["class_names"] FEATURE_SELECTOR = data.get("feature_selector") EXPLAINER = ContractExplainer( MODEL, VECTORIZER, CLASS_NAMES, FEATURE_SELECTOR) logger.info(f"Loaded model from: {path}") return True, f"Loaded model: {os.path.basename(path)}" except Exception as e: last_error = str(e) logger.exception("Failed loading model") return False, last_error or "Model file not found. Place model under web/models/." # ------------------------------ # Text extraction # ------------------------------ def extract_text_from_pdf(file_path: str) -> str: text = "" try: with pdfplumber.open(file_path) as pdf: for page in pdf.pages: page_text = page.extract_text() if page_text: text += page_text + "\n" except Exception as e: logger.warning(f"pdfplumber failed: {e}") if text.strip(): return text.strip() # OCR fallback if not PYMUPDF_AVAILABLE: return text.strip() try: doc = fitz.open(file_path) for page_index in range(len(doc)): page = doc.load_page(page_index) pix = page.get_pixmap(matrix=fitz.Matrix(2, 2)) img = Image.open(io.BytesIO(pix.tobytes("png"))) text += pytesseract.image_to_string(img, lang="eng") + "\n" doc.close() except Exception as e: logger.warning(f"OCR fallback failed: {e}") return text.strip() def extract_text_from_docx(file_path: str) -> str: try: doc = DocxDocument(file_path) return "\n".join(p.text for p in doc.paragraphs).strip() except Exception as e: logger.warning(f"DOCX extraction failed: {e}") return "" def extract_text_from_doc(file_path: str) -> str: # Best-effort: try antiword try: import subprocess result = subprocess.run(["antiword", file_path], capture_output=True, text=True) if result.returncode == 0: return result.stdout.strip() except Exception: pass return "" def preprocess_text(text: str) -> str: if not text: raise ValueError("Empty text") text = text.strip() text = " ".join(text.split()) if len(text) < 10: raise ValueError("Text too short for classification") return text # ------------------------------ # Inference and explanation # ------------------------------ def classify_text(text: str, num_features: int = 1) -> Dict[str, Any]: ok, msg = load_model_if_needed() if not ok: raise RuntimeError(f"Model not available: {msg}") text = preprocess_text(text) explanation = EXPLAINER.explain_prediction(text, num_features=num_features) if not explanation.get("success"): raise RuntimeError(explanation.get("error", "Explanation failed")) # Compute prediction using the same preprocessing as the model (no full probs for speed) features = VECTORIZER.transform([text]) if FEATURE_SELECTOR is not None: features = FEATURE_SELECTOR.transform(features) probs = MODEL.predict_proba(features)[0] # Align predicted class using model.classes_ model_classes = list(getattr(MODEL, "classes_", CLASS_NAMES)) predicted_index = int(np.argmax(probs)) explanation["prediction"] = model_classes[predicted_index] explanation["confidence"] = float(probs[predicted_index]) return explanation def classify_text_fast(text: str) -> Dict[str, Any]: """Fast prediction without LIME (used for batch).""" ok, msg = load_model_if_needed() if not ok: raise RuntimeError(f"Model not available: {msg}") text = preprocess_text(text) features = VECTORIZER.transform([text]) if FEATURE_SELECTOR is not None: features = FEATURE_SELECTOR.transform(features) probs = MODEL.predict_proba(features)[0] # Use model-provided class order to avoid misalignment model_classes = list(getattr(MODEL, "classes_", CLASS_NAMES)) predicted_index = int(np.argmax(probs)) predicted_class = model_classes[predicted_index] confidence = float(probs[predicted_index]) return { "prediction": predicted_class, "confidence": confidence, "class_probabilities": {cls: float(probs[i]) for i, cls in enumerate(model_classes)}, "text": text[:200] + "..." if len(text) > 200 else text, } def classify_file(tmp_path: str, mime_type: str, num_features: int = 10) -> Dict[str, Any]: if mime_type == "application/pdf": text = extract_text_from_pdf(tmp_path) elif mime_type == "application/vnd.openxmlformats-officedocument.wordprocessingml.document": text = extract_text_from_docx(tmp_path) elif mime_type == "application/msword": text = extract_text_from_doc(tmp_path) else: # Treat as plain text with open(tmp_path, "r", encoding="utf-8", errors="ignore") as f: text = f.read() return classify_text(text, num_features=num_features) def _extract_key_phrase_fast(text: str) -> str: """Approximate influential phrase quickly using top TF-IDF term and context.""" try: tokens = VECTORIZER.transform([text]) if hasattr(tokens, "toarray"): arr = tokens.toarray()[0] else: arr = tokens.A[0] if arr.sum() == 0: return "" top_idx = int(arr.argmax()) feature_names = getattr(VECTORIZER, "get_feature_names_out", None) if feature_names is None: return "" feat = VECTORIZER.get_feature_names_out()[top_idx] # Build phrase around first occurrence words = text.split() feat_lower = feat.lower() for i, w in enumerate(words): if feat_lower in w.lower(): start_idx = max(0, i - 2) end_idx = min(len(words), i + 4) phrase = " ".join(words[start_idx:end_idx]).strip( '.,!?;:"()[]{}') if len(phrase.split()) >= 3: return phrase # fallback to sentence-level break # fallback: first sentence for sep in [". ", "\n", "? ", "! "]: if sep in text: return text.split(sep, 1)[0].strip() return text[:120] except Exception: return "" def classify_file_fast(tmp_path: str, mime_type: str) -> Dict[str, Any]: if mime_type == "application/pdf": text = extract_text_from_pdf(tmp_path) elif mime_type == "application/vnd.openxmlformats-officedocument.wordprocessingml.document": text = extract_text_from_docx(tmp_path) elif mime_type == "application/msword": text = extract_text_from_doc(tmp_path) else: with open(tmp_path, "r", encoding="utf-8", errors="ignore") as f: text = f.read() result = classify_text_fast(text) # Add fast key phrase extraction result["key_phrase"] = _extract_key_phrase_fast(text) return result # ------------------------------ # Gradio UI callbacks # ------------------------------ def predict_single(file_path: str): if not file_path: return "No file uploaded", None, None, None try: mime, _ = mimetypes.guess_type(file_path) mime = mime or "text/plain" result = classify_file(file_path, mime, num_features=1) pred = f"Prediction: {result['prediction']} (confidence: {result['confidence']:.3f})" # One-line influential statement top_feats = result.get("important_features", []) key_phrase = top_feats[0][0] if top_feats else _extract_key_phrase_fast( result.get("full_text", "")) html = result.get("explanation_html", "") key_line = key_phrase return pred, html, key_line except Exception as e: return f"Error: {e}", None, None def predict_batch(file_paths: List[str], num_features: int): if not file_paths: return None, None rows = [] for fp in file_paths: try: mime, _ = mimetypes.guess_type(fp) mime = mime or "text/plain" # Use fast prediction (no LIME) for batch speed result = classify_file_fast(fp, mime) key_phrase = "" rows.append({ "filename": os.path.basename(fp), "prediction": result["prediction"], "confidence": float(result["confidence"]), "key_phrase": result.get("key_phrase", key_phrase), }) except Exception as e: rows.append({ "filename": os.path.basename(fp), "prediction": "", "confidence": 0.0, "key_phrase": f"Error: {e}", }) df = pd.DataFrame(rows) # Write CSV to a temporary file and return the path for DownloadButton tmp_csv = tempfile.NamedTemporaryFile( delete=False, suffix="_batch_results.csv") try: with open(tmp_csv.name, "w", encoding="utf-8", newline="") as f: df.to_csv(f, index=False) finally: pass return df, tmp_csv.name # ------------------------------ # Build UI # ------------------------------ with gr.Blocks(title="Contract Classifier") as demo: gr.Markdown(""" **Contract Classification Dashboard** - Upload single or multiple documents - View prediction, probabilities, and highlighted influential text - Download CSV for batch results """) with gr.Tab("Single Document"): with gr.Row(): file_in = gr.File( label="Upload document (PDF/DOCX/DOC/TXT)", type="filepath") with gr.Row(): predict_btn = gr.Button("Predict") with gr.Row(): pred_out = gr.Textbox(label="Prediction", lines=1) with gr.Row(): html_out = gr.HTML(label="LIME Explanation (highlighted text)") with gr.Row(): preview_out = gr.Textbox(label="Key Phrase", lines=6) predict_btn.click( predict_single, inputs=[file_in], outputs=[pred_out, html_out, preview_out] ) with gr.Tab("Batch"): with gr.Row(): files_in = gr.File( label="Upload multiple documents", file_count="multiple", type="filepath") with gr.Row(): batch_btn = gr.Button("Run Batch") with gr.Row(): table_out = gr.Dataframe(label="Batch Results", interactive=False) with gr.Row(): download_btn = gr.DownloadButton( label="Download CSV") def _batch_and_prepare(files): df, csv_path = predict_batch(files, num_features=3) return df, gr.update(value=csv_path) batch_btn.click( _batch_and_prepare, inputs=[files_in], outputs=[table_out, download_btn] ) # Ensure model loads at launch for quicker first prediction def _warmup(): ok, msg = load_model_if_needed() return f"Model: {'ready' if ok else 'not ready'} — {msg}" warmup_status = gr.Markdown() demo.load( _warmup, inputs=None, outputs=warmup_status ) if __name__ == "__main__": # Let Gradio pick an available port automatically demo.launch(server_name="0.0.0.0", show_api=False)