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, ...).
pip install pyterrier-pylateimport 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')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()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],
)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')| 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.