INF-WSE is a series of word-level sparse embedding models developed by
INF TECH.
These models are optimized to generate sparse, high-dimensional text embeddings that excel in capturing the most
relevant information for search and retrieval, particularly in Chinese text.
1import torch
2from transformers import AutoTokenizer, AutoModel
3
4queries = ['电脑一体机由什么构成?', '什么是掌上电脑?']
5documents = [
6 '电脑一体机,是由一台显示器、一个电脑键盘和一个鼠标组成的电脑。',
7 '掌上电脑是一种运行在嵌入式操作系统和内嵌式应用软件之上的、小巧、轻便、易带、实用、价廉的手持式计算设备。',
8]
9input_texts = queries + documents
10
11tokenizer = AutoTokenizer.from_pretrained("infly/inf-wse-v1-base-zh", trust_remote_code=True, use_fast=False) # Fast tokenizer has not been supported yet
12model = AutoModel.from_pretrained("infly/inf-wse-v1-base-zh", trust_remote_code=True)
13model.eval()
14
15max_length = 512
16
17input_batch = tokenizer(input_texts, padding=True, max_length=max_length, truncation=True, return_tensors="pt")
18with torch.no_grad():
19 embeddings = model(input_batch['input_ids'], input_batch['attention_mask'], return_sparse=False) # if return_sparse=True, return sparse tensor, else return dense tensor
20
21scores = embeddings[:2] @ embeddings[2:].T
22print(scores.tolist())
23# [[21.224790573120117, 4.520412921905518], [10.290857315063477, 19.359437942504883]]
1from collections import OrderedDict
2def convert_embeddings_to_weights(embeddings, tokenizer):
3 values, indices = torch.sort(embeddings, dim=-1, descending=True)
4
5 token2weight = []
6 for i in range(embeddings.size(0)):
7 token2weight.append(OrderedDict())
8
9 non_zero_mask = values[i] != 0
10 tokens = tokenizer.convert_ids_to_tokens(indices[i][non_zero_mask])
11 weights = values[i][non_zero_mask].tolist()
12
13 for token, weight in zip(tokens, weights):
14 token2weight[i][token] = weight
15
16 return token2weight
17
18token2weight = convert_embeddings_to_weights(embeddings, tokenizer)
19print(token2weight[1])
20# OrderedDict([('掌上', 3.4572525024414062), ('电脑', 2.6253132820129395), ('是', 2.0787220001220703), ('什么', 1.2899624109268188)])
All results, except for BM25, are measured by building the sparse index via
Qdrant.