Skip to content

SPLADE Reranker

rankify.models.splade_reranker

BaseRanking

Bases: ABC

An abstract base class for implementing different ranking models.

This class defines the interface for all ranking models, ensuring that all subclasses implement the required methods.

Attributes:

Name Type Description
method str

The name of the ranking method.

model_name str

The name of the model being used for ranking.

api_key str

An optional API key for accessing remote models or services.

Source code in rankify/models/base.py
class BaseRanking(ABC):
    """
    An abstract base class for implementing different ranking models.

    This class defines the interface for all ranking models, ensuring that all subclasses implement the required methods.

    Attributes:
        method (str): The name of the ranking method.
        model_name (str): The name of the model being used for ranking.
        api_key (str, optional): An optional API key for accessing remote models or services.
    """

    @abstractmethod
    def __init__(self, method: str= None, model_name: str= None, api_key: str= None, **kwargs) ->None:
        """
        Initializes the base ranking model.

        Args:
            method (str, optional): The name of the ranking method. Defaults to None.
            model_name (str, optional): The name of the model being used for ranking. Defaults to None.
            api_key (str, optional): An optional API key for accessing remote models or services. Defaults to None.

        Example:
            ```python
            class MyRanking(BaseRanking):
                def __init__(self, method, model_name):
                    super().__init__(method, model_name)
            ```
        """
        pass

    @abstractmethod
    def rank(self, documents: list[Document] ):
        """
        Abstract method to rank a list of documents.

        Args:
            documents (list[Document]): A list of Document instances that need to be ranked.

        Raises:
            NotImplementedError: This method must be implemented by subclasses.

        Example:
            ```python
            class MyRanking(BaseRanking):
                def __init__(self, method, model_name):
                    super().__init__(method, model_name)

                def rank(self, documents):
                    # Ranking implementation here
                    pass
            ```
        """
        pass

__init__(method=None, model_name=None, api_key=None, **kwargs) abstractmethod

Initializes the base ranking model.

Parameters:

Name Type Description Default
method str

The name of the ranking method. Defaults to None.

None
model_name str

The name of the model being used for ranking. Defaults to None.

None
api_key str

An optional API key for accessing remote models or services. Defaults to None.

None
Example
class MyRanking(BaseRanking):
    def __init__(self, method, model_name):
        super().__init__(method, model_name)
Source code in rankify/models/base.py
@abstractmethod
def __init__(self, method: str= None, model_name: str= None, api_key: str= None, **kwargs) ->None:
    """
    Initializes the base ranking model.

    Args:
        method (str, optional): The name of the ranking method. Defaults to None.
        model_name (str, optional): The name of the model being used for ranking. Defaults to None.
        api_key (str, optional): An optional API key for accessing remote models or services. Defaults to None.

    Example:
        ```python
        class MyRanking(BaseRanking):
            def __init__(self, method, model_name):
                super().__init__(method, model_name)
        ```
    """
    pass

rank(documents) abstractmethod

Abstract method to rank a list of documents.

Parameters:

Name Type Description Default
documents list[Document]

A list of Document instances that need to be ranked.

required

Raises:

Type Description
NotImplementedError

This method must be implemented by subclasses.

Example
class MyRanking(BaseRanking):
    def __init__(self, method, model_name):
        super().__init__(method, model_name)

    def rank(self, documents):
        # Ranking implementation here
        pass
Source code in rankify/models/base.py
@abstractmethod
def rank(self, documents: list[Document] ):
    """
    Abstract method to rank a list of documents.

    Args:
        documents (list[Document]): A list of Document instances that need to be ranked.

    Raises:
        NotImplementedError: This method must be implemented by subclasses.

    Example:
        ```python
        class MyRanking(BaseRanking):
            def __init__(self, method, model_name):
                super().__init__(method, model_name)

            def rank(self, documents):
                # Ranking implementation here
                pass
        ```
    """
    pass

Document

Represents a document consisting of a question, answers, and contexts.

Attributes:

Name Type Description
question Question

The question associated with the document.

answers Answer

The answers to the question.

contexts list[Context]

A list of related contexts.

reorder_contexts list[Context] or None

