Views
No views yet
1import torch
2from transformers import AutoModel, AutoTokenizer
3from peft import PeftModel, PeftConfig
4
5def get_model(peft_model_name):
6 config = PeftConfig.from_pretrained(peft_model_name)
7 base_model = AutoModel.from_pretrained(config.base_model_name_or_path)
8 model = PeftModel.from_pretrained(base_model, peft_model_name)
9 model = model.merge_and_unload()
10 model.eval()
11 return model
12
13# Load the tokenizer and model
14tokenizer = AutoTokenizer.from_pretrained('meta-llama/Llama-2-7b-hf')
15model = get_model('castorini/repllama-v1-7b-lora-passage')
16
17# Define query and passage inputs
18query = "What is llama?"
19title = "Llama"
20passage = "The llama is a domesticated South American camelid, widely used as a meat and pack animal by Andean cultures since the pre-Columbian era."
21query_input = tokenizer(f'query: {query}</s>', return_tensors='pt')
22passage_input = tokenizer(f'passage: {title} {passage}</s>', return_tensors='pt')
23
24# Run the model forward to compute embeddings and query-passage similarity score
25with torch.no_grad():
26 # compute query embedding
27 query_outputs = model(**query_input)
28 query_embedding = query_outputs.last_hidden_state[0][-1]
29 query_embedding = torch.nn.functional.normalize(query_embedding, p=2, dim=0)
30
31 # compute passage embedding
32 passage_outputs = model(**passage_input)
33 passage_embeddings = passage_outputs.last_hidden_state[0][-1]
34 passage_embeddings = torch.nn.functional.normalize(passage_embeddings, p=2, dim=0)
35
36 # compute similarity score
37 score = torch.dot(query_embedding, passage_embeddings)
38 print(score)
39@article{rankllama,
title={Fine-Tuning LLaMA for Multi-Stage Text Retrieval},
author={Xueguang Ma and Liang Wang and Nan Yang and Furu Wei and Jimmy Lin},
year={2023},
journal={arXiv:2310.08319},
}