Team Ai
Apppublic

Hanan-Alnakhal/Lab-test-decoder

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
build_vector_db.py212 linesDownload Raw Back to root
1"""
2Build Vector Database for Lab Report Decoder
3Uses Hugging Face sentence-transformers for embeddings
4"""
5
6import os
7from pathlib import Path
8from sentence_transformers import SentenceTransformer
9import chromadb
10from chromadb.config import Settings
11import glob
12
13def load_documents_from_directory(directory: str) -> list:
14    """Load all text files from a directory"""
15    documents = []
16    
17    if not os.path.exists(directory):
18        print(f"โš ๏ธ  Directory not found: {directory}")
19        return documents
20    
21    # Find all .txt files
22    txt_files = glob.glob(os.path.join(directory, "**", "*.txt"), recursive=True)
23    
24    for filepath in txt_files:
25        try:
26            with open(filepath, 'r', encoding='utf-8') as f:
27                content = f.read()
28                if content.strip():
29                    documents.append({
30                        'content': content,
31                        'source': filepath,
32                        'filename': os.path.basename(filepath)
33                    })
34        except Exception as e:
35            print(f"Error reading {filepath}: {e}")
36    
37    return documents
38
39def chunk_text(text: str, chunk_size: int = 1000, overlap: int = 200) -> list:
40    """Split text into overlapping chunks"""
41    chunks = []
42    start = 0
43    
44    while start < len(text):
45        end = start + chunk_size
46        chunk = text[start:end]
47        
48        # Try to break at sentence boundary
49        if end < len(text):
50            last_period = chunk.rfind('.')
51            last_newline = chunk.rfind('\n')
52            break_point = max(last_period, last_newline)
53            
54            if break_point > chunk_size * 0.5:  # Only if break point is reasonable
55                chunk = chunk[:break_point + 1]
56                end = start + break_point + 1
57        
58        chunks.append(chunk.strip())
59        start = end - overlap
60    
61    return chunks
62
63def build_knowledge_base():
64    """Build the vector database from medical documents"""
65    
66    print("๐Ÿ“š Loading medical documents...")
67    
68    # Load documents from data directory
69    data_dir = 'data/'
70    all_documents = []
71    
72    if not os.path.exists(data_dir):
73        print(f"โš ๏ธ  Creating data directory: {data_dir}")
74        os.makedirs(data_dir, exist_ok=True)
75        os.makedirs(os.path.join(data_dir, 'lab_markers'), exist_ok=True)
76        os.makedirs(os.path.join(data_dir, 'nutrition'), exist_ok=True)
77        os.makedirs(os.path.join(data_dir, 'conditions'), exist_ok=True)
78        print("โš ๏ธ  Please add medical reference documents to the data/ folder")
79        return None
80    
81    # Load from all subdirectories
82    for subdir in ['lab_markers', 'nutrition', 'conditions']:
83        subdir_path = os.path.join(data_dir, subdir)
84        docs = load_documents_from_directory(subdir_path)
85        all_documents.extend(docs)
86    
87    if not all_documents:
88        print("โš ๏ธ  No documents found in data/ directory")
89        print("Please add .txt files with medical information")
90        return None
91    
92    print(f"โœ… Loaded {len(all_documents)} documents")
93    
94    # Chunk documents
95    print("โœ‚๏ธ  Splitting documents into chunks...")
96    all_chunks = []
97    all_metadata = []
98    
99    for doc in all_documents:
100        chunks = chunk_text(doc['content'], chunk_size=1000, overlap=200)
101        for i, chunk in enumerate(chunks):
102            all_chunks.append(chunk)
103            all_metadata.append({
104                'source': doc['source'],
105                'filename': doc['filename'],
106                'chunk_id': i
107            })
108    
109    print(f"โœ… Created {len(all_chunks)} text chunks")
110    
111    # Load embedding model
112    print("๐Ÿง  Loading embedding model (this may take a moment)...")
113    embedding_model = SentenceTransformer('all-MiniLM-L6-v2')
114    print("โœ… Embedding model loaded")
115    
116    # Create embeddings
117    print("๐Ÿ”„ Creating embeddings (this may take a few minutes)...")
118    embeddings = embedding_model.encode(
119        all_chunks,
120        show_progress_bar=True,
121        convert_to_numpy=True
122    )
123    print(f"โœ… Created {len(embeddings)} embeddings")
124    
125    # Create ChromaDB collection
126    print("๐Ÿ’พ Building ChromaDB vector store...")
127    
128    # Initialize client
129    db_path = "./chroma_db"
130    client = chromadb.PersistentClient(path=db_path)
131    
132    # Delete existing collection if it exists
133    try:
134        client.delete_collection("lab_reports")
135        print("๐Ÿ—‘๏ธ  Deleted existing collection")
136    except:
137        pass
138    
139    # Create new collection
140    collection = client.create_collection(
141        name="lab_reports",
142        metadata={"description": "Medical lab report information"}
143    )
144    
145    # Add documents in batches
146    batch_size = 100
147    for i in range(0, len(all_chunks), batch_size):
148        batch_chunks = all_chunks[i:i + batch_size]
149        batch_embeddings = embeddings[i:i + batch_size].tolist()
150        batch_ids = [f"doc_{j}" for j in range(i, i + len(batch_chunks))]
151        batch_metadata = all_metadata[i:i + batch_size]
152        
153        collection.add(
154            documents=batch_chunks,
155            embeddings=batch_embeddings,
156            ids=batch_ids,
157            metadatas=batch_metadata
158        )
159    
160    print("โœ… Vector database built successfully!")
161    print(f"๐Ÿ“ Database location: {db_path}")
162    print(f"๐Ÿ“Š Total vectors: {len(all_chunks)}")
163    
164    return collection
165
166def test_retrieval(collection):
167    """Test the retrieval system"""
168    if collection is None:
169        print("\nโš ๏ธ  No collection to test")
170        return
171    
172    print("\n๐Ÿ” Testing retrieval system...")
173    
174    # Load embedding model for queries
175    embedding_model = SentenceTransformer('all-MiniLM-L6-v2')
176    
177    test_queries = [
178        "What does low hemoglobin mean?",
179        "What foods are high in iron?",
180        "Normal range for glucose"
181    ]
182    
183    for query in test_queries:
184        print(f"\n๐Ÿ“ Query: {query}")
185        
186        # Create query embedding
187        query_embedding = embedding_model.encode(query).tolist()
188        
189        # Search
190        results = collection.query(
191            query_embeddings=[query_embedding],
192            n_results=2
193        )
194        
195        if results and results['documents']:
196            print(f"  โœ… Found {len(results['documents'][0])} relevant documents")
197            print(f"  ๐Ÿ“„ Top result preview: {results['documents'][0][0][:150]}...")
198        else:
199            print("  โŒ No results found")
200
201if __name__ == "__main__":
202    print("๐Ÿš€ Building Lab Report Decoder Vector Database\n")
203    
204    # Build the database
205    collection = build_knowledge_base()
206    
207    # Test it
208    if collection:
209        test_retrieval(collection)
210        print("\n๐ŸŽ‰ Setup complete! You can now run the Flask application.")
211    else:
212        print("\nโš ๏ธ  Please add medical documents to the data/ folder and run again.")