Spaces:
Runtime error
Runtime error
Sean MacAvaney commited on
Commit ·
5354363
1
Parent(s): d1035ed
initial commit
Browse files- Dockerfile +45 -0
- README.md +9 -5
- app.py +107 -0
- build.sh +15 -0
- doc.md +10 -0
- packages.txt +7 -0
- query.md +10 -0
- requirements.txt +5 -0
- wrapup.md +46 -0
Dockerfile
ADDED
|
@@ -0,0 +1,45 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Base image
|
| 2 |
+
FROM python:3.10-slim
|
| 3 |
+
|
| 4 |
+
# Avoid interactive prompts
|
| 5 |
+
ENV DEBIAN_FRONTEND=noninteractive
|
| 6 |
+
|
| 7 |
+
# System deps (for building Python + Rust crates)
|
| 8 |
+
RUN apt-get update && apt-get install -y \
|
| 9 |
+
curl \
|
| 10 |
+
build-essential \
|
| 11 |
+
default-jre \
|
| 12 |
+
default-jre-headless \
|
| 13 |
+
default-jdk \
|
| 14 |
+
default-jdk-headless \
|
| 15 |
+
debianutils \
|
| 16 |
+
git \
|
| 17 |
+
openssl \
|
| 18 |
+
libssl-dev \
|
| 19 |
+
pkg-config \
|
| 20 |
+
&& rm -rf /var/lib/apt/lists/*
|
| 21 |
+
|
| 22 |
+
# Install Rust (latest stable via rustup)
|
| 23 |
+
RUN curl https://sh.rustup.rs -sSf | sh -s -- -y
|
| 24 |
+
ENV PATH="/root/.cargo/bin:${PATH}"
|
| 25 |
+
|
| 26 |
+
# Verify versions (optional but useful for logs)
|
| 27 |
+
RUN rustc --version && cargo --version
|
| 28 |
+
|
| 29 |
+
# Set workdir
|
| 30 |
+
WORKDIR /app
|
| 31 |
+
|
| 32 |
+
# Copy dependency files first (for caching)
|
| 33 |
+
COPY requirements.txt .
|
| 34 |
+
|
| 35 |
+
# Install Python deps
|
| 36 |
+
RUN pip install --no-cache-dir -r requirements.txt
|
| 37 |
+
|
| 38 |
+
# Copy the rest of the app
|
| 39 |
+
COPY . .
|
| 40 |
+
|
| 41 |
+
# Expose port (HF expects 7860)
|
| 42 |
+
EXPOSE 7860
|
| 43 |
+
|
| 44 |
+
# Run the app
|
| 45 |
+
CMD ["python", "app.py"]
|
README.md
CHANGED
|
@@ -1,12 +1,16 @@
|
|
| 1 |
---
|
| 2 |
-
title:
|
| 3 |
-
emoji:
|
| 4 |
-
colorFrom:
|
| 5 |
colorTo: green
|
| 6 |
sdk: gradio
|
| 7 |
-
sdk_version:
|
|
|
|
| 8 |
app_file: app.py
|
| 9 |
pinned: false
|
| 10 |
---
|
| 11 |
|
| 12 |
-
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
---
|
| 2 |
+
title: PyTerrier SPLADE
|
| 3 |
+
emoji: 🐕
|
| 4 |
+
colorFrom: green
|
| 5 |
colorTo: green
|
| 6 |
sdk: gradio
|
| 7 |
+
sdk_version: 5.7.1
|
| 8 |
+
python_version: '3.10'
|
| 9 |
app_file: app.py
|
| 10 |
pinned: false
|
| 11 |
---
|
| 12 |
|
| 13 |
+
# 🐕 PyTerrier: SPLADE
|
| 14 |
+
|
| 15 |
+
This is a demonstration of [PyTerrier's SPLADE package](https://github.com/cmacdonald/pyt_splade). The SPLADE model encodes queries and documents
|
| 16 |
+
into sparse representations, which can then be used for indexing and retrieval.
|
app.py
ADDED
|
@@ -0,0 +1,107 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import base64
|
| 2 |
+
import re
|
| 3 |
+
import json
|
| 4 |
+
import pandas as pd
|
| 5 |
+
import gradio as gr
|
| 6 |
+
import pyterrier as pt
|
| 7 |
+
pt.init()
|
| 8 |
+
import pyt_splade
|
| 9 |
+
from pyterrier_gradio import Demo, MarkdownFile, interface, df2code, code2md, EX_Q, EX_D, df2list
|
| 10 |
+
factory_max = pyt_splade.Splade(agg='max')
|
| 11 |
+
factory_sum = pyt_splade.Splade(agg='sum')
|
| 12 |
+
|
| 13 |
+
COLAB_NAME = 'pyterrier_splade.ipynb'
|
| 14 |
+
COLAB_INSTALL = '''
|
| 15 |
+
!pip install -q git+https://github.com/naver/splade
|
| 16 |
+
!pip install -q git+https://github.com/cmacdonald/pyt_splade
|
| 17 |
+
'''.strip()
|
| 18 |
+
|
| 19 |
+
def generate_vis(df, mode='Document'):
|
| 20 |
+
if len(df) == 0:
|
| 21 |
+
return ''
|
| 22 |
+
result = []
|
| 23 |
+
if mode == 'Document':
|
| 24 |
+
max_score = max(max(t.values()) for t in df['toks'])
|
| 25 |
+
for row in df.itertuples(index=False):
|
| 26 |
+
if mode == 'Query':
|
| 27 |
+
tok_scores = row.query_toks
|
| 28 |
+
orig_tokens = factory_max.tokenizer.tokenize(row.query)
|
| 29 |
+
max_score = max(tok_scores.values())
|
| 30 |
+
id = row.qid
|
| 31 |
+
else:
|
| 32 |
+
tok_scores = row.toks
|
| 33 |
+
orig_tokens = factory_max.tokenizer.tokenize(row.text)
|
| 34 |
+
id = row.docno
|
| 35 |
+
def toks2span(toks):
|
| 36 |
+
return '<kbd> </kbd>'.join(f'<kbd style="background-color: rgba(66, 135, 245, {tok_scores.get(t, 0)/max_score});">{t}</kbd>' for t in toks)
|
| 37 |
+
orig_tokens_set = set(orig_tokens)
|
| 38 |
+
exp_tokens = [t for t, v in sorted(tok_scores.items(), key=lambda x: (-x[1], x[0])) if t not in orig_tokens_set]
|
| 39 |
+
result.append(f'''
|
| 40 |
+
<div style="font-size: 1.2em;">{mode}: <strong>{id}</strong></div>
|
| 41 |
+
<div style="margin: 4px 0 16px; padding: 4px; border: 1px solid black;">
|
| 42 |
+
<div>
|
| 43 |
+
{toks2span(orig_tokens)}
|
| 44 |
+
</div>
|
| 45 |
+
<div><strong>Expansion Tokens:</strong> {toks2span(exp_tokens)}</div>
|
| 46 |
+
</div>
|
| 47 |
+
''')
|
| 48 |
+
return '\n'.join(result)
|
| 49 |
+
|
| 50 |
+
def predict_query(input, agg):
|
| 51 |
+
code = f'''import pyt_splade
|
| 52 |
+
|
| 53 |
+
splade = pyt_splade.Splade(agg={agg!r})
|
| 54 |
+
|
| 55 |
+
query_pipeline = splade.query_encoder()
|
| 56 |
+
|
| 57 |
+
query_pipeline({df2list(input)})
|
| 58 |
+
'''
|
| 59 |
+
pipeline = {
|
| 60 |
+
'max': factory_max,
|
| 61 |
+
'sum': factory_sum
|
| 62 |
+
}[agg].query_encoder()
|
| 63 |
+
res = pipeline(input)
|
| 64 |
+
vis = generate_vis(res, mode='Query')
|
| 65 |
+
res['query_toks'] = [json.dumps({k: round(v, 4) for k, v in t.items()}) for t in res['query_toks']]
|
| 66 |
+
return (res, code2md(code, COLAB_INSTALL, COLAB_NAME), vis)
|
| 67 |
+
|
| 68 |
+
def predict_doc(input, agg):
|
| 69 |
+
code = f'''import pyt_splade
|
| 70 |
+
|
| 71 |
+
splade = pyt_splade.Splade(agg={repr(agg)})
|
| 72 |
+
|
| 73 |
+
doc_pipeline = splade.doc_encoder()
|
| 74 |
+
|
| 75 |
+
doc_pipeline({df2list(input)})
|
| 76 |
+
'''
|
| 77 |
+
pipeline = {
|
| 78 |
+
'max': factory_max,
|
| 79 |
+
'sum': factory_sum
|
| 80 |
+
}[agg].doc_encoder()
|
| 81 |
+
res = pipeline(input)
|
| 82 |
+
vis = generate_vis(res, mode='Document')
|
| 83 |
+
res['toks'] = [json.dumps({k: round(v, 4) for k, v in t.items()}) for t in res['toks']]
|
| 84 |
+
return (res, code2md(code, COLAB_INSTALL, COLAB_NAME), vis)
|
| 85 |
+
|
| 86 |
+
interface(
|
| 87 |
+
MarkdownFile('README.md'),
|
| 88 |
+
MarkdownFile('query.md'),
|
| 89 |
+
Demo(
|
| 90 |
+
predict_query,
|
| 91 |
+
EX_Q,
|
| 92 |
+
[
|
| 93 |
+
gr.Dropdown(choices=['max', 'sum'], value='max', label='Aggregation'),
|
| 94 |
+
],
|
| 95 |
+
scale=2/3
|
| 96 |
+
),
|
| 97 |
+
MarkdownFile('doc.md'),
|
| 98 |
+
Demo(
|
| 99 |
+
predict_doc,
|
| 100 |
+
EX_D,
|
| 101 |
+
[
|
| 102 |
+
gr.Dropdown(choices=['max', 'sum'], value='max', label='Aggregation'),
|
| 103 |
+
],
|
| 104 |
+
scale=2/3
|
| 105 |
+
),
|
| 106 |
+
MarkdownFile('wrapup.md'),
|
| 107 |
+
).launch(share=False, server_name="0.0.0.0", server_port=7860)
|
build.sh
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
set -e
|
| 3 |
+
|
| 4 |
+
# Install rustup (latest stable toolchain)
|
| 5 |
+
curl https://sh.rustup.rs -sSf | sh -s -- -y
|
| 6 |
+
|
| 7 |
+
# Load cargo into PATH
|
| 8 |
+
source $HOME/.cargo/env
|
| 9 |
+
|
| 10 |
+
# Ensure latest stable (>= 1.88)
|
| 11 |
+
rustup update stable
|
| 12 |
+
rustup default stable
|
| 13 |
+
|
| 14 |
+
rustc --version
|
| 15 |
+
cargo --version
|
doc.md
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
### Document Encoding
|
| 2 |
+
|
| 3 |
+
The document encoder works similarly to the query encoder: it is a `D→D` (document rewriting, doc-to-doc) transformer, and can be used in pipelines accordingly.
|
| 4 |
+
It maps a document's text into a dictionary with terms from the document re-weighted and weighted expansion terms added.
|
| 5 |
+
|
| 6 |
+
<div class="pipeline">
|
| 7 |
+
<div class="df" title="Document Frame">D</div>
|
| 8 |
+
<div class="transformer attn" title="SPLADE Indexing Transformer">SPLADE</div>
|
| 9 |
+
<div class="df" title="Document Frame">D</div>
|
| 10 |
+
</div>
|
packages.txt
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
default-jre
|
| 2 |
+
default-jre-headless
|
| 3 |
+
default-jdk
|
| 4 |
+
default-jdk-headless
|
| 5 |
+
debianutils
|
| 6 |
+
build-essential
|
| 7 |
+
curl
|
query.md
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
### Query Encoding
|
| 2 |
+
|
| 3 |
+
Let's start by exploring SPLADE's query encoder. The query encoder is a `Q→Q` (query rewriting, query-to-query) transformer, and can be used in pipelines accordingly.
|
| 4 |
+
It maps a query string into a sparse `dict`-formatted query vector, with the tokens as the keys and the weights as values.
|
| 5 |
+
|
| 6 |
+
<div class="pipeline">
|
| 7 |
+
<div class="df" title="Query Frame">Q</div>
|
| 8 |
+
<div class="transformer attn" title="SPLADE Query Transformer">SPLADE</div>
|
| 9 |
+
<div class="df" title="Query Frame">Q</div>
|
| 10 |
+
</div>
|
requirements.txt
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
torch
|
| 2 |
+
python-terrier>=0.11.0
|
| 3 |
+
git+https://github.com/seanmacavaney/pyterrier_gradio@v0.0.9
|
| 4 |
+
git+https://github.com/naver/splade
|
| 5 |
+
git+https://github.com/cmacdonald/pyt_splade
|
wrapup.md
ADDED
|
@@ -0,0 +1,46 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
### Putting it all together
|
| 2 |
+
|
| 3 |
+
When you use the document encoder in an indexing pipeline, the rewritten document contents are indexed:
|
| 4 |
+
|
| 5 |
+
<div class="pipeline">
|
| 6 |
+
<div class="df" title="Document Frame">D</div>
|
| 7 |
+
<div class="transformer attn" title="SPLADE Indexing Transformer">SPLADE</div>
|
| 8 |
+
<div class="df" title="Document Frame">D</div>
|
| 9 |
+
<div class="transformer" title="Indexer">Indexer</div>
|
| 10 |
+
<div class="artefact" title="SPLADE Index">IDX</div>
|
| 11 |
+
</div>
|
| 12 |
+
|
| 13 |
+
```python
|
| 14 |
+
import pyterrier as pt
|
| 15 |
+
import pyt_splade
|
| 16 |
+
|
| 17 |
+
dataset = pt.get_dataset('irds:msmarco-passage')
|
| 18 |
+
splade = pyt_splade.Splade()
|
| 19 |
+
|
| 20 |
+
indexer = pt.IterDictIndexer('./msmarco_psg', pretokenised=True)
|
| 21 |
+
|
| 22 |
+
indxer_pipe = splade.doc_encoder() >> indexer
|
| 23 |
+
indxer_pipe.index(dataset.get_corpus_iter())
|
| 24 |
+
```
|
| 25 |
+
|
| 26 |
+
Once you built an index, you can build a retrieval pipeline that first encodes the query,
|
| 27 |
+
and then performs retrieval:
|
| 28 |
+
|
| 29 |
+
<div class="pipeline">
|
| 30 |
+
<div class="df" title="Query Frame">Q</div>
|
| 31 |
+
<div class="transformer attn" title="SPLADE Query Transformer">SPLADE</div>
|
| 32 |
+
<div class="df" title="Query Frame">Q</div>
|
| 33 |
+
<div class="transformer" title="Term Frequency Transformer">TF Retriever <div class="artefact" title="SPLADE Index">IDX</div></div>
|
| 34 |
+
<div class="df" title="Result Frame">R</div>
|
| 35 |
+
</div>
|
| 36 |
+
|
| 37 |
+
```python
|
| 38 |
+
splade_retr = splade.query_encoder() >> pt.terrier.Retriever('./msmarco_psg', wmodel='Tf')
|
| 39 |
+
```
|
| 40 |
+
|
| 41 |
+
### References & Credits
|
| 42 |
+
|
| 43 |
+
This package uses [Naver's SPLADE repository](https://github.com/naver/splade).
|
| 44 |
+
|
| 45 |
+
- Thibault Formal, Benjamin Piwowarski, Stéphane Clinchant. [SPLADE: Sparse Lexical and Expansion Model for First Stage Ranking](https://arxiv.org/abs/2107.05720). SIGIR 2021.
|
| 46 |
+
- Craig Macdonald, Nicola Tonellotto, Sean MacAvaney, Iadh Ounis. [PyTerrier: Declarative Experimentation in Python from BM25 to Dense Retrieval](https://dl.acm.org/doi/abs/10.1145/3459637.3482013). CIKM 2021.
|