Views
No views yet
64, 128, 256, 384, 512 and the full size of 768. It's important to note while this method saves space, the same computational resources are used regardless of the dimension size.1import txtai
2
3# New embeddings with requested dimensionality
4embeddings = txtai.Embeddings(
5 path="neuml/pubmedbert-base-embeddings-matryoshka",
6 content=True,
7 dimensionality=256
8)
9embeddings.index(documents())
10
11# Run a query
12embeddings.search("query to run")1from sentence_transformers import SentenceTransformer
2sentences = ["This is an example sentence", "Each sentence is converted"]
3
4model = SentenceTransformer("neuml/pubmedbert-base-embeddings-matryoshka")
5embeddings = model.encode(sentences)
6
7# Requested dimensionality
8dimensionality = 256
9
10print(embeddings[:, :dimensionality])1from transformers import AutoTokenizer, AutoModel
2import torch
3
4# Mean Pooling - Take attention mask into account for correct averaging
5def meanpooling(output, mask):
6 embeddings = output[0] # First element of model_output contains all token embeddings
7 mask = mask.unsqueeze(-1).expand(embeddings.size()).float()
8 return torch.sum(embeddings * mask, 1) / torch.clamp(mask.sum(1), min=1e-9)
9
10# Sentences we want sentence embeddings for
11sentences = ['This is an example sentence', 'Each sentence is converted']
12
13# Load model from HuggingFace Hub
14tokenizer = AutoTokenizer.from_pretrained("neuml/pubmedbert-base-embeddings-matryoshka")
15model = AutoModel.from_pretrained("neuml/pubmedbert-base-embeddings-matryoshka")
16
17# Tokenize sentences
18inputs = tokenizer(sentences, padding=True, truncation=True, return_tensors='pt')
19
20# Compute token embeddings
21with torch.no_grad():
22 output = model(**inputs)
23
24# Perform pooling. In this case, mean pooling.
25embeddings = meanpooling(output, inputs['attention_mask'])
26
27# Requested dimensionality
28dimensionality = 256
29
30print("Sentence embeddings:")
31print(embeddings[:, :dimensionality])| Model | PubMed QA | PubMed Subset | PubMed Summary | Average |
|---|---|---|---|---|
| all-MiniLM-L6-v2 | 90.40 | 95.92 | 94.07 | 93.46 |
| bge-base-en-v1.5 | 91.02 | 95.82 | 94.49 | 93.78 |
| gte-base | 92.97 | 96.90 | 96.24 | 95.37 |
| pubmedbert-base-embeddings | 93.27 | 97.00 | 96.58 | 95.62 |
| S-PubMedBert-MS-MARCO | 90.86 | 93.68 | 93.54 | 92.69 |
pubmedbert-base-embeddings-matryoshka.| Model | PubMed QA | PubMed Subset | PubMed Summary | Average |
|---|---|---|---|---|
| Dimensions = 64 | 92.16 | 96.14 | 95.67 | 94.66 |
| Dimensions = 128 | 92.80 | 96.58 | 96.22 | 95.20 |
| Dimensions = 256 | 93.11 | 96.82 | 96.53 | 95.49 |
| Dimensions = 384 | 93.42 | 97.00 | 96.61 | 95.68 |
| Dimensions = 512 | 93.37 | 97.07 | 96.61 | 95.68 |
| Dimensions = 768 | 93.53 | 97.13 | 96.70 | 95.79 |
Dimensions = 256 performs better than all the other models originally tested above. Even Dimensions = 64 performs better than all-MiniLM-L6-v2 and bge-base-en-v1.5.torch.utils.data.dataloader.DataLoader of length 20191 with parameters:{'batch_size': 24, 'sampler': 'torch.utils.data.sampler.RandomSampler', 'batch_sampler': 'torch.utils.data.sampler.BatchSampler'}sentence_transformers.losses.MatryoshkaLoss.MatryoshkaLoss with parameters:{'loss': 'MultipleNegativesRankingLoss', 'matryoshka_dims': [768, 512, 384, 256, 128, 64], 'matryoshka_weights': [1, 1, 1, 1, 1, 1]}{
"epochs": 1,
"evaluation_steps": 500,
"evaluator": "sentence_transformers.evaluation.EmbeddingSimilarityEvaluator.EmbeddingSimilarityEvaluator",
"max_grad_norm": 1,
"optimizer_class": "<class 'torch.optim.adamw.AdamW'>",
"optimizer_params": {
"lr": 2e-05
},
"scheduler": "WarmupLinear",
"steps_per_epoch": null,
"warmup_steps": 10000,
"weight_decay": 0.01
}SentenceTransformer(
(0): Transformer({'max_seq_length': 512, 'do_lower_case': False}) with Transformer model: BertModel
(1): Pooling({'word_embedding_dimension': 768, 'pooling_mode_cls_token': False, 'pooling_mode_mean_tokens': True, 'pooling_mode_max_tokens': False, 'pooling_mode_mean_sqrt_len_tokens': False, 'pooling_mode_weightedmean_tokens': False, 'pooling_mode_lasttoken': False, 'include_prompt': True})
)