A reordered list of contexts based on relevance.

Source code in rankify/dataset/dataset.py
class Document:
    """
    Represents a document consisting of a question, answers, and contexts.

    Attributes:
        question (Question): The question associated with the document.
        answers (Answer): The answers to the question.
        contexts (list[Context]): A list of related contexts.
        reorder_contexts (list[Context] or None): A reordered list of contexts based on relevance.
    """
    def __init__(self, question: Question, answers: Answer, contexts: list = None , id: int = None) -> None:
        """
        Initializes a Document instance.

        Args:
            question (Question): The question associated with the document.
            answers (Answer): The answers to the question.
            contexts (list[Context], optional): A list of contexts related to the question.

        Example:
            ```python
            q = Question("What is the capital of France?")
            a = Answer(["Paris"])
            c1 = Context(score=0.9, has_answer=True, id=1, title="Paris", text="The capital of France is Paris.")
            c2 = Context(score=0.5, has_answer=False, id=2, title="Berlin", text="Berlin is the capital of Germany.")
            d = Document(question=q, answers=a, contexts=[c1, c2])
            print(d)
            ```
        """
        self.question: Question = question
        self.answers: Answer = answers
        self.contexts: List[Context] = contexts
        self.reorder_contexts: List[Context] = None
        self.id = str(id) 

    @classmethod
    def from_dict(cls, data: dict,n_docs:int=100) -> 'Document':
        """
        Creates a Document instance from a dictionary.

        Args:
            data (dict): A dictionary containing the question, answers, and contexts.
            n_docs (int, optional): The number of contexts to include. Defaults to 100.

        Returns:
            Document: A new Document instance.

        Example:
            ```python
            data = {
                "question": "What is the capital of France?",
                "answers": ["Paris"],
                "ctxs": [
                    {"score": 0.9, "has_answer": True, "id": 1, "title": "Paris", "text": "The capital of France is Paris."},
                    {"score": 0.5, "has_answer": False, "id": 2, "title": "Berlin", "text": "Berlin is the capital of Germany."}
                ]
            }
            d = Document.from_dict(data)
            print(d.question)
            ```
        """
        question = Question(data["question"])
        if "answers" in data:
            answers = Answer(data["answers"])
        else:
            answers =Answer('')

        if "query_id" in data:
            id = data["query_id"]
        else:
            id = None
        contexts = [Context(**ctx) for ctx in data["ctxs"][:n_docs]]
        return cls(question, answers, contexts, id=id)

    def to_dict(self) -> Dict[str, Optional[object]]:
        """
        Converts the document into a dictionary representation.

        Returns:
            dict: A dictionary containing the question, answers, and contexts.
        """
        return {
            "question": self.question.question,
            "answers": self.answers.answers,
            "contexts": [ctx.to_dict() for ctx in self.contexts]
        }
    def to_dict_reoreder(self) -> Dict[str,Optional[object]]:
        return {
            "question" : self.question.question,
            "answers" : self.answers.answers,
            "contexts" : [ctx.to_dict() for ctx in self.reorder_contexts]
        }
    def __str__(self) -> str:
        """
        Returns a string representation of the Document instance.

        Returns:
            str: The formatted document information.

        Example:
            ```python
            d = Document(Question("What is the capital of France?"), Answer(["Paris"]))
            print(d)
            ```
        """
        contexts_str = "\n\n".join([str(ctx) for ctx in self.contexts])
        reorder_contexts_str= ''
        if self.reorder_contexts is not None:
            reorder_contexts_str = "\n\n".join([str(ctx) for ctx in self.reorder_contexts])
        return f"{self.question}\n\n{self.answers}\n\nContext: \n\n{contexts_str}\nReorder contexts: \n\n{reorder_contexts_str}"

__init__(question, answers, contexts=None, id=None)

Initializes a Document instance.

Parameters:

Name Type Description Default
question Question

The question associated with the document.

required
answers Answer

The answers to the question.

required
contexts list[Context]

A list of contexts related to the question.

