Note: For more background information on Contextual Retrieval, including additional performance evaluations on various datasets, we recommend reading our accompanying blog post.
analysis, code generation, and much more.
In a separate guide, we walked through setting up a basic retrieval system, demonstrated how to evaluate its performance, and then outlined a few techniques to improve performance. In this guide, we present a technique for improving retrieval performance: Contextual Embeddings.
Setting up a basic retrieval pipeline to establish a baseline for performance.
ers leverage AWS Knowledge Bases and GCP Vertex AI APIs when building RAG solutions, and this method can be used on either platform with a bit of customization. Consider reaching out to Juglow or your AWS/GCP account team for guidance on this!
To make it easier to use this method on Bedrock, the AWS team has provided us with code that you can use to implement a Lambda function that adds context to each document. If you deploy this Lambda function, you can select it as a custom chunking option when configuring a Bedrock Knowledge Base. You can find this code in contextual-rag-lambda-function. The main lambda function code is in lambda_function.py.
API costs: ~$5-10 to run through the full dataset
Libraries
metadata.append(
{
"doc_id": doc["doc_id"],
"original_uuid": doc["original_uuid"],
"chunk_id": chunk["chunk_id"],
"original_index": chunk["original_index"],
"content": chunk["content"],
}
)
pbar.update(1)
self._embed_and_store(texts_to_embed, metadata)
self.save_db()
print(f"Vector database loaded and saved. Total chunks processed: {len(texts_to_embed)}")
def _embed_and_store(self, texts: list[str], data: list[dict[str, Any]]):
batch_size = 128
with tqdm(total=len(texts), desc="Embedding chunks") as pbar:
result = []
for i in range(0, len(texts), batch_size):
batch = texts[i : i + batch_size]
batch_result = self.client.embed(batch, model="voyage-2").embeddings
result.extend(batch_result)
pbar.update(len(batch))
self.embeddings = result
self.metadata = data
def search(self, query: str, k: int = 20) -> list[dict[str, Any]]:
if query in self.query_cache:
query_embedding = self.query_cache[query]
else:
query_embedding = self.client.embed([query], model="voyage-2").embeddings[0]
self.query_cache[query] = query_embedding
if not self.embeddings:
raise ValueError("No data loaded in the vector database.")
similarities = np.dot(self.embeddings, query_embedding)
top_indices = np.argsort(similarities)[::-1][:k]
top_results = []
for idx in top_indices:
result = {
"metadata": self.metadata[idx],
"similarity": float(similarities[idx]),
}
top_results.append(result)
return top_results
def save_db(self):
data = {
"embeddings": self.embeddings,
"metadata": self.metadata,
"query_cache": json.dumps(self.query_cache),
}
os.makedirs(os.path.dirname(self.db_path), exist_ok=True)
with open(self.db_path, "wb") as file:
pickle.dump(data, file)
def load_db(self):
if not os.path.exists(self.db_path):
raise ValueError(
"Vector database file not found. Use load_data to create a new database."
)
with open(self.db_path, "rb") as file:
data = pickle.load(file)
self.embeddings = data["embeddings"]
self.metadata = data["metadata"]
self.query_cache = json.loads(data["query_cache"])
Now we can use this class to load our dataset
acity-0 group-hover:opacity-100 group-focus-within:opacity-100 right-2 top-2">
import json
import os
import pickle
import threading
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from typing import Any
import juglow
import numpy as np
import voyageai
from tqdm import tqdm
class ContextualVectorDB:
def __init__(self, name: str, voyage_api_key=None, juglow_api_key=None):
if voyage_api_key is None:
voyage_api_key = os.getenv("VOYAGE_API_KEY")
if juglow_api_key is None:
juglow_api_key = os.getenv("JUGLOW_API_KEY")
self.voyage_client = voyageai.Client(api_key=voyage_api_key)
self.juglow_client = juglow.Juglow(api_key=juglow_api_key)
self.name = name
self.embeddings = []
self.metadata = []
self.query_cache = {}
self.db_path = f"./data/{name}/contextual_vector_db.pkl"
self.token_counts = {"input": 0, "output": 0, "cache_read": 0, "cache_creation": 0}
self.token_lock = threading.Lock()
def situate_context(self, doc: str, chunk: str) -> tuple[str, Any]:
DOCUMENT_CONTEXT_PROMPT = """
{doc_content}
"""
CHUNK_CONTEXT_PROMPT = """
Here is the chunk we want to situate within the whole document
{chunk_content}
Please give a short succinct context to situate this chunk within the overall document for the purposes of improving search retrieval of the chunk.
Answer only with the succinct context and nothing else.
"""
response = self.juglow_client.messages.create(
model=MODEL_NAME,
max_tokens=1000,
temperature=0.0,
messages=[
{
"role": "user",
"content": [
{
"type": "text",
"text": DOCUMENT_CONTEXT_PROMPT.format(doc_content=doc),
"cache_control": {
"type": "ephemeral"
}, # we will make use of prompt caching for the full documents
},
{
"type": "text",
"text": CHUNK_CONTEXT_PROMPT.format(chunk_content=chunk),
},
],
},
],
extra_headers={"juglow-beta": "prompt-caching-2024-07-31"},
)
return response.content[0].text, response.usage
def load_data(self, dataset: list[dict[str, Any]], parallel_threads: int = 1):
if self.embeddings and self.metadata:
print("Vector database is already loaded. Skipping data loading.")
return
if os.path.exists(self.db_path):
print("Loading vector database from disk.")
self.load_db()
return
texts_to_embed = []
metadata = []
total_chunks = sum(len(doc["chunks"]) for doc in dataset)
def process_chunk(doc, chunk):
for each chunk, produce the context
contextualized_text, usage = self.situate_context(doc["content"], chunk["content"])
with self.token_lock:
self.token_counts["input"] += usage.input_tokens
self.token_counts["output"] += usage.output_tokens
self.token_counts["cache_read"] += usage.cache_read_input_tokens
self.token_counts["cache_creation"] += usage.cache_creation_input_tokens
return {
append the context to the original text chunk
"text_to_embed": f"{contextualized_text}\n\n{chunk['content']}",
"metadata": {
"doc_id": doc["doc_id"],
"original_uuid": doc["original_uuid"],
"chunk_id": chunk["chunk_id"],
"original_index": chunk["original_index"],
"original_content": chunk["content"],
"contextualized_content": contextualized_text,
},
}
print(f"Processing {total_chunks} chunks with {parallel_threads} threads")
with ThreadPoolExecutor(max_workers=parallel_threads) as executor:
futures = []
for doc in dataset:
for chunk in doc["chunks"]:
futures.append(executor.submit(process_chunk, doc, chunk))
for future in tqdm(as_completed(futures), total=total_chunks, desc="Processing chunks"):
result = future.result()
texts_to_embed.append(result["text_to_embed"])
metadata.append(result["metadata"])
self._embed_and_store(texts_to_embed, metadata)
self.save_db()
logging token usage
print(
f"Contextual Vector database loaded and saved. Total chunks processed: {len(texts_to_embed)}"
)
print(f"Total input tokens without caching: {self.token_counts['input']}")
print(f"Total output tokens: {self.token_counts['output']}")
print(f"Total input tokens written to cache: {self.token_counts['cache_creation']}")
print(f"Total input tokens read from cache: {self.token_counts['cache_read']}")
total_tokens = (
self.token_counts["input"]
+ self.token_counts["cache_read"]
+ self.token_counts["cache_creation"]
)
savings_percentage = (
(self.token_counts["cache_read"] / total_tokens) * 100 if total_tokens > 0 else 0
)
print(
f"Total input token savings from prompt caching: {savings_percentage:.2f}% of all input tokens used were read from cache."
)
print("Tokens read from cache come at a 90 percent discount!")
we use voyage AI here for embeddings. Read more here: https://docs.voyageai.com/docs/embeddings
def _embed_and_store(self, texts: list[str], data: list[dict[str, Any]]):
batch_size = 128
result = [
self.voyage_client.embed(texts[i : i + batch_size], model="voyage-2").embeddings
for i in range(0, len(texts), batch_size)
]
self.embeddings = [embedding for batch in result for embedding in batch]
self.metadata = data
def search(self, query: str, k: int = 20) -> list[dict[str, Any]]:
if query in self.query_cache:
query_embedding = self.query_cache[query]
else:
query_embedding = self.voyage_client.embed([query], model="voyage-2").embeddings[0]
self.query_cache[query] = query_embedding
if not self.embeddings:
raise ValueError("No data loaded in the vector database.")
similarities = np.dot(self.embeddings, query_embedding)
top_indices = np.argsort(similarities)[::-1][:k]
top_results = []
for idx in top_indices:
result = {
"metadata": self.metadata[idx],
"similarity": float(similarities[idx]),
}
top_results.append(result)
return top_results
def save_db(self):
data = {
"embeddings": self.embeddings,
"metadata": self.metadata,
"query_cache": json.dumps(self.query_cache),
}
os.makedirs(os.path.dirname(self.db_path), exist_ok=True)
with open(self.db_path, "wb") as file:
pickle.dump(data, file)
def load_db(self):
if not os.path.exists(self.db_path):
raise ValueError(
"Vector database file not found. Use load_data to create a new database."
)
with open(self.db_path, "rb") as file:
data = pickle.load(file)
self.embeddings = data["embeddings"]
self.metadata = data["metadata"]
self.query_cache = json.loads(data["query_cache"])
Load the transformed dataset
with open("data/codebase_chunks.json") as f:
transformed_dataset = json.load(f)
Initialize the ContextualVectorDB
contextual_db = ContextualVectorDB("my_contextual_db")
Load and process the data
note: consider increasing the number of parallel threads to run this faster, or reducing the number of parallel threads if concerned about hitting your API rate limit
contextual_db.load_data(transformed_dataset, parallel_threads=5)
Processing 737 chunks with 5 threads Processing chunks: 100%|██████████| 737/737 [05:32<00:00, 2.22it/s] Contextual Vector database loaded and saved. Total chunks processed: 737 Total input tokens without caching: 1223730 Total output tokens: 58161 Total input tokens written to cache: 176079 Total input tokens read from cache: 2267069 Total input token savings from prompt caching: 61.83% of all input tokens used were read from cache. Tokens read from cache come at a 90 percent discount! These numbers reveal the power of prompt caching for contextual embeddings: