Repository navigation
Expand file tree
/
Copy pathingest_reports.py
More file actions
313 lines (258 loc) · 9.85 KB
/
Copy pathingest_reports.py
File metadata and controls
313 lines (258 loc) · 9.85 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
"""
Ingestion script for medical reports RAG system.
This script:
1. Loads parsed medical report data from JSON
2. Deduplicates by unique findings text
3. Generates embeddings (supports Ollama, Gemini, OpenAI)
4. Stores vectors in FAISS with full metadata
Usage:
# Using local Ollama (default, no rate limits!)
uv run python ingest_reports.py --input ./data/reports/all_v2.json
# Using cloud providers
uv run python ingest_reports.py --provider gemini
uv run python ingest_reports.py --provider openai
Environment variables:
RAG_EMBEDDINGS_MODEL: Required for Ollama (e.g., 'nomic-embed-text')
GEMINI_API_KEY: Required for Gemini provider
OPENAI_API_KEY: Required for OpenAI provider
Ollama Setup:
1. Install: curl -fsSL https://ollama.com/install.sh | sh
2. Start server: ollama serve
3. Pull model: ollama pull nomic-embed-text
4. Set env: RAG_EMBEDDINGS_MODEL=nomic-embed-text
"""
import os
import json
import time
import argparse
from typing import List, Dict, Any
from dotenv import load_dotenv
from spoon_ai.rag.embeddings import get_embedding_client
from spoon_ai.rag.vectorstores.faiss_store import FaissVectorStore
def parse_mesh_terms(mesh: Dict[str, List[str]]) -> List[str]:
"""
Parse structured MeSH major terms into individual terms.
Example:
Input: "Opacity/lung/upper lobe/right"
Output: ["opacity", "lung", "upper lobe", "right"]
"""
terms = set()
for term in mesh.get("major", []):
# Split on "/" and normalize
parts = [p.strip().lower() for p in term.split("/")]
terms.update(p for p in parts if p)
# Also include automatic terms
for term in mesh.get("automatic", []):
terms.add(term.strip().lower())
return sorted(terms)
def compute_difficulty(mesh: Dict[str, List[str]]) -> str:
"""
Compute difficulty level based on MeSH terms.
Difficulty Heuristic:
- Easy: 1 condition or contains "normal"
- Medium: 2-3 conditions
- Hard: 4+ conditions OR contains complex terms like "mass", "consolidation",
"cardiomegaly", "effusion", "pneumothorax", "atelectasis"
This helps learners select appropriate cases for their skill level.
"""
major_terms = mesh.get("major", [])
# Check for normal/healthy cases
if any("normal" in t.lower() for t in major_terms):
return "easy"
# Check for complex/serious conditions (indicates harder case)
complex_indicators = [
"mass", "consolidation", "cardiomegaly", "effusion",
"pneumothorax", "atelectasis", "nodule", "tumor", "opacity"
]
has_complex = any(
indicator in t.lower()
for t in major_terms
for indicator in complex_indicators
)
num_conditions = len(major_terms)
if num_conditions >= 4 or has_complex:
return "hard"
elif num_conditions >= 2:
return "medium"
else:
return "easy"
def format_embedding_text(record: Dict[str, Any]) -> str:
"""
Format record fields for embedding using labeled format.
Format (Option B):
Findings: {findings}
Impression: {impression}
Conditions: {parsed mesh terms}
"""
findings = record.get("findings", "")
impression = record.get("impression", "")
mesh_terms = parse_mesh_terms(record.get("mesh", {}))
parts = []
if findings:
parts.append(f"Findings: {findings}")
if impression:
parts.append(f"Impression: {impression}")
if mesh_terms:
parts.append(f"Conditions: {', '.join(mesh_terms)}")
return "\n".join(parts)
def deduplicate_records(records: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""
Deduplicate records by unique findings text.
Keeps the first occurrence (representative image) for each unique finding.
"""
seen_findings = set()
unique_records = []
for record in records:
findings = record.get("findings", "")
if findings and findings not in seen_findings:
seen_findings.add(findings)
unique_records.append(record)
return unique_records
def ingest_reports(
input_path: str,
collection: str = "medical_reports",
limit: int | None = None,
provider: str = "ollama",
) -> int:
"""
Main ingestion function.
Args:
input_path: Path to the JSON file with parsed reports
collection: FAISS collection name
limit: Optional limit on number of records (for testing)
provider: Embedding provider ("ollama", "gemini", "openai", or "hash" for offline)
Returns the number of records ingested.
"""
# Load environment variables
load_dotenv()
# Validate provider and API key
if provider == "openai" and not os.getenv("OPENAI_API_KEY"):
raise ValueError("OPENAI_API_KEY environment variable is required for OpenAI provider")
elif provider == "gemini" and not os.getenv("GEMINI_API_KEY"):
raise ValueError("GEMINI_API_KEY environment variable is required for Gemini provider")
elif provider == "ollama" and not os.getenv("RAG_EMBEDDINGS_MODEL"):
raise ValueError("RAG_EMBEDDINGS_MODEL environment variable is required for Ollama provider (e.g., 'nomic-embed-text')")
# Load data
print(f"Loading data from {input_path}...")
with open(input_path, "r", encoding="utf-8") as f:
records = json.load(f)
print(f"Loaded {len(records)} records")
# Deduplicate
records = deduplicate_records(records)
print(f"After deduplication: {len(records)} unique records")
# Apply limit if specified
if limit and limit < len(records):
records = records[:limit]
print(f"Limited to {limit} records for testing")
if not records:
print("No records to ingest")
return 0
# Prepare embedding texts and metadata
print("Preparing embedding texts...")
texts = []
ids = []
metadatas = []
for record in records:
embedding_text = format_embedding_text(record)
difficulty = compute_difficulty(record.get("mesh", {}))
mesh_terms = parse_mesh_terms(record.get("mesh", {}))
texts.append(embedding_text)
ids.append(record["image_id"])
metadatas.append({
"image_id": record["image_id"],
"image_path": record["image_path"],
"findings": record.get("findings", ""),
"impression": record.get("impression", ""),
"mesh_terms": mesh_terms,
"difficulty": difficulty,
"text": embedding_text, # Store for retrieval
})
# Generate embeddings
print(f"Generating embeddings using {provider} (this may take a few minutes)...")
if provider == "ollama":
# Local Ollama embeddings - no rate limits!
embedding_client = get_embedding_client(
provider="ollama",
openai_model=os.getenv("RAG_EMBEDDINGS_MODEL"),
)
elif provider == "openai":
embedding_client = get_embedding_client(
provider="openai",
openai_api_key=os.getenv("OPENAI_API_KEY"),
openai_model="text-embedding-3-small",
)
elif provider == "gemini":
# Spoon OS uses openai_model param for Gemini model name (confusingly)
embedding_client = get_embedding_client(
provider="gemini",
openai_model="gemini-embedding-001",
)
else:
# Hash-based offline embeddings (for testing)
embedding_client = get_embedding_client(provider="hash")
# Process in batches with rate limit handling
batch_size = 50 # Smaller batches to avoid rate limits
all_embeddings = []
max_retries = 5
for i in range(0, len(texts), batch_size):
batch = texts[i:i + batch_size]
# Retry with exponential backoff
for attempt in range(max_retries):
try:
batch_embeddings = embedding_client.embed(batch)
all_embeddings.extend(batch_embeddings)
print(f" Processed {min(i + batch_size, len(texts))}/{len(texts)} records")
break
except Exception as e:
if "429" in str(e) or "rate" in str(e).lower():
wait_time = 2 ** attempt * 5 # 5, 10, 20, 40, 80 seconds
print(f" Rate limited, waiting {wait_time}s (attempt {attempt + 1}/{max_retries})")
time.sleep(wait_time)
else:
raise
else:
raise RuntimeError(f"Failed to process batch after {max_retries} retries")
# Small delay between batches to avoid rate limits
time.sleep(0.5)
# Store in FAISS
print(f"Storing vectors in FAISS collection '{collection}'...")
store = FaissVectorStore()
# Delete existing collection if it exists
store.delete_collection(collection)
# Add new vectors
store.add(
collection=collection,
ids=ids,
embeddings=all_embeddings,
metadatas=metadatas,
)
print(f"Successfully ingested {len(records)} records into '{collection}'")
return len(records)
def main():
parser = argparse.ArgumentParser(description="Ingest medical reports into RAG system")
parser.add_argument(
"--input",
default="./data/reports/all_v2.json",
help="Path to parsed reports JSON file"
)
parser.add_argument(
"--collection",
default="medical_reports",
help="FAISS collection name"
)
parser.add_argument(
"--limit",
type=int,
default=None,
help="Limit number of records to ingest (for testing)"
)
parser.add_argument(
"--provider",
choices=["ollama", "gemini", "openai", "hash"],
default="ollama",
help="Embedding provider: ollama (default, local), gemini, openai, or hash (offline)"
)
args = parser.parse_args()
ingest_reports(args.input, args.collection, limit=args.limit, provider=args.provider)
if __name__ == "__main__":
main()