None
Example
q = Question("What is the capital of France?")
a = Answer(["Paris"])
c1 = Context(score=0.9, has_answer=True, id=1, title="Paris", text="The capital of France is Paris.")
c2 = Context(score=0.5, has_answer=False, id=2, title="Berlin", text="Berlin is the capital of Germany.")
d = Document(question=q, answers=a, contexts=[c1, c2])
print(d)
Source code in rankify/dataset/dataset.py
def __init__(self, question: Question, answers: Answer, contexts: list = None , id: int = None) -> None:
    """
    Initializes a Document instance.

    Args:
        question (Question): The question associated with the document.
        answers (Answer): The answers to the question.
        contexts (list[Context], optional): A list of contexts related to the question.

    Example:
        ```python
        q = Question("What is the capital of France?")
        a = Answer(["Paris"])
        c1 = Context(score=0.9, has_answer=True, id=1, title="Paris", text="The capital of France is Paris.")
        c2 = Context(score=0.5, has_answer=False, id=2, title="Berlin", text="Berlin is the capital of Germany.")
        d = Document(question=q, answers=a, contexts=[c1, c2])
        print(d)
        ```
    """
    self.question: Question = question
    self.answers: Answer = answers
    self.contexts: List[Context] = contexts
    self.reorder_contexts: List[Context] = None
    self.id = str(id) 

from_dict(data, n_docs=100) classmethod

Creates a Document instance from a dictionary.

Parameters:

Name Type Description Default
data dict

A dictionary containing the question, answers, and contexts.

required
n_docs int

The number of contexts to include. Defaults to 100.

100

Returns:

Name Type Description
Document Document

A new Document instance.

Example
data = {
    "question": "What is the capital of France?",
    "answers": ["Paris"],
    "ctxs": [
        {"score": 0.9, "has_answer": True, "id": 1, "title": "Paris", "text": "The capital of France is Paris."},
        {"score": 0.5, "has_answer": False, "id": 2, "title": "Berlin", "text": "Berlin is the capital of Germany."}
    ]
}
d = Document.from_dict(data)
print(d.question)
Source code in rankify/dataset/dataset.py
@classmethod
def from_dict(cls, data: dict,n_docs:int=100) -> 'Document':
    """
    Creates a Document instance from a dictionary.

    Args:
        data (dict): A dictionary containing the question, answers, and contexts.
        n_docs (int, optional): The number of contexts to include. Defaults to 100.

    Returns:
        Document: A new Document instance.

    Example:
        ```python
        data = {
            "question": "What is the capital of France?",
            "answers": ["Paris"],
            "ctxs": [
                {"score": 0.9, "has_answer": True, "id": 1, "title": "Paris", "text": "The capital of France is Paris."},
                {"score": 0.5, "has_answer": False, "id": 2, "title": "Berlin", "text": "Berlin is the capital of Germany."}
            ]
        }
        d = Document.from_dict(data)
        print(d.question)
        ```
    """
    question = Question(data["question"])
    if "answers" in data:
        answers = Answer(data["answers"])
    else:
        answers =Answer('')

    if "query_id" in data:
        id = data["query_id"]
    else:
        id = None
    contexts = [Context(**ctx) for ctx in data["ctxs"][:n_docs]]
    return cls(question, answers, contexts, id=id)

to_dict()

Converts the document into a dictionary representation.

Returns:

Name Type Description
dict Dict[str, Optional[object]]

A dictionary containing the question, answers, and contexts.

Source code in rankify/dataset/dataset.py
def to_dict(self) -> Dict[str, Optional[object]]:
    """
    Converts the document into a dictionary representation.

    Returns:
        dict: A dictionary containing the question, answers, and contexts.
    """
    return {
        "question": self.question.question,
        "answers": self.answers.answers,
        "contexts": [ctx.to_dict() for ctx in self.contexts]
    }

__str__()

Returns a string representation of the Document instance.

Returns:

Name Type Description
str str

The formatted document information.

Example
d = Document(Question("What is the capital of France?"), Answer(["Paris"]))
print(d)
Source code in rankify/dataset/dataset.py
def __str__(self) -> str:
    """
    Returns a string representation of the Document instance.

    Returns:
        str: The formatted document information.

    Example:
        ```python
        d = Document(Question("What is the capital of France?"), Answer(["Paris"]))
        print(d)
        ```
    """
    contexts_str = "\n\n".join([str(ctx) for ctx in self.contexts])
    reorder_contexts_str= ''
    if self.reorder_contexts is not None:
        reorder_contexts_str = "\n\n".join([str(ctx) for ctx in self.reorder_contexts])
    return f"{self.question}\n\n{self.answers}\n\nContext: \n\n{contexts_str}\nReorder contexts: \n\n{reorder_contexts_str}"

SpladeReranker

Bases: BaseRanking

Implements SpladeReranker, a sparse lexical and expansion model for first-stage ranking using masked language models (MLMs).

SPLADE employs sparse representations of queries and documents, performing lexical matching while expanding terms in an interpretable way.

References
  • Formal et al. (2021): SPLADE: Sparse Lexical and Expansion Model for First Stage Ranking. Paper
  • Formal et al. (2021): SPLADE v2: Sparse Lexical and Expansion Model for Information Retrieval. Paper

Attributes:

Name Type Description
method str

The name of the reranking method.

model_name str

The name or path to the SPLADE model.

device str

The device (CPU/GPU) used for inference.

query_max_length int

The maximum token length for queries.

document_max_length int

The maximum token length for documents.

batch_size int

The batch size for document encoding.

model AutoModelForMaskedLM

The Masked Language Model (MLM) used for sparse representation.

tokenizer AutoTokenizer

The tokenizer for encoding query and document texts.

Examples:

Basic Usage:

from rankify.dataset.dataset import Document, Question, Context
from rankify.models.reranking import Reranking

# Define a query and contexts
question = Question("What are the advantages of renewable energy?")
contexts = [
    Context(text="Renewable energy reduces carbon emissions and is sustainable.", id=0),
    Context(text="Fossil fuels have been the primary source of energy for centuries.", id=1),
    Context(text="Solar and wind power are prominent forms of renewable energy.", id=2),
]
document = Document(question=question, contexts=contexts)

# Initialize SPLADE Reranker
model = Reranking(method='splade', model_name='splade-cocondenser')
model.rank([document])

# Print reordered contexts
print("Reordered Contexts:")
for context in document.reorder_contexts:
    print(context.text)

Source code in rankify/models/splade_reranker.py
class SpladeReranker(BaseRanking):
    """
    Implements **SpladeReranker**, a **sparse lexical and expansion model**
    for first-stage ranking using **masked language models (MLMs).**



    SPLADE employs **sparse representations** of queries and documents,
    performing lexical matching while **expanding terms** in an interpretable way.

    References:
        - **Formal et al. (2021)**: *SPLADE: Sparse Lexical and Expansion Model for First Stage Ranking*. [Paper](https://arxiv.org/abs/2107.05720)
        - **Formal et al. (2021)**: *SPLADE v2: Sparse Lexical and Expansion Model for Information Retrieval*. [Paper](https://arxiv.org/abs/2109.10086)

    Attributes:
        method (str): The name of the reranking method.
        model_name (str): The name or path to the **SPLADE** model.
        device (str): The device (CPU/GPU) used for inference.
        query_max_length (int): The maximum token length for **queries**.
        document_max_length (int): The maximum token length for **documents**.
        batch_size (int): The batch size for document encoding.
        model (AutoModelForMaskedLM): The **Masked Language Model (MLM)** used for sparse representation.
        tokenizer (AutoTokenizer): The tokenizer for encoding query and document texts.

    Examples:
        **Basic Usage:**
        ```python
        from rankify.dataset.dataset import Document, Question, Context
        from rankify.models.reranking import Reranking

        # Define a query and contexts
        question = Question("What are the advantages of renewable energy?")
        contexts = [
            Context(text="Renewable energy reduces carbon emissions and is sustainable.", id=0),
            Context(text="Fossil fuels have been the primary source of energy for centuries.", id=1),
            Context(text="Solar and wind power are prominent forms of renewable energy.", id=2),
        ]
        document = Document(question=question, contexts=contexts)

        # Initialize SPLADE Reranker
        model = Reranking(method='splade', model_name='splade-cocondenser')
        model.rank([document])

        # Print reordered contexts
        print("Reordered Contexts:")
        for context in document.reorder_contexts:
            print(context.text)
        ```
    """

    def __init__(
        self,
        method: str = None,
        model_name: str = "naver/splade-cocondenser-ensembledistil",
        **kwargs
    ):
        """
        Initializes **SPLADE Reranker** for reranking tasks.

        Args:
            method (str): The name of the reranking method.
            model_name (str): The name or path to the **SPLADE** model.
            **kwargs: Additional parameters:
                - device (str, optional): Computation device (`"auto"`, `"cuda"`, `"cpu"`). Default is `"auto"`.
                - use_fp16 (bool, optional): Whether to use **FP16** inference. Default is `True`.
                - batch_size (int, optional): Batch size for document encoding. Default is `16`.
                - query_max_length (int, optional): Maximum token length for **queries**. Default is `512`.
                - document_max_length (int, optional): Maximum token length for **documents**. Default is `512`.
        """
        super().__init__(method)
        self.device = self._detect_device(kwargs.get("device", "auto"))
        self.model = AutoModelForMaskedLM.from_pretrained(model_name).to(self.device)
        self.tokenizer = AutoTokenizer.from_pretrained(model_name)
        self.model.eval()

        if kwargs.get("use_fp16", True) and "cuda" in self.device:
            self.model.half()

        self.query_max_length = kwargs.get("query_max_length", 512)
        self.document_max_length = kwargs.get("document_max_length", 512)
        self.batch_size = kwargs.get("batch_size", 16)

    def rank(self, documents: List[Document]) -> List[Document]:
        """
        Reranks a list of **Document** instances using the **SPLADE** model.

        Args:
            documents (List[Document]): A list of `Document` instances to rerank.

        Returns:
            List[Document]: The documents with updated `reorder_contexts`.
        """
        for document in tqdm(documents, desc="Reranking Documents"):
            document = self._rerank_document(document)
        return documents

    def _rerank_document(self, document: Document) -> Document:
        """
        Reranks a single document's contexts using the **SPLADE** model.

        Args:
            document (Document): A **Document** instance to rerank.

        Returns:
            Document: The reranked **Document** with updated `reorder_contexts`.
        """
        query = document.question.question  # Extract query text
        contexts = copy.deepcopy(document.contexts)

        # Extract context texts
        context_texts = [ctx.text for ctx in contexts]
        if not context_texts:
            document.reorder_contexts = []
            return document

        # Compute query and document scores
        scores = self._rerank(query, context_texts)
        # Map scores back to contexts and update the scores
        for ctx, score in zip(contexts, scores):
            ctx.score = score

        # Map scores back to contexts
        scored_contexts = list(zip(contexts, scores))
        scored_contexts.sort(key=lambda x: x[1], reverse=True)

        # Reorder contexts in the document
        document.reorder_contexts = [ctx for ctx, _ in scored_contexts]
        return document

    def _compute_vector(self, texts: List[str], max_length: int) -> torch.Tensor:
        """
        Compute **SPLADE-style** sparse embeddings for a list of texts.

        Args:
            texts (List[str]): A list of **texts** to compute embeddings for.
            max_length (int): The **maximum length** for tokenization.

        Returns:
            torch.Tensor: Sparse **document representations**.
        """
        tokens = self.tokenizer(
            texts,
            return_tensors="pt",
            padding=True,
            truncation=True,
            max_length=max_length,
        ).to(self.device)

        with torch.no_grad():
            output = self.model(**tokens)
            logits, attention_mask = output.logits, tokens["attention_mask"]

        return splade_max_pooling(logits, attention_mask)

    def _rerank(self, query: str, documents: List[str]) -> List[float]:
        """
        Computes relevance scores between a **query** and **documents**.

        Args:
            query (str): The **query** text.
            documents (List[str]): A list of **document texts**.

        Returns:
            List[float]: Relevance scores for each document.
        """
        # Compute query embedding
        query_emb = self._compute_vector([query], max_length=self.query_max_length)[0]

        # Compute document embeddings
        doc_embs = []
        for i in range(0, len(documents), self.batch_size):
            doc_embs.append(
                self._compute_vector(
                    documents[i : i + self.batch_size],
                    max_length=self.document_max_length,
                )
            )
        doc_embs = torch.cat(doc_embs, dim=0)

        # Compute similarity scores
        scores = torch.matmul(query_emb.unsqueeze(0), doc_embs.t()).squeeze(0)
        return scores.tolist()

    @staticmethod
    def _detect_device(device: str) -> str:
        """
        Detects the appropriate device for computation.

        Args:
            device (str): Desired device (`"auto"`, `"cuda"`, or `"cpu"`).

        Returns:
            str: The detected device.
        """
        if device == "auto":
            return "cuda" if torch.cuda.is_available() else "cpu"
        return device

