RAG / phase1_ingestion.py
sumitnewold's picture
Upload 10 files
76bd1fc verified
Raw
History Blame Contribute Delete
8.31 kB
"""
Phase 1 β€” Document Ingestion Pipeline
Loads the Infosys AR PDF, extracts tables, tags entities, builds
Parent-Child ChromaDB retriever, and runs a smoke-test retrieval.
"""
import json
import os
import time
import pandas as pd
import pdfplumber
import spacy
from dotenv import load_dotenv
from langchain_community.document_loaders import PyPDFLoader
from langchain_community.vectorstores import Chroma
from langchain_core.stores import InMemoryStore
from langchain_experimental.text_splitter import SemanticChunker
from langchain_groq import ChatGroq
from langchain_huggingface import HuggingFaceEmbeddings
from langchain_text_splitters import RecursiveCharacterTextSplitter, TextSplitter
from langchain_classic.retrievers import ParentDocumentRetriever
load_dotenv()
# ── Paths ────────────────────────────────────────────────────────────────────
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
PDF_PATH = os.path.join(BASE_DIR, "infosys-ar-25.pdf")
TABLE_STORE_PATH = os.path.join(BASE_DIR, "table_store.json")
CHROMA_DIR = os.path.join(BASE_DIR, "finance_db")
# ── Step 1: LLM connection check ─────────────────────────────────────────────
def build_llm(verify: bool = False):
llm = ChatGroq(
model="llama-3.3-70b-versatile",
temperature=0,
api_key=os.environ["GROQ_API_KEY"],
)
# Only ping Groq when explicitly asked (standalone phase scripts).
# The app skips this to avoid burning free-tier tokens on every launch.
if verify:
response = llm.invoke("Say exactly: Groq connected successfully")
print(response.content)
return llm
# ── Step 2: Embeddings ───────────────────────────────────────────────────────
def build_embeddings():
print("Loading embedding model... (first load downloads ~420MB, wait for it)")
embeddings = HuggingFaceEmbeddings(
model_name="sentence-transformers/all-mpnet-base-v2",
model_kwargs={"device": "cpu"},
encode_kwargs={"normalize_embeddings": True},
)
test_vector = embeddings.embed_query("What is Infosys revenue?")
print(f"Embeddings working. Vector length: {len(test_vector)}")
print(f"First 5 values: {test_vector[:5]}")
return embeddings
# ── Step 3: Load PDF ─────────────────────────────────────────────────────────
def load_pdf(pdf_path):
loader = PyPDFLoader(pdf_path)
pages = loader.load()
print(f"Total pages loaded: {len(pages)}")
print(f"\nPage 1 metadata: {pages[0].metadata}")
print(f"\nPage 1 preview:\n{pages[0].page_content[:500]}")
return pages
# ── Step 4: Table extraction ─────────────────────────────────────────────────
def extract_tables_from_pdf(pdf_path):
extracted_tables = []
with pdfplumber.open(pdf_path) as pdf:
for page_num, page in enumerate(pdf.pages, start=1):
for table_idx, table in enumerate(page.extract_tables()):
if not table or len(table) < 2:
continue
try:
df = pd.DataFrame(table[1:], columns=table[0]).fillna("")
table_text = f"Table on page {page_num}:\n" + df.to_string(
index=False
)
extracted_tables.append(
{
"page": page_num,
"table_index": table_idx,
"source_file": pdf_path.split("/")[-1],
"text_representation": table_text,
"headers": [str(c) for c in df.columns],
}
)
except Exception as e:
print(
f"[TABLE EXTRACTION] Skipped malformed table on page {page_num}: {e}"
)
print(f"[INGESTION] Extracted {len(extracted_tables)} tables")
return extracted_tables
# ── Step 5: spaCy NER entity tagging ─────────────────────────────────────────
def tag_financial_entities(documents):
nlp = spacy.load("en_core_web_sm")
for doc in documents:
ents = nlp(doc.page_content).ents
doc.metadata["entities"] = json.dumps({
"organizations": list(set(e.text for e in ents if e.label_ == "ORG")),
"monetary_values": list(set(e.text for e in ents if e.label_ == "MONEY")),
"percentages": list(set(e.text for e in ents if e.label_ == "PERCENT")),
"dates": list(set(e.text for e in ents if e.label_ == "DATE")),
})
return documents
# ── Step 6: Parent-Child Retriever ───────────────────────────────────────────
class SemanticParentSplitter(TextSplitter):
"""Wraps SemanticChunker to satisfy ParentDocumentRetriever's TextSplitter type check."""
def __init__(self, semantic_chunker, **kwargs):
super().__init__(**kwargs)
self._chunker = semantic_chunker
def split_text(self, text):
return self._chunker.split_text(text)
def build_retriever(embeddings, chroma_dir):
parent_splitter = SemanticParentSplitter(
SemanticChunker(
embeddings,
breakpoint_threshold_type="percentile",
breakpoint_threshold_amount=95,
)
)
child_splitter = RecursiveCharacterTextSplitter(chunk_size=400, chunk_overlap=50)
vectorstore = Chroma(
collection_name="child_chunks",
embedding_function=embeddings,
persist_directory=chroma_dir,
)
store = InMemoryStore()
retriever = ParentDocumentRetriever(
vectorstore=vectorstore,
docstore=store,
child_splitter=child_splitter,
parent_splitter=parent_splitter,
search_kwargs={"k": 5},
)
print("Retriever built successfully")
return retriever, store
# ── Step 7: Smoke-test indexing & retrieval ──────────────────────────────────
def run_smoke_test(retriever, store, tagged_pages):
test_subset = tagged_pages[20:40]
print(f"\n[TEST] Indexing {len(test_subset)} pages...")
t0 = time.time()
retriever.add_documents(test_subset)
print(f"[TEST] Done in {time.time() - t0:.1f}s")
print(f"[TEST] Parent docs stored: {len(list(store.yield_keys()))}")
results = retriever.invoke("What are the key risk factors?")
print(f"\nRetrieved {len(results)} parent chunks")
for r in results[:2]:
print(f"--- Page {r.metadata.get('page')} | len={len(r.page_content)} ---")
print(r.page_content[:300], "\n")
# ── Main ─────────────────────────────────────────────────────────────────────
if __name__ == "__main__":
# 1. LLM
llm = build_llm(verify=True)
# 2. Embeddings
embeddings = build_embeddings()
# 3. Load PDF
assert os.path.exists(PDF_PATH), f"PDF not found at {PDF_PATH}"
pages = load_pdf(PDF_PATH)
# 4. Tables
tables = extract_tables_from_pdf(PDF_PATH)
with open(TABLE_STORE_PATH, "w") as f:
json.dump(tables, f, indent=2)
print(f"Saved {len(tables)} tables to {TABLE_STORE_PATH}")
# 5. NER tagging
tagged_pages = tag_financial_entities(pages)
print(f"Tagged {len(tagged_pages)} pages")
print(f"Sample (page 5) entities: {tagged_pages[4].metadata['entities']}")
# 6. Retriever
retriever, store = build_retriever(embeddings, CHROMA_DIR)
# 7. Smoke test
run_smoke_test(retriever, store, tagged_pages)