mirror of
https://github.com/schwartz1375/genai-essentials
synced 2026-08-09 13:08:20 +00:00
- Switch from nomic-embed-text to embeddinggemma for embeddings - Add directory verification cells to ensure required directories exist - Add error handling for corrupted ChromaDB databases with auto-recovery - Create persist directories before ChromaDB initialization
24 KiB
24 KiB
In [1]:
# Install necessary packages
#!pip install --upgrade pip --quiet
#!pip install --upgrade langchain langchain-ollama langchain-chroma chromadb pymupdf "unstructured[local-inference]" ollama --quiet
# Required Ollama models - ensure these are installed:
# ollama pull gemma3
# ollama pull embeddinggemmaIn [2]:
# --- Configuration ---
import os
from langchain.text_splitter import RecursiveCharacterTextSplitter
from langchain_ollama import OllamaEmbeddings, OllamaLLM
from langchain_chroma import Chroma
# Prevent Chroma telemetry issues
os.environ["CHROMADB_DISABLE_TELEMETRY"] = "1"
# Define directories
DATA_DIR = "./data"
CHROMA_DIR = "./chromadb_store"
SUPPORTED_FORMATS = {
"pdfs": "PyMuPDFLoader",
"markdowns": "UnstructuredMarkdownLoader",
"txt": "TextLoader",
}
EMBEDDING_MODEL = "embeddinggemma"
LLM_MODEL = "gemma3"
# --- Utility: Write Permission Check ---
def assert_directory_writable(path):
os.makedirs(path, exist_ok=True)
test_file = os.path.join(path, "test_write.tmp")
try:
with open(test_file, "w") as f:
f.write("test")
os.remove(test_file)
print(f"✅ Confirmed write access to: {path}")
except Exception as e:
raise PermissionError(f"❌ Cannot write to {path}: {e}")
# --- Document Loaders ---
from langchain_community.document_loaders import (
DirectoryLoader,
PyMuPDFLoader,
TextLoader,
UnstructuredMarkdownLoader,
)
def load_documents(base_dir, formats_dict):
glob_patterns = {
PyMuPDFLoader: "*.pdf",
UnstructuredMarkdownLoader: "*.md",
TextLoader: "*.txt",
}
loaders = {
"PyMuPDFLoader": PyMuPDFLoader,
"UnstructuredMarkdownLoader": UnstructuredMarkdownLoader,
"TextLoader": TextLoader,
}
docs = []
for subdir, loader_name in formats_dict.items():
dir_path = os.path.join(base_dir, subdir)
os.makedirs(dir_path, exist_ok=True)
loader_cls = loaders[loader_name]
pattern = glob_patterns[loader_cls]
loader = DirectoryLoader(
dir_path,
glob=pattern,
loader_cls=loader_cls,
show_progress=True,
)
loaded = loader.load()
print(f"📂 Loaded {len(loaded)} from {dir_path}")
if loaded:
print(f"📝 Sample:\n{loaded[0].page_content[:500]}")
docs.extend(loaded)
print(f"📄 Total docs loaded: {len(docs)}")
return docs
# --- Text Splitting ---
def split_documents(docs):
# Smaller chunks to keep embeddings safe for Ollama 0.13.x
splitter = RecursiveCharacterTextSplitter(
chunk_size=500, # was 1000
chunk_overlap=100, # was 200
separators=["\n\n", "\n", ".", "!", "?", ",", " ", ""],
length_function=len,
)
chunks = splitter.split_documents(docs)
print(f"🔪 {len(chunks)} chunks created")
return chunks
# --- Vector Store Build ---
def build_vectordb(docs, embedding_model, persist_dir):
import shutil
assert_directory_writable(persist_dir)
chunks = split_documents(docs)
try:
vectordb = Chroma.from_documents(
documents=chunks,
embedding=embedding_model,
persist_directory=persist_dir,
)
except Exception as e:
if "Database error" in str(e) or "readonly" in str(e):
print(f"⚠️ Corrupted database, recreating...")
shutil.rmtree(persist_dir, ignore_errors=True)
os.makedirs(persist_dir, exist_ok=True)
vectordb = Chroma.from_documents(
documents=chunks,
embedding=embedding_model,
persist_directory=persist_dir,
)
else:
raise
print(f"✅ VectorDB built at: {persist_dir}")
return vectordb
# --- Vector Store Load ---
def load_vectordb(persist_dir, embedding_model):
if not os.path.exists(persist_dir):
raise FileNotFoundError(f"❌ VectorDB directory not found: {persist_dir}. Run with build=True first.")
vectordb = Chroma(
persist_directory=persist_dir,
embedding_function=embedding_model,
)
print(f"✅ VectorDB loaded from: {persist_dir}")
return vectordb
# --- QA Chain Setup ---
from langchain.prompts import PromptTemplate
from langchain.chains import RetrievalQA
qa_template = """
You are a helpful assistant answering questions based on the provided context.
Use the following pieces of context to answer the user's question. If unsure, state you don't know.
Context:
{context}
Question: {question}
Answer:
"""
prompt = PromptTemplate.from_template(qa_template)
def create_qa_chain(vectordb):
try:
llm = OllamaLLM(
model=LLM_MODEL,
base_url="http://localhost:11434", # make endpoint explicit
)
# Test LLM model
llm.invoke("test")
print(f"🤖 LLM ready: {LLM_MODEL}")
except Exception as e:
raise RuntimeError(
f"❌ LLM model '{LLM_MODEL}' not available. "
f"Run: ollama pull {LLM_MODEL}. Error: {e}"
)
return RetrievalQA.from_chain_type(
llm=llm,
retriever=vectordb.as_retriever(search_kwargs={"k": 5}),
return_source_documents=True,
chain_type="stuff",
chain_type_kwargs={"prompt": prompt},
)
# --- Build Pipeline ---
def initialize_pipeline(build=True):
"""
Initializes the RAG pipeline using ChromaDB and Ollama.
Parameters:
build (bool): If True, (re)build the vector database from the document source.
If False, load the existing persisted vector store.
Returns:
RetrievalQA: A ready-to-use QA chain for question answering.
"""
print("🚀 Initializing pipeline...")
# Validate Ollama embedding model is available
try:
embedding_model = OllamaEmbeddings(
model=EMBEDDING_MODEL,
base_url="http://localhost:11434",
num_ctx=2048, # keep context size modest and explicit
# keep_alive can be added as seconds (int) if you like, e.g. keep_alive=300
)
# Test embedding model
embedding_model.embed_query("test")
print(f"✅ Embedding model '{EMBEDDING_MODEL}' is ready")
except Exception as e:
raise RuntimeError(
f"❌ Embedding model '{EMBEDDING_MODEL}' not available. "
f"Run: ollama pull {EMBEDDING_MODEL}. Error: {e}"
)
if build or not os.path.exists(os.path.join(CHROMA_DIR, "chroma.sqlite3")):
docs = load_documents(DATA_DIR, SUPPORTED_FORMATS)
vectordb = build_vectordb(docs, embedding_model, CHROMA_DIR)
else:
vectordb = load_vectordb(CHROMA_DIR, embedding_model)
qa_chain = create_qa_chain(vectordb)
return qa_chain
In [3]:
qa_chain = initialize_pipeline(
build=True
) # Use build=True to (re)build the vector database from the document source.🚀 Initializing pipeline... ✅ Embedding model 'embeddinggemma' is ready
100%|██████████| 20/20 [00:00<00:00, 21.88it/s]
📂 Loaded 406 from ./data/pdfs 📝 Sample: ETSI TS 103 994-1 V1.1.1 (2024-03) Cyber Security (CYBER); Privileged Access Workstations; Part 1: Physical Device TECHNICAL SPECIFICATION
100%|██████████| 8/8 [00:02<00:00, 3.53it/s]
📂 Loaded 8 from ./data/markdowns 📝 Sample: Python Networking Expansion Reference source files: ClientPython.py, ServerPython.py This tutorial will complete the functionality of the client, making it usable (when paired with an equally functional server, which will be written later) for real-world capabilities by adding the networking and logic code necessary for it to exchange query and log data with a server. This will also allow us to explore additional features and nuances of networking in Python. Using SSL to secure connections B
0it [00:00, ?it/s]
📂 Loaded 0 from ./data/txt 📄 Total docs loaded: 414 ✅ Confirmed write access to: ./chromadb_store 🔪 3400 chunks created ✅ VectorDB built at: ./chromadb_store 🤖 LLM ready: gemma3
In [4]:
from IPython.display import Markdown, display
qa_chain = initialize_pipeline(
build=False
) # Use build=False to load the existing persisted vector store
# Ask a question
def ask(question):
print(f"\n❓ Question: {question}")
response = qa_chain.invoke({"query": question})
display(Markdown(f"**Answer:**\n\n{response['result']}"))
print("\n📄 Sources:")
seen_sources = set()
for doc in response["source_documents"][:2]:
source_info = f"{os.path.basename(doc.metadata.get('source', 'unknown'))}, Page {doc.metadata.get('page', 'N/A')}"
if source_info not in seen_sources:
print(source_info)
seen_sources.add(source_info)
ask("Why are pickles bad?")
ask("What is a ClickFix?")
ask("What is LummaC2?")🚀 Initializing pipeline... ✅ Embedding model 'embeddinggemma' is ready ✅ VectorDB loaded from: ./chromadb_store 🤖 LLM ready: gemma3 ❓ Question: Why are pickles bad?
<IPython.core.display.Markdown object>
📄 Sources: Tutorial-4-MultipleConnections.md, Page N/A ❓ Question: What is a ClickFix?
<IPython.core.display.Markdown object>
📄 Sources: clickfix-attacks-sector-alert-tlpclear.pdf, Page 0 ❓ Question: What is LummaC2?
<IPython.core.display.Markdown object>
📄 Sources: aa25-141b-threat-actors-deploy-lummac2-malware-to-exfiltrate-sensitive-data-from-organizations.pdf, Page 1 aa25-141b-threat-actors-deploy-lummac2-malware-to-exfiltrate-sensitive-data-from-organizations.pdf, Page 7
In [5]:
from langchain_ollama import OllamaLLM, OllamaEmbeddings
from langchain_chroma import Chroma
from langchain.prompts import PromptTemplate
from langchain.chains import RetrievalQA
from IPython.display import Markdown, display
import os
# Set env for Chroma (just in case)
os.environ["CHROMADB_DISABLE_TELEMETRY"] = "1"
# Config
CHROMA_DIR = "./chromadb_store"
EMBEDDING_MODEL = "embeddinggemma"
LLM_MODEL = "gemma3"
# Prompt template
qa_template = """
You are a helpful assistant answering questions based on the provided context.
Use the following pieces of context to answer the user's question. If unsure, state you don't know.
Context:
{context}
Question: {question}
Answer:
"""
prompt = PromptTemplate.from_template(qa_template)
# Reload vector DB + embedding
embedding_model = OllamaEmbeddings(model=EMBEDDING_MODEL)
vectordb = Chroma(persist_directory=CHROMA_DIR, embedding_function=embedding_model)
# Recreate QA chain
llm = OllamaLLM(model=LLM_MODEL)
qa_chain = RetrievalQA.from_chain_type(
llm=llm,
retriever=vectordb.as_retriever(search_kwargs={"k": 5}),
return_source_documents=True,
chain_type="stuff",
chain_type_kwargs={"prompt": prompt},
)
# Ask a question
def ask(question):
print(f"\n❓ Question: {question}")
response = qa_chain.invoke({"query": question})
display(Markdown(f"**Answer:**\n\n{response['result']}"))
print("\n📄 Sources:")
seen_sources = set()
for doc in response["source_documents"][:2]:
source_info = f"{os.path.basename(doc.metadata.get('source', 'unknown'))}, Page {doc.metadata.get('page', 'N/A')}"
if source_info not in seen_sources:
print(source_info)
seen_sources.add(source_info)
In [6]:
ask("What is a ClickFix?")
ask("What is LummaC2?")
ask("Why is the sky blue?")❓ Question: What is a ClickFix?
<IPython.core.display.Markdown object>
📄 Sources: clickfix-attacks-sector-alert-tlpclear.pdf, Page 0 ❓ Question: What is LummaC2?
<IPython.core.display.Markdown object>
📄 Sources: aa25-141b-threat-actors-deploy-lummac2-malware-to-exfiltrate-sensitive-data-from-organizations.pdf, Page 1 aa25-141b-threat-actors-deploy-lummac2-malware-to-exfiltrate-sensitive-data-from-organizations.pdf, Page 7 ❓ Question: Why is the sky blue?
<IPython.core.display.Markdown object>
📄 Sources: Tutorial-4-MultipleConnections.md, Page N/A Tutorial-7-PythonNetworkingExpansion.md, Page N/A