Views
No views yet
1
2from transformers import AutoModel, AutoTokenizer, default_data_collator
3from torch.utils.data import Dataset, DataLoader
4import torch
5import numpy as np
6from tqdm import tqdm
7
8# Define model and tokenizer
9model_name_or_path = "Dash00/retriever-SapBERT-mondo"
10encoder = AutoModel.from_pretrained(model_name_or_path)
11tokenizer = AutoTokenizer.from_pretrained(model_name_or_path)
12
13# Optional: move model to GPU if available
14use_cuda = torch.cuda.is_available()
15if use_cuda:
16 encoder = encoder.cuda()
17
18# Define parameters
19max_length = 128
20batch_size = 16
21show_progress = True
22
23# Define input names
24names = ["covid-19", "Coronavirus infection", "high fever", "Tumor of posterior wall of oropharynx"]
25
26# Encode names
27name_encodings = tokenizer(names, padding="max_length", max_length=max_length, truncation=True, return_tensors="pt")
28if use_cuda:
29 name_encodings = {k: v.cuda() for k, v in name_encodings.items()}
30
31# Create dataset and dataloader
32class NamesDataset(Dataset):
33 def __init__(self, encodings):
34 self.encodings = encodings
35
36 def __len__(self):
37 return self.encodings['input_ids'].shape[0]
38
39 def __getitem__(self, idx):
40 return {k: v[idx] for k, v in self.encodings.items()}
41
42name_dataset = NamesDataset(name_encodings)
43name_dataloader = DataLoader(name_dataset, shuffle=False, collate_fn=default_data_collator, batch_size=batch_size)
44
45# Encode and collect embeddings
46dense_embeds = []
47encoder.eval()
48with torch.no_grad():
49 for batch in tqdm(name_dataloader, disable=not show_progress, desc='embedding dictionary'):
50 outputs = encoder(**batch)
51 batch_dense_embeds = outputs.last_hidden_state[:, 0].cpu().numpy() # CLS token
52 dense_embeds.append(batch_dense_embeds)
53
54dense_embeds = np.concatenate(dense_embeds, axis=0)