simplify code

This commit is contained in:
generall
2025-01-13 21:39:48 +01:00
parent 1e428bf25d
commit a4e534d7a3
@@ -37,7 +37,7 @@ To understand the impact, consider the construction of an [**HNSW index**](/arti
- **Vectors per page:** ~700 (ColQwen) or ~1,000 (ColPali)
- **[ef_construct](/documentation/concepts/indexing/#vector-index):** 100 (default)
The number of comparisons required is:
The lower bound estimation for the number of vector comparisions comparisons would be:
$$
700 \times 700 \times 100 = 49 \, \text{millions}
@@ -51,7 +51,7 @@ For ColPali, this number doubles. The result is **extremely slow index construct
We recommend reducing the number of vectors in a PDF page representation for the **first-stage retrieval**. After the first stage retrieval with a reduced amount of vectors, we propose to **rerank** retrieved subset with the original uncompressed representation.
<aside role="status"> You might consider using <b>quantization</b> (e.g., binary quantization) to reduce computational resources. However, as you can see above, quantization does not impact the parameters that determine the number of comparisons, so its effect in this context would be <b>minimal.</b> </aside>
<aside role="status"> You might consider using <b>quantization</b> (e.g., binary quantization) to reduce computational resources. However, as you can see above, quantization does not impact the parameters that determine the number of comparisons, so it will only affect memory consumption.</aside>
The reduction of vectors can be achieved by applying a **mean pooling operation** to the multivector VLLM-generated outputs. Mean pooling averages the values across all vectors within a selected subgroup, condensing multiple vectors into a single representative vector. If done right, it allows the preservation of important information from the original page while significantly reducing the number of vectors.
@@ -87,15 +87,13 @@ In the following sections, we will demonstrate an optimized retrieval algorithm
Install & import required libraries
```python
from colpali_engine.models import ColPali, ColPaliProcessor, ColQwen2, ColQwen2Processor
from datasets import load_dataset
# pip install colpali_engine>=0.3.1
from colpali_engine.models import ColPali, ColPaliProcessor
# pip install qdrant-client>=1.12.0
from qdrant_client import QdrantClient, models
import torch
from tqdm import tqdm
import uuid
```
To run these experiments, we’re using a **Qdrant cluster**. If you’re just getting started, you can set up a **free-tier cluster** for testing and exploration. Follow the instructions in the documentation ["How to Create a Free-Tier Qdrant Cluster"](/documentation/cloud/create-cluster/?q=free+tier#free-clusters)
To run these experiments, we’re using a **Qdrant cluster**. If you’re just getting started, you can set up a **free-tier cluster** for testing and exploration. Follow the instructions in the documentation ["How to Create a Free-Tier Qdrant Cluster"](/documentation/cloud/create-cluster/#free-clusters)
```python
client = QdrantClient(
@@ -104,17 +102,9 @@ client = QdrantClient(
)
```
Download **ColQwen** and **ColPali** models along with their input processors. Make sure to select the backend that suits your setup.
Download **ColPali** model along with its input processors. Make sure to select the backend that suits your setup.
```python
colqwen_model = ColQwen2.from_pretrained(
"vidore/colqwen2-v0.1",
torch_dtype=torch.bfloat16,
device_map="mps", # Use "cuda:0" for GPU, "cpu" for CPU, or "mps" for Apple Silicon
).eval()
colqwen_processor = ColQwen2Processor.from_pretrained("vidore/colqwen2-v0.1")
colpali_model = ColPali.from_pretrained(
"vidore/colpali-v1.3",
torch_dtype=torch.bfloat16,
@@ -124,50 +114,70 @@ colpali_model = ColPali.from_pretrained(
colpali_processor = ColPaliProcessor.from_pretrained("vidore/colpali-v1.3")
```
<details>
<summary> For <b>ColQwen</b> model </summary>
```python
from colpali_engine.models import ColQwen2, ColQwen2Processor
colqwen_model = ColQwen2.from_pretrained(
"vidore/colqwen2-v0.1",
torch_dtype=torch.bfloat16,
device_map="mps", # Use "cuda:0" for GPU, "cpu" for CPU, or "mps" for Apple Silicon
).eval()
colqwen_processor = ColQwen2Processor.from_pretrained("vidore/colqwen2-v0.1")
```
</details>
## Create Qdrant Collections
We will create two separate collections: one for the **ColQwen** model and one for the **ColPali** model. Each collection will include **mean pooled** by rows and columns representations of a PDF page.
We can now create a collection in Qdrant to store the multivector representations of PDF pages generated by **ColPali** or **ColQwen**.
Collection will include **mean pooled** by rows and columns representations of a PDF page, as well as the **original** multivector representation.
<aside role="status"> For the original multivectors generated by the models, we will disable HNSW index construction </aside>
```python
for collection_name in ["colpali_tutorial", "colqwen_tutorial"]:
client.create_collection(
collection_name=collection_name,
vectors_config={
"original":
models.VectorParams( #switch off HNSW
size=128,
distance=models.Distance.COSINE,
multivector_config=models.MultiVectorConfig(
comparator=models.MultiVectorComparator.MAX_SIM
),
hnsw_config=models.HnswConfigDiff(
m=0 #switching off HNSW
)
),
"mean_pooling_columns": models.VectorParams(
size=128,
distance=models.Distance.COSINE,
multivector_config=models.MultiVectorConfig(
comparator=models.MultiVectorComparator.MAX_SIM
)
),
"mean_pooling_rows": models.VectorParams(
client.create_collection(
collection_name=collection_name,
vectors_config={
"original":
models.VectorParams( #switch off HNSW
size=128,
distance=models.Distance.COSINE,
multivector_config=models.MultiVectorConfig(
comparator=models.MultiVectorComparator.MAX_SIM
),
hnsw_config=models.HnswConfigDiff(
m=0 #switching off HNSW
)
),
"mean_pooling_columns": models.VectorParams(
size=128,
distance=models.Distance.COSINE,
multivector_config=models.MultiVectorConfig(
comparator=models.MultiVectorComparator.MAX_SIM
)
}
)
),
"mean_pooling_rows": models.VectorParams(
size=128,
distance=models.Distance.COSINE,
multivector_config=models.MultiVectorConfig(
comparator=models.MultiVectorComparator.MAX_SIM
)
)
}
)
```
## Choose a dataset
We’ll use the **UFO Dataset** by Daniel van Strien for this tutorial. It’s available on Hugging Face; you can download it directly from there.
```python
from datasets import load_dataset
ufo_dataset = "davanstrien/ufo-ColPali"
dataset = load_dataset(ufo_dataset, split="train")
```
@@ -191,262 +201,135 @@ The `get_patches` function is to get the number of `x_patches` (rows) and `y_pat
For ColPali, the numbers will always be 32 by 32; ColQwen will define them dynamically based on the PDF page size.
```python
def get_patches(image_size, model_processor, model, model_name):
if model_name == "colPali":
return model_processor.get_n_patches(
image_size,
patch_size=model.patch_size
)
elif model_name == "colQwen":
return model_processor.get_n_patches(
image_size,
patch_size=model.patch_size,
spatial_merge_size=model.spatial_merge_size
)
return None, None
x_patches, y_patches = model_processor.get_n_patches(
image_size,
patch_size=model.patch_size
)
```
<details>
<summary> For <b>ColQwen</b> model </summary>
```python
model_processor.get_n_patches(
image_size,
patch_size=model.patch_size,
spatial_merge_size=model.spatial_merge_size
)
```
</details>
We choose to **preserve prefix and postfix multivectors**. Our **pooling** operation compresses the multivectors representing **the image tokens** based on the number of rows and columns determined by the model (static 32x32 for ColPali, dynamic XxY for ColQwen). Function retains and integrates the additional multivectors produced by the model back to pooled representations.
That's an illustration of this process:
![Our mean pooling strategy](/documentation/tutorials/pdf-retrieval-at-scale/mean-pooling.png)
Simplified version of pooling for **ColQwen** / **ColPali** models:
`embed_and_mean_pool_batch` is the universal function which embeds a batch of PDF pages with ColPali (ColQwen) and then does the explained mean pooling processing.
<details>
<summary> <span style="background-color: gray; color: black;"> Code of the function </span> </summary>
(see the full version in the [tutorial notebook](https://githubtocolab.com/qdrant/examples/blob/master/pdf-retrieval-at-scale/ColPali_ColQwen_Tutorial.ipynb))
```python
def embed_and_mean_pool_batch(image_batch, model_processor, model, model_name):
#embed
with torch.no_grad():
processed_images = model_processor.process_images(image_batch).to(model.device)
image_embeddings = model(**processed_images)
image_embeddings_batch = image_embeddings.cpu().float().numpy().tolist()
#mean pooling
pooled_by_rows_batch = []
pooled_by_columns_batch = []
for image_embedding, tokenized_image, image in zip(image_embeddings,
processed_images.input_ids,
image_batch):
x_patches, y_patches = get_patches(image.size, model_processor, model, model_name)
#print(f"{model_name} model divided this PDF page in {x_patches} rows and {y_patches} columns")
processed_images = model_processor.process_images(image_batch)
# Image embeddings of shape (batch_size, 1030, 128)
image_embeddings = model(**processed_images)
image_tokens_mask = (tokenized_image == model_processor.image_token_id)
image_tokens = image_embedding[image_tokens_mask].view(x_patches, y_patches, model.dim)
pooled_by_rows = torch.mean(image_tokens, dim=0)
pooled_by_columns = torch.mean(image_tokens, dim=1)
# (1030, 128)
image_embedding = image_embeddings[0] # take the first element of the batch
image_token_idxs = torch.nonzero(image_tokens_mask.int(), as_tuple=False)
first_image_token_idx = image_token_idxs[0].cpu().item()
last_image_token_idx = image_token_idxs[-1].cpu().item()
prefix_tokens = image_embedding[:first_image_token_idx]
postfix_tokens = image_embedding[last_image_token_idx + 1:]
#print(f"There are {len(prefix_tokens)} prefix tokens and {len(postfix_tokens)} in a {model_name} PDF page embedding")
# Now we need to find identify sub-vectors that correspond to the image tokens
# It can be done by selecting tokens corresponding to special `image_token_id`
#adding back prefix and postfix special tokens
pooled_by_rows = torch.cat((prefix_tokens, pooled_by_rows, postfix_tokens), dim=0).cpu().float().numpy().tolist()
pooled_by_columns = torch.cat((prefix_tokens, pooled_by_columns, postfix_tokens), dim=0).cpu().float().numpy().tolist()
pooled_by_rows_batch.append(pooled_by_rows)
pooled_by_columns_batch.append(pooled_by_columns)
# (1030, ) - boolean mask that is True for image tokens
mask = processed_images.input_ids[0] == model_processor.image_token_id
# For convenience we now select only image tokens
# and reshape them to (x_patches, y_patches, dim)
return image_embeddings_batch, pooled_by_rows_batch, pooled_by_columns_batch
# (x_patches, y_patches, 128)
image_tokens = image_embedding[mask].view(x_patches, y_patches, model.dim)
# Now we can apply mean pooling by rows and columns
# (x_patches, 128)
pooled_by_rows = image_tokens.mean(dim=0)
# (y_patches, 128)
pooled_by_columns = image_tokens.mean(dim=1)
# [Optionally] we can also concatenate special tokens to the pooled representations
# (x_patches + 6, 128)
pooled_by_rows = torch.cat([pooled_by_rows, image_embedding[~mask]])
# (y_patches + 6, 128)
pooled_by_columns = torch.cat([pooled_by_columns, image_embedding[~mask]])
```
</details>
## Batch uploading to Qdrant
Below is the function to batch upload multivectors into the collections created earlier to Qdrant.
```python
def upload_batch(original_batch, pooled_by_rows_batch, pooled_by_columns_batch, payload_batch, collection_name):
client.upload_collection(
collection_name=collection_name,
vectors={
"mean_pooling_columns": pooled_by_columns_batch,
"original": original_batch,
"mean_pooling_rows": pooled_by_rows_batch
},
payload=payload_batch,
ids=[str(uuid.uuid4()) for i in range(len(original_batch))]
)
```
Now you can test the uploading process of the **UFO dataset**, pre-processed according to our approach by `embed_and_mean_pool_batch` function.
## Upload to Qdrant
<details>
<summary> <span style="background-color: gray; color: black;"> Uploading "UFO dataset" to "colpali_tutorial" collection </span> </summary>
Upload process is trivial, the only thing to pay attention to is the compute cost for ColPali and ColQwen2 models.
In low-resource environments, it's recommended to use a smaller batch size for embedding and mean pooling.
```python
batch_size = 1 #based on available compute
dataset_source = ufo_dataset
collection_name = "colpali_tutorial"
Full version of the upload code is available in the [tutorial notebook](https://githubtocolab.com/qdrant/examples/blob/master/pdf-retrieval-at-scale/ColPali_ColQwen_Tutorial.ipynb)
with tqdm(total=len(dataset), desc=f"Uploading progress of \"{dataset_source}\" dataset to \"{collection_name}\" collection") as pbar:
for i in range(0, len(dataset), batch_size):
batch = dataset[i : i + batch_size]
image_batch = batch["image"]
current_batch_size = len(image_batch)
original_batch, pooled_by_rows_batch, pooled_by_columns_batch = embed_and_mean_pool_batch(
image_batch,
colpali_processor,
colpali_model,
"colPali"
)
upload_batch(
np.asarray(original_batch, dtype=np.float32),
np.asarray(pooled_by_rows_batch, dtype=np.float32),
np.asarray(pooled_by_columns_batch, dtype=np.float32),
[
{
"source": dataset_source, #HF dataset handle
"index": j #order of an image in the dataset
}
for j in range(i, i + current_batch_size)
],
collection_name
)
# Update the progress bar
pbar.update(current_batch_size)
print("Uploading complete!")
```
</details>
<details>
<summary> <span style="background-color: gray; color: black;"> Uploading "UFO dataset" to "colqwen_tutorial" collection </span> </summary>
```python
batch_size = 1 #based on available compute
dataset_source = ufo_dataset
collection_name = "colqwen_tutorial"
with tqdm(total=len(dataset), desc=f"Uploading progress of \"{dataset_source}\" dataset to \"{collection_name}\" collection") as pbar:
for i in range(0, len(dataset), batch_size):
batch = dataset[i : i + batch_size]
image_batch = batch["image"]
current_batch_size = len(image_batch)
original_batch, pooled_by_rows_batch, pooled_by_columns_batch = embed_and_mean_pool_batch(
image_batch,
colqwen_processor,
colqwen_model,
"colQwen"
)
upload_batch(
np.asarray(original_batch, dtype=np.float32),
np.asarray(pooled_by_rows_batch, dtype=np.float32),
np.asarray(pooled_by_columns_batch, dtype=np.float32),
[
{
"source": dataset_source, #HF dataset handle
"index": j #order of an image in the dataset
}
for j in range(i, i + current_batch_size)
],
collection_name
)
# Update the progress bar
pbar.update(current_batch_size)
print("Uploading complete!")
```
</details>
## Querying PDFs
After indexing PDF documents, we can move on to querying them using our two-stage retrieval approach.
```python
def batch_embed_query(query_batch, model_processor, model):
with torch.no_grad():
processed_queries = model_processor.process_queries(query_batch).to(model.device)
query_embeddings_batch = model(**processed_queries)
return query_embeddings_batch.cpu().float().numpy()
```
Let's use some query from the "UFO dataset" for testing.
```python
query = "Lee Harvey Oswald's involvement in the JFK assassination"
colpali_query = batch_embed_query([query], colpali_processor, colpali_model)
colqwen_query = batch_embed_query([query], colqwen_processor, colqwen_model)
processed_queries = model_processor.process_queries([query]).to(model.device)
print(f"ColPali embedded query \"{query}\" with {len(colpali_query[0])} multivectors of dim {len(colpali_query[0][0])}")
print(f"ColQwen embedded query \"{query}\" with {len(colqwen_query[0])} multivectors of dim {len(colqwen_query[0][0])}")
```
For example, our random batch had this query within:
```bash
ColPali embedded query "Lee Harvey Oswald's involvement in the JFK assassination" with 22 multivectors of dim 128
ColQwen embedded query "Lee Harvey Oswald's involvement in the JFK assassination" with 21 multivectors of dim 128
# Resulting query embedding is a tensor of shape (2, 128)
query_embedding = model(**processed_queries)[0]
```
Now let's design a function for the two-stage retrieval with multivectors produced by VLLMs:
- **Step 1:** Prefetch results using a compressed multivector representation & HNSW index.
- **Step 2:** Re-rank the prefetched results using the original multivector representation.
This function supports batch querying.
Let's query our collections using combined mean pooled representations for the first stage of retrieval.
```python
def reranking_search_batch(query_batch,
collection_name,
search_limit=20,
prefetch_limit=200):
search_queries = [
models.QueryRequest(
query=query,
prefetch=[
models.Prefetch(
query=query,
limit=prefetch_limit,
using="mean_pooling_columns"
),
models.Prefetch(
query=query,
limit=prefetch_limit,
using="mean_pooling_rows"
),
],
limit=search_limit,
with_payload=True,
with_vector=False,
using="original"
) for query in query_batch
]
return client.query_batch_points(
collection_name=collection_name,
requests=search_queries
)
```
Let's query our collections using combined mean pooled representations for the first stage of retrieval.
```python
answer_colpali = reranking_search_batch(colpali_query, "colpali_tutorial")
answer_colqwen = reranking_search_batch(colqwen_query, "colqwen_tutorial")
# Final amount of results to return
search_limit = 10
# Amount of results to prefetch for reranking
prefetch_limit = 100
response = client.query_points(
collection_name=collection_name,
query=query_embedding,
prefetch=[
models.Prefetch(
query=query_embedding,
limit=prefetch_limit,
using="mean_pooling_columns"
),
models.Prefetch(
query=query_embedding,
limit=prefetch_limit,
using="mean_pooling_rows"
),
],
limit=search_limit,
with_payload=True,
with_vector=False,
using="original"
)
```
And check the top retrieved result to our query *"Lee Harvey Oswald's involvement in the JFK assassination"*.
### Results, ColPali
```python
dataset[answer_colpali[0].points[0].payload['index']]['image']
dataset[response.points[0].payload['index']]['image']
```
<details>
<summary> <span style="background-color: gray; color: black;"> Result </span> </summary>
![Results, ColPali](/documentation/tutorials/pdf-retrieval-at-scale/result-VLLMs.png)
</details>
### Results, ColQwen2
```python
dataset[answer_colqwen[0].points[0].payload['index']]['image']
```
<details>
<summary> <span style="background-color: gray; color: black;"> Result </span> </summary>
![Results, ColQwen2](/documentation/tutorials/pdf-retrieval-at-scale/result-VLLMs.png)
</details>
## Conclusion