Views
No views yet
bert-base-uncased[CLS] token representation)1from transformers import AutoTokenizer, AutoModel
2import torch
3
4# Load query encoder
5q_tokenizer = AutoTokenizer.from_pretrained("liyongkang/dragon-query-encoder")
6q_model = AutoModel.from_pretrained("liyongkang/dragon-query-encoder")
7
8# Load context encoder
9p_tokenizer = AutoTokenizer.from_pretrained("liyongkang/dragon-context-encoder")
10p_model = AutoModel.from_pretrained("liyongkang/dragon-context-encoder")
11
12query = "What is Dragon in NLP?"
13passage = "A dual-encoder retrieval model for dense passage retrieval."
14
15
16# Tokenize. In fact, the two tokenizers are the same.
17q_inputs = q_tokenizer(query, return_tensors="pt", truncation=True, padding=True)
18p_inputs = p_tokenizer(passage, return_tensors="pt", truncation=True, padding=True)
19
20with torch.no_grad():
21 q_vec = q_model(**q_inputs).last_hidden_state[:, 0] # CLS pooling
22 p_vec = p_model(**p_inputs).last_hidden_state[:, 0] # CLS pooling
23 score = (q_vec * p_vec).sum(dim=-1)
24 print("Dot product similarity:", score.item())