Views
No views yet
1import torch, math
2import numpy as np
3from transformers import AutoTokenizer, AutoModelForMaskedLM
4
5model_name = "sdadas/polish-splade"
6tokenizer = AutoTokenizer.from_pretrained(model_name)
7model = AutoModelForMaskedLM.from_pretrained(model_name)
8vocab = {v: k for k, v in tokenizer.get_vocab().items()}
9
10def encode_splade(text: str):
11 input = tokenizer([text], padding="longest", truncation=True, return_tensors="pt", max_length=512)
12 output = model(**input)
13 logits, attention_mask = output["logits"].detach(), input["attention_mask"].detach()
14 attention_mask = attention_mask.unsqueeze(-1)
15 vector = torch.max(torch.log(torch.add(torch.relu(logits), 1)) * attention_mask, dim=1)
16 vector = vector[0].detach().squeeze()
17 idx = np.nonzero(vector.cpu().numpy())[0]
18 vector = vector[idx]
19 return {vocab[k]: float(v) for k, v in zip(list(idx), list(vector))}
20
21def cos_sim(vec1, vec2):
22 intersection = set(vec1.keys()) & set(vec2.keys())
23 numerator = sum([vec1[x] * vec2[x] for x in intersection])
24 sum1 = sum([vec1[x] ** 2 for x in list(vec1.keys())])
25 sum2 = sum([vec2[x] ** 2 for x in list(vec2.keys())])
26 denominator = math.sqrt(sum1) * math.sqrt(sum2)
27 return (numerator / denominator) if denominator else 0.0
28
29question = encode_splade("Jak dożyć 100 lat?")
30answer = encode_splade("Trzeba zdrowo się odżywiać i uprawiać sport.")
31print(cos_sim(question, answer))1from search import SpladeEncoder
2from sentence_transformers.util import cos_sim
3
4config = {"name": "sdadas/polish-splade", "fp16": True}
5encoder = SpladeEncoder(config, True)
6results = encoder.encode_batch(["Jak dożyć 100 lat?", "Trzeba zdrowo się odżywiać i uprawiać sport."])
7print(cos_sim(results[0], results[1]))1@article{dadas2024pirb,
2 title={{PIRB}: A Comprehensive Benchmark of Polish Dense and Hybrid Text Retrieval Methods},
3 author={Sławomir Dadas and Michał Perełkiewicz and Rafał Poświata},
4 year={2024},
5 eprint={2402.13350},
6 archivePrefix={arXiv},
7 primaryClass={cs.CL}
8}