Team Ai
Apppublic

Hanan-Alnakhal/Lab-test-decoder

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
rag_engine.py231 linesDownload Raw Back to root
1from sentence_transformers import SentenceTransformer2from transformers import pipeline3import chromadb4from typing import List, Dict5from pdf_extractor import LabResult6import os7 8class LabReportRAG:9    """RAG system for explaining lab results - Fast and efficient"""10    11    def __init__(self, db_path: str = "./chroma_db"):12        """Initialize the RAG system with fast models"""13        14        print("๐Ÿ”„ Loading models (optimized for speed)...")15        16        # Fast embedding model17        self.embedding_model = SentenceTransformer('all-MiniLM-L6-v2')18        print("โœ… Embeddings loaded")19        20        # Use FAST text generation model21        print("๐Ÿ”„ Loading text generation model...")22        try:23            # Use Flan-T5 - efficient instruction tuned model24            self.text_generator = pipeline(25                "text2text-generation",26                model="google/flan-t5-base", 27                max_length=512, # Increased slightly for better answers28                device=-1  # Force CPU29            )30            print("โœ… Text generation model loaded (Flan-T5-base)")31        except Exception as e:32            print(f"โš ๏ธ Model loading error: {e}")33            self.text_generator = None34        35        # Load vector store36        try:37            self.client = chromadb.PersistentClient(path=db_path)38            self.collection = self.client.get_collection("lab_reports")39            print("โœ… Vector database loaded")40        except Exception as e:41            print(f"โš ๏ธ Vector database not found: {e}")42            self.collection = None43    44    def _retrieve_context(self, query: str, k: int = 2) -> str:45        """Retrieve relevant context from vector database"""46        if self.collection is None:47            return None48        49        try:50            # Create query embedding51            query_embedding = self.embedding_model.encode(query).tolist()52            53            # Query the collection54            results = self.collection.query(55                query_embeddings=[query_embedding],56                n_results=k57            )58            59            # Combine documents60            if results and results['documents'] and len(results['documents'][0]) > 0:61                # Calculate simple relevance check (if distances available)62                # For now, we assume if Chroma returns it, it's the best it has.63                context = "\n".join(results['documents'][0])64                return context[:1500] # Increased context window for the LLM65            else:66                return None67        except Exception as e:68            print(f"Retrieval error: {e}")69            return None70    71    def _generate_text(self, prompt: str) -> str:72        """Generate text using Flan-T5"""73        if self.text_generator is None:74            return "AI model not available."75        76        try:77            result = self.text_generator(78                prompt,79                max_length=256,80                do_sample=True,81                temperature=0.5, # Lower temperature for more factual answers82                num_return_sequences=183            )84            return result[0]['generated_text'].strip()85        except Exception as e:86            print(f"Generation error: {e}")87            return "Unable to generate explanation."88 89    def explain_result(self, result: LabResult) -> str:90        """Generate explanation for a single lab result using LLM"""91        92        print(f"  Explaining: {result.test_name} ({result.status})...")93        94        # 1. Base Message (Deterministic)95        base_msg = f"Your {result.test_name} is {result.value} {result.unit} ({result.status}). "96        97        # 2. Retrieve Context98        query = f"What does {result.status} {result.test_name} mean? causes and simple explanation"99        context = self._retrieve_context(query, k=1)100        101        if not context:102            return base_msg + "Please consult your doctor for interpretation."103 104        # 3. Generate Simple Explanation via LLM105        # We instruct Flan-T5 to summarize the context simply106        prompt = f"""107        Explain simply in 2 sentences what {result.status} {result.test_name} means based on this medical info.108        109        Medical Info: {context}110        111        Explanation:"""112        113        ai_explanation = self._generate_text(prompt)114        115        # 4. Construct Final Output116        final_output = f"""{base_msg}117        118๐Ÿ’ก Analysis: {ai_explanation}119 120(Reference Range: {result.reference_range})"""121        122        return final_output123 124    def answer_followup_question(self, question: str, lab_results: List[LabResult]) -> str:125        """126        Answer questions using RAG + LLM. 127        Handles out-of-context questions strictly.128        """129        print(f"๐Ÿ’ฌ Processing question: {question[:50]}...")130        131        # 1. Retrieve Medical Context132        medical_context = self._retrieve_context(question, k=2)133        print(medical_context)134        # 2. Check for "Out of Context"135        # If no documents found in DB, or query seems totally unrelated136        if not medical_context:137            return "Sorry I can't answer this Question"138 139        # 3. Prepare Patient Context (Current Results)140        # We give the model a snapshot of the patient's actual data141        relevant_results = [f"{r.test_name}: {r.value} ({r.status})" for r in lab_results]142        patient_data = ", ".join(relevant_results[:5]) # Top 5 results to save space143        144        # 4. Construct Strict Prompt145        # Flan-T5 instruction to enforce the fallback phrase146        prompt = f"""147        Answer the question based strictly on the Medical Context and Patient Results provided below.148        If the answer cannot be found in the context, or if the question is not about health/labs, reply exactly "Sorry I can't answer this Question".149 150        Medical Context: {medical_context}151        152        Patient Results: {patient_data}153        154        Question: {question}155        156        Answer:"""157        158        # 5. Generate159        answer = self._generate_text(prompt)160        161        # Double check: sometimes models hallucinate. 162        # If the context was extremely short/weak, we might want to override, 163        # but relying on the prompt instructions is standard for T5.164        return answer165 166    def generate_summary(self, results: List[LabResult]) -> str:167        """Generate a summary using the LLM"""168        print("๐Ÿ“Š Generating summary...")169        170        abnormal = [r for r in results if r.status in ['high', 'low']]171        172        if not abnormal:173            return "โœ… All results are normal. Great job maintaining your health!"174            175        # Create a prompt for the summary176        abnormal_text = ", ".join([f"{r.test_name} is {r.status}" for r in abnormal])177        178        # Get general context about these specific abnormal tests179        context = self._retrieve_context(f"health implications of {abnormal_text}", k=1)180        181        prompt = f"""182        The patient has these abnormal lab results: {abnormal_text}.183        Based on this medical info: {context}184        185        Write a short, encouraging 2-sentence summary advising them to see a doctor.186        187        Summary:"""188        189        ai_summary = self._generate_text(prompt)190        191        return f"โš ๏ธ **Abnormal Results Detected**\n\n{ai_summary}\n\nDetailed changes:\n" + "\n".join([f"- {r.test_name}: {r.value} {r.unit}" for r in abnormal])192 193    # Keep the other helper methods if needed or rely on the new logic194    # The explain_all_results wrapper is still useful195    def explain_all_results(self, results: List[LabResult]) -> Dict[str, str]:196        explanations = {}197        for result in results:198            explanations[result.test_name] = self.explain_result(result)199        return explanations200 201# Testing block202if __name__ == "__main__":203    print("Testing RAG system...")204    try:205        rag = LabReportRAG()206        207        # Mock Data208        from pdf_extractor import LabResult209        results = [210            LabResult("Hemoglobin", "10.5", "g/dL", "12.0-15.5", "low"),211            LabResult("Glucose", "95", "mg/dL", "70-100", "normal")212        ]213        214        # Test 1: Explanation using LLM215        print("\n--- Test Explanation ---")216        print(rag.explain_result(results[0]))217        218        # Test 2: Follow up (Valid)219        print("\n--- Test Valid Question ---")220        q1 = "What foods should I eat for low hemoglobin?"221        print(f"Q: {q1}")222        print(f"A: {rag.answer_followup_question(q1, results)}")223        224        # Test 3: Follow up (Out of Context)225        print("\n--- Test Invalid Question ---")226        q2 = "Who is the president of the USA?"227        print(f"Q: {q2}")228        print(f"A: {rag.answer_followup_question(q2, results)}")229 230    except Exception as e:231        print(f"\nโŒ Error: {e}")