Haijun Platform Docs
ID

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:

On this page
Libraries