Hanan-Alnakhal/Lab-test-decoder
0
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}")