__init__(method=None, model_name='naver/splade-cocondenser-ensembledistil', **kwargs)

Initializes SPLADE Reranker for reranking tasks.

Parameters:

Name Type Description Default
method str

The name of the reranking method.

None
model_name str

The name or path to the SPLADE model.

'naver/splade-cocondenser-ensembledistil'
**kwargs

Additional parameters: - device (str, optional): Computation device ("auto", "cuda", "cpu"). Default is "auto". - use_fp16 (bool, optional): Whether to use FP16 inference. Default is True. - batch_size (int, optional): Batch size for document encoding. Default is 16. - query_max_length (int, optional): Maximum token length for queries. Default is 512. - document_max_length (int, optional): Maximum token length for documents. Default is 512.

{}
Source code in rankify/models/splade_reranker.py
def __init__(
    self,
    method: str = None,
    model_name: str = "naver/splade-cocondenser-ensembledistil",
    **kwargs
):
    """
    Initializes **SPLADE Reranker** for reranking tasks.

    Args:
        method (str): The name of the reranking method.
        model_name (str): The name or path to the **SPLADE** model.
        **kwargs: Additional parameters:
            - device (str, optional): Computation device (`"auto"`, `"cuda"`, `"cpu"`). Default is `"auto"`.
            - use_fp16 (bool, optional): Whether to use **FP16** inference. Default is `True`.
            - batch_size (int, optional): Batch size for document encoding. Default is `16`.
            - query_max_length (int, optional): Maximum token length for **queries**. Default is `512`.
            - document_max_length (int, optional): Maximum token length for **documents**. Default is `512`.
    """
    super().__init__(method)
    self.device = self._detect_device(kwargs.get("device", "auto"))
    self.model = AutoModelForMaskedLM.from_pretrained(model_name).to(self.device)
    self.tokenizer = AutoTokenizer.from_pretrained(model_name)
    self.model.eval()

    if kwargs.get("use_fp16", True) and "cuda" in self.device:
        self.model.half()

    self.query_max_length = kwargs.get("query_max_length", 512)
    self.document_max_length = kwargs.get("document_max_length", 512)
    self.batch_size = kwargs.get("batch_size", 16)

rank(documents)

Reranks a list of Document instances using the SPLADE model.

Parameters:

Name Type Description Default
documents List[Document]

A list of Document instances to rerank.

required

Returns:

Type Description
List[Document]

List[Document]: The documents with updated reorder_contexts.

Source code in rankify/models/splade_reranker.py
def rank(self, documents: List[Document]) -> List[Document]:
    """
    Reranks a list of **Document** instances using the **SPLADE** model.

    Args:
        documents (List[Document]): A list of `Document` instances to rerank.

    Returns:
        List[Document]: The documents with updated `reorder_contexts`.
    """
    for document in tqdm(documents, desc="Reranking Documents"):
        document = self._rerank_document(document)
    return documents

splade_max_pooling(logits, attention_mask)

Perform Splade-style max pooling with log scaling.

Source code in rankify/models/splade_reranker.py
def splade_max_pooling(logits, attention_mask):
    """
    Perform Splade-style max pooling with log scaling.
    """
    relu_log = torch.log(1 + torch.relu(logits))
    weighted_log = relu_log * attention_mask.unsqueeze(-1)
    max_val, _ = torch.max(weighted_log, dim=1)
    return max_val