Views
No views yet
| Virtual token embeddings file | Backbone Model |
|---|---|
| mistral.7b.instruct.added_token_embeddings.pt | Mistral-7B-Instruct-v0.1 |
| mistral.7b.base.added_token_embeddings.pt | Mistral-7b-v0.1 |
| llama2.7b.chat.added_token_embeddings.pt | LLaMA-2-7b-chat |
| llama2.7b.base.added_token_embeddings.pt | LLaMA-2-7b |
| llama2.13b.chat.added_token_embeddings.pt | LLaMA-2-7b-chat |
| llama2.13b.base.added_token_embeddings.pt | LLaMA-2-7b-base |
torch 2.0.0
transformers 4.37.01def load_tokens(model, tokenizer, token_embedding_path=""):
2 new_tokens_weights = torch.load(token_embedding_path)
3 new_tokens_length = new_tokens_weights.shape[0]
4
5 # expand vocabulary
6 new_tokens = [f"[ref{i+1}]" for i in range(new_tokens_length)]
7 tokenizer.add_tokens(new_tokens)
8
9 # get original embedding weight matrix
10 embedding_layer = model.get_input_embeddings()
11 embedding_weights = embedding_layer.weight
12 original_vocab_size, embedding_dim = embedding_weights.shape
13
14 # create new embedding matrix
15 new_vocab_size = original_vocab_size + new_tokens_length
16 new_embedding_weights = torch.zeros(new_vocab_size, embedding_dim)
17
18 # copy original embeddings to the new weights
19 new_embedding_weights[:original_vocab_size, :] = embedding_weights
20
21 # append virtual token embeddings to the new weights
22 for token, embedding in zip(new_tokens, new_tokens_weights):
23 token_id = tokenizer.convert_tokens_to_ids(token)
24 new_embedding_weights[token_id] = embedding
25
26 # update the embedding table
27 # note: we should avoid using the function resize_token_embeddings() because this function will also change the lm_head of the model
28 embedding_layer.weight.data = new_embedding_weights
29
30 # model.resize_token_embeddings(len(tokenizer))
31
32 return model, tokenizer
33
34model_path = "path/to/Mistral-7B-Instruct-v0.1"
35model = AutoModelForCausalLM.from_pretrained(model_path)
36tokenizer = AutoTokenizer.from_pretrained(model_path)
37model, tokenizer = load_tokens(model, tokenizer, token_embedding_path="/path/to/mistral.7b.instruct.added_token_embeddings.pt")1# using 50 tokens as an example
2added_tokens = [f" [ref{i}]" for i in range(1, 51)]
3added_tokens = "".join(added_tokens)
4retrieved_results = "..."
5question = "..."
6text = [f"{retrieved_results}{added_tokens}Question: {question}\nAnswer:"]
7
8...
9
10outputs = model.generate(...)
111@article{SPRING,
2 author = {Yutao Zhu and
3 Zhaoheng Huang and
4 Zhicheng Dou and
5 Ji{-}Rong Wen},
6 title = {One Token Can Help! Learning Scalable and Pluggable Virtual Tokens for Retrieval-Augmented Large Language Models},
7 journal = {CoRR},
8 volume = {abs/2405.19670},
9 year = {2024},
10 url = {https://doi.org/10.48550/arXiv.2405.19670},
11 doi = {10.48550/ARXIV.2405.19670},
12 eprinttype = {arXiv},
13 eprint = {2405.19670}
14}