Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

pyterrier-pylate

Late interaction (multi-vector) retrieval for PyTerrier, backed by PyLate.

PyLate provides ColBERT-style models built on Sentence Transformers and an efficient PLAID index. This package exposes them as PyTerrier transformers so they compose with the rest of the PyTerrier ecosystem (BM25 first stages, pt.Experiment, artifact sharing on HuggingFace Hub, ...).

Installation

pip install pyterrier-pylate

Usage

Indexing and end-to-end retrieval

import pyterrier as pt
from pyterrier_pylate import PyLateBiEncoder, PlaidIndex

model = PyLateBiEncoder('lightonai/GTE-ModernColBERT-v1')
index = PlaidIndex('./msmarco-passage.plaid')

dataset = pt.get_dataset('irds:msmarco-passage')

# index: encode documents, then store their token embeddings in PLAID
(model.doc_encoder() >> index).index(dataset.get_corpus_iter())

# retrieve: encode queries, then search the PLAID index
retriever = model.query_encoder() >> index.retriever(k=1000)
retriever.search('chemical reactions')

Re-ranking a first stage

Two options, depending on whether the documents are already indexed:

# 1. Score from the indexed embeddings (no document re-encoding)
pipeline = bm25 % 100 >> model.query_encoder() >> index.scorer()

# 2. Score from document text (no index needed at all)
pipeline = bm25 % 100 >> pt.text.get_text(dataset, 'text') >> model.text_scorer()

Evaluation

from pyterrier.measures import nDCG, RR

pt.Experiment(
    [bm25, retriever, bm25 % 100 >> model.query_encoder() >> index.scorer()],
    dataset.get_topics(),
    dataset.get_qrels(),
    [nDCG @ 10, RR @ 10],
)

Sharing indexes

PlaidIndex is a PyTerrier artifact, so built indexes can be shared and reused:

index.to_hf('username/msmarco-passage.plaid')
index = PlaidIndex.from_hf('username/msmarco-passage.plaid')

Transformers

Transformer Input Output
model.query_encoder() qid, query qid, query, query_embs
model.doc_encoder() docno, text docno, text, doc_embs
model.text_scorer() qid, query, docno, text + score, rank
index.retriever(k) qid, query_embs qid, docno, score, rank
index.scorer() qid, query_embs, docno + score, rank

query_embs / doc_embs columns hold one (num_tokens, dim) matrix per row, following the convention of pyterrier_dr's multi-vector encoders.

About

Late interaction (multi-vector) retrieval for PyTerrier, backed by PyLate